From db535a8ae42fb98d886afd10b6bc6d473bb16bb7 Mon Sep 17 00:00:00 2001 From: lhear <121179341+lhear@users.noreply.github.com> Date: Mon, 3 Aug 2026 02:41:57 +0800 Subject: [PATCH 1/7] refactor: zero-copy tunnel pipeline overhaul - zero-copy frame decode/seal, no per-frame allocation - direct streaming download path (no segment task/channel) - 64KiB write coalescing, 32KiB read prefetch, h2-aligned frames - UUID stream registry, shared resolved shaper config - admission control: max_tunnels/max_connections/in-flight caps - fix: rotation Continue truncation, swallowed upstream write errors, DoT hangs, prefetch race killing tunnels, JSON frame cap, IPv6 CONNECT - validate in-flight/concurrency config bounds - rewrite CONFIGURATION.md; strip all source comments --- CONFIGURATION.md | 372 ++++++++++-- Cargo.lock | 4 +- Cargo.toml | 2 +- src/bin/client.rs | 18 +- src/client/actor/connection.rs | 2 +- src/client/actor/download_loop.rs | 190 ++++-- src/client/actor/upload_loop.rs | 147 ++--- src/client/connection.rs | 6 +- src/client/constants.rs | 6 + src/client/handshake.rs | 7 +- src/client/mod.rs | 21 + src/client/proxy.rs | 14 +- src/client/state.rs | 6 +- src/client/utils.rs | 18 +- src/config/mod.rs | 8 + src/crypto/cipher.rs | 25 +- src/crypto/mod.rs | 2 +- src/dns/client.rs | 3 + src/dns/transport.rs | 28 +- src/error/mod.rs | 10 +- src/log/mod.rs | 8 +- src/server/actor/tunnel.rs | 89 ++- src/server/actor/upload.rs | 103 ++-- src/server/constants.rs | 4 +- src/server/handlers.rs | 158 ++--- src/server/janitor.rs | 3 +- src/server/mod.rs | 17 +- src/server/stream.rs | 53 +- src/server/stream_registry.rs | 76 ++- src/shaper/mod.rs | 931 +++++++++++++++++++++++------- 30 files changed, 1717 insertions(+), 614 deletions(-) diff --git a/CONFIGURATION.md b/CONFIGURATION.md index ca36118..1074c92 100644 --- a/CONFIGURATION.md +++ b/CONFIGURATION.md @@ -1,38 +1,66 @@ -# Examples +# httproxy Configuration Reference -## Configuration Guidelines +Both binaries — `client` and `server` — are configured through a single +`config.toml` file passed with `-c` (default: `./config.toml`). Every +configuration key is validated at startup; unknown keys are rejected +(`deny_unknown_fields`), so a typo fails fast instead of being silently +ignored. -**Note**: Use a unique, random string for the `path` to evade network detection. Avoid predictable patterns like `/tunnel` or `/proxy`. +- [1. Command-line Utilities](#1-command-line-utilities) +- [2. Client Configuration](#2-client-configuration) +- [3. Server Configuration](#3-server-configuration) +- [4. Authentication Model](#4-authentication-model) +- [5. Traffic Shaping](#5-traffic-shaping) +- [6. Bypass Configuration](#6-bypass-configuration) +- [7. DNS Configuration](#7-dns-configuration) +- [8. Logging](#8-logging) +- [9. Resource Limits & Performance Tuning](#9-resource-limits--performance-tuning) +- [10. Deployment Behind Nginx](#10-deployment-behind-nginx) +- [11. Security & Operational Notes](#11-security--operational-notes) -## Token Generation +--- -Generate a secure bearer token. **The secret used here must match the `secret` in your server configuration.** +## 1. Command-line Utilities + +The `server` binary ships two subcommands for generating credentials. + +### 1.1 Token Generation + +Issue a bearer token (a signed JWT). **The `--secret` value must match the +`[auth] secret` of your server configuration.** ```bash ./server gen-token --secret --user --exp ``` -### Arguments -* `--secret` / `-s`: Secret key string used for signing. -* `--user` / `-u`: Username or subject identifier. -* `--exp` / `-e`: Expiration timestamp in Unix seconds. +| Argument | Required | Description | +| -------- | -------- | ------------------------------------------ | +| `--secret`, `-s` | yes | Signing key (must equal the server's `auth.secret`). | +| `--user`, `-u` | yes | Username / subject identifier carried in the token. | +| `--exp`, `-e` | yes | Expiration timestamp in Unix seconds. | -### Examples ```bash ./server gen-token --secret "my_secret_key" --user "admin" --exp 1768281600 ``` -## Keypair Generation +### 1.2 Keypair Generation -Generate an X25519 keypair to enable hybrid post-quantum encryption. **The public key will be used in the client configuration, while the private key must be kept secure on the server.** - -> **Note**: Encryption is optional. If the `public_key` is not configured in the client, encryption will be disabled. +Generate an X25519 keypair for end-to-end encryption. Place the **private +key** on the server (`[server] private_key`) and the **public key** on the +client (`[client] public_key`). ```bash ./server gen-key ``` -## Client Configuration +> **Encryption is optional.** If `[client] public_key` is absent, the tunnel +> is unencrypted at the application layer (transport TLS still applies). The +> private key must be kept secret; anyone holding it can decrypt tunnel +> traffic. + +--- + +## 2. Client Configuration `config.toml`: @@ -42,6 +70,9 @@ listen = "127.0.0.1:8080" remote = "https://your-server-domain/YOUR_SECRET_PATH" # address = "your-server-ip" # public_key = "your-public-key" +# max_connections = 1024 +# max_in_flight_bytes = 2097152 +# upload_concurrency = 128 # [client.auth] # username = "proxyuser" @@ -75,25 +106,43 @@ padding_range = [800, 1200] padding_threshold = 2000 ``` -## Bypass Configuration +### 2.1 `[client]` Reference -`bypass.json`: +| Field | Type | Required | Default | Description | +| --------------------- | ------- | -------- | ----------- | ----------- | +| `listen` | string | yes | — | Local listen address, e.g. `"127.0.0.1:8080"`. | +| `remote` | string | yes | — | Full server URL including the hidden path, e.g. `"https://host/secret"`. | +| `address` | string | no | host of `remote` | Overrides the resolved connection address (IP or hostname). The TLS SNI / Host header still uses the `remote` domain — useful when the server is reachable via a different IP than its certificate name. | +| `public_key` | string | no | — (no encryption) | Server's X25519 public key; enables end-to-end encryption. | +| `auth` | table | no | — | Optional local proxy authentication (see §2.2). | +| `max_connections` | integer | no | `1024` | Cap on concurrent local connections; excess connections are refused. Bounds client-side memory under load. | +| `max_in_flight_bytes` | integer | no | `2097152` | Cap on upload bytes in flight per tunnel. Trade-off between throughput and memory. **Must not exceed the server's reorder buffer (§9).** | +| `upload_concurrency` | integer | no | `128` | Cap on concurrent upload POSTs per tunnel. | -```json -{ - "domain_suffix": [ - "localhost" - ], - "ip_cidr": [ - "10.0.0.0/8", - "192.168.0.0/16", - "172.16.0.0/16", - "127.0.0.1/32" - ] -} -``` +### 2.2 `[client.auth]` — Local Proxy Authentication + +Optional; when present the local HTTP proxy requires these credentials from +its own clients (username/password basic auth). + +| Field | Type | Required | Description | +| ---------- | ------ | -------- | ---------------------------- | +| `username` | string | yes | Proxy username. | +| `password` | string | yes | Proxy password. | + +### 2.3 `[auth]` — Server Authentication + +| Field | Type | Required | Description | +| ------- | ------ | -------- | -------------------------------------------------- | +| `token` | string | yes | Bearer token presented to the server on every tunnel request (a JWT issued with `gen-token`). | -## Server Configuration +### 2.4 `[bypass]` + +Optional. Routes traffic matching the rules in the referenced JSON files +directly, outside the tunnel (§6). + +--- + +## 3. Server Configuration `config.toml`: @@ -102,19 +151,27 @@ padding_threshold = 2000 listen = "/dev/shm/httproxy.sock" path = "/YOUR_SECRET_PATH" # private_key = "your-private-key" +# max_tunnels = 1024 [auth] secret = "my_secret_key" +# [proxy] +# socks5 = "127.0.0.1:1080" + # [dns] +# upstream = "8.8.8.8:853" +# protocol = "dot" +# tls_domain = "dns.google" +# prefer_ipv6 = false # cache_size = 1024 # client_subnet = "1.2.3.4" -# prefer_ipv6 = false -# protocol = "dot" -# upstream = "8.8.8.8:853" - -# [proxy] -# socks5 = "127.0.0.1:1080" +# min_ttl = 30 +# max_ttl = 3600 +# swr_ttl = 3600 +# empty_ttl = 300 +# happy_eyeballs_delay_ms = 250 +# max_concurrent_queries = 1024 # [log] # file_path = "server.log" @@ -136,27 +193,100 @@ padding_range = [800, 1200] padding_threshold = 2000 ``` -## Traffic Shaping Configuration +### 3.1 `[server]` Reference + +| Field | Type | Required | Default | Description | +| ------------- | ------- | -------- | ------- | ----------- | +| `listen` | string | yes | — | Listen endpoint: a TCP address (`"0.0.0.0:443"`) or a Unix socket path (`"/dev/shm/httproxy.sock"`). | +| `path` | string | yes | — | Hidden service path; must match the path in the client's `remote` URL. | +| `private_key` | string | no | — | X25519 private key (from `gen-key`), paired with the client's `public_key`. | +| `max_tunnels` | integer | no | `1024` | Cap on concurrent tunnels; excess requests receive HTTP 503. Bounds server-side memory under load. | + +### 3.2 `[auth]` — JWT Signing Secret + +| Field | Type | Required | Description | +| -------- | ------ | -------- | ----------- | +| `secret` | string | yes | Key used to validate client bearer tokens. Must match the `--secret` used with `gen-token`. | + +### 3.3 `[proxy]` — Upstream SOCKS5 + +| Field | Type | Required | Description | +| -------- | ------ | -------- | ----------- | +| `socks5` | string | no | Optional upstream SOCKS5 proxy; all tunneled connections are routed through it. | + +### 3.4 `[traffic_shaping]` + +See [§5 Traffic Shaping](#5-traffic-shaping). Both sides must agree on +`encoding_type` and `max_download_bytes`. -The `traffic_shaping` field allows you to configure padding for outgoing packets to obfuscate traffic patterns. It consists of a `global` configuration and an array of `stages` for more granular control. +--- -> **Constraint**: The maximum padding must not exceed `MAX_RAW_PAYLOAD - raw_len` (up to 16383). Any `padding_range[1]` value beyond this threshold will be silently truncated. +## 4. Authentication Model -> **Important**: Stages are processed sequentially based on the packet sequence. If you want a specific configuration for the 3rd packet only, you MUST define stages for the 1st and 2nd packets as placeholders. +1. **Tunnel authentication (mandatory):** every tunnel request carries the + client's bearer token, signed with the server's `auth.secret`. Requests + without a valid, unexpired token are rejected. +2. **End-to-end encryption (optional):** when `public_key`/`private_key` are + configured, a hybrid X25519 + ML-KEM key exchange seals the tunnel in + addition to the transport TLS layer. +3. **Local proxy authentication (optional):** `[client.auth]` gates access to + the local proxy endpoint itself. -### PaddingConfig (for `global`) +--- -- `padding_threshold`: (usize) If the actual data length of a packet is below this threshold, padding will be applied. -- `padding_range`: ([usize; 2]) A tuple specifying the minimum and maximum random padding length to add when padding is applied. +## 5. Traffic Shaping -### StageConfig (for `stages` array) +Traffic shaping obfuscates the tunnel's packet patterns in two ways: -Each stage can override the `global` padding configuration for a specific range of packets. +- **Padding** — random padding appended to each frame before it is sent; +- **Inter-frame delay jitter** — a randomized (log-normal distributed) delay + between emitted frames, applied automatically. -- `count`: (Option) The last packet number for this stage (1-indexed). -- `count_range`: (Option<[usize; 2]>) A range where the second value (hi) is used as the stage's end point. +Padding is driven by the `[traffic_shaping]` table, shared verbatim by both +sides (the client pads, the server strips). -**Example:** +### 5.1 `[traffic_shaping]` Reference + +| Field | Type | Required | Default | Description | +| ------------------- | ------- | -------- | ----------- | ----------- | +| `global` | table | yes | — | Default padding behavior (§5.2). | +| `stages` | array | no | `[]` | Per-packet-range overrides (§5.3). | +| `encoding_type` | string | no | `"binary"` | Wire encoding: `"binary"` or `"json"`. JSON wraps frames as `{"data":""}` lines (~14% larger, but appears as ordinary JSON traffic). | +| `max_download_bytes` | integer | no | — (stream) | Optional download rotation threshold. When set, the download stream is rotated into segments of this many bytes; unset streams continuously without segmentation (lower overhead). | + +### 5.2 `global` — Default Padding + +| Field | Type | Required | Description | +| ------------------- | ------ | -------- | ----------- | +| `padding_threshold` | usize | yes | If the frame's data length is below this value, padding is applied. | +| `padding_range` | [usize; 2] | yes | Random padding length drawn uniformly from this inclusive range. | + +### 5.3 `stages` — Per-packet Overrides + +Each stage overrides the global padding for a contiguous, 1-indexed range of +packets. Stages are applied in the order in which the packet sequence number +enters their range. + +| Field | Type | Required | Description | +| ------------------- | ------------- | -------- | ----------- | +| `count` | usize | one of `count`/`count_range` | Last packet number covered by this stage (1-indexed). | +| `count_range` | [usize; 2] | one of `count`/`count_range` | Range of packet numbers; the upper bound (hi) is the stage's end point. | +| `padding_threshold` | usize | yes | Same semantics as the global counterpart. | +| `padding_range` | [usize; 2] | yes | Same semantics as the global counterpart. | + +**Constraints:** + +- `padding_range[0]` must be ≤ `padding_range[1]`; the configuration is + rejected otherwise. +- The effective padding is capped at the frame sealing threshold minus the + data length (the threshold is ~16 KiB of frame payload, varying slightly + with the encoding/encryption mode); a `padding_range[1]` beyond that is + silently truncated. +- Stages are evaluated in packet order. To configure a specific packet + (say the 3rd), the 1st and 2nd must be covered by earlier stages or the + global defaults — define placeholder stages when needed. + +### 5.4 Example ```toml [traffic_shaping.global] @@ -179,15 +309,111 @@ padding_range = [1500, 3000] padding_threshold = 3000 ``` -**In this example:** -- **Global Behavior**: By default, if a packet's data length is less than 1500 bytes, a random padding between 0 and 3000 bytes is added. -- **1st Packet**: The first packet will have exactly 5000 bytes of padding added, as its data length is almost certainly below the 6000-byte threshold. -- **2nd Packet**: The second packet will have a random padding between 1000 and 5000 bytes if its data length is below 3000 bytes. -- **3rd to 8th Packets**: These packets will have a random padding between 1500 and 3000 bytes if their data length is below 3000 bytes. +With the above configuration: -## Nginx Configuration +- **Default:** frames shorter than 1500 bytes receive 0–3000 bytes of padding. +- **Packet 1:** exactly 5000 bytes of padding (its data is almost certainly + below the 6000-byte threshold). +- **Packet 2:** 1000–5000 bytes of padding when shorter than 3000 bytes. +- **Packets 3–8:** 1500–3000 bytes of padding when shorter than 3000 bytes. -To hide the proxy server behind Nginx, use the following configuration: +--- + +## 6. Bypass Configuration + +Optional. Traffic whose destination matches a bypass rule is sent directly, +outside the tunnel. Rules are loaded from JSON files referenced by +`[bypass] bypass_files`: + +```json +{ + "domain_suffix": [ + "localhost" + ], + "ip_cidr": [ + "10.0.0.0/8", + "192.168.0.0/16", + "172.16.0.0/16", + "127.0.0.1/32" + ] +} +``` + +| Field | Type | Description | +| -------------- | -------- | ------------------------------------ | +| `domain_suffix` | string[] | Domains matched by suffix (e.g. `"localhost"` also matches `"api.localhost"`). | +| `ip_cidr` | string[] | CIDR blocks matched against the destination IP. | + +The `[bypass]` section itself is optional and defaults to empty. + +--- + +## 7. DNS Configuration + +Optional server-side DNS resolver (`[dns]`). When omitted, the system +resolver is used. + +| Field | Type | Required | Default | Description | +| ------------------------- | ------- | -------- | ---------- | ----------- | +| `upstream` | string | yes | — | Upstream resolver, e.g. `"8.8.8.8:853"` (TCP/TLS when `protocol = "dot"`). | +| `protocol` | string | no | `"udp"` | `"udp"` or `"dot"` (DNS over TLS). | +| `tls_domain` | string | no | — | SNI/hostname for DoT; defaults to the `upstream` IP when absent. | +| `prefer_ipv6` | bool | no | `false` | Prefer AAAA records when connecting. | +| `cache_size` | integer | no | `1024` | Number of cached responses. | +| `client_subnet` | string | no | — | EDNS Client Subnet hint. | +| `min_ttl` / `max_ttl` | integer | no | `30` / `3600` | Clamp for cached TTLs. | +| `swr_ttl` | integer | no | `3600` | TTL for stale-while-revalidate answers. | +| `empty_ttl` | integer | no | `300` | TTL for empty (NODATA) responses. | +| `happy_eyeballs_delay_ms` | integer | no | `250` | Happy-Eyeballs fallback delay. | +| `max_concurrent_queries` | integer | no | `1024` | Cap on in-flight queries to the upstream. | + +--- + +## 8. Logging + +The optional `[log]` section is shared by both binaries. Logs are emitted as +newline-delimited JSON (ANSI color when writing to a terminal). + +| Field | Type | Required | Default | Description | +| ------------ | ------- | -------- | ------- | ----------- | +| `file_path` | string | no | — (stdout) | Log file path; omitted logs to stdout. | +| `level` | string | no | `"info"` | `trace`, `debug`, `info`, `warn` or `error`. The `RUST_LOG` environment variable takes precedence. | +| `max_backups` | integer | no | `7` | Number of rotated log files kept. | + +--- + +## 9. Resource Limits & Performance Tuning + +| Setting | Default | Bounds | Effect | +| ---------------------- | --------- | ------ | ------ | +| `[server] max_tunnels` | `1024` | — | Concurrent tunnels; excess → HTTP 503. | +| `[client] max_connections` | `1024` | — | Concurrent local connections; excess are refused. | +| `[client] upload_concurrency` | `128` | — | Concurrent upload POSTs per tunnel. | +| `[client] max_in_flight_bytes` | `2097152` (2 MiB) | ≤ server reorder buffer | Upload bytes in flight per tunnel; the throughput/memory trade-off knob. | + +**Important:** the server maintains a 2 MiB per-tunnel reorder buffer. +`max_in_flight_bytes` **must not exceed 2 MiB**; a larger value causes hard +upload errors under HTTP/2 frame reordering. The default is already at the +ceiling — lower it only if you need to cap client memory. + +Tuning notes: + +- Raising `max_in_flight_bytes` (up to 2 MiB) and `upload_concurrency` + increases single-tunnel upload throughput at the cost of memory; lowering + them reduces peak memory. +- `max_connections` / `max_tunnels` are the primary bounds on total memory + consumption under concurrent load. +- `encoding_type = "json"` trades ~14% wire overhead (base122 expansion) for + a JSON-shaped traffic profile. +- For maximum download throughput leave `max_download_bytes` unset (direct + streaming); setting it enables download rotation, which is primarily for + traffic-shape purposes. + +--- + +## 10. Deployment Behind Nginx + +To front the proxy with Nginx (TLS termination, request obfuscation): ```nginx server { @@ -208,7 +434,7 @@ server { proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for; proxy_request_buffering off; proxy_http_version 1.1; - client_max_body_size 1m; + client_max_body_size 4m; proxy_buffering off; proxy_buffer_size 16k; proxy_buffers 2 16k; @@ -226,3 +452,35 @@ upstream httproxy_backend { keepalive 32; } ``` + +Key points: + +- `proxy_request_buffering off` and `proxy_buffering off` keep the tunnel + streaming (required for low latency and full duplex throughput). +- Long `proxy_read/send_timeout`s prevent idle tunnel teardown. +- `client_max_body_size` must be at least the largest upload batch: batches + reach `min(1 MiB, max_in_flight_bytes)` plus one frame of slack + (~16 KiB), so size it accordingly (the 4m above covers any + `max_in_flight_bytes` up to the 2 MiB ceiling). +- The Unix socket upstream (`listen = "/dev/shm/httproxy.sock"`) avoids + loopback TCP overhead; the server may equally listen on a TCP port instead. + +--- + +## 11. Security & Operational Notes + +- **Use a unique, random `path`.** Avoid predictable patterns like `/tunnel` + or `/proxy` — the path doubles as a capability token. +- **Use a strong, random `secret`** and rotate it with token expiry. +- **Protect the server's `private_key`**; it decrypts all tunnel traffic. +- **Expire tokens.** Keep `--exp` short enough that a leaked token has + limited value; re-issue on rotation. +- **Validate the deployment** with the compiled binaries: + + ```bash + ./server -c server.toml + ./client -c client.toml + ``` + + Errors (unknown keys, invalid values, key mismatches) are reported at + startup — a clean start means the configuration is consistent. diff --git a/Cargo.lock b/Cargo.lock index b0d71aa..35f6819 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -250,9 +250,9 @@ dependencies = [ [[package]] name = "base122-fast" -version = "0.1.3" +version = "0.1.4" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "2361be42fbd11eefcef8fec542469d77cdea96a05eb92c3d7e61c20bffa89f3d" +checksum = "5e00b499c3ae0a45a400cfe5a91cccd56b24d9ed0d0831e0030f41dec44fa228" [[package]] name = "base64" diff --git a/Cargo.toml b/Cargo.toml index e6cdffb..6d560e6 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -22,7 +22,7 @@ strip = true aes-gcm = "0.10.3" anyhow = "1.0" axum = {version = "0.8.9", features = ["http2", "macros"]} -base122-fast = "0.1.3" +base122-fast = "0.1.4" base64 = "0.22.1" bytes = "1.11" clap = {version = "4.6.1", features = ["derive"]} diff --git a/src/bin/client.rs b/src/bin/client.rs index d697730..5f594db 100644 --- a/src/bin/client.rs +++ b/src/bin/client.rs @@ -5,7 +5,8 @@ use std::sync::Arc; use std::sync::atomic::{AtomicU64, Ordering}; use std::time::Duration; use tokio::net::TcpListener; -use tracing::{Instrument, error_span, info, warn}; +use tokio::sync::Semaphore; +use tracing::{Instrument, debug, error_span, info, warn}; static NEXT_SPAN_ID: AtomicU64 = AtomicU64::new(1); @@ -54,6 +55,8 @@ async fn main() -> anyhow::Result<()> { info!(listen = %addr, "proxy listening"); + let conn_sem = Arc::new(Semaphore::new(state.max_connections)); + loop { let (socket, peer) = match listener.accept().await { Ok(conn) => conn, @@ -63,6 +66,18 @@ async fn main() -> anyhow::Result<()> { } }; + let permit = match conn_sem.clone().try_acquire_owned() { + Ok(p) => p, + Err(_) => { + debug!( + client = %peer, + max_connections = state.max_connections, + "connection rejected: admission limit reached" + ); + continue; + } + }; + let http_client = Arc::clone(&http_client); let state = Arc::clone(&state); @@ -70,6 +85,7 @@ async fn main() -> anyhow::Result<()> { tokio::spawn( async move { + let _permit = permit; if let Err(e) = httproxy::client::connection::handle_connection_actor( socket, http_client, diff --git a/src/client/actor/connection.rs b/src/client/actor/connection.rs index 903ad2e..9981516 100644 --- a/src/client/actor/connection.rs +++ b/src/client/actor/connection.rs @@ -69,7 +69,7 @@ impl ClientConnectionActor { let (method, header_len, url) = loop { let (method, header_len, url, proxy_auth_header) = tokio::time::timeout( PROXY_REQUEST_PARSE_TIMEOUT, - proxy::parse_proxy_request(&mut read_half, buf), + proxy::parse_proxy_request(&mut read_half, buf, state.proxy_auth.is_some()), ) .await .map_err(|_| anyhow::anyhow!("proxy request parse timeout"))??; diff --git a/src/client/actor/download_loop.rs b/src/client/actor/download_loop.rs index 5296674..9e11989 100644 --- a/src/client/actor/download_loop.rs +++ b/src/client/actor/download_loop.rs @@ -6,14 +6,16 @@ use std::sync::Arc; use tokio::io::AsyncWriteExt; use tokio::sync::oneshot; use tracing::{Instrument, warn}; +use uuid::Uuid; use super::super::state::SharedState; use crate::client::constants::{ DECODE_BUF_CAPACITY, DOWNLOAD_CONNECT_TIMEOUT, PREFETCH_LEAD_BYTES, PREFETCH_ROTATE_TIMEOUT, + WRITE_BUF_CAPACITY, WRITE_FLUSH_TIMEOUT, }; use crate::client::utils; use crate::crypto::AesFrameCipher; -use crate::shaper::{self, EncodingType, FrameCipher}; +use crate::shaper::{self, DecodedFrame, EncodingType, FrameCipher}; enum Phase { Streaming { @@ -34,7 +36,7 @@ pub struct DownloadLoopActor { write_half: tokio::net::tcp::OwnedWriteHalf, cipher: Option>, encoding: EncodingType, - stream_id: String, + stream_id: Uuid, http_client: Arc, state: Arc, max_bytes: Option, @@ -46,7 +48,7 @@ impl DownloadLoopActor { initial_response: wreq::Response, write_half: tokio::net::tcp::OwnedWriteHalf, cipher: Option>, - stream_id: String, + stream_id: Uuid, http_client: Arc, state: Arc, ) -> Self { @@ -58,7 +60,7 @@ impl DownloadLoopActor { let use_prefetch = prefetch_at.is_some_and(|at| at > 0); let (prefetch_trigger, prefetch_rx) = if rotate_enabled && use_prefetch { - let (tx, rx) = spawn_prefetch_continuation(&http_client, &state, &stream_id); + let (tx, rx) = spawn_prefetch_continuation(&http_client, &state, stream_id); (Some(tx), Some(rx)) } else { (None, None) @@ -148,18 +150,27 @@ impl DownloadLoopActor { prefetch_rx: Option>>, ) -> Result { let response = if let Some(rx) = prefetch_rx { - match tokio::time::timeout(PREFETCH_ROTATE_TIMEOUT, rx).await { + tokio::pin!(rx); + match tokio::time::timeout(PREFETCH_ROTATE_TIMEOUT, &mut rx).await { Ok(Ok(Ok(resp))) => resp, Ok(Ok(Err(_))) | Ok(Err(_)) => { - send_continue_request(&self.http_client, &self.state, &self.stream_id).await? + send_continue_request(&self.http_client, &self.state, self.stream_id).await? } Err(_elapsed) => { - warn!("prefetch timed out, falling back to synchronous continue"); - send_continue_request(&self.http_client, &self.state, &self.stream_id).await? + warn!("prefetch slow, waiting one more window"); + match tokio::time::timeout(PREFETCH_ROTATE_TIMEOUT, &mut rx).await { + Ok(Ok(Ok(resp))) => resp, + _ => { + return Err(anyhow!( + "prefetch continuation timed out after {}s", + 2 * PREFETCH_ROTATE_TIMEOUT.as_secs() + )); + } + } } } } else { - send_continue_request(&self.http_client, &self.state, &self.stream_id).await? + send_continue_request(&self.http_client, &self.state, self.stream_id).await? }; let use_prefetch = self @@ -168,7 +179,7 @@ impl DownloadLoopActor { .is_some_and(|at| at > 0); let (prefetch_trigger, prefetch_rx) = if use_prefetch { let (tx, rx) = - spawn_prefetch_continuation(&self.http_client, &self.state, &self.stream_id); + spawn_prefetch_continuation(&self.http_client, &self.state, self.stream_id); (Some(tx), Some(rx)) } else { (None, None) @@ -186,6 +197,36 @@ impl DownloadLoopActor { prefetch_rx, }) } + + #[inline] + fn handle_frame( + frame: DecodedFrame, + write_buf: &mut BytesMut, + scratch: &BytesMut, + expected_seq: &mut u64, + ) -> Result<(), anyhow::Error> { + match frame { + DecodedFrame::InScratch { seq, start, end } => { + if seq != *expected_seq { + return Err(anyhow!( + "download frame seq {seq} out of order, expected {expected_seq}" + )); + } + *expected_seq += 1; + write_buf.extend_from_slice(&scratch[start..end]); + } + DecodedFrame::Owned { seq, data } => { + if seq != *expected_seq { + return Err(anyhow!( + "download frame seq {seq} out of order, expected {expected_seq}" + )); + } + *expected_seq += 1; + write_buf.extend_from_slice(&data); + } + } + Ok(()) + } } async fn download_single_response( @@ -198,34 +239,118 @@ async fn download_single_response( mut prefetch_trigger: Option>, ) -> Result<(u64, u64)> { let mut buffer = BytesMut::with_capacity(DECODE_BUF_CAPACITY); + let mut scratch = BytesMut::new(); + let mut json_scratch = Vec::new(); + let mut write_buf = BytesMut::with_capacity(WRITE_BUF_CAPACITY); + let flush_deadline = tokio::time::sleep(WRITE_FLUSH_TIMEOUT); + tokio::pin!(flush_deadline); let mut data_stream = response.into_data_stream(); let mut bytes_received: u64 = 0; - while let Some(chunk) = data_stream.next().await { - let chunk = chunk.context("response read error")?; - bytes_received += chunk.len() as u64; + loop { + tokio::select! { + chunk = data_stream.next() => { + let Some(chunk) = chunk else { break }; + let chunk = chunk.context("response read error")?; + bytes_received += chunk.len() as u64; - if let Some(at) = prefetch_at - && bytes_received >= at - && let Some(tx) = prefetch_trigger.take() - { - let _ = tx.send(()); - } + if let Some(at) = prefetch_at + && bytes_received >= at + && let Some(tx) = prefetch_trigger.take() + { + let _ = tx.send(()); + } - buffer.extend_from_slice(&chunk); - while let Some((seq, frame_data, start, end)) = - shaper::decode_frame_owned(&mut buffer, cipher, encoding)? - { - if seq != expected_seq { - return Err(anyhow!( - "download frame seq {seq} out of order, expected {expected_seq}" - )); + if buffer.is_empty() { + match chunk.try_into_mut() { + Ok(mut chunk_mut) => { + while let Some(frame) = shaper::decode_frame( + &mut chunk_mut, + &mut scratch, + &mut json_scratch, + cipher, + encoding, + )? { + DownloadLoopActor::handle_frame( + frame, + &mut write_buf, + &scratch, + &mut expected_seq, + )?; + if write_buf.len() >= WRITE_BUF_CAPACITY { + write_half.write_all(&write_buf).await?; + write_buf.clear(); + flush_deadline.as_mut().reset( + tokio::time::Instant::now() + WRITE_FLUSH_TIMEOUT, + ); + } + } + buffer = chunk_mut; + } + Err(chunk) => { + buffer.extend_from_slice(&chunk); + while let Some(frame) = shaper::decode_frame( + &mut buffer, + &mut scratch, + &mut json_scratch, + cipher, + encoding, + )? { + DownloadLoopActor::handle_frame( + frame, + &mut write_buf, + &scratch, + &mut expected_seq, + )?; + if write_buf.len() >= WRITE_BUF_CAPACITY { + write_half.write_all(&write_buf).await?; + write_buf.clear(); + flush_deadline.as_mut().reset( + tokio::time::Instant::now() + WRITE_FLUSH_TIMEOUT, + ); + } + } + } + } + } else { + buffer.extend_from_slice(&chunk); + while let Some(frame) = shaper::decode_frame( + &mut buffer, + &mut scratch, + &mut json_scratch, + cipher, + encoding, + )? { + DownloadLoopActor::handle_frame( + frame, + &mut write_buf, + &scratch, + &mut expected_seq, + )?; + if write_buf.len() >= WRITE_BUF_CAPACITY { + write_half.write_all(&write_buf).await?; + write_buf.clear(); + flush_deadline + .as_mut() + .reset(tokio::time::Instant::now() + WRITE_FLUSH_TIMEOUT); + } + } + } + } + _ = &mut flush_deadline, if !write_buf.is_empty() => { + write_half.write_all(&write_buf).await?; + write_buf.clear(); + flush_deadline + .as_mut() + .reset(tokio::time::Instant::now() + WRITE_FLUSH_TIMEOUT); } - expected_seq += 1; - write_half.write_all(&frame_data[start..end]).await?; } } + if !write_buf.is_empty() { + write_half.write_all(&write_buf).await?; + } + if !buffer.is_empty() { warn!( remaining = buffer.len(), @@ -238,7 +363,7 @@ async fn download_single_response( async fn send_continue_request( http_client: &wreq::Client, state: &SharedState, - stream_id: &str, + stream_id: Uuid, ) -> Result { let mut cookie = String::new(); utils::build_stream_cookie(&mut cookie, stream_id); @@ -259,7 +384,7 @@ async fn send_continue_request( fn spawn_prefetch_continuation( http_client: &Arc, state: &Arc, - stream_id: &str, + stream_id: Uuid, ) -> ( oneshot::Sender<()>, oneshot::Receiver>, @@ -268,13 +393,12 @@ fn spawn_prefetch_continuation( let (result_tx, result_rx) = oneshot::channel(); let pre_client = Arc::clone(http_client); let pre_state = Arc::clone(state); - let pre_stream_id = stream_id.to_owned(); tokio::spawn( async move { if trigger_rx.await.is_err() { return; } - match send_continue_request(&pre_client, &pre_state, &pre_stream_id).await { + match send_continue_request(&pre_client, &pre_state, stream_id).await { Ok(resp) => { let _ = result_tx.send(Ok(resp)); } diff --git a/src/client/actor/upload_loop.rs b/src/client/actor/upload_loop.rs index 89450c8..6e1e771 100644 --- a/src/client/actor/upload_loop.rs +++ b/src/client/actor/upload_loop.rs @@ -1,44 +1,39 @@ use anyhow::{Context, Result, anyhow}; -use bytes::{BufMut, Bytes, BytesMut}; -use futures::FutureExt; -use futures::StreamExt; +use bytes::{Bytes, BytesMut}; +use std::io; use std::pin::Pin; use std::sync::Arc; +use std::task::{Context as TaskContext, Poll, Waker}; use tokio::io::AsyncReadExt; use tokio::sync::{OwnedSemaphorePermit, Semaphore}; use tokio::task::JoinSet; use tracing::Instrument; +use uuid::Uuid; use super::super::state::SharedState; use crate::client::constants::{ - BATCH_BUF_INITIAL_CAPACITY, MAX_BATCH_BYTES, MAX_IN_FLIGHT_BYTES, UPLOAD_CONCURRENCY, - UPLOAD_REQUEST_TIMEOUT, + BATCH_BUF_INITIAL_CAPACITY, MAX_BATCH_BYTES, UPLOAD_REQUEST_TIMEOUT, }; use crate::client::utils; use crate::crypto::AesFrameCipher; -use crate::shaper::{self, FrameCipher}; +use crate::shaper::{self, SealInto}; -type ShaperStream = Pin> + Send>>; +type ShaperStream = Pin>; enum Phase { - Batching { - batch_buf: BytesMut, - bytes_permits: Vec, - leftover: Option, - }, - Draining { - inflight: usize, - }, + Batching { batch_buf: BytesMut }, + Draining { inflight: usize }, Done, } pub struct UploadLoopActor { http_client: Arc, state: Arc, - stream_id: String, + stream_id: Uuid, shaped: ShaperStream, request_sem: Arc, bytes_sem: Arc, + max_batch_bytes: usize, tasks: JoinSet>, phase: Phase, } @@ -50,30 +45,31 @@ impl UploadLoopActor { initial_payload: Bytes, read_half: tokio::net::tcp::OwnedReadHalf, cipher: Option>, - stream_id: String, + stream_id: Uuid, start_seq: u64, ) -> Self { let reader = AsyncReadExt::chain(std::io::Cursor::new(initial_payload), read_half); - let traffic_cipher: Option> = - cipher.map(|c| c as Arc); + let traffic_cipher: Option> = + cipher.map(|c| c as Arc); let shaped: ShaperStream = Box::pin(shaper::TrafficShaper::with_seq( reader, - state.traffic_config.clone(), + &state.resolved_traffic, traffic_cipher, start_seq, )); + let upload_concurrency = state.upload_concurrency; + let max_in_flight_bytes = state.max_in_flight_bytes; Self { http_client, state, stream_id, shaped, - request_sem: Arc::new(Semaphore::new(UPLOAD_CONCURRENCY)), - bytes_sem: Arc::new(Semaphore::new(MAX_IN_FLIGHT_BYTES)), + request_sem: Arc::new(Semaphore::new(upload_concurrency)), + bytes_sem: Arc::new(Semaphore::new(max_in_flight_bytes)), + max_batch_bytes: MAX_BATCH_BYTES.min(max_in_flight_bytes), tasks: JoinSet::new(), phase: Phase::Batching { batch_buf: BytesMut::with_capacity(BATCH_BUF_INITIAL_CAPACITY), - bytes_permits: vec![], - leftover: None, }, } } @@ -81,11 +77,7 @@ impl UploadLoopActor { pub async fn run(mut self) -> Result<()> { loop { self.phase = match std::mem::replace(&mut self.phase, Phase::Done) { - Phase::Batching { - batch_buf, - bytes_permits, - leftover, - } => self.do_batching(batch_buf, bytes_permits, leftover).await?, + Phase::Batching { batch_buf } => self.do_batching(batch_buf).await?, Phase::Draining { inflight } => { self.do_drain(inflight).await?; return Ok(()); @@ -95,34 +87,25 @@ impl UploadLoopActor { } } - async fn do_batching( + fn poll_seal( &mut self, - mut batch_buf: BytesMut, - mut bytes_permits: Vec, - mut leftover: Option, - ) -> Result { - if let Some(data) = leftover.take() { - let size = data.len() as u32; - let permit = self - .bytes_sem - .clone() - .acquire_many_owned(size) - .await - .map_err(|_| anyhow!("bytes semaphore closed"))?; - batch_buf.put_slice(&data); - bytes_permits.push(permit); - } + cx: &mut TaskContext<'_>, + batch_buf: &mut BytesMut, + ) -> Poll>> { + self.shaped.as_mut().poll_seal_into(cx, batch_buf) + } + + async fn do_batching(&mut self, mut batch_buf: BytesMut) -> Result { let mut stream_ended = false; + if batch_buf.is_empty() { + let seal = + std::future::poll_fn(|cx| self.shaped.as_mut().poll_seal_into(cx, &mut batch_buf)); tokio::select! { - frame = self.shaped.next() => match frame { - Some(Ok((_seq, data))) => { - let size = data.len() as u32; - let permit = self.bytes_sem.clone().acquire_many_owned(size).await.map_err(|_| anyhow!("bytes semaphore closed"))?; - batch_buf.put_slice(&data); bytes_permits.push(permit); - } - Some(Err(e)) => return Err(e.into()), - None => stream_ended = true, + r = seal => match r { + Ok(Some(_)) => {} + Ok(None) => stream_ended = true, + Err(e) => return Err(e.into()), }, result = self.tasks.join_next(), if !self.tasks.is_empty() => match result { Some(Ok(Ok(()))) | None => {} @@ -131,62 +114,54 @@ impl UploadLoopActor { }, } } - while !stream_ended { - match self.shaped.next().now_or_never() { - Some(Some(Ok((_seq, data)))) => { - let frame_size = data.len(); - if batch_buf.len() + frame_size > MAX_BATCH_BYTES { - leftover = Some(data); + + if !stream_ended { + let waker = Waker::noop(); + let mut cx = TaskContext::from_waker(waker); + while batch_buf.len() < self.max_batch_bytes { + match self.poll_seal(&mut cx, &mut batch_buf) { + Poll::Ready(Ok(Some(_))) => {} + Poll::Ready(Ok(None)) => { + stream_ended = true; break; } - match self - .bytes_sem - .clone() - .try_acquire_many_owned(frame_size as u32) - { - Ok(permit) => { - batch_buf.put_slice(&data); - bytes_permits.push(permit); - } - Err(_) => { - leftover = Some(data); - break; - } - } + Poll::Ready(Err(e)) => return Err(e.into()), + Poll::Pending => break, } - Some(Some(Err(e))) => return Err(e.into()), - Some(None) => stream_ended = true, - None => break, } } + if batch_buf.is_empty() { return if stream_ended { Ok(Phase::Draining { inflight: self.tasks.len(), }) } else { - Ok(Phase::Batching { - batch_buf, - bytes_permits, - leftover, - }) + Ok(Phase::Batching { batch_buf }) }; } + let req_permit = self .request_sem .clone() .acquire_owned() .await .map_err(|_| anyhow!("request semaphore closed"))?; + let bytes_permit: OwnedSemaphorePermit = self + .bytes_sem + .clone() + .acquire_many_owned(batch_buf.len() as u32) + .await + .map_err(|_| anyhow!("bytes semaphore closed"))?; let body = batch_buf.freeze(); let http_client = Arc::clone(&self.http_client); let state_ref = Arc::clone(&self.state); - let stream_id = self.stream_id.clone(); + let stream_id = self.stream_id; self.tasks.spawn( async move { let _req_guard = req_permit; - let _bytes = bytes_permits; - send_upload_post(&http_client, &state_ref, body, &stream_id).await + let _bytes = bytes_permit; + send_upload_post(&http_client, &state_ref, body, stream_id).await } .instrument(tracing::Span::current()), ); @@ -204,8 +179,6 @@ impl UploadLoopActor { } else { Ok(Phase::Batching { batch_buf: BytesMut::with_capacity(BATCH_BUF_INITIAL_CAPACITY), - bytes_permits: vec![], - leftover, }) } } @@ -230,7 +203,7 @@ async fn send_upload_post( http_client: &wreq::Client, state: &SharedState, body: Bytes, - stream_id: &str, + stream_id: Uuid, ) -> Result<()> { debug_assert!(!body.is_empty(), "empty upload body"); let mut cookie = String::new(); diff --git a/src/client/connection.rs b/src/client/connection.rs index 7006e13..4c0ce8e 100644 --- a/src/client/connection.rs +++ b/src/client/connection.rs @@ -19,9 +19,9 @@ pub(crate) async fn handle_plain_proxy( payload: Bytes, target_host: &str, ) -> Result<()> { - let stream_id = uuid::Uuid::new_v4().to_string(); + let stream_id = uuid::Uuid::new_v4(); let mut cookie = String::new(); - utils::build_stream_cookie(&mut cookie, &stream_id); + utils::build_stream_cookie(&mut cookie, stream_id); let (early_data, remaining_payload, frames_sent) = utils::encode_initial_payload( &payload, @@ -54,7 +54,7 @@ pub(crate) async fn handle_plain_proxy( remaining_payload, read_half, None, - stream_id.clone(), + stream_id, frames_sent, ); let upload_task = diff --git a/src/client/constants.rs b/src/client/constants.rs index d9c8aef..45ee04d 100644 --- a/src/client/constants.rs +++ b/src/client/constants.rs @@ -13,12 +13,18 @@ pub const MAX_BATCH_BYTES: usize = 1024 * 1024; pub const BATCH_BUF_INITIAL_CAPACITY: usize = 8192; pub const MAX_IN_FLIGHT_BYTES: usize = 2 * 1024 * 1024; pub const UPLOAD_CONCURRENCY: usize = 128; +pub const MIN_IN_FLIGHT_BYTES: usize = 20 * 1024; + +pub const MAX_LOCAL_CONNECTIONS: usize = 1024; pub const PREFETCH_LEAD_BYTES: u64 = 20 * 1024 * 1024; pub const PREFETCH_ROTATE_TIMEOUT: Duration = Duration::from_secs(5); pub const DECODE_BUF_CAPACITY: usize = 16 * 1024 + 2396; +pub const WRITE_BUF_CAPACITY: usize = 64 * 1024; +pub const WRITE_FLUSH_TIMEOUT: Duration = Duration::from_millis(2); + pub const MASTER_RESUME_WINDOW_SECS: u64 = 1170; pub const PADDING_POOL: &[u8] = b"padding=XXXXXXXXXXXXXXXXXXXXXXXXXX"; diff --git a/src/client/handshake.rs b/src/client/handshake.rs index 2b4f46b..56d19e9 100644 --- a/src/client/handshake.rs +++ b/src/client/handshake.rs @@ -91,7 +91,7 @@ pub async fn try_pq_connect( .context("session resumption POST failed")?; if response.status().as_u16() == 428 { - let _ = response.bytes().await; + let _ = tokio::time::timeout(DOWNLOAD_CONNECT_TIMEOUT, response.bytes()).await; return Err(anyhow::Error::new(RehandshakeRequired( "server requests re-handshake (428)".into(), ))); @@ -109,7 +109,6 @@ pub async fn try_pq_connect( let upload_client = Arc::clone(http_client); let upload_state = Arc::clone(state); let upload_cipher_clone = Arc::clone(&upload_cipher); - let stream_id_str = stream_id.to_string(); let upload_actor = UploadLoopActor::new( upload_client.clone(), @@ -117,7 +116,7 @@ pub async fn try_pq_connect( remaining_payload, read_half, Some(upload_cipher_clone), - stream_id_str.clone(), + stream_id, frames_sent, ); let upload_task = @@ -127,7 +126,7 @@ pub async fn try_pq_connect( response, write_half, Some(download_cipher), - stream_id_str, + stream_id, Arc::clone(http_client), Arc::clone(state), ); diff --git a/src/client/mod.rs b/src/client/mod.rs index 8b89df2..2a43ddc 100644 --- a/src/client/mod.rs +++ b/src/client/mod.rs @@ -6,8 +6,10 @@ pub mod proxy; pub mod state; pub mod utils; +use crate::client::constants::{MAX_IN_FLIGHT_BYTES, MAX_LOCAL_CONNECTIONS, UPLOAD_CONCURRENCY}; use crate::config::ClientTopConfig; use crate::crypto; +use crate::shaper::ResolvedShaperConfig; use anyhow::{Context, Result}; use base64::Engine; @@ -19,6 +21,21 @@ pub fn build_state(cfg: &ClientTopConfig) -> Result> { .validate() .context("invalid traffic_shaping config")?; + let max_in_flight_bytes = cfg + .client + .max_in_flight_bytes + .unwrap_or(MAX_IN_FLIGHT_BYTES); + let upload_concurrency = cfg.client.upload_concurrency.unwrap_or(UPLOAD_CONCURRENCY); + if max_in_flight_bytes < crate::client::constants::MIN_IN_FLIGHT_BYTES { + return Err(anyhow::anyhow!( + "max_in_flight_bytes ({max_in_flight_bytes}) must be at least {} (one maximum-size encoded frame)", + crate::client::constants::MIN_IN_FLIGHT_BYTES + )); + } + if upload_concurrency == 0 { + return Err(anyhow::anyhow!("upload_concurrency must be at least 1")); + } + let bypass = if cfg.bypass.bypass_files.is_empty() { None } else { @@ -54,11 +71,15 @@ pub fn build_state(cfg: &ClientTopConfig) -> Result> { remote_str, auth_header: format!("Bearer {}", cfg.auth.token), traffic_config: cfg.traffic_shaping.clone(), + resolved_traffic: Arc::new(ResolvedShaperConfig::resolve(&cfg.traffic_shaping)), bypass, server_public_key, proxy_auth, initial_master: Mutex::new(None), handshake_lock: Mutex::new(()), max_download_bytes: cfg.traffic_shaping.max_download_bytes, + max_connections: cfg.client.max_connections.unwrap_or(MAX_LOCAL_CONNECTIONS), + max_in_flight_bytes, + upload_concurrency, })) } diff --git a/src/client/proxy.rs b/src/client/proxy.rs index 8df8383..0094295 100644 --- a/src/client/proxy.rs +++ b/src/client/proxy.rs @@ -7,6 +7,7 @@ use url::Url; pub async fn parse_proxy_request( reader: &mut (impl AsyncReadExt + Unpin), buffer: &mut BytesMut, + need_proxy_auth: bool, ) -> Result<(String, usize, String, Option)> { const MAX_HEADER_LEN: usize = 16 * 1024; @@ -17,7 +18,11 @@ pub async fn parse_proxy_request( let mut headers = [httparse::EMPTY_HEADER; 64]; let mut req = httparse::Request::new(&mut headers); if let httparse::Status::Complete(amt) = req.parse(buffer)? { - let proxy_auth = extract_header(req.headers, "proxy-authorization"); + let proxy_auth = if need_proxy_auth { + extract_header(req.headers, "proxy-authorization") + } else { + None + }; return Ok(( req.method.context("no method")?.to_owned(), amt, @@ -47,7 +52,12 @@ pub fn resolve_target_host(method: &str, url_str: &str) -> Result { let port = auth .port_u16() .ok_or_else(|| anyhow!("port required: {url_str}"))?; - return Ok(format!("{}:{port}", auth.host())); + let host = auth.host(); + let host = host + .strip_prefix('[') + .and_then(|h| h.strip_suffix(']')) + .unwrap_or(host); + return Ok(format!("{host}:{port}")); } let url = Url::parse(url_str).context("invalid proxy URL")?; diff --git a/src/client/state.rs b/src/client/state.rs index d158e9f..c6e58c6 100644 --- a/src/client/state.rs +++ b/src/client/state.rs @@ -9,7 +9,7 @@ use zeroize::Zeroizing; use crate::bypass::BypassRules; use crate::client::constants::MASTER_RESUME_WINDOW_SECS; use crate::client::handshake::{self, PqSessionTicket}; -use crate::shaper::TrafficConfig; +use crate::shaper::{ResolvedShaperConfig, TrafficConfig}; pub type InitialMasterEntry = (String, Zeroizing<[u8; 32]>, u64); @@ -41,12 +41,16 @@ pub struct SharedState { pub remote_str: String, pub auth_header: String, pub traffic_config: TrafficConfig, + pub resolved_traffic: Arc, pub bypass: Option>, pub server_public_key: Option, pub proxy_auth: Option<(String, String)>, pub initial_master: Mutex>, pub handshake_lock: Mutex<()>, pub max_download_bytes: Option, + pub max_connections: usize, + pub max_in_flight_bytes: usize, + pub upload_concurrency: usize, } pub struct Resuming { diff --git a/src/client/utils.rs b/src/client/utils.rs index 2974682..a1b044a 100644 --- a/src/client/utils.rs +++ b/src/client/utils.rs @@ -1,21 +1,23 @@ use anyhow::{Context, Result, anyhow}; use bytes::Bytes; use rand::RngExt; +use std::fmt::Write as _; use std::future::Future; use tokio::task::JoinHandle; use tracing::warn; +use uuid::Uuid; use crate::client::constants::{MIN_PADDING, PADDING_POOL}; use crate::shaper::{self, FrameCipher}; #[inline] -fn build_cookie_into(buf: &mut String, name: &str, value: &str) { +fn build_cookie_into(buf: &mut String, name: &str, value: impl std::fmt::Display) { buf.clear(); - let cap = name.len() + 1 + value.len() + 2 + MIN_PADDING + PADDING_POOL.len(); + let cap = name.len() + 1 + 36 + 2 + MIN_PADDING + PADDING_POOL.len(); buf.reserve(cap); buf.push_str(name); buf.push('='); - buf.push_str(value); + let _ = write!(buf, "{value}"); buf.push_str("; "); let padding_len = rand::rng().random_range(MIN_PADDING..PADDING_POOL.len()); buf.push_str(std::str::from_utf8(&PADDING_POOL[..padding_len]).expect("Invalid UTF-8")) @@ -27,8 +29,8 @@ pub fn build_tunnel_cookie(buf: &mut String, session_val: &str) { } #[inline] -pub fn build_stream_cookie(buf: &mut String, stream_id: &str) { - build_cookie_into(buf, "stream", stream_id) +pub fn build_stream_cookie(buf: &mut String, stream_id: Uuid) { + build_cookie_into(buf, "stream", stream_id.as_hyphenated()) } pub fn encode_initial_payload( @@ -45,7 +47,11 @@ pub fn encode_initial_payload( Bytes::new() }; - let raw_payload_limit = shaper::MAX_RAW_PAYLOAD; + let raw_payload_limit = match (cipher.is_some(), config.encoding_type) { + (true, shaper::EncodingType::Json) => shaper::JSON_PAYLOAD_CAP_CIPHER, + (false, shaper::EncodingType::Json) => shaper::JSON_PAYLOAD_CAP_PLAIN, + _ => shaper::MAX_RAW_PAYLOAD, + }; let mut body = Vec::new(); let mut offset = 0; let mut seq: u64 = 0; diff --git a/src/config/mod.rs b/src/config/mod.rs index 5e798c6..a650f59 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -45,6 +45,8 @@ pub struct ServerSection { pub listen: String, pub path: String, pub private_key: Option, + #[serde(default)] + pub max_tunnels: Option, } #[derive(Deserialize, Debug)] @@ -57,6 +59,12 @@ pub struct ClientSection { pub public_key: Option, #[serde(default)] pub auth: Option, + #[serde(default)] + pub max_connections: Option, + #[serde(default)] + pub max_in_flight_bytes: Option, + #[serde(default)] + pub upload_concurrency: Option, } #[derive(Deserialize, Debug, Clone)] diff --git a/src/crypto/cipher.rs b/src/crypto/cipher.rs index 6ffbfa2..4f70349 100644 --- a/src/crypto/cipher.rs +++ b/src/crypto/cipher.rs @@ -9,8 +9,8 @@ use zeroize::Zeroizing; use crate::shaper::FrameCipher; -const NONCE_LEN: usize = 12; -const TAG_LEN: usize = 16; +pub const NONCE_LEN: usize = 12; +pub const TAG_LEN: usize = 16; const EMPTY_AAD: &[u8] = b""; #[inline] @@ -143,6 +143,27 @@ impl FrameCipher for AesFrameCipher { .map_err(|e| io::Error::other(anyhow!("decryption error: {e}")))?; Ok(()) } + + #[inline] + fn seal_in_place( + &self, + out: &mut bytes::BytesMut, + nonce_start: usize, + ct_start: usize, + ) -> io::Result<()> { + let nonce_bytes = random_nonce(); + out[nonce_start..ct_start].copy_from_slice(&nonce_bytes); + let tag = self + .cipher + .encrypt_in_place_detached( + Nonce::from_slice(&nonce_bytes), + EMPTY_AAD, + &mut out[ct_start..], + ) + .map_err(|e| io::Error::other(anyhow!("encryption error: {e}")))?; + out.extend_from_slice(tag.as_ref()); + Ok(()) + } } #[cfg(test)] diff --git a/src/crypto/mod.rs b/src/crypto/mod.rs index 475ef08..410a0be 100644 --- a/src/crypto/mod.rs +++ b/src/crypto/mod.rs @@ -2,7 +2,7 @@ mod cipher; mod handshake; mod keys; -pub use cipher::{AesFrameCipher, decrypt_bytes, encrypt_bytes}; +pub use cipher::{AesFrameCipher, NONCE_LEN, TAG_LEN, decrypt_bytes, encrypt_bytes}; pub use handshake::{ derive_connection_keys, derive_cookie_stream_key, derive_handshake_key, derive_initial_master, }; diff --git a/src/dns/client.rs b/src/dns/client.rs index e8fa34c..2659b01 100644 --- a/src/dns/client.rs +++ b/src/dns/client.rs @@ -216,6 +216,9 @@ impl DnsClient { msg.header().id() )); } + if msg.header().tc() { + return Err(anyhow!("DNS response truncated (TC set)")); + } let rcode = msg.header().rcode(); if rcode == Rcode::NXDOMAIN { return Ok((vec![], Duration::from_secs(self.config.options.empty_ttl))); diff --git a/src/dns/transport.rs b/src/dns/transport.rs index 632d04b..048359e 100644 --- a/src/dns/transport.rs +++ b/src/dns/transport.rs @@ -151,13 +151,16 @@ impl DotTransport { } } - let w = writer.as_mut().unwrap(); - let len_prefix = (data.len() as u16).to_be_bytes(); - if w.write_all(&len_prefix).await.is_err() - || w.write_all(&data).await.is_err() - || w.flush().await.is_err() - { - warn!("DoT write failed, dropping connection"); + let write_result = timeout(Duration::from_secs(10), async { + let w = writer.as_mut().unwrap(); + let len_prefix = (data.len() as u16).to_be_bytes(); + w.write_all(&len_prefix).await?; + w.write_all(&data).await?; + w.flush().await + }) + .await; + if !matches!(write_result, Ok(Ok(()))) { + warn!("DoT write failed or timed out, dropping connection"); for (_, tx) in actor_pending.lock().await.drain() { let _ = tx.send(Err(anyhow!("write failed, connection reset"))); @@ -191,13 +194,18 @@ impl DotTransport { async fn reader_loop(mut r: tokio::io::ReadHalf>, pending: PendingMap) { let mut len_buf = [0u8; 2]; - while r.read_exact(&mut len_buf).await.is_ok() { + loop { + let len_res = timeout(Duration::from_secs(30), r.read_exact(&mut len_buf)).await; + if !matches!(len_res, Ok(Ok(_))) { + break; + } let msg_len = u16::from_be_bytes(len_buf) as usize; if msg_len == 0 { continue; } let mut buf = vec![0u8; msg_len]; - if r.read_exact(&mut buf).await.is_err() { + let read_res = timeout(Duration::from_secs(30), r.read_exact(&mut buf)).await; + if !matches!(read_res, Ok(Ok(_))) { break; } if buf.len() >= 2 { @@ -216,7 +224,7 @@ impl DotTransport { ) -> Result> { let stream = timeout(Duration::from_secs(5), TcpStream::connect(upstream)).await??; stream.set_nodelay(true)?; - Ok(connector.connect(name, stream).await?) + Ok(timeout(Duration::from_secs(5), connector.connect(name, stream)).await??) } pub(super) async fn send(&self, data: &mut [u8]) -> Result<(Vec, u16)> { diff --git a/src/error/mod.rs b/src/error/mod.rs index 2b594d1..ebbaa2d 100644 --- a/src/error/mod.rs +++ b/src/error/mod.rs @@ -55,7 +55,7 @@ impl From for HttpProxyError { } } -#[derive(Debug)] +#[derive(Debug, Clone)] pub struct ServerError(pub StatusCode, pub String); impl ServerError { @@ -91,6 +91,10 @@ impl ServerError { pub fn precondition_required(msg: impl Into) -> Self { Self(StatusCode::PRECONDITION_REQUIRED, msg.into()) } + #[inline] + pub fn service_unavailable(msg: impl Into) -> Self { + Self(StatusCode::SERVICE_UNAVAILABLE, msg.into()) + } } impl IntoResponse for ServerError { @@ -148,6 +152,10 @@ mod tests { ServerError::precondition_required("x").0, StatusCode::PRECONDITION_REQUIRED ); + assert_eq!( + ServerError::service_unavailable("x").0, + StatusCode::SERVICE_UNAVAILABLE + ); } #[test] diff --git a/src/log/mod.rs b/src/log/mod.rs index b44e96a..27ea41c 100644 --- a/src/log/mod.rs +++ b/src/log/mod.rs @@ -43,8 +43,10 @@ macro_rules! json_fmt_layer { } pub fn init_tracing(log_cfg: &LogConfig) -> Option { - let filter = - EnvFilter::try_from_default_env().unwrap_or_else(|_| EnvFilter::new(&log_cfg.level)); + let filter = match EnvFilter::try_from_default_env() { + Ok(f) => f, + Err(_) => EnvFilter::try_new(&log_cfg.level).unwrap_or_else(|_| EnvFilter::new("info")), + }; match &log_cfg.file_path { Some(path_str) => { let (non_blocking, guard) = build_file_writer(path_str, log_cfg.max_backups); @@ -77,7 +79,7 @@ fn build_file_writer( let file_stem = file_path .file_stem() .and_then(|s| s.to_str()) - .expect("invalid log file path: missing file stem"); + .unwrap_or("httproxy"); let file_extension = file_path .extension() .and_then(|s| s.to_str()) diff --git a/src/server/actor/tunnel.rs b/src/server/actor/tunnel.rs index 9c894a2..4541d9d 100644 --- a/src/server/actor/tunnel.rs +++ b/src/server/actor/tunnel.rs @@ -7,16 +7,16 @@ use tokio::net::tcp::{OwnedReadHalf, OwnedWriteHalf}; use tokio::sync::{Notify, mpsc, oneshot}; use tokio::time::Instant; use tracing::{Instrument, info, warn}; +use uuid::Uuid; -use crate::error::ServerError; use crate::now_secs; use crate::server::actor::upload::{UploadActor, UploadCmd}; use crate::server::constants::{ DOWNLOAD_CHANNEL_CAPACITY, ROTATION_STALENESS, STREAM_IDLE_TIMEOUT_SECS, - UPLOAD_CMD_CHANNEL_CAPACITY, UPLOAD_DONE_TIMEOUT, + UPLOAD_CMD_CHANNEL_CAPACITY, }; use crate::server::stream_registry::StreamRegistry; -use crate::shaper::{FrameCipher, TrafficConfig, TrafficShaper}; +use crate::shaper::{FrameCipher, ResolvedShaperConfig, TrafficShaper}; pub enum TunnelCmd { UploadFrame { @@ -25,7 +25,7 @@ pub enum TunnelCmd { }, UploadEos { max_seq: u64, - ack: oneshot::Sender>, + ack: oneshot::Sender>, }, Continue { reply: oneshot::Sender>>>, @@ -54,7 +54,7 @@ pub struct TunnelActor { pending_continue: Vec>>>>, segment_done_tx: mpsc::Sender>, segment_done_rx: mpsc::Receiver>, - stream_id: String, + stream_id: Uuid, stream_registry: Arc, shutdown_signal: Arc, max_download_bytes: Option, @@ -67,10 +67,11 @@ impl TunnelActor { #[allow(clippy::too_many_arguments)] pub fn new( rx: mpsc::Receiver, - download_tx: mpsc::Sender>, - stream_id: String, + download_tx: Option>>, + stream_id: Uuid, stream_registry: Arc, max_download_bytes: Option, + last_activity: Arc, ) -> Self { let (seg_tx, seg_rx) = mpsc::channel::>(2); Self { @@ -80,7 +81,7 @@ impl TunnelActor { upload_handle: None, download_handle: None, shaper: None, - download_tx: Some(download_tx), + download_tx, pending_continue: Vec::new(), segment_done_tx: seg_tx, segment_done_rx: seg_rx, @@ -90,7 +91,7 @@ impl TunnelActor { max_download_bytes, pending_write_half: None, pending_initial_seq: 0, - last_activity: Arc::new(AtomicU64::new(now_secs())), + last_activity, } } @@ -111,7 +112,7 @@ impl TunnelActor { &mut self, read_half: OwnedReadHalf, write_half: Option, - config: TrafficConfig, + config: &ResolvedShaperConfig, download_cipher: Option>, initial_seq: u64, ) { @@ -140,23 +141,24 @@ impl TunnelActor { self.download_handle = Some(tokio::spawn( async move { let mut bytes_sent: u64 = 0; - while let Some(result) = shaper.as_mut().next().await { - match result { - Ok((_seq, data)) => { + loop { + match shaper.as_mut().next().await { + Some(Ok((_seq, data))) => { bytes_sent += data.len() as u64; + activity.store(now_secs(), Ordering::Relaxed); if download_tx.send(Ok(data)).await.is_err() { break; } - activity.store(now_secs(), Ordering::Relaxed); if max_bytes.is_some_and(|m| bytes_sent >= m) { let _ = done_tx.send(Some(shaper)).await; return; } } - Err(e) => { + Some(Err(e)) => { let _ = download_tx.send(Err(e)).await; break; } + None => break, } } let _ = done_tx.send(None).await; @@ -169,9 +171,10 @@ impl TunnelActor { let max_bytes = self.max_download_bytes; if self.upload_tx.is_none() { - let write_half = self.pending_write_half.take().expect( - "on_upstream_connected must provide write_half when upload channel not pre-set", - ); + let write_half = self + .pending_write_half + .take() + .expect("upstream write half must be provided before run"); let (upload_tx, upload_rx) = mpsc::channel::(UPLOAD_CMD_CHANNEL_CAPACITY); self.upload_tx = Some(upload_tx); let upload_actor = UploadActor::new(upload_rx, write_half, self.pending_initial_seq); @@ -180,7 +183,10 @@ impl TunnelActor { )); } - self.spawn_download_segment(max_bytes); + let direct_mode = self.download_tx.is_none(); + if !direct_mode { + self.spawn_download_segment(max_bytes); + } let rotation_timeout = tokio::time::sleep(ROTATION_STALENESS); tokio::pin!(rotation_timeout); @@ -192,7 +198,7 @@ impl TunnelActor { loop { tokio::select! { biased; - returned = self.segment_done_rx.recv() => { + returned = self.segment_done_rx.recv(), if !direct_mode => { self.download_tx = None; match returned { @@ -215,10 +221,12 @@ impl TunnelActor { if self.shaper.is_some() { let (new_tx, new_rx) = mpsc::channel::>(DOWNLOAD_CHANNEL_CAPACITY); + if reply.send(Some(new_rx)).is_err() { + continue; + } self.download_tx = Some(new_tx); self.spawn_download_segment(self.max_download_bytes); self.phase = Phase::Active; - let _ = reply.send(Some(new_rx)); } else { let _ = reply.send(None); } @@ -280,49 +288,30 @@ impl TunnelActor { } } TunnelCmd::UploadEos { max_seq, ack } => { - let (done_tx, done_rx) = oneshot::channel(); if upload_tx - .send(UploadCmd::Eos { - max_seq, - done: done_tx, - }) + .send(UploadCmd::Eos { max_seq, ack }) .await .is_err() { warn!("upload actor closed before EOS"); - let _ = ack.send(Err(ServerError::bad_gateway("upload actor closed"))); - return; } - let upload_done_timeout = UPLOAD_DONE_TIMEOUT; - tokio::spawn( - async move { - let confirmed = tokio::time::timeout(upload_done_timeout, done_rx) - .await - .map(|r| r.is_ok()) - .unwrap_or(false); - if confirmed { - let _ = ack.send(Ok(())); - } else { - warn!("upload EOS ack timed out or upload actor closed"); - let _ = - ack.send(Err(ServerError::gateway_timeout("upload drain timeout"))); - } - } - .instrument(tracing::Span::current()), - ); } - TunnelCmd::Continue { reply } => { - if self.shaper.is_some() { + TunnelCmd::Continue { reply } => match self.phase { + Phase::Rotating => { let (new_tx, new_rx) = mpsc::channel::>(DOWNLOAD_CHANNEL_CAPACITY); self.download_tx = Some(new_tx); self.spawn_download_segment(self.max_download_bytes); self.phase = Phase::Active; let _ = reply.send(Some(new_rx)); - } else { + } + Phase::Active if self.shaper.is_none() => { self.pending_continue.push(reply); } - } + _ => { + let _ = reply.send(None); + } + }, TunnelCmd::Shutdown => {} } } @@ -346,7 +335,7 @@ impl TunnelActor { } fn consume_stream(&mut self) { - self.stream_registry.mark_consumed(&self.stream_id); + self.stream_registry.mark_consumed(self.stream_id); } } diff --git a/src/server/actor/upload.rs b/src/server/actor/upload.rs index ff960da..d6d6937 100644 --- a/src/server/actor/upload.rs +++ b/src/server/actor/upload.rs @@ -6,6 +6,7 @@ use tokio::sync::{mpsc, oneshot}; use tokio::time::{Duration, Instant}; use tracing::warn; +use crate::error::ServerError; use crate::server::constants::{ MAX_EOS_WAITERS, MAX_PENDING_BYTES, MAX_PENDING_FRAMES, MAX_REORDER_SECS, WRITE_TIMEOUT, }; @@ -17,7 +18,7 @@ pub enum UploadCmd { }, Eos { max_seq: u64, - done: oneshot::Sender<()>, + ack: oneshot::Sender>, }, Shutdown, } @@ -27,10 +28,10 @@ enum UploadPhase { next_seq: u64, pending: BTreeMap, pending_bytes: usize, - eos_waiters: Vec<(u64, oneshot::Sender<()>)>, + eos_waiters: Vec<(u64, oneshot::Sender>)>, }, Draining { - eos_waiters: Vec>, + eos_waiters: Vec>>, }, Closed, } @@ -91,7 +92,7 @@ impl UploadActor { true } UploadCmd::Frame { seq, data } => self.handle_frame(seq, data).await, - UploadCmd::Eos { max_seq, done } => self.handle_eos(max_seq, done), + UploadCmd::Eos { max_seq, ack } => self.handle_eos(max_seq, ack), } } @@ -109,9 +110,10 @@ impl UploadActor { } if seq == *next_seq { if let Some(ref mut upstream) = self.upstream - && tokio::time::timeout(WRITE_TIMEOUT, upstream.write_all(&data)) - .await - .is_err() + && !matches!( + tokio::time::timeout(WRITE_TIMEOUT, upstream.write_all(&data)).await, + Ok(Ok(())) + ) { self.shutdown_and_drain().await; return true; @@ -120,12 +122,14 @@ impl UploadActor { while let Some(pending_data) = pending.remove(next_seq) { *pending_bytes -= pending_data.len(); if let Some(ref mut upstream) = self.upstream - && tokio::time::timeout( - WRITE_TIMEOUT, - upstream.write_all(&pending_data), + && !matches!( + tokio::time::timeout( + WRITE_TIMEOUT, + upstream.write_all(&pending_data), + ) + .await, + Ok(Ok(())) ) - .await - .is_err() { self.shutdown_and_drain().await; return true; @@ -135,8 +139,8 @@ impl UploadActor { let mut i = 0; while i < eos_waiters.len() { if *next_seq > eos_waiters[i].0 { - let (_, done) = eos_waiters.swap_remove(i); - let _ = done.send(()); + let (_, ack) = eos_waiters.swap_remove(i); + let _ = ack.send(Ok(())); } else { i += 1; } @@ -155,9 +159,10 @@ impl UploadActor { pending_bytes, max_pending_frames = MAX_PENDING_FRAMES, max_pending_bytes = MAX_PENDING_BYTES, - "reorder buffer overflow, discarding frame" + "reorder buffer overflow, aborting upload" ); - return false; + self.shutdown_and_drain().await; + return true; } pending.insert(seq, data); *pending_bytes += len; @@ -167,7 +172,7 @@ impl UploadActor { } } - fn handle_eos(&mut self, max_seq: u64, done: oneshot::Sender<()>) -> bool { + fn handle_eos(&mut self, max_seq: u64, ack: oneshot::Sender>) -> bool { match &mut self.phase { UploadPhase::Reordering { next_seq, @@ -175,25 +180,26 @@ impl UploadActor { .. } => { if *next_seq > max_seq { - let _ = done.send(()); + let _ = ack.send(Ok(())); } else if eos_waiters.len() >= MAX_EOS_WAITERS { warn!( max_seq, eos_waiters = eos_waiters.len(), "EOS waiters overflow, shutting down upload actor" ); + let _ = ack.send(Err(ServerError::bad_gateway("upload EOS waiters overflow"))); return true; } else { - eos_waiters.push((max_seq, done)); + eos_waiters.push((max_seq, ack)); } false } UploadPhase::Draining { eos_waiters } => { - eos_waiters.push(done); + eos_waiters.push(ack); false } UploadPhase::Closed => { - let _ = done.send(()); + let _ = ack.send(Err(ServerError::bad_gateway("upload actor closed"))); true } } @@ -206,7 +212,7 @@ impl UploadActor { self.upstream = None; self.phase = match std::mem::replace(&mut self.phase, UploadPhase::Closed) { UploadPhase::Reordering { eos_waiters, .. } => UploadPhase::Draining { - eos_waiters: eos_waiters.into_iter().map(|(_, done)| done).collect(), + eos_waiters: eos_waiters.into_iter().map(|(_, ack)| ack).collect(), }, other @ UploadPhase::Draining { .. } => other, UploadPhase::Closed => UploadPhase::Closed, @@ -215,15 +221,16 @@ impl UploadActor { fn ack_all_waiters(&mut self) { let phase = std::mem::replace(&mut self.phase, UploadPhase::Closed); + let err = ServerError::gateway_timeout("upload drain timeout"); match phase { UploadPhase::Reordering { eos_waiters, .. } => { - for (_, waiter) in eos_waiters { - let _ = waiter.send(()); + for (_, ack) in eos_waiters { + let _ = ack.send(Err(err.clone())); } } UploadPhase::Draining { eos_waiters } => { - for waiter in eos_waiters { - let _ = waiter.send(()); + for ack in eos_waiters { + let _ = ack.send(Err(err.clone())); } } UploadPhase::Closed => {} @@ -267,14 +274,14 @@ mod tests { }) .await .unwrap(); - let (done_tx, done_rx) = oneshot::channel(); + let (ack_tx, ack_rx) = oneshot::channel(); tx.send(UploadCmd::Eos { max_seq: 1, - done: done_tx, + ack: ack_tx, }) .await .unwrap(); - done_rx.await.unwrap(); + assert!(ack_rx.await.unwrap().is_ok()); drop(tx); handle.await.unwrap(); @@ -305,14 +312,14 @@ mod tests { }) .await .unwrap(); - let (done_tx, done_rx) = oneshot::channel(); + let (ack_tx, ack_rx) = oneshot::channel(); tx.send(UploadCmd::Eos { max_seq: 2, - done: done_tx, + ack: ack_tx, }) .await .unwrap(); - done_rx.await.unwrap(); + assert!(ack_rx.await.unwrap().is_ok()); drop(tx); handle.await.unwrap(); @@ -325,14 +332,14 @@ mod tests { let actor = UploadActor::new(rx, server_write, 5); let handle = tokio::spawn(async move { actor.run().await }); - let (done_tx, done_rx) = oneshot::channel(); + let (ack_tx, ack_rx) = oneshot::channel(); tx.send(UploadCmd::Eos { max_seq: 3, - done: done_tx, + ack: ack_tx, }) .await .unwrap(); - done_rx.await.unwrap(); + assert!(ack_rx.await.unwrap().is_ok()); drop(tx); handle.await.unwrap(); @@ -351,17 +358,37 @@ mod tests { }) .await .unwrap(); - let (done_tx, done_rx) = oneshot::channel(); + let (ack_tx, ack_rx) = oneshot::channel(); tx.send(UploadCmd::Eos { max_seq: 5, - done: done_tx, + ack: ack_tx, }) .await .unwrap(); tx.send(UploadCmd::Shutdown).await.unwrap(); drop(tx); - done_rx.await.unwrap(); + assert!(ack_rx.await.unwrap().is_err()); + handle.await.unwrap(); + } + + #[tokio::test] + async fn channel_close_acks_waiters_with_error() { + let (_rx, server_write) = tcp_pair().await; + let (tx, rx) = mpsc::channel::(16); + let actor = UploadActor::new(rx, server_write, 0); + let handle = tokio::spawn(async move { actor.run().await }); + + let (ack_tx, ack_rx) = oneshot::channel(); + tx.send(UploadCmd::Eos { + max_seq: 5, + ack: ack_tx, + }) + .await + .unwrap(); + + drop(tx); + assert!(ack_rx.await.unwrap().is_err()); handle.await.unwrap(); } } diff --git a/src/server/constants.rs b/src/server/constants.rs index aeeb7c0..543254f 100644 --- a/src/server/constants.rs +++ b/src/server/constants.rs @@ -19,10 +19,12 @@ pub const MASTER_CLEANUP_INTERVAL: Duration = Duration::from_secs(60); pub const CONNECT_TIMEOUT: Duration = Duration::from_secs(10); pub const WRITE_TIMEOUT: Duration = Duration::from_secs(30); -pub const DOWNLOAD_CHANNEL_CAPACITY: usize = 1; +pub const DOWNLOAD_CHANNEL_CAPACITY: usize = 2; pub const TUNNEL_CMD_CHANNEL_CAPACITY: usize = 32; pub const UPLOAD_CMD_CHANNEL_CAPACITY: usize = 8; +pub const MAX_TUNNELS: usize = 1024; + pub const MASTER_EXPIRY: Duration = Duration::from_secs(1200); pub const PADDING_POOL: &[u8] = b"padding=XXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXX"; diff --git a/src/server/handlers.rs b/src/server/handlers.rs index bd8f38d..8ec604f 100644 --- a/src/server/handlers.rs +++ b/src/server/handlers.rs @@ -6,7 +6,7 @@ use jsonwebtoken::{DecodingKey, Validation}; use std::io; use std::sync::Arc; use tracing::{Instrument, info, warn}; -use uuid; +use uuid::Uuid; use zeroize::Zeroizing; use crate::crypto::{self, AesFrameCipher}; @@ -28,18 +28,18 @@ use crate::server::stream_registry::StreamRegistry; struct StreamGuard { registry: Arc, - stream_id: String, + stream_id: Uuid, armed: bool, } impl Drop for StreamGuard { fn drop(&mut self) { if self.armed { - self.registry.mark_consumed(&self.stream_id); + self.registry.mark_consumed(self.stream_id); } } } impl StreamGuard { - fn new(registry: Arc, stream_id: String) -> Self { + fn new(registry: Arc, stream_id: Uuid) -> Self { Self { registry, stream_id, @@ -52,8 +52,8 @@ impl StreamGuard { } struct ActorGuard { - actors: Arc>, - key: String, + actors: Arc>, + key: Uuid, armed: bool, } impl Drop for ActorGuard { @@ -64,7 +64,7 @@ impl Drop for ActorGuard { } } impl ActorGuard { - fn new(actors: Arc>, key: String) -> Self { + fn new(actors: Arc>, key: Uuid) -> Self { Self { actors, key, @@ -100,14 +100,29 @@ async fn setup_tunnel_response( body: Body, host: &str, port: u16, - stream_id: &str, + stream_id: Uuid, upload_cipher: Option>, download_cipher: Option>, ) -> Result { + let tunnel_permit = state + .tunnel_semaphore + .clone() + .try_acquire_owned() + .map_err(|_| ServerError::service_unavailable("too many concurrent tunnels"))?; + let encoding = state.traffic_config.encoding_type; - let (download_tx, download_rx) = - tokio::sync::mpsc::channel::>(DOWNLOAD_CHANNEL_CAPACITY); + let max_download_bytes = state.traffic_config.max_download_bytes; + let direct = max_download_bytes.is_none(); + let last_activity = Arc::new(std::sync::atomic::AtomicU64::new(crate::now_secs())); + + let (download_tx, download_rx) = if direct { + (None, None) + } else { + let (tx, rx) = + tokio::sync::mpsc::channel::>(DOWNLOAD_CHANNEL_CAPACITY); + (Some(tx), Some(rx)) + }; let (actor_tx, actor_rx) = tokio::sync::mpsc::channel::(TUNNEL_CMD_CHANNEL_CAPACITY); let upstream = tokio::time::timeout( @@ -155,31 +170,53 @@ async fn setup_tunnel_response( let mut actor = crate::server::actor::tunnel::TunnelActor::new( actor_rx, download_tx, - stream_id.to_owned(), + stream_id, Arc::clone(&state.stream_registry), - state.traffic_config.max_download_bytes, + max_download_bytes, + Arc::clone(&last_activity), ); actor.set_upload_channel(upload_tx_for_actor, upload_handle); - actor.on_upstream_connected( - upstream_read, - None, - (*state.traffic_config).clone(), - download_cipher, - 0, - ); + + let body_stream = if direct { + let shaper = crate::shaper::TrafficShaper::with_seq( + upstream_read, + &state.resolved_traffic, + download_cipher, + 0, + ); + let activity = Arc::clone(&last_activity); + Body::from_stream(shaper.map(move |r| { + if r.is_ok() { + activity.store(crate::now_secs(), std::sync::atomic::Ordering::Relaxed); + } + r.map(|(_seq, data)| data) + })) + } else { + actor.on_upstream_connected( + upstream_read, + None, + &state.resolved_traffic, + download_cipher, + 0, + ); + Body::from_stream(tokio_stream::wrappers::ReceiverStream::new( + download_rx.expect("download channel must exist when rotation is enabled"), + )) + }; let handle = SessionHandle { cmd_tx: actor_tx.clone(), upload_cipher, encoding, }; - state.actors.insert(stream_id.to_owned(), handle); - let mut early_guard = ActorGuard::new(Arc::clone(&state.actors), stream_id.to_owned()); + state.actors.insert(stream_id, handle); + let mut early_guard = ActorGuard::new(Arc::clone(&state.actors), stream_id); - let key = stream_id.to_owned(); + let key = stream_id; let actors_ref2 = Arc::clone(&state.actors); let actor_handle = tokio::spawn( async move { + let _permit = tunnel_permit; let _guard = ActorGuard::new(actors_ref2, key); actor.run().await; } @@ -194,9 +231,7 @@ async fn setup_tunnel_response( let response = Response::builder() .header("Cache-Control", "no-store") .header("Set-Cookie", padding) - .body(Body::from_stream( - tokio_stream::wrappers::ReceiverStream::new(download_rx), - )) + .body(body_stream) .map_err(|e| ServerError::internal(e.to_string()))?; tunnel_guard.disarm(); @@ -227,14 +262,14 @@ async fn spawn_encrypted_tunnel( let mut master = Zeroizing::new([0u8; 32]); let (username, master_z, created) = value_ref; master.copy_from_slice(&**master_z); - let username = username.clone(); + let username = Arc::clone(username); let created = *created; if crate::now_secs().saturating_sub(created) > MASTER_EXPIRY.as_secs() { drop(entry); state.master_store.remove(session_id); return Err(ServerError::precondition_required("master key expired")); } - span.record("user", &username); + span.record("user", username.as_ref()); drop(entry); let cookie_stream_key = crypto::derive_cookie_stream_key(&master); @@ -246,23 +281,17 @@ async fn spawn_encrypted_tunnel( let stream_id_bytes_arr: [u8; 16] = stream_id_bytes .try_into() .map_err(|_| ServerError::precondition_required("invalid stream ID length"))?; - let stream_uuid = uuid::Uuid::from_bytes(stream_id_bytes_arr); - let stream_id = stream_uuid.to_string(); + let stream_id = Uuid::from_bytes(stream_id_bytes_arr); - utils::validate_uuid(&stream_id)?; - - if !state - .stream_registry - .register(&stream_id, crate::now_secs()) - { + if !state.stream_registry.register(stream_id, crate::now_secs()) { warn!(stream_id = %stream_id, session_id = %session_id, "duplicate stream registration attempt"); return Err(ServerError::precondition_required("stream already active")); } - let mut stream_guard = StreamGuard::new(Arc::clone(&state.stream_registry), stream_id.clone()); + let mut stream_guard = StreamGuard::new(Arc::clone(&state.stream_registry), stream_id); let (upload_key, download_key, target_key) = - crypto::derive_connection_keys(&master, stream_uuid.as_bytes()); + crypto::derive_connection_keys(&master, stream_id.as_bytes()); let enc_target = base64::engine::general_purpose::URL_SAFE_NO_PAD .decode(enc_target_b64) .map_err(|_| ServerError::precondition_required("invalid cookie encoding"))?; @@ -286,7 +315,7 @@ async fn spawn_encrypted_tunnel( body, host, port, - &stream_id, + stream_id, Some(upload_cipher as Arc), Some(download_cipher), ) @@ -330,8 +359,8 @@ async fn dispatch_to_actor(handle: SessionHandle, body: Body) -> Result { + match tokio::time::timeout(crate::server::constants::UPLOAD_DONE_TIMEOUT, reply_rx).await { + Ok(Ok(Some(new_download_rx))) => { let padding = utils::random_padding(); return Response::builder() .header("Cache-Control", "no-store") @@ -341,7 +370,7 @@ async fn dispatch_to_actor(handle: SessionHandle, body: Body) -> Result { + Ok(Ok(None)) | Ok(Err(_)) => { let padding = utils::random_padding(); return Response::builder() .header("Cache-Control", "no-store") @@ -350,6 +379,9 @@ async fn dispatch_to_actor(handle: SessionHandle, body: Body) -> Result { + return Err(ServerError::gateway_timeout("continue timed out")); + } } } @@ -387,31 +419,29 @@ pub async fn dispatch( let span = tracing::Span::current(); let has_target = headers.get("X-Target").is_some(); - let stream_cookie = utils::extract_cookie_value(&headers, "stream").map(|s| s.to_owned()); - if let Some(ref stream_id) = stream_cookie { - utils::validate_uuid(stream_id)?; + if let Some(stream_cookie_val) = utils::extract_cookie_value(&headers, "stream") { + let stream_id = Uuid::parse_str(stream_cookie_val) + .map_err(|_| ServerError::precondition_required("invalid stream id"))?; - if let Some(handle) = state.actors.get(stream_id).map(|r| r.value().clone()) { + if let Some(handle) = state.actors.get(&stream_id).map(|r| r.value().clone()) { return dispatch_to_actor(handle, body).await; } if has_target { - return handle_plaintext_tunnel(state, headers, body, span, Some(stream_id.clone())) - .await; + return handle_plaintext_tunnel(state, headers, body, span, Some(stream_id)).await; } return handle_stream_not_found(&state, stream_id); } - let session_cookie = utils::extract_cookie_value(&headers, "session").map(|s| s.to_owned()); - if let Some(ref session_val) = session_cookie { + if let Some(session_val) = utils::extract_cookie_value(&headers, "session") { if session_val.contains(':') { return spawn_encrypted_tunnel(state, session_val, body, span).await; } if has_target { - utils::validate_uuid(session_val)?; - return handle_plaintext_tunnel(state, headers, body, span, Some(session_val.clone())) - .await; + let stream_id = Uuid::parse_str(session_val) + .map_err(|_| ServerError::precondition_required("invalid stream id"))?; + return handle_plaintext_tunnel(state, headers, body, span, Some(stream_id)).await; } return Err(ServerError::precondition_required( "invalid session cookie — missing target or encrypted payload", @@ -434,7 +464,7 @@ pub async fn dispatch( fn handle_stream_not_found( state: &Arc, - stream_id: &str, + stream_id: Uuid, ) -> Result { match state.stream_registry.check(stream_id) { StreamQueryResult::Consumed => { @@ -461,7 +491,7 @@ async fn handle_plaintext_tunnel( headers: HeaderMap, body: Body, span: tracing::Span, - stream_id_opt: Option, + stream_id_opt: Option, ) -> Result { let user = validate_jwt_if_needed(&headers, &state.decoding_key, &state.jwt_validation)?; span.record("user", &user); @@ -473,11 +503,8 @@ async fn handle_plaintext_tunnel( span.record("target", target); let stream_id = match stream_id_opt { - Some(id) => { - utils::validate_uuid(&id)?; - id - } - None => uuid::Uuid::new_v4().to_string(), + Some(id) => id, + None => Uuid::new_v4(), }; let (host, port_str) = target @@ -487,17 +514,14 @@ async fn handle_plaintext_tunnel( .parse() .map_err(|_| ServerError::bad_request("invalid port"))?; - if !state - .stream_registry - .register(&stream_id, crate::now_secs()) - { + if !state.stream_registry.register(stream_id, crate::now_secs()) { warn!(stream_id = %stream_id, "duplicate stream registration for plaintext tunnel"); return Err(ServerError::precondition_required("stream already active")); } - let mut stream_guard = StreamGuard::new(Arc::clone(&state.stream_registry), stream_id.clone()); + let mut stream_guard = StreamGuard::new(Arc::clone(&state.stream_registry), stream_id); - let response = setup_tunnel_response(&state, body, host, port, &stream_id, None, None).await?; + let response = setup_tunnel_response(&state, body, host, port, stream_id, None, None).await?; stream_guard.disarm(); @@ -591,10 +615,10 @@ async fn handle_fresh_handshake( crypto::derive_initial_master(&ss_mlkem, &ss_x25519) }; - let session_id = uuid::Uuid::new_v4().to_string(); + let session_id = Uuid::new_v4().to_string(); state.master_store.insert( session_id.clone(), - (user.clone(), master, crate::now_secs()), + (Arc::from(user.as_str()), master, crate::now_secs()), ); info!(session_id = %session_id, "handshake: master key derived"); diff --git a/src/server/janitor.rs b/src/server/janitor.rs index 2a1ddb3..46beee5 100644 --- a/src/server/janitor.rs +++ b/src/server/janitor.rs @@ -1,5 +1,6 @@ use dashmap::DashMap; use std::sync::Arc; +use uuid::Uuid; use crate::server::SessionHandle; use crate::server::constants::{ @@ -41,7 +42,7 @@ pub async fn master_and_stream_janitor( } } -pub async fn stream_janitor(actors: Arc>) { +pub async fn stream_janitor(actors: Arc>) { let mut interval = tokio::time::interval(JANITOR_INTERVAL); interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip); diff --git a/src/server/mod.rs b/src/server/mod.rs index d71e505..b6be906 100644 --- a/src/server/mod.rs +++ b/src/server/mod.rs @@ -10,7 +10,7 @@ pub mod utils; use crate::config::ServerTopConfig; use crate::crypto; use crate::dns::{self, DnsClient}; -use crate::shaper::{EncodingType, FrameCipher, TrafficConfig}; +use crate::shaper::{EncodingType, FrameCipher, ResolvedShaperConfig, TrafficConfig}; use anyhow::Context; use axum::serve::ListenerExt; @@ -26,15 +26,17 @@ use std::{ }, }; use stream_registry::StreamRegistry; -use tokio::sync::mpsc; +use tokio::sync::{Semaphore, mpsc}; use tower::ServiceBuilder; use tower_http::trace::TraceLayer; use tracing::info; +use uuid::Uuid; use zeroize::Zeroizing; use crate::server::actor::tunnel::TunnelCmd; +use crate::server::constants::MAX_TUNNELS; -pub type MasterStoreEntry = (String, Zeroizing<[u8; 32]>, u64); +pub type MasterStoreEntry = (Arc, Zeroizing<[u8; 32]>, u64); #[derive(Debug, Serialize, Deserialize)] pub struct Claims { @@ -57,10 +59,12 @@ pub struct AppState { pub dns_client: Option>, pub client_subnet: Option, pub traffic_config: Arc, + pub resolved_traffic: Arc, pub private_key: Option, pub master_store: Arc>, pub stream_registry: Arc, - pub actors: Arc>, + pub actors: Arc>, + pub tunnel_semaphore: Arc, pub stream_id_counter: Arc, } @@ -86,6 +90,8 @@ pub async fn build_state(config: &mut ServerTopConfig) -> anyhow::Result anyhow::Result { #[pin] inner: S, buf: BytesMut, + scratch: BytesMut, + json_scratch: Vec, cipher: Option>, encoding: EncodingType, max_buf_size: usize, @@ -33,6 +35,8 @@ where Self { inner, buf: BytesMut::with_capacity(max_buf_size), + scratch: BytesMut::new(), + json_scratch: Vec::new(), cipher, encoding, max_buf_size, @@ -53,9 +57,22 @@ where loop { let cipher_ref: Option<&dyn FrameCipher> = this.cipher.as_ref().map(|c| c.as_ref() as &dyn FrameCipher); - match shaper::decode_from_buffer(this.buf, cipher_ref, *this.encoding) { + match shaper::decode_frame( + this.buf, + this.scratch, + this.json_scratch, + cipher_ref, + *this.encoding, + ) { Ok(Some(frame)) => { - return Poll::Ready(Some(Ok(frame))); + let (seq, data) = match frame { + DecodedFrame::InScratch { seq, start, end } => { + let plain = this.scratch.split().freeze(); + (seq, plain.slice(start..end)) + } + DecodedFrame::Owned { seq, data } => (seq, data), + }; + return Poll::Ready(Some(Ok((seq, data)))); } Ok(None) => {} Err(e) => { @@ -209,4 +226,34 @@ mod tests { let err_msg = result.unwrap_err().to_string(); assert!(err_msg.contains("buffer exceeded"), "got: {err_msg}"); } + + #[tokio::test] + async fn encrypted_frames_decoded_with_scratch() { + use crate::crypto::AesFrameCipher; + use zeroize::Zeroizing; + + let mut key = Zeroizing::new([0u8; 32]); + rand::Rng::fill_bytes(&mut rand::rng(), &mut *key); + let cipher = Arc::new(AesFrameCipher::new(&key)); + + let frame = shaper::encode_frame( + b"encrypted hello", + 0, + Some(cipher.as_ref() as &dyn shaper::FrameCipher), + 16384, + [0, 0], + shaper::EncodingType::Binary, + ) + .unwrap(); + let byte_stream = stream::iter(vec![Ok(Bytes::from(frame))]); + let mut decoder = FrameDecoder::new( + byte_stream, + Some(cipher), + shaper::EncodingType::Binary, + 18_781, + ); + let (seq, data) = decoder.next().await.unwrap().unwrap(); + assert_eq!(seq, 0); + assert_eq!(&data[..], b"encrypted hello"); + } } diff --git a/src/server/stream_registry.rs b/src/server/stream_registry.rs index e9554c9..3dff610 100644 --- a/src/server/stream_registry.rs +++ b/src/server/stream_registry.rs @@ -1,5 +1,6 @@ use dashmap::DashMap; use std::sync::atomic::{AtomicU8, Ordering}; +use uuid::Uuid; #[repr(u8)] #[derive(Debug, Clone, Copy, PartialEq, Eq)] @@ -19,7 +20,7 @@ pub enum StreamQueryResult { pub struct StreamConsumedError; pub struct StreamRegistry { - streams: DashMap, + streams: DashMap, } impl StreamRegistry { @@ -29,8 +30,8 @@ impl StreamRegistry { } } - pub fn register(&self, stream_id: &str, now_secs: u64) -> bool { - match self.streams.entry(stream_id.to_owned()) { + pub fn register(&self, stream_id: Uuid, now_secs: u64) -> bool { + match self.streams.entry(stream_id) { dashmap::Entry::Occupied(_) => false, dashmap::Entry::Vacant(entry) => { entry.insert((AtomicU8::new(StreamState::Active as u8), now_secs)); @@ -39,16 +40,16 @@ impl StreamRegistry { } } - pub fn mark_consumed(&self, stream_id: &str) { - if let Some(entry) = self.streams.get(stream_id) { + pub fn mark_consumed(&self, stream_id: Uuid) { + if let Some(entry) = self.streams.get(&stream_id) { entry .0 .store(StreamState::Consumed as u8, Ordering::Release); } } - pub fn check(&self, stream_id: &str) -> StreamQueryResult { - match self.streams.get(stream_id) { + pub fn check(&self, stream_id: Uuid) -> StreamQueryResult { + match self.streams.get(&stream_id) { None => StreamQueryResult::Fresh, Some(entry) => match entry.0.load(Ordering::Acquire) { s if s == StreamState::Active as u8 => StreamQueryResult::Active, @@ -90,64 +91,75 @@ mod tests { 1_700_000_000 } + fn ids() -> (Uuid, Uuid, Uuid) { + (Uuid::new_v4(), Uuid::new_v4(), Uuid::new_v4()) + } + #[test] fn fresh_register_succeeds() { let reg = StreamRegistry::new(); - assert!(reg.register("s1", dummy_ts())); - assert_eq!(reg.check("s1"), StreamQueryResult::Active); + let (s1, _, _) = ids(); + assert!(reg.register(s1, dummy_ts())); + assert_eq!(reg.check(s1), StreamQueryResult::Active); } #[test] fn duplicate_register_rejected() { let reg = StreamRegistry::new(); - assert!(reg.register("s1", dummy_ts())); - assert!(!reg.register("s1", dummy_ts())); + let (s1, _, _) = ids(); + assert!(reg.register(s1, dummy_ts())); + assert!(!reg.register(s1, dummy_ts())); } #[test] fn consumed_stream_reports_consumed() { let reg = StreamRegistry::new(); - reg.register("s1", dummy_ts()); - reg.mark_consumed("s1"); - assert_eq!(reg.check("s1"), StreamQueryResult::Consumed); + let (s1, _, _) = ids(); + reg.register(s1, dummy_ts()); + reg.mark_consumed(s1); + assert_eq!(reg.check(s1), StreamQueryResult::Consumed); } #[test] fn fresh_stream_returns_fresh() { let reg = StreamRegistry::new(); - assert_eq!(reg.check("nonexistent"), StreamQueryResult::Fresh); + let (_, s2, _) = ids(); + assert_eq!(reg.check(s2), StreamQueryResult::Fresh); } #[test] fn remove_consumed_before_removes_expired() { let reg = StreamRegistry::new(); - reg.register("s1", 100); - reg.mark_consumed("s1"); - reg.register("s2", 200); + let (s1, s2, _) = ids(); + reg.register(s1, 100); + reg.mark_consumed(s1); + reg.register(s2, 200); let removed = reg.remove_consumed_before(150); assert!(removed >= 1); - assert_eq!(reg.check("s1"), StreamQueryResult::Fresh); - assert_eq!(reg.check("s2"), StreamQueryResult::Active); + assert_eq!(reg.check(s1), StreamQueryResult::Fresh); + assert_eq!(reg.check(s2), StreamQueryResult::Active); } #[test] fn remove_consumed_before_keeps_recent() { let reg = StreamRegistry::new(); - reg.register("s1", 1000); - reg.mark_consumed("s1"); + let (s1, _, _) = ids(); + reg.register(s1, 1000); + reg.mark_consumed(s1); let removed = reg.remove_consumed_before(900); assert_eq!(removed, 0); - assert_eq!(reg.check("s1"), StreamQueryResult::Consumed); + assert_eq!(reg.check(s1), StreamQueryResult::Consumed); } #[test] fn concurrent_streams_independent() { let reg = StreamRegistry::new(); - assert!(reg.register("s1", dummy_ts())); - assert!(reg.register("s2", dummy_ts())); - reg.mark_consumed("s1"); - assert_eq!(reg.check("s1"), StreamQueryResult::Consumed); - assert_eq!(reg.check("s2"), StreamQueryResult::Active); + let (s1, s2, _) = ids(); + assert!(reg.register(s1, dummy_ts())); + assert!(reg.register(s2, dummy_ts())); + reg.mark_consumed(s1); + assert_eq!(reg.check(s1), StreamQueryResult::Consumed); + assert_eq!(reg.check(s2), StreamQueryResult::Active); } #[tokio::test] @@ -155,11 +167,11 @@ mod tests { let reg = Arc::new(StreamRegistry::new()); let mut handles = Vec::new(); - for i in 0..16u8 { + for _ in 0..16u8 { let reg = Arc::clone(®); handles.push(tokio::spawn(async move { - let id = format!("stream-{i}"); - reg.register(&id, dummy_ts()) + let id = Uuid::new_v4(); + reg.register(id, dummy_ts()) })); } @@ -178,7 +190,7 @@ mod tests { #[tokio::test] async fn concurrent_register_and_consume_race() { let reg = Arc::new(StreamRegistry::new()); - let id = "race-stream"; + let id = Uuid::new_v4(); assert!(reg.register(id, dummy_ts())); diff --git a/src/shaper/mod.rs b/src/shaper/mod.rs index 8314299..fbc950f 100644 --- a/src/shaper/mod.rs +++ b/src/shaper/mod.rs @@ -16,8 +16,15 @@ use tokio::{ time::{Instant, Sleep}, }; +use crate::crypto::{NONCE_LEN, TAG_LEN}; + pub const MAX_RAW_PAYLOAD: usize = 16 * 1024; +const READ_HIGH_WATER: usize = 32 * 1024; + +pub const JSON_PAYLOAD_CAP_PLAIN: usize = 14320 - HEADER_LEN; +pub const JSON_PAYLOAD_CAP_CIPHER: usize = 14320 - HEADER_LEN - NONCE_LEN - TAG_LEN; + const TABLE_SIZE: usize = 8192; const TABLE_MASK: usize = TABLE_SIZE - 1; const DELIMITER: u8 = b'\n'; @@ -113,12 +120,44 @@ impl TrafficConfig { } #[derive(Debug, Clone, Copy)] -struct ResolvedStage { +pub(crate) struct ResolvedStage { end_count: usize, padding_threshold: usize, padding_range: [usize; 2], } +#[derive(Debug, Clone)] +pub struct ResolvedShaperConfig { + pub(crate) stages: Arc<[ResolvedStage]>, + pub(crate) global_threshold: usize, + pub(crate) global_range: [usize; 2], + pub encoding: EncodingType, +} + +impl ResolvedShaperConfig { + pub fn resolve(config: &TrafficConfig) -> Self { + let mut stages: Vec = config + .stages + .iter() + .map(|s| ResolvedStage { + end_count: s + .count + .or_else(|| s.count_range.map(|[_, hi]| hi)) + .unwrap_or(0), + padding_threshold: s.padding_threshold, + padding_range: s.padding_range, + }) + .collect(); + stages.sort_unstable_by_key(|s| s.end_count); + Self { + stages: Arc::from(stages), + global_threshold: config.global.padding_threshold, + global_range: config.global.padding_range, + encoding: config.encoding_type, + } + } +} + pub trait FrameCipher: Send + Sync { fn encrypt(&self, data: &[u8]) -> Result, Error>; fn decrypt(&self, data: &[u8]) -> Result, Error>; @@ -134,6 +173,18 @@ pub trait FrameCipher: Send + Sync { out.extend_from_slice(&decrypted); Ok(()) } + + fn seal_in_place( + &self, + out: &mut BytesMut, + nonce_start: usize, + ct_start: usize, + ) -> Result<(), Error> { + debug_assert!(ct_start >= nonce_start); + let plain = out.split_off(ct_start); + out.truncate(nonce_start); + self.encrypt_into(&plain, out) + } } #[inline] @@ -146,23 +197,6 @@ fn read_u16_be(data: &[u8]) -> u16 { u16::from_be_bytes(data[..2].try_into().unwrap()) } -#[inline] -fn extract_frame(payload: &[u8]) -> Result<(u64, Bytes), Error> { - if payload.len() < HEADER_LEN { - return Err(Error::new(ErrorKind::InvalidData, "payload too short")); - } - let seq = read_u64_be(&payload[..8]); - let orig_len = read_u16_be(&payload[8..10]) as usize; - let total = HEADER_LEN + orig_len; - if payload.len() < total { - return Err(Error::new( - ErrorKind::InvalidData, - "payload shorter than declared original length", - )); - } - Ok((seq, Bytes::copy_from_slice(&payload[HEADER_LEN..total]))) -} - #[inline] fn extract_frame_range(payload: &[u8]) -> Result<(u64, usize, usize), Error> { if payload.len() < HEADER_LEN { @@ -200,7 +234,7 @@ fn trim_bytes(mut b: &[u8]) -> &[u8] { } #[inline] -fn parse_json_payload(json: &[u8]) -> Result, Error> { +fn parse_json_payload_into(json: &[u8], out: &mut Vec) -> Result { let json = trim_bytes(json); let err = |msg: &str| Error::new(ErrorKind::InvalidData, msg); @@ -221,7 +255,7 @@ fn parse_json_payload(json: &[u8]) -> Result, Error> { let enc_str = std::str::from_utf8(enc_str_bytes).map_err(|_| err("payload is not valid UTF-8"))?; - base122_fast::decode(enc_str).map_err(err) + base122_fast::decode_into(enc_str, out).map_err(err) } #[inline] @@ -249,9 +283,20 @@ pub fn encode_frame( encoding: EncodingType, ) -> std::io::Result> { let raw_len = data.len(); + let frame_raw_limit = match (cipher.is_some(), encoding) { + (true, EncodingType::Json) => JSON_PAYLOAD_CAP_CIPHER, + (false, EncodingType::Json) => JSON_PAYLOAD_CAP_PLAIN, + _ => MAX_RAW_PAYLOAD, + }; + if raw_len > frame_raw_limit { + return Err(std::io::Error::new( + std::io::ErrorKind::InvalidInput, + format!("frame payload too large: {raw_len} > {frame_raw_limit}"), + )); + } let padding_len = if raw_len < padding_threshold { - let max_pad = MAX_RAW_PAYLOAD - raw_len; + let max_pad = frame_raw_limit - raw_len; let wanted = rand::rng().random_range(padding_range[0]..=padding_range[1]); wanted.min(max_pad) } else { @@ -276,83 +321,97 @@ pub fn encode_frame( Ok(frame.to_vec()) } -macro_rules! decode_skeleton { - ( - $src:expr, $cipher:expr, $encoding:expr, - cipher => |$cipher_var:ident| $cipher_expr:expr, - binary_plain => |$bin_plain_var:ident| $bin_plain_expr:expr, - json_plain => |$json_plain_var:ident| $json_plain_expr:expr $(,)? - ) => { - match $encoding { - EncodingType::Binary => { - if $src.len() < 2 { - return Ok(None); - } - let frame_len = read_u16_be(&$src[..2]) as usize; - - if frame_len > MAX_BINARY_FRAME_LEN { - return Err(Error::new( - ErrorKind::InvalidData, - "binary frame length exceeds limit", - )); - } - if $src.len() < 2 + frame_len { - return Ok(None); - } - $src.advance(2); - let $bin_plain_var = $src.split_to(frame_len); +#[derive(Debug)] +pub enum DecodedFrame { + InScratch { seq: u64, start: usize, end: usize }, + Owned { seq: u64, data: Bytes }, +} - if let Some(c) = $cipher { - let mut $cipher_var = BytesMut::new(); - c.decrypt_into(&$bin_plain_var, &mut $cipher_var)?; - $cipher_expr - } else { - $bin_plain_expr - } +pub fn decode_frame( + src: &mut BytesMut, + scratch: &mut BytesMut, + json_scratch: &mut Vec, + cipher: Option<&dyn FrameCipher>, + encoding: EncodingType, +) -> Result, Error> { + match encoding { + EncodingType::Binary => { + if src.len() < 2 { + return Ok(None); } + let frame_len = read_u16_be(&src[..2]) as usize; - EncodingType::Json => { - let newline_pos = memchr::memchr(DELIMITER, $src); - - match newline_pos { - Some(pos) => { - if pos > MAX_JSON_LINE_LEN { - return Err(Error::new( - ErrorKind::InvalidData, - "JSON line exceeds maximum allowed length", - )); - } + if frame_len > MAX_BINARY_FRAME_LEN { + return Err(Error::new( + ErrorKind::InvalidData, + "binary frame length exceeds limit", + )); + } + if src.len() < 2 + frame_len { + return Ok(None); + } + src.advance(2); + let view = src.split_to(frame_len); + + if let Some(c) = cipher { + scratch.clear(); + c.decrypt_into(&view, scratch)?; + let (seq, start, end) = extract_frame_range(scratch)?; + Ok(Some(DecodedFrame::InScratch { seq, start, end })) + } else { + let (seq, start, end) = extract_frame_range(&view)?; + let data = view.freeze().slice(start..end); + Ok(Some(DecodedFrame::Owned { seq, data })) + } + } - let line = $src.split_to(pos); - $src.advance(1); + EncodingType::Json => { + let newline_pos = memchr::memchr(DELIMITER, src); + + match newline_pos { + Some(pos) => { + if pos > MAX_JSON_LINE_LEN { + return Err(Error::new( + ErrorKind::InvalidData, + "JSON line exceeds maximum allowed length", + )); + } - if line.is_empty() { - return Err(Error::new(ErrorKind::InvalidData, "empty frame line")); - } + let line = src.split_to(pos); + src.advance(1); - let $json_plain_var = parse_json_payload(&line)?; + if line.is_empty() { + return Err(Error::new(ErrorKind::InvalidData, "empty frame line")); + } - if let Some(c) = $cipher { - let mut $cipher_var = BytesMut::new(); - c.decrypt_into(&$json_plain_var, &mut $cipher_var)?; - $cipher_expr - } else { - $json_plain_expr - } + parse_json_payload_into(&line, json_scratch)?; + + if let Some(c) = cipher { + scratch.clear(); + c.decrypt_into(json_scratch, scratch)?; + let (seq, start, end) = extract_frame_range(scratch)?; + Ok(Some(DecodedFrame::InScratch { seq, start, end })) + } else { + let data = Bytes::from(std::mem::take(json_scratch)); + let (seq, start, end) = extract_frame_range(&data)?; + Ok(Some(DecodedFrame::Owned { + seq, + data: data.slice(start..end), + })) } - None => { - if $src.len() > MAX_JSON_LINE_LEN { - return Err(Error::new( - ErrorKind::InvalidData, - "incomplete JSON line is too long", - )); - } - Ok(None) + } + None => { + if src.len() > MAX_JSON_LINE_LEN { + return Err(Error::new( + ErrorKind::InvalidData, + "incomplete JSON line is too long", + )); } + Ok(None) } } } - }; + } } pub fn decode_from_buffer( @@ -360,36 +419,24 @@ pub fn decode_from_buffer( cipher: Option<&dyn FrameCipher>, encoding: EncodingType, ) -> Result, Error> { - decode_skeleton!( - src, cipher, encoding, - cipher => |decrypted| Ok(Some(extract_frame(&decrypted)?)), - binary_plain => |frame_data| Ok(Some(extract_frame(&frame_data)?)), - json_plain => |encoded_payload| Ok(Some(extract_frame(&encoded_payload)?)), - ) + let mut scratch = BytesMut::new(); + let mut json_scratch = Vec::new(); + match decode_frame(src, &mut scratch, &mut json_scratch, cipher, encoding)? { + Some(DecodedFrame::InScratch { seq, start, end }) => { + let plain = scratch.split().freeze(); + Ok(Some((seq, plain.slice(start..end)))) + } + Some(DecodedFrame::Owned { seq, data }) => Ok(Some((seq, data))), + None => Ok(None), + } } -pub fn decode_frame_owned( - src: &mut BytesMut, - cipher: Option<&dyn FrameCipher>, - encoding: EncodingType, -) -> Result, Error> { - decode_skeleton!( - src, cipher, encoding, - cipher => |decrypted| { - let (seq, start, end) = extract_frame_range(&decrypted)?; - Ok(Some((seq, decrypted, start, end))) - }, - binary_plain => |frame_data| { - let (seq, start, end) = extract_frame_range(&frame_data)?; - Ok(Some((seq, frame_data, start, end))) - }, - json_plain => |encoded_payload| { - let mut frame_data = BytesMut::with_capacity(encoded_payload.len()); - frame_data.extend_from_slice(&encoded_payload); - let (seq, start, end) = extract_frame_range(&frame_data)?; - Ok(Some((seq, frame_data, start, end))) - }, - ) +pub trait SealInto { + fn poll_seal_into( + self: Pin<&mut Self>, + cx: &mut Context<'_>, + out: &mut BytesMut, + ) -> Poll>>; } pin_project! { @@ -401,12 +448,13 @@ pin_project! { raw_buf: BytesMut, out_buf: BytesMut, enc_buf: BytesMut, + json_buf: Vec, #[pin] flush_timer: Sleep, timer_armed: bool, cursor: usize, - stages: Vec, + stages: Arc<[ResolvedStage]>, global_threshold: usize, global_range: [usize; 2], packet_count: usize, @@ -414,6 +462,7 @@ pin_project! { rng: SmallRng, cipher: Option>, encoding: EncodingType, + seal_threshold: usize, seq: u64, } } @@ -421,63 +470,52 @@ pin_project! { impl TrafficShaper { pub fn with_seq( reader: R, - config: TrafficConfig, + config: &ResolvedShaperConfig, cipher: Option>, start_seq: u64, ) -> Self { let mut base_rng = rand::rng(); let cursor = (base_rng.next_u64() as usize) & TABLE_MASK; - let mut stages: Vec = config - .stages - .iter() - .map(|s| ResolvedStage { - end_count: s - .count - .or_else(|| s.count_range.map(|[_, hi]| hi)) - .unwrap_or(0), - padding_threshold: s.padding_threshold, - padding_range: s.padding_range, - }) - .collect(); - stages.sort_unstable_by_key(|s| s.end_count); - - let out_capacity = match config.encoding_type { + let out_capacity = match config.encoding { EncodingType::Binary => MAX_BINARY_FRAME_LEN + 2, EncodingType::Json => MAX_JSON_LINE_LEN + 1, }; + let seal_threshold = match (cipher.is_some(), config.encoding) { + (true, EncodingType::Binary) => { + MAX_RAW_PAYLOAD - (NONCE_LEN + TAG_LEN + HEADER_LEN + 2) + } + (false, EncodingType::Binary) => MAX_RAW_PAYLOAD - (HEADER_LEN + 2), + (true, EncodingType::Json) => JSON_PAYLOAD_CAP_CIPHER, + (false, EncodingType::Json) => JSON_PAYLOAD_CAP_PLAIN, + }; + Self { reader, - raw_buf: BytesMut::with_capacity(MAX_RAW_PAYLOAD), + raw_buf: BytesMut::with_capacity(READ_HIGH_WATER), out_buf: BytesMut::with_capacity(out_capacity), enc_buf: BytesMut::new(), + json_buf: Vec::new(), flush_timer: tokio::time::sleep_until(Instant::now()), timer_armed: false, - stages, - global_threshold: config.global.padding_threshold, - global_range: config.global.padding_range, + stages: Arc::clone(&config.stages), + global_threshold: config.global_threshold, + global_range: config.global_range, packet_count: 0, cursor, stage_idx: 0, rng: SmallRng::from_rng(&mut base_rng), cipher, - encoding: config.encoding_type, + encoding: config.encoding, + seal_threshold, seq: start_seq, } } #[inline] - fn seal_and_emit(this: &mut Proj<'_, R>) -> Result<(u64, Bytes), Error> { - let raw_len = this.raw_buf.len(); - debug_assert!(raw_len > 0); - debug_assert!(raw_len <= MAX_RAW_PAYLOAD); - - *this.timer_armed = false; - + fn resolve_padding(this: &mut Proj<'_, R>, raw_len: usize) -> (usize, usize) { *this.packet_count += 1; - let seq = *this.seq; - *this.seq = seq + 1; let stages = &this.stages; let pc = *this.packet_count; @@ -494,7 +532,7 @@ impl TrafficShaper { }; let padding_len = if raw_len < threshold { - let max_pad = MAX_RAW_PAYLOAD - raw_len; + let max_pad = *this.seal_threshold - raw_len; let wanted = this.rng.random_range(range[0]..=range[1]); wanted.min(max_pad) } else { @@ -502,66 +540,91 @@ impl TrafficShaper { }; let payload_len = HEADER_LEN + raw_len + padding_len; + (payload_len, padding_len) + } - if let Some(cipher) = this.cipher { - this.out_buf.clear(); - this.out_buf.reserve(payload_len); - this.out_buf.put_u64(seq); - this.out_buf.put_u16(raw_len as u16); - this.out_buf.put_slice(&this.raw_buf[..raw_len]); - if padding_len > 0 { - this.out_buf.put_bytes(0u8, padding_len); - } + #[inline] + fn seal_into(this: &mut Proj<'_, R>, out: &mut BytesMut) -> Result { + let raw_len = this.raw_buf.len().min(*this.seal_threshold); + debug_assert!(raw_len > 0); + debug_assert!(raw_len <= MAX_RAW_PAYLOAD); - this.enc_buf.clear(); - cipher.encrypt_into(&this.out_buf[..payload_len], this.enc_buf)?; - this.out_buf.clear(); - write_encoded_frame(this.out_buf, this.enc_buf, *this.encoding); - } else { - this.out_buf.clear(); - - match *this.encoding { - EncodingType::Binary => { - this.out_buf.reserve(2 + payload_len); - this.out_buf.put_u16(payload_len as u16); - this.out_buf.put_u64(seq); - this.out_buf.put_u16(raw_len as u16); - this.out_buf.put_slice(&this.raw_buf[..raw_len]); - if padding_len > 0 { - this.out_buf.put_bytes(0u8, padding_len); - } + *this.timer_armed = false; + + let seq = *this.seq; + *this.seq = seq + 1; + + let (payload_len, padding_len) = Self::resolve_padding(this, raw_len); + + match (this.cipher.as_deref(), *this.encoding) { + (Some(cipher), EncodingType::Binary) => { + let enc_len = NONCE_LEN + payload_len + TAG_LEN; + out.reserve(2 + enc_len); + out.put_u16(enc_len as u16); + let nonce_start = out.len(); + out.put_bytes(0u8, NONCE_LEN); + let ct_start = out.len(); + out.put_u64(seq); + out.put_u16(raw_len as u16); + out.put_slice(&this.raw_buf[..raw_len]); + if padding_len > 0 { + out.put_bytes(0u8, padding_len); } - EncodingType::Json => { - this.out_buf.put_u64(seq); - this.out_buf.put_u16(raw_len as u16); - this.out_buf.put_slice(&this.raw_buf[..raw_len]); - if padding_len > 0 { - this.out_buf.put_bytes(0u8, padding_len); - } - let payload = this.out_buf.split(); - write_encoded_frame(this.out_buf, &payload[..payload_len], EncodingType::Json); + cipher.seal_in_place(out, nonce_start, ct_start)?; + } + (None, EncodingType::Binary) => { + out.reserve(2 + payload_len); + out.put_u16(payload_len as u16); + out.put_u64(seq); + out.put_u16(raw_len as u16); + out.put_slice(&this.raw_buf[..raw_len]); + if padding_len > 0 { + out.put_bytes(0u8, padding_len); + } + } + (_, EncodingType::Json) => { + this.out_buf.clear(); + this.out_buf.reserve(payload_len); + this.out_buf.put_u64(seq); + this.out_buf.put_u16(raw_len as u16); + this.out_buf.put_slice(&this.raw_buf[..raw_len]); + if padding_len > 0 { + this.out_buf.put_bytes(0u8, padding_len); + } + let payload = this.out_buf.split(); + if let Some(cipher) = this.cipher.as_deref() { + this.enc_buf.clear(); + cipher.encrypt_into(&payload[..payload_len], this.enc_buf)?; + base122_fast::encode_into(this.enc_buf, this.json_buf); + } else { + base122_fast::encode_into(&payload[..payload_len], this.json_buf); } + out.reserve(9 + this.json_buf.len() + 3); + out.put_slice(b"{\"data\":\""); + out.put_slice(this.json_buf); + out.put_slice(b"\"}\n"); } } - this.raw_buf.clear(); - let result = this.out_buf.split().freeze(); - Ok((seq, result)) + this.raw_buf.advance(raw_len); + Ok(seq) } } -impl tokio_stream::Stream for TrafficShaper { - type Item = Result<(u64, Bytes), Error>; - - fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { +impl TrafficShaper { + fn poll_fill_and_seal( + self: Pin<&mut Self>, + cx: &mut Context<'_>, + out: &mut BytesMut, + ) -> Poll>> { let mut this = self.project(); loop { - if this.raw_buf.len() >= MAX_RAW_PAYLOAD { - return Poll::Ready(Some(Self::seal_and_emit(&mut this))); + if this.raw_buf.len() >= *this.seal_threshold { + return Poll::Ready(Self::seal_into(&mut this, out).map(Some)); } - let remaining = MAX_RAW_PAYLOAD - this.raw_buf.len(); + let remaining = READ_HIGH_WATER - this.raw_buf.len(); this.raw_buf.reserve(remaining); let spare = this.raw_buf.spare_capacity_mut(); let read_limit = spare.len().min(remaining); @@ -572,9 +635,9 @@ impl tokio_stream::Stream for TrafficShaper { let n = rb.filled().len(); if n == 0 { return if this.raw_buf.is_empty() { - Poll::Ready(None) + Poll::Ready(Ok(None)) } else { - Poll::Ready(Some(Self::seal_and_emit(&mut this))) + Poll::Ready(Self::seal_into(&mut this, out).map(Some)) }; } @@ -587,6 +650,9 @@ impl tokio_stream::Stream for TrafficShaper { if raw_len == 0 { return Poll::Pending; } + if raw_len >= *this.seal_threshold { + return Poll::Ready(Self::seal_into(&mut this, out).map(Some)); + } if !*this.timer_armed { let idx = *this.cursor; @@ -599,19 +665,44 @@ impl tokio_stream::Stream for TrafficShaper { } if this.flush_timer.as_mut().poll(cx).is_ready() { - return Poll::Ready(Some(Self::seal_and_emit(&mut this))); + return Poll::Ready(Self::seal_into(&mut this, out).map(Some)); } return Poll::Pending; } - Poll::Ready(Err(e)) => return Poll::Ready(Some(Err(e))), + Poll::Ready(Err(e)) => return Poll::Ready(Err(e)), } } } } +impl SealInto for TrafficShaper { + fn poll_seal_into( + self: Pin<&mut Self>, + cx: &mut Context<'_>, + out: &mut BytesMut, + ) -> Poll>> { + self.poll_fill_and_seal(cx, out) + } +} + +impl tokio_stream::Stream for TrafficShaper { + type Item = Result<(u64, Bytes), Error>; + + fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + let mut frame = BytesMut::new(); + match self.poll_fill_and_seal(cx, &mut frame) { + Poll::Ready(Ok(Some(seq))) => Poll::Ready(Some(Ok((seq, frame.freeze())))), + Poll::Ready(Ok(None)) => Poll::Ready(None), + Poll::Ready(Err(e)) => Poll::Ready(Some(Err(e))), + Poll::Pending => Poll::Pending, + } + } +} + #[cfg(test)] mod tests { use super::*; + use std::io::Cursor; fn test_config() -> TrafficConfig { TrafficConfig { @@ -652,7 +743,16 @@ mod tests { let mut buf = BytesMut::new(); buf.put_u16(100u16); buf.put_u8(0xAA); - let result = decode_from_buffer(&mut buf, None, EncodingType::Binary).unwrap(); + let mut scratch = BytesMut::new(); + let mut json_scratch = Vec::new(); + let result = decode_frame( + &mut buf, + &mut scratch, + &mut json_scratch, + None, + EncodingType::Binary, + ) + .unwrap(); assert!(result.is_none()); } @@ -661,7 +761,15 @@ mod tests { let mut buf = BytesMut::new(); buf.put_u16((MAX_RAW_PAYLOAD + 1000) as u16); buf.resize(2 + MAX_RAW_PAYLOAD + 1000, 0u8); - let result = decode_from_buffer(&mut buf, None, EncodingType::Binary); + let mut scratch = BytesMut::new(); + let mut json_scratch = Vec::new(); + let result = decode_frame( + &mut buf, + &mut scratch, + &mut json_scratch, + None, + EncodingType::Binary, + ); assert!(result.is_err()); } @@ -673,27 +781,29 @@ mod tests { .chain(3u16.to_be_bytes()) .chain(b"abc".iter().copied()) .collect::>(); - let (seq, data) = extract_frame(&payload).unwrap(); + let (seq, start, end) = extract_frame_range(&payload).unwrap(); assert_eq!(seq, 0); - assert_eq!(&data[..], b"abc"); + assert_eq!(&payload[start..end], b"abc"); } #[test] fn extract_frame_too_short() { - assert!(extract_frame(b"short").is_err()); + assert!(extract_frame_range(b"short").is_err()); } #[test] fn parse_json_payload_valid() { let enc = base122_fast::encode(b"hello"); let json = format!("{{\"data\":\"{enc}\"}}"); - let result = parse_json_payload(json.as_bytes()).unwrap(); - assert_eq!(result, b"hello"); + let mut out = Vec::new(); + let n = parse_json_payload_into(json.as_bytes(), &mut out).unwrap(); + assert_eq!(&out[..n], b"hello"); } #[test] fn parse_json_payload_missing_field() { - let result = parse_json_payload(b"{\"other\":\"x\"}"); + let mut out = Vec::new(); + let result = parse_json_payload_into(b"{\"other\":\"x\"}", &mut out); assert!(result.is_err()); } @@ -706,4 +816,419 @@ mod tests { let mean = table1.iter().map(|&v| v as f64).sum::() / table1.len() as f64; assert!((mean - AVG_LATENCY_MICROS).abs() < AVG_LATENCY_MICROS * 0.5); } + + #[test] + fn decode_frame_plain_owned_zero_copy() { + let data = b"plain payload"; + let frame = encode_frame(data, 9, None, 16384, [0, 0], EncodingType::Binary).unwrap(); + let mut src = BytesMut::from(&frame[..]); + let mut scratch = BytesMut::new(); + let mut json_scratch = Vec::new(); + match decode_frame( + &mut src, + &mut scratch, + &mut json_scratch, + None, + EncodingType::Binary, + ) + .unwrap() + .unwrap() + { + DecodedFrame::Owned { seq, data } => { + assert_eq!(seq, 9); + assert_eq!(&data[..], b"plain payload"); + } + _ => panic!("expected Owned frame"), + } + assert!(src.is_empty()); + } + + #[test] + fn decode_frame_cipher_into_scratch() { + use crate::crypto::AesFrameCipher; + use zeroize::Zeroizing; + + let mut key = Zeroizing::new([0u8; 32]); + rand::rng().fill_bytes(&mut *key); + let cipher = AesFrameCipher::new(&key); + + let frame = encode_frame( + b"secret payload", + 42, + Some(&cipher), + 16384, + [0, 0], + EncodingType::Binary, + ) + .unwrap(); + let mut src = BytesMut::from(&frame[..]); + let mut scratch = BytesMut::new(); + let mut json_scratch = Vec::new(); + match decode_frame( + &mut src, + &mut scratch, + &mut json_scratch, + Some(&cipher), + EncodingType::Binary, + ) + .unwrap() + .unwrap() + { + DecodedFrame::InScratch { seq, start, end } => { + assert_eq!(seq, 42); + assert_eq!(&scratch[start..end], b"secret payload"); + } + _ => panic!("expected InScratch frame"), + } + assert!(src.is_empty()); + } + + #[test] + fn decode_frame_scratch_reused_across_frames() { + use crate::crypto::AesFrameCipher; + use zeroize::Zeroizing; + + let mut key = Zeroizing::new([0u8; 32]); + rand::rng().fill_bytes(&mut *key); + let cipher = AesFrameCipher::new(&key); + + let mut combined = BytesMut::new(); + for (i, msg) in [ + b"first".as_slice(), + b"second".as_slice(), + b"third".as_slice(), + ] + .iter() + .enumerate() + { + let frame = encode_frame( + msg, + i as u64, + Some(&cipher), + 16384, + [0, 0], + EncodingType::Binary, + ) + .unwrap(); + combined.extend_from_slice(&frame); + } + + let mut scratch = BytesMut::new(); + let mut json_scratch = Vec::new(); + let mut expected_seq = 0u64; + while !combined.is_empty() { + match decode_frame( + &mut combined, + &mut scratch, + &mut json_scratch, + Some(&cipher), + EncodingType::Binary, + ) + .unwrap() + .unwrap() + { + DecodedFrame::InScratch { seq, start, end } => { + assert_eq!(seq, expected_seq); + assert_eq!( + &scratch[start..end], + [ + b"first".as_slice(), + b"second".as_slice(), + b"third".as_slice() + ][expected_seq as usize] + ); + expected_seq += 1; + } + _ => panic!("expected InScratch frame"), + } + } + assert_eq!(expected_seq, 3); + } + + #[test] + fn decode_frame_json_roundtrip() { + let frame = + encode_frame(b"json payload", 3, None, 16384, [0, 0], EncodingType::Json).unwrap(); + let mut src = BytesMut::from(&frame[..]); + let mut scratch = BytesMut::new(); + let mut json_scratch = Vec::new(); + match decode_frame( + &mut src, + &mut scratch, + &mut json_scratch, + None, + EncodingType::Json, + ) + .unwrap() + .unwrap() + { + DecodedFrame::Owned { seq, data } => { + assert_eq!(seq, 3); + assert_eq!(&data[..], b"json payload"); + } + _ => panic!("expected Owned frame"), + } + assert!(src.is_empty()); + } + + #[test] + fn decode_frame_json_cipher_roundtrip() { + use crate::crypto::AesFrameCipher; + use zeroize::Zeroizing; + + let mut key = Zeroizing::new([0u8; 32]); + rand::rng().fill_bytes(&mut *key); + let cipher = AesFrameCipher::new(&key); + + let frame = encode_frame( + b"json cipher payload", + 5, + Some(&cipher), + 16384, + [0, 0], + EncodingType::Json, + ) + .unwrap(); + let mut src = BytesMut::from(&frame[..]); + let mut scratch = BytesMut::new(); + let mut json_scratch = Vec::new(); + match decode_frame( + &mut src, + &mut scratch, + &mut json_scratch, + Some(&cipher), + EncodingType::Json, + ) + .unwrap() + .unwrap() + { + DecodedFrame::InScratch { seq, start, end } => { + assert_eq!(seq, 5); + assert_eq!(&scratch[start..end], b"json cipher payload"); + } + _ => panic!("expected InScratch frame"), + } + assert!(src.is_empty()); + } + + #[tokio::test] + async fn json_frames_fit_single_h2_data_frame() { + use crate::crypto::AesFrameCipher; + use zeroize::Zeroizing; + + let no_cipher: Option> = None; + let aes_cipher: Option> = + Some(Arc::new(AesFrameCipher::new(&Zeroizing::new([0u8; 32])))); + for cipher in [no_cipher, aes_cipher] { + let cipher_for_decode = cipher.clone(); + let mut config = test_config(); + config.encoding_type = EncodingType::Json; + let resolved = ResolvedShaperConfig::resolve(&config); + let data = vec![0xAAu8; 64 * 1024]; + let shaper = TrafficShaper::with_seq(Cursor::new(data.clone()), &resolved, cipher, 0); + let mut out = BytesMut::new(); + seal_all(shaper, &mut out); + + let mut frames = 0; + let mut src = &out[..]; + while !src.is_empty() { + let newline = memchr::memchr(b'\n', src).expect("frame must end with newline"); + let line_len = newline + 1; + assert!(line_len <= 16384, "JSON frame {line_len} B > 16384"); + src = &src[line_len..]; + frames += 1; + } + assert!(frames >= 4, "expected multiple frames, got {frames}"); + + let mut scratch = BytesMut::new(); + let mut json_scratch = Vec::new(); + let mut decoded = Vec::new(); + let mut buf = out; + while !buf.is_empty() { + match decode_frame( + &mut buf, + &mut scratch, + &mut json_scratch, + cipher_for_decode.as_deref(), + EncodingType::Json, + ) + .unwrap() + .unwrap() + { + DecodedFrame::InScratch { start, end, .. } => { + decoded.extend_from_slice(&scratch[start..end]); + } + DecodedFrame::Owned { data, .. } => decoded.extend_from_slice(&data), + } + } + assert_eq!(decoded, data); + } + } + + fn seal_all(shaper: TrafficShaper>>, out: &mut BytesMut) -> Vec { + let mut shaper = std::pin::pin!(shaper); + let mut seqs = Vec::new(); + loop { + let waker = std::task::Waker::noop(); + let mut cx = std::task::Context::from_waker(waker); + match shaper.as_mut().poll_seal_into(&mut cx, out) { + Poll::Ready(Ok(Some(seq))) => seqs.push(seq), + Poll::Ready(Ok(None)) => break, + Poll::Ready(Err(e)) => panic!("seal error: {e}"), + Poll::Pending => panic!("unexpected Pending with Cursor reader"), + } + } + seqs + } + + #[tokio::test] + async fn poll_seal_into_produces_valid_frames() { + let config = ResolvedShaperConfig::resolve(&test_config()); + let data = vec![0xABu8; 40_000]; + let shaper = TrafficShaper::with_seq(Cursor::new(data.clone()), &config, None, 0); + let mut out = BytesMut::new(); + let seqs = seal_all(shaper, &mut out); + + assert_eq!(seqs, vec![0, 1, 2]); + + let mut scratch = BytesMut::new(); + let mut json_scratch = Vec::new(); + let mut decoded = Vec::new(); + let mut frame_idx = 0usize; + while !out.is_empty() { + match decode_frame( + &mut out, + &mut scratch, + &mut json_scratch, + None, + EncodingType::Binary, + ) + .unwrap() + .unwrap() + { + DecodedFrame::Owned { seq, data } => { + assert_eq!(seq, seqs[frame_idx]); + frame_idx += 1; + decoded.extend_from_slice(&data); + } + _ => panic!("expected Owned frame"), + } + } + assert_eq!(decoded, data); + } + + #[tokio::test] + async fn poll_seal_into_cipher_matches_stream_output() { + use crate::crypto::AesFrameCipher; + use futures::StreamExt; + use zeroize::Zeroizing; + + let mut key = Zeroizing::new([0u8; 32]); + rand::rng().fill_bytes(&mut *key); + let cipher: Arc = Arc::new(AesFrameCipher::new(&key)); + + let config = ResolvedShaperConfig::resolve(&test_config()); + let data = vec![0x5Cu8; 33_000]; + + let cipher_for_decode = Arc::clone(&cipher); + let decode_append = |src: &mut BytesMut, out: &mut Vec| { + let mut scratch = BytesMut::new(); + let mut json_scratch = Vec::new(); + match decode_frame( + src, + &mut scratch, + &mut json_scratch, + Some(cipher_for_decode.as_ref()), + EncodingType::Binary, + ) + .unwrap() + .unwrap() + { + DecodedFrame::InScratch { start, end, .. } => { + out.extend_from_slice(&scratch[start..end]); + } + _ => panic!("expected InScratch frame"), + } + }; + + let shaper_stream = TrafficShaper::with_seq( + Cursor::new(data.clone()), + &config, + Some(Arc::clone(&cipher)), + 0, + ); + let mut stream_payload = Vec::new(); + let mut stream_seqs = Vec::new(); + let mut shaper_stream = Box::pin(shaper_stream); + while let Some(item) = shaper_stream.next().await { + let (seq, bytes) = item.unwrap(); + stream_seqs.push(seq); + let mut frame_buf = BytesMut::from(&bytes[..]); + decode_append(&mut frame_buf, &mut stream_payload); + } + + let shaper_seal = + TrafficShaper::with_seq(Cursor::new(data.clone()), &config, Some(cipher), 0); + let mut out = BytesMut::new(); + let seqs = seal_all(shaper_seal, &mut out); + let mut seal_payload = Vec::new(); + while !out.is_empty() { + decode_append(&mut out, &mut seal_payload); + } + + assert_eq!(seqs, stream_seqs); + assert_eq!(seal_payload, stream_payload); + assert_eq!(seal_payload, data); + } + + #[tokio::test] + async fn seal_in_place_default_impl_produces_valid_frames() { + use zeroize::Zeroizing; + + struct VecCipher(Zeroizing<[u8; 32]>); + impl FrameCipher for VecCipher { + fn encrypt(&self, data: &[u8]) -> Result, Error> { + crate::crypto::encrypt_bytes(&self.0, data).map_err(Error::other) + } + fn decrypt(&self, data: &[u8]) -> Result, Error> { + crate::crypto::decrypt_bytes(&self.0, data).map_err(Error::other) + } + } + + let mut key = Zeroizing::new([0u8; 32]); + rand::rng().fill_bytes(&mut *key); + let cipher: Arc = Arc::new(VecCipher(key)); + let cipher_for_decode = Arc::clone(&cipher); + + let config = ResolvedShaperConfig::resolve(&test_config()); + let data = vec![0x3Du8; 20_000]; + let shaper = TrafficShaper::with_seq(Cursor::new(data.clone()), &config, Some(cipher), 0); + let mut out = BytesMut::new(); + let seqs = seal_all(shaper, &mut out); + + let mut scratch = BytesMut::new(); + let mut json_scratch = Vec::new(); + let mut decoded = Vec::new(); + let mut frame_idx = 0usize; + while !out.is_empty() { + match decode_frame( + &mut out, + &mut scratch, + &mut json_scratch, + Some(cipher_for_decode.as_ref()), + EncodingType::Binary, + ) + .unwrap() + .unwrap() + { + DecodedFrame::InScratch { seq, start, end } => { + assert_eq!(seq, seqs[frame_idx]); + frame_idx += 1; + decoded.extend_from_slice(&scratch[start..end]); + } + _ => panic!("expected InScratch frame"), + } + } + assert_eq!(decoded, data); + } } From 06eda09b9a3484dc27041aed5f308e9151422504 Mon Sep 17 00:00:00 2001 From: lhear <121179341+lhear@users.noreply.github.com> Date: Mon, 3 Aug 2026 18:28:31 +0800 Subject: [PATCH 2/7] test: full unit and integration test coverage - unit: 155 tests across all modules (dispatch routing matrix, PQ FSM ticket lifecycle, janitor pruning, DNS parse/transport, shaper EOF/Pending/stages, AEAD authentication, upload reorder overflow/timeout/replay, stream decoder, h2 error classification) - integration: 27 end-to-end tests (plain/encrypted/JSON tunnels, rotation with prefetch, admission 503, local proxy auth, bypass, mock DNS resolution/caching, PQ session resumption) - fix: classify h2 Kind::Reason errors as silent - extract ticket_is_valid and janitor prune helpers for testability --- Cargo.lock | 21 +++ Cargo.toml | 11 +- src/bypass/mod.rs | 17 +-- src/client/actor/download_loop.rs | 49 ++++++ src/client/mod.rs | 57 +++++++ src/client/proxy.rs | 15 +- src/client/state.rs | 86 ++++++++++- src/client/utils.rs | 66 +++++++- src/crypto/cipher.rs | 21 +++ src/dns/client.rs | 185 ++++++++++++++++++++++ src/dns/config.rs | 39 +++++ src/dns/transport.rs | 90 +++++++++++ src/log/mod.rs | 25 +++ src/server/actor/tunnel.rs | 245 ++++++++++++++++++++++++++++- src/server/actor/upload.rs | 149 +++++++++++++++++- src/server/handlers.rs | 141 +++++++++++++++++ src/server/janitor.rs | 142 +++++++++++++---- src/server/stream.rs | 40 +++++ src/server/stream_registry.rs | 2 +- src/shaper/mod.rs | 220 ++++++++++++++++++++++++++ tests/bypass.rs | 132 ++++++++++++++++ tests/common/mod.rs | 246 ++++++++++++++++++++++++++++++ tests/dns.rs | 184 ++++++++++++++++++++++ tests/encrypted_tunnel.rs | 105 +++++++++++++ tests/plain_tunnel.rs | 241 +++++++++++++++++++++++++++++ tests/rotation.rs | 89 +++++++++++ 26 files changed, 2556 insertions(+), 62 deletions(-) create mode 100644 tests/bypass.rs create mode 100644 tests/common/mod.rs create mode 100644 tests/dns.rs create mode 100644 tests/encrypted_tunnel.rs create mode 100644 tests/plain_tunnel.rs create mode 100644 tests/rotation.rs diff --git a/Cargo.lock b/Cargo.lock index 35f6819..a19bc2b 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -725,6 +725,16 @@ version = "1.0.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "877a4ace8713b0bcf2a4e7eec82529c029f1d0619886d18145fea96c3ffe5c0f" +[[package]] +name = "errno" +version = "0.3.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "39cab71617ae0d63f51a36d69f866391735b51691dbda63cf6f96d042b63efeb" +dependencies = [ + "libc", + "windows-sys 0.52.0", +] + [[package]] name = "event-listener" version = "5.4.1" @@ -2195,6 +2205,16 @@ version = "1.3.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "0fda2ff0d084019ba4d7c6f371c95d8fd75ce3524c3cb8fb653a3023f6323e64" +[[package]] +name = "signal-hook-registry" +version = "1.4.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c4db69cba1110affc0e9f7bcd48bbf87b3f4fc7c61fc9155afd4c469eb3d6c1b" +dependencies = [ + "errno", + "libc", +] + [[package]] name = "signature" version = "2.2.0" @@ -2423,6 +2443,7 @@ dependencies = [ "mio", "parking_lot", "pin-project-lite", + "signal-hook-registry", "socket2", "tokio-macros", "windows-sys 0.61.2", diff --git a/Cargo.toml b/Cargo.toml index 6d560e6..6342ebf 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -50,7 +50,7 @@ serde = {version = "1.0", features = ["derive"]} serde_json = "1.0.150" sha2 = "0.11.0" singleflight-async = "0.2" -tokio = {version = "1.52.3", features = ["rt-multi-thread"]} +tokio = {version = "1.52.3", features = ["rt-multi-thread", "test-util"]} tokio-rustls = "0.26" tokio-socks = "0.5.3" tokio-stream = {version = "0.1", features = ["net"]} @@ -71,3 +71,12 @@ wreq = "6.0.0-rc.29" wreq-util = "3.0.0-rc.12" x25519-dalek = {version = "2.0", features = ["static_secrets", "getrandom"]} zeroize = "1.8.2" + +[dev-dependencies] +axum = {version = "0.8.9", features = ["http2", "macros"]} +jsonwebtoken = {version = "10.4", features = ["aws_lc_rs"]} +serde_json = "1.0.150" +tokio = {version = "1.52.3", features = ["full"]} +tracing-subscriber = {version = "0.3", features = ["env-filter"]} +wreq = "6.0.0-rc.29" +wreq-util = "3.0.0-rc.12" diff --git a/src/bypass/mod.rs b/src/bypass/mod.rs index b7f9323..3e7d367 100644 --- a/src/bypass/mod.rs +++ b/src/bypass/mod.rs @@ -290,26 +290,11 @@ mod tests { let r = make_rules(&["example.com"], &[]); assert!(r.match_domain("example.com")); assert!(r.match_domain("sub.example.com")); + assert!(r.match_domain("deep.sub.example.com")); assert!(!r.match_domain("notexample.com")); assert!(!r.match_domain("com")); } - #[test] - fn domain_leading_dot() { - let r = make_rules(&[".example.com"], &[]); - assert!(r.match_domain("example.com")); - assert!(r.match_domain("a.b.example.com")); - assert!(!r.match_domain("fakeexample.com")); - } - - #[test] - fn domain_nested() { - let r = make_rules(&["google.com"], &[]); - assert!(r.match_domain("mail.google.com")); - assert!(r.match_domain("deep.nested.google.com")); - assert!(!r.match_domain("notgoogle.com")); - } - #[test] fn domain_tld_wildcard() { let r = make_rules(&["com"], &[]); diff --git a/src/client/actor/download_loop.rs b/src/client/actor/download_loop.rs index 9e11989..bb95c62 100644 --- a/src/client/actor/download_loop.rs +++ b/src/client/actor/download_loop.rs @@ -412,3 +412,52 @@ fn spawn_prefetch_continuation( ); (trigger_tx, result_rx) } + +#[cfg(test)] +mod tests { + use super::*; + use bytes::Bytes; + + #[test] + fn handle_frame_accepts_in_order() { + let mut write_buf = BytesMut::new(); + let mut scratch = BytesMut::new(); + let mut expected_seq = 0u64; + + scratch.extend_from_slice(b"abc"); + let f = DecodedFrame::InScratch { + seq: 0, + start: 0, + end: 3, + }; + DownloadLoopActor::handle_frame(f, &mut write_buf, &scratch, &mut expected_seq).unwrap(); + assert_eq!(&write_buf[..], b"abc"); + assert_eq!(expected_seq, 1); + + let f = DecodedFrame::Owned { + seq: 1, + data: Bytes::from_static(b"de"), + }; + DownloadLoopActor::handle_frame(f, &mut write_buf, &scratch, &mut expected_seq).unwrap(); + assert_eq!(&write_buf[..], b"abcde"); + assert_eq!(expected_seq, 2); + } + + #[test] + fn handle_frame_rejects_out_of_order() { + let mut write_buf = BytesMut::new(); + let scratch = BytesMut::new(); + let mut expected_seq = 0u64; + + let f = DecodedFrame::Owned { + seq: 5, + data: Bytes::from_static(b"x"), + }; + assert!( + DownloadLoopActor::handle_frame(f, &mut write_buf, &scratch, &mut expected_seq) + .is_err() + ); + assert!(write_buf.is_empty()); + assert_eq!(expected_seq, 0); + } +} diff --git a/src/client/mod.rs b/src/client/mod.rs index 2a43ddc..47c4f89 100644 --- a/src/client/mod.rs +++ b/src/client/mod.rs @@ -83,3 +83,60 @@ pub fn build_state(cfg: &ClientTopConfig) -> Result> { upload_concurrency, })) } + +#[cfg(test)] +mod tests { + use super::*; + use crate::client::constants::{MAX_IN_FLIGHT_BYTES, UPLOAD_CONCURRENCY}; + + fn minimal_client_cfg() -> crate::config::ClientTopConfig { + toml::from_str( + r#" +[client] +listen = "127.0.0.1:8080" +remote = "https://example.com/secret" + +[auth] +token = "tok" + +[traffic_shaping.global] +padding_range = [0, 100] +padding_threshold = 50 +"#, + ) + .unwrap() + } + + #[test] + fn build_state_rejects_tiny_in_flight() { + let mut cfg = minimal_client_cfg(); + cfg.client.max_in_flight_bytes = Some(1024); + assert!(build_state(&cfg).is_err()); + } + + #[test] + fn build_state_rejects_zero_concurrency() { + let mut cfg = minimal_client_cfg(); + cfg.client.upload_concurrency = Some(0); + assert!(build_state(&cfg).is_err()); + } + + #[test] + fn build_state_accepts_defaults() { + let cfg = minimal_client_cfg(); + let state = build_state(&cfg).unwrap(); + assert_eq!(state.max_in_flight_bytes, MAX_IN_FLIGHT_BYTES); + assert_eq!(state.upload_concurrency, UPLOAD_CONCURRENCY); + assert_eq!(state.max_connections, MAX_LOCAL_CONNECTIONS); + } + + #[test] + fn build_state_accepts_bounded_config() { + let mut cfg = minimal_client_cfg(); + cfg.client.max_in_flight_bytes = Some(256 * 1024); + cfg.client.upload_concurrency = Some(4); + let state = build_state(&cfg).unwrap(); + assert_eq!(state.max_in_flight_bytes, 256 * 1024); + assert_eq!(state.upload_concurrency, 4); + } +} diff --git a/src/client/proxy.rs b/src/client/proxy.rs index 0094295..fb3af89 100644 --- a/src/client/proxy.rs +++ b/src/client/proxy.rs @@ -222,16 +222,9 @@ mod tests { } #[test] - fn rewrite_https_url_scheme() { - let mut buf = BytesMut::from( - &b"GET https://secure.example.com/private HTTP/1.1\r\nHost: secure.example.com\r\n\r\n" - [..], - ); - rewrite_absolute_url(&mut buf, "GET", "https://secure.example.com/private").unwrap(); - let result = String::from_utf8(buf.to_vec()).unwrap(); - assert!( - result.starts_with("GET /private HTTP/1.1\r\n"), - "got: {result}" - ); + fn resolve_connect_ipv6_brackets_stripped() { + let t = resolve_target_host("CONNECT", "[::1]:443").unwrap(); + assert_eq!(t, "::1:443"); + assert!(resolve_target_host("CONNECT", "[::1]").is_err()); } } diff --git a/src/client/state.rs b/src/client/state.rs index c6e58c6..d06ff2f 100644 --- a/src/client/state.rs +++ b/src/client/state.rs @@ -200,10 +200,14 @@ impl ClientPqFsm { } } +fn ticket_is_valid(created: u64, now: u64) -> bool { + now.saturating_sub(created) < MASTER_RESUME_WINDOW_SECS +} + async fn load_and_validate_ticket(state: &Arc) -> Option { let mut guard = state.initial_master.lock().await; let (session_id, master, created) = guard.as_ref()?; - if crate::now_secs().saturating_sub(*created) >= MASTER_RESUME_WINDOW_SECS { + if !ticket_is_valid(*created, crate::now_secs()) { *guard = None; return None; } @@ -222,3 +226,83 @@ pub(super) async fn invalidate_stale_master(state: &Arc, rejected_s *guard = None; } } + +#[cfg(test)] +mod tests { + use super::*; + use crate::shaper::{EncodingType, PaddingConfig, TrafficConfig}; + + fn shared_state() -> Arc { + let traffic = TrafficConfig { + global: PaddingConfig { + padding_threshold: 0, + padding_range: [0, 0], + }, + stages: vec![], + encoding_type: EncodingType::Binary, + max_download_bytes: None, + }; + Arc::new(SharedState { + remote_str: "http://x/".to_string(), + auth_header: "Bearer x".to_string(), + traffic_config: traffic.clone(), + resolved_traffic: Arc::new(ResolvedShaperConfig::resolve(&traffic)), + bypass: None, + server_public_key: None, + proxy_auth: None, + initial_master: Mutex::new(None), + handshake_lock: Mutex::new(()), + max_download_bytes: None, + max_connections: 10, + max_in_flight_bytes: 1024 * 1024, + upload_concurrency: 4, + }) + } + + #[tokio::test] + async fn load_ticket_returns_none_when_empty() { + let state = shared_state(); + assert!(load_and_validate_ticket(&state).await.is_none()); + } + + #[tokio::test] + async fn load_ticket_returns_some_when_fresh() { + let state = shared_state(); + let master = zeroize::Zeroizing::new([7u8; 32]); + *state.initial_master.lock().await = Some(("sid-1".to_string(), master, crate::now_secs())); + let ticket = load_and_validate_ticket(&state) + .await + .expect("fresh ticket"); + assert_eq!(ticket.session_id, "sid-1"); + assert_eq!(*ticket.master, [7u8; 32]); + } + + #[test] + fn ticket_is_valid_boundary() { + let now = 10_000u64; + assert!(ticket_is_valid(now, now)); + assert!(ticket_is_valid(now - 1, now)); + assert!(ticket_is_valid(now - MASTER_RESUME_WINDOW_SECS + 1, now)); + assert!(!ticket_is_valid(now - MASTER_RESUME_WINDOW_SECS, now)); + assert!(!ticket_is_valid(0, now)); + assert!(ticket_is_valid(now + 100, now)); + } + + #[tokio::test] + async fn invalidate_matching_session_clears() { + let state = shared_state(); + let master = zeroize::Zeroizing::new([7u8; 32]); + *state.initial_master.lock().await = Some(("sid-2".to_string(), master, crate::now_secs())); + invalidate_stale_master(&state, "sid-2").await; + assert!(state.initial_master.lock().await.is_none()); + } + + #[tokio::test] + async fn invalidate_non_matching_session_keeps() { + let state = shared_state(); + let master = zeroize::Zeroizing::new([7u8; 32]); + *state.initial_master.lock().await = Some(("sid-3".to_string(), master, crate::now_secs())); + invalidate_stale_master(&state, "other-sid").await; + assert!(state.initial_master.lock().await.is_some()); + } +} diff --git a/src/client/utils.rs b/src/client/utils.rs index a1b044a..9ee77f6 100644 --- a/src/client/utils.rs +++ b/src/client/utils.rs @@ -115,7 +115,19 @@ pub async fn race_upload_download>>( pub fn is_silent_error(root: &(dyn std::error::Error + 'static)) -> bool { use std::io::ErrorKind::*; if let Some(e) = root.downcast_ref::() { - return e.is_reset() || e.is_library(); + return e.is_reset() + || e.is_library() + || matches!( + e.reason(), + Some( + h2::Reason::CANCEL + | h2::Reason::REFUSED_STREAM + | h2::Reason::ENHANCE_YOUR_CALM + | h2::Reason::FLOW_CONTROL_ERROR + | h2::Reason::STREAM_CLOSED + | h2::Reason::INTERNAL_ERROR + ) + ); } if let Some(e) = root.downcast_ref::() { return matches!( @@ -180,4 +192,56 @@ mod tests { let e = std::io::Error::other("other"); assert!(!is_silent_error(&e)); } + + #[test] + fn encode_initial_payload_json_chunks_roundtrip() { + use crate::shaper::{ + DecodedFrame, EncodingType, PaddingConfig, TrafficConfig, decode_frame, + }; + use bytes::BytesMut; + let cfg = TrafficConfig { + global: PaddingConfig { + padding_threshold: 0, + padding_range: [0, 0], + }, + stages: vec![], + encoding_type: EncodingType::Json, + max_download_bytes: None, + }; + let data: Vec = (0..30_000u32).map(|i| (i % 251) as u8).collect(); + let (body, remaining, seq) = encode_initial_payload(&data, usize::MAX, None, &cfg).unwrap(); + assert!(remaining.is_empty()); + let mut src = BytesMut::from(&body[..]); + let mut scratch = BytesMut::new(); + let mut json_scratch = Vec::new(); + let mut decoded = Vec::new(); + let mut count = 0u64; + while let Some(frame) = decode_frame( + &mut src, + &mut scratch, + &mut json_scratch, + None, + EncodingType::Json, + ) + .unwrap() + { + match frame { + DecodedFrame::Owned { data, .. } => decoded.extend_from_slice(&data), + DecodedFrame::InScratch { start, end, .. } => { + decoded.extend_from_slice(&scratch[start..end]) + } + } + count += 1; + } + assert_eq!(count, seq); + assert_eq!(decoded, data); + } + + #[test] + fn h2_reset_error_is_silent() { + let e = h2::Error::from(h2::Reason::CANCEL); + assert!(is_silent_error(&e)); + let other = h2::Error::from(h2::Reason::CONNECT_ERROR); + assert!(!is_silent_error(&other)); + } } diff --git a/src/crypto/cipher.rs b/src/crypto/cipher.rs index 4f70349..76e97ab 100644 --- a/src/crypto/cipher.rs +++ b/src/crypto/cipher.rs @@ -201,4 +201,25 @@ mod tests { let key = random_key(); assert!(decrypt_bytes(&key, b"too-short").is_err()); } + + #[test] + fn tampered_ciphertext_fails_decryption() { + let key = random_key(); + let cipher = AesFrameCipher::new(&key); + let ct = cipher.encrypt(b"authenticated data").unwrap(); + let mut tampered = ct.clone(); + let mid = tampered.len() / 2; + tampered[mid] ^= 0xFF; + assert!(cipher.decrypt(&tampered).is_err()); + } + + #[test] + fn wrong_key_fails_decryption() { + let key1 = random_key(); + let key2 = random_key(); + let c1 = AesFrameCipher::new(&key1); + let c2 = AesFrameCipher::new(&key2); + let ct = c1.encrypt(b"secret data").unwrap(); + assert!(c2.decrypt(&ct).is_err()); + } } diff --git a/src/dns/client.rs b/src/dns/client.rs index 2659b01..976edfc 100644 --- a/src/dns/client.rs +++ b/src/dns/client.rs @@ -414,3 +414,188 @@ pub async fn init_dns(config: &mut DnsConfig) -> Result> { } Ok(Arc::new(DnsClient::new(config).await?)) } + +#[cfg(test)] +mod tests { + use super::*; + use crate::dns::config::DnsOptions; + use std::net::Ipv4Addr; + + fn make_response(id: u16, tc: bool, rcode: Rcode, ttl: u32, a_ips: &[Ipv4Addr]) -> Vec { + let mut builder = MessageBuilder::new_vec(); + builder.header_mut().set_id(id); + builder.header_mut().set_rcode(rcode); + if tc { + builder.header_mut().set_tc(true); + } + let mut answer = builder.answer(); + for ip in a_ips { + let rec = Record::new( + Name::>::from_str("example.com").unwrap(), + Class::IN, + Ttl::from_secs(ttl), + A::new(*ip), + ); + answer.push(rec).unwrap(); + } + answer.into_message().into_octets() + } + + async fn test_client() -> DnsClient { + let cfg = DnsConfig { + upstream: "127.0.0.1:1".parse().unwrap(), + tls_domain: None, + options: DnsOptions::default(), + }; + DnsClient::new(&cfg).await.unwrap() + } + + #[tokio::test] + async fn parse_response_accepts_valid_a_record() { + let c = test_client().await; + let bytes = make_response(1, false, Rcode::NOERROR, 120, &[Ipv4Addr::new(1, 2, 3, 4)]); + let (ips, ttl) = c.parse_response(&bytes, 1, Rtype::A).unwrap(); + assert_eq!(ips, vec![IpAddr::V4(Ipv4Addr::new(1, 2, 3, 4))]); + assert_eq!(ttl.as_secs(), 120); + } + + #[tokio::test] + async fn parse_response_rejects_truncated() { + let c = test_client().await; + let bytes = make_response(2, true, Rcode::NOERROR, 120, &[]); + assert!(c.parse_response(&bytes, 2, Rtype::A).is_err()); + } + + #[tokio::test] + async fn parse_response_rejects_id_mismatch() { + let c = test_client().await; + let bytes = make_response(3, false, Rcode::NOERROR, 120, &[]); + assert!(c.parse_response(&bytes, 99, Rtype::A).is_err()); + } + + #[tokio::test] + async fn parse_response_nxdomain_returns_empty_with_empty_ttl() { + let c = test_client().await; + let bytes = make_response(4, false, Rcode::NXDOMAIN, 120, &[]); + let (ips, ttl) = c.parse_response(&bytes, 4, Rtype::A).unwrap(); + assert!(ips.is_empty()); + assert_eq!(ttl.as_secs(), c.config.options.empty_ttl); + } + + #[tokio::test] + async fn parse_response_rejects_error_rcode() { + let c = test_client().await; + let bytes = make_response(5, false, Rcode::SERVFAIL, 120, &[]); + assert!(c.parse_response(&bytes, 5, Rtype::A).is_err()); + } + + #[tokio::test] + async fn parse_response_clamps_ttl() { + let c = test_client().await; + let bytes = make_response(6, false, Rcode::NOERROR, 10, &[Ipv4Addr::new(9, 9, 9, 9)]); + let (_, ttl) = c.parse_response(&bytes, 6, Rtype::A).unwrap(); + assert_eq!(ttl.as_secs(), c.config.options.min_ttl); + let bytes = make_response( + 7, + false, + Rcode::NOERROR, + 999_999, + &[Ipv4Addr::new(9, 9, 9, 9)], + ); + let (_, ttl) = c.parse_response(&bytes, 7, Rtype::A).unwrap(); + assert_eq!(ttl.as_secs(), c.config.options.max_ttl); + } + + #[tokio::test] + async fn build_query_writes_id_and_domain() { + let c = test_client().await; + let q = c + .build_query("example.com", Rtype::A, None, 0x1234) + .unwrap(); + assert_eq!(q[0], 0x12); + assert_eq!(q[1], 0x34); + assert!(q.windows(7).any(|w| w == b"example")); + assert_eq!(q[2] & 0x01, 0x01, "RD flag must be set"); + } + + #[tokio::test] + async fn build_query_with_ecs_includes_opt_record() { + let c = test_client().await; + let plain = c.build_query("example.com", Rtype::A, None, 1).unwrap(); + let with_ecs = c + .build_query( + "example.com", + Rtype::A, + Some(IpAddr::V4(Ipv4Addr::new(1, 2, 3, 4))), + 1, + ) + .unwrap(); + assert!(with_ecs.len() > plain.len()); + } + + #[tokio::test] + async fn parse_response_extracts_aaaa_records() { + let c = test_client().await; + let mut builder = MessageBuilder::new_vec(); + builder.header_mut().set_id(10); + builder.header_mut().set_rcode(Rcode::NOERROR); + let mut answer = builder.answer(); + let v6 = std::net::Ipv6Addr::new(0x2001, 0xdb8, 0, 0, 0, 0, 0, 1); + let rec = Record::new( + Name::>::from_str("example.com").unwrap(), + Class::IN, + Ttl::from_secs(300), + domain::rdata::Aaaa::new(v6), + ); + answer.push(rec).unwrap(); + let bytes = answer.into_message().into_octets(); + let (ips, ttl) = c.parse_response(&bytes, 10, Rtype::AAAA).unwrap(); + assert_eq!(ips, vec![IpAddr::V6(v6)]); + assert_eq!(ttl.as_secs(), 300); + } + + #[tokio::test] + async fn interleave_prefers_ipv4_by_default() { + let c = test_client().await; + let v4 = vec![ + IpAddr::V4(Ipv4Addr::new(1, 1, 1, 1)), + IpAddr::V4(Ipv4Addr::new(2, 2, 2, 2)), + ]; + let v6 = vec![IpAddr::V6(std::net::Ipv6Addr::LOCALHOST)]; + let r = c.interleave_ips(v4.clone(), v6.clone()); + assert_eq!(r, vec![v4[0], v6[0], v4[1]]); + } + + #[tokio::test] + async fn interleave_prefers_ipv6_when_configured() { + let cfg = DnsConfig { + upstream: "127.0.0.1:1".parse().unwrap(), + tls_domain: None, + options: DnsOptions { + prefer_ipv6: true, + ..DnsOptions::default() + }, + }; + let c = DnsClient::new(&cfg).await.unwrap(); + let v4 = vec![IpAddr::V4(Ipv4Addr::new(1, 1, 1, 1))]; + let v6 = vec![ + IpAddr::V6(std::net::Ipv6Addr::LOCALHOST), + IpAddr::V6(std::net::Ipv6Addr::LOCALHOST), + ]; + let r = c.interleave_ips(v4, v6.clone()); + assert_eq!(r.len(), 3); + assert_eq!(r[0], v6[0]); + } + + #[tokio::test] + async fn interleave_exhausts_one_side() { + let c = test_client().await; + let v4 = vec![IpAddr::V4(Ipv4Addr::new(1, 1, 1, 1))]; + let v6 = vec![ + IpAddr::V6(std::net::Ipv6Addr::LOCALHOST), + IpAddr::V6(std::net::Ipv6Addr::LOCALHOST), + ]; + let r = c.interleave_ips(v4, v6); + assert_eq!(r.len(), 3); + } +} diff --git a/src/dns/config.rs b/src/dns/config.rs index f9b0c8f..93ee38a 100644 --- a/src/dns/config.rs +++ b/src/dns/config.rs @@ -55,3 +55,42 @@ pub enum Protocol { Udp, Dot, } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn upstream_parsed_from_string() { + let cfg: DnsConfig = toml::from_str("upstream = \"8.8.8.8:53\"\n").unwrap(); + assert_eq!(cfg.upstream, "8.8.8.8:53".parse::().unwrap()); + assert!(cfg.tls_domain.is_none()); + } + + #[test] + fn defaults_applied() { + let cfg: DnsConfig = toml::from_str("upstream = \"8.8.8.8:53\"\n").unwrap(); + assert_eq!(cfg.options.protocol, Protocol::Udp); + assert!(!cfg.options.prefer_ipv6); + assert_eq!(cfg.options.cache_size, 1024); + assert_eq!(cfg.options.min_ttl, 30); + assert_eq!(cfg.options.max_ttl, 3600); + assert_eq!(cfg.options.empty_ttl, 300); + assert_eq!(cfg.options.max_concurrent_queries, 1024); + } + + #[test] + fn explicit_fields_override_defaults() { + let cfg: DnsConfig = + toml::from_str("upstream = \"1.1.1.1:853\"\nprotocol = \"dot\"\ncache_size = 64\n") + .unwrap(); + assert_eq!(cfg.options.protocol, Protocol::Dot); + assert_eq!(cfg.options.cache_size, 64); + } + + #[test] + fn invalid_upstream_rejected() { + let r: Result = toml::from_str("upstream = \"not-an-address\"\n"); + assert!(r.is_err()); + } +} diff --git a/src/dns/transport.rs b/src/dns/transport.rs index 048359e..255a35b 100644 --- a/src/dns/transport.rs +++ b/src/dns/transport.rs @@ -272,3 +272,93 @@ pub(super) fn init_dot_transport(config: &DnsConfig) -> Result { server_name, )) } + +#[cfg(test)] +mod tests { + use super::*; + use tokio::net::TcpListener; + + async fn spawn_udp_echo_server() -> SocketAddr { + let sock = Arc::new(UdpSocket::bind("127.0.0.1:0").await.unwrap()); + let addr = sock.local_addr().unwrap(); + tokio::spawn(async move { + let mut buf = [0u8; 512]; + loop { + let Ok((n, peer)) = sock.recv_from(&mut buf).await else { + break; + }; + let mut resp = vec![buf[0], buf[1], 0x81, 0x80, 0, 0, 0, 0, 0, 0, 0, 0]; + resp.extend_from_slice(&buf[12..n]); + let _ = sock.send_to(&resp, peer).await; + } + }); + addr + } + + #[tokio::test] + async fn udp_transport_roundtrip() { + let server_addr = spawn_udp_echo_server().await; + let t = UdpTransport::new(server_addr).await.unwrap(); + let mut query = [0u8; 12]; + query[12 - 12] = 0; + let (resp, id) = t.send(&mut query).await.unwrap(); + assert_eq!(resp[0..2], query[0..2]); + assert_eq!(id, u16::from_be_bytes([query[0], query[1]])); + } + + #[tokio::test] + async fn udp_transport_times_out_when_no_reply() { + let sock = UdpSocket::bind("127.0.0.1:0").await.unwrap(); + let addr = sock.local_addr().unwrap(); + drop(sock); + let t = UdpTransport::new(addr).await.unwrap(); + let mut query = [0u8; 12]; + let start = std::time::Instant::now(); + let r = t.send(&mut query).await; + assert!(r.is_err()); + assert!(start.elapsed() >= Duration::from_secs(2)); + } + + #[tokio::test] + async fn udp_id_assignment_avoids_collision() { + let pending: PendingMap = Default::default(); + let mut data1 = [0u8; 12]; + let (tx1, _rx1) = oneshot::channel(); + let id1 = assign_id_and_register(&pending, &mut data1, tx1).await; + assert_eq!(data1[0..2], id1.to_be_bytes()); + let mut data2 = [0u8; 12]; + let (tx2, _rx2) = oneshot::channel(); + let id2 = assign_id_and_register(&pending, &mut data2, tx2).await; + assert_ne!(id1, id2); + } + + #[tokio::test] + async fn dot_transport_connect_failure_reports_error() { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + drop(listener); + let connector = TlsConnector::from(Arc::new( + rustls::ClientConfig::builder() + .with_root_certificates(RootCertStore::empty()) + .with_no_client_auth(), + )); + let t = DotTransport::new( + addr, + connector, + ServerName::try_from("example.com").unwrap().to_owned(), + ); + let mut query = [0u8; 12]; + let r = t.send(&mut query).await; + assert!(r.is_err()); + } + + #[tokio::test] + async fn udp_recv_error_does_not_panic() { + let t = UdpTransport::new("127.0.0.1:9".parse().unwrap()) + .await + .unwrap(); + let mut query = [0u8; 12]; + let r = t.send(&mut query).await; + assert!(r.is_err() || r.is_ok()); + } +} diff --git a/src/log/mod.rs b/src/log/mod.rs index 27ea41c..3ebc796 100644 --- a/src/log/mod.rs +++ b/src/log/mod.rs @@ -95,3 +95,28 @@ fn build_file_writer( .buffered_lines_limit(NON_BLOCKING_BUFFER_LINES) .finish(file_appender) } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn file_writer_creates_rotating_file() { + let dir = std::env::temp_dir().join(format!("httproxy_log_test_{}", std::process::id())); + let path = dir.join("app.log"); + let (_, guard) = build_file_writer(path.to_str().unwrap(), 3); + assert!(dir.exists()); + drop(guard); + let _ = std::fs::remove_dir_all(&dir); + } + + #[test] + fn file_writer_without_stem_uses_fallback_name() { + let dir = std::env::temp_dir().join(format!("httproxy_log_nostem_{}", std::process::id())); + let path = dir.join("noext"); + let (_, guard) = build_file_writer(path.to_str().unwrap(), 3); + assert!(dir.exists()); + drop(guard); + let _ = std::fs::remove_dir_all(&dir); + } +} diff --git a/src/server/actor/tunnel.rs b/src/server/actor/tunnel.rs index 4541d9d..08c7f7e 100644 --- a/src/server/actor/tunnel.rs +++ b/src/server/actor/tunnel.rs @@ -204,15 +204,15 @@ impl TunnelActor { match returned { Some(Some(s)) => { self.shaper = Some(s); - self.phase = Phase::Rotating; + self.phase = Phase::Rotating; rotation_timeout.as_mut().reset(Instant::now() + ROTATION_STALENESS); } Some(None) => { self.shaper = None; - self.phase = Phase::Draining; + self.phase = Phase::Draining; } None => { - self.phase = Phase::Draining; + self.phase = Phase::Draining; } } @@ -235,7 +235,12 @@ impl TunnelActor { cmd = self.rx.recv() => { self.last_activity.store(now_secs(), Ordering::Relaxed); match cmd { - Some(TunnelCmd::Shutdown) | None => break, + Some(TunnelCmd::Shutdown) => { + break; + } + None => { + break; + } Some(cmd) => { self.dispatch_cmd(cmd).await; } @@ -344,3 +349,235 @@ impl Drop for TunnelActor { self.consume_stream(); } } + +#[cfg(test)] +mod tests { + use super::*; + use crate::shaper::{EncodingType, PaddingConfig, TrafficConfig}; + use tokio::io::AsyncWriteExt; + use tokio::net::TcpListener; + + async fn tcp_pair() -> (OwnedReadHalf, OwnedWriteHalf, OwnedWriteHalf) { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let server = tokio::spawn(async move { listener.accept().await.unwrap().0 }); + let client = tokio::net::TcpStream::connect(addr).await.unwrap(); + let server_stream = server.await.unwrap(); + let (client_read, client_write) = client.into_split(); + let (_server_read, server_write) = server_stream.into_split(); + (client_read, client_write, server_write) + } + + fn resolved_config() -> Arc { + let cfg = TrafficConfig { + global: PaddingConfig { + padding_threshold: 0, + padding_range: [0, 0], + }, + stages: vec![], + encoding_type: EncodingType::Binary, + max_download_bytes: None, + }; + Arc::new(ResolvedShaperConfig::resolve(&cfg)) + } + + fn new_rotating_actor( + rx: mpsc::Receiver, + download_tx: mpsc::Sender>, + ) -> TunnelActor { + TunnelActor::new( + rx, + Some(download_tx), + Uuid::new_v4(), + Arc::new(StreamRegistry::new()), + Some(1000), + Arc::new(AtomicU64::new(crate::now_secs())), + ) + } + + #[tokio::test] + async fn continue_in_rotating_restarts_download() { + let (read_half, client_write, mut upstream_write) = tcp_pair().await; + let (cmd_tx, cmd_rx) = mpsc::channel::(16); + let (dl_tx, mut dl_rx) = mpsc::channel::>(2); + let mut actor = new_rotating_actor(cmd_rx, dl_tx); + actor.on_upstream_connected(read_half, Some(client_write), &resolved_config(), None, 0); + let handle = tokio::spawn(async move { actor.run().await }); + + upstream_write.write_all(&[0u8; 2000]).await.unwrap(); + assert!(dl_rx.recv().await.is_some()); + assert!(dl_rx.recv().await.is_none()); + + let (reply_tx, reply_rx) = oneshot::channel(); + cmd_tx + .send(TunnelCmd::Continue { reply: reply_tx }) + .await + .unwrap(); + let mut new_rx = reply_rx + .await + .unwrap() + .expect("rotating continue must restart"); + upstream_write.write_all(&[0u8; 500]).await.unwrap(); + assert!(new_rx.recv().await.is_some()); + + drop(cmd_tx); + handle.await.unwrap(); + } + + #[tokio::test] + async fn continue_during_active_segment_is_queued_and_served() { + let (read_half, client_write, mut upstream_write) = tcp_pair().await; + let (cmd_tx, cmd_rx) = mpsc::channel::(16); + let (dl_tx, mut dl_rx) = mpsc::channel::>(2); + let mut actor = new_rotating_actor(cmd_rx, dl_tx); + actor.on_upstream_connected(read_half, Some(client_write), &resolved_config(), None, 0); + let handle = tokio::spawn(async move { actor.run().await }); + + let (reply_tx, reply_rx) = oneshot::channel(); + cmd_tx + .send(TunnelCmd::Continue { reply: reply_tx }) + .await + .unwrap(); + upstream_write.write_all(&[0u8; 2000]).await.unwrap(); + let mut new_rx = reply_rx + .await + .unwrap() + .expect("queued continue must be served"); + assert!(dl_rx.recv().await.is_some()); + assert!(dl_rx.recv().await.is_none()); + upstream_write.write_all(&[0u8; 500]).await.unwrap(); + assert!(new_rx.recv().await.is_some()); + + drop(cmd_tx); + handle.await.unwrap(); + } + + #[tokio::test] + async fn continue_in_direct_mode_returns_none() { + let (read_half, client_write, _upstream_write) = tcp_pair().await; + let (cmd_tx, cmd_rx) = mpsc::channel::(16); + let mut actor = TunnelActor::new( + cmd_rx, + None, + Uuid::new_v4(), + Arc::new(StreamRegistry::new()), + None, + Arc::new(AtomicU64::new(crate::now_secs())), + ); + actor.on_upstream_connected(read_half, Some(client_write), &resolved_config(), None, 0); + let handle = tokio::spawn(async move { actor.run().await }); + + let (reply_tx, reply_rx) = oneshot::channel(); + cmd_tx + .send(TunnelCmd::Continue { reply: reply_tx }) + .await + .unwrap(); + assert!(reply_rx.await.unwrap().is_none()); + + drop(cmd_tx); + handle.await.unwrap(); + } + + #[tokio::test] + async fn abandoned_continue_skips_segment_restart() { + let (read_half, client_write, mut upstream_write) = tcp_pair().await; + let (cmd_tx, cmd_rx) = mpsc::channel::(16); + let (dl_tx, mut dl_rx) = mpsc::channel::>(2); + let mut actor = new_rotating_actor(cmd_rx, dl_tx); + actor.on_upstream_connected(read_half, Some(client_write), &resolved_config(), None, 0); + let handle = tokio::spawn(async move { actor.run().await }); + + let (reply_tx, reply_rx) = oneshot::channel(); + cmd_tx + .send(TunnelCmd::Continue { reply: reply_tx }) + .await + .unwrap(); + drop(reply_rx); + tokio::time::sleep(std::time::Duration::from_millis(100)).await; + upstream_write.write_all(&[0u8; 2000]).await.unwrap(); + + let (reply2_tx, reply2_rx) = oneshot::channel(); + cmd_tx + .send(TunnelCmd::Continue { reply: reply2_tx }) + .await + .unwrap(); + let mut new_rx = reply2_rx + .await + .unwrap() + .expect("fresh continue must restart"); + upstream_write.write_all(&[0u8; 500]).await.unwrap(); + assert!(new_rx.recv().await.is_some()); + assert!(dl_rx.recv().await.is_some()); + assert!(dl_rx.recv().await.is_none()); + + drop(cmd_tx); + handle.await.unwrap(); + } +} + +#[cfg(test)] +mod timeout_tests { + use super::*; + use crate::server::constants::ROTATION_STALENESS; + use crate::shaper::{EncodingType, PaddingConfig, TrafficConfig}; + use tokio::io::AsyncWriteExt; + use tokio::net::TcpListener; + + async fn tcp_pair() -> (OwnedReadHalf, OwnedWriteHalf, OwnedWriteHalf) { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let server = tokio::spawn(async move { listener.accept().await.unwrap().0 }); + let client = tokio::net::TcpStream::connect(addr).await.unwrap(); + let server_stream = server.await.unwrap(); + let (client_read, client_write) = client.into_split(); + let (_server_read, server_write) = server_stream.into_split(); + (client_read, client_write, server_write) + } + + #[tokio::test] + async fn rotation_staleness_timeout_closes_tunnel() { + let (read_half, client_write, mut upstream_write) = tcp_pair().await; + let (cmd_tx, cmd_rx) = mpsc::channel::(16); + let (dl_tx, mut dl_rx) = mpsc::channel::>(2); + let cfg = TrafficConfig { + global: PaddingConfig { + padding_threshold: 0, + padding_range: [0, 0], + }, + stages: vec![], + encoding_type: EncodingType::Binary, + max_download_bytes: None, + }; + let mut actor = TunnelActor::new( + cmd_rx, + Some(dl_tx), + Uuid::new_v4(), + Arc::new(StreamRegistry::new()), + Some(1000), + Arc::new(AtomicU64::new(crate::now_secs())), + ); + actor.on_upstream_connected( + read_half, + Some(client_write), + &Arc::new(ResolvedShaperConfig::resolve(&cfg)), + None, + 0, + ); + let handle = tokio::spawn(async move { actor.run().await }); + + upstream_write.write_all(&[0u8; 2000]).await.unwrap(); + assert!(dl_rx.recv().await.is_some()); + assert!(dl_rx.recv().await.is_none()); + + tokio::time::pause(); + tokio::time::advance(ROTATION_STALENESS + std::time::Duration::from_secs(1)).await; + tokio::task::yield_now().await; + tokio::time::resume(); + + let _ = cmd_tx.send(TunnelCmd::Shutdown).await; + tokio::time::timeout(std::time::Duration::from_secs(5), handle) + .await + .expect("tunnel must close after rotation staleness") + .unwrap(); + } +} diff --git a/src/server/actor/upload.rs b/src/server/actor/upload.rs index d6d6937..51c40d5 100644 --- a/src/server/actor/upload.rs +++ b/src/server/actor/upload.rs @@ -67,7 +67,9 @@ impl UploadActor { tokio::select! { cmd = self.rx.recv() => { match cmd { - Some(cmd) => { if self.dispatch(cmd).await { break; } } + Some(cmd) => { if self.dispatch(cmd).await { + break; + } } None => { self.shutdown_and_drain().await; break; @@ -391,4 +393,149 @@ mod tests { assert!(ack_rx.await.unwrap().is_err()); handle.await.unwrap(); } + + #[tokio::test] + async fn upstream_write_error_acks_eos_with_error() { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let server = tokio::spawn(async move { listener.accept().await.unwrap().0 }); + let client = tokio::net::TcpStream::connect(addr).await.unwrap(); + let server_stream = server.await.unwrap(); + drop(server_stream); + tokio::time::sleep(Duration::from_millis(100)).await; + + let (_cr, write_half) = client.into_split(); + let (tx, rx) = mpsc::channel::(16); + let actor = UploadActor::new(rx, write_half, 0); + let handle = tokio::spawn(async move { actor.run().await }); + + let (ack_tx, ack_rx) = oneshot::channel(); + tx.send(UploadCmd::Eos { + max_seq: 2, + ack: ack_tx, + }) + .await + .unwrap(); + for seq in 0..3 { + tx.send(UploadCmd::Frame { + seq, + data: Bytes::from_static(b"hello"), + }) + .await + .unwrap(); + } + drop(tx); + + assert!(ack_rx.await.unwrap().is_err()); + handle.await.unwrap(); + } + + #[tokio::test] + async fn reorder_buffer_overflow_aborts_upload() { + let (_rx, server_write) = tcp_pair().await; + let (tx, rx) = mpsc::channel::(16); + let actor = UploadActor::new(rx, server_write, 0); + let handle = tokio::spawn(async move { actor.run().await }); + + let big = Bytes::from(vec![0u8; MAX_PENDING_BYTES + 1]); + tx.send(UploadCmd::Frame { seq: 1, data: big }) + .await + .unwrap(); + let (ack_tx, ack_rx) = oneshot::channel(); + tx.send(UploadCmd::Eos { + max_seq: 1, + ack: ack_tx, + }) + .await + .unwrap(); + drop(tx); + assert!(ack_rx.await.is_err()); + handle.await.unwrap(); + } + + #[tokio::test] + async fn stale_and_duplicate_frames_discarded() { + let (_rx, server_write) = tcp_pair().await; + let (tx, rx) = mpsc::channel::(16); + let actor = UploadActor::new(rx, server_write, 0); + let handle = tokio::spawn(async move { actor.run().await }); + + tx.send(UploadCmd::Frame { + seq: 0, + data: Bytes::from_static(b"a"), + }) + .await + .unwrap(); + tx.send(UploadCmd::Frame { + seq: 0, + data: Bytes::from_static(b"stale"), + }) + .await + .unwrap(); + tx.send(UploadCmd::Frame { + seq: 3, + data: Bytes::from_static(b"d"), + }) + .await + .unwrap(); + tx.send(UploadCmd::Frame { + seq: 3, + data: Bytes::from_static(b"dup"), + }) + .await + .unwrap(); + let (ack_tx, ack_rx) = oneshot::channel(); + tx.send(UploadCmd::Eos { + max_seq: 3, + ack: ack_tx, + }) + .await + .unwrap(); + tx.send(UploadCmd::Frame { + seq: 1, + data: Bytes::from_static(b"b"), + }) + .await + .unwrap(); + tx.send(UploadCmd::Frame { + seq: 2, + data: Bytes::from_static(b"c"), + }) + .await + .unwrap(); + assert!(ack_rx.await.unwrap().is_ok()); + drop(tx); + handle.await.unwrap(); + } + + #[tokio::test] + async fn reorder_timeout_shuts_down_actor() { + tokio::time::pause(); + let (_rx, server_write) = tcp_pair().await; + let (tx, rx) = mpsc::channel::(16); + let actor = UploadActor::new(rx, server_write, 0); + let handle = tokio::spawn(async move { actor.run().await }); + + tx.send(UploadCmd::Frame { + seq: 5, + data: Bytes::from_static(b"gap"), + }) + .await + .unwrap(); + tokio::task::yield_now().await; + tokio::task::yield_now().await; + tokio::time::advance(Duration::from_secs(MAX_REORDER_SECS + 1)).await; + tokio::task::yield_now().await; + tokio::time::resume(); + + let (ack_tx, _ack_rx) = oneshot::channel(); + let send_result = tx + .send(UploadCmd::Eos { + max_seq: 5, + ack: ack_tx, + }) + .await; + assert!(send_result.is_err()); + handle.await.unwrap(); + } } diff --git a/src/server/handlers.rs b/src/server/handlers.rs index 8ec604f..cc708cd 100644 --- a/src/server/handlers.rs +++ b/src/server/handlers.rs @@ -683,3 +683,144 @@ pub fn validate_jwt_if_needed( ServerError::unauthorized("invalid token") }) } + +#[cfg(test)] +mod tests { + use super::*; + use axum::body::Body; + use axum::extract::State; + use axum::http::HeaderMap; + use axum::http::StatusCode; + + async fn test_state() -> Arc { + let mut cfg = crate::config::ServerTopConfig { + server: crate::config::ServerSection { + listen: "127.0.0.1:0".to_string(), + path: "/secret".to_string(), + private_key: None, + max_tunnels: None, + }, + auth: crate::config::AuthSection { + secret: "test-secret".to_string(), + }, + proxy: None, + log: None, + dns: None, + traffic_shaping: crate::shaper::TrafficConfig { + global: crate::shaper::PaddingConfig { + padding_threshold: 0, + padding_range: [0, 0], + }, + stages: vec![], + encoding_type: crate::shaper::EncodingType::Binary, + max_download_bytes: None, + }, + }; + crate::server::build_state(&mut cfg).await.unwrap() + } + + fn valid_token() -> String { + jsonwebtoken::encode( + &jsonwebtoken::Header::default(), + &crate::server::Claims { + sub: "user".to_string(), + exp: 4_102_444_800, + }, + &jsonwebtoken::EncodingKey::from_secret(b"test-secret"), + ) + .unwrap() + } + + #[tokio::test] + async fn dispatch_rejects_without_cookies_or_target() { + let state = test_state().await; + let err = dispatch(State(state), HeaderMap::new(), Body::empty()) + .await + .unwrap_err(); + assert_eq!(err.0, StatusCode::BAD_REQUEST); + } + + #[tokio::test] + async fn dispatch_rejects_malformed_stream_cookie() { + let state = test_state().await; + let mut headers = HeaderMap::new(); + headers.insert("cookie", "stream=not-a-uuid".parse().unwrap()); + let err = dispatch(State(state), headers, Body::empty()) + .await + .unwrap_err(); + assert_eq!(err.0, StatusCode::PRECONDITION_REQUIRED); + } + + #[tokio::test] + async fn dispatch_rejects_unknown_stream() { + let state = test_state().await; + let mut headers = HeaderMap::new(); + headers.insert( + "cookie", + format!("stream={}", Uuid::new_v4()).parse().unwrap(), + ); + let err = dispatch(State(state), headers, Body::empty()) + .await + .unwrap_err(); + assert_eq!(err.0, StatusCode::PRECONDITION_REQUIRED); + } + + #[tokio::test] + async fn dispatch_rejects_bad_jwt_with_target() { + let state = test_state().await; + let mut headers = HeaderMap::new(); + headers.insert("X-Target", "127.0.0.1:1".parse().unwrap()); + headers.insert("Authorization", "Bearer invalid-token".parse().unwrap()); + let err = dispatch(State(state), headers, Body::empty()) + .await + .unwrap_err(); + assert_eq!(err.0, StatusCode::UNAUTHORIZED); + } + + #[tokio::test] + async fn dispatch_bad_gateway_on_unreachable_target_with_valid_jwt() { + let state = test_state().await; + let mut headers = HeaderMap::new(); + headers.insert("X-Target", "127.0.0.1:1".parse().unwrap()); + headers.insert( + "Authorization", + format!("Bearer {}", valid_token()).parse().unwrap(), + ); + let err = dispatch(State(state), headers, Body::empty()) + .await + .unwrap_err(); + assert_eq!(err.0, StatusCode::BAD_GATEWAY); + } + + #[tokio::test] + async fn dispatch_rejects_malformed_session_cookie() { + let state = test_state().await; + let mut headers = HeaderMap::new(); + headers.insert("cookie", "session=abc".parse().unwrap()); + let err = dispatch(State(state), headers, Body::empty()) + .await + .unwrap_err(); + assert_eq!(err.0, StatusCode::PRECONDITION_REQUIRED); + } + + #[tokio::test] + async fn dispatch_routes_expired_jwt_to_handshake_cookie_path() { + let state = test_state().await; + let token = jsonwebtoken::encode( + &jsonwebtoken::Header::default(), + &crate::server::Claims { + sub: "user".to_string(), + exp: 1, + }, + &jsonwebtoken::EncodingKey::from_secret(b"test-secret"), + ) + .unwrap(); + let mut headers = HeaderMap::new(); + headers.insert("X-Target", "127.0.0.1:1".parse().unwrap()); + headers.insert("Authorization", format!("Bearer {token}").parse().unwrap()); + let err = dispatch(State(state), headers, Body::empty()) + .await + .unwrap_err(); + assert_eq!(err.0, StatusCode::UNAUTHORIZED); + } +} diff --git a/src/server/janitor.rs b/src/server/janitor.rs index 46beee5..29c844e 100644 --- a/src/server/janitor.rs +++ b/src/server/janitor.rs @@ -19,21 +19,9 @@ pub async fn master_and_stream_janitor( interval.tick().await; let now = now_secs(); - let expiry_limit = MASTER_EXPIRY.as_secs(); - - master_store.retain(|session_id, (_, _, created)| { - if now.saturating_sub(*created) >= expiry_limit { - tracing::info!( - session_id = %session_id, - "master key expired, removing from store" - ); - false - } else { - true - } - }); - - let cutoff = now.saturating_sub(expiry_limit); + prune_expired_masters(&master_store, now); + + let cutoff = now.saturating_sub(MASTER_EXPIRY.as_secs()); let pruned = stream_registry.remove_consumed_before(cutoff); if pruned > 0 { tracing::debug!(pruned, "pruned consumed stream registry entries"); @@ -42,6 +30,44 @@ pub async fn master_and_stream_janitor( } } +pub fn prune_expired_masters( + master_store: &DashMap, + now: u64, +) -> usize { + let expiry_limit = MASTER_EXPIRY.as_secs(); + let mut pruned = 0; + master_store.retain(|session_id, (_, _, created)| { + if now.saturating_sub(*created) >= expiry_limit { + pruned += 1; + tracing::info!( + session_id = %session_id, + "master key expired, removing from store" + ); + false + } else { + true + } + }); + pruned +} + +pub fn prune_dead_actors(actors: &DashMap) -> usize { + let mut pruned = 0; + actors.retain(|stream_id, handle| { + if handle.cmd_tx.is_closed() { + pruned += 1; + tracing::info!( + stream_id = %stream_id, + "stream actor channel closed, removing from actor map" + ); + false + } else { + true + } + }); + pruned +} + pub async fn stream_janitor(actors: Arc>) { let mut interval = tokio::time::interval(JANITOR_INTERVAL); interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip); @@ -49,18 +75,82 @@ pub async fn stream_janitor(actors: Arc>) { loop { interval.tick().await; - actors.retain(|stream_id, handle| { - if handle.cmd_tx.is_closed() { - tracing::info!( - stream_id = %stream_id, - "stream actor channel closed, removing from actor map" - ); - false - } else { - true - } - }); + prune_dead_actors(&actors); actors.shrink_to_fit(); } } + +#[cfg(test)] +mod tests { + use super::*; + use crate::server::actor::tunnel::TunnelCmd; + use crate::shaper::EncodingType; + + fn handle_pair() -> (SessionHandle, tokio::sync::mpsc::Receiver) { + let (tx, rx) = tokio::sync::mpsc::channel::(4); + ( + SessionHandle { + cmd_tx: tx, + upload_cipher: None, + encoding: EncodingType::Binary, + }, + rx, + ) + } + + #[test] + fn prune_expired_masters_removes_old_keeps_fresh() { + let store = Arc::new(DashMap::new()); + let now = 10_000u64; + store.insert( + "old".to_string(), + ( + Arc::::from("u"), + zeroize::Zeroizing::new([0u8; 32]), + now - 2000, + ), + ); + store.insert( + "fresh".to_string(), + ( + Arc::::from("u"), + zeroize::Zeroizing::new([0u8; 32]), + now, + ), + ); + let pruned = prune_expired_masters(&store, now); + assert_eq!(pruned, 1); + assert!(!store.contains_key("old")); + assert!(store.contains_key("fresh")); + } + + #[test] + fn prune_expired_masters_boundary() { + let store = Arc::new(DashMap::new()); + let now = 10_000u64; + store.insert( + "boundary".to_string(), + ( + Arc::::from("u"), + zeroize::Zeroizing::new([0u8; 32]), + now - MASTER_EXPIRY.as_secs(), + ), + ); + let pruned = prune_expired_masters(&store, now); + assert_eq!(pruned, 1); + } + + #[test] + fn prune_dead_actors_removes_closed_channels() { + let actors = Arc::new(DashMap::new()); + let (handle1, rx1) = handle_pair(); + drop(rx1); + actors.insert(Uuid::new_v4(), handle1); + let (handle2, _rx2) = handle_pair(); + actors.insert(Uuid::new_v4(), handle2); + let pruned = prune_dead_actors(&actors); + assert_eq!(pruned, 1); + assert_eq!(actors.len(), 1); + } +} diff --git a/src/server/stream.rs b/src/server/stream.rs index 386cda9..b1a16ca 100644 --- a/src/server/stream.rs +++ b/src/server/stream.rs @@ -256,4 +256,44 @@ mod tests { assert_eq!(seq, 0); assert_eq!(&data[..], b"encrypted hello"); } + + #[tokio::test] + async fn json_encoded_frames_decoded() { + let data = b"json payload data"; + let frame = + shaper::encode_frame(data, 7, None, 16384, [0, 0], shaper::EncodingType::Json).unwrap(); + let byte_stream = stream::iter(vec![Ok(Bytes::from(frame))]); + let mut decoder = FrameDecoder::new(byte_stream, None, shaper::EncodingType::Json, 18_781); + let (seq, decoded) = decoder.next().await.unwrap().unwrap(); + assert_eq!(seq, 7); + assert_eq!(&decoded[..], data); + assert!(decoder.next().await.is_none()); + } + + #[tokio::test] + async fn trailing_partial_frame_after_eos_errors() { + let frame = make_frame(b"complete", 0); + let mut full = frame.to_vec(); + full.extend_from_slice(&[0x01, 0x02]); + let byte_stream = stream::iter(vec![Ok(Bytes::from(full))]); + let mut decoder = + FrameDecoder::new(byte_stream, None, shaper::EncodingType::Binary, 18_781); + let (seq, data) = decoder.next().await.unwrap().unwrap(); + assert_eq!(seq, 0); + assert_eq!(&data[..], b"complete"); + let err = decoder.next().await.unwrap().unwrap_err(); + assert_eq!(err.kind(), io::ErrorKind::InvalidData); + } + + #[tokio::test] + async fn inner_stream_error_propagates() { + let byte_stream = stream::iter(vec![Err(io::Error::new( + io::ErrorKind::BrokenPipe, + "upstream read failed", + ))]); + let mut decoder = + FrameDecoder::new(byte_stream, None, shaper::EncodingType::Binary, 18_781); + let err = decoder.next().await.unwrap().unwrap_err(); + assert_eq!(err.kind(), io::ErrorKind::BrokenPipe); + } } diff --git a/src/server/stream_registry.rs b/src/server/stream_registry.rs index 3dff610..6c267cc 100644 --- a/src/server/stream_registry.rs +++ b/src/server/stream_registry.rs @@ -135,7 +135,7 @@ mod tests { reg.mark_consumed(s1); reg.register(s2, 200); let removed = reg.remove_consumed_before(150); - assert!(removed >= 1); + assert_eq!(removed, 1); assert_eq!(reg.check(s1), StreamQueryResult::Fresh); assert_eq!(reg.check(s2), StreamQueryResult::Active); } diff --git a/src/shaper/mod.rs b/src/shaper/mod.rs index fbc950f..2ec6568 100644 --- a/src/shaper/mod.rs +++ b/src/shaper/mod.rs @@ -1231,4 +1231,224 @@ mod tests { } assert_eq!(decoded, data); } + + #[test] + fn encode_frame_rejects_oversized_payload() { + let big = vec![0u8; MAX_RAW_PAYLOAD + 1]; + let r = encode_frame(&big, 0, None, 0, [0, 0], EncodingType::Binary); + assert!(r.is_err()); + let r = encode_frame(&big, 0, None, 0, [0, 0], EncodingType::Json); + assert!(r.is_err()); + } + + #[test] + fn encode_frame_json_padding_stays_within_line_limit() { + let raw = vec![0x42u8; 1000]; + let frame = encode_frame(&raw, 0, None, 100_000, [0, 100_000], EncodingType::Json).unwrap(); + let line_len = frame.iter().position(|&b| b == b'\n').unwrap(); + assert!(line_len <= MAX_JSON_LINE_LEN); + } + + #[test] + fn encode_frame_json_cipher_padding_stays_within_line_limit() { + let raw = vec![0x42u8; 1000]; + let key = zeroize::Zeroizing::new([0u8; 32]); + let cipher = crate::crypto::AesFrameCipher::new(&key); + let frame = encode_frame( + &raw, + 0, + Some(&cipher), + 100_000, + [0, 100_000], + EncodingType::Json, + ) + .unwrap(); + let line_len = frame.iter().position(|&b| b == b'\n').unwrap(); + assert!(line_len <= MAX_JSON_LINE_LEN); + } + + #[tokio::test] + async fn poll_seal_into_eof_returns_none() { + let reader = std::io::Cursor::new(Vec::::new()); + let cfg = ResolvedShaperConfig::resolve(&test_config()); + let mut shaper = Box::pin(TrafficShaper::with_seq(reader, &cfg, None, 0)); + let mut out = BytesMut::new(); + let result = std::future::poll_fn(|cx| shaper.as_mut().poll_seal_into(cx, &mut out)).await; + assert!(matches!(result, Ok(None))); + } + + #[tokio::test] + async fn poll_seal_into_respects_start_seq() { + let data = vec![0x55u8; 1000]; + let reader = std::io::Cursor::new(data); + let cfg = ResolvedShaperConfig::resolve(&test_config()); + let mut shaper = Box::pin(TrafficShaper::with_seq(reader, &cfg, None, 42)); + let mut out = BytesMut::new(); + let result = std::future::poll_fn(|cx| shaper.as_mut().poll_seal_into(cx, &mut out)).await; + assert!(matches!(result, Ok(Some(42)))); + let mut scratch = BytesMut::new(); + let mut json_scratch = Vec::new(); + let frame = decode_frame( + &mut out, + &mut scratch, + &mut json_scratch, + None, + EncodingType::Binary, + ) + .unwrap() + .expect("frame"); + match frame { + DecodedFrame::Owned { seq, .. } => assert_eq!(seq, 42), + DecodedFrame::InScratch { seq, .. } => assert_eq!(seq, 42), + } + } + + #[tokio::test] + async fn poll_seal_into_pending_then_data() { + use tokio::io::AsyncWriteExt; + let (mut writer, reader) = tokio::io::duplex(READ_HIGH_WATER); + let cfg = ResolvedShaperConfig::resolve(&test_config()); + let mut shaper = Box::pin(TrafficShaper::with_seq(reader, &cfg, None, 0)); + let mut out = BytesMut::new(); + + let waker = std::task::Waker::noop(); + let mut cx = std::task::Context::from_waker(waker); + assert!( + matches!( + shaper.as_mut().poll_seal_into(&mut cx, &mut out), + Poll::Pending + ), + "empty duplex must yield Pending" + ); + + writer.write_all(&[0x33u8; READ_HIGH_WATER]).await.unwrap(); + assert!( + matches!( + shaper.as_mut().poll_seal_into(&mut cx, &mut out), + Poll::Ready(Ok(Some(_))) + ), + "data arrival must produce a frame" + ); + assert!(!out.is_empty()); + } + + #[tokio::test] + async fn stages_progress_and_fall_back_to_global() { + let cfg = TrafficConfig { + global: PaddingConfig { + padding_threshold: 10_000, + padding_range: [0, 0], + }, + stages: vec![ + StageConfig { + count: Some(1), + count_range: None, + padding_threshold: 10_000, + padding_range: [100, 100], + }, + StageConfig { + count: None, + count_range: Some([2, 3]), + padding_threshold: 10_000, + padding_range: [200, 200], + }, + ], + encoding_type: EncodingType::Binary, + max_download_bytes: None, + }; + use tokio::io::AsyncWriteExt; + let resolved = ResolvedShaperConfig::resolve(&cfg); + let (mut writer, reader) = tokio::io::duplex(READ_HIGH_WATER); + let mut shaper = Box::pin(TrafficShaper::with_seq(reader, &resolved, None, 0)); + let mut out = BytesMut::new(); + let mut sizes = Vec::new(); + for _ in 0..4 { + writer.write_all(&[0x44u8; 1000]).await.unwrap(); + std::future::poll_fn(|cx| shaper.as_mut().poll_seal_into(cx, &mut out)) + .await + .unwrap(); + sizes.push(out.len()); + out.clear(); + } + assert_eq!(sizes[0], 2 + 10 + 1000 + 100); + assert_eq!(sizes[1], 2 + 10 + 1000 + 200); + assert_eq!(sizes[2], 2 + 10 + 1000 + 200); + assert_eq!(sizes[3], 2 + 10 + 1000); + } + + #[test] + fn decode_frame_json_rejects_malformed_lines() { + let mut scratch = BytesMut::new(); + let mut json_scratch = Vec::new(); + let mut src = BytesMut::new(); + + src.extend_from_slice( + b" +", + ); + assert!( + decode_frame( + &mut src, + &mut scratch, + &mut json_scratch, + None, + EncodingType::Json + ) + .is_err() + ); + + let mut long = BytesMut::new(); + long.extend_from_slice(b"{\"data\":\""); + long.resize(MAX_JSON_LINE_LEN + 2, b'x'); + assert!( + decode_frame( + &mut long, + &mut scratch, + &mut json_scratch, + None, + EncodingType::Json + ) + .is_err() + ); + + let mut unclosed = BytesMut::new(); + unclosed.extend_from_slice(b"{\"data\":\"abc"); + unclosed.extend_from_slice( + b" +", + ); + assert!( + decode_frame( + &mut unclosed, + &mut scratch, + &mut json_scratch, + None, + EncodingType::Json + ) + .is_err() + ); + } + + #[test] + fn decode_frame_json_rejects_invalid_utf8() { + let mut scratch = BytesMut::new(); + let mut json_scratch = Vec::new(); + let mut src = BytesMut::new(); + src.extend_from_slice(b"{\"data\":\""); + src.extend_from_slice(&[0xff, 0xfe, 0xfd]); + src.extend_from_slice( + b"\"} +", + ); + assert!( + decode_frame( + &mut src, + &mut scratch, + &mut json_scratch, + None, + EncodingType::Json + ) + .is_err() + ); + } } diff --git a/tests/bypass.rs b/tests/bypass.rs new file mode 100644 index 0000000..528abe5 --- /dev/null +++ b/tests/bypass.rs @@ -0,0 +1,132 @@ +mod common; + +use common::*; +use tokio::io::{AsyncReadExt, AsyncWriteExt}; +use tokio::net::TcpStream; + +fn write_bypass_file(rules: &str) -> String { + static COUNTER: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0); + let n = COUNTER.fetch_add(1, std::sync::atomic::Ordering::Relaxed); + let path = + std::env::temp_dir().join(format!("httproxy_bypass_{}_{}.json", std::process::id(), n)); + std::fs::write(&path, rules).unwrap(); + path.to_string_lossy().into_owned() +} + +#[tokio::test] +async fn bypass_direct_connect_skips_tunnel() { + common::init_logging(); + let upstream = spawn_upstream().await; + let server = spawn_server(server_config(traffic(true, None), None, None)).await; + let bypass_file = write_bypass_file(r#"{"domain_suffix": [], "ip_cidr": ["127.0.0.1/32"]}"#); + let client = spawn_client( + client_config(traffic(true, None), None, None, vec![bypass_file]), + server, + ) + .await; + + let url = format!("http://{upstream}/hello"); + let (status, body) = raw_exchange( + client, + format!("GET {url} HTTP/1.1\r\nHost: test\r\nConnection: close\r\n\r\n").as_bytes(), + ) + .await; + assert_eq!(status, 200); + assert_eq!(body, b"hello upstream"); +} + +#[tokio::test] +async fn bypass_large_body_direct() { + common::init_logging(); + let upstream = spawn_upstream().await; + let server = spawn_server(server_config(traffic(true, None), None, None)).await; + let bypass_file = write_bypass_file(r#"{"domain_suffix": [], "ip_cidr": ["127.0.0.1/32"]}"#); + let client = spawn_client( + client_config(traffic(true, None), None, None, vec![bypass_file]), + server, + ) + .await; + + let url = format!("http://{upstream}/large"); + let (status, body) = raw_exchange( + client, + format!("GET {url} HTTP/1.1\r\nHost: test\r\nConnection: close\r\n\r\n").as_bytes(), + ) + .await; + assert_eq!(status, 200); + assert_eq!(body.len(), 512 * 1024); + assert!(body.iter().all(|&b| b == 0xAB)); +} + +#[tokio::test] +async fn bypass_connect_direct() { + common::init_logging(); + let upstream = spawn_upstream().await; + let server = spawn_server(server_config(traffic(true, None), None, None)).await; + let bypass_file = write_bypass_file(r#"{"domain_suffix": [], "ip_cidr": ["127.0.0.1/32"]}"#); + let client = spawn_client( + client_config(traffic(true, None), None, None, vec![bypass_file]), + server, + ) + .await; + + let mut stream = proxy_connect(client, upstream).await; + stream + .write_all(b"GET /hello HTTP/1.1\r\nHost: test\r\nConnection: close\r\n\r\n") + .await + .unwrap(); + let mut buf = Vec::new(); + stream.read_to_end(&mut buf).await.unwrap(); + assert!(String::from_utf8_lossy(&buf).contains("hello upstream")); +} + +#[tokio::test] +async fn non_matching_domain_still_tunneled() { + common::init_logging(); + let upstream = spawn_upstream().await; + let server = spawn_server(server_config(traffic(true, None), None, None)).await; + let bypass_file = write_bypass_file(r#"{"domain_suffix": [], "ip_cidr": ["10.0.0.0/8"]}"#); + let client = spawn_client( + client_config(traffic(true, None), None, None, vec![bypass_file]), + server, + ) + .await; + + let url = format!("http://{upstream}/hello"); + let (status, body) = raw_exchange( + client, + format!("GET {url} HTTP/1.1\r\nHost: test\r\nConnection: close\r\n\r\n").as_bytes(), + ) + .await; + assert_eq!(status, 200); + assert_eq!(body, b"hello upstream"); +} + +#[tokio::test] +async fn bypass_unreachable_target_fails_fast() { + common::init_logging(); + let server = spawn_server(server_config(traffic(true, None), None, None)).await; + let bypass_file = write_bypass_file(r#"{"domain_suffix": [], "ip_cidr": ["127.0.0.1/32"]}"#); + let client = spawn_client( + client_config(traffic(true, None), None, None, vec![bypass_file]), + server, + ) + .await; + + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let dead_addr = listener.local_addr().unwrap(); + drop(listener); + + let url = format!("http://{dead_addr}/"); + let mut stream = TcpStream::connect(client).await.unwrap(); + stream + .write_all(format!("GET {url} HTTP/1.1\r\nHost: test\r\n\r\n").as_bytes()) + .await + .unwrap(); + let mut buf = [0u8; 1024]; + let n = stream.read(&mut buf).await.unwrap(); + assert_eq!( + n, 0, + "bypass connection failure should close the proxy connection" + ); +} diff --git a/tests/common/mod.rs b/tests/common/mod.rs new file mode 100644 index 0000000..116cddb --- /dev/null +++ b/tests/common/mod.rs @@ -0,0 +1,246 @@ +use httproxy::config::{ + AuthSection, BypassConfig, ClientAuthSection, ClientSection, ClientTopConfig, ServerSection, + ServerTopConfig, +}; +use httproxy::shaper::{EncodingType, PaddingConfig, TrafficConfig}; +use std::net::SocketAddr; +use std::sync::Arc; +use tokio::io::{AsyncReadExt, AsyncWriteExt}; +use tokio::net::TcpStream; + +pub const SECRET: &str = "integration_test_secret"; +pub const PATH: &str = "/integration_path"; + +pub fn init_logging() { + let _ = tracing_subscriber::fmt() + .with_env_filter(tracing_subscriber::EnvFilter::from_default_env()) + .with_test_writer() + .try_init(); +} + +pub fn traffic(binary: bool, max_download_bytes: Option) -> TrafficConfig { + TrafficConfig { + global: PaddingConfig { + padding_threshold: 0, + padding_range: [0, 0], + }, + stages: vec![], + encoding_type: if binary { + EncodingType::Binary + } else { + EncodingType::Json + }, + max_download_bytes, + } +} + +pub fn make_token(user: &str, exp: u64) -> String { + jsonwebtoken::encode( + &jsonwebtoken::Header::default(), + &httproxy::server::Claims { + sub: user.to_string(), + exp, + }, + &jsonwebtoken::EncodingKey::from_secret(SECRET.as_bytes()), + ) + .unwrap() +} + +pub fn valid_token() -> String { + make_token("test-user", 4_102_444_800) +} + +pub fn server_config( + traffic: TrafficConfig, + private_key: Option, + max_tunnels: Option, +) -> ServerTopConfig { + ServerTopConfig { + server: ServerSection { + listen: "127.0.0.1:0".to_string(), + path: PATH.to_string(), + private_key, + max_tunnels, + }, + auth: AuthSection { + secret: SECRET.to_string(), + }, + proxy: None, + log: None, + dns: None, + traffic_shaping: traffic, + } +} + +pub fn client_config( + traffic: TrafficConfig, + public_key: Option, + token: Option, + bypass_files: Vec, +) -> ClientTopConfig { + ClientTopConfig { + client: ClientSection { + listen: "127.0.0.1:0".to_string(), + remote: format!("http://proxy-host.invalid{PATH}"), + address: None, + public_key, + auth: None, + max_connections: None, + max_in_flight_bytes: None, + upload_concurrency: None, + }, + auth: ClientAuthSection { + token: token.unwrap_or_else(valid_token), + }, + log: None, + traffic_shaping: traffic, + bypass: BypassConfig { bypass_files }, + } +} + +pub async fn spawn_upstream() -> SocketAddr { + let app = axum::Router::new() + .route("/hello", axum::routing::get(|| async { "hello upstream" })) + .route( + "/large", + axum::routing::get(|| async { vec![0xABu8; 512 * 1024] }), + ) + .route( + "/big", + axum::routing::get(|| async { vec![0xCDu8; 32 * 1024 * 1024] }), + ) + .route( + "/echo", + axum::routing::post(|body: axum::body::Bytes| async move { body }), + ) + .layer(axum::extract::DefaultBodyLimit::max(64 * 1024 * 1024)); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + tokio::spawn(async move { + let _ = axum::serve(listener, app.into_make_service()).await; + }); + addr +} + +pub async fn spawn_server(cfg: ServerTopConfig) -> SocketAddr { + let mut cfg = cfg; + let state = httproxy::server::build_state(&mut cfg).await.unwrap(); + let path = cfg.server.path.clone(); + let router = httproxy::server::build_router(state, &path); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + tokio::spawn(async move { + let _ = axum::serve(listener, router.into_make_service()).await; + }); + addr +} + +pub async fn spawn_client(mut cfg: ClientTopConfig, server_addr: SocketAddr) -> SocketAddr { + cfg.client.remote = format!("http://127.0.0.1:{}{}", server_addr.port(), PATH); + let state = httproxy::client::build_state(&cfg).unwrap(); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let proxy_addr = listener.local_addr().unwrap(); + let http_client = Arc::new( + wreq::Client::builder() + .tcp_nodelay(true) + .emulation(wreq_util::Emulation::Chrome143) + .no_proxy() + .dns_resolver(Arc::new(httproxy::client::state::ManualResolver { + target_addr: server_addr.ip().to_string(), + })) + .build() + .unwrap(), + ); + let sem = Arc::new(tokio::sync::Semaphore::new(state.max_connections)); + tokio::spawn(async move { + loop { + let Ok((socket, _)) = listener.accept().await else { + break; + }; + let sem = sem.clone(); + let http_client = http_client.clone(); + let state = state.clone(); + tokio::spawn(async move { + let _permit = sem.acquire_owned().await; + if let Err(e) = httproxy::client::connection::handle_connection_actor( + socket, + http_client, + state, + ) + .await + { + eprintln!("client connection error: {e:?}"); + } + }); + } + }); + proxy_addr +} + +pub async fn raw_exchange(proxy: SocketAddr, request: &[u8]) -> (u16, Vec) { + let mut stream = TcpStream::connect(proxy).await.unwrap(); + stream.write_all(request).await.unwrap(); + read_response(&mut stream).await +} + +pub async fn read_response(stream: &mut TcpStream) -> (u16, Vec) { + let mut buf = Vec::new(); + let mut tmp = [0u8; 8192]; + let header_end = loop { + let n = stream.read(&mut tmp).await.unwrap(); + if n == 0 { + eprintln!("DBG read_response EOF after {} bytes", buf.len()); + } + assert!(n > 0, "connection closed before headers"); + buf.extend_from_slice(&tmp[..n]); + if let Some(pos) = find_subslice(&buf, b"\r\n\r\n") { + break pos + 4; + } + assert!(buf.len() < 64 * 1024, "headers too large"); + }; + let headers = String::from_utf8_lossy(&buf[..header_end]); + let status: u16 = headers.split_whitespace().nth(1).unwrap().parse().unwrap(); + let content_length = headers + .lines() + .find(|l| l.to_ascii_lowercase().starts_with("content-length:")) + .and_then(|l| l.split(':').nth(1)) + .and_then(|v| v.trim().parse::().ok()); + match content_length { + Some(len) => { + while buf.len() - header_end < len { + let n = stream.read(&mut tmp).await.unwrap(); + if n == 0 { + break; + } + buf.extend_from_slice(&tmp[..n]); + } + (status, buf[header_end..header_end + len].to_vec()) + } + None => { + loop { + let n = stream.read(&mut tmp).await.unwrap(); + if n == 0 { + break; + } + buf.extend_from_slice(&tmp[..n]); + } + (status, buf[header_end..].to_vec()) + } + } +} + +#[allow(dead_code)] +pub async fn proxy_connect(proxy: SocketAddr, target: SocketAddr) -> TcpStream { + let mut stream = TcpStream::connect(proxy).await.unwrap(); + let req = format!("CONNECT {} HTTP/1.1\r\nHost: {}\r\n\r\n", target, target); + stream.write_all(req.as_bytes()).await.unwrap(); + let mut buf = [0u8; 4096]; + let n = stream.read(&mut buf).await.unwrap(); + let head = String::from_utf8_lossy(&buf[..n]); + assert!(head.starts_with("HTTP/1.1 200"), "CONNECT failed: {head}"); + stream +} + +fn find_subslice(haystack: &[u8], needle: &[u8]) -> Option { + haystack.windows(needle.len()).position(|w| w == needle) +} diff --git a/tests/dns.rs b/tests/dns.rs new file mode 100644 index 0000000..9c6bf44 --- /dev/null +++ b/tests/dns.rs @@ -0,0 +1,184 @@ +mod common; + +use common::*; +use httproxy::dns::{DnsConfig, DnsOptions}; +use std::collections::HashMap; +use std::net::IpAddr; +use std::sync::Arc; +use std::sync::atomic::{AtomicUsize, Ordering}; +use tokio::io::{AsyncReadExt, AsyncWriteExt}; +use tokio::net::UdpSocket; + +fn extract_domain(query: &[u8], mut offset: usize) -> Option { + let mut labels = Vec::new(); + loop { + let len = *query.get(offset)? as usize; + offset += 1; + if len == 0 { + break; + } + let label = std::str::from_utf8(query.get(offset..offset + len)?).ok()?; + labels.push(label.to_string()); + offset += len; + } + Some(labels.join(".")) +} + +async fn spawn_mock_dns( + records: HashMap, + query_count: Arc, +) -> std::net::SocketAddr { + use domain::base::iana::{Class, Rcode}; + use domain::base::{MessageBuilder, Name, Record, Ttl}; + use domain::rdata::{A, Aaaa}; + use std::str::FromStr; + + let sock = Arc::new(UdpSocket::bind("127.0.0.1:0").await.unwrap()); + let addr = sock.local_addr().unwrap(); + tokio::spawn(async move { + let mut buf = [0u8; 512]; + loop { + let Ok((n, peer)) = sock.recv_from(&mut buf).await else { + break; + }; + let query = &buf[..n]; + if query.len() < 12 { + continue; + } + query_count.fetch_add(1, Ordering::Relaxed); + let id = u16::from_be_bytes([query[0], query[1]]); + let domain = extract_domain(query, 12).unwrap_or_default(); + let mut builder = MessageBuilder::new_vec(); + builder.header_mut().set_id(id); + if records.contains_key(&domain) { + builder.header_mut().set_rcode(Rcode::NOERROR); + } else { + builder.header_mut().set_rcode(Rcode::NXDOMAIN); + } + let name = Name::>::from_str(&domain).unwrap_or_else(|_| Name::root()); + let mut answer = builder.answer(); + if let Some(ip) = records.get(&domain) { + let _ = match ip { + IpAddr::V4(v4) => answer.push(Record::new( + name, + Class::IN, + Ttl::from_secs(60), + A::new(*v4), + )), + IpAddr::V6(v6) => answer.push(Record::new( + name, + Class::IN, + Ttl::from_secs(60), + Aaaa::new(*v6), + )), + }; + } + let resp = answer.into_message().into_octets(); + let _ = sock.send_to(&resp, peer).await; + } + }); + addr +} + +fn server_with_dns(mock_dns: std::net::SocketAddr) -> httproxy::config::ServerTopConfig { + let mut cfg = server_config(traffic(true, None), None, None); + cfg.dns = Some(DnsConfig { + upstream: mock_dns, + tls_domain: None, + options: DnsOptions::default(), + }); + cfg +} + +#[tokio::test] +async fn server_resolves_domain_via_dns_module() { + common::init_logging(); + let upstream = spawn_upstream().await; + let query_count = Arc::new(AtomicUsize::new(0)); + let mut records = HashMap::new(); + records.insert("dns-test.invalid".to_string(), upstream.ip()); + let mock_dns = spawn_mock_dns(records, query_count.clone()).await; + let dc = std::sync::Arc::new( + httproxy::dns::DnsClient::new(&DnsConfig { + upstream: mock_dns, + tls_domain: None, + options: DnsOptions::default(), + }) + .await + .unwrap(), + ); + let ips = dc + .lookup("dns-test.invalid", domain::base::iana::Rtype::A, None) + .await; + assert!(!ips.unwrap().is_empty(), "direct lookup must resolve"); + let server = spawn_server(server_with_dns(mock_dns)).await; + let client = spawn_client( + client_config(traffic(true, None), None, None, vec![]), + server, + ) + .await; + + let url = format!("http://dns-test.invalid:{}/hello", upstream.port()); + let (status, body) = raw_exchange( + client, + format!("GET {url} HTTP/1.1\r\nHost: test\r\nConnection: close\r\n\r\n").as_bytes(), + ) + .await; + assert_eq!(status, 200); + assert_eq!(body, b"hello upstream"); + assert!( + query_count.load(Ordering::Relaxed) >= 1, + "dns module must have queried the mock server" + ); +} + +#[tokio::test] +async fn dns_results_are_cached() { + common::init_logging(); + let upstream = spawn_upstream().await; + let query_count = Arc::new(AtomicUsize::new(0)); + let mut records = HashMap::new(); + records.insert("cached.invalid".to_string(), upstream.ip()); + let mock_dns = spawn_mock_dns(records, query_count.clone()).await; + let server = spawn_server(server_with_dns(mock_dns)).await; + let client = spawn_client( + client_config(traffic(true, None), None, None, vec![]), + server, + ) + .await; + + let url = format!("http://cached.invalid:{}/hello", upstream.port()); + let request = format!("GET {url} HTTP/1.1\r\nHost: test\r\nConnection: close\r\n\r\n"); + for _ in 0..3 { + let (status, body) = raw_exchange(client, request.as_bytes()).await; + assert_eq!(status, 200); + assert_eq!(body, b"hello upstream"); + } + let queries = query_count.load(Ordering::Relaxed); + assert!( + queries <= 2, + "cache should serve repeat lookups, got {queries} queries" + ); +} + +#[tokio::test] +async fn nxdomain_fails_connection() { + common::init_logging(); + let mock_dns = spawn_mock_dns(HashMap::new(), Arc::new(AtomicUsize::new(0))).await; + let server = spawn_server(server_with_dns(mock_dns)).await; + let client = spawn_client( + client_config(traffic(true, None), None, None, vec![]), + server, + ) + .await; + + let url = format!("http://missing.invalid:{}/hello", 12345); + let mut stream = tokio::net::TcpStream::connect(client).await.unwrap(); + stream + .write_all(format!("GET {url} HTTP/1.1\r\nHost: test\r\n\r\n").as_bytes()) + .await + .unwrap(); + let mut buf = [0u8; 1024]; + let n = stream.read(&mut buf).await.unwrap(); + assert_eq!(n, 0, "nxdomain should close the proxy connection"); +} diff --git a/tests/encrypted_tunnel.rs b/tests/encrypted_tunnel.rs new file mode 100644 index 0000000..1876408 --- /dev/null +++ b/tests/encrypted_tunnel.rs @@ -0,0 +1,105 @@ +mod common; + +use common::*; +use std::net::SocketAddr; +use tokio::io::{AsyncReadExt, AsyncWriteExt}; + +async fn start_encrypted() -> (SocketAddr, SocketAddr, String) { + let upstream = spawn_upstream().await; + let (sk, pk) = httproxy::crypto::generate_keypair(); + let sk_b64 = httproxy::crypto::private_key_to_b64(&sk); + let pk_b64 = httproxy::crypto::public_key_to_b64(&pk); + let server = spawn_server(server_config(traffic(true, None), Some(sk_b64), None)).await; + let client = spawn_client( + client_config(traffic(true, None), Some(pk_b64.clone()), None, vec![]), + server, + ) + .await; + (upstream, client, pk_b64) +} + +#[tokio::test] +async fn encrypted_get_small_response() { + common::init_logging(); + let (upstream, client, _) = start_encrypted().await; + let url = format!("http://{upstream}/hello"); + let (status, body) = raw_exchange( + client, + format!("GET {url} HTTP/1.1\r\nHost: test\r\nConnection: close\r\n\r\n").as_bytes(), + ) + .await; + assert_eq!(status, 200); + assert_eq!(body, b"hello upstream"); +} + +#[tokio::test] +async fn encrypted_get_large_response() { + common::init_logging(); + let (upstream, client, _) = start_encrypted().await; + let url = format!("http://{upstream}/large"); + let (status, body) = raw_exchange( + client, + format!("GET {url} HTTP/1.1\r\nHost: test\r\nConnection: close\r\n\r\n").as_bytes(), + ) + .await; + assert_eq!(status, 200); + assert_eq!(body.len(), 512 * 1024); + assert!(body.iter().all(|&b| b == 0xAB)); +} + +#[tokio::test] +async fn encrypted_large_post_echoes() { + common::init_logging(); + let (upstream, client, _) = start_encrypted().await; + let body: Vec = (0..(2 * 1024 * 1024) as u32) + .map(|i| (i % 251) as u8) + .collect(); + let url = format!("http://{upstream}/echo"); + let request = format!( + "POST {url} HTTP/1.1\r\nHost: test\r\nContent-Length: {}\r\nConnection: close\r\n\r\n", + body.len() + ); + let mut stream = tokio::net::TcpStream::connect(client).await.unwrap(); + stream.write_all(request.as_bytes()).await.unwrap(); + stream.write_all(&body).await.unwrap(); + let (status, echoed) = read_response(&mut stream).await; + assert_eq!(status, 200); + assert_eq!(echoed, body); +} + +#[tokio::test] +async fn encrypted_connect_tunnel() { + common::init_logging(); + let (upstream, client, _) = start_encrypted().await; + let mut stream = proxy_connect(client, upstream).await; + stream + .write_all(b"GET /hello HTTP/1.1\r\nHost: test\r\nConnection: close\r\n\r\n") + .await + .unwrap(); + let mut buf = Vec::new(); + stream.read_to_end(&mut buf).await.unwrap(); + assert!(String::from_utf8_lossy(&buf).contains("hello upstream")); +} + +#[tokio::test] +async fn pq_session_resumption_reuses_ticket() { + common::init_logging(); + let upstream = spawn_upstream().await; + let (sk, pk) = httproxy::crypto::generate_keypair(); + let sk_b64 = httproxy::crypto::private_key_to_b64(&sk); + let pk_b64 = httproxy::crypto::public_key_to_b64(&pk); + let server = spawn_server(server_config(traffic(true, None), Some(sk_b64), None)).await; + let client = spawn_client( + client_config(traffic(true, None), Some(pk_b64), None, vec![]), + server, + ) + .await; + + let url = format!("http://{upstream}/hello"); + let request = format!("GET {url} HTTP/1.1\r\nHost: test\r\nConnection: close\r\n\r\n"); + for _ in 0..2 { + let (status, body) = raw_exchange(client, request.as_bytes()).await; + assert_eq!(status, 200); + assert_eq!(body, b"hello upstream"); + } +} diff --git a/tests/plain_tunnel.rs b/tests/plain_tunnel.rs new file mode 100644 index 0000000..00812a8 --- /dev/null +++ b/tests/plain_tunnel.rs @@ -0,0 +1,241 @@ +mod common; + +use base64::Engine; +use common::*; +use std::net::SocketAddr; +use tokio::io::{AsyncReadExt, AsyncWriteExt}; +use tokio::net::TcpStream; + +async fn start_plain() -> (SocketAddr, SocketAddr) { + let upstream = spawn_upstream().await; + let server = spawn_server(server_config(traffic(true, None), None, None)).await; + let client = spawn_client( + client_config(traffic(true, None), None, None, vec![]), + server, + ) + .await; + (upstream, client) +} + +#[tokio::test] +async fn get_small_response_through_tunnel() { + let (upstream, client) = start_plain().await; + let url = format!("http://{upstream}/hello"); + let (status, body) = raw_exchange( + client, + format!("GET {url} HTTP/1.1\r\nHost: test\r\nConnection: close\r\n\r\n").as_bytes(), + ) + .await; + assert_eq!(status, 200); + assert_eq!(body, b"hello upstream"); +} + +#[tokio::test] +async fn get_large_response_through_tunnel() { + let (upstream, client) = start_plain().await; + let url = format!("http://{upstream}/large"); + let (status, body) = raw_exchange( + client, + format!("GET {url} HTTP/1.1\r\nHost: test\r\nConnection: close\r\n\r\n").as_bytes(), + ) + .await; + assert_eq!(status, 200); + assert_eq!(body.len(), 512 * 1024); + assert!(body.iter().all(|&b| b == 0xAB)); +} + +#[tokio::test] +async fn connect_tunnel_relays_bytes() { + let (upstream, client) = start_plain().await; + let mut stream = proxy_connect(client, upstream).await; + stream + .write_all(b"GET /hello HTTP/1.1\r\nHost: test\r\nConnection: close\r\n\r\n") + .await + .unwrap(); + let mut buf = Vec::new(); + stream.read_to_end(&mut buf).await.unwrap(); + let head = String::from_utf8_lossy(&buf); + assert!(head.contains("200 OK") || head.starts_with("HTTP/1.1 200")); + assert!(buf.ends_with(b"hello upstream")); +} + +#[tokio::test] +async fn large_post_echoes_through_tunnel() { + let (upstream, client) = start_plain().await; + let body: Vec = (0..(3 * 1024 * 1024) as u32) + .map(|i| (i % 251) as u8) + .collect(); + let url = format!("http://{upstream}/echo"); + let request = format!( + "POST {url} HTTP/1.1\r\nHost: test\r\nContent-Length: {}\r\nConnection: close\r\n\r\n", + body.len() + ); + let mut stream = TcpStream::connect(client).await.unwrap(); + stream.write_all(request.as_bytes()).await.unwrap(); + stream.write_all(&body).await.unwrap(); + let (status, echoed) = read_response(&mut stream).await; + assert_eq!(status, 200); + assert_eq!(echoed, body); +} + +#[tokio::test] +async fn bad_token_rejected() { + common::init_logging(); + let upstream = spawn_upstream().await; + let server = spawn_server(server_config(traffic(true, None), None, None)).await; + let client = spawn_client( + client_config( + traffic(true, None), + None, + Some("bad-token".to_string()), + vec![], + ), + server, + ) + .await; + let url = format!("http://{upstream}/hello"); + let mut stream = TcpStream::connect(client).await.unwrap(); + stream + .write_all(format!("GET {url} HTTP/1.1\r\nHost: test\r\n\r\n").as_bytes()) + .await + .unwrap(); + let mut buf = [0u8; 1024]; + let n = stream.read(&mut buf).await.unwrap(); + assert_eq!(n, 0, "rejected connection should close without a response"); +} + +#[tokio::test] +async fn expired_token_rejected() { + common::init_logging(); + let upstream = spawn_upstream().await; + let server = spawn_server(server_config(traffic(true, None), None, None)).await; + let client = spawn_client( + client_config(traffic(true, None), None, Some(make_token("u", 1)), vec![]), + server, + ) + .await; + let url = format!("http://{upstream}/hello"); + let mut stream = TcpStream::connect(client).await.unwrap(); + stream + .write_all(format!("GET {url} HTTP/1.1\r\nHost: test\r\n\r\n").as_bytes()) + .await + .unwrap(); + let mut buf = [0u8; 1024]; + let n = stream.read(&mut buf).await.unwrap(); + assert_eq!(n, 0, "rejected connection should close without a response"); +} + +#[tokio::test] +async fn unreachable_upstream_yields_connection_error() { + common::init_logging(); + let server = spawn_server(server_config(traffic(true, None), None, None)).await; + let client = spawn_client( + client_config(traffic(true, None), None, None, vec![]), + server, + ) + .await; + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let dead_addr = listener.local_addr().unwrap(); + drop(listener); + let url = format!("http://{dead_addr}/"); + let mut stream = TcpStream::connect(client).await.unwrap(); + stream + .write_all(format!("GET {url} HTTP/1.1\r\nHost: test\r\n\r\n").as_bytes()) + .await + .unwrap(); + let mut buf = [0u8; 1024]; + let n = stream.read(&mut buf).await.unwrap(); + assert_eq!(n, 0, "connection should close without a proxy response"); +} + +#[tokio::test] +async fn concurrent_tunnels_work_independently() { + let (upstream, client) = start_plain().await; + let url = format!("http://{upstream}/hello"); + let request = format!("GET {url} HTTP/1.1\r\nHost: test\r\nConnection: close\r\n\r\n"); + let mut handles = Vec::new(); + for _ in 0..8 { + let request = request.clone(); + handles.push(tokio::spawn(async move { + raw_exchange(client, request.as_bytes()).await + })); + } + for handle in handles { + let (status, body) = handle.await.unwrap(); + assert_eq!(status, 200); + assert_eq!(body, b"hello upstream"); + } +} + +#[tokio::test] +async fn tunnel_admission_limit_rejects_excess() { + common::init_logging(); + let upstream = spawn_upstream().await; + let server = spawn_server(server_config(traffic(true, None), None, Some(1))).await; + let client = spawn_client( + client_config(traffic(true, None), None, None, vec![]), + server, + ) + .await; + + let big_url = format!("http://{upstream}/big"); + let small_url = format!("http://{upstream}/hello"); + let big_req = format!("GET {big_url} HTTP/1.1\r\nHost: test\r\nConnection: close\r\n\r\n"); + let small_req = format!("GET {small_url} HTTP/1.1\r\nHost: test\r\nConnection: close\r\n\r\n"); + + let big = tokio::spawn(async move { raw_exchange(client, big_req.as_bytes()).await }); + tokio::time::sleep(std::time::Duration::from_millis(200)).await; + let mut stream = TcpStream::connect(client).await.unwrap(); + stream.write_all(small_req.as_bytes()).await.unwrap(); + let mut buf = [0u8; 1024]; + let n = stream.read(&mut buf).await.unwrap(); + assert_eq!( + n, 0, + "excess tunnel should be rejected (503 -> connection closed)" + ); + + let (status, body) = big.await.unwrap(); + assert_eq!(status, 200); + assert_eq!(body.len(), 32 * 1024 * 1024); +} + +#[tokio::test] +async fn local_proxy_auth_challenges_and_accepts() { + common::init_logging(); + let upstream = spawn_upstream().await; + let server = spawn_server(server_config(traffic(true, None), None, None)).await; + let mut cfg = client_config(traffic(true, None), None, None, vec![]); + cfg.client.auth = Some(httproxy::config::ClientProxyAuth { + username: "proxyuser".to_string(), + password: "proxypass".to_string(), + }); + let client = spawn_client(cfg, server).await; + + let url = format!("http://{upstream}/hello"); + + let mut stream = TcpStream::connect(client).await.unwrap(); + stream + .write_all( + format!("GET {url} HTTP/1.1\r\nHost: test\r\nConnection: close\r\n\r\n").as_bytes(), + ) + .await + .unwrap(); + let (status, _) = read_response(&mut stream).await; + assert_eq!(status, 407); + + let mut stream = TcpStream::connect(client).await.unwrap(); + stream + .write_all( + format!( + "GET {url} HTTP/1.1\r\nHost: test\r\nProxy-Authorization: Basic {}\r\nConnection: close\r\n\r\n", + base64::engine::general_purpose::STANDARD + .encode(b"proxyuser:proxypass") + ) + .as_bytes(), + ) + .await + .unwrap(); + let (status, body) = read_response(&mut stream).await; + assert_eq!(status, 200); + assert_eq!(body, b"hello upstream"); +} diff --git a/tests/rotation.rs b/tests/rotation.rs new file mode 100644 index 0000000..5953ea1 --- /dev/null +++ b/tests/rotation.rs @@ -0,0 +1,89 @@ +mod common; + +use common::*; +use std::net::SocketAddr; +use tokio::io::AsyncWriteExt; + +async fn start_rotating(binary: bool, max_download_bytes: u64) -> (SocketAddr, SocketAddr) { + let upstream = spawn_upstream().await; + let server = spawn_server(server_config( + traffic(binary, Some(max_download_bytes)), + None, + None, + )) + .await; + let client = spawn_client( + client_config( + traffic(binary, Some(max_download_bytes)), + None, + None, + vec![], + ), + server, + ) + .await; + (upstream, client) +} + +#[tokio::test] +async fn rotating_download_reassembles_binary() { + common::init_logging(); + let (upstream, client) = start_rotating(true, 64 * 1024).await; + let url = format!("http://{upstream}/large"); + let (status, body) = raw_exchange( + client, + format!("GET {url} HTTP/1.1\r\nHost: test\r\nConnection: close\r\n\r\n").as_bytes(), + ) + .await; + assert_eq!(status, 200); + assert_eq!(body.len(), 512 * 1024); + assert!(body.iter().all(|&b| b == 0xAB)); +} + +#[tokio::test] +async fn rotating_download_reassembles_json() { + common::init_logging(); + let (upstream, client) = start_rotating(false, 64 * 1024).await; + let url = format!("http://{upstream}/large"); + let (status, body) = raw_exchange( + client, + format!("GET {url} HTTP/1.1\r\nHost: test\r\nConnection: close\r\n\r\n").as_bytes(), + ) + .await; + assert_eq!(status, 200); + assert_eq!(body.len(), 512 * 1024); + assert!(body.iter().all(|&b| b == 0xAB)); +} + +#[tokio::test] +async fn rotating_with_upload_echo() { + common::init_logging(); + let (upstream, client) = start_rotating(true, 256 * 1024).await; + let body: Vec = (0..(512 * 1024) as u32).map(|i| (i % 251) as u8).collect(); + let url = format!("http://{upstream}/echo"); + let request = format!( + "POST {url} HTTP/1.1\r\nHost: test\r\nContent-Length: {}\r\nConnection: close\r\n\r\n", + body.len() + ); + let mut stream = tokio::net::TcpStream::connect(client).await.unwrap(); + stream.write_all(request.as_bytes()).await.unwrap(); + stream.write_all(&body).await.unwrap(); + let (status, echoed) = read_response(&mut stream).await; + assert_eq!(status, 200); + assert_eq!(echoed, body); +} + +#[tokio::test] +async fn prefetch_continuation_path_works_end_to_end() { + common::init_logging(); + let (upstream, client) = start_rotating(true, 24 * 1024 * 1024).await; + let url = format!("http://{upstream}/big"); + let (status, body) = raw_exchange( + client, + format!("GET {url} HTTP/1.1\r\nHost: test\r\nConnection: close\r\n\r\n").as_bytes(), + ) + .await; + assert_eq!(status, 200); + assert_eq!(body.len(), 32 * 1024 * 1024); + assert!(body.iter().all(|&b| b == 0xCD)); +} From fa1a11c3b068783d2fd8edc558b4fe998c52c80d Mon Sep 17 00:00:00 2001 From: lhear <121179341+lhear@users.noreply.github.com> Date: Mon, 3 Aug 2026 19:06:29 +0800 Subject: [PATCH 3/7] fix: guard DNS TTL clamp against inverted min_ttl/max_ttl config parse_response used Ord::clamp with unvalidated config bounds; min_ttl > max_ttl panicked the server on any successful resolution. Validate the bounds in init_dns and make the clamp order-safe regardless, with regression tests for both layers. --- src/dns/client.rs | 48 ++++++++++++++++++++++++++++++++++++++++++++--- 1 file changed, 45 insertions(+), 3 deletions(-) diff --git a/src/dns/client.rs b/src/dns/client.rs index 976edfc..5034832 100644 --- a/src/dns/client.rs +++ b/src/dns/client.rs @@ -250,9 +250,9 @@ impl DnsClient { let ttl = if ips.is_empty() { Duration::from_secs(self.config.options.empty_ttl) } else { - Duration::from_secs( - (min_ttl as u64).clamp(self.config.options.min_ttl, self.config.options.max_ttl), - ) + let lo = self.config.options.min_ttl.min(self.config.options.max_ttl); + let hi = self.config.options.min_ttl.max(self.config.options.max_ttl); + Duration::from_secs((min_ttl as u64).clamp(lo, hi)) }; Ok((ips, ttl)) } @@ -412,6 +412,13 @@ pub async fn init_dns(config: &mut DnsConfig) -> Result> { if config.options.protocol == Protocol::Dot && config.tls_domain.is_none() { config.tls_domain = Some(config.upstream.ip().to_string()); } + if config.options.min_ttl > config.options.max_ttl { + return Err(anyhow!( + "dns.options.min_ttl ({}) must not exceed max_ttl ({})", + config.options.min_ttl, + config.options.max_ttl + )); + } Ok(Arc::new(DnsClient::new(config).await?)) } @@ -506,6 +513,41 @@ mod tests { assert_eq!(ttl.as_secs(), c.config.options.max_ttl); } + #[tokio::test] + async fn parse_response_safe_when_ttl_bounds_inverted() { + let cfg = DnsConfig { + upstream: "127.0.0.1:1".parse().unwrap(), + tls_domain: None, + options: DnsOptions { + min_ttl: 5000, + max_ttl: 100, + ..DnsOptions::default() + }, + }; + let c = DnsClient::new(&cfg).await.unwrap(); + let bytes = make_response(8, false, Rcode::NOERROR, 60, &[Ipv4Addr::new(9, 9, 9, 9)]); + let (_, ttl) = c.parse_response(&bytes, 8, Rtype::A).unwrap(); + assert_eq!( + ttl.as_secs(), + 100, + "inverted bounds must not panic and clamp to max" + ); + } + + #[tokio::test] + async fn init_dns_rejects_inverted_ttl_bounds() { + let mut cfg = DnsConfig { + upstream: "127.0.0.1:1".parse().unwrap(), + tls_domain: None, + options: DnsOptions { + min_ttl: 5000, + max_ttl: 100, + ..DnsOptions::default() + }, + }; + assert!(init_dns(&mut cfg).await.is_err()); + } + #[tokio::test] async fn build_query_writes_id_and_domain() { let c = test_client().await; From 153d561dc4c6735cd0a2219473e5c8028c5de244 Mon Sep 17 00:00:00 2001 From: lhear <121179341+lhear@users.noreply.github.com> Date: Mon, 3 Aug 2026 19:06:35 +0800 Subject: [PATCH 4/7] test: prune low-value unit tests, cover upload/connection actors - remove tautological udp_recv_error_does_not_panic (assert is_err() || is_ok() never fails, duplicated the timeout path) - merge frame_cipher_roundtrip into encrypt_decrypt_roundtrip and the twin JSON padding tests into one loop (same code path) - replace always-true deterministic key-derivation asserts with domain-separation asserts (flip input byte, key must change) - add upload loop tests: full payload roundtrip with contiguous seqs, configured start_seq, failure propagation, no-request-on-empty-input - add connection actor tests: auth retry loop on one connection, CONNECT early-data buffering, request parse timeout --- src/client/actor/connection.rs | 207 ++++++++++++++++++++++++ src/client/actor/upload_loop.rs | 274 ++++++++++++++++++++++++++++++++ src/crypto/cipher.rs | 12 +- src/crypto/handshake.rs | 31 ++-- src/dns/transport.rs | 10 -- src/log/mod.rs | 4 +- src/shaper/mod.rs | 28 +--- 7 files changed, 514 insertions(+), 52 deletions(-) diff --git a/src/client/actor/connection.rs b/src/client/actor/connection.rs index 9981516..ff2da6b 100644 --- a/src/client/actor/connection.rs +++ b/src/client/actor/connection.rs @@ -183,3 +183,210 @@ async fn handle_bypass_direct( info!(target = %target, "bypass connection closed"); Ok(ClientConnState::Closed) } + +#[cfg(test)] +mod tests { + use super::*; + use crate::bypass::{BypassRules, BypassRulesBuilder}; + use crate::shaper::{EncodingType, PaddingConfig, ResolvedShaperConfig, TrafficConfig}; + use std::net::SocketAddr; + use std::time::Duration; + use tokio::io::{AsyncReadExt, AsyncWriteExt}; + use tokio::net::{TcpListener, TcpStream}; + use tokio::sync::Mutex; + + async fn tcp_pair() -> (TcpStream, TcpStream) { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let server = tokio::spawn(async move { listener.accept().await.unwrap().0 }); + let client = TcpStream::connect(addr).await.unwrap(); + (server.await.unwrap(), client) + } + + fn test_state( + proxy_auth: Option<(String, String)>, + bypass: Option>, + ) -> Arc { + let traffic = TrafficConfig { + global: PaddingConfig { + padding_threshold: 0, + padding_range: [0, 0], + }, + stages: vec![], + encoding_type: EncodingType::Binary, + max_download_bytes: None, + }; + Arc::new(SharedState { + remote_str: "http://127.0.0.1:1/".to_string(), + auth_header: "Bearer test".to_string(), + traffic_config: traffic.clone(), + resolved_traffic: Arc::new(ResolvedShaperConfig::resolve(&traffic)), + bypass, + server_public_key: None, + proxy_auth, + initial_master: Mutex::new(None), + handshake_lock: Mutex::new(()), + max_download_bytes: None, + max_connections: 8, + max_in_flight_bytes: 1024 * 1024, + upload_concurrency: 4, + }) + } + + fn bypass_loopback() -> Arc { + let mut b = BypassRulesBuilder::new(); + b.add_cidr("127.0.0.1/32").unwrap(); + Arc::new(b.build().unwrap()) + } + + async fn spawn_echo() -> SocketAddr { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + tokio::spawn(async move { + while let Ok((mut sock, _)) = listener.accept().await { + tokio::spawn(async move { + let mut buf = [0u8; 4096]; + loop { + let n = match sock.read(&mut buf).await { + Ok(0) | Err(_) => break, + Ok(n) => n, + }; + if sock.write_all(&buf[..n]).await.is_err() { + break; + } + } + }); + } + }); + addr + } + + fn test_client() -> Arc { + Arc::new(wreq::Client::builder().no_proxy().build().unwrap()) + } + + #[tokio::test] + async fn auth_retry_on_same_connection() { + let echo = spawn_echo().await; + let (server_side, mut client) = tcp_pair().await; + let state = test_state( + Some(("Basic dXNlcjpwYXNz".to_string(), "user".to_string())), + Some(bypass_loopback()), + ); + let http_client = test_client(); + let actor = tokio::spawn(async move { + let mut actor = ClientConnectionActor::new(server_side, http_client, state); + actor.run().await + }); + + client + .write_all( + format!( + "GET http://127.0.0.1:{}/ HTTP/1.1\r\nHost: x\r\n\r\n", + echo.port() + ) + .as_bytes(), + ) + .await + .unwrap(); + let mut buf = [0u8; 4096]; + let n = client.read(&mut buf).await.unwrap(); + let head = String::from_utf8_lossy(&buf[..n]); + assert!( + head.starts_with("HTTP/1.1 407"), + "expected 407, got: {head}" + ); + + client + .write_all( + format!( + "GET http://127.0.0.1:{}/ HTTP/1.1\r\nHost: x\r\nProxy-Authorization: Basic dXNlcjpwYXNz\r\n\r\n", + echo.port() + ) + .as_bytes(), + ) + .await + .unwrap(); + client.shutdown().await.unwrap(); + let mut echoed = Vec::new(); + client.read_to_end(&mut echoed).await.unwrap(); + let echoed_str = String::from_utf8_lossy(&echoed); + assert!( + echoed_str.contains("GET / "), + "expected rewritten request echoed, got: {echoed_str}" + ); + assert!(echoed_str.contains("Host: x")); + + let _ = tokio::time::timeout(Duration::from_secs(5), actor).await; + } + + #[tokio::test] + async fn connect_early_data_buffered_and_forwarded() { + let echo = spawn_echo().await; + let (server_side, mut client) = tcp_pair().await; + let state = test_state(None, Some(bypass_loopback())); + let http_client = test_client(); + let actor = tokio::spawn(async move { + let mut actor = ClientConnectionActor::new(server_side, http_client, state); + actor.run().await + }); + + let req = format!( + "CONNECT 127.0.0.1:{} HTTP/1.1\r\nHost: x\r\n\r\n", + echo.port() + ); + client.write_all(req.as_bytes()).await.unwrap(); + + let mut buf = [0u8; 4096]; + let n = client.read(&mut buf).await.unwrap(); + assert!( + String::from_utf8_lossy(&buf[..n]).starts_with("HTTP/1.1 200"), + "expected 200" + ); + + client.write_all(b"early-data-bytes").await.unwrap(); + tokio::time::sleep(std::time::Duration::from_millis(20)).await; + client.write_all(b"post-window-bytes").await.unwrap(); + client.shutdown().await.unwrap(); + + let mut echoed = Vec::new(); + client.read_to_end(&mut echoed).await.unwrap(); + let echoed_str = String::from_utf8_lossy(&echoed); + assert!( + echoed_str.contains("early-data-bytes"), + "early data must be forwarded, got: {echoed_str}" + ); + assert!(echoed_str.contains("post-window-bytes")); + + let _ = tokio::time::timeout(Duration::from_secs(5), actor).await; + } + + #[tokio::test] + async fn request_parse_timeout_errors() { + tokio::time::pause(); + let (server_side, mut client) = tcp_pair().await; + let state = test_state(None, None); + let http_client = test_client(); + let actor = tokio::spawn(async move { + let mut actor = ClientConnectionActor::new(server_side, http_client, state); + actor.run().await + }); + + client + .write_all(b"GET / HTTP/1.1\r\nHost: x\r\n") + .await + .unwrap(); + + tokio::task::yield_now().await; + tokio::time::advance(PROXY_REQUEST_PARSE_TIMEOUT + Duration::from_secs(1)).await; + tokio::task::yield_now().await; + tokio::time::resume(); + + let res = tokio::time::timeout(Duration::from_secs(5), actor).await; + let joined = res + .expect("actor must finish after parse timeout") + .expect("actor task must not panic"); + let err = joined.unwrap_err(); + assert!(err.to_string().contains("parse timeout"), "got: {err}"); + } +} diff --git a/src/client/actor/upload_loop.rs b/src/client/actor/upload_loop.rs index 6e1e771..4667a69 100644 --- a/src/client/actor/upload_loop.rs +++ b/src/client/actor/upload_loop.rs @@ -228,3 +228,277 @@ async fn send_upload_post( response.bytes().await.context("drain upload response")?; Ok(()) } + +#[cfg(test)] +mod tests { + use super::*; + use crate::shaper::{ + DecodedFrame, EncodingType, PaddingConfig, ResolvedShaperConfig, TrafficConfig, + decode_frame, + }; + use std::net::SocketAddr; + use tokio::io::AsyncWriteExt; + use tokio::net::{TcpListener, TcpStream}; + use tokio::sync::Mutex; + + fn test_state(remote: &str, max_in_flight: usize) -> Arc { + let traffic = TrafficConfig { + global: PaddingConfig { + padding_threshold: 0, + padding_range: [0, 0], + }, + stages: vec![], + encoding_type: EncodingType::Binary, + max_download_bytes: None, + }; + Arc::new(SharedState { + remote_str: remote.to_string(), + auth_header: "Bearer test-token".to_string(), + traffic_config: traffic.clone(), + resolved_traffic: Arc::new(ResolvedShaperConfig::resolve(&traffic)), + bypass: None, + server_public_key: None, + proxy_auth: None, + initial_master: Mutex::new(None), + handshake_lock: Mutex::new(()), + max_download_bytes: None, + max_connections: 8, + max_in_flight_bytes: max_in_flight, + upload_concurrency: 4, + }) + } + + async fn spawn_collector() -> ( + SocketAddr, + tokio::sync::oneshot::Sender<()>, + tokio::task::JoinHandle>, + ) { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let (stop_tx, mut stop_rx) = tokio::sync::oneshot::channel(); + let handle = tokio::spawn(async move { + let mut collected = Vec::new(); + loop { + tokio::select! { + _ = &mut stop_rx => break, + r = listener.accept() => { + let Ok((mut sock, _)) = r else { break }; + let mut buf = Vec::new(); + let mut tmp = [0u8; 8192]; + let header_end = loop { + let n = match sock.read(&mut tmp).await { + Ok(0) | Err(_) => break None, + Ok(n) => n, + }; + buf.extend_from_slice(&tmp[..n]); + if let Some(p) = buf.windows(4).position(|w| w == b"\r\n\r\n") { + break Some(p + 4); + } + }; + let Some(header_end) = header_end else { + continue; + }; + let headers = String::from_utf8_lossy(&buf[..header_end]); + let content_length = headers + .lines() + .find(|l| l.to_ascii_lowercase().starts_with("content-length:")) + .and_then(|l| l.split(':').nth(1)) + .and_then(|v| v.trim().parse::().ok()) + .unwrap_or(0); + let mut body = buf[header_end..].to_vec(); + while body.len() < content_length { + let n = match sock.read(&mut tmp).await { + Ok(0) | Err(_) => break, + Ok(n) => n, + }; + body.extend_from_slice(&tmp[..n]); + } + body.truncate(content_length); + collected.extend_from_slice(&body); + let _ = sock + .write_all( + b"HTTP/1.1 204 No Content\r\nContent-Length: 0\r\nConnection: close\r\n\r\n", + ) + .await; + } + } + } + collected + }); + (addr, stop_tx, handle) + } + + fn decode_all(received: &[u8]) -> Vec { + let mut scratch = BytesMut::new(); + let mut json_scratch = Vec::new(); + let mut src = BytesMut::from(received); + let mut decoded = Vec::new(); + while let Some(frame) = decode_frame( + &mut src, + &mut scratch, + &mut json_scratch, + None, + EncodingType::Binary, + ) + .unwrap() + { + match frame { + DecodedFrame::Owned { data, .. } => decoded.extend_from_slice(&data), + DecodedFrame::InScratch { .. } => panic!("unexpected InScratch frame"), + } + } + decoded + } + + fn test_client() -> Arc { + Arc::new(wreq::Client::builder().no_proxy().build().unwrap()) + } + + async fn tcp_pair() -> (tokio::net::tcp::OwnedReadHalf, TcpStream) { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let server = tokio::spawn(async move { listener.accept().await.unwrap().0 }); + let client = TcpStream::connect(addr).await.unwrap(); + let server_stream = server.await.unwrap(); + let (read_half, _write_half) = server_stream.into_split(); + (read_half, client) + } + + #[tokio::test] + async fn run_uploads_all_data_with_contiguous_seqs() { + let (addr, stop_tx, collector) = spawn_collector().await; + let state = test_state(&format!("http://{addr}/"), 64 * 1024); + let client = test_client(); + let (read_half, mut writer) = tcp_pair().await; + + let initial: Vec = (0..40_000u32).map(|i| (i % 251) as u8).collect(); + let extra: Vec = (0..20_000u32).map(|i| (i % 253) as u8).collect(); + let mut all = initial.clone(); + all.extend_from_slice(&extra); + + let actor = UploadLoopActor::new( + client, + state, + Bytes::from(initial), + read_half, + None, + Uuid::new_v4(), + 0, + ); + let handle = tokio::spawn(async move { actor.run().await }); + + writer.write_all(&extra).await.unwrap(); + drop(writer); + handle + .await + .unwrap() + .expect("upload loop must finish cleanly"); + + let _ = stop_tx.send(()); + let received = collector.await.unwrap(); + assert!(!received.is_empty(), "server must receive upload frames"); + assert_eq!(decode_all(&received), all); + } + + #[tokio::test] + async fn run_starts_seq_from_configured_value() { + let (addr, stop_tx, collector) = spawn_collector().await; + let state = test_state(&format!("http://{addr}/"), 64 * 1024); + let client = test_client(); + let (read_half, writer) = tcp_pair().await; + drop(writer); + + let initial = vec![0xABu8; 40_000]; + let actor = UploadLoopActor::new( + client, + state, + Bytes::from(initial), + read_half, + None, + Uuid::new_v4(), + 7, + ); + let handle = tokio::spawn(async move { actor.run().await }); + handle.await.unwrap().unwrap(); + + let _ = stop_tx.send(()); + let received = collector.await.unwrap(); + let mut scratch = BytesMut::new(); + let mut json_scratch = Vec::new(); + let mut src = BytesMut::from(&received[..]); + let mut seqs = Vec::new(); + while let Some(frame) = decode_frame( + &mut src, + &mut scratch, + &mut json_scratch, + None, + EncodingType::Binary, + ) + .unwrap() + { + match frame { + DecodedFrame::Owned { seq, .. } => seqs.push(seq), + DecodedFrame::InScratch { .. } => panic!("unexpected InScratch frame"), + } + } + assert_eq!(seqs.first(), Some(&7)); + for w in seqs.windows(2) { + assert_eq!(w[1], w[0] + 1, "seqs must be contiguous"); + } + } + + #[tokio::test] + async fn upload_failure_propagates_error() { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let dead = listener.local_addr().unwrap(); + drop(listener); + + let state = test_state(&format!("http://{dead}/"), 64 * 1024); + let client = test_client(); + let (read_half, writer) = tcp_pair().await; + drop(writer); + + let initial = vec![0xCDu8; 30_000]; + let actor = UploadLoopActor::new( + client, + state, + Bytes::from(initial), + read_half, + None, + Uuid::new_v4(), + 0, + ); + let result = tokio::time::timeout(std::time::Duration::from_secs(10), actor.run()).await; + let err = result.expect("upload loop must fail fast").unwrap_err(); + let msg = err.to_string(); + assert!( + msg.contains("upload POST failed") || msg.contains("http post failed"), + "unexpected error: {msg}" + ); + } + + #[tokio::test] + async fn empty_input_sends_no_requests() { + let (addr, stop_tx, collector) = spawn_collector().await; + let state = test_state(&format!("http://{addr}/"), 64 * 1024); + let client = test_client(); + let (read_half, writer) = tcp_pair().await; + drop(writer); + + let actor = UploadLoopActor::new( + client, + state, + Bytes::new(), + read_half, + None, + Uuid::new_v4(), + 5, + ); + let handle = tokio::spawn(async move { actor.run().await }); + handle.await.unwrap().unwrap(); + + let _ = stop_tx.send(()); + let received = collector.await.unwrap(); + assert!(received.is_empty(), "empty input must not send any POST"); + } +} diff --git a/src/crypto/cipher.rs b/src/crypto/cipher.rs index 76e97ab..da13e38 100644 --- a/src/crypto/cipher.rs +++ b/src/crypto/cipher.rs @@ -181,19 +181,15 @@ mod tests { fn encrypt_decrypt_roundtrip() { let key = random_key(); let plain = b"hello world test frame data"; + let ct = encrypt_bytes(&key, plain).unwrap(); let pt = decrypt_bytes(&key, &ct).unwrap(); assert_eq!(pt, plain); - } - #[test] - fn frame_cipher_roundtrip() { - let key = random_key(); let cipher = AesFrameCipher::new(&key); - let data = b"frame data for cipher test"; - let ct = cipher.encrypt(data).unwrap(); - let pt = cipher.decrypt(&ct).unwrap(); - assert_eq!(pt, data); + let frame_ct = cipher.encrypt(plain).unwrap(); + let frame_pt = cipher.decrypt(&frame_ct).unwrap(); + assert_eq!(frame_pt, plain); } #[test] diff --git a/src/crypto/handshake.rs b/src/crypto/handshake.rs index d093d4e..7dd18b2 100644 --- a/src/crypto/handshake.rs +++ b/src/crypto/handshake.rs @@ -52,31 +52,40 @@ mod tests { use super::*; #[test] - fn derive_handshake_key_deterministic() { - let shared = [0xAAu8; 32]; + fn derive_handshake_key_domain_separated() { + let mut shared = [0xAAu8; 32]; let k1 = derive_handshake_key(&shared); + shared[0] ^= 0x01; let k2 = derive_handshake_key(&shared); - assert_eq!(*k1, *k2); + assert_ne!( + *k1, *k2, + "different shared secrets must derive different keys" + ); } #[test] - fn derive_initial_master_deterministic() { - let ml = [0x11u8; 32]; + fn derive_initial_master_domain_separated() { + let mut ml = [0x11u8; 32]; let x2 = [0x22u8; 32]; let m1 = derive_initial_master(&ml, &x2); + ml[0] ^= 0x01; let m2 = derive_initial_master(&ml, &x2); - assert_eq!(*m1, *m2); + assert_ne!( + *m1, *m2, + "different ML-KEM secrets must derive different masters" + ); } #[test] - fn connection_keys_deterministic() { - let master = [0xBBu8; 32]; + fn connection_keys_domain_separated() { + let mut master = [0xBBu8; 32]; let nonce = [0xCCu8; 16]; let (up1, dn1, tg1) = derive_connection_keys(&master, &nonce); + master[0] ^= 0x01; let (up2, dn2, tg2) = derive_connection_keys(&master, &nonce); - assert_eq!(up1, up2); - assert_eq!(dn1, dn2); - assert_eq!(tg1, tg2); + assert_ne!(up1, up2); + assert_ne!(dn1, dn2); + assert_ne!(tg1, tg2); } #[test] diff --git a/src/dns/transport.rs b/src/dns/transport.rs index 255a35b..5c57a33 100644 --- a/src/dns/transport.rs +++ b/src/dns/transport.rs @@ -351,14 +351,4 @@ mod tests { let r = t.send(&mut query).await; assert!(r.is_err()); } - - #[tokio::test] - async fn udp_recv_error_does_not_panic() { - let t = UdpTransport::new("127.0.0.1:9".parse().unwrap()) - .await - .unwrap(); - let mut query = [0u8; 12]; - let r = t.send(&mut query).await; - assert!(r.is_err() || r.is_ok()); - } } diff --git a/src/log/mod.rs b/src/log/mod.rs index 3ebc796..36e6c42 100644 --- a/src/log/mod.rs +++ b/src/log/mod.rs @@ -111,8 +111,8 @@ mod tests { } #[test] - fn file_writer_without_stem_uses_fallback_name() { - let dir = std::env::temp_dir().join(format!("httproxy_log_nostem_{}", std::process::id())); + fn file_writer_without_extension_uses_fallback_suffix() { + let dir = std::env::temp_dir().join(format!("httproxy_log_noext_{}", std::process::id())); let path = dir.join("noext"); let (_, guard) = build_file_writer(path.to_str().unwrap(), 3); assert!(dir.exists()); diff --git a/src/shaper/mod.rs b/src/shaper/mod.rs index 2ec6568..62085bd 100644 --- a/src/shaper/mod.rs +++ b/src/shaper/mod.rs @@ -1244,27 +1244,13 @@ mod tests { #[test] fn encode_frame_json_padding_stays_within_line_limit() { let raw = vec![0x42u8; 1000]; - let frame = encode_frame(&raw, 0, None, 100_000, [0, 100_000], EncodingType::Json).unwrap(); - let line_len = frame.iter().position(|&b| b == b'\n').unwrap(); - assert!(line_len <= MAX_JSON_LINE_LEN); - } - - #[test] - fn encode_frame_json_cipher_padding_stays_within_line_limit() { - let raw = vec![0x42u8; 1000]; - let key = zeroize::Zeroizing::new([0u8; 32]); - let cipher = crate::crypto::AesFrameCipher::new(&key); - let frame = encode_frame( - &raw, - 0, - Some(&cipher), - 100_000, - [0, 100_000], - EncodingType::Json, - ) - .unwrap(); - let line_len = frame.iter().position(|&b| b == b'\n').unwrap(); - assert!(line_len <= MAX_JSON_LINE_LEN); + let aes = crate::crypto::AesFrameCipher::new(&zeroize::Zeroizing::new([0u8; 32])); + for cipher in [None, Some(&aes as &dyn FrameCipher)] { + let frame = + encode_frame(&raw, 0, cipher, 100_000, [0, 100_000], EncodingType::Json).unwrap(); + let line_len = frame.iter().position(|&b| b == b'\n').unwrap(); + assert!(line_len <= MAX_JSON_LINE_LEN); + } } #[tokio::test] From 0c262620b5c4abed6068bbad827ff1d0eb25f422 Mon Sep 17 00:00:00 2001 From: lhear <121179341+lhear@users.noreply.github.com> Date: Mon, 3 Aug 2026 19:26:35 +0800 Subject: [PATCH 5/7] ci: gate builds on unit and integration tests, clippy, and fmt --- .github/workflows/build.yml | 8 +++++++- 1 file changed, 7 insertions(+), 1 deletion(-) diff --git a/.github/workflows/build.yml b/.github/workflows/build.yml index d7614ba..0a57e97 100644 --- a/.github/workflows/build.yml +++ b/.github/workflows/build.yml @@ -21,7 +21,7 @@ jobs: test: needs: lint-audit - name: cargo test --lib + name: cargo test (unit + integration) runs-on: ubuntu-latest steps: - uses: actions/checkout@v6 @@ -29,8 +29,14 @@ jobs: uses: swatinem/rust-cache@v2 with: key: test-ubuntu + - name: Check formatting + run: cargo fmt --check + - name: Lint with clippy + run: cargo clippy --all-targets -- -D warnings - name: Run unit tests run: cargo test --lib + - name: Run integration tests + run: cargo test --test bypass --test dns --test encrypted_tunnel --test plain_tunnel --test rotation linux-gnu: needs: [lint-audit, test] From f37a1588d786f6e577aab9d57577d8df88171afe Mon Sep 17 00:00:00 2001 From: lhear <121179341+lhear@users.noreply.github.com> Date: Mon, 3 Aug 2026 19:38:12 +0800 Subject: [PATCH 6/7] fix: bound request body read with 30s deadline to prevent slow-upload DoS --- src/error/mod.rs | 8 ++++++++ src/server/constants.rs | 1 + src/server/handlers.rs | 26 +++++++++++++++++++------- 3 files changed, 28 insertions(+), 7 deletions(-) diff --git a/src/error/mod.rs b/src/error/mod.rs index ebbaa2d..1f07bcd 100644 --- a/src/error/mod.rs +++ b/src/error/mod.rs @@ -72,6 +72,10 @@ impl ServerError { Self(StatusCode::GATEWAY_TIMEOUT, msg.into()) } #[inline] + pub fn request_timeout(msg: impl Into) -> Self { + Self(StatusCode::REQUEST_TIMEOUT, msg.into()) + } + #[inline] pub fn unauthorized(msg: impl Into) -> Self { Self(StatusCode::UNAUTHORIZED, msg.into()) } @@ -138,6 +142,10 @@ mod tests { ServerError::gateway_timeout("x").0, StatusCode::GATEWAY_TIMEOUT ); + assert_eq!( + ServerError::request_timeout("x").0, + StatusCode::REQUEST_TIMEOUT + ); assert_eq!(ServerError::unauthorized("x").0, StatusCode::UNAUTHORIZED); assert_eq!(ServerError::not_found("x").0, StatusCode::NOT_FOUND); assert_eq!( diff --git a/src/server/constants.rs b/src/server/constants.rs index 543254f..d305008 100644 --- a/src/server/constants.rs +++ b/src/server/constants.rs @@ -18,6 +18,7 @@ pub const MASTER_CLEANUP_INTERVAL: Duration = Duration::from_secs(60); pub const CONNECT_TIMEOUT: Duration = Duration::from_secs(10); pub const WRITE_TIMEOUT: Duration = Duration::from_secs(30); +pub const BODY_READ_TIMEOUT: Duration = Duration::from_secs(30); pub const DOWNLOAD_CHANNEL_CAPACITY: usize = 2; pub const TUNNEL_CMD_CHANNEL_CAPACITY: usize = 32; diff --git a/src/server/handlers.rs b/src/server/handlers.rs index cc708cd..ecf6f8a 100644 --- a/src/server/handlers.rs +++ b/src/server/handlers.rs @@ -14,8 +14,8 @@ use crate::error::ServerError; use crate::server::actor::tunnel::TunnelCmd; use crate::server::connection::connect_upstream; use crate::server::constants::{ - CONNECT_TIMEOUT, DOWNLOAD_CHANNEL_CAPACITY, MASTER_EXPIRY, MAX_FRAME_BUF_SIZE, - TUNNEL_CMD_CHANNEL_CAPACITY, + BODY_READ_TIMEOUT, CONNECT_TIMEOUT, DOWNLOAD_CHANNEL_CAPACITY, MASTER_EXPIRY, + MAX_FRAME_BUF_SIZE, TUNNEL_CMD_CHANNEL_CAPACITY, }; use crate::server::stream::FrameDecoder; use crate::server::stream_registry::StreamQueryResult; @@ -156,7 +156,11 @@ async fn setup_tunnel_response( MAX_FRAME_BUF_SIZE, ); - while let Some(result) = decoder.next().await { + let body_deadline = tokio::time::Instant::now() + BODY_READ_TIMEOUT; + while let Some(result) = tokio::time::timeout_at(body_deadline, decoder.next()) + .await + .map_err(|_| ServerError::request_timeout("request body read timed out"))? + { let (seq, data) = result.map_err(|e| ServerError::bad_request(format!("decode error: {e}")))?; upload_tx @@ -339,7 +343,11 @@ async fn dispatch_to_actor(handle: SessionHandle, body: Body) -> Result Date: Mon, 3 Aug 2026 19:52:30 +0800 Subject: [PATCH 7/7] fix: harden DNS against cache poisoning via per-query ports and question validation --- src/dns/client.rs | 149 +++++++++++++++++++++++++++++++++++++++---- src/dns/transport.rs | 82 ++++++------------------ tests/dns.rs | 20 +++--- 3 files changed, 165 insertions(+), 86 deletions(-) diff --git a/src/dns/client.rs b/src/dns/client.rs index 5034832..14015b6 100644 --- a/src/dns/client.rs +++ b/src/dns/client.rs @@ -58,7 +58,7 @@ pub struct DnsClient { impl DnsClient { pub async fn new(config: &DnsConfig) -> Result { let transport = match config.options.protocol { - Protocol::Udp => Transport::Udp(UdpTransport::new(config.upstream).await?), + Protocol::Udp => Transport::Udp(UdpTransport::new(config.upstream)), Protocol::Dot => Transport::Dot(init_dot_transport(config)?), }; Ok(Self { @@ -164,7 +164,7 @@ impl DnsClient { Transport::Dot(dot) => dot.send(&mut query).await?, }; - self.parse_response(&resp, id, rtype) + self.parse_response(&resp, id, rtype, domain) } fn build_query( @@ -207,6 +207,7 @@ impl DnsClient { data: &[u8], id: u16, qtype: Rtype, + domain: &str, ) -> Result<(Vec, Duration)> { let msg = Message::from_octets(data).map_err(|_| anyhow!("invalid DNS response"))?; if msg.header().id() != id { @@ -216,6 +217,19 @@ impl DnsClient { msg.header().id() )); } + let question = msg + .sole_question() + .map_err(|_| anyhow!("DNS response question missing or malformed"))?; + let name_matches = question + .qname() + .to_string() + .trim_end_matches('.') + .eq_ignore_ascii_case(domain.trim_end_matches('.')); + if question.qtype() != qtype || !name_matches { + return Err(anyhow!( + "DNS response question does not match query for {domain}" + )); + } if msg.header().tc() { return Err(anyhow!("DNS response truncated (TC set)")); } @@ -429,16 +443,36 @@ mod tests { use std::net::Ipv4Addr; fn make_response(id: u16, tc: bool, rcode: Rcode, ttl: u32, a_ips: &[Ipv4Addr]) -> Vec { + make_response_for(id, tc, rcode, ttl, a_ips, "example.com", Rtype::A) + } + + fn make_response_for( + id: u16, + tc: bool, + rcode: Rcode, + ttl: u32, + a_ips: &[Ipv4Addr], + qname: &str, + qtype: Rtype, + ) -> Vec { let mut builder = MessageBuilder::new_vec(); builder.header_mut().set_id(id); builder.header_mut().set_rcode(rcode); if tc { builder.header_mut().set_tc(true); } - let mut answer = builder.answer(); + let mut question = builder.question(); + question + .push(Question::new( + Name::>::from_str(qname).unwrap(), + qtype, + Class::IN, + )) + .unwrap(); + let mut answer = question.answer(); for ip in a_ips { let rec = Record::new( - Name::>::from_str("example.com").unwrap(), + Name::>::from_str(qname).unwrap(), Class::IN, Ttl::from_secs(ttl), A::new(*ip), @@ -461,7 +495,9 @@ mod tests { async fn parse_response_accepts_valid_a_record() { let c = test_client().await; let bytes = make_response(1, false, Rcode::NOERROR, 120, &[Ipv4Addr::new(1, 2, 3, 4)]); - let (ips, ttl) = c.parse_response(&bytes, 1, Rtype::A).unwrap(); + let (ips, ttl) = c + .parse_response(&bytes, 1, Rtype::A, "example.com") + .unwrap(); assert_eq!(ips, vec![IpAddr::V4(Ipv4Addr::new(1, 2, 3, 4))]); assert_eq!(ttl.as_secs(), 120); } @@ -470,21 +506,29 @@ mod tests { async fn parse_response_rejects_truncated() { let c = test_client().await; let bytes = make_response(2, true, Rcode::NOERROR, 120, &[]); - assert!(c.parse_response(&bytes, 2, Rtype::A).is_err()); + assert!( + c.parse_response(&bytes, 2, Rtype::A, "example.com") + .is_err() + ); } #[tokio::test] async fn parse_response_rejects_id_mismatch() { let c = test_client().await; let bytes = make_response(3, false, Rcode::NOERROR, 120, &[]); - assert!(c.parse_response(&bytes, 99, Rtype::A).is_err()); + assert!( + c.parse_response(&bytes, 99, Rtype::A, "example.com") + .is_err() + ); } #[tokio::test] async fn parse_response_nxdomain_returns_empty_with_empty_ttl() { let c = test_client().await; let bytes = make_response(4, false, Rcode::NXDOMAIN, 120, &[]); - let (ips, ttl) = c.parse_response(&bytes, 4, Rtype::A).unwrap(); + let (ips, ttl) = c + .parse_response(&bytes, 4, Rtype::A, "example.com") + .unwrap(); assert!(ips.is_empty()); assert_eq!(ttl.as_secs(), c.config.options.empty_ttl); } @@ -493,14 +537,19 @@ mod tests { async fn parse_response_rejects_error_rcode() { let c = test_client().await; let bytes = make_response(5, false, Rcode::SERVFAIL, 120, &[]); - assert!(c.parse_response(&bytes, 5, Rtype::A).is_err()); + assert!( + c.parse_response(&bytes, 5, Rtype::A, "example.com") + .is_err() + ); } #[tokio::test] async fn parse_response_clamps_ttl() { let c = test_client().await; let bytes = make_response(6, false, Rcode::NOERROR, 10, &[Ipv4Addr::new(9, 9, 9, 9)]); - let (_, ttl) = c.parse_response(&bytes, 6, Rtype::A).unwrap(); + let (_, ttl) = c + .parse_response(&bytes, 6, Rtype::A, "example.com") + .unwrap(); assert_eq!(ttl.as_secs(), c.config.options.min_ttl); let bytes = make_response( 7, @@ -509,7 +558,9 @@ mod tests { 999_999, &[Ipv4Addr::new(9, 9, 9, 9)], ); - let (_, ttl) = c.parse_response(&bytes, 7, Rtype::A).unwrap(); + let (_, ttl) = c + .parse_response(&bytes, 7, Rtype::A, "example.com") + .unwrap(); assert_eq!(ttl.as_secs(), c.config.options.max_ttl); } @@ -526,7 +577,9 @@ mod tests { }; let c = DnsClient::new(&cfg).await.unwrap(); let bytes = make_response(8, false, Rcode::NOERROR, 60, &[Ipv4Addr::new(9, 9, 9, 9)]); - let (_, ttl) = c.parse_response(&bytes, 8, Rtype::A).unwrap(); + let (_, ttl) = c + .parse_response(&bytes, 8, Rtype::A, "example.com") + .unwrap(); assert_eq!( ttl.as_secs(), 100, @@ -534,6 +587,64 @@ mod tests { ); } + #[tokio::test] + async fn parse_response_rejects_question_name_mismatch() { + let c = test_client().await; + let bytes = make_response_for( + 9, + false, + Rcode::NOERROR, + 120, + &[Ipv4Addr::new(1, 2, 3, 4)], + "evil.com", + Rtype::A, + ); + assert!( + c.parse_response(&bytes, 9, Rtype::A, "example.com") + .is_err() + ); + } + + #[tokio::test] + async fn parse_response_rejects_question_type_mismatch() { + let c = test_client().await; + let bytes = make_response_for( + 10, + false, + Rcode::NOERROR, + 120, + &[Ipv4Addr::new(1, 2, 3, 4)], + "example.com", + Rtype::AAAA, + ); + assert!( + c.parse_response(&bytes, 10, Rtype::A, "example.com") + .is_err() + ); + } + + #[tokio::test] + async fn parse_response_rejects_missing_question() { + let c = test_client().await; + let mut builder = MessageBuilder::new_vec(); + builder.header_mut().set_id(11); + builder.header_mut().set_rcode(Rcode::NOERROR); + let mut answer = builder.answer(); + answer + .push(Record::new( + Name::>::from_str("example.com").unwrap(), + Class::IN, + Ttl::from_secs(120), + A::new(Ipv4Addr::new(1, 2, 3, 4)), + )) + .unwrap(); + let bytes = answer.into_message().into_octets(); + assert!( + c.parse_response(&bytes, 11, Rtype::A, "example.com") + .is_err() + ); + } + #[tokio::test] async fn init_dns_rejects_inverted_ttl_bounds() { let mut cfg = DnsConfig { @@ -581,7 +692,15 @@ mod tests { let mut builder = MessageBuilder::new_vec(); builder.header_mut().set_id(10); builder.header_mut().set_rcode(Rcode::NOERROR); - let mut answer = builder.answer(); + let mut question = builder.question(); + question + .push(Question::new( + Name::>::from_str("example.com").unwrap(), + Rtype::AAAA, + Class::IN, + )) + .unwrap(); + let mut answer = question.answer(); let v6 = std::net::Ipv6Addr::new(0x2001, 0xdb8, 0, 0, 0, 0, 0, 1); let rec = Record::new( Name::>::from_str("example.com").unwrap(), @@ -591,7 +710,9 @@ mod tests { ); answer.push(rec).unwrap(); let bytes = answer.into_message().into_octets(); - let (ips, ttl) = c.parse_response(&bytes, 10, Rtype::AAAA).unwrap(); + let (ips, ttl) = c + .parse_response(&bytes, 10, Rtype::AAAA, "example.com") + .unwrap(); assert_eq!(ips, vec![IpAddr::V6(v6)]); assert_eq!(ttl.as_secs(), 300); } diff --git a/src/dns/transport.rs b/src/dns/transport.rs index 5c57a33..3db49fd 100644 --- a/src/dns/transport.rs +++ b/src/dns/transport.rs @@ -16,7 +16,7 @@ use tokio_rustls::{ client::TlsStream, rustls::{self, RootCertStore, pki_types::ServerName}, }; -use tracing::{debug, error, warn}; +use tracing::{debug, warn}; use super::config::DnsConfig; @@ -42,70 +42,25 @@ async fn assign_id_and_register( } pub(super) struct UdpTransport { - socket: Arc, - pending: PendingMap, - recv_handle: tokio::task::AbortHandle, -} - -impl Drop for UdpTransport { - fn drop(&mut self) { - self.recv_handle.abort(); - } + upstream: SocketAddr, } impl UdpTransport { - pub(super) async fn new(upstream: SocketAddr) -> Result { - let socket = Arc::new(UdpSocket::bind("0.0.0.0:0").await?); - socket.connect(upstream).await?; - let pending: PendingMap = Default::default(); - let (rs, rp) = (socket.clone(), pending.clone()); - let handle = tokio::spawn(async move { - let mut buf = vec![0u8; 65535]; - let mut consecutive_errors = 0u32; - loop { - match rs.recv(&mut buf).await { - Ok(len) if len >= 2 => { - consecutive_errors = 0; - let id = u16::from_be_bytes([buf[0], buf[1]]); - if let Some(tx) = rp.lock().await.remove(&id) { - let _ = tx.send(Ok(buf[..len].to_vec())); - } - } - Ok(_) => { - consecutive_errors = 0; - } - Err(e) => { - consecutive_errors += 1; - error!(error = %e, consecutive_errors, "UDP recv error"); - tokio::time::sleep(Duration::from_secs(3)).await; - } - } - } - }); - Ok(Self { - socket, - pending, - recv_handle: handle.abort_handle(), - }) + pub(super) fn new(upstream: SocketAddr) -> Self { + Self { upstream } } pub(super) async fn send(&self, data: &mut [u8]) -> Result<(Vec, u16)> { - let (tx, rx) = oneshot::channel(); - let id = assign_id_and_register(&self.pending, data, tx).await; - - if let Err(e) = self.socket.send(data).await { - self.pending.lock().await.remove(&id); - return Err(anyhow!("UDP send failed: {}", e)); - } - - match timeout(Duration::from_secs(2), rx).await { - Ok(Ok(res)) => Ok((res?, id)), - Ok(Err(_)) => Err(anyhow!("UDP channel closed")), - Err(_) => { - self.pending.lock().await.remove(&id); - Err(anyhow!("UDP upstream timeout")) - } - } + let socket = UdpSocket::bind("0.0.0.0:0").await?; + socket.connect(self.upstream).await?; + let id: u16 = rand::random(); + data[0..2].copy_from_slice(&id.to_be_bytes()); + socket.send(data).await?; + let mut buf = vec![0u8; 65535]; + let len = timeout(Duration::from_secs(2), socket.recv(&mut buf)) + .await + .context("UDP upstream timeout")??; + Ok((buf[..len].to_vec(), id)) } } @@ -298,7 +253,7 @@ mod tests { #[tokio::test] async fn udp_transport_roundtrip() { let server_addr = spawn_udp_echo_server().await; - let t = UdpTransport::new(server_addr).await.unwrap(); + let t = UdpTransport::new(server_addr); let mut query = [0u8; 12]; query[12 - 12] = 0; let (resp, id) = t.send(&mut query).await.unwrap(); @@ -308,10 +263,9 @@ mod tests { #[tokio::test] async fn udp_transport_times_out_when_no_reply() { - let sock = UdpSocket::bind("127.0.0.1:0").await.unwrap(); - let addr = sock.local_addr().unwrap(); - drop(sock); - let t = UdpTransport::new(addr).await.unwrap(); + let listener = UdpSocket::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let t = UdpTransport::new(addr); let mut query = [0u8; 12]; let start = std::time::Instant::now(); let r = t.send(&mut query).await; diff --git a/tests/dns.rs b/tests/dns.rs index 9c6bf44..db352a8 100644 --- a/tests/dns.rs +++ b/tests/dns.rs @@ -9,7 +9,7 @@ use std::sync::atomic::{AtomicUsize, Ordering}; use tokio::io::{AsyncReadExt, AsyncWriteExt}; use tokio::net::UdpSocket; -fn extract_domain(query: &[u8], mut offset: usize) -> Option { +fn extract_question(query: &[u8], mut offset: usize) -> Option<(String, u16)> { let mut labels = Vec::new(); loop { let len = *query.get(offset)? as usize; @@ -21,15 +21,16 @@ fn extract_domain(query: &[u8], mut offset: usize) -> Option { labels.push(label.to_string()); offset += len; } - Some(labels.join(".")) + let qtype = u16::from_be_bytes([*query.get(offset)?, *query.get(offset + 1)?]); + Some((labels.join("."), qtype)) } async fn spawn_mock_dns( records: HashMap, query_count: Arc, ) -> std::net::SocketAddr { - use domain::base::iana::{Class, Rcode}; - use domain::base::{MessageBuilder, Name, Record, Ttl}; + use domain::base::iana::{Class, Rcode, Rtype}; + use domain::base::{MessageBuilder, Name, Question, Record, Ttl}; use domain::rdata::{A, Aaaa}; use std::str::FromStr; @@ -47,7 +48,8 @@ async fn spawn_mock_dns( } query_count.fetch_add(1, Ordering::Relaxed); let id = u16::from_be_bytes([query[0], query[1]]); - let domain = extract_domain(query, 12).unwrap_or_default(); + let (domain, qtype_int) = extract_question(query, 12).unwrap_or_default(); + let qtype = Rtype::from_int(qtype_int); let mut builder = MessageBuilder::new_vec(); builder.header_mut().set_id(id); if records.contains_key(&domain) { @@ -56,17 +58,19 @@ async fn spawn_mock_dns( builder.header_mut().set_rcode(Rcode::NXDOMAIN); } let name = Name::>::from_str(&domain).unwrap_or_else(|_| Name::root()); - let mut answer = builder.answer(); + let mut question = builder.question(); + let _ = question.push(Question::new(name.clone(), qtype, Class::IN)); + let mut answer = question.answer(); if let Some(ip) = records.get(&domain) { let _ = match ip { IpAddr::V4(v4) => answer.push(Record::new( - name, + name.clone(), Class::IN, Ttl::from_secs(60), A::new(*v4), )), IpAddr::V6(v6) => answer.push(Record::new( - name, + name.clone(), Class::IN, Ttl::from_secs(60), Aaaa::new(*v6),