Compare commits

...

60 Commits

Author SHA1 Message Date
Alexey d851200e47 Merge pull request #883 from telemt/flow/3.4.25
Flow/3.4.25
2026-07-20 10:36:20 +03:00
Alexey a5216d77fb Rustfmt 2026-07-19 18:41:23 +03:00
Alexey 5b5cd952c8 Bump -> 3.4.25 2026-07-19 18:36:04 +03:00
Alexey a40578b278 Merge pull request #882 from telemt/flow-runtime-reload
In-Runtime Reload
2026-07-19 17:26:43 +03:00
Alexey bd99493622 Merge pull request #881 from astronaut808/feature/handshake-stage-accounting
add handshake failure stage accounting
2026-07-19 17:20:28 +03:00
Alexey d862deebb2 Fix handshake stage attribution and bound counters
Co-Authored-By: brekotis <93345790+brekotis@users.noreply.github.com>
2026-07-19 17:20:18 +03:00
Alexey 0d869d716c Merge branch 'flow-runtime-reload' into feature/handshake-stage-accounting 2026-07-19 17:18:18 +03:00
Alexey 1e85b91ad3 LCOV artifact for pull request + manual start
Co-Authored-By: brekotis <93345790+brekotis@users.noreply.github.com>
2026-07-19 16:11:31 +03:00
Alexey 4679bdcfd5 Atomic Maestro Sessions + Shutdown gate
Co-Authored-By: brekotis <93345790+brekotis@users.noreply.github.com>
2026-07-19 16:10:42 +03:00
astronaut808 7291c3192c docs handshake failure stage diagnostics 2026-07-18 21:25:40 +05:00
astronaut808 4a5dc0b21b add handshake failure stage accounting 2026-07-18 20:46:00 +05:00
Alexey fabd98ce89 Cfg Test for handle_bad_client 2026-07-18 14:31:43 +03:00
Alexey 7df3dab5e8 Update API.md
Co-Authored-By: brekotis <93345790+brekotis@users.noreply.github.com>
2026-07-18 14:30:03 +03:00
Alexey c6f40e3717 Harden Maestro reload lifecycle and readiness barriers
Co-Authored-By: brekotis <93345790+brekotis@users.noreply.github.com>
2026-07-18 14:27:24 +03:00
Alexey 91e05265be Docs for API In-Proccess Reload
Co-Authored-By: brekotis <93345790+brekotis@users.noreply.github.com>
2026-07-17 23:56:42 +03:00
Alexey 991d5b2c38 Maestro: add in-process runtime generation reload
Co-Authored-By: brekotis <93345790+brekotis@users.noreply.github.com>
2026-07-17 23:56:00 +03:00
Alexey 1f9c82c924 Merge pull request #875 from telemt/flow-aescr
Bound ME writer queues by resident payload bytes + Enable expanded AES key schedule zeroization
2026-07-14 22:31:27 +03:00
Alexey 73afeccae1 Rustfmt 2026-07-13 12:20:24 +03:00
Alexey feb51cbf57 Fix RustCrypto zeroize feature wiring 2026-07-12 09:27:15 +03:00
Alexey 8c65cd868c Enable expanded AES key schedule zeroization
Co-Authored-By: brekotis <93345790+brekotis@users.noreply.github.com>
2026-07-11 21:43:01 +03:00
Alexey ea296bbdc8 Replace per-session pool trimming with pressure hysteresis
Co-Authored-By: brekotis <93345790+brekotis@users.noreply.github.com>
2026-07-11 21:36:01 +03:00
Alexey fb042f826e Optimize crypto and Fake-TLS buffer residency
Co-Authored-By: brekotis <93345790+brekotis@users.noreply.github.com>
2026-07-11 21:22:17 +03:00
Alexey 96425f15c8 Bound Direct relay buffers with an adaptive global memory envelope
Co-Authored-By: brekotis <93345790+brekotis@users.noreply.github.com>
2026-07-11 20:56:54 +03:00
Alexey d4c4980e5a Bound ME writer queues by resident payload bytes
Co-Authored-By: brekotis <93345790+brekotis@users.noreply.github.com>
2026-07-11 18:43:42 +03:00
Alexey 893ce0cf36 Hold C2ME byte permits through ME writer completion
Co-Authored-By: brekotis <93345790+brekotis@users.noreply.github.com>
2026-07-10 16:35:39 +03:00
Alexey 2ac93c6d49 Update Cargo.lock 2026-07-10 11:47:09 +03:00
Alexey a51e58009b Update Cargo.toml 2026-07-10 11:46:33 +03:00
Alexey d523406c0a Merge pull request #869 from telemt/flow
Regression coverage + Synlimit netfilter rules per-target + CidrRateLimitKey with IpNetwork + Telemt-Referenz der Konfigurationsparameter
2026-07-10 11:39:22 +03:00
Alexey 5b3ad0096b Create CONFIG_PARAMS.de.md 2026-07-09 22:42:07 +03:00
Alexey b587fdbf94 Update CONFIG_PARAMS.ru.md 2026-07-08 20:52:26 +03:00
Alexey 77a45e509a Update CONFIG_PARAMS.en.md 2026-07-07 22:46:03 +03:00
Alexey b8be805aed Rustfmt 2026-07-06 20:12:36 +03:00
Alexey a1ebd44cee Update load_basic_tests.rs 2026-07-05 19:58:37 +03:00
Alexey 25d02a8e0e Update traffic_limiter.rs 2026-07-04 15:45:51 +03:00
Alexey 3375017460 Add the new key to the hot-reload snapshot type 2026-07-04 13:11:08 +03:00
Alexey 25e0abae8a Validate duplicate normalized auto-templates 2026-07-04 12:35:22 +03:00
Alexey 50538d234e CidrRateLimitKey with IpNetwork parsing and serialization added 2026-07-04 10:54:55 +03:00
Alexey e3a7be6786 Update CONTRIBUTING.md 2026-07-03 23:39:46 +03:00
Alexey 3fc2877205 Reset ports configuration in merge docker-compose file: merge pull request #866 from dvrass/patch-1
Reset ports configuration in merge docker-compose file
2026-07-03 23:34:04 +03:00
Alexey cd1dc2f4c9 Update docker-compose.host-netfilter.yml 2026-07-03 23:32:30 +03:00
Alexey 451227da60 Namespace synlimit netfilter rules per target set
Co-Authored-By: brekotis <93345790+brekotis@users.noreply.github.com>
2026-07-01 01:47:14 +03:00
dvrass f55d8479e3 Reset ports configuration in merge docker-compose file
Reset ports configuration in docker-compose file for network_mode:host
2026-06-30 13:23:56 +03:00
Alexey 81ae483201 Add regression coverage for ME routing, D2C padding, synlimit, and MSS bulk validation
Co-Authored-By: brekotis <93345790+brekotis@users.noreply.github.com>
2026-06-30 13:13:11 +03:00
Alexey ed1895d6df Merge pull request #865 from telemt/flow
Secure + VersionD Outbound Paddings Fix
2026-06-29 17:37:13 +03:00
Alexey 88d161a5e9 Bump -> 3.4.22 2026-06-29 16:38:09 +03:00
Alexey a0ac108807 Secure + VersionD Outbound Paddings Fix 2026-06-29 13:56:16 +03:00
Alexey 809352fac5 Merge pull request #864 from telemt/flow
Restore ME writer source IP for initial proxy request binding
2026-06-29 12:51:04 +03:00
Alexey 22627b498d Bump -> 3.4.21 2026-06-29 12:44:03 +03:00
Alexey b9c5c71dbc Restore ME writer source IP for initial proxy request binding 2026-06-29 12:37:31 +03:00
Alexey 7aee991416 Synlimit V2 + Split config loader helpers into modules + Restore ME D2C queue-drain flushing: merge pull request #863 from telemt/flow
Synlimit V2 + Split config loader helpers into modules + Restore ME D2C queue-drain flushing
2026-06-29 11:44:29 +03:00
Alexey 9a9fd3f55d Update d2c.rs 2026-06-29 11:37:30 +03:00
Alexey 3a5fe31262 Rustfmt 2026-06-29 11:01:36 +03:00
Alexey 82f63d0d8a Split config loader helpers into focused modules 2026-06-28 17:25:03 +03:00
Alexey fce75163b0 Rustfmt 2026-06-28 15:45:53 +03:00
Alexey fe56621a83 Delete synlimit_control.rs 2026-06-28 15:43:26 +03:00
Alexey 1f2910f5bc Update crypto_bench.rs 2026-06-28 15:42:53 +03:00
Alexey d67e7c5a6f Update mod.rs 2026-06-28 15:29:54 +03:00
Alexey 558f352a57 Synlimit V2 2026-06-28 12:53:28 +03:00
Alexey 1ee9a234d7 Merge pull request #859 from CryZFix/main
fixed a docker build bug for CI or docker-compose build variations
2026-06-25 19:24:08 +03:00
CryZFix 2e13f89f6d fixed a docker build bug for CI or docker-compose build variations 2026-06-25 16:24:42 +04:00
103 changed files with 16589 additions and 5292 deletions
+49
View File
@@ -0,0 +1,49 @@
name: Coverage
on:
pull_request:
branches: [ "*" ]
workflow_dispatch:
env:
CARGO_TERM_COLOR: always
jobs:
coverage:
name: LLVM coverage report
runs-on: ubuntu-latest
permissions:
contents: read
steps:
- uses: actions/checkout@v4
- uses: dtolnay/rust-toolchain@stable
with:
components: llvm-tools-preview
- name: Cache cargo
uses: actions/cache@v4
with:
path: |
~/.cargo/bin
~/.cargo/registry
~/.cargo/git
target
key: ${{ runner.os }}-cargo-llvm-cov-${{ hashFiles('**/Cargo.lock') }}
restore-keys: |
${{ runner.os }}-cargo-llvm-cov-
${{ runner.os }}-cargo-
- name: Install cargo-llvm-cov
run: cargo install --locked cargo-llvm-cov || true
- name: Generate LCOV report
run: cargo llvm-cov --locked --lcov --output-path lcov.info
- name: Upload LCOV report
uses: actions/upload-artifact@v4
with:
name: telemt-lcov
path: lcov.info
+13 -11
View File
@@ -1,18 +1,20 @@
# Issues # Issues
## Warnung ### Warnung
Before opening Issue, if it is more question than problem or bug - ask about that [in our chat](https://t.me/telemtrs) Before opening Issue, if it is more question than problem or bug - ask about that [in our chat](https://t.me/telemtrs)
## What it is not
- NOT Question and Answer
- NOT Helpdesk
***Each of your Issues triggers attempts to reproduce problems and analyze them, which are done manually by people*** ***Each of your Issues triggers attempts to reproduce problems and analyze them, which are done manually by people***
Issues is **NOT** about:
- Question and Answer
- Helpdesk
- Configuration or Intergraion Support
--- ---
# Pull Requests # Pull Requests
## General ### General
- ONLY signed and verified commits - ONLY signed and verified commits
- ONLY from your name - ONLY from your name
- DO NOT commit with `codex`, `claude`, or other AI tools as author/committer - DO NOT commit with `codex`, `claude`, or other AI tools as author/committer
@@ -20,7 +22,7 @@ Before opening Issue, if it is more question than problem or bug - ask about tha
--- ---
## Definition of Ready (MANDATORY) ### Definition of Ready (MANDATORY)
A Pull Request WILL be ignored or closed if: A Pull Request WILL be ignored or closed if:
@@ -32,14 +34,14 @@ A Pull Request WILL be ignored or closed if:
--- ---
## Blessed Principles ### Blessed Principles
- PR must build - PR must build
- PR must pass tests - PR must pass tests
- PR must be understood by author - PR must be understood by author
--- ---
## AI Usage Policy ### AI Usage Policy
AI tools (Claude, ChatGPT, Codex, DeepSeek, etc.) are allowed as **assistants**, NOT as decision-makers. AI tools (Claude, ChatGPT, Codex, DeepSeek, etc.) are allowed as **assistants**, NOT as decision-makers.
@@ -60,7 +62,7 @@ PRs that look like unverified AI dumps WILL be closed
--- ---
## Maintainer Policy ### Maintainer Policy
Maintainers reserve the right to: Maintainers reserve the right to:
@@ -72,7 +74,7 @@ Respect the reviewers time
--- ---
## Enforcement ### Enforcement
Pull Requests that violate project standards may be closed without review. Pull Requests that violate project standards may be closed without review.
Generated
+2 -1
View File
@@ -21,6 +21,7 @@ dependencies = [
"cfg-if", "cfg-if",
"cipher", "cipher",
"cpufeatures 0.2.17", "cpufeatures 0.2.17",
"zeroize",
] ]
[[package]] [[package]]
@@ -2899,7 +2900,7 @@ checksum = "7b2093cf4c8eb1e67749a6762251bc9cd836b6fc171623bd0a9d324d37af2417"
[[package]] [[package]]
name = "telemt" name = "telemt"
version = "3.4.19" version = "3.4.25"
dependencies = [ dependencies = [
"aes", "aes",
"anyhow", "anyhow",
+3 -3
View File
@@ -1,6 +1,6 @@
[package] [package]
name = "telemt" name = "telemt"
version = "3.4.19" version = "3.4.25"
edition = "2024" edition = "2024"
[features] [features]
@@ -15,8 +15,8 @@ tokio = { version = "1.52.3", features = ["full", "tracing"] }
tokio-util = { version = "0.7.18", features = ["full"] } tokio-util = { version = "0.7.18", features = ["full"] }
# Crypto # Crypto
aes = "0.8.4" aes = { version = "0.8.4", features = ["zeroize"] }
ctr = "0.9.2" ctr = { version = "0.9.2", features = ["zeroize"] }
cbc = "0.1.2" cbc = "0.1.2"
sha2 = "0.10.9" sha2 = "0.10.9"
sha1 = "0.10.6" sha1 = "0.10.6"
+12 -2
View File
@@ -55,6 +55,16 @@ RUN set -eux; \
strip --strip-unneeded /telemt || true; \ strip --strip-unneeded /telemt || true; \
rm -f "/tmp/${ASSET}" "/tmp/${ASSET}.sha256" /tmp/telemt rm -f "/tmp/${ASSET}" "/tmp/${ASSET}.sha256" /tmp/telemt
RUN --mount=type=bind,target=/tmp \
mkdir -p /app && \
if [ -f /tmp/config.toml ]; then \
cp /tmp/config.toml /app/config.toml; \
elif [ -f /tmp/config/config.toml ]; then \
cp /tmp/config/config.toml /app/config.toml; \
else \
echo "Config file not found" && exit 1; \
fi
# ========================== # ==========================
# Debug Image # Debug Image
# ========================== # ==========================
@@ -99,7 +109,7 @@ RUN set -eux; \
WORKDIR /app WORKDIR /app
COPY --from=minimal /telemt /app/telemt COPY --from=minimal /telemt /app/telemt
COPY ./config/config.toml /app/config.toml COPY --from=minimal /app/config.toml /app/config.toml
EXPOSE 443 9090 9091 EXPOSE 443 9090 9091
@@ -116,7 +126,7 @@ FROM gcr.io/distroless/static-debian12 AS prod
WORKDIR /app WORKDIR /app
COPY --from=minimal /telemt /app/telemt COPY --from=minimal /telemt /app/telemt
COPY ./config/config.toml /app/config.toml COPY --from=minimal /app/config.toml /app/config.toml
USER nonroot:nonroot USER nonroot:nonroot
+15 -3
View File
@@ -1,12 +1,24 @@
// Cryptobench use criterion::{Criterion, criterion_group, criterion_main};
use criterion::{Criterion, black_box, criterion_group}; use std::hint::black_box;
#[allow(unused_imports)]
#[path = "../src/crypto/aes.rs"]
mod aes_impl;
#[allow(unused_imports)]
#[path = "../src/error.rs"]
mod error;
use aes_impl::AesCtr;
fn bench_aes_ctr(c: &mut Criterion) { fn bench_aes_ctr(c: &mut Criterion) {
c.bench_function("aes_ctr_encrypt_64kb", |b| { c.bench_function("aes_ctr_encrypt_64kb", |b| {
let data = vec![0u8; 65536]; let data = vec![0u8; 65536];
b.iter(|| { b.iter(|| {
let mut enc = AesCtr::new(&[0u8; 32], 0); let mut enc = AesCtr::new(&[0u8; 32], 0);
black_box(enc.encrypt(&data)) black_box(enc.encrypt(black_box(data.as_slice())))
}) })
}); });
} }
criterion_group!(benches, bench_aes_ctr);
criterion_main!(benches);
+1 -2
View File
@@ -4,7 +4,6 @@ services:
context: . context: .
target: prod-netfilter target: prod-netfilter
network_mode: host network_mode: host
ports: [] ports: !reset []
cap_add: cap_add:
- NET_BIND_SERVICE
- NET_ADMIN - NET_ADMIN
+106 -9
View File
@@ -107,7 +107,9 @@ Notes:
| `GET` | `/v1/stats/users/active-ips` | none | `200` | `UserActiveIps[]` | | `GET` | `/v1/stats/users/active-ips` | none | `200` | `UserActiveIps[]` |
| `GET` | `/v1/stats/users` | none | `200` | `UserInfo[]` | | `GET` | `/v1/stats/users` | none | `200` | `UserInfo[]` |
| `GET` | `/v1/config` | none | `200` | `ConfigData` | | `GET` | `/v1/config` | none | `200` | `ConfigData` |
| `PATCH` | `/v1/config` | sparse JSON object | `200` | `PatchConfigResponse` | | `PATCH` | `/v1/config` | sparse JSON object; optional reload query | `200` or `202` | `PatchConfigResponse` |
| `POST` | `/v1/system/reload` | `ReloadRequest` or empty body | `202` | `ReloadAccepted` |
| `GET` | `/v1/system/reload/{id}` | none | `200` | `ReloadStatus` |
| `GET` | `/v1/users` | none | `200` | `UserInfo[]` | | `GET` | `/v1/users` | none | `200` | `UserInfo[]` |
| `POST` | `/v1/users` | `CreateUserRequest` | `201` or `202` | `CreateUserResponse` | | `POST` | `/v1/users` | `CreateUserRequest` | `201` or `202` | `CreateUserResponse` |
| `GET` | `/v1/users/{username}` | none | `200` | `UserInfo` | | `GET` | `/v1/users/{username}` | none | `200` | `UserInfo` |
@@ -146,7 +148,9 @@ Notes:
| `GET /v1/stats/users/active-ips` | Returns users that currently have non-empty active source-IP lists. | | `GET /v1/stats/users/active-ips` | Returns users that currently have non-empty active source-IP lists. |
| `GET /v1/stats/users` | Alias of `GET /v1/users`; returns disk-first user views with runtime lag flag. | | `GET /v1/stats/users` | Alias of `GET /v1/users`; returns disk-first user views with runtime lag flag. |
| `GET /v1/config` | Returns the current editable config sections as JSON (no `access.*`) plus the revision. | | `GET /v1/config` | Returns the current editable config sections as JSON (no `access.*`) plus the revision. |
| `PATCH /v1/config` | Applies a sparse patch to editable config sections; validates, writes, and reports restart impact. | | `PATCH /v1/config` | Applies a sparse patch and optionally submits an in-process runtime reload to Maestro. |
| `POST /v1/system/reload` | Loads and validates the current on-disk config, then asks Maestro to prepare and activate a new runtime generation. |
| `GET /v1/system/reload/{id}` | Returns one retained reload status; the coordinator retains the most recent 32 operations. |
| `GET /v1/users` | Returns disk-first user views sorted by username. | | `GET /v1/users` | Returns disk-first user views sorted by username. |
| `POST /v1/users` | Creates a user and returns the effective user view plus secret. | | `POST /v1/users` | Creates a user and returns the effective user view plus secret. |
| `GET /v1/users/{username}` | Returns one disk-first user view or `404` when absent. | | `GET /v1/users/{username}` | Returns one disk-first user view or `404` when absent. |
@@ -170,11 +174,13 @@ Notes:
| `404` | `not_found` | Unknown route, unknown user, or unsupported sub-route. | | `404` | `not_found` | Unknown route, unknown user, or unsupported sub-route. |
| `405` | `method_not_allowed` | Unsupported method for `/v1/users/{username}` route shape. | | `405` | `method_not_allowed` | Unsupported method for `/v1/users/{username}` route shape. |
| `409` | `revision_conflict` | `If-Match` revision mismatch. | | `409` | `revision_conflict` | `If-Match` revision mismatch. |
| `409` | `reload_in_progress` | Another reload operation is non-terminal. |
| `409` | `user_exists` | User already exists on create. | | `409` | `user_exists` | User already exists on create. |
| `409` | `last_user_forbidden` | Attempt to delete last configured user. | | `409` | `last_user_forbidden` | Attempt to delete last configured user. |
| `413` | `payload_too_large` | Body exceeds `request_body_limit_bytes`. | | `413` | `payload_too_large` | Body exceeds `request_body_limit_bytes`. |
| `500` | `internal_error` | Internal error (I/O, serialization, config load/save). | | `500` | `internal_error` | Internal error (I/O, serialization, config load/save). |
| `503` | `api_disabled` | API disabled in config. | | `503` | `api_disabled` | API disabled in config. |
| `503` | `maestro_unavailable` | Maestro's reload command channel is unavailable. |
## Routing and Method Edge Cases ## Routing and Method Edge Cases
@@ -292,13 +298,17 @@ Sections absent from the config file are absent from the response (not `null`).
### `PatchConfigResponse` ### `PatchConfigResponse`
Returned by `PATCH /v1/config` on success (`200`). Returned by `PATCH /v1/config` on success (`200`, or `202` when a reload was accepted).
| Field | Type | Description | | Field | Type | Description |
| --- | --- | --- | | --- | --- | --- |
| `revision` | `string` | SHA-256 hex of the config file after the patch was written. | | `revision` | `string` | SHA-256 hex of the config file after the patch was written. |
| `restart_required` | `bool` | `true` when one or more changed fields require a process restart to take effect. Hot-reloadable fields (e.g. `general.log_level`) are applied automatically by the config file watcher; restart-required fields (e.g. any `censorship.*`, `timeouts.*`, `upstreams`, or `general.modes` change) are written to disk but only take effect after the Telemt process is restarted. The caller is responsible for triggering a restart when this flag is `true`. | | `restart_required` | `bool` | Legacy classifier result: `true` when the old file watcher alone cannot apply every changed field. Use `runtime_reload_required` and `process_restart_required` for new integrations. |
| `runtime_reload_required` | `bool` | `true` when full effect requires a Maestro runtime-generation reload rather than the legacy hot-field overlay. |
| `process_restart_required` | `bool` | `true` when process-owned sockets or paths changed and remain deferred after an in-process reload. |
| `deferred_process_fields` | `string[]` | Process-owned fields that the active process cannot rebind during generation activation. |
| `changed` | `string[]` | Top-level section names that differed between the old and new config (e.g. `["censorship"]`). | | `changed` | `string[]` | Top-level section names that differed between the old and new config (e.g. `["censorship"]`). |
| `reload` | `ReloadAccepted?` | Present only when the patch included a valid reload query and Maestro accepted the operation. |
### `HealthData` ### `HealthData`
| Field | Type | Description | | Field | Type | Description |
@@ -324,6 +334,7 @@ Returned by `PATCH /v1/config` on success (`200`).
| `connections_bad_total` | `u64` | Failed/invalid client connections. | | `connections_bad_total` | `u64` | Failed/invalid client connections. |
| `connections_bad_by_class` | `ClassCount[]` | Failed/invalid connections grouped by class. | | `connections_bad_by_class` | `ClassCount[]` | Failed/invalid connections grouped by class. |
| `handshake_failures_by_class` | `ClassCount[]` | Handshake failures grouped by class. | | `handshake_failures_by_class` | `ClassCount[]` | Handshake failures grouped by class. |
| `handshake_failures_by_stage` | `StageCount[]` | Handshake failures grouped by state-machine stage. |
| `handshake_timeouts_total` | `u64` | Handshake timeout count. | | `handshake_timeouts_total` | `u64` | Handshake timeout count. |
| `configured_users` | `usize` | Number of configured users in config. | | `configured_users` | `usize` | Number of configured users in config. |
@@ -333,6 +344,38 @@ Returned by `PATCH /v1/config` on success (`200`).
| `class` | `string` | Failure class label. | | `class` | `string` | Failure class label. |
| `total` | `u64` | Counter value for this class. | | `total` | `u64` | Counter value for this class. |
#### `StageCount`
| Field | Type | Description |
| --- | --- | --- |
| `stage` | `string` | State-machine stage label. |
| `total` | `u64` | Counter value for this stage. |
#### Handshake failure stage diagnostics
`handshake_failures_by_class` and `telemt_handshake_failures_by_class_total` describe the error kind. `handshake_failures_by_stage` and `telemt_handshake_failures_by_stage_total` describe where the same failure happened in the handshake state machine.
This does not add a DPI verdict or any protocol decision. The stage is derived from the existing Telemt handshake control flow and is counted only when the existing handshake failure or timeout accounting path is reached.
Fixed stage labels:
| Stage | Meaning |
| --- | --- |
| `first_packet_prelude` | Reading the first 5 bytes before selecting the TLS or direct branch. |
| `tls_clienthello_body` | Reading the TLS ClientHello body after the TLS record header. |
| `tls_core` | Running the TLS-F handshake/auth flow. |
| `tls_post_serverhello_mtproto` | Waiting for the 64-byte MTProto handshake after TLS ServerHello. |
| `direct_mtproto` | Reading the direct classic/secure 64-byte MTProto handshake. |
Example:
```text
telemt_handshake_failures_by_class_total{class="expected_64_got_0_unexpected_eof"} 3
telemt_handshake_failures_by_stage_total{stage="direct_mtproto"} 1
telemt_handshake_failures_by_stage_total{stage="tls_post_serverhello_mtproto"} 2
```
This means the same EOF-while-reading-64-bytes failure happened once in the direct MTProto path and twice after TLS ServerHello.
### `SystemInfoData` ### `SystemInfoData`
| Field | Type | Description | | Field | Type | Description |
| --- | --- | --- | | --- | --- | --- |
@@ -917,6 +960,7 @@ JA3 follows the Salesforce ClientHello field order. JA4 follows the FoxIO TLS-cl
| `connections_bad_total` | `u64` | Failed/invalid connections. | | `connections_bad_total` | `u64` | Failed/invalid connections. |
| `connections_bad_by_class` | `ClassCount[]` | Failed/invalid connections grouped by class. | | `connections_bad_by_class` | `ClassCount[]` | Failed/invalid connections grouped by class. |
| `handshake_failures_by_class` | `ClassCount[]` | Handshake failures grouped by class. | | `handshake_failures_by_class` | `ClassCount[]` | Handshake failures grouped by class. |
| `handshake_failures_by_stage` | `StageCount[]` | Handshake failures grouped by state-machine stage. |
| `handshake_timeouts_total` | `u64` | Handshake timeouts. | | `handshake_timeouts_total` | `u64` | Handshake timeouts. |
| `accept_permit_timeout_total` | `u64` | Listener admission permit acquisition timeouts. | | `accept_permit_timeout_total` | `u64` | Listener admission permit acquisition timeouts. |
| `configured_users` | `usize` | Configured user count. | | `configured_users` | `usize` | Configured user count. |
@@ -1376,24 +1420,50 @@ Applies a sparse patch to the editable config sections. The merged config is ful
**Read-only mode:** returns `403 read_only` when the API runs with `read_only = true`. **Read-only mode:** returns `403 read_only` when the API runs with `read_only = true`.
**Success `200` response body** (`data` field of the standard envelope): **Optional in-process reload query:**
| Query | Required | Description |
| --- | --- | --- |
| `reload=instant` | no | Activates a new generation and cancels sessions owned by the previous generation. |
| `reload=drain` | no | Activates a new generation and lets old sessions finish until `timeout_secs`. |
| `timeout_secs=1..3600` | for `reload=drain` | Bounded old-generation drain interval. Invalid with `reload=instant`. |
| `failure_policy=keep_new\|rollback` | no | Defaults to `keep_new`. `rollback` applies only through the activation barrier, before old-generation teardown. |
Without a `reload` query parameter, the endpoint preserves the legacy behavior: it writes the patch and the file watcher applies only supported hot fields.
**Success `200` or `202` response body** (`data` field of the standard envelope):
```json ```json
{ {
"revision": "<new-sha256-hex>", "revision": "<new-sha256-hex>",
"restart_required": true, "restart_required": true,
"changed": ["censorship"] "runtime_reload_required": true,
"process_restart_required": false,
"deferred_process_fields": [],
"changed": ["censorship"],
"reload": {
"reload_id": 7,
"target_generation": 2,
"config_revision": "<new-sha256-hex>",
"state": "accepted",
"mode": "instant",
"failure_policy": "keep_new"
}
} }
``` ```
- `revision` — SHA-256 hex of the config file after the write. - `revision` — SHA-256 hex of the config file after the write.
- `restart_required``true` when the change affects a field that Telemt cannot hot-reload (e.g. `censorship.*`, `timeouts.*`, `upstreams`, `general.modes`). Hot-reloadable fields (e.g. `general.log_level`) are applied automatically by the config file watcher. Restart-required fields are written to disk but only take effect after the Telemt process is restarted; the caller is responsible for triggering the restart. - `restart_required`legacy file-watcher classification retained for compatibility.
- `runtime_reload_required` — reports whether a full Maestro generation reload is needed for runtime effect.
- `process_restart_required` and `deferred_process_fields` — report process-owned sockets or paths that remain unchanged by an in-process reload.
- `changed` — list of top-level section names that differed. - `changed` — list of top-level section names that differed.
- `reload` — accepted operation metadata; omitted when no reload query was supplied.
**Status codes:** **Status codes:**
| HTTP | `error.code` | Condition | | HTTP | `error.code` | Condition |
| --- | --- | --- | | --- | --- | --- |
| `200` | — | Patch applied successfully. | | `200` | — | Patch applied successfully. |
| `202` | — | Patch applied and runtime reload accepted. |
| `400` | `bad_request` | Invalid JSON, empty patch, or config validation/deserialization failure. | | `400` | `bad_request` | Invalid JSON, empty patch, or config validation/deserialization failure. |
| `400` | `access_not_editable` | Patch contains an `access` key. | | `400` | `access_not_editable` | Patch contains an `access` key. |
| `400` | `section_not_editable` | Patch contains `server`, `network`, or an unknown top-level key. | | `400` | `section_not_editable` | Patch contains `server`, `network`, or an unknown top-level key. |
@@ -1401,6 +1471,7 @@ Applies a sparse patch to the editable config sections. The merged config is ful
| `403` | `read_only` | API is in read-only mode. | | `403` | `read_only` | API is in read-only mode. |
| `405` | `method_not_allowed` | Method other than `GET` or `PATCH` used on `/v1/config`. | | `405` | `method_not_allowed` | Method other than `GET` or `PATCH` used on `/v1/config`. |
| `409` | `revision_conflict` | `If-Match` header supplied but does not match current revision. | | `409` | `revision_conflict` | `If-Match` header supplied but does not match current revision. |
| `409` | `reload_in_progress` | Another runtime reload is active; the patch is not written. |
| `500` | `internal_error` | I/O or serialization failure. | | `500` | `internal_error` | I/O or serialization failure. |
**curl example:** **curl example:**
@@ -1412,14 +1483,40 @@ curl -s -H "Authorization: <token>" http://127.0.0.1:<api>/v1/system/info | jq -
curl -s -X PATCH -H "Authorization: <token>" -H "If-Match: <revision>" \ curl -s -X PATCH -H "Authorization: <token>" -H "If-Match: <revision>" \
-H "Content-Type: application/json" \ -H "Content-Type: application/json" \
-d '{"censorship":{"tls_domain":"front.example.com"}}' \ -d '{"censorship":{"tls_domain":"front.example.com"}}' \
http://127.0.0.1:<api>/v1/config 'http://127.0.0.1:<api>/v1/config?reload=instant'
``` ```
## Runtime Reload Endpoints
### `POST /v1/system/reload`
Loads the current on-disk config under the API mutation lock and submits an immutable config snapshot to Maestro. `If-Match` is optional and uses the same revision contract as `PATCH /v1/config`. An empty body defaults to `{"mode":"instant","failure_policy":"keep_new"}`.
```json
{
"mode": "drain",
"timeout_secs": 30,
"failure_policy": "rollback"
}
```
The endpoint returns `202` with `ReloadAccepted`. A concurrent non-terminal reload returns `409 reload_in_progress`. Config parsing or validation failure is reported before a command is submitted.
### `GET /v1/system/reload/{id}`
Returns `ReloadStatus` with `state` equal to `accepted`, `preparing`, `activating`, `draining`, `succeeded`, `rolled_back`, or `failed`. Terminal statuses include `finished_at_epoch_secs`; failures include `error`. Successful activation may include `warnings` for old-generation cleanup failures and `deferred_process_fields` for process-owned settings.
Runtime generation activation rebuilds statistics, upstream routing, replay and buffer state, TLS-front cache, IP tracking, admission/route state, and Middle-End orchestration. Per-user quota accounting is process-scoped and remains continuous across generations. API, metrics, client TCP/Unix listeners, PID ownership, and logging remain process-scoped; changed bind/path fields are reported as deferred and do not cause Maestro to invoke systemd, containerd, or another process supervisor.
Reload preparation requires every configured TLS-front domain to have a non-default cached profile and requires a ready Middle-End pool when direct fallback is disabled. A candidate that does not satisfy either readiness condition fails without replacing the active generation.
The revision is verified again after preparation. With `failure_policy=rollback`, a changed revision or revision read failure rolls the candidate back; with `failure_policy=keep_new`, the condition is reported in `warnings` and activation continues.
## Mutation Semantics ## Mutation Semantics
| Endpoint | Notes | | Endpoint | Notes |
| --- | --- | | --- | --- |
| `PATCH /v1/config` | Deep-merges the patch into editable config sections (tables merged per-field; arrays/scalars replaced wholesale). Validates the merged result before writing. Writes only the touched sections via atomic `tmp + rename`. Returns the new revision and which sections changed. | | `PATCH /v1/config` | Deep-merges and validates the patch, writes touched sections via atomic `tmp + rename`, and optionally submits the exact written revision for an in-process Maestro reload. |
| `POST /v1/users` | Creates user, validates config, then atomically updates only affected `access.*` TOML tables (`access.users` always, plus optional per-user tables present in request). | | `POST /v1/users` | Creates user, validates config, then atomically updates only affected `access.*` TOML tables (`access.users` always, plus optional per-user tables present in request). |
| `PATCH /v1/users/{username}` | Partial update of provided fields only. Missing fields remain unchanged; explicit `null` removes optional per-user entries. The write path updates only affected `access.*` TOML tables. | | `PATCH /v1/users/{username}` | Partial update of provided fields only. Missing fields remain unchanged; explicit `null` removes optional per-user entries. The write path updates only affected `access.*` TOML tables. |
| `POST /v1/users/{username}/rotate-secret` | Replaces the user's secret with a provided valid 32-hex value or a generated value, then returns the effective secret in `CreateUserResponse`. | | `POST /v1/users/{username}/rotate-secret` | Replaces the user's secret with a provided valid 32-hex value or a generated value, then returns the effective secret in `CreateUserResponse`. |
File diff suppressed because it is too large Load Diff
+85 -16
View File
@@ -1972,6 +1972,7 @@ This document lists all configuration keys accepted by `config.toml`.
## client_mss ## client_mss
- **Constraints / validation**: `String`. Empty or omitted means do not change kernel MSS. Presets: `"extreme-low"` = `88`, `"tspu"` = `92`, `"2in8"` = `256`. Custom decimal strings must be within `88..=4096`. - **Constraints / validation**: `String`. Empty or omitted means do not change kernel MSS. Presets: `"extreme-low"` = `88`, `"tspu"` = `92`, `"2in8"` = `256`. Custom decimal strings must be within `88..=4096`.
- **Description**: Client-facing TCP MSS applied to TCP listener sockets before `listen(2)`, so Linux can announce it in SYN/ACK. This affects only proxy client TCP listeners, not API, metrics, Unix sockets, Telegram upstreams, ME sockets, or mask backend connections. Changes require listener restart/rebind. - **Description**: Client-facing TCP MSS applied to TCP listener sockets before `listen(2)`, so Linux can announce it in SYN/ACK. This affects only proxy client TCP listeners, not API, metrics, Unix sockets, Telegram upstreams, ME sockets, or mask backend connections. Changes require listener restart/rebind.
- **Operator note**: The two-tier `synlimit` profile does not require Telemt to disable MSS automatically. Operators that follow external host-tuning recipes should decide explicitly whether to leave MSS shaping enabled for handshake fragmentation or disable it for higher media throughput.
- **Performance note**: Low MSS increases packet count predictably. Approximate segment multiplier is `ceil(1460 / client_mss)`. - **Performance note**: Low MSS increases packet count predictably. Approximate segment multiplier is `ceil(1460 / client_mss)`.
- **Example**: - **Example**:
@@ -2311,9 +2312,14 @@ Note: This section also accepts the legacy alias `[server.admin_api]` (same sche
| [`port`](#port-serverlisteners) | `u16` | `server.port` | `✘` | | [`port`](#port-serverlisteners) | `u16` | `server.port` | `✘` |
| [`client_mss`](#client_mss-serverlisteners) | `String` | `[server].client_mss` | `✘` | | [`client_mss`](#client_mss-serverlisteners) | `String` | `[server].client_mss` | `✘` |
| [`synlimit`](#synlimit-serverlisteners) | `false`, `"iptables"`, or `"nftables"` | `false` | `✔` | | [`synlimit`](#synlimit-serverlisteners) | `false`, `"iptables"`, or `"nftables"` | `false` | `✔` |
| [`synlimit_seconds`](#synlimit_seconds-serverlisteners) | `u32` | `1` | `✔` | | [`synlimit_seconds`](#synlimit_seconds-serverlisteners) | `u32` | `60` | `✔` |
| [`synlimit_hitcount`](#synlimit_hitcount-serverlisteners) | `u32` | `1` | `✔` | | [`synlimit_hitcount`](#synlimit_hitcount-serverlisteners) | `u32` | `48` | `✔` |
| [`synlimit_burst`](#synlimit_burst-serverlisteners) | `u32` | `2` | `✔` | | [`synlimit_burst`](#synlimit_burst-serverlisteners) | `u32` | `1` | `✔` |
| [`synlimit_ios_seconds`](#synlimit_ios_seconds-serverlisteners) | `u32` | `1` | `✔` |
| [`synlimit_ios_hitcount`](#synlimit_ios_hitcount-serverlisteners) | `u32` | `12` | `✔` |
| [`synlimit_ios_burst`](#synlimit_ios_burst-serverlisteners) | `u32` | `24` | `✔` |
| [`synlimit_hashlimit_expire_ms`](#synlimit_hashlimit_expire_ms-serverlisteners) | `u32` | `60000` | `✔` |
| [`synlimit_hashlimit_size`](#synlimit_hashlimit_size-serverlisteners) | `u32` | `32768` | `✔` |
| [`announce`](#announce) | `String` | — | `✘` | | [`announce`](#announce) | `String` | — | `✘` |
| [`announce_ip`](#announce_ip) | `IpAddr` | — | `✘` | | [`announce_ip`](#announce_ip) | `IpAddr` | — | `✘` |
| [`proxy_protocol`](#proxy_protocol) | `bool` | — | `✘` | | [`proxy_protocol`](#proxy_protocol) | `bool` | — | `✘` |
@@ -2351,7 +2357,8 @@ Note: This section also accepts the legacy alias `[server.admin_api]` (same sche
``` ```
## synlimit (server.listeners) ## synlimit (server.listeners)
- **Constraints / validation**: `false`, `"iptables"`, or `"nftables"`. Omitted or `false` disables SYN limiting for this listener. - **Constraints / validation**: `false`, `"iptables"`, or `"nftables"`. Omitted or `false` disables SYN limiting for this listener.
- **Description**: Installs per-listener Linux netfilter SYN limiter rules for the listener port. `"iptables"` uses `iptables`/`ip6tables` filter rules with the `hashlimit` match as a per-source token bucket. `"nftables"` uses per-source `meter` rules with `limit rate over` and auto-detects whether the host already uses `inet`, `ip`, or `ip6` table families before creating Telemt-owned tables. The token-bucket rate is `synlimit_hitcount / synlimit_seconds`; `synlimit_burst` controls the burst size. Rules are reconciled at runtime and removed during graceful Telemt shutdown; `SIGKILL` cannot be cleaned up by the process. Requires CAP_NET_ADMIN. `synlimit*` changes hot-reload for existing listener endpoints; changing listener `ip` or `port` still requires restart/rebind. - **Description**: Installs per-listener Linux netfilter two-tier SYN-fix rules for the listener port. `"iptables"` uses `iptables`/`ip6tables` filter rules with the `hashlimit`, `length`, and TTL/hop-limit matches. `"nftables"` uses Telemt-owned tables with per-source `meter` rules and equivalent IPv4/IPv6 classifiers. Rules are inserted early in `INPUT`, accept under-limit SYN packets, and reject over-limit SYN packets with TCP RST so clients retry promptly instead of waiting for a silent DROP timeout. The generic bucket is controlled by `synlimit_seconds`, `synlimit_hitcount`, and `synlimit_burst`; the iOS-like TTL/length bucket is controlled by `synlimit_ios_*`. Rules are reconciled at runtime and removed during graceful Telemt shutdown; `SIGKILL` cannot be cleaned up by the process. Requires CAP_NET_ADMIN. `synlimit*` changes hot-reload for existing listener endpoints; changing listener `ip` or `port` still requires restart/rebind.
- **Operator note**: Telemt does not persist rules with `iptables-persistent`, write `/etc/sysctl.d`, edit systemd limits, or modify `client_mss`. Apply host-level tuning manually if your deployment policy requires it.
- **Example**: - **Example**:
```toml ```toml
@@ -2366,8 +2373,8 @@ Note: This section also accepts the legacy alias `[server.admin_api]` (same sche
synlimit = "nftables" synlimit = "nftables"
``` ```
## synlimit_seconds (server.listeners) ## synlimit_seconds (server.listeners)
- **Constraints / validation**: `u32`, must be `> 0`. Default is `1`. - **Constraints / validation**: `u32`, must be `> 0`. Default is `60`.
- **Description**: Token-bucket interval for both SYN limiter backends. The rate is `synlimit_hitcount / synlimit_seconds` and is rendered to native netfilter rate units (`second`, `minute`, `hour`, or `day`). - **Description**: Generic SYN-fix token-bucket interval. The rate is `synlimit_hitcount / synlimit_seconds` and is rendered to native netfilter rate units (`second`, `minute`, `hour`, or `day`). This bucket handles SYN packets that do not match the iOS-like TTL/length classifier.
- **Example**: - **Example**:
```toml ```toml
@@ -2375,11 +2382,11 @@ Note: This section also accepts the legacy alias `[server.admin_api]` (same sche
ip = "0.0.0.0" ip = "0.0.0.0"
port = 443 port = 443
synlimit = "iptables" synlimit = "iptables"
synlimit_seconds = 1 synlimit_seconds = 60
``` ```
## synlimit_hitcount (server.listeners) ## synlimit_hitcount (server.listeners)
- **Constraints / validation**: `u32`, must be `> 0`. Default is `1`. - **Constraints / validation**: `u32`, must be `> 0`. Default is `48`.
- **Description**: Token-bucket rate amount for both SYN limiter backends. Together with `synlimit_seconds`, it defines the allowed source-IP SYN rate before excess SYN packets are dropped. - **Description**: Generic SYN-fix token-bucket rate amount. Together with `synlimit_seconds`, it defines the allowed source-IP SYN rate before excess SYN packets receive TCP RST.
- **Example**: - **Example**:
```toml ```toml
@@ -2387,11 +2394,11 @@ Note: This section also accepts the legacy alias `[server.admin_api]` (same sche
ip = "0.0.0.0" ip = "0.0.0.0"
port = 443 port = 443
synlimit = "iptables" synlimit = "iptables"
synlimit_hitcount = 1 synlimit_hitcount = 48
``` ```
## synlimit_burst (server.listeners) ## synlimit_burst (server.listeners)
- **Constraints / validation**: `u32`, must be `> 0`. Default is `2`. - **Constraints / validation**: `u32`, must be `> 0`. Default is `1`.
- **Description**: Token-bucket burst size for both SYN limiter backends. Higher values allow short connection bursts from the same source IP before the steady-state `synlimit_hitcount / synlimit_seconds` rate is enforced. - **Description**: Generic SYN-fix token-bucket burst size. Higher values allow short connection bursts from the same source IP before the steady-state `synlimit_hitcount / synlimit_seconds` rate is enforced.
- **Example**: - **Example**:
```toml ```toml
@@ -2399,7 +2406,67 @@ Note: This section also accepts the legacy alias `[server.admin_api]` (same sche
ip = "0.0.0.0" ip = "0.0.0.0"
port = 443 port = 443
synlimit = "iptables" synlimit = "iptables"
synlimit_burst = 2 synlimit_burst = 1
```
## synlimit_ios_seconds (server.listeners)
- **Constraints / validation**: `u32`, must be `> 0`. Default is `1`.
- **Description**: Token-bucket interval for SYN packets matching the iOS-like classifier. IPv4 matches packet length `64` and TTL `< 65`; IPv6 matches packet length `84` and hop limit `< 65`.
- **Example**:
```toml
[[server.listeners]]
ip = "0.0.0.0"
port = 443
synlimit = "iptables"
synlimit_ios_seconds = 1
```
## synlimit_ios_hitcount (server.listeners)
- **Constraints / validation**: `u32`, must be `> 0`. Default is `12`.
- **Description**: Token-bucket rate amount for the iOS-like SYN classifier.
- **Example**:
```toml
[[server.listeners]]
ip = "0.0.0.0"
port = 443
synlimit = "iptables"
synlimit_ios_hitcount = 12
```
## synlimit_ios_burst (server.listeners)
- **Constraints / validation**: `u32`, must be `> 0`. Default is `24`.
- **Description**: Token-bucket burst size for the iOS-like SYN classifier.
- **Example**:
```toml
[[server.listeners]]
ip = "0.0.0.0"
port = 443
synlimit = "iptables"
synlimit_ios_burst = 24
```
## synlimit_hashlimit_expire_ms (server.listeners)
- **Constraints / validation**: `u32`, must be `> 0`. Default is `60000`.
- **Description**: Entry expiration in milliseconds for iptables/ip6tables hashlimit buckets. nftables meters use kernel-managed state and do not expose this exact knob.
- **Example**:
```toml
[[server.listeners]]
ip = "0.0.0.0"
port = 443
synlimit = "iptables"
synlimit_hashlimit_expire_ms = 60000
```
## synlimit_hashlimit_size (server.listeners)
- **Constraints / validation**: `u32`, must be `> 0`. Default is `32768`.
- **Description**: Hash table size for iptables/ip6tables hashlimit buckets. nftables meters use kernel-managed state and do not expose this exact knob.
- **Example**:
```toml
[[server.listeners]]
ip = "0.0.0.0"
port = 443
synlimit = "iptables"
synlimit_hashlimit_size = 32768
``` ```
## announce ## announce
- **Constraints / validation**: `String` (optional). Must not be empty when set. - **Constraints / validation**: `String` (optional). Must not be empty when set.
@@ -3119,7 +3186,7 @@ If your backend or network is very bandwidth-constrained, reduce cap first. If p
| [`replay_window_secs`](#replay_window_secs) | `u64` | `120` | `✘` | | [`replay_window_secs`](#replay_window_secs) | `u64` | `120` | `✘` |
| [`ignore_time_skew`](#ignore_time_skew) | `bool` | `false` | `✘` | | [`ignore_time_skew`](#ignore_time_skew) | `bool` | `false` | `✘` |
| [`user_rate_limits`](#user_rate_limits) | `Map<String, RateLimitBps>` | `{}` | `✔` | | [`user_rate_limits`](#user_rate_limits) | `Map<String, RateLimitBps>` | `{}` | `✔` |
| [`cidr_rate_limits`](#cidr_rate_limits) | `Map<IpNetwork, RateLimitBps>` | `{}` | `✔` | | [`cidr_rate_limits`](#cidr_rate_limits) | `Map<CidrRateLimitKey, RateLimitBps>` | `{}` | `✔` |
## users ## users
- **Constraints / validation**: Must not be empty (at least one user must exist). Each value must be **exactly 32 hex characters**. - **Constraints / validation**: Must not be empty (at least one user must exist). Each value must be **exactly 32 hex characters**.
@@ -3282,13 +3349,15 @@ If your backend or network is very bandwidth-constrained, reduce cap first. If p
alice = { up_bps = 1048576, down_bps = 2097152 } alice = { up_bps = 1048576, down_bps = 2097152 }
``` ```
## cidr_rate_limits ## cidr_rate_limits
- **Constraints / validation**: Table `CIDR -> { up_bps, down_bps }`. CIDR must parse as `IpNetwork`; at least one direction must be non-zero. - **Constraints / validation**: Table `CIDR or auto-template -> { up_bps, down_bps }`. Explicit CIDR keys must parse as `IpNetwork`; auto-template keys must be `*4/N` (`N=0..32`), `*6/N` (`N=0..128`), or `*/N` (`N=0..32`). At least one direction must be non-zero. Duplicate normalized auto-templates are rejected.
- **Description**: Source-subnet bandwidth caps applied alongside per-user limits. - **Description**: Source-subnet bandwidth caps applied alongside per-user limits. Explicit CIDR rules use longest-prefix-wins and take priority over auto-templates. Auto-templates create buckets lazily per matched source subnet: `*4/N` for IPv4, `*6/N` for IPv6, and `*/N` as a dual-stack shorthand where IPv4 uses `/N` and IPv6 uses `/(N * 4)`.
- **Example**: - **Example**:
```toml ```toml
[access.cidr_rate_limits] [access.cidr_rate_limits]
"203.0.113.0/24" = { up_bps = 0, down_bps = 1048576 } "203.0.113.0/24" = { up_bps = 0, down_bps = 1048576 }
"*4/32" = { up_bps = 262144, down_bps = 1048576 }
"*6/64" = { up_bps = 262144, down_bps = 1048576 }
``` ```
# [[upstreams]] # [[upstreams]]
+85 -16
View File
@@ -1894,6 +1894,7 @@
## client_mss ## client_mss
- **Ограничения / валидация**: `String`. Пустое значение или отсутствие параметра означает, что Telemt не изменяет MSS, выбранный ядром. Поддерживаемые presets: `"extreme-low"` = `88`, `"tspu"` = `92`, `"2in8"` = `256`. Пользовательское десятичное значение должно быть строкой в диапазоне `88..=4096`. - **Ограничения / валидация**: `String`. Пустое значение или отсутствие параметра означает, что Telemt не изменяет MSS, выбранный ядром. Поддерживаемые presets: `"extreme-low"` = `88`, `"tspu"` = `92`, `"2in8"` = `256`. Пользовательское десятичное значение должно быть строкой в диапазоне `88..=4096`.
- **Описание**: MSS для входящих TCP-соединений клиентов. Значение применяется к TCP listener-сокетам до `listen(2)`, чтобы Linux мог объявить его в SYN/ACK. Параметр влияет только на proxy client TCP listeners и не применяется к API, metrics, Unix sockets, Telegram upstreams, ME sockets или mask backend connections. Изменение требует restart/rebind listener’ов. - **Описание**: MSS для входящих TCP-соединений клиентов. Значение применяется к TCP listener-сокетам до `listen(2)`, чтобы Linux мог объявить его в SYN/ACK. Параметр влияет только на proxy client TCP listeners и не применяется к API, metrics, Unix sockets, Telegram upstreams, ME sockets или mask backend connections. Изменение требует restart/rebind listener’ов.
- **Operator note**: Two-tier `synlimit` profile больше не требует автоматического отключения MSS внутри Telemt. Оператор должен сам решить, оставлять MSS shaping для handshake fragmentation или отключать его ради более высокой скорости media.
- **Performance note**: Низкий MSS предсказуемо увеличивает количество TCP-сегментов. Приблизительный multiplier: `ceil(1460 / client_mss)`. - **Performance note**: Низкий MSS предсказуемо увеличивает количество TCP-сегментов. Приблизительный multiplier: `ceil(1460 / client_mss)`.
- **Пример**: - **Пример**:
@@ -2237,9 +2238,14 @@
| [`port`](#port-serverlisteners) | `u16` | `server.port` | `✘` | | [`port`](#port-serverlisteners) | `u16` | `server.port` | `✘` |
| [`client_mss`](#client_mss-serverlisteners) | `String` | `[server].client_mss` | `✘` | | [`client_mss`](#client_mss-serverlisteners) | `String` | `[server].client_mss` | `✘` |
| [`synlimit`](#synlimit-serverlisteners) | `false`, `"iptables"` или `"nftables"` | `false` | `✔` | | [`synlimit`](#synlimit-serverlisteners) | `false`, `"iptables"` или `"nftables"` | `false` | `✔` |
| [`synlimit_seconds`](#synlimit_seconds-serverlisteners) | `u32` | `1` | `✔` | | [`synlimit_seconds`](#synlimit_seconds-serverlisteners) | `u32` | `60` | `✔` |
| [`synlimit_hitcount`](#synlimit_hitcount-serverlisteners) | `u32` | `1` | `✔` | | [`synlimit_hitcount`](#synlimit_hitcount-serverlisteners) | `u32` | `48` | `✔` |
| [`synlimit_burst`](#synlimit_burst-serverlisteners) | `u32` | `2` | `✔` | | [`synlimit_burst`](#synlimit_burst-serverlisteners) | `u32` | `1` | `✔` |
| [`synlimit_ios_seconds`](#synlimit_ios_seconds-serverlisteners) | `u32` | `1` | `✔` |
| [`synlimit_ios_hitcount`](#synlimit_ios_hitcount-serverlisteners) | `u32` | `12` | `✔` |
| [`synlimit_ios_burst`](#synlimit_ios_burst-serverlisteners) | `u32` | `24` | `✔` |
| [`synlimit_hashlimit_expire_ms`](#synlimit_hashlimit_expire_ms-serverlisteners) | `u32` | `60000` | `✔` |
| [`synlimit_hashlimit_size`](#synlimit_hashlimit_size-serverlisteners) | `u32` | `32768` | `✔` |
| [`announce`](#announce) | `String` | — | `✘` | | [`announce`](#announce) | `String` | — | `✘` |
| [`announce_ip`](#announce_ip) | `IpAddr` | — | `✘` | | [`announce_ip`](#announce_ip) | `IpAddr` | — | `✘` |
| [`proxy_protocol`](#proxy_protocol) | `bool` | — | `✘` | | [`proxy_protocol`](#proxy_protocol) | `bool` | — | `✘` |
@@ -2277,7 +2283,8 @@
``` ```
## synlimit (server.listeners) ## synlimit (server.listeners)
- **Ограничения / валидация**: `false`, `"iptables"` или `"nftables"`. Если параметр не задан или задан как `false`, SYN limiter для этого listener’а выключен. - **Ограничения / валидация**: `false`, `"iptables"` или `"nftables"`. Если параметр не задан или задан как `false`, SYN limiter для этого listener’а выключен.
- **Описание**: Устанавливает per-listener Linux netfilter SYN limiter rules для порта listener’а. `"iptables"` использует `iptables`/`ip6tables` filter rules с `hashlimit` match как per-source token bucket. `"nftables"` использует per-source `meter` rules с `limit rate over` и автоматически определяет, какие table families уже используются на хосте (`inet`, `ip`, `ip6`), перед созданием Telemt-owned tables. Token-bucket rate равен `synlimit_hitcount / synlimit_seconds`; `synlimit_burst` управляет burst size. Rules reconciled at runtime и удаляются при graceful shutdown Telemt; `SIGKILL` процессом не очищается. Требует CAP_NET_ADMIN. Изменения `synlimit*` hot-reload’ятся для существующих listener endpoints; изменение listener `ip` или `port` по-прежнему требует restart/rebind. - **Описание**: Устанавливает per-listener Linux netfilter two-tier SYN-fix rules для порта listener’а. `"iptables"` использует `iptables`/`ip6tables` filter rules с `hashlimit`, `length` и TTL/hop-limit matches. `"nftables"` использует Telemt-owned tables с per-source `meter` rules и эквивалентными IPv4/IPv6 classifiers. Rules вставляются рано в `INPUT`, принимают under-limit SYN packets и отвечают TCP RST на over-limit SYN packets, чтобы клиент быстро переподключался вместо ожидания silent DROP timeout. Generic bucket управляется `synlimit_seconds`, `synlimit_hitcount` и `synlimit_burst`; iOS-like TTL/length bucket управляется `synlimit_ios_*`. Rules reconciled at runtime и удаляются при graceful shutdown Telemt; `SIGKILL` процессом не очищается. Требует CAP_NET_ADMIN. Изменения `synlimit*` hot-reload’ятся для существующих listener endpoints; изменение listener `ip` или `port` по-прежнему требует restart/rebind.
- **Operator note**: Telemt не сохраняет rules через `iptables-persistent`, не пишет `/etc/sysctl.d`, не меняет systemd limits и не модифицирует `client_mss`. Host-level tuning применяется оператором вручную.
- **Пример**: - **Пример**:
```toml ```toml
@@ -2292,8 +2299,8 @@
synlimit = "nftables" synlimit = "nftables"
``` ```
## synlimit_seconds (server.listeners) ## synlimit_seconds (server.listeners)
- **Ограничения / валидация**: `u32`, должно быть `> 0`. Значение по умолчанию: `1`. - **Ограничения / валидация**: `u32`, должно быть `> 0`. Значение по умолчанию: `60`.
- **Описание**: Token-bucket interval для обоих SYN limiter backends. Rate равен `synlimit_hitcount / synlimit_seconds` и рендерится в native netfilter rate units (`second`, `minute`, `hour` или `day`). - **Описание**: Generic SYN-fix token-bucket interval. Rate равен `synlimit_hitcount / synlimit_seconds` и рендерится в native netfilter rate units (`second`, `minute`, `hour` или `day`). Этот bucket обрабатывает SYN packets, которые не совпали с iOS-like TTL/length classifier.
- **Пример**: - **Пример**:
```toml ```toml
@@ -2301,11 +2308,11 @@
ip = "0.0.0.0" ip = "0.0.0.0"
port = 443 port = 443
synlimit = "iptables" synlimit = "iptables"
synlimit_seconds = 1 synlimit_seconds = 60
``` ```
## synlimit_hitcount (server.listeners) ## synlimit_hitcount (server.listeners)
- **Ограничения / валидация**: `u32`, должно быть `> 0`. Значение по умолчанию: `1`. - **Ограничения / валидация**: `u32`, должно быть `> 0`. Значение по умолчанию: `48`.
- **Описание**: Token-bucket rate amount для обоих SYN limiter backends. Вместе с `synlimit_seconds` задает разрешенный source-IP SYN rate до того, как excess SYN packets начнут drop’аться. - **Описание**: Generic SYN-fix token-bucket rate amount. Вместе с `synlimit_seconds` задает разрешенный source-IP SYN rate до того, как excess SYN packets получат TCP RST.
- **Пример**: - **Пример**:
```toml ```toml
@@ -2313,11 +2320,11 @@
ip = "0.0.0.0" ip = "0.0.0.0"
port = 443 port = 443
synlimit = "iptables" synlimit = "iptables"
synlimit_hitcount = 1 synlimit_hitcount = 48
``` ```
## synlimit_burst (server.listeners) ## synlimit_burst (server.listeners)
- **Ограничения / валидация**: `u32`, должно быть `> 0`. Значение по умолчанию: `2`. - **Ограничения / валидация**: `u32`, должно быть `> 0`. Значение по умолчанию: `1`.
- **Описание**: Token-bucket burst size для обоих SYN limiter backends. Более высокие значения разрешают short connection bursts с одного source IP перед применением steady-state rate `synlimit_hitcount / synlimit_seconds`. - **Описание**: Generic SYN-fix token-bucket burst size. Более высокие значения разрешают short connection bursts с одного source IP перед применением steady-state rate `synlimit_hitcount / synlimit_seconds`.
- **Пример**: - **Пример**:
```toml ```toml
@@ -2325,7 +2332,67 @@
ip = "0.0.0.0" ip = "0.0.0.0"
port = 443 port = 443
synlimit = "iptables" synlimit = "iptables"
synlimit_burst = 2 synlimit_burst = 1
```
## synlimit_ios_seconds (server.listeners)
- **Ограничения / валидация**: `u32`, должно быть `> 0`. Значение по умолчанию: `1`.
- **Описание**: Token-bucket interval для SYN packets, совпавших с iOS-like classifier. IPv4 match: packet length `64` и TTL `< 65`; IPv6 match: packet length `84` и hop limit `< 65`.
- **Пример**:
```toml
[[server.listeners]]
ip = "0.0.0.0"
port = 443
synlimit = "iptables"
synlimit_ios_seconds = 1
```
## synlimit_ios_hitcount (server.listeners)
- **Ограничения / валидация**: `u32`, должно быть `> 0`. Значение по умолчанию: `12`.
- **Описание**: Token-bucket rate amount для iOS-like SYN classifier.
- **Пример**:
```toml
[[server.listeners]]
ip = "0.0.0.0"
port = 443
synlimit = "iptables"
synlimit_ios_hitcount = 12
```
## synlimit_ios_burst (server.listeners)
- **Ограничения / валидация**: `u32`, должно быть `> 0`. Значение по умолчанию: `24`.
- **Описание**: Token-bucket burst size для iOS-like SYN classifier.
- **Пример**:
```toml
[[server.listeners]]
ip = "0.0.0.0"
port = 443
synlimit = "iptables"
synlimit_ios_burst = 24
```
## synlimit_hashlimit_expire_ms (server.listeners)
- **Ограничения / валидация**: `u32`, должно быть `> 0`. Значение по умолчанию: `60000`.
- **Описание**: Entry expiration в миллисекундах для iptables/ip6tables hashlimit buckets. nftables meters используют kernel-managed state и не имеют точного аналога этого knob.
- **Пример**:
```toml
[[server.listeners]]
ip = "0.0.0.0"
port = 443
synlimit = "iptables"
synlimit_hashlimit_expire_ms = 60000
```
## synlimit_hashlimit_size (server.listeners)
- **Ограничения / валидация**: `u32`, должно быть `> 0`. Значение по умолчанию: `32768`.
- **Описание**: Hash table size для iptables/ip6tables hashlimit buckets. nftables meters используют kernel-managed state и не имеют точного аналога этого knob.
- **Пример**:
```toml
[[server.listeners]]
ip = "0.0.0.0"
port = 443
synlimit = "iptables"
synlimit_hashlimit_size = 32768
``` ```
## announce ## announce
- **Ограничения / валидация**: `String` (необязательный параметр). Не должен быть пустым, если задан. - **Ограничения / валидация**: `String` (необязательный параметр). Не должен быть пустым, если задан.
@@ -3045,7 +3112,7 @@
| [`replay_window_secs`](#replay_window_secs) | `u64` | `120` | `✘` | | [`replay_window_secs`](#replay_window_secs) | `u64` | `120` | `✘` |
| [`ignore_time_skew`](#ignore_time_skew) | `bool` | `false` | `✘` | | [`ignore_time_skew`](#ignore_time_skew) | `bool` | `false` | `✘` |
| [`user_rate_limits`](#user_rate_limits) | `Map<String, RateLimitBps>` | `{}` | `✔` | | [`user_rate_limits`](#user_rate_limits) | `Map<String, RateLimitBps>` | `{}` | `✔` |
| [`cidr_rate_limits`](#cidr_rate_limits) | `Map<IpNetwork, RateLimitBps>` | `{}` | `✔` | | [`cidr_rate_limits`](#cidr_rate_limits) | `Map<CidrRateLimitKey, RateLimitBps>` | `{}` | `✔` |
## users ## users
- **Ограничения / валидация**: Не должно быть пустым (должен существовать хотя бы один пользователь). Каждое значение должно состоять **ровно из 32 шестнадцатеричных символов**. - **Ограничения / валидация**: Не должно быть пустым (должен существовать хотя бы один пользователь). Каждое значение должно состоять **ровно из 32 шестнадцатеричных символов**.
@@ -3198,13 +3265,15 @@
alice = { up_bps = 1048576, down_bps = 2097152 } alice = { up_bps = 1048576, down_bps = 2097152 }
``` ```
## cidr_rate_limits ## cidr_rate_limits
- **Ограничения / валидация**: Таблица `CIDR -> { up_bps, down_bps }`. CIDR должен корректно разбираться как `IpNetwork`; хотя бы одно направление должно быть ненулевым. - **Ограничения / валидация**: Таблица `CIDR или auto-template -> { up_bps, down_bps }`. Explicit CIDR-ключи должны корректно разбираться как `IpNetwork`; auto-template ключи должны иметь вид `*4/N` (`N=0..32`), `*6/N` (`N=0..128`) или `*/N` (`N=0..32`). Хотя бы одно направление должно быть ненулевым. Дублирующиеся нормализованные auto-template отклоняются.
- **Описание**: Лимиты скорости для подсетей источников, применяются поверх пользовательских ограничений. - **Описание**: Лимиты скорости для подсетей источников, применяются поверх пользовательских ограничений. Explicit CIDR-правила используют longest-prefix-wins и имеют приоритет над auto-template. Auto-template создают bucket’ы лениво по matched source subnet: `*4/N` для IPv4, `*6/N` для IPv6, а `*/N` является dual-stack shorthand, где IPv4 использует `/N`, а IPv6 — `/(N * 4)`.
- **Example**: - **Example**:
```toml ```toml
[access.cidr_rate_limits] [access.cidr_rate_limits]
"203.0.113.0/24" = { up_bps = 0, down_bps = 1048576 } "203.0.113.0/24" = { up_bps = 0, down_bps = 1048576 }
"*4/32" = { up_bps = 262144, down_bps = 1048576 }
"*6/64" = { up_bps = 262144, down_bps = 1048576 }
``` ```
# [[upstreams]] # [[upstreams]]
+51 -2
View File
@@ -6,19 +6,28 @@ use toml::Value as Toml;
use super::ApiShared; use super::ApiShared;
use super::config_store::{ use super::config_store::{
EDITABLE_SECTIONS, compute_revision, current_revision, save_sections_to_disk, EDITABLE_SECTIONS, compute_revision, current_revision, load_config_from_disk,
save_sections_to_disk,
}; };
use super::model::ApiFailure; use super::model::ApiFailure;
use crate::config::ProxyConfig; use crate::config::ProxyConfig;
use crate::config::hot_reload::classify_config_changes; use crate::config::hot_reload::classify_config_changes;
use crate::maestro::reload::{ReloadAccepted, ReloadRequest, ReloadSubmitError};
use crate::maestro::runtime_build::deferred_process_fields;
use serde::Serialize; use serde::Serialize;
use std::path::Path; use std::path::Path;
use std::sync::Arc;
#[derive(Debug, Serialize)] #[derive(Debug, Serialize)]
pub(super) struct PatchConfigResponse { pub(super) struct PatchConfigResponse {
pub revision: String, pub revision: String,
pub restart_required: bool, pub restart_required: bool,
pub runtime_reload_required: bool,
pub process_restart_required: bool,
pub deferred_process_fields: Vec<String>,
pub changed: Vec<String>, pub changed: Vec<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub reload: Option<ReloadAccepted>,
} }
/// Shared-state wrapper around [`apply_patch_to_path`]: serializes config /// Shared-state wrapper around [`apply_patch_to_path`]: serializes config
@@ -27,10 +36,40 @@ pub(super) struct PatchConfigResponse {
pub(super) async fn patch_config( pub(super) async fn patch_config(
patch_json: Json, patch_json: Json,
expected_revision: Option<String>, expected_revision: Option<String>,
reload_request: Option<ReloadRequest>,
shared: &ApiShared, shared: &ApiShared,
) -> Result<PatchConfigResponse, ApiFailure> { ) -> Result<PatchConfigResponse, ApiFailure> {
let _guard = shared.mutation_lock.lock().await; let _guard = shared.mutation_lock.lock().await;
let resp = apply_patch_to_path(&shared.config_path, &patch_json, expected_revision).await?; if reload_request.is_some()
&& let Some(reload_id) = shared.reload_control.in_progress().await
{
return Err(ApiFailure::new(
hyper::StatusCode::CONFLICT,
"reload_in_progress",
format!("Reload {} is already in progress", reload_id),
));
}
let mut resp = apply_patch_to_path(&shared.config_path, &patch_json, expected_revision).await?;
if let Some(request) = reload_request {
let config = Arc::new(load_config_from_disk(&shared.config_path).await?);
let accepted = shared
.reload_control
.submit(config, resp.revision.clone(), request)
.await
.map_err(|error| match error {
ReloadSubmitError::InProgress(reload_id) => ApiFailure::new(
hyper::StatusCode::CONFLICT,
"reload_in_progress",
format!("Reload {} is already in progress", reload_id),
),
ReloadSubmitError::MaestroUnavailable => ApiFailure::new(
hyper::StatusCode::SERVICE_UNAVAILABLE,
"maestro_unavailable",
"Maestro reload coordinator is unavailable",
),
})?;
resp.reload = Some(accepted);
}
drop(_guard); drop(_guard);
shared shared
.runtime_events .runtime_events
@@ -114,6 +153,7 @@ pub(super) async fn apply_patch_to_path(
// 4. classify changes (Telemt's own hot/restart rule) // 4. classify changes (Telemt's own hot/restart rule)
let class = classify_config_changes(&old_cfg, &new_cfg); let class = classify_config_changes(&old_cfg, &new_cfg);
let deferred_process_fields = deferred_process_fields(&old_cfg, &new_cfg);
// 5. write only the touched top-level sections // 5. write only the touched top-level sections
let revision = save_sections_to_disk(config_path, &new_cfg, &touched).await?; let revision = save_sections_to_disk(config_path, &new_cfg, &touched).await?;
@@ -121,7 +161,11 @@ pub(super) async fn apply_patch_to_path(
Ok(PatchConfigResponse { Ok(PatchConfigResponse {
revision, revision,
restart_required: class.restart_required, restart_required: class.restart_required,
runtime_reload_required: class.restart_required,
process_restart_required: !deferred_process_fields.is_empty(),
deferred_process_fields,
changed: class.changed, changed: class.changed,
reload: None,
}) })
} }
@@ -266,6 +310,9 @@ mod tests {
let patch: Json = serde_json::json!({"censorship": {"tls_domain": "b.com"}}); let patch: Json = serde_json::json!({"censorship": {"tls_domain": "b.com"}});
let resp = apply_patch_to_path(&path, &patch, None).await.unwrap(); let resp = apply_patch_to_path(&path, &patch, None).await.unwrap();
assert!(resp.restart_required); assert!(resp.restart_required);
assert!(resp.runtime_reload_required);
assert!(!resp.process_restart_required);
assert!(resp.deferred_process_fields.is_empty());
assert!(resp.changed.iter().any(|c| c == "censorship")); assert!(resp.changed.iter().any(|c| c == "censorship"));
let written = std::fs::read_to_string(&path).unwrap(); let written = std::fs::read_to_string(&path).unwrap();
assert!(written.contains("tls_domain = \"b.com\"")); assert!(written.contains("tls_domain = \"b.com\""));
@@ -406,6 +453,8 @@ mod tests {
let patch: Json = serde_json::json!({"general": {"log_level": "debug"}}); let patch: Json = serde_json::json!({"general": {"log_level": "debug"}});
let resp = apply_patch_to_path(&path, &patch, None).await.unwrap(); let resp = apply_patch_to_path(&path, &patch, None).await.unwrap();
assert!(!resp.restart_required); assert!(!resp.restart_required);
assert!(!resp.runtime_reload_required);
assert!(!resp.process_restart_required);
assert!(resp.changed.iter().any(|c| c == "general")); assert!(resp.changed.iter().any(|c| c == "general"));
} }
} }
+15
View File
@@ -72,6 +72,13 @@ pub(super) async fn current_revision(config_path: &Path) -> Result<String, ApiFa
Ok(compute_revision(&content)) Ok(compute_revision(&content))
} }
pub(crate) async fn current_revision_for_maestro(config_path: &Path) -> Result<String, String> {
let content = tokio::fs::read_to_string(config_path)
.await
.map_err(|error| format!("failed to read config: {}", error))?;
Ok(compute_revision(&content))
}
pub(super) fn compute_revision(content: &str) -> String { pub(super) fn compute_revision(content: &str) -> String {
let mut hasher = Sha256::new(); let mut hasher = Sha256::new();
hasher.update(content.as_bytes()); hasher.update(content.as_bytes());
@@ -86,6 +93,14 @@ pub(super) async fn load_config_from_disk(config_path: &Path) -> Result<ProxyCon
.map_err(|e| ApiFailure::internal(format!("failed to load config: {}", e))) .map_err(|e| ApiFailure::internal(format!("failed to load config: {}", e)))
} }
pub(super) async fn load_config_for_reload(config_path: &Path) -> Result<ProxyConfig, ApiFailure> {
let config_path = config_path.to_path_buf();
tokio::task::spawn_blocking(move || ProxyConfig::load(config_path))
.await
.map_err(|error| ApiFailure::internal(format!("failed to join config loader: {}", error)))?
.map_err(|error| ApiFailure::bad_request(format!("invalid runtime config: {}", error)))
}
#[allow(dead_code)] #[allow(dead_code)]
pub(super) async fn save_config_to_disk( pub(super) async fn save_config_to_disk(
config_path: &Path, config_path: &Path,
+170 -13
View File
@@ -7,6 +7,7 @@ use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering}; use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::time::Duration; use std::time::Duration;
use arc_swap::ArcSwap;
use http_body_util::Full; use http_body_util::Full;
use hyper::body::{Bytes, Incoming}; use hyper::body::{Bytes, Incoming};
use hyper::header::AUTHORIZATION; use hyper::header::AUTHORIZATION;
@@ -19,8 +20,10 @@ use tokio::sync::{Mutex, RwLock, Semaphore, watch};
use tokio::time::timeout; use tokio::time::timeout;
use tracing::{debug, info, warn}; use tracing::{debug, info, warn};
use crate::config::{ApiGrayAction, ProxyConfig}; use crate::config::ApiGrayAction;
use crate::ip_tracker::UserIpTracker; use crate::ip_tracker::UserIpTracker;
use crate::maestro::generation::{RuntimeGeneration, RuntimeWatchState};
use crate::maestro::reload::{ReloadAccepted, ReloadControl, ReloadRequest, ReloadSubmitError};
use crate::proxy::route_mode::RouteRuntimeController; use crate::proxy::route_mode::RouteRuntimeController;
use crate::proxy::shared_state::ProxySharedState; use crate::proxy::shared_state::ProxySharedState;
use crate::startup::StartupTracker; use crate::startup::StartupTracker;
@@ -29,11 +32,13 @@ use crate::transport::UpstreamManager;
use crate::transport::middle_proxy::MePool; use crate::transport::middle_proxy::MePool;
mod config_edit; mod config_edit;
mod config_store; pub(crate) mod config_store;
mod events; mod events;
mod http_utils; mod http_utils;
mod model; mod model;
mod patch; mod patch;
#[cfg(test)]
mod reload_tests;
mod runtime_edge; mod runtime_edge;
mod runtime_init; mod runtime_init;
mod runtime_min; mod runtime_min;
@@ -44,7 +49,8 @@ mod runtime_zero;
mod users; mod users;
use config_store::{ use config_store::{
current_revision, ensure_expected_revision, load_config_from_disk, parse_if_match, current_revision, ensure_expected_revision, load_config_for_reload, load_config_from_disk,
parse_if_match,
}; };
use events::ApiEventStore; use events::ApiEventStore;
use http_utils::{error_response, read_json, read_optional_json, success_response}; use http_utils::{error_response, read_json, read_optional_json, success_response};
@@ -107,12 +113,15 @@ pub(super) struct ApiShared {
pub(super) minimal_cache: Arc<Mutex<Option<MinimalCacheEntry>>>, pub(super) minimal_cache: Arc<Mutex<Option<MinimalCacheEntry>>>,
pub(super) runtime_edge_connections_cache: Arc<Mutex<Option<EdgeConnectionsCacheEntry>>>, pub(super) runtime_edge_connections_cache: Arc<Mutex<Option<EdgeConnectionsCacheEntry>>>,
pub(super) runtime_edge_recompute_lock: Arc<Mutex<()>>, pub(super) runtime_edge_recompute_lock: Arc<Mutex<()>>,
pub(super) cache_generation: Arc<AtomicU64>,
pub(super) runtime_events: Arc<ApiEventStore>, pub(super) runtime_events: Arc<ApiEventStore>,
pub(super) request_id: Arc<AtomicU64>, pub(super) request_id: Arc<AtomicU64>,
pub(super) runtime_state: Arc<ApiRuntimeState>, pub(super) runtime_state: Arc<ApiRuntimeState>,
pub(super) startup_tracker: Arc<StartupTracker>, pub(super) startup_tracker: Arc<StartupTracker>,
pub(super) route_runtime: Arc<RouteRuntimeController>, pub(super) route_runtime: Arc<RouteRuntimeController>,
pub(super) proxy_shared: Arc<ProxySharedState>, pub(super) proxy_shared: Arc<ProxySharedState>,
pub(super) reload_control: ReloadControl,
pub(super) active_runtime: Arc<ArcSwap<RuntimeGeneration>>,
} }
impl ApiShared { impl ApiShared {
@@ -123,6 +132,31 @@ impl ApiShared {
fn detected_link_ips(&self) -> (Option<IpAddr>, Option<IpAddr>) { fn detected_link_ips(&self) -> (Option<IpAddr>, Option<IpAddr>) {
*self.detected_ips_rx.borrow() *self.detected_ips_rx.borrow()
} }
fn for_runtime(&self, runtime: &RuntimeGeneration) -> Self {
Self {
stats: runtime.stats.clone(),
ip_tracker: runtime.ip_tracker.clone(),
me_pool: runtime.me_pool_runtime.clone(),
upstream_manager: runtime.upstream_manager.clone(),
config_path: self.config_path.clone(),
quota_state_path: self.quota_state_path.clone(),
detected_ips_rx: self.detected_ips_rx.clone(),
mutation_lock: self.mutation_lock.clone(),
minimal_cache: self.minimal_cache.clone(),
runtime_edge_connections_cache: self.runtime_edge_connections_cache.clone(),
runtime_edge_recompute_lock: self.runtime_edge_recompute_lock.clone(),
cache_generation: self.cache_generation.clone(),
runtime_events: self.runtime_events.clone(),
request_id: self.request_id.clone(),
runtime_state: self.runtime_state.clone(),
startup_tracker: self.startup_tracker.clone(),
route_runtime: runtime.route_runtime.clone(),
proxy_shared: runtime.proxy_shared.clone(),
reload_control: self.reload_control.clone(),
active_runtime: self.active_runtime.clone(),
}
}
} }
fn auth_header_matches(actual: &str, expected: &str) -> bool { fn auth_header_matches(actual: &str, expected: &str) -> bool {
@@ -144,6 +178,41 @@ fn user_action_route_matches(path: &str, suffix: &str) -> bool {
.unwrap_or(false) .unwrap_or(false)
} }
fn reload_status_route_id(path: &str) -> Option<u64> {
path.strip_prefix("/v1/system/reload/")
.filter(|id| !id.is_empty() && !id.contains('/'))
.and_then(|id| id.parse().ok())
}
async fn submit_reload_from_disk(
config_path: &std::path::Path,
mutation_lock: &Mutex<()>,
reload_control: &ReloadControl,
expected_revision: Option<&str>,
request: ReloadRequest,
) -> Result<(ReloadAccepted, String), ApiFailure> {
let _guard = mutation_lock.lock().await;
ensure_expected_revision(config_path, expected_revision).await?;
let revision = current_revision(config_path).await?;
let config = Arc::new(load_config_for_reload(config_path).await?);
let accepted = reload_control
.submit(config, revision.clone(), request)
.await
.map_err(|error| match error {
ReloadSubmitError::InProgress(reload_id) => ApiFailure::new(
StatusCode::CONFLICT,
"reload_in_progress",
format!("Reload {} is already in progress", reload_id),
),
ReloadSubmitError::MaestroUnavailable => ApiFailure::new(
StatusCode::SERVICE_UNAVAILABLE,
"maestro_unavailable",
"Maestro reload coordinator is unavailable",
),
})?;
Ok((accepted, revision))
}
fn allowed_methods_for_path(path: &str) -> Option<&'static str> { fn allowed_methods_for_path(path: &str) -> Option<&'static str> {
match path { match path {
"/v1/health" "/v1/health"
@@ -175,12 +244,14 @@ fn allowed_methods_for_path(path: &str) -> Option<&'static str> {
| "/v1/stats/users/active-ips" | "/v1/stats/users/active-ips"
| "/v1/stats/users/quota" | "/v1/stats/users/quota"
| "/v1/stats/users" => Some(ALLOW_GET), | "/v1/stats/users" => Some(ALLOW_GET),
"/v1/system/reload" => Some(ALLOW_POST),
"/v1/users" => Some(ALLOW_GET_POST), "/v1/users" => Some(ALLOW_GET_POST),
"/v1/config" => Some(ALLOW_GET_PATCH), "/v1/config" => Some(ALLOW_GET_PATCH),
_ if user_action_route_matches(path, "/reset-quota") => Some(ALLOW_POST), _ if user_action_route_matches(path, "/reset-quota") => Some(ALLOW_POST),
_ if user_action_route_matches(path, "/rotate-secret") => Some(ALLOW_POST), _ if user_action_route_matches(path, "/rotate-secret") => Some(ALLOW_POST),
_ if user_action_route_matches(path, "/enable") => Some(ALLOW_POST), _ if user_action_route_matches(path, "/enable") => Some(ALLOW_POST),
_ if user_action_route_matches(path, "/disable") => Some(ALLOW_POST), _ if user_action_route_matches(path, "/disable") => Some(ALLOW_POST),
_ if reload_status_route_id(path).is_some() => Some(ALLOW_GET),
_ if path _ if path
.strip_prefix("/v1/users/") .strip_prefix("/v1/users/")
.map(|user| !user.is_empty() && !user.contains('/')) .map(|user| !user.is_empty() && !user.contains('/'))
@@ -200,14 +271,35 @@ pub async fn serve(
route_runtime: Arc<RouteRuntimeController>, route_runtime: Arc<RouteRuntimeController>,
proxy_shared: Arc<ProxySharedState>, proxy_shared: Arc<ProxySharedState>,
upstream_manager: Arc<UpstreamManager>, upstream_manager: Arc<UpstreamManager>,
config_rx: watch::Receiver<Arc<ProxyConfig>>,
admission_rx: watch::Receiver<bool>,
config_path: PathBuf, config_path: PathBuf,
quota_state_path: PathBuf, quota_state_path: PathBuf,
detected_ips_rx: watch::Receiver<(Option<IpAddr>, Option<IpAddr>)>, detected_ips_rx: watch::Receiver<(Option<IpAddr>, Option<IpAddr>)>,
process_started_at_epoch_secs: u64, process_started_at_epoch_secs: u64,
startup_tracker: Arc<StartupTracker>, startup_tracker: Arc<StartupTracker>,
reload_control: ReloadControl,
mut active_runtime_rx: watch::Receiver<Option<Arc<ArcSwap<RuntimeGeneration>>>>,
mut runtime_watch_rx: watch::Receiver<Option<RuntimeWatchState>>,
) { ) {
let active_runtime = loop {
if let Some(active_runtime) = active_runtime_rx.borrow().clone() {
break active_runtime;
}
if active_runtime_rx.changed().await.is_err() {
warn!("Runtime generation channel closed before API bootstrap");
return;
}
};
let initial_watch_state = loop {
if let Some(watch_state) = runtime_watch_rx.borrow().clone() {
break watch_state;
}
if runtime_watch_rx.changed().await.is_err() {
warn!("Runtime watch channel closed before API bootstrap");
return;
}
};
let config_rx = initial_watch_state.config_rx.clone();
let admission_rx = initial_watch_state.admission_rx.clone();
let listener = match TcpListener::bind(listen).await { let listener = match TcpListener::bind(listen).await {
Ok(listener) => listener, Ok(listener) => listener,
Err(error) => { Err(error) => {
@@ -241,6 +333,7 @@ pub async fn serve(
minimal_cache: Arc::new(Mutex::new(None)), minimal_cache: Arc::new(Mutex::new(None)),
runtime_edge_connections_cache: Arc::new(Mutex::new(None)), runtime_edge_connections_cache: Arc::new(Mutex::new(None)),
runtime_edge_recompute_lock: Arc::new(Mutex::new(())), runtime_edge_recompute_lock: Arc::new(Mutex::new(())),
cache_generation: Arc::new(AtomicU64::new(1)),
runtime_events: Arc::new(ApiEventStore::new( runtime_events: Arc::new(ApiEventStore::new(
config_rx.borrow().server.api.runtime_edge_events_capacity, config_rx.borrow().server.api.runtime_edge_events_capacity,
)), )),
@@ -249,11 +342,12 @@ pub async fn serve(
startup_tracker, startup_tracker,
route_runtime, route_runtime,
proxy_shared, proxy_shared,
reload_control,
active_runtime,
}); });
spawn_runtime_watchers( spawn_runtime_watchers(
config_rx.clone(), runtime_watch_rx,
admission_rx.clone(),
runtime_state.clone(), runtime_state.clone(),
shared.runtime_events.clone(), shared.runtime_events.clone(),
); );
@@ -282,13 +376,11 @@ pub async fn serve(
}; };
let shared_conn = shared.clone(); let shared_conn = shared.clone();
let config_rx_conn = config_rx.clone();
tokio::spawn(async move { tokio::spawn(async move {
let _connection_permit = connection_permit; let _connection_permit = connection_permit;
let svc = service_fn(move |req: Request<Incoming>| { let svc = service_fn(move |req: Request<Incoming>| {
let shared_req = shared_conn.clone(); let shared_req = shared_conn.clone();
let config_rx_req = config_rx_conn.clone(); async move { handle(req, peer, shared_req).await }
async move { handle(req, peer, shared_req, config_rx_req).await }
}); });
match timeout( match timeout(
API_HTTP_CONNECTION_TIMEOUT, API_HTTP_CONNECTION_TIMEOUT,
@@ -318,8 +410,19 @@ async fn handle(
req: Request<Incoming>, req: Request<Incoming>,
peer: SocketAddr, peer: SocketAddr,
shared: Arc<ApiShared>, shared: Arc<ApiShared>,
config_rx: watch::Receiver<Arc<ProxyConfig>>,
) -> Result<Response<Full<Bytes>>, IoError> { ) -> Result<Response<Full<Bytes>>, IoError> {
let runtime = shared.active_runtime.load_full();
let previous_cache_generation = shared.cache_generation.swap(runtime.id, Ordering::AcqRel);
if previous_cache_generation != runtime.id {
*shared.minimal_cache.lock().await = None;
*shared.runtime_edge_connections_cache.lock().await = None;
}
let shared = Arc::new(shared.for_runtime(runtime.as_ref()));
let config_rx = runtime.config_rx.clone();
shared
.runtime_state
.admission_open
.store(*runtime.admission_rx.borrow(), Ordering::Relaxed);
let request_id = shared.next_request_id(); let request_id = shared.next_request_id();
let cfg = config_rx.borrow().clone(); let cfg = config_rx.borrow().clone();
let api_cfg = &cfg.server.api; let api_cfg = &cfg.server.api;
@@ -651,6 +754,33 @@ async fn handle(
config_edit::read_managed_config(&shared.config_path).await?; config_edit::read_managed_config(&shared.config_path).await?;
Ok(success_response(StatusCode::OK, value, revision)) Ok(success_response(StatusCode::OK, value, revision))
} }
("POST", "/v1/system/reload") => {
if api_cfg.read_only {
return Ok(error_response(
request_id,
ApiFailure::new(
StatusCode::FORBIDDEN,
"read_only",
"API runs in read-only mode",
),
));
}
let expected_revision = parse_if_match(req.headers());
let request = read_optional_json::<ReloadRequest>(req.into_body(), body_limit)
.await?
.unwrap_or_default();
request.validate().map_err(ApiFailure::bad_request)?;
let (accepted, revision) = submit_reload_from_disk(
&shared.config_path,
shared.mutation_lock.as_ref(),
&shared.reload_control,
expected_revision.as_deref(),
request,
)
.await?;
Ok(success_response(StatusCode::ACCEPTED, accepted, revision))
}
("PATCH", "/v1/config") => { ("PATCH", "/v1/config") => {
if api_cfg.read_only { if api_cfg.read_only {
return Ok(error_response( return Ok(error_response(
@@ -663,11 +793,20 @@ async fn handle(
)); ));
} }
let expected_revision = parse_if_match(req.headers()); let expected_revision = parse_if_match(req.headers());
let reload_request =
ReloadRequest::from_query(query.as_deref()).map_err(ApiFailure::bad_request)?;
let body = read_json::<serde_json::Value>(req.into_body(), body_limit).await?; let body = read_json::<serde_json::Value>(req.into_body(), body_limit).await?;
match config_edit::patch_config(body, expected_revision, &shared).await { match config_edit::patch_config(body, expected_revision, reload_request, &shared)
.await
{
Ok(resp) => { Ok(resp) => {
let revision = resp.revision.clone(); let revision = resp.revision.clone();
Ok(success_response(StatusCode::OK, resp, revision)) let status = if resp.reload.is_some() {
StatusCode::ACCEPTED
} else {
StatusCode::OK
};
Ok(success_response(status, resp, revision))
} }
Err(error) => { Err(error) => {
shared shared
@@ -678,6 +817,24 @@ async fn handle(
} }
} }
_ => { _ => {
if method == Method::GET
&& let Some(reload_id) = reload_status_route_id(normalized_path)
{
let revision = current_revision(&shared.config_path).await?;
let status =
shared
.reload_control
.status(reload_id)
.await
.ok_or_else(|| {
ApiFailure::new(
StatusCode::NOT_FOUND,
"reload_not_found",
format!("Reload {} was not found", reload_id),
)
})?;
return Ok(success_response(StatusCode::OK, status, revision));
}
if method == Method::POST if method == Method::POST
&& let Some(base_user) = normalized_path && let Some(base_user) = normalized_path
.strip_prefix("/v1/users/") .strip_prefix("/v1/users/")
+119
View File
@@ -0,0 +1,119 @@
use super::*;
use crate::config::ProxyConfig;
async fn config_file() -> (tempfile::TempDir, PathBuf, String) {
let directory = tempfile::tempdir().unwrap();
let path = directory.path().join("config.toml");
let mut config = ProxyConfig::default();
config.server.max_connections = 4_242;
let body = toml::to_string_pretty(&config).unwrap();
tokio::fs::write(&path, &body).await.unwrap();
let revision = config_store::compute_revision(&body);
(directory, path, revision)
}
#[tokio::test]
async fn reload_submission_uses_matching_disk_revision_and_snapshot() {
let (_directory, path, revision) = config_file().await;
let mutation_lock = Mutex::new(());
let (control, mut commands) = ReloadControl::channel(1);
let request = ReloadRequest::default();
let (accepted, response_revision) = submit_reload_from_disk(
&path,
&mutation_lock,
&control,
Some(&revision),
request.clone(),
)
.await
.unwrap();
let command = commands.recv().await.unwrap();
assert_eq!(response_revision, revision);
assert_eq!(accepted.config_revision, revision);
assert_eq!(command.config_revision, revision);
assert_eq!(command.request, request);
assert_eq!(command.config.server.max_connections, 4_242);
}
#[tokio::test]
async fn revision_conflict_rejects_without_enqueuing_reload() {
let (_directory, path, _revision) = config_file().await;
let mutation_lock = Mutex::new(());
let (control, _commands) = ReloadControl::channel(1);
let error = submit_reload_from_disk(
&path,
&mutation_lock,
&control,
Some("stale-revision"),
ReloadRequest::default(),
)
.await
.unwrap_err();
assert_eq!(error.status, StatusCode::CONFLICT);
assert_eq!(error.code, "revision_conflict");
assert_eq!(control.in_progress().await, None);
}
#[tokio::test]
async fn reload_conflict_and_closed_coordinator_map_to_http_contract() {
let (_directory, path, _revision) = config_file().await;
let mutation_lock = Mutex::new(());
let (control, mut commands) = ReloadControl::channel(1);
let _accepted = submit_reload_from_disk(
&path,
&mutation_lock,
&control,
None,
ReloadRequest::default(),
)
.await
.unwrap();
let _command = commands.recv().await.unwrap();
let conflict = submit_reload_from_disk(
&path,
&mutation_lock,
&control,
None,
ReloadRequest::default(),
)
.await
.unwrap_err();
assert_eq!(conflict.status, StatusCode::CONFLICT);
assert_eq!(conflict.code, "reload_in_progress");
control.fail(1, "test cleanup").await;
drop(commands);
let unavailable = submit_reload_from_disk(
&path,
&mutation_lock,
&control,
None,
ReloadRequest::default(),
)
.await
.unwrap_err();
assert_eq!(unavailable.status, StatusCode::SERVICE_UNAVAILABLE);
assert_eq!(unavailable.code, "maestro_unavailable");
}
#[test]
fn reload_routes_expose_only_documented_methods_and_ids() {
assert_eq!(
allowed_methods_for_path("/v1/system/reload"),
Some(ALLOW_POST)
);
assert_eq!(
allowed_methods_for_path("/v1/system/reload/42"),
Some(ALLOW_GET)
);
assert_eq!(reload_status_route_id("/v1/system/reload/42"), Some(42));
assert_eq!(
reload_status_route_id("/v1/system/reload/not-a-number"),
None
);
}
+282 -29
View File
@@ -4,58 +4,184 @@ use std::time::{SystemTime, UNIX_EPOCH};
use tokio::sync::watch; use tokio::sync::watch;
use crate::config::ProxyConfig; use crate::maestro::generation::RuntimeWatchState;
use super::ApiRuntimeState; use super::ApiRuntimeState;
use super::events::ApiEventStore; use super::events::ApiEventStore;
pub(super) fn spawn_runtime_watchers( pub(super) fn spawn_runtime_watchers(
config_rx: watch::Receiver<Arc<ProxyConfig>>, runtime_watch_rx: watch::Receiver<Option<RuntimeWatchState>>,
admission_rx: watch::Receiver<bool>,
runtime_state: Arc<ApiRuntimeState>, runtime_state: Arc<ApiRuntimeState>,
runtime_events: Arc<ApiEventStore>, runtime_events: Arc<ApiEventStore>,
) { ) {
let mut config_rx_reload = config_rx; let _config_watcher = spawn_config_watcher(
let runtime_state_reload = runtime_state.clone(); runtime_watch_rx.clone(),
let runtime_events_reload = runtime_events.clone(); runtime_state.clone(),
runtime_events.clone(),
);
let _admission_watcher =
spawn_admission_watcher(runtime_watch_rx, runtime_state, runtime_events);
}
fn spawn_config_watcher(
mut runtime_watch_rx: watch::Receiver<Option<RuntimeWatchState>>,
runtime_state: Arc<ApiRuntimeState>,
runtime_events: Arc<ApiEventStore>,
) -> tokio::task::JoinHandle<()> {
tokio::spawn(async move { tokio::spawn(async move {
let Some(mut current) = runtime_watch_rx.borrow().clone() else {
return;
};
loop { loop {
if config_rx_reload.changed().await.is_err() { tokio::select! {
biased;
changed = runtime_watch_rx.changed() => {
if changed.is_err() {
break; break;
} }
runtime_state_reload let Some(next) = runtime_watch_rx.borrow().clone() else {
continue;
};
if next.generation_id != current.generation_id {
current = next;
record_config_reload(
&runtime_state,
&runtime_events,
format!("runtime generation {} activated", current.generation_id),
);
}
}
changed = current.config_rx.changed() => {
if changed.is_err() {
let Some(next) = wait_for_new_generation(
&mut runtime_watch_rx,
current.generation_id,
).await else {
break;
};
current = next;
record_config_reload(
&runtime_state,
&runtime_events,
format!("runtime generation {} activated", current.generation_id),
);
continue;
}
if active_generation_id(&runtime_watch_rx) != Some(current.generation_id) {
continue;
}
record_config_reload(
&runtime_state,
&runtime_events,
format!("generation {} config receiver updated", current.generation_id),
);
}
}
}
})
}
fn spawn_admission_watcher(
mut runtime_watch_rx: watch::Receiver<Option<RuntimeWatchState>>,
runtime_state: Arc<ApiRuntimeState>,
runtime_events: Arc<ApiEventStore>,
) -> tokio::task::JoinHandle<()> {
tokio::spawn(async move {
let Some(mut current) = runtime_watch_rx.borrow().clone() else {
return;
};
record_admission_state(&runtime_state, &runtime_events, &current);
loop {
tokio::select! {
biased;
changed = runtime_watch_rx.changed() => {
if changed.is_err() {
break;
}
let Some(next) = runtime_watch_rx.borrow().clone() else {
continue;
};
if next.generation_id != current.generation_id {
current = next;
record_admission_state(&runtime_state, &runtime_events, &current);
}
}
changed = current.admission_rx.changed() => {
if changed.is_err() {
let Some(next) = wait_for_new_generation(
&mut runtime_watch_rx,
current.generation_id,
).await else {
break;
};
current = next;
record_admission_state(&runtime_state, &runtime_events, &current);
continue;
}
if active_generation_id(&runtime_watch_rx) == Some(current.generation_id) {
record_admission_state(&runtime_state, &runtime_events, &current);
}
}
}
}
})
}
fn active_generation_id(
runtime_watch_rx: &watch::Receiver<Option<RuntimeWatchState>>,
) -> Option<u64> {
runtime_watch_rx
.borrow()
.as_ref()
.map(|state| state.generation_id)
}
async fn wait_for_new_generation(
runtime_watch_rx: &mut watch::Receiver<Option<RuntimeWatchState>>,
previous_generation_id: u64,
) -> Option<RuntimeWatchState> {
loop {
if let Some(state) = runtime_watch_rx.borrow().clone()
&& state.generation_id != previous_generation_id
{
return Some(state);
}
if runtime_watch_rx.changed().await.is_err() {
return None;
}
}
}
fn record_config_reload(
runtime_state: &ApiRuntimeState,
runtime_events: &ApiEventStore,
context: String,
) {
runtime_state
.config_reload_count .config_reload_count
.fetch_add(1, Ordering::Relaxed); .fetch_add(1, Ordering::Relaxed);
runtime_state_reload runtime_state
.last_config_reload_epoch_secs .last_config_reload_epoch_secs
.store(now_epoch_secs(), Ordering::Relaxed); .store(now_epoch_secs(), Ordering::Relaxed);
runtime_events_reload.record("config.reload.applied", "config receiver updated"); runtime_events.record("config.reload.applied", context);
} }
});
let mut admission_rx_watch = admission_rx; fn record_admission_state(
tokio::spawn(async move { runtime_state: &ApiRuntimeState,
runtime_state runtime_events: &ApiEventStore,
.admission_open current: &RuntimeWatchState,
.store(*admission_rx_watch.borrow(), Ordering::Relaxed); ) {
runtime_events.record( let admission_open = *current.admission_rx.borrow();
"admission.state",
format!("accepting_new_connections={}", *admission_rx_watch.borrow()),
);
loop {
if admission_rx_watch.changed().await.is_err() {
break;
}
let admission_open = *admission_rx_watch.borrow();
runtime_state runtime_state
.admission_open .admission_open
.store(admission_open, Ordering::Relaxed); .store(admission_open, Ordering::Relaxed);
runtime_events.record( runtime_events.record(
"admission.state", "admission.state",
format!("accepting_new_connections={}", admission_open), format!(
"generation={} accepting_new_connections={}",
current.generation_id, admission_open
),
); );
}
});
} }
fn now_epoch_secs() -> u64 { fn now_epoch_secs() -> u64 {
@@ -64,3 +190,130 @@ fn now_epoch_secs() -> u64 {
.unwrap_or_default() .unwrap_or_default()
.as_secs() .as_secs()
} }
#[cfg(test)]
mod tests {
use super::*;
use crate::config::ProxyConfig;
use std::sync::atomic::{AtomicBool, AtomicU64};
use std::time::Duration;
fn state(
generation_id: u64,
) -> (
RuntimeWatchState,
watch::Sender<Arc<ProxyConfig>>,
watch::Sender<bool>,
) {
let (config_tx, config_rx) = watch::channel(Arc::new(ProxyConfig::default()));
let (admission_tx, admission_rx) = watch::channel(true);
(
RuntimeWatchState {
generation_id,
config_rx,
admission_rx,
},
config_tx,
admission_tx,
)
}
fn runtime_state() -> Arc<ApiRuntimeState> {
Arc::new(ApiRuntimeState {
process_started_at_epoch_secs: 1,
config_reload_count: AtomicU64::new(0),
last_config_reload_epoch_secs: AtomicU64::new(0),
admission_open: AtomicBool::new(false),
})
}
async fn wait_for_count(runtime_state: &ApiRuntimeState, expected: u64) {
tokio::time::timeout(Duration::from_secs(1), async {
loop {
if runtime_state.config_reload_count.load(Ordering::Relaxed) == expected {
break;
}
tokio::task::yield_now().await;
}
})
.await
.unwrap();
}
#[tokio::test]
async fn watchers_follow_only_the_active_generation() {
let (initial, initial_config_tx, initial_admission_tx) = state(1);
let (runtime_watch_tx, runtime_watch_rx) = watch::channel(Some(initial));
let runtime_state = runtime_state();
let events = Arc::new(ApiEventStore::new(16));
spawn_runtime_watchers(runtime_watch_rx, runtime_state.clone(), events.clone());
tokio::task::yield_now().await;
assert_eq!(runtime_state.config_reload_count.load(Ordering::Relaxed), 0);
initial_config_tx.send_replace(Arc::new(ProxyConfig::default()));
wait_for_count(&runtime_state, 1).await;
let (next, next_config_tx, next_admission_tx) = state(2);
runtime_watch_tx.send_replace(Some(next));
wait_for_count(&runtime_state, 2).await;
initial_config_tx.send_replace(Arc::new(ProxyConfig::default()));
initial_admission_tx.send_replace(false);
tokio::task::yield_now().await;
assert_eq!(runtime_state.config_reload_count.load(Ordering::Relaxed), 2);
assert!(runtime_state.admission_open.load(Ordering::Relaxed));
next_config_tx.send_replace(Arc::new(ProxyConfig::default()));
next_admission_tx.send_replace(false);
wait_for_count(&runtime_state, 3).await;
tokio::time::timeout(Duration::from_secs(1), async {
while runtime_state.admission_open.load(Ordering::Relaxed) {
tokio::task::yield_now().await;
}
})
.await
.unwrap();
let snapshot = events.snapshot(16);
assert_eq!(
snapshot
.events
.iter()
.filter(|event| event.event_type == "config.reload.applied")
.count(),
3
);
}
#[tokio::test]
async fn watcher_recovers_from_closed_generation_and_exits_with_process_channel() {
let (initial, initial_config_tx, _initial_admission_tx) = state(1);
let (runtime_watch_tx, runtime_watch_rx) = watch::channel(Some(initial));
let runtime_state = runtime_state();
let events = Arc::new(ApiEventStore::new(16));
let watcher = spawn_config_watcher(runtime_watch_rx, runtime_state.clone(), events.clone());
drop(initial_config_tx);
tokio::task::yield_now().await;
let (next, next_config_tx, _next_admission_tx) = state(2);
runtime_watch_tx.send_replace(Some(next));
wait_for_count(&runtime_state, 1).await;
next_config_tx.send_replace(Arc::new(ProxyConfig::default()));
wait_for_count(&runtime_state, 2).await;
drop(runtime_watch_tx);
tokio::time::timeout(Duration::from_secs(1), watcher)
.await
.unwrap()
.unwrap();
assert_eq!(
events
.snapshot(16)
.events
.iter()
.filter(|event| event.event_type == "config.reload.applied")
.count(),
2
);
}
}
+25 -1
View File
@@ -1,4 +1,5 @@
use std::sync::atomic::Ordering; use std::sync::atomic::Ordering;
use std::time::{SystemTime, UNIX_EPOCH};
use serde::Serialize; use serde::Serialize;
@@ -162,7 +163,7 @@ pub(super) fn build_system_info_data(
build_time_utc, build_time_utc,
rustc_version, rustc_version,
process_started_at_epoch_secs: shared.runtime_state.process_started_at_epoch_secs, process_started_at_epoch_secs: shared.runtime_state.process_started_at_epoch_secs,
uptime_seconds: shared.stats.uptime_secs(), uptime_seconds: process_uptime_seconds(shared.runtime_state.process_started_at_epoch_secs),
config_path: shared.config_path.display().to_string(), config_path: shared.config_path.display().to_string(),
config_hash: revision.to_string(), config_hash: revision.to_string(),
config_reload_count: shared config_reload_count: shared
@@ -173,6 +174,18 @@ pub(super) fn build_system_info_data(
} }
} }
fn process_uptime_seconds(process_started_at_epoch_secs: u64) -> f64 {
let now_epoch_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
process_uptime_seconds_at(process_started_at_epoch_secs, now_epoch_secs)
}
fn process_uptime_seconds_at(process_started_at_epoch_secs: u64, now_epoch_secs: u64) -> f64 {
now_epoch_secs.saturating_sub(process_started_at_epoch_secs) as f64
}
pub(super) async fn build_runtime_gates_data( pub(super) async fn build_runtime_gates_data(
shared: &ApiShared, shared: &ApiShared,
cfg: &ProxyConfig, cfg: &ProxyConfig,
@@ -339,3 +352,14 @@ fn me_writer_pick_mode_label(mode: MeWriterPickMode) -> &'static str {
MeWriterPickMode::P2c => "p2c", MeWriterPickMode::P2c => "p2c",
} }
} }
#[cfg(test)]
mod tests {
use super::process_uptime_seconds_at;
#[test]
fn process_uptime_is_monotonic_and_saturating() {
assert_eq!(process_uptime_seconds_at(100, 135), 35.0);
assert_eq!(process_uptime_seconds_at(135, 100), 0.0);
}
}
+50 -3
View File
@@ -24,6 +24,10 @@ const DEFAULT_ME_ADAPTIVE_FLOOR_MAX_WARM_WRITERS_GLOBAL: u32 = 256;
const DEFAULT_ME_ROUTE_BACKPRESSURE_ENABLED: bool = false; const DEFAULT_ME_ROUTE_BACKPRESSURE_ENABLED: bool = false;
const DEFAULT_ME_ROUTE_FAIRSHARE_ENABLED: bool = false; const DEFAULT_ME_ROUTE_FAIRSHARE_ENABLED: bool = false;
const DEFAULT_ME_WRITER_CMD_CHANNEL_CAPACITY: usize = 4096; const DEFAULT_ME_WRITER_CMD_CHANNEL_CAPACITY: usize = 4096;
pub(crate) const ME_WRITER_BYTE_PERMIT_UNIT_BYTES: usize = 16 * 1024;
pub(crate) const ME_WRITER_FRAME_OVERHEAD_RESERVE_BYTES: usize = 256;
const DEFAULT_ME_WRITER_BYTE_BUDGET_BYTES: usize =
32 * 1024 * 1024 + ME_WRITER_BYTE_PERMIT_UNIT_BYTES;
const DEFAULT_ME_ROUTE_CHANNEL_CAPACITY: usize = 768; const DEFAULT_ME_ROUTE_CHANNEL_CAPACITY: usize = 768;
const DEFAULT_ME_C2ME_CHANNEL_CAPACITY: usize = 1024; const DEFAULT_ME_C2ME_CHANNEL_CAPACITY: usize = 1024;
const DEFAULT_ME_READER_ROUTE_DATA_WAIT_MS: u64 = 2; const DEFAULT_ME_READER_ROUTE_DATA_WAIT_MS: u64 = 2;
@@ -35,6 +39,8 @@ const DEFAULT_ME_QUOTA_SOFT_OVERSHOOT_BYTES: u64 = 64 * 1024;
const DEFAULT_ME_D2C_FRAME_BUF_SHRINK_THRESHOLD_BYTES: usize = 256 * 1024; const DEFAULT_ME_D2C_FRAME_BUF_SHRINK_THRESHOLD_BYTES: usize = 256 * 1024;
const DEFAULT_DIRECT_RELAY_COPY_BUF_C2S_BYTES: usize = 64 * 1024; const DEFAULT_DIRECT_RELAY_COPY_BUF_C2S_BYTES: usize = 64 * 1024;
const DEFAULT_DIRECT_RELAY_COPY_BUF_S2C_BYTES: usize = 256 * 1024; const DEFAULT_DIRECT_RELAY_COPY_BUF_S2C_BYTES: usize = 256 * 1024;
pub(crate) const DIRECT_RELAY_BUFFER_BUDGET_UNIT_BYTES: usize = 4 * 1024;
const DEFAULT_DIRECT_RELAY_BUFFER_BUDGET_MAX_BYTES: usize = 0;
const DEFAULT_ME_WRITER_PICK_SAMPLE_SIZE: u8 = 3; const DEFAULT_ME_WRITER_PICK_SAMPLE_SIZE: u8 = 3;
const DEFAULT_ME_HEALTH_INTERVAL_MS_UNHEALTHY: u64 = 1000; const DEFAULT_ME_HEALTH_INTERVAL_MS_UNHEALTHY: u64 = 1000;
const DEFAULT_ME_HEALTH_INTERVAL_MS_HEALTHY: u64 = 3000; const DEFAULT_ME_HEALTH_INTERVAL_MS_HEALTHY: u64 = 3000;
@@ -54,9 +60,14 @@ const DEFAULT_CONNTRACK_CONTROL_ENABLED: bool = true;
const DEFAULT_CONNTRACK_PRESSURE_HIGH_WATERMARK_PCT: u8 = 85; const DEFAULT_CONNTRACK_PRESSURE_HIGH_WATERMARK_PCT: u8 = 85;
const DEFAULT_CONNTRACK_PRESSURE_LOW_WATERMARK_PCT: u8 = 70; const DEFAULT_CONNTRACK_PRESSURE_LOW_WATERMARK_PCT: u8 = 70;
const DEFAULT_CONNTRACK_DELETE_BUDGET_PER_SEC: u64 = 4096; const DEFAULT_CONNTRACK_DELETE_BUDGET_PER_SEC: u64 = 4096;
const DEFAULT_SYNLIMIT_SECONDS: u32 = 1; const DEFAULT_SYNLIMIT_SECONDS: u32 = 60;
const DEFAULT_SYNLIMIT_HITCOUNT: u32 = 1; const DEFAULT_SYNLIMIT_HITCOUNT: u32 = 48;
const DEFAULT_SYNLIMIT_BURST: u32 = 2; const DEFAULT_SYNLIMIT_BURST: u32 = 1;
const DEFAULT_SYNLIMIT_IOS_SECONDS: u32 = 1;
const DEFAULT_SYNLIMIT_IOS_HITCOUNT: u32 = 12;
const DEFAULT_SYNLIMIT_IOS_BURST: u32 = 24;
const DEFAULT_SYNLIMIT_HASHLIMIT_EXPIRE_MS: u32 = 60_000;
const DEFAULT_SYNLIMIT_HASHLIMIT_SIZE: u32 = 32_768;
const DEFAULT_UPSTREAM_CONNECT_RETRY_ATTEMPTS: u32 = 2; const DEFAULT_UPSTREAM_CONNECT_RETRY_ATTEMPTS: u32 = 2;
const DEFAULT_UPSTREAM_UNHEALTHY_FAIL_THRESHOLD: u32 = 5; const DEFAULT_UPSTREAM_UNHEALTHY_FAIL_THRESHOLD: u32 = 5;
const DEFAULT_UPSTREAM_CONNECT_BUDGET_MS: u64 = 3000; const DEFAULT_UPSTREAM_CONNECT_BUDGET_MS: u64 = 3000;
@@ -258,6 +269,26 @@ pub(crate) fn default_synlimit_burst() -> u32 {
DEFAULT_SYNLIMIT_BURST DEFAULT_SYNLIMIT_BURST
} }
pub(crate) fn default_synlimit_ios_seconds() -> u32 {
DEFAULT_SYNLIMIT_IOS_SECONDS
}
pub(crate) fn default_synlimit_ios_hitcount() -> u32 {
DEFAULT_SYNLIMIT_IOS_HITCOUNT
}
pub(crate) fn default_synlimit_ios_burst() -> u32 {
DEFAULT_SYNLIMIT_IOS_BURST
}
pub(crate) fn default_synlimit_hashlimit_expire_ms() -> u32 {
DEFAULT_SYNLIMIT_HASHLIMIT_EXPIRE_MS
}
pub(crate) fn default_synlimit_hashlimit_size() -> u32 {
DEFAULT_SYNLIMIT_HASHLIMIT_SIZE
}
pub(crate) fn default_prefer_4() -> u8 { pub(crate) fn default_prefer_4() -> u8 {
4 4
} }
@@ -430,6 +461,18 @@ pub(crate) fn default_me_writer_cmd_channel_capacity() -> usize {
DEFAULT_ME_WRITER_CMD_CHANNEL_CAPACITY DEFAULT_ME_WRITER_CMD_CHANNEL_CAPACITY
} }
pub(crate) fn default_me_writer_byte_budget_bytes() -> usize {
DEFAULT_ME_WRITER_BYTE_BUDGET_BYTES
}
pub(crate) fn minimum_me_writer_byte_budget_bytes(max_client_frame: usize) -> usize {
max_client_frame
.saturating_mul(2)
.saturating_add(ME_WRITER_FRAME_OVERHEAD_RESERVE_BYTES)
.div_ceil(ME_WRITER_BYTE_PERMIT_UNIT_BYTES)
.saturating_mul(ME_WRITER_BYTE_PERMIT_UNIT_BYTES)
}
pub(crate) fn default_me_route_channel_capacity() -> usize { pub(crate) fn default_me_route_channel_capacity() -> usize {
DEFAULT_ME_ROUTE_CHANNEL_CAPACITY DEFAULT_ME_ROUTE_CHANNEL_CAPACITY
} }
@@ -474,6 +517,10 @@ pub(crate) fn default_direct_relay_copy_buf_s2c_bytes() -> usize {
DEFAULT_DIRECT_RELAY_COPY_BUF_S2C_BYTES DEFAULT_DIRECT_RELAY_COPY_BUF_S2C_BYTES
} }
pub(crate) fn default_direct_relay_buffer_budget_max_bytes() -> usize {
DEFAULT_DIRECT_RELAY_BUFFER_BUDGET_MAX_BYTES
}
pub(crate) fn default_me_writer_pick_sample_size() -> u8 { pub(crate) fn default_me_writer_pick_sample_size() -> u8 {
DEFAULT_ME_WRITER_PICK_SAMPLE_SIZE DEFAULT_ME_WRITER_PICK_SAMPLE_SIZE
} }
+71 -6
View File
@@ -36,8 +36,8 @@ use tracing::{error, info, warn};
use super::load::{LoadedConfig, ProxyConfig}; use super::load::{LoadedConfig, ProxyConfig};
use crate::config::{ use crate::config::{
ListenerConfig, LogLevel, MeBindStaleMode, MeFloorMode, MeSocksKdfPolicy, MeTelemetryLevel, CidrRateLimitKey, ListenerConfig, LogLevel, MeBindStaleMode, MeFloorMode, MeSocksKdfPolicy,
MeWriterPickMode, SynLimitMode, MeTelemetryLevel, MeWriterPickMode, SynLimitMode,
}; };
const HOT_RELOAD_DEBOUNCE: Duration = Duration::from_millis(50); const HOT_RELOAD_DEBOUNCE: Duration = Duration::from_millis(50);
@@ -128,8 +128,7 @@ pub struct HotFields {
pub user_expirations: std::collections::HashMap<String, chrono::DateTime<chrono::Utc>>, pub user_expirations: std::collections::HashMap<String, chrono::DateTime<chrono::Utc>>,
pub user_data_quota: std::collections::HashMap<String, u64>, pub user_data_quota: std::collections::HashMap<String, u64>,
pub user_rate_limits: std::collections::HashMap<String, crate::config::RateLimitBps>, pub user_rate_limits: std::collections::HashMap<String, crate::config::RateLimitBps>,
pub cidr_rate_limits: pub cidr_rate_limits: std::collections::HashMap<CidrRateLimitKey, crate::config::RateLimitBps>,
std::collections::HashMap<ipnetwork::IpNetwork, crate::config::RateLimitBps>,
pub user_max_unique_ips: std::collections::HashMap<String, usize>, pub user_max_unique_ips: std::collections::HashMap<String, usize>,
pub user_max_unique_ips_global_each: usize, pub user_max_unique_ips_global_each: usize,
pub user_max_unique_ips_mode: crate::config::UserMaxUniqueIpsMode, pub user_max_unique_ips_mode: crate::config::UserMaxUniqueIpsMode,
@@ -145,6 +144,11 @@ pub struct ListenerSynLimitHotFields {
pub synlimit_seconds: u32, pub synlimit_seconds: u32,
pub synlimit_hitcount: u32, pub synlimit_hitcount: u32,
pub synlimit_burst: u32, pub synlimit_burst: u32,
pub synlimit_ios_seconds: u32,
pub synlimit_ios_hitcount: u32,
pub synlimit_ios_burst: u32,
pub synlimit_hashlimit_expire_ms: u32,
pub synlimit_hashlimit_size: u32,
} }
impl HotFields { impl HotFields {
@@ -293,6 +297,11 @@ impl ListenerSynLimitHotFields {
synlimit_seconds: listener.synlimit_seconds, synlimit_seconds: listener.synlimit_seconds,
synlimit_hitcount: listener.synlimit_hitcount, synlimit_hitcount: listener.synlimit_hitcount,
synlimit_burst: listener.synlimit_burst, synlimit_burst: listener.synlimit_burst,
synlimit_ios_seconds: listener.synlimit_ios_seconds,
synlimit_ios_hitcount: listener.synlimit_ios_hitcount,
synlimit_ios_burst: listener.synlimit_ios_burst,
synlimit_hashlimit_expire_ms: listener.synlimit_hashlimit_expire_ms,
synlimit_hashlimit_size: listener.synlimit_hashlimit_size,
} }
} }
} }
@@ -620,6 +629,11 @@ fn overlay_listener_synlimit_fields(old: &mut [ListenerConfig], new: &[ListenerC
old_listener.synlimit_seconds = new_listener.synlimit_seconds; old_listener.synlimit_seconds = new_listener.synlimit_seconds;
old_listener.synlimit_hitcount = new_listener.synlimit_hitcount; old_listener.synlimit_hitcount = new_listener.synlimit_hitcount;
old_listener.synlimit_burst = new_listener.synlimit_burst; old_listener.synlimit_burst = new_listener.synlimit_burst;
old_listener.synlimit_ios_seconds = new_listener.synlimit_ios_seconds;
old_listener.synlimit_ios_hitcount = new_listener.synlimit_ios_hitcount;
old_listener.synlimit_ios_burst = new_listener.synlimit_ios_burst;
old_listener.synlimit_hashlimit_expire_ms = new_listener.synlimit_hashlimit_expire_ms;
old_listener.synlimit_hashlimit_size = new_listener.synlimit_hashlimit_size;
} }
} }
@@ -1389,11 +1403,13 @@ fn reload_config(
/// `detected_ip_v4` / `detected_ip_v6` are the IPs discovered during the /// `detected_ip_v4` / `detected_ip_v6` are the IPs discovered during the
/// startup probe — used when generating proxy links for newly added users, /// startup probe — used when generating proxy links for newly added users,
/// matching the same logic as the startup output. /// matching the same logic as the startup output.
/// The watcher releases its notify and signal resources when `cancellation` fires.
pub fn spawn_config_watcher( pub fn spawn_config_watcher(
config_path: PathBuf, config_path: PathBuf,
initial: Arc<ProxyConfig>, initial: Arc<ProxyConfig>,
detected_ip_v4: Option<IpAddr>, detected_ip_v4: Option<IpAddr>,
detected_ip_v6: Option<IpAddr>, detected_ip_v6: Option<IpAddr>,
cancellation: tokio_util::sync::CancellationToken,
) -> (watch::Receiver<Arc<ProxyConfig>>, watch::Receiver<LogLevel>) { ) -> (watch::Receiver<Arc<ProxyConfig>>, watch::Receiver<LogLevel>) {
let initial_level = initial.general.log_level.clone(); let initial_level = initial.general.log_level.clone();
let (config_tx, config_rx) = watch::channel(initial); let (config_tx, config_rx) = watch::channel(initial);
@@ -1501,10 +1517,14 @@ pub fn spawn_config_watcher(
_ = sighup.recv() => { _ = sighup.recv() => {
info!("SIGHUP received — reloading {:?}", config_path); info!("SIGHUP received — reloading {:?}", config_path);
} }
_ = cancellation.cancelled() => break,
} }
#[cfg(not(unix))] #[cfg(not(unix))]
if notify_rx.recv().await.is_none() { tokio::select! {
break; msg = notify_rx.recv() => {
if msg.is_none() { break; }
}
_ = cancellation.cancelled() => break,
} }
// Debounce: drain extra events that arrive within a short quiet window. // Debounce: drain extra events that arrive within a short quiet window.
@@ -1710,6 +1730,51 @@ mod tests {
assert!(!config_equal(&applied, &new)); assert!(!config_equal(&applied, &new));
} }
#[test]
fn listener_synlimit_extended_fields_are_hot() {
let mut old = sample_config();
old.server.listeners.push(ListenerConfig {
ip: "0.0.0.0".parse().unwrap(),
port: Some(443),
client_mss: None,
synlimit: SynLimitMode::Iptables,
synlimit_seconds: 60,
synlimit_hitcount: 48,
synlimit_burst: 1,
synlimit_ios_seconds: 1,
synlimit_ios_hitcount: 12,
synlimit_ios_burst: 24,
synlimit_hashlimit_expire_ms: 60_000,
synlimit_hashlimit_size: 32_768,
announce: None,
announce_ip: None,
proxy_protocol: None,
reuse_allow: false,
});
let mut new = old.clone();
new.server.port = 8443;
new.server.listeners[0].synlimit_seconds = 120;
new.server.listeners[0].synlimit_hitcount = 96;
new.server.listeners[0].synlimit_burst = 2;
new.server.listeners[0].synlimit_ios_seconds = 2;
new.server.listeners[0].synlimit_ios_hitcount = 18;
new.server.listeners[0].synlimit_ios_burst = 36;
new.server.listeners[0].synlimit_hashlimit_expire_ms = 90_000;
new.server.listeners[0].synlimit_hashlimit_size = 65_536;
let applied = overlay_hot_fields(&old, &new);
let listener = &applied.server.listeners[0];
assert_eq!(applied.server.port, old.server.port);
assert_eq!(listener.synlimit_seconds, 120);
assert_eq!(listener.synlimit_hitcount, 96);
assert_eq!(listener.synlimit_burst, 2);
assert_eq!(listener.synlimit_ios_seconds, 2);
assert_eq!(listener.synlimit_ios_hitcount, 18);
assert_eq!(listener.synlimit_ios_burst, 36);
assert_eq!(listener.synlimit_hashlimit_expire_ms, 90_000);
assert_eq!(listener.synlimit_hashlimit_size, 65_536);
}
#[test] #[test]
fn reload_applies_hot_change_on_first_observed_snapshot() { fn reload_applies_hot_change_on_first_observed_snapshot() {
let initial_tag = "11111111111111111111111111111111"; let initial_tag = "11111111111111111111111111111111";
+123 -3339
View File
File diff suppressed because it is too large Load Diff
+60
View File
@@ -0,0 +1,60 @@
use std::collections::BTreeSet;
use std::hash::{DefaultHasher, Hash, Hasher};
use std::path::{Path, PathBuf};
use crate::error::{ProxyError, Result};
pub(super) fn normalize_config_path(path: &Path) -> PathBuf {
path.canonicalize().unwrap_or_else(|_| {
if path.is_absolute() {
path.to_path_buf()
} else {
std::env::current_dir()
.map(|cwd| cwd.join(path))
.unwrap_or_else(|_| path.to_path_buf())
}
})
}
pub(super) fn hash_rendered_snapshot(rendered: &str) -> u64 {
let mut hasher = DefaultHasher::new();
rendered.hash(&mut hasher);
hasher.finish()
}
pub(super) fn preprocess_includes(
content: &str,
base_dir: &Path,
depth: u8,
source_files: &mut BTreeSet<PathBuf>,
) -> Result<String> {
if depth > 10 {
return Err(ProxyError::Config("Include depth > 10".into()));
}
let mut output = String::with_capacity(content.len());
for line in content.lines() {
let trimmed = line.trim();
if let Some(rest) = trimmed.strip_prefix("include") {
let rest = rest.trim();
if let Some(rest) = rest.strip_prefix('=') {
let path_str = rest.trim().trim_matches('"');
let resolved = base_dir.join(path_str);
source_files.insert(normalize_config_path(&resolved));
let included = std::fs::read_to_string(&resolved)
.map_err(|e| ProxyError::Config(e.to_string()))?;
let included_dir = resolved.parent().unwrap_or(base_dir);
output.push_str(&preprocess_includes(
&included,
included_dir,
depth + 1,
source_files,
)?);
output.push('\n');
continue;
}
}
output.push_str(line);
output.push('\n');
}
Ok(output)
}
+115
View File
@@ -0,0 +1,115 @@
use crate::error::{ProxyError, Result};
use tracing::warn;
pub(super) fn is_valid_tls_domain_name(domain: &str) -> bool {
!domain.is_empty()
&& !domain
.chars()
.any(|ch| ch.is_whitespace() || matches!(ch, '/' | '\\'))
}
pub(super) fn normalize_domain_to_ascii(domain: &str, field: &str) -> Result<String> {
let domain = domain.trim();
if !is_valid_tls_domain_name(domain) {
return Err(ProxyError::Config(format!(
"Invalid {field}: '{}'. Must be a valid domain name",
domain
)));
}
let parsed = url::Url::parse(&format!("https://{domain}/")).map_err(|error| {
ProxyError::Config(format!(
"Invalid {field}: '{}'. IDNA conversion failed: {error}",
domain
))
})?;
let host = parsed.host_str().ok_or_else(|| {
ProxyError::Config(format!("Invalid {field}: '{}'. Host is empty", domain))
})?;
Ok(host.to_ascii_lowercase())
}
pub(super) fn normalize_mask_host_to_ascii(host: &str, field: &str) -> Result<String> {
let host = host.trim();
if host.starts_with('[') && host.ends_with(']') {
let inner = &host[1..host.len() - 1];
let ip = inner.parse::<std::net::IpAddr>().map_err(|_| {
ProxyError::Config(format!(
"Invalid {field}: '{}'. IPv6 literal is invalid",
host
))
})?;
return match ip {
std::net::IpAddr::V6(v6) => Ok(format!("[{v6}]")),
std::net::IpAddr::V4(v4) => Ok(v4.to_string()),
};
}
if let Ok(ip) = host.parse::<std::net::IpAddr>() {
return match ip {
std::net::IpAddr::V4(v4) => Ok(v4.to_string()),
std::net::IpAddr::V6(v6) => Ok(format!("[{v6}]")),
};
}
normalize_domain_to_ascii(host, field)
}
pub(super) fn parse_exclusive_mask_target(target: &str) -> Option<(&str, u16)> {
let target = target.trim();
if target.is_empty() {
return None;
}
if target.starts_with('[') {
let end = target.find(']')?;
if target.get(end + 1..end + 2)? != ":" {
return None;
}
let host = &target[..=end];
let port = target[end + 2..].parse::<u16>().ok()?;
return (port > 0).then_some((host, port));
}
let (host, port) = target.rsplit_once(':')?;
if host.is_empty() || host.contains(':') {
return None;
}
let port = port.parse::<u16>().ok()?;
(port > 0).then_some((host, port))
}
pub(super) fn normalize_exclusive_mask_target(target: &str, field: &str) -> Result<String> {
let (host, port) = parse_exclusive_mask_target(target).ok_or_else(|| {
ProxyError::Config(format!(
"Invalid {field}: '{}'. Expected host:port with port > 0",
target
))
})?;
let host = normalize_mask_host_to_ascii(host, field)?;
Ok(format!("{host}:{port}"))
}
pub(super) fn push_unique_nonempty(target: &mut Vec<String>, value: String) {
let trimmed = value.trim();
if trimmed.is_empty() {
return;
}
if !target.iter().any(|existing| existing == trimmed) {
target.push(trimmed.to_string());
}
}
pub(super) fn is_valid_ad_tag(tag: &str) -> bool {
tag.len() == 32 && tag.chars().all(|ch| ch.is_ascii_hexdigit())
}
pub(super) fn sanitize_ad_tag(ad_tag: &mut Option<String>) {
let Some(tag) = ad_tag.as_ref() else {
return;
};
if !is_valid_ad_tag(tag) {
warn!("Invalid general.ad_tag value, expected exactly 32 hex chars; ad_tag is disabled");
*ad_tag = None;
}
}
+112
View File
@@ -0,0 +1,112 @@
use std::collections::HashMap;
use std::collections::hash_map::DefaultHasher;
use std::hash::Hasher;
use crate::error::{ProxyError, Result};
const ACCESS_SECRET_BYTES: usize = 16;
/// Precomputed, immutable user authentication data used by handshake hot paths.
#[derive(Debug, Clone, Default)]
pub(crate) struct UserAuthSnapshot {
entries: Vec<UserAuthEntry>,
by_name: HashMap<String, u32>,
sni_index: HashMap<u64, Vec<u32>>,
sni_initial_index: HashMap<u8, Vec<u32>>,
}
#[derive(Debug, Clone)]
pub(crate) struct UserAuthEntry {
pub(crate) user: String,
pub(crate) secret: [u8; ACCESS_SECRET_BYTES],
}
impl UserAuthSnapshot {
pub(super) fn from_users(users: &HashMap<String, String>) -> Result<Self> {
let mut entries = Vec::with_capacity(users.len());
let mut by_name = HashMap::with_capacity(users.len());
let mut sni_index = HashMap::with_capacity(users.len());
let mut sni_initial_index = HashMap::with_capacity(users.len());
for (user, secret_hex) in users {
let decoded = hex::decode(secret_hex).map_err(|_| ProxyError::InvalidSecret {
user: user.clone(),
reason: "Must be 32 hex characters".to_string(),
})?;
if decoded.len() != ACCESS_SECRET_BYTES {
return Err(ProxyError::InvalidSecret {
user: user.clone(),
reason: "Must be 32 hex characters".to_string(),
});
}
let user_id = u32::try_from(entries.len()).map_err(|_| {
ProxyError::Config("Too many users for runtime auth snapshot".to_string())
})?;
let mut secret = [0u8; ACCESS_SECRET_BYTES];
secret.copy_from_slice(&decoded);
entries.push(UserAuthEntry {
user: user.clone(),
secret,
});
by_name.insert(user.clone(), user_id);
sni_index
.entry(Self::sni_lookup_hash(user))
.or_insert_with(Vec::new)
.push(user_id);
if let Some(initial) = user
.as_bytes()
.first()
.map(|byte| byte.to_ascii_lowercase())
{
sni_initial_index
.entry(initial)
.or_insert_with(Vec::new)
.push(user_id);
}
}
Ok(Self {
entries,
by_name,
sni_index,
sni_initial_index,
})
}
pub(crate) fn entries(&self) -> &[UserAuthEntry] {
&self.entries
}
pub(crate) fn user_id_by_name(&self, user: &str) -> Option<u32> {
self.by_name.get(user).copied()
}
pub(crate) fn entry_by_id(&self, user_id: u32) -> Option<&UserAuthEntry> {
let idx = usize::try_from(user_id).ok()?;
self.entries.get(idx)
}
pub(crate) fn sni_candidates(&self, sni: &str) -> Option<&[u32]> {
self.sni_index
.get(&Self::sni_lookup_hash(sni))
.map(Vec::as_slice)
}
pub(crate) fn sni_initial_candidates(&self, sni: &str) -> Option<&[u32]> {
let initial = sni
.as_bytes()
.first()
.map(|byte| byte.to_ascii_lowercase())?;
self.sni_initial_index.get(&initial).map(Vec::as_slice)
}
fn sni_lookup_hash(value: &str) -> u64 {
let mut hasher = DefaultHasher::new();
for byte in value.bytes() {
hasher.write_u8(byte.to_ascii_lowercase());
}
hasher.finish()
}
}
+695
View File
@@ -0,0 +1,695 @@
use tracing::warn;
use crate::error::{ProxyError, Result};
const TOP_LEVEL_CONFIG_KEYS: &[&str] = &[
"general",
"logging",
"network",
"server",
"timeouts",
"censorship",
"access",
"upstreams",
"show_link",
"dc_overrides",
"default_dc",
"beobachten",
"beobachten_minutes",
"beobachten_flush_secs",
"beobachten_file",
"include",
];
const GENERAL_CONFIG_KEYS: &[&str] = &[
"data_path",
"quota_state_path",
"config_strict",
"modes",
"prefer_ipv6",
"fast_mode",
"use_middle_proxy",
"proxy_secret_path",
"proxy_secret_url",
"proxy_config_v4_cache_path",
"proxy_config_v4_url",
"proxy_config_v6_cache_path",
"proxy_config_v6_url",
"ad_tag",
"middle_proxy_nat_ip",
"middle_proxy_nat_probe",
"middle_proxy_nat_stun",
"middle_proxy_nat_stun_servers",
"stun_nat_probe_concurrency",
"middle_proxy_pool_size",
"middle_proxy_warm_standby",
"me_init_retry_attempts",
"me2dc_fallback",
"me2dc_fast",
"me_keepalive_enabled",
"me_keepalive_interval_secs",
"me_keepalive_jitter_secs",
"me_keepalive_payload_random",
"rpc_proxy_req_every",
"me_writer_cmd_channel_capacity",
"me_writer_byte_budget_bytes",
"me_route_channel_capacity",
"me_c2me_channel_capacity",
"me_c2me_send_timeout_ms",
"me_reader_route_data_wait_ms",
"me_d2c_flush_batch_max_frames",
"me_d2c_flush_batch_max_bytes",
"me_d2c_flush_batch_max_delay_us",
"me_d2c_ack_flush_immediate",
"me_quota_soft_overshoot_bytes",
"me_d2c_frame_buf_shrink_threshold_bytes",
"direct_relay_copy_buf_c2s_bytes",
"direct_relay_copy_buf_s2c_bytes",
"direct_relay_buffer_budget_max_bytes",
"crypto_pending_buffer",
"max_client_frame",
"desync_all_full",
"beobachten",
"beobachten_minutes",
"beobachten_flush_secs",
"beobachten_file",
"hardswap",
"me_warmup_stagger_enabled",
"me_warmup_step_delay_ms",
"me_warmup_step_jitter_ms",
"me_reconnect_max_concurrent_per_dc",
"me_reconnect_backoff_base_ms",
"me_reconnect_backoff_cap_ms",
"me_reconnect_fast_retry_count",
"me_single_endpoint_shadow_writers",
"me_single_endpoint_outage_mode_enabled",
"me_single_endpoint_outage_disable_quarantine",
"me_single_endpoint_outage_backoff_min_ms",
"me_single_endpoint_outage_backoff_max_ms",
"me_single_endpoint_shadow_rotate_every_secs",
"me_floor_mode",
"me_adaptive_floor_idle_secs",
"me_adaptive_floor_min_writers_single_endpoint",
"me_adaptive_floor_min_writers_multi_endpoint",
"me_adaptive_floor_recover_grace_secs",
"me_adaptive_floor_writers_per_core_total",
"me_adaptive_floor_cpu_cores_override",
"me_adaptive_floor_max_extra_writers_single_per_core",
"me_adaptive_floor_max_extra_writers_multi_per_core",
"me_adaptive_floor_max_active_writers_per_core",
"me_adaptive_floor_max_warm_writers_per_core",
"me_adaptive_floor_max_active_writers_global",
"me_adaptive_floor_max_warm_writers_global",
"upstream_connect_retry_attempts",
"upstream_connect_retry_backoff_ms",
"upstream_connect_budget_ms",
"tg_connect",
"upstream_unhealthy_fail_threshold",
"upstream_connect_failfast_hard_errors",
"stun_iface_mismatch_ignore",
"unknown_dc_log_path",
"unknown_dc_file_log_enabled",
"log_level",
"disable_colors",
"telemetry",
"me_socks_kdf_policy",
"me_route_backpressure_enabled",
"me_route_fairshare_enabled",
"me_route_backpressure_base_timeout_ms",
"me_route_backpressure_high_timeout_ms",
"me_route_backpressure_high_watermark_pct",
"me_health_interval_ms_unhealthy",
"me_health_interval_ms_healthy",
"me_admission_poll_ms",
"me_warn_rate_limit_ms",
"me_route_no_writer_mode",
"me_route_no_writer_wait_ms",
"me_route_hybrid_max_wait_ms",
"me_route_blocking_send_timeout_ms",
"me_route_inline_recovery_attempts",
"me_route_inline_recovery_wait_ms",
"links",
"fast_mode_min_tls_record",
"update_every",
"me_reinit_every_secs",
"me_hardswap_warmup_delay_min_ms",
"me_hardswap_warmup_delay_max_ms",
"me_hardswap_warmup_extra_passes",
"me_hardswap_warmup_pass_backoff_base_ms",
"me_config_stable_snapshots",
"me_config_apply_cooldown_secs",
"me_snapshot_require_http_2xx",
"me_snapshot_reject_empty_map",
"me_snapshot_min_proxy_for_lines",
"proxy_secret_stable_snapshots",
"proxy_secret_rotate_runtime",
"me_secret_atomic_snapshot",
"proxy_secret_len_max",
"me_pool_drain_ttl_secs",
"me_instadrain",
"me_pool_drain_threshold",
"me_pool_drain_soft_evict_enabled",
"me_pool_drain_soft_evict_grace_secs",
"me_pool_drain_soft_evict_per_writer",
"me_pool_drain_soft_evict_budget_per_core",
"me_pool_drain_soft_evict_cooldown_ms",
"me_bind_stale_mode",
"me_bind_stale_ttl_secs",
"me_pool_min_fresh_ratio",
"me_reinit_drain_timeout_secs",
"proxy_secret_auto_reload_secs",
"proxy_config_auto_reload_secs",
"me_reinit_singleflight",
"me_reinit_trigger_channel",
"me_reinit_coalesce_window_ms",
"me_deterministic_writer_sort",
"me_writer_pick_mode",
"me_writer_pick_sample_size",
"ntp_check",
"ntp_servers",
"auto_degradation_enabled",
"degradation_min_unavailable_dc_groups",
"rst_on_close",
];
const NETWORK_CONFIG_KEYS: &[&str] = &[
"ipv4",
"ipv6",
"prefer",
"multipath",
"stun_use",
"stun_servers",
"stun_tcp_fallback",
"http_ip_detect_urls",
"cache_public_ip_path",
"dns_overrides",
];
const SERVER_CONFIG_KEYS: &[&str] = &[
"port",
"listen_addr_ipv4",
"listen_addr_ipv6",
"listen_unix_sock",
"listen_unix_sock_perm",
"listen_tcp",
"client_mss",
"client_mss_bulk",
"proxy_protocol",
"proxy_protocol_header_timeout_ms",
"proxy_protocol_trusted_cidrs",
"metrics_port",
"metrics_listen",
"metrics_whitelist",
"api",
"admin_api",
"listeners",
"listen_backlog",
"max_connections",
"accept_permit_timeout_ms",
"conntrack_control",
];
const API_CONFIG_KEYS: &[&str] = &[
"enabled",
"listen",
"whitelist",
"gray_action",
"auth_header",
"request_body_limit_bytes",
"minimal_runtime_enabled",
"minimal_runtime_cache_ttl_ms",
"runtime_edge_enabled",
"runtime_edge_cache_ttl_ms",
"runtime_edge_top_n",
"runtime_edge_events_capacity",
"read_only",
];
const CONNTRACK_CONTROL_CONFIG_KEYS: &[&str] = &[
"inline_conntrack_control",
"mode",
"backend",
"profile",
"hybrid_listener_ips",
"pressure_high_watermark_pct",
"pressure_low_watermark_pct",
"delete_budget_per_sec",
];
const LISTENER_CONFIG_KEYS: &[&str] = &[
"ip",
"port",
"client_mss",
"synlimit",
"synlimit_seconds",
"synlimit_hitcount",
"synlimit_burst",
"synlimit_ios_seconds",
"synlimit_ios_hitcount",
"synlimit_ios_burst",
"synlimit_hashlimit_expire_ms",
"synlimit_hashlimit_size",
"announce",
"announce_ip",
"proxy_protocol",
"reuse_allow",
];
const TIMEOUTS_CONFIG_KEYS: &[&str] = &[
"client_first_byte_idle_secs",
"client_handshake",
"relay_idle_policy_v2_enabled",
"relay_client_idle_soft_secs",
"relay_client_idle_hard_secs",
"relay_idle_grace_after_downstream_activity_secs",
"client_keepalive",
"client_ack",
"me_one_retry",
"me_one_timeout_ms",
];
const CENSORSHIP_CONFIG_KEYS: &[&str] = &[
"tls_domain",
"tls_domains",
"unknown_sni_action",
"tls_fetch_scope",
"tls_fetch",
"mask",
"mask_dynamic",
"mask_host",
"mask_port",
"exclusive_mask",
"mask_unix_sock",
"fake_cert_len",
"tls_emulation",
"tls_front_dir",
"server_hello_delay_min_ms",
"server_hello_delay_max_ms",
"tls_new_session_tickets",
"serverhello_compact",
"tls_full_cert_ttl_secs",
"alpn_enforce",
"mask_proxy_protocol",
"mask_shape_hardening",
"mask_shape_hardening_aggressive_mode",
"mask_shape_bucket_floor_bytes",
"mask_shape_bucket_cap_bytes",
"mask_shape_above_cap_blur",
"mask_shape_above_cap_blur_max_bytes",
"mask_relay_max_bytes",
"mask_relay_timeout_ms",
"mask_relay_idle_timeout_ms",
"mask_classifier_prefetch_timeout_ms",
"mask_timing_normalization_enabled",
"mask_timing_normalization_floor_ms",
"mask_timing_normalization_ceiling_ms",
];
const TLS_FETCH_CONFIG_KEYS: &[&str] = &[
"profiles",
"strict_route",
"attempt_timeout_ms",
"total_budget_ms",
"grease_enabled",
"deterministic",
"profile_cache_ttl_secs",
];
const ACCESS_CONFIG_KEYS: &[&str] = &[
"users",
"user_enabled",
"user_ad_tags",
"user_max_tcp_conns",
"user_max_tcp_conns_global_each",
"user_expirations",
"user_data_quota",
"user_rate_limits",
"cidr_rate_limits",
"user_max_unique_ips",
"user_max_unique_ips_global_each",
"user_max_unique_ips_mode",
"user_max_unique_ips_window_secs",
"replay_check_len",
"replay_window_secs",
"ignore_time_skew",
];
const RATE_LIMIT_BPS_CONFIG_KEYS: &[&str] = &["up_bps", "down_bps"];
const UPSTREAM_CONFIG_KEYS: &[&str] = &[
"type",
"interface",
"bind_addresses",
"bindtodevice",
"force_bind",
"address",
"user_id",
"username",
"password",
"url",
"weight",
"enabled",
"scopes",
"ipv4",
"ipv6",
];
const PROXY_MODES_CONFIG_KEYS: &[&str] = &["classic", "secure", "tls"];
const TELEMETRY_CONFIG_KEYS: &[&str] = &["core_enabled", "user_enabled", "me_level"];
const LINKS_CONFIG_KEYS: &[&str] = &["show", "public_host", "public_port"];
const LOGGING_CONFIG_KEYS: &[&str] = &[
"destination",
"path",
"rotation",
"max_size_bytes",
"max_files",
"max_age_secs",
];
#[derive(Debug)]
struct UnknownConfigKey {
path: String,
suggestion: Option<String>,
}
fn table_at<'a>(value: &'a toml::Value, path: &[&str]) -> Option<&'a toml::Table> {
let mut current = value;
for segment in path {
current = current.get(*segment)?;
}
current.as_table()
}
fn is_strict_config(parsed_toml: &toml::Value) -> bool {
table_at(parsed_toml, &["general"])
.and_then(|table| table.get("config_strict"))
.and_then(toml::Value::as_bool)
.unwrap_or(false)
}
fn known_config_keys_for_suggestion() -> Vec<&'static str> {
let mut keys = Vec::new();
for group in [
TOP_LEVEL_CONFIG_KEYS,
GENERAL_CONFIG_KEYS,
NETWORK_CONFIG_KEYS,
SERVER_CONFIG_KEYS,
API_CONFIG_KEYS,
CONNTRACK_CONTROL_CONFIG_KEYS,
LISTENER_CONFIG_KEYS,
TIMEOUTS_CONFIG_KEYS,
CENSORSHIP_CONFIG_KEYS,
TLS_FETCH_CONFIG_KEYS,
ACCESS_CONFIG_KEYS,
RATE_LIMIT_BPS_CONFIG_KEYS,
UPSTREAM_CONFIG_KEYS,
PROXY_MODES_CONFIG_KEYS,
TELEMETRY_CONFIG_KEYS,
LINKS_CONFIG_KEYS,
LOGGING_CONFIG_KEYS,
] {
keys.extend_from_slice(group);
}
keys
}
fn levenshtein_distance(a: &str, b: &str) -> usize {
let b_chars: Vec<char> = b.chars().collect();
let mut prev: Vec<usize> = (0..=b_chars.len()).collect();
let mut curr = vec![0usize; b_chars.len() + 1];
for (i, ca) in a.chars().enumerate() {
curr[0] = i + 1;
for (j, cb) in b_chars.iter().enumerate() {
let replace = if ca == *cb { prev[j] } else { prev[j] + 1 };
curr[j + 1] = (prev[j + 1] + 1).min(curr[j] + 1).min(replace);
}
std::mem::swap(&mut prev, &mut curr);
}
prev[b_chars.len()]
}
fn unknown_key_suggestion(key: &str, known_keys: &[&'static str]) -> Option<String> {
let normalized = key.to_ascii_lowercase();
let mut best: Option<(&str, usize)> = None;
for known in known_keys {
let distance = levenshtein_distance(&normalized, known);
let is_better = match best {
Some((_, best_distance)) => distance < best_distance,
None => true,
};
if distance <= 4 && is_better {
best = Some((known, distance));
}
}
best.map(|(known, _)| known.to_string())
}
fn push_unknown_keys(
unknown: &mut Vec<UnknownConfigKey>,
known_for_suggestion: &[&'static str],
path: &str,
table: &toml::Table,
allowed: &[&str],
) {
for key in table.keys() {
if !allowed.contains(&key.as_str()) {
let full_path = if path.is_empty() {
key.clone()
} else {
format!("{path}.{key}")
};
unknown.push(UnknownConfigKey {
path: full_path,
suggestion: unknown_key_suggestion(key, known_for_suggestion),
});
}
}
}
fn check_known_table(
parsed_toml: &toml::Value,
unknown: &mut Vec<UnknownConfigKey>,
known_for_suggestion: &[&'static str],
path: &[&str],
allowed: &[&str],
) {
if let Some(table) = table_at(parsed_toml, path) {
push_unknown_keys(
unknown,
known_for_suggestion,
&path.join("."),
table,
allowed,
);
}
}
fn check_nested_table_value(
unknown: &mut Vec<UnknownConfigKey>,
known_for_suggestion: &[&'static str],
path: String,
value: &toml::Value,
allowed: &[&str],
) {
if let Some(table) = value.as_table() {
push_unknown_keys(unknown, known_for_suggestion, &path, table, allowed);
}
}
fn collect_unknown_config_keys(parsed_toml: &toml::Value) -> Vec<UnknownConfigKey> {
let known_for_suggestion = known_config_keys_for_suggestion();
let mut unknown = Vec::new();
if let Some(root) = parsed_toml.as_table() {
push_unknown_keys(
&mut unknown,
&known_for_suggestion,
"",
root,
TOP_LEVEL_CONFIG_KEYS,
);
}
check_known_table(
parsed_toml,
&mut unknown,
&known_for_suggestion,
&["general"],
GENERAL_CONFIG_KEYS,
);
check_known_table(
parsed_toml,
&mut unknown,
&known_for_suggestion,
&["general", "modes"],
PROXY_MODES_CONFIG_KEYS,
);
check_known_table(
parsed_toml,
&mut unknown,
&known_for_suggestion,
&["general", "telemetry"],
TELEMETRY_CONFIG_KEYS,
);
check_known_table(
parsed_toml,
&mut unknown,
&known_for_suggestion,
&["general", "links"],
LINKS_CONFIG_KEYS,
);
check_known_table(
parsed_toml,
&mut unknown,
&known_for_suggestion,
&["logging"],
LOGGING_CONFIG_KEYS,
);
check_known_table(
parsed_toml,
&mut unknown,
&known_for_suggestion,
&["network"],
NETWORK_CONFIG_KEYS,
);
check_known_table(
parsed_toml,
&mut unknown,
&known_for_suggestion,
&["server"],
SERVER_CONFIG_KEYS,
);
check_known_table(
parsed_toml,
&mut unknown,
&known_for_suggestion,
&["server", "api"],
API_CONFIG_KEYS,
);
check_known_table(
parsed_toml,
&mut unknown,
&known_for_suggestion,
&["server", "admin_api"],
API_CONFIG_KEYS,
);
check_known_table(
parsed_toml,
&mut unknown,
&known_for_suggestion,
&["server", "conntrack_control"],
CONNTRACK_CONTROL_CONFIG_KEYS,
);
check_known_table(
parsed_toml,
&mut unknown,
&known_for_suggestion,
&["timeouts"],
TIMEOUTS_CONFIG_KEYS,
);
check_known_table(
parsed_toml,
&mut unknown,
&known_for_suggestion,
&["censorship"],
CENSORSHIP_CONFIG_KEYS,
);
check_known_table(
parsed_toml,
&mut unknown,
&known_for_suggestion,
&["censorship", "tls_fetch"],
TLS_FETCH_CONFIG_KEYS,
);
check_known_table(
parsed_toml,
&mut unknown,
&known_for_suggestion,
&["access"],
ACCESS_CONFIG_KEYS,
);
if let Some(listeners) = table_at(parsed_toml, &["server"])
.and_then(|table| table.get("listeners"))
.and_then(toml::Value::as_array)
{
for (idx, listener) in listeners.iter().enumerate() {
check_nested_table_value(
&mut unknown,
&known_for_suggestion,
format!("server.listeners[{idx}]"),
listener,
LISTENER_CONFIG_KEYS,
);
}
}
if let Some(upstreams) = parsed_toml.get("upstreams").and_then(toml::Value::as_array) {
for (idx, upstream) in upstreams.iter().enumerate() {
check_nested_table_value(
&mut unknown,
&known_for_suggestion,
format!("upstreams[{idx}]"),
upstream,
UPSTREAM_CONFIG_KEYS,
);
}
}
for access_map in ["user_rate_limits", "cidr_rate_limits"] {
if let Some(table) = table_at(parsed_toml, &["access"])
.and_then(|access| access.get(access_map))
.and_then(toml::Value::as_table)
{
for (entry_name, value) in table {
check_nested_table_value(
&mut unknown,
&known_for_suggestion,
format!("access.{access_map}.{entry_name}"),
value,
RATE_LIMIT_BPS_CONFIG_KEYS,
);
}
}
}
unknown
}
pub(super) fn handle_unknown_config_keys(parsed_toml: &toml::Value) -> Result<()> {
let unknown = collect_unknown_config_keys(parsed_toml);
if unknown.is_empty() {
return Ok(());
}
for item in &unknown {
if let Some(suggestion) = item.suggestion.as_deref() {
warn!(
key = %item.path,
suggestion = %suggestion,
"Unknown config key ignored; did you mean the suggested key?"
);
} else {
warn!(key = %item.path, "Unknown config key ignored");
}
}
if is_strict_config(parsed_toml) {
let mut paths = Vec::with_capacity(unknown.len());
for item in unknown {
if let Some(suggestion) = item.suggestion {
paths.push(format!("{} (did you mean `{}`?)", item.path, suggestion));
} else {
paths.push(item.path);
}
}
return Err(ProxyError::Config(format!(
"unknown config keys are not allowed when general.config_strict=true: {}",
paths.join(", ")
)));
}
Ok(())
}
+111
View File
@@ -0,0 +1,111 @@
use shadowsocks::config::ServerConfig as ShadowsocksServerConfig;
use tracing::warn;
use crate::error::{ProxyError, Result};
use super::super::types::{LoggingConfig, LoggingDestination, NetworkConfig, UpstreamType};
use super::ProxyConfig;
pub(super) fn validate_network_cfg(net: &mut NetworkConfig) -> Result<()> {
if !net.ipv4 && matches!(net.ipv6, Some(false)) {
return Err(ProxyError::Config(
"Both ipv4 and ipv6 are disabled in [network]".to_string(),
));
}
if net.prefer != 4 && net.prefer != 6 {
return Err(ProxyError::Config(
"network.prefer must be 4 or 6".to_string(),
));
}
if !net.ipv4 && net.prefer == 4 {
warn!("prefer=4 but ipv4=false; forcing prefer=6");
net.prefer = 6;
}
if matches!(net.ipv6, Some(false)) && net.prefer == 6 {
warn!("prefer=6 but ipv6=false; forcing prefer=4");
net.prefer = 4;
}
Ok(())
}
pub(super) fn validate_logging_config(logging: &LoggingConfig) -> Result<()> {
if let Some(path) = logging.path.as_ref()
&& path.trim().is_empty()
{
return Err(ProxyError::Config(
"logging.path cannot be empty when provided".to_string(),
));
}
if matches!(logging.destination, LoggingDestination::File) && logging.path.is_none() {
return Err(ProxyError::Config(
"logging.path must be set when logging.destination=\"file\"".to_string(),
));
}
Ok(())
}
pub(super) fn validate_upstreams(config: &ProxyConfig) -> Result<()> {
let has_enabled_shadowsocks = config.upstreams.iter().any(|upstream| {
upstream.enabled && matches!(upstream.upstream_type, UpstreamType::Shadowsocks { .. })
});
if has_enabled_shadowsocks && config.general.use_middle_proxy {
return Err(ProxyError::Config(
"shadowsocks upstreams require general.use_middle_proxy = false".to_string(),
));
}
for upstream in &config.upstreams {
if matches!(upstream.ipv4, Some(false)) && matches!(upstream.ipv6, Some(false)) {
return Err(ProxyError::Config(
"upstream.ipv4 and upstream.ipv6 cannot both be false".to_string(),
));
}
if let Some(prefer) = upstream.prefer
&& prefer != 4
&& prefer != 6
{
return Err(ProxyError::Config(
"upstream.prefer must be 4 or 6".to_string(),
));
}
if let UpstreamType::Shadowsocks { url, .. } = &upstream.upstream_type {
let parsed = ShadowsocksServerConfig::from_url(url)
.map_err(|error| ProxyError::Config(format!("invalid shadowsocks url: {error}")))?;
if parsed.plugin().is_some() {
return Err(ProxyError::Config(
"shadowsocks plugins are not supported".to_string(),
));
}
}
}
Ok(())
}
pub(super) fn normalize_upstream_family_policy(config: &mut ProxyConfig) {
for (idx, upstream) in config.upstreams.iter_mut().enumerate() {
if matches!(upstream.ipv4, Some(false)) && upstream.prefer == Some(4) {
warn!(
upstream = idx,
"upstream.prefer=4 but upstream.ipv4=false; forcing prefer=6"
);
upstream.prefer = Some(6);
}
if matches!(upstream.ipv6, Some(false)) && upstream.prefer == Some(6) {
warn!(
upstream = idx,
"upstream.prefer=6 but upstream.ipv6=false; forcing prefer=4"
);
upstream.prefer = Some(4);
}
}
}
File diff suppressed because it is too large Load Diff
@@ -95,6 +95,101 @@ max_client_frame = 16777217
remove_temp_config(&path); remove_temp_config(&path);
} }
#[test]
fn load_rejects_writer_byte_budget_below_frame_residency_minimum() {
let path = write_temp_config(
r#"
[general]
max_client_frame = 16777216
me_writer_byte_budget_bytes = 33554432
"#,
);
let err = ProxyConfig::load(&path)
.expect_err("writer byte budget below frame residency minimum must fail");
let msg = err.to_string();
assert!(
msg.contains("general.me_writer_byte_budget_bytes must be within [33570816, 268435456]"),
"error must explain writer byte budget minimum, got: {msg}"
);
remove_temp_config(&path);
}
#[test]
fn load_rejects_unaligned_writer_byte_budget() {
let path = write_temp_config(
r#"
[general]
me_writer_byte_budget_bytes = 33570817
"#,
);
let err = ProxyConfig::load(&path)
.expect_err("writer byte budget outside permit granularity must fail");
let msg = err.to_string();
assert!(
msg.contains("general.me_writer_byte_budget_bytes must be a multiple of 16384"),
"error must explain writer byte budget alignment, got: {msg}"
);
remove_temp_config(&path);
}
#[test]
fn load_rejects_writer_byte_budget_above_hard_cap() {
let path = write_temp_config(
r#"
[general]
me_writer_byte_budget_bytes = 268451840
"#,
);
let err = ProxyConfig::load(&path).expect_err("writer byte budget above hard cap must fail");
let msg = err.to_string();
assert!(
msg.contains("general.me_writer_byte_budget_bytes must be within [33570816, 268435456]"),
"error must explain writer byte budget hard cap, got: {msg}"
);
remove_temp_config(&path);
}
#[test]
fn load_rejects_unaligned_direct_relay_buffer_budget() {
let path = write_temp_config(
r#"
[general]
direct_relay_buffer_budget_max_bytes = 16777217
"#,
);
let err = ProxyConfig::load(&path).expect_err("unaligned direct relay buffer budget must fail");
assert!(
err.to_string().contains(
"general.direct_relay_buffer_budget_max_bytes must be 0 or a multiple of 4096"
)
);
remove_temp_config(&path);
}
#[test]
fn load_rejects_direct_relay_buffer_budget_above_hard_cap() {
let path = write_temp_config(
r#"
[general]
direct_relay_buffer_budget_max_bytes = 2147487744
"#,
);
let err =
ProxyConfig::load(&path).expect_err("direct relay buffer budget above hard cap must fail");
assert!(err.to_string().contains(
"general.direct_relay_buffer_budget_max_bytes must be 0 or within [16777216, 2147483648]"
));
remove_temp_config(&path);
}
#[test] #[test]
fn load_rejects_listen_backlog_above_i32_upper_bound() { fn load_rejects_listen_backlog_above_i32_upper_bound() {
let path = write_temp_config( let path = write_temp_config(
@@ -139,16 +234,23 @@ fn load_accepts_memory_limits_at_hard_upper_bounds() {
r#" r#"
[general] [general]
me_writer_cmd_channel_capacity = 16384 me_writer_cmd_channel_capacity = 16384
me_writer_byte_budget_bytes = 268435456
me_route_channel_capacity = 8192 me_route_channel_capacity = 8192
me_c2me_channel_capacity = 8192 me_c2me_channel_capacity = 8192
direct_relay_buffer_budget_max_bytes = 2147483648
max_client_frame = 16777216 max_client_frame = 16777216
"#, "#,
); );
let cfg = ProxyConfig::load(&path).expect("hard upper bound values must be accepted"); let cfg = ProxyConfig::load(&path).expect("hard upper bound values must be accepted");
assert_eq!(cfg.general.me_writer_cmd_channel_capacity, 16384); assert_eq!(cfg.general.me_writer_cmd_channel_capacity, 16384);
assert_eq!(cfg.general.me_writer_byte_budget_bytes, 256 * 1024 * 1024);
assert_eq!(cfg.general.me_route_channel_capacity, 8192); assert_eq!(cfg.general.me_route_channel_capacity, 8192);
assert_eq!(cfg.general.me_c2me_channel_capacity, 8192); assert_eq!(cfg.general.me_c2me_channel_capacity, 8192);
assert_eq!(
cfg.general.direct_relay_buffer_budget_max_bytes,
2 * 1024 * 1024 * 1024
);
assert_eq!(cfg.general.max_client_frame, 16 * 1024 * 1024); assert_eq!(cfg.general.max_client_frame, 16 * 1024 * 1024);
remove_temp_config(&path); remove_temp_config(&path);
+186 -11
View File
@@ -2,6 +2,7 @@ use chrono::{DateTime, Utc};
use ipnetwork::IpNetwork; use ipnetwork::IpNetwork;
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use std::collections::HashMap; use std::collections::HashMap;
use std::fmt;
use std::net::IpAddr; use std::net::IpAddr;
use std::path::PathBuf; use std::path::PathBuf;
@@ -578,6 +579,10 @@ pub struct GeneralConfig {
#[serde(default = "default_me_writer_cmd_channel_capacity")] #[serde(default = "default_me_writer_cmd_channel_capacity")]
pub me_writer_cmd_channel_capacity: usize, pub me_writer_cmd_channel_capacity: usize,
/// Resident-memory budget in bytes for each ME writer data queue.
#[serde(default = "default_me_writer_byte_budget_bytes")]
pub me_writer_byte_budget_bytes: usize,
/// Capacity of per-connection ME response route channel. /// Capacity of per-connection ME response route channel.
#[serde(default = "default_me_route_channel_capacity")] #[serde(default = "default_me_route_channel_capacity")]
pub me_route_channel_capacity: usize, pub me_route_channel_capacity: usize,
@@ -621,7 +626,7 @@ pub struct GeneralConfig {
#[serde(default = "default_me_d2c_frame_buf_shrink_threshold_bytes")] #[serde(default = "default_me_d2c_frame_buf_shrink_threshold_bytes")]
pub me_d2c_frame_buf_shrink_threshold_bytes: usize, pub me_d2c_frame_buf_shrink_threshold_bytes: usize,
/// Copy buffer size for client->DC direction in direct relay. /// Copy buffer ceiling for client->DC direction in direct relay.
/// ///
/// This is also the upper bound for one amortized upload rate-limit burst: /// This is also the upper bound for one amortized upload rate-limit burst:
/// upload debt is settled before the next relay read instead of blocking /// upload debt is settled before the next relay read instead of blocking
@@ -629,13 +634,18 @@ pub struct GeneralConfig {
#[serde(default = "default_direct_relay_copy_buf_c2s_bytes")] #[serde(default = "default_direct_relay_copy_buf_c2s_bytes")]
pub direct_relay_copy_buf_c2s_bytes: usize, pub direct_relay_copy_buf_c2s_bytes: usize,
/// Copy buffer size for DC->client direction in direct relay. /// Copy buffer ceiling for DC->client direction in direct relay.
/// ///
/// This bounds one direct download rate-limit grant because writes are /// This bounds one direct download rate-limit grant because writes are
/// clipped to the currently available shaper budget. /// clipped to the currently available shaper budget.
#[serde(default = "default_direct_relay_copy_buf_s2c_bytes")] #[serde(default = "default_direct_relay_copy_buf_s2c_bytes")]
pub direct_relay_copy_buf_s2c_bytes: usize, pub direct_relay_copy_buf_s2c_bytes: usize,
/// Process-wide hard ceiling for Direct relay copy buffers.
/// `0` derives the ceiling from host and cgroup memory limits.
#[serde(default = "default_direct_relay_buffer_budget_max_bytes")]
pub direct_relay_buffer_budget_max_bytes: usize,
/// Max pending ciphertext buffer per client writer (bytes). /// Max pending ciphertext buffer per client writer (bytes).
/// Controls FakeTLS backpressure vs throughput. /// Controls FakeTLS backpressure vs throughput.
#[serde(default = "default_crypto_pending_buffer")] #[serde(default = "default_crypto_pending_buffer")]
@@ -1102,6 +1112,7 @@ impl Default for GeneralConfig {
me_keepalive_payload_random: default_true(), me_keepalive_payload_random: default_true(),
rpc_proxy_req_every: default_rpc_proxy_req_every(), rpc_proxy_req_every: default_rpc_proxy_req_every(),
me_writer_cmd_channel_capacity: default_me_writer_cmd_channel_capacity(), me_writer_cmd_channel_capacity: default_me_writer_cmd_channel_capacity(),
me_writer_byte_budget_bytes: default_me_writer_byte_budget_bytes(),
me_route_channel_capacity: default_me_route_channel_capacity(), me_route_channel_capacity: default_me_route_channel_capacity(),
me_c2me_channel_capacity: default_me_c2me_channel_capacity(), me_c2me_channel_capacity: default_me_c2me_channel_capacity(),
me_c2me_send_timeout_ms: default_me_c2me_send_timeout_ms(), me_c2me_send_timeout_ms: default_me_c2me_send_timeout_ms(),
@@ -1115,6 +1126,7 @@ impl Default for GeneralConfig {
default_me_d2c_frame_buf_shrink_threshold_bytes(), default_me_d2c_frame_buf_shrink_threshold_bytes(),
direct_relay_copy_buf_c2s_bytes: default_direct_relay_copy_buf_c2s_bytes(), direct_relay_copy_buf_c2s_bytes: default_direct_relay_copy_buf_c2s_bytes(),
direct_relay_copy_buf_s2c_bytes: default_direct_relay_copy_buf_s2c_bytes(), direct_relay_copy_buf_s2c_bytes: default_direct_relay_copy_buf_s2c_bytes(),
direct_relay_buffer_budget_max_bytes: default_direct_relay_buffer_budget_max_bytes(),
me_warmup_stagger_enabled: default_true(), me_warmup_stagger_enabled: default_true(),
me_warmup_step_delay_ms: default_warmup_step_delay_ms(), me_warmup_step_delay_ms: default_warmup_step_delay_ms(),
me_warmup_step_jitter_ms: default_warmup_step_jitter_ms(), me_warmup_step_jitter_ms: default_warmup_step_jitter_ms(),
@@ -1455,9 +1467,9 @@ pub enum SynLimitMode {
/// Disable SYN limiting for this listener. /// Disable SYN limiting for this listener.
#[default] #[default]
Off, Off,
/// Use iptables/ip6tables filter rules with the hashlimit match. /// Use iptables/ip6tables two-tier SYN-fix rules with the hashlimit match.
Iptables, Iptables,
/// Use nftables rules with per-source token-bucket meters. /// Use nftables two-tier SYN-fix rules with per-source token-bucket meters.
Nftables, Nftables,
} }
@@ -2098,11 +2110,13 @@ pub struct AccessConfig {
/// Per-CIDR aggregate transport rate limits in bits-per-second. /// Per-CIDR aggregate transport rate limits in bits-per-second.
/// ///
/// Matching uses longest-prefix-wins semantics. A value of `0` in one /// Explicit CIDR keys use longest-prefix-wins semantics. Auto-template
/// direction means "unlimited" for that direction. Limits are amortized /// keys (`*4/N`, `*6/N`, `*/N`) lazily create per-source-subnet buckets
/// with the same bounded-burst contract as per-user rate limits. /// after explicit CIDR matching misses. A value of `0` in one direction
/// means "unlimited" for that direction. Limits are amortized with the
/// same bounded-burst contract as per-user rate limits.
#[serde(default)] #[serde(default)]
pub cidr_rate_limits: HashMap<IpNetwork, RateLimitBps>, pub cidr_rate_limits: HashMap<CidrRateLimitKey, RateLimitBps>,
/// Per-username client source IP/CIDR deny list. Checked after successful /// Per-username client source IP/CIDR deny list. Checked after successful
/// authentication; matching IPs get the same rejection path as invalid auth /// authentication; matching IPs get the same rejection path as invalid auth
@@ -2171,10 +2185,156 @@ impl AccessConfig {
} }
} }
/// Key used by `access.cidr_rate_limits`.
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub enum CidrRateLimitKey {
/// Explicit source CIDR rule.
Network(IpNetwork),
/// IPv4 auto-template that creates one bucket for each matching `/N`.
AutoV4(u8),
/// IPv6 auto-template that creates one bucket for each matching `/N`.
AutoV6(u8),
/// Dual-stack auto-template; IPv4 uses `/N`, IPv6 uses `/(N * 4)`.
AutoDual(u8),
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub(crate) enum CidrAutoTemplateFamily {
V4,
V6,
}
impl CidrAutoTemplateFamily {
pub(crate) fn marker(self) -> &'static str {
match self {
Self::V4 => "*4",
Self::V6 => "*6",
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub(crate) struct CidrAutoTemplate {
pub(crate) family: CidrAutoTemplateFamily,
pub(crate) prefix_len: u8,
}
impl fmt::Display for CidrAutoTemplate {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(formatter, "{}/{}", self.family.marker(), self.prefix_len)
}
}
impl CidrRateLimitKey {
pub(crate) fn auto_templates(&self) -> [Option<CidrAutoTemplate>; 2] {
match *self {
Self::Network(_) => [None, None],
Self::AutoV4(prefix_len) => [
Some(CidrAutoTemplate {
family: CidrAutoTemplateFamily::V4,
prefix_len,
}),
None,
],
Self::AutoV6(prefix_len) => [
Some(CidrAutoTemplate {
family: CidrAutoTemplateFamily::V6,
prefix_len,
}),
None,
],
Self::AutoDual(prefix_len) => [
Some(CidrAutoTemplate {
family: CidrAutoTemplateFamily::V4,
prefix_len,
}),
Some(CidrAutoTemplate {
family: CidrAutoTemplateFamily::V6,
prefix_len: prefix_len.saturating_mul(4),
}),
],
}
}
}
impl fmt::Display for CidrRateLimitKey {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Network(cidr) => write!(formatter, "{cidr}"),
Self::AutoV4(prefix_len) => write!(formatter, "*4/{prefix_len}"),
Self::AutoV6(prefix_len) => write!(formatter, "*6/{prefix_len}"),
Self::AutoDual(prefix_len) => write!(formatter, "*/{prefix_len}"),
}
}
}
impl Serialize for CidrRateLimitKey {
fn serialize<S>(&self, serializer: S) -> std::result::Result<S::Ok, S::Error>
where
S: serde::Serializer,
{
serializer.collect_str(self)
}
}
impl<'de> Deserialize<'de> for CidrRateLimitKey {
fn deserialize<D>(deserializer: D) -> std::result::Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
let value = String::deserialize(deserializer)?;
parse_cidr_rate_limit_key(&value).map_err(serde::de::Error::custom)
}
}
fn parse_cidr_rate_limit_key(value: &str) -> std::result::Result<CidrRateLimitKey, String> {
if let Some(prefix) = value.strip_prefix("*4/") {
return parse_cidr_auto_prefix(value, prefix, 32).map(CidrRateLimitKey::AutoV4);
}
if let Some(prefix) = value.strip_prefix("*6/") {
return parse_cidr_auto_prefix(value, prefix, 128).map(CidrRateLimitKey::AutoV6);
}
if let Some(prefix) = value.strip_prefix("*/") {
return parse_cidr_auto_prefix(value, prefix, 32).map(CidrRateLimitKey::AutoDual);
}
if value.starts_with('*') {
return Err(format!(
"invalid CIDR rate limit key {value:?}; expected CIDR, *4/N, *6/N, or */N"
));
}
value
.parse::<IpNetwork>()
.map(CidrRateLimitKey::Network)
.map_err(|error| {
format!(
"invalid CIDR rate limit key {value:?}: {error}; expected CIDR, *4/N, *6/N, or */N"
)
})
}
fn parse_cidr_auto_prefix(
key: &str,
prefix: &str,
max_prefix: u8,
) -> std::result::Result<u8, String> {
let prefix = prefix.parse::<u8>().map_err(|_| {
format!("invalid CIDR auto-template key {key:?}; prefix must be within 0..={max_prefix}")
})?;
if prefix > max_prefix {
return Err(format!(
"invalid CIDR auto-template key {key:?}; prefix must be within 0..={max_prefix}"
));
}
Ok(prefix)
}
/// Transport rate limit in bits-per-second.
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)] #[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
pub struct RateLimitBps { pub struct RateLimitBps {
/// Upload direction limit in bits-per-second; `0` means unlimited.
#[serde(default)] #[serde(default)]
pub up_bps: u64, pub up_bps: u64,
/// Download direction limit in bits-per-second; `0` means unlimited.
#[serde(default)] #[serde(default)]
pub down_bps: u64, pub down_bps: u64,
} }
@@ -2266,15 +2426,30 @@ pub struct ListenerConfig {
/// Per-listener SYN limiter mode. /// Per-listener SYN limiter mode.
#[serde(default)] #[serde(default)]
pub synlimit: SynLimitMode, pub synlimit: SynLimitMode,
/// Token-bucket rate interval for the per-listener SYN limiter. /// Generic SYN-fix token-bucket rate interval.
#[serde(default = "default_synlimit_seconds")] #[serde(default = "default_synlimit_seconds")]
pub synlimit_seconds: u32, pub synlimit_seconds: u32,
/// Token-bucket rate amount for the per-listener SYN limiter. /// Generic SYN-fix token-bucket rate amount.
#[serde(default = "default_synlimit_hitcount")] #[serde(default = "default_synlimit_hitcount")]
pub synlimit_hitcount: u32, pub synlimit_hitcount: u32,
/// Token-bucket burst size for the per-listener SYN limiter. /// Generic SYN-fix token-bucket burst size.
#[serde(default = "default_synlimit_burst")] #[serde(default = "default_synlimit_burst")]
pub synlimit_burst: u32, pub synlimit_burst: u32,
/// iOS-like SYN-fix token-bucket rate interval.
#[serde(default = "default_synlimit_ios_seconds")]
pub synlimit_ios_seconds: u32,
/// iOS-like SYN-fix token-bucket rate amount.
#[serde(default = "default_synlimit_ios_hitcount")]
pub synlimit_ios_hitcount: u32,
/// iOS-like SYN-fix token-bucket burst size.
#[serde(default = "default_synlimit_ios_burst")]
pub synlimit_ios_burst: u32,
/// Hashlimit entry expiration in milliseconds for iptables/ip6tables rules.
#[serde(default = "default_synlimit_hashlimit_expire_ms")]
pub synlimit_hashlimit_expire_ms: u32,
/// Hashlimit table size for iptables/ip6tables rules.
#[serde(default = "default_synlimit_hashlimit_size")]
pub synlimit_hashlimit_size: u32,
/// IP address or hostname to announce in proxy links. /// IP address or hostname to announce in proxy links.
/// Takes precedence over `announce_ip` if both are set. /// Takes precedence over `announce_ip` if both are set.
#[serde(default)] #[serde(default)]
+7 -5
View File
@@ -7,6 +7,7 @@ use std::time::Duration;
use tokio::io::AsyncWriteExt; use tokio::io::AsyncWriteExt;
use tokio::process::Command; use tokio::process::Command;
use tokio::sync::{mpsc, watch}; use tokio::sync::{mpsc, watch};
use tokio_util::sync::CancellationToken;
use tracing::{debug, info, warn}; use tracing::{debug, info, warn};
use crate::config::{ConntrackBackend, ConntrackMode, ProxyConfig}; use crate::config::{ConntrackBackend, ConntrackMode, ProxyConfig};
@@ -57,10 +58,11 @@ impl PressureState {
} }
} }
pub(crate) fn spawn_conntrack_controller( pub(crate) async fn run_conntrack_controller(
config_rx: watch::Receiver<Arc<ProxyConfig>>, config_rx: watch::Receiver<Arc<ProxyConfig>>,
stats: Arc<Stats>, stats: Arc<Stats>,
shared: Arc<ProxySharedState>, shared: Arc<ProxySharedState>,
cancellation: CancellationToken,
) { ) {
if !cfg!(target_os = "linux") { if !cfg!(target_os = "linux") {
let cfg = config_rx.borrow(); let cfg = config_rx.borrow();
@@ -87,16 +89,15 @@ pub(crate) fn spawn_conntrack_controller(
let (tx, rx) = mpsc::channel(CONNTRACK_EVENT_QUEUE_CAPACITY); let (tx, rx) = mpsc::channel(CONNTRACK_EVENT_QUEUE_CAPACITY);
shared.set_conntrack_close_sender(tx); shared.set_conntrack_close_sender(tx);
tokio::spawn(async move { run_conntrack_controller_worker(config_rx, stats, shared, rx, cancellation).await;
run_conntrack_controller(config_rx, stats, shared, rx).await;
});
} }
async fn run_conntrack_controller( async fn run_conntrack_controller_worker(
mut config_rx: watch::Receiver<Arc<ProxyConfig>>, mut config_rx: watch::Receiver<Arc<ProxyConfig>>,
stats: Arc<Stats>, stats: Arc<Stats>,
shared: Arc<ProxySharedState>, shared: Arc<ProxySharedState>,
mut close_rx: mpsc::Receiver<ConntrackCloseEvent>, mut close_rx: mpsc::Receiver<ConntrackCloseEvent>,
cancellation: CancellationToken,
) { ) {
let mut cfg = config_rx.borrow().clone(); let mut cfg = config_rx.borrow().clone();
let mut pressure_state = PressureState::new(stats.as_ref()); let mut pressure_state = PressureState::new(stats.as_ref());
@@ -115,6 +116,7 @@ async fn run_conntrack_controller(
loop { loop {
tokio::select! { tokio::select! {
_ = cancellation.cancelled() => break,
changed = config_rx.changed() => { changed = config_rx.changed() => {
if changed.is_err() { if changed.is_err() {
break; break;
+18 -12
View File
@@ -4,12 +4,12 @@
//! //!
//! ## Zeroize policy //! ## Zeroize policy
//! //!
//! - `AesCbc` stores raw key/IV bytes and zeroizes them on drop. //! - `AesCbc` stores raw key/IV bytes and zeroizes them on drop. Temporary
//! - `AesCtr` wraps an opaque `Aes256Ctr` cipher from the `ctr` crate. //! expanded key schedules are also zeroized by the RustCrypto backend.
//! The expanded key schedule lives inside that type and cannot be //! - `AesCtr` uses the RustCrypto `zeroize` contract to clear its expanded
//! zeroized from outside. Callers that hold raw key material (e.g. //! key schedule, counter, and buffered keystream on drop.
//! `HandshakeSuccess`, `ObfuscationParams`) are responsible for //! - Callers that hold raw key material (e.g. `HandshakeSuccess`,
//! zeroizing their own copies. //! `ObfuscationParams`) remain responsible for zeroizing their own copies.
#![allow(dead_code)] #![allow(dead_code)]
@@ -19,24 +19,28 @@ use ctr::{
Ctr128BE, Ctr128BE,
cipher::{KeyIvInit, StreamCipher}, cipher::{KeyIvInit, StreamCipher},
}; };
use zeroize::Zeroize; use zeroize::{Zeroize, ZeroizeOnDrop};
type Aes256Ctr = Ctr128BE<Aes256>; type Aes256Ctr = Ctr128BE<Aes256>;
static_assertions::assert_impl_all!(Aes256: ZeroizeOnDrop);
static_assertions::assert_impl_all!(Aes256Ctr: ZeroizeOnDrop);
// ============= AES-256-CTR ============= // ============= AES-256-CTR =============
/// AES-256-CTR encryptor/decryptor /// AES-256-CTR encryptor/decryptor
/// ///
/// CTR mode is symmetric — encryption and decryption are the same operation. /// CTR mode is symmetric — encryption and decryption are the same operation.
/// ///
/// **Zeroize note:** The inner `Aes256Ctr` cipher state (expanded key schedule /// **Zeroize note:** The inner `Aes256Ctr` zeroizes its expanded key schedule,
/// + counter) is opaque and cannot be zeroized. If you need to protect key /// counter, and buffered keystream on drop. Callers remain responsible for
/// material, zeroize the `[u8; 32]` key and `u128` IV at the call site /// zeroizing their own raw key and IV copies.
/// before dropping them.
pub struct AesCtr { pub struct AesCtr {
cipher: Aes256Ctr, cipher: Aes256Ctr,
} }
impl ZeroizeOnDrop for AesCtr {}
impl AesCtr { impl AesCtr {
/// Create new AES-CTR cipher with key and IV /// Create new AES-CTR cipher with key and IV
pub fn new(key: &[u8; 32], iv: u128) -> Self { pub fn new(key: &[u8; 32], iv: u128) -> Self {
@@ -92,7 +96,7 @@ impl AesCtr {
/// are different operations. This implementation handles CBC chaining /// are different operations. This implementation handles CBC chaining
/// correctly across multiple blocks. /// correctly across multiple blocks.
/// ///
/// Key and IV are zeroized on drop. /// Key, IV, and temporary expanded key schedules are zeroized on drop.
pub struct AesCbc { pub struct AesCbc {
key: [u8; 32], key: [u8; 32],
iv: [u8; 16], iv: [u8; 16],
@@ -105,6 +109,8 @@ impl Drop for AesCbc {
} }
} }
impl ZeroizeOnDrop for AesCbc {}
impl AesCbc { impl AesCbc {
/// AES block size /// AES block size
const BLOCK_SIZE: usize = 16; const BLOCK_SIZE: usize = 16;
+3 -3
View File
@@ -5,7 +5,7 @@ pub mod hash;
pub mod random; pub mod random;
pub use aes::{AesCbc, AesCtr}; pub use aes::{AesCbc, AesCtr};
pub use hash::{ #[allow(unused_imports)]
build_middleproxy_prekey, crc32, crc32c, derive_middleproxy_keys, sha256, sha256_hmac, pub use hash::build_middleproxy_prekey;
}; pub use hash::{crc32, crc32c, derive_middleproxy_keys, sha256, sha256_hmac};
pub use random::SecureRandom; pub use random::SecureRandom;
+4 -1
View File
@@ -8,6 +8,8 @@ use crate::config::ProxyConfig;
use crate::proxy::route_mode::{RelayRouteMode, RouteRuntimeController}; use crate::proxy::route_mode::{RelayRouteMode, RouteRuntimeController};
use crate::transport::middle_proxy::MePool; use crate::transport::middle_proxy::MePool;
use super::generation::RuntimeTaskScope;
const STARTUP_FALLBACK_AFTER: Duration = Duration::from_secs(80); const STARTUP_FALLBACK_AFTER: Duration = Duration::from_secs(80);
const RUNTIME_FALLBACK_AFTER: Duration = Duration::from_secs(6); const RUNTIME_FALLBACK_AFTER: Duration = Duration::from_secs(6);
@@ -19,6 +21,7 @@ pub(crate) async fn configure_admission_gate(
admission_tx: &watch::Sender<bool>, admission_tx: &watch::Sender<bool>,
config_rx: watch::Receiver<Arc<ProxyConfig>>, config_rx: watch::Receiver<Arc<ProxyConfig>>,
me_ready_rx: watch::Receiver<u64>, me_ready_rx: watch::Receiver<u64>,
task_scope: RuntimeTaskScope,
) { ) {
if config.general.use_middle_proxy { if config.general.use_middle_proxy {
if me_pool.is_some() || config.general.me2dc_fallback { if me_pool.is_some() || config.general.me2dc_fallback {
@@ -64,7 +67,7 @@ pub(crate) async fn configure_admission_gate(
let mut config_rx_gate = config_rx.clone(); let mut config_rx_gate = config_rx.clone();
let mut me_ready_rx_gate = me_ready_rx; let mut me_ready_rx_gate = me_ready_rx;
let mut admission_poll_ms = config.general.me_admission_poll_ms.max(1); let mut admission_poll_ms = config.general.me_admission_poll_ms.max(1);
tokio::spawn(async move { task_scope.spawn(async move {
let mut gate_open = initial_gate_open; let mut gate_open = initial_gate_open;
let mut route_mode = initial_route_mode; let mut route_mode = initial_route_mode;
let mut ready_observed = initial_ready; let mut ready_observed = initial_ready;
+435
View File
@@ -0,0 +1,435 @@
use std::future::Future;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::time::Duration;
use tokio::sync::{RwLock, Semaphore, watch};
use tokio_util::sync::CancellationToken;
use tokio_util::task::TaskTracker;
use crate::config::ProxyConfig;
use crate::crypto::SecureRandom;
use crate::ip_tracker::UserIpTracker;
#[cfg(test)]
use crate::proxy::route_mode::RelayRouteMode;
use crate::proxy::route_mode::RouteRuntimeController;
use crate::proxy::shared_state::ProxySharedState;
use crate::stats::beobachten::BeobachtenStore;
use crate::stats::{ReplayChecker, Stats};
use crate::stream::BufferPool;
use crate::tls_front::TlsFrontCache;
use crate::transport::UpstreamManager;
use crate::transport::middle_proxy::MePool;
const SESSION_STOP_TIMEOUT: Duration = Duration::from_secs(5);
const BACKGROUND_STOP_TIMEOUT: Duration = Duration::from_secs(5);
const SESSION_ADMISSION_CLOSED: usize = 1 << (usize::BITS - 1);
const SESSION_REGISTRATION_COUNT: usize = SESSION_ADMISSION_CLOSED - 1;
struct SessionAdmission {
state: AtomicUsize,
}
struct SessionRegistration<'a> {
admission: &'a SessionAdmission,
}
impl SessionAdmission {
fn new() -> Self {
Self {
state: AtomicUsize::new(0),
}
}
fn try_register(&self) -> Option<SessionRegistration<'_>> {
let mut state = self.state.load(Ordering::Acquire);
loop {
if state & SESSION_ADMISSION_CLOSED != 0
|| state & SESSION_REGISTRATION_COUNT == SESSION_REGISTRATION_COUNT
{
return None;
}
match self.state.compare_exchange_weak(
state,
state + 1,
Ordering::AcqRel,
Ordering::Acquire,
) {
Ok(_) => return Some(SessionRegistration { admission: self }),
Err(observed) => state = observed,
}
}
}
fn close(&self) {
self.state
.fetch_or(SESSION_ADMISSION_CLOSED, Ordering::AcqRel);
}
fn reopen(&self) {
self.state
.fetch_and(!SESSION_ADMISSION_CLOSED, Ordering::AcqRel);
}
async fn wait_for_registrations(&self) {
while self.state.load(Ordering::Acquire) & SESSION_REGISTRATION_COUNT != 0 {
tokio::task::yield_now().await;
}
}
}
impl Drop for SessionRegistration<'_> {
fn drop(&mut self) {
self.admission.state.fetch_sub(1, Ordering::Release);
}
}
/// Process-visible control-plane receivers for one active runtime generation.
#[derive(Clone)]
pub(crate) struct RuntimeWatchState {
pub(crate) generation_id: u64,
pub(crate) config_rx: watch::Receiver<Arc<ProxyConfig>>,
pub(crate) admission_rx: watch::Receiver<bool>,
}
/// Cancellation and join ownership for one generation's background tasks.
#[derive(Clone)]
pub(crate) struct RuntimeTaskScope {
tracker: TaskTracker,
cancel: CancellationToken,
}
impl RuntimeTaskScope {
/// Creates an open generation-owned task scope.
pub(crate) fn new() -> Self {
Self {
tracker: TaskTracker::new(),
cancel: CancellationToken::new(),
}
}
/// Spawns one task that is cancelled when the generation stops.
pub(crate) fn spawn<F>(&self, future: F)
where
F: Future<Output = ()> + Send + 'static,
{
let cancel = self.cancel.clone();
self.tracker.spawn(async move {
tokio::select! {
_ = cancel.cancelled() => {}
_ = future => {}
}
});
}
/// Returns the cancellation signal shared by generation-owned controllers.
pub(crate) fn cancellation_token(&self) -> CancellationToken {
self.cancel.clone()
}
/// Cancels the scope and waits within the bounded background-task budget.
pub(crate) async fn stop(&self) {
self.cancel.cancel();
self.tracker.close();
let _ = tokio::time::timeout(BACKGROUND_STOP_TIMEOUT, self.tracker.wait()).await;
}
}
/// Runtime-owned data plane and control-plane dependencies for one generation.
pub(crate) struct RuntimeGeneration {
pub(crate) id: u64,
pub(crate) config_rx: watch::Receiver<Arc<ProxyConfig>>,
pub(crate) admission_rx: watch::Receiver<bool>,
pub(crate) stats: Arc<Stats>,
pub(crate) upstream_manager: Arc<UpstreamManager>,
pub(crate) replay_checker: Arc<ReplayChecker>,
pub(crate) buffer_pool: Arc<BufferPool>,
pub(crate) rng: Arc<SecureRandom>,
pub(crate) me_pool: Option<Arc<MePool>>,
pub(crate) me_pool_runtime: Arc<RwLock<Option<Arc<MePool>>>>,
pub(crate) route_runtime: Arc<RouteRuntimeController>,
pub(crate) tls_cache: Option<Arc<TlsFrontCache>>,
pub(crate) ip_tracker: Arc<UserIpTracker>,
pub(crate) beobachten: Arc<BeobachtenStore>,
pub(crate) proxy_shared: Arc<ProxySharedState>,
pub(crate) max_connections: Arc<Semaphore>,
background_tasks: RuntimeTaskScope,
sessions: TaskTracker,
session_cancel: CancellationToken,
session_admission: SessionAdmission,
}
impl RuntimeGeneration {
#[allow(clippy::too_many_arguments)]
/// Builds one fully owned runtime generation.
pub(crate) fn new(
id: u64,
config_rx: watch::Receiver<Arc<ProxyConfig>>,
admission_rx: watch::Receiver<bool>,
stats: Arc<Stats>,
upstream_manager: Arc<UpstreamManager>,
replay_checker: Arc<ReplayChecker>,
buffer_pool: Arc<BufferPool>,
rng: Arc<SecureRandom>,
me_pool: Option<Arc<MePool>>,
me_pool_runtime: Arc<RwLock<Option<Arc<MePool>>>>,
route_runtime: Arc<RouteRuntimeController>,
tls_cache: Option<Arc<TlsFrontCache>>,
ip_tracker: Arc<UserIpTracker>,
beobachten: Arc<BeobachtenStore>,
proxy_shared: Arc<ProxySharedState>,
max_connections: Arc<Semaphore>,
background_tasks: RuntimeTaskScope,
) -> Arc<Self> {
Arc::new(Self {
id,
config_rx,
admission_rx,
stats,
upstream_manager,
replay_checker,
buffer_pool,
rng,
me_pool,
me_pool_runtime,
route_runtime,
tls_cache,
ip_tracker,
beobachten,
proxy_shared,
max_connections,
background_tasks,
sessions: TaskTracker::new(),
session_cancel: CancellationToken::new(),
session_admission: SessionAdmission::new(),
})
}
/// Returns the latest hot-reloaded configuration for this generation.
pub(crate) fn config(&self) -> Arc<ProxyConfig> {
self.config_rx.borrow().clone()
}
/// Returns receivers used by process-scoped observers of this generation.
pub(crate) fn watch_state(&self) -> RuntimeWatchState {
RuntimeWatchState {
generation_id: self.id,
config_rx: self.config_rx.clone(),
admission_rx: self.admission_rx.clone(),
}
}
/// Returns the initial or asynchronously published Middle-End pool.
pub(crate) async fn current_me_pool(&self) -> Option<Arc<MePool>> {
if let Some(pool) = &self.me_pool {
return Some(pool.clone());
}
self.me_pool_runtime.read().await.clone()
}
/// Registers a session only while admission remains open.
pub(crate) fn spawn_session<F>(&self, future: F) -> bool
where
F: Future<Output = ()> + Send + 'static,
{
let Some(_registration) = self.session_admission.try_register() else {
return false;
};
let cancel = self.session_cancel.clone();
self.sessions.spawn(async move {
tokio::select! {
_ = cancel.cancelled() => {}
_ = future => {}
}
});
true
}
/// Closes admission while preserving already registered sessions.
pub(crate) fn stop_accepting_sessions(&self) {
self.session_admission.close();
}
/// Reopens admission after a candidate activation rolls back.
pub(crate) fn resume_accepting_sessions(&self) {
self.session_admission.reopen();
}
/// Waits for registered sessions and cancels them when the deadline expires.
pub(crate) async fn drain_sessions(&self, timeout: Duration) -> bool {
self.stop_accepting_sessions();
self.session_admission.wait_for_registrations().await;
self.sessions.close();
if tokio::time::timeout(timeout, self.sessions.wait())
.await
.is_ok()
{
return true;
}
self.stop_sessions().await;
false
}
/// Cancels all sessions and waits within the bounded session-stop budget.
pub(crate) async fn stop_sessions(&self) {
self.stop_accepting_sessions();
self.session_admission.wait_for_registrations().await;
self.session_cancel.cancel();
self.sessions.close();
let _ = tokio::time::timeout(SESSION_STOP_TIMEOUT, self.sessions.wait()).await;
}
/// Stops all background tasks owned by this generation.
pub(crate) async fn stop_background_tasks(&self) {
self.background_tasks.stop().await;
}
}
#[cfg(test)]
/// Builds a lightweight runtime generation without network startup tasks.
pub(super) fn test_runtime_generation(id: u64, config: ProxyConfig) -> Arc<RuntimeGeneration> {
let (config_tx, config_rx) = watch::channel(Arc::new(config.clone()));
let (_admission_tx, admission_rx) = watch::channel(true);
let stats = Arc::new(Stats::new());
let upstream_manager = Arc::new(UpstreamManager::new(
config.upstreams,
config.general.upstream_connect_retry_attempts,
config.general.upstream_connect_retry_backoff_ms,
config.general.upstream_connect_budget_ms,
config.general.tg_connect,
config.general.upstream_unhealthy_fail_threshold,
config.general.upstream_connect_failfast_hard_errors,
stats.clone(),
));
let _config_tx = config_tx;
RuntimeGeneration::new(
id,
config_rx,
admission_rx,
stats,
upstream_manager,
Arc::new(ReplayChecker::new(128, Duration::from_secs(60))),
Arc::new(BufferPool::with_config(4096, 16)),
Arc::new(SecureRandom::new()),
None,
Arc::new(RwLock::new(None)),
Arc::new(RouteRuntimeController::new(RelayRouteMode::Direct)),
None,
Arc::new(UserIpTracker::new()),
Arc::new(BeobachtenStore::new()),
ProxySharedState::new(),
Arc::new(Semaphore::new(64)),
RuntimeTaskScope::new(),
)
}
#[cfg(test)]
mod tests {
use super::*;
use tokio::sync::Barrier;
#[tokio::test]
async fn stop_sessions_cancels_tracked_future() {
let generation = test_runtime_generation(1, ProxyConfig::default());
let started = Arc::new(tokio::sync::Notify::new());
let dropped = Arc::new(tokio::sync::Notify::new());
let started_task = started.clone();
let dropped_task = dropped.clone();
assert!(generation.spawn_session(async move {
struct DropSignal(Arc<tokio::sync::Notify>);
impl Drop for DropSignal {
fn drop(&mut self) {
self.0.notify_one();
}
}
let _drop_signal = DropSignal(dropped_task);
started_task.notify_one();
std::future::pending::<()>().await;
}));
started.notified().await;
generation.stop_sessions().await;
tokio::time::timeout(Duration::from_secs(1), dropped.notified())
.await
.unwrap();
assert!(!generation.spawn_session(async {}));
}
#[tokio::test]
async fn runtime_task_scope_joins_cancelled_background_task() {
let scope = RuntimeTaskScope::new();
scope.spawn(std::future::pending());
tokio::time::timeout(Duration::from_secs(1), scope.stop())
.await
.unwrap();
}
#[tokio::test]
async fn session_admission_waits_for_registration_started_before_cutover() {
let admission = Arc::new(SessionAdmission::new());
let registration = admission.try_register().unwrap();
admission.close();
assert!(admission.try_register().is_none());
let wait_admission = admission.clone();
let waiter = tokio::spawn(async move {
wait_admission.wait_for_registrations().await;
});
tokio::task::yield_now().await;
assert!(!waiter.is_finished());
drop(registration);
tokio::time::timeout(Duration::from_secs(1), waiter)
.await
.unwrap()
.unwrap();
admission.reopen();
assert!(admission.try_register().is_some());
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn cutover_never_leaves_late_session_registrations() {
const ATTEMPTS: usize = 10_000;
let admission = Arc::new(SessionAdmission::new());
let tracker = TaskTracker::new();
let cancel = CancellationToken::new();
let start = Arc::new(Barrier::new(ATTEMPTS + 1));
let live = Arc::new(AtomicUsize::new(0));
let mut attempts = tokio::task::JoinSet::new();
for _ in 0..ATTEMPTS {
let admission = admission.clone();
let tracker = tracker.clone();
let cancel = cancel.clone();
let start = start.clone();
let live = live.clone();
attempts.spawn(async move {
start.wait().await;
let Some(_registration) = admission.try_register() else {
return;
};
tracker.spawn(async move {
live.fetch_add(1, Ordering::AcqRel);
cancel.cancelled().await;
live.fetch_sub(1, Ordering::AcqRel);
});
});
}
start.wait().await;
admission.close();
admission.wait_for_registrations().await;
tracker.close();
cancel.cancel();
while attempts.join_next().await.is_some() {}
tokio::time::timeout(Duration::from_secs(1), tracker.wait())
.await
.unwrap();
assert_eq!(live.load(Ordering::Acquire), 0);
assert!(admission.try_register().is_none());
}
}
+43 -196
View File
@@ -3,35 +3,33 @@ use std::net::{IpAddr, SocketAddr};
use std::sync::Arc; use std::sync::Arc;
use std::time::Duration; use std::time::Duration;
use arc_swap::ArcSwap;
use tokio::net::TcpListener; use tokio::net::TcpListener;
#[cfg(unix)] #[cfg(unix)]
use tokio::net::UnixListener; use tokio::net::UnixListener;
use tokio::sync::{RwLock, Semaphore, watch};
use tracing::{debug, error, info, warn}; use tracing::{debug, error, info, warn};
use crate::config::{ProxyConfig, RstOnCloseMode}; use crate::config::{ProxyConfig, RstOnCloseMode};
use crate::crypto::SecureRandom;
use crate::ip_tracker::UserIpTracker;
use crate::proxy::ClientHandler; use crate::proxy::ClientHandler;
use crate::proxy::route_mode::RouteRuntimeController;
use crate::proxy::shared_state::ProxySharedState;
use crate::startup::{COMPONENT_LISTENERS_BIND, StartupTracker}; use crate::startup::{COMPONENT_LISTENERS_BIND, StartupTracker};
use crate::stats::beobachten::BeobachtenStore;
use crate::stats::{ReplayChecker, Stats};
use crate::stream::BufferPool;
use crate::tls_front::TlsFrontCache;
use crate::transport::middle_proxy::MePool;
use crate::transport::socket::set_linger_zero; use crate::transport::socket::set_linger_zero;
use crate::transport::{ListenOptions, UpstreamManager, create_listener, find_listener_processes}; use crate::transport::{ListenOptions, create_listener, find_listener_processes};
use super::generation::RuntimeGeneration;
use super::helpers::{ use super::helpers::{
expected_handshake_close_description, is_expected_handshake_eof, peer_close_description, expected_handshake_close_description, is_expected_handshake_eof, peer_close_description,
print_proxy_links, print_proxy_links,
}; };
#[cfg(unix)]
mod unix;
#[cfg(unix)]
pub(crate) use unix::spawn_unix_accept_loop;
pub(crate) struct BoundListeners { pub(crate) struct BoundListeners {
pub(crate) listeners: Vec<(TcpListener, bool)>, pub(crate) listeners: Vec<(TcpListener, bool)>,
pub(crate) has_unix_listener: bool, #[cfg(unix)]
pub(crate) unix_listener: Option<UnixListener>,
} }
fn listener_port_or_legacy(listener: &crate::config::ListenerConfig, config: &ProxyConfig) -> u16 { fn listener_port_or_legacy(listener: &crate::config::ListenerConfig, config: &ProxyConfig) -> u16 {
@@ -59,21 +57,6 @@ pub(crate) async fn bind_listeners(
detected_ip_v4: Option<IpAddr>, detected_ip_v4: Option<IpAddr>,
detected_ip_v6: Option<IpAddr>, detected_ip_v6: Option<IpAddr>,
startup_tracker: &Arc<StartupTracker>, startup_tracker: &Arc<StartupTracker>,
config_rx: watch::Receiver<Arc<ProxyConfig>>,
admission_rx: watch::Receiver<bool>,
stats: Arc<Stats>,
upstream_manager: Arc<UpstreamManager>,
replay_checker: Arc<ReplayChecker>,
buffer_pool: Arc<BufferPool>,
rng: Arc<SecureRandom>,
me_pool: Option<Arc<MePool>>,
me_pool_runtime: Arc<RwLock<Option<Arc<MePool>>>>,
route_runtime: Arc<RouteRuntimeController>,
tls_cache: Option<Arc<TlsFrontCache>>,
ip_tracker: Arc<UserIpTracker>,
beobachten: Arc<BeobachtenStore>,
shared: Arc<ProxySharedState>,
max_connections: Arc<Semaphore>,
) -> Result<BoundListeners, Box<dyn Error>> { ) -> Result<BoundListeners, Box<dyn Error>> {
startup_tracker startup_tracker
.start_component( .start_component(
@@ -218,7 +201,8 @@ pub(crate) async fn bind_listeners(
print_proxy_links(&host, port, config); print_proxy_links(&host, port, config);
} }
let mut has_unix_listener = false; #[cfg(unix)]
let mut unix_listener_out = None;
#[cfg(unix)] #[cfg(unix)]
if let Some(ref unix_path) = config.server.listen_unix_sock { if let Some(ref unix_path) = config.server.listen_unix_sock {
let _ = tokio::fs::remove_file(unix_path).await; let _ = tokio::fs::remove_file(unix_path).await;
@@ -251,122 +235,13 @@ pub(crate) async fn bind_listeners(
info!("Listening on unix:{}", unix_path); info!("Listening on unix:{}", unix_path);
} }
has_unix_listener = true; unix_listener_out = Some(unix_listener);
}
let mut config_rx_unix: watch::Receiver<Arc<ProxyConfig>> = config_rx.clone(); #[cfg(unix)]
let admission_rx_unix = admission_rx.clone(); let has_unix_listener = unix_listener_out.is_some();
let stats = stats.clone(); #[cfg(not(unix))]
let upstream_manager = upstream_manager.clone(); let has_unix_listener = false;
let replay_checker = replay_checker.clone();
let buffer_pool = buffer_pool.clone();
let rng = rng.clone();
let me_pool = me_pool.clone();
let me_pool_runtime = me_pool_runtime.clone();
let route_runtime = route_runtime.clone();
let tls_cache = tls_cache.clone();
let ip_tracker = ip_tracker.clone();
let beobachten = beobachten.clone();
let shared = shared.clone();
let max_connections_unix = max_connections.clone();
tokio::spawn(async move {
let unix_conn_counter = Arc::new(std::sync::atomic::AtomicU64::new(1));
loop {
match unix_listener.accept().await {
Ok((stream, _)) => {
if !*admission_rx_unix.borrow() {
drop(stream);
continue;
}
let accept_permit_timeout_ms =
config_rx_unix.borrow().server.accept_permit_timeout_ms;
let permit = if accept_permit_timeout_ms == 0 {
match max_connections_unix.clone().acquire_owned().await {
Ok(permit) => permit,
Err(_) => {
error!("Connection limiter is closed");
break;
}
}
} else {
match tokio::time::timeout(
Duration::from_millis(accept_permit_timeout_ms),
max_connections_unix.clone().acquire_owned(),
)
.await
{
Ok(Ok(permit)) => permit,
Ok(Err(_)) => {
error!("Connection limiter is closed");
break;
}
Err(_) => {
stats.increment_accept_permit_timeout_total();
debug!(
timeout_ms = accept_permit_timeout_ms,
"Dropping accepted unix connection: permit wait timeout"
);
drop(stream);
continue;
}
}
};
let conn_id =
unix_conn_counter.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
let fake_peer =
SocketAddr::from(([127, 0, 0, 1], (conn_id % 65535) as u16));
let config = config_rx_unix.borrow_and_update().clone();
let stats = stats.clone();
let upstream_manager = upstream_manager.clone();
let replay_checker = replay_checker.clone();
let buffer_pool = buffer_pool.clone();
let rng = rng.clone();
let me_pool = me_pool.clone();
let me_pool_runtime = me_pool_runtime.clone();
let route_runtime = route_runtime.clone();
let tls_cache = tls_cache.clone();
let ip_tracker = ip_tracker.clone();
let beobachten = beobachten.clone();
let shared = shared.clone();
let proxy_protocol_enabled = config.server.proxy_protocol;
tokio::spawn(async move {
let _permit = permit;
if let Err(e) =
crate::proxy::client::handle_client_stream_with_shared_and_pool_runtime(
stream,
fake_peer,
config,
stats,
upstream_manager,
replay_checker,
buffer_pool,
rng,
me_pool,
Some(me_pool_runtime),
route_runtime,
tls_cache,
ip_tracker,
beobachten,
shared,
proxy_protocol_enabled,
)
.await
{
debug!(error = %e, "Unix socket connection error");
}
});
}
Err(e) => {
error!("Unix socket accept error: {}", e);
tokio::time::sleep(Duration::from_millis(100)).await;
}
}
}
});
}
startup_tracker startup_tracker
.complete_component( .complete_component(
@@ -381,51 +256,25 @@ pub(crate) async fn bind_listeners(
Ok(BoundListeners { Ok(BoundListeners {
listeners, listeners,
has_unix_listener, #[cfg(unix)]
unix_listener: unix_listener_out,
}) })
} }
#[allow(clippy::too_many_arguments)]
pub(crate) fn spawn_tcp_accept_loops( pub(crate) fn spawn_tcp_accept_loops(
listeners: Vec<(TcpListener, bool)>, listeners: Vec<(TcpListener, bool)>,
config_rx: watch::Receiver<Arc<ProxyConfig>>, active_runtime: Arc<ArcSwap<RuntimeGeneration>>,
admission_rx: watch::Receiver<bool>,
stats: Arc<Stats>,
upstream_manager: Arc<UpstreamManager>,
replay_checker: Arc<ReplayChecker>,
buffer_pool: Arc<BufferPool>,
rng: Arc<SecureRandom>,
me_pool: Option<Arc<MePool>>,
me_pool_runtime: Arc<RwLock<Option<Arc<MePool>>>>,
route_runtime: Arc<RouteRuntimeController>,
tls_cache: Option<Arc<TlsFrontCache>>,
ip_tracker: Arc<UserIpTracker>,
beobachten: Arc<BeobachtenStore>,
shared: Arc<ProxySharedState>,
max_connections: Arc<Semaphore>,
) { ) {
for (listener, listener_proxy_protocol) in listeners { for (listener, listener_proxy_protocol) in listeners {
let mut config_rx: watch::Receiver<Arc<ProxyConfig>> = config_rx.clone(); let active_runtime = active_runtime.clone();
let admission_rx_tcp = admission_rx.clone();
let stats = stats.clone();
let upstream_manager = upstream_manager.clone();
let replay_checker = replay_checker.clone();
let buffer_pool = buffer_pool.clone();
let rng = rng.clone();
let me_pool = me_pool.clone();
let me_pool_runtime = me_pool_runtime.clone();
let route_runtime = route_runtime.clone();
let tls_cache = tls_cache.clone();
let ip_tracker = ip_tracker.clone();
let beobachten = beobachten.clone();
let shared = shared.clone();
let max_connections_tcp = max_connections.clone();
tokio::spawn(async move { tokio::spawn(async move {
loop { loop {
match listener.accept().await { match listener.accept().await {
Ok((stream, peer_addr)) => { Ok((stream, peer_addr)) => {
let rst_mode = config_rx.borrow().general.rst_on_close; let runtime = active_runtime.load_full();
let config = runtime.config();
let rst_mode = config.general.rst_on_close;
#[cfg(unix)] #[cfg(unix)]
let raw_fd = { let raw_fd = {
use std::os::unix::io::AsRawFd; use std::os::unix::io::AsRawFd;
@@ -434,15 +283,14 @@ pub(crate) fn spawn_tcp_accept_loops(
if matches!(rst_mode, RstOnCloseMode::Errors | RstOnCloseMode::Always) { if matches!(rst_mode, RstOnCloseMode::Errors | RstOnCloseMode::Always) {
let _ = set_linger_zero(&stream); let _ = set_linger_zero(&stream);
} }
if !*admission_rx_tcp.borrow() { if !*runtime.admission_rx.borrow() {
debug!(peer = %peer_addr, "Admission gate closed, dropping connection"); debug!(peer = %peer_addr, "Admission gate closed, dropping connection");
drop(stream); drop(stream);
continue; continue;
} }
let accept_permit_timeout_ms = let accept_permit_timeout_ms = config.server.accept_permit_timeout_ms;
config_rx.borrow().server.accept_permit_timeout_ms;
let permit = if accept_permit_timeout_ms == 0 { let permit = if accept_permit_timeout_ms == 0 {
match max_connections_tcp.clone().acquire_owned().await { match runtime.max_connections.clone().acquire_owned().await {
Ok(permit) => permit, Ok(permit) => permit,
Err(_) => { Err(_) => {
error!("Connection limiter is closed"); error!("Connection limiter is closed");
@@ -452,7 +300,7 @@ pub(crate) fn spawn_tcp_accept_loops(
} else { } else {
match tokio::time::timeout( match tokio::time::timeout(
Duration::from_millis(accept_permit_timeout_ms), Duration::from_millis(accept_permit_timeout_ms),
max_connections_tcp.clone().acquire_owned(), runtime.max_connections.clone().acquire_owned(),
) )
.await .await
{ {
@@ -462,7 +310,7 @@ pub(crate) fn spawn_tcp_accept_loops(
break; break;
} }
Err(_) => { Err(_) => {
stats.increment_accept_permit_timeout_total(); runtime.stats.increment_accept_permit_timeout_total();
debug!( debug!(
peer = %peer_addr, peer = %peer_addr,
timeout_ms = accept_permit_timeout_ms, timeout_ms = accept_permit_timeout_ms,
@@ -473,24 +321,23 @@ pub(crate) fn spawn_tcp_accept_loops(
} }
} }
}; };
let config = config_rx.borrow_and_update().clone(); let stats = runtime.stats.clone();
let stats = stats.clone(); let upstream_manager = runtime.upstream_manager.clone();
let upstream_manager = upstream_manager.clone(); let replay_checker = runtime.replay_checker.clone();
let replay_checker = replay_checker.clone(); let buffer_pool = runtime.buffer_pool.clone();
let buffer_pool = buffer_pool.clone(); let rng = runtime.rng.clone();
let rng = rng.clone(); let me_pool = runtime.me_pool.clone();
let me_pool = me_pool.clone(); let me_pool_runtime = runtime.me_pool_runtime.clone();
let me_pool_runtime = me_pool_runtime.clone(); let route_runtime = runtime.route_runtime.clone();
let route_runtime = route_runtime.clone(); let tls_cache = runtime.tls_cache.clone();
let tls_cache = tls_cache.clone(); let ip_tracker = runtime.ip_tracker.clone();
let ip_tracker = ip_tracker.clone(); let beobachten = runtime.beobachten.clone();
let beobachten = beobachten.clone(); let shared = runtime.proxy_shared.clone();
let shared = shared.clone();
let proxy_protocol_enabled = listener_proxy_protocol; let proxy_protocol_enabled = listener_proxy_protocol;
let real_peer_report = Arc::new(std::sync::Mutex::new(None)); let real_peer_report = Arc::new(std::sync::Mutex::new(None));
let real_peer_report_for_handler = real_peer_report.clone(); let real_peer_report_for_handler = real_peer_report.clone();
tokio::spawn(async move { let _ = runtime.spawn_session(async move {
let _permit = permit; let _permit = permit;
if let Err(e) = ClientHandler::new_with_shared( if let Err(e) = ClientHandler::new_with_shared(
stream, stream,
+117
View File
@@ -0,0 +1,117 @@
use std::net::SocketAddr;
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::Duration;
use arc_swap::ArcSwap;
use tokio::net::UnixListener;
use tracing::{debug, error};
use super::RuntimeGeneration;
pub(crate) fn spawn_unix_accept_loop(
listener: Option<UnixListener>,
active_runtime: Arc<ArcSwap<RuntimeGeneration>>,
) {
let Some(listener) = listener else {
return;
};
tokio::spawn(async move {
let connection_counter = AtomicU64::new(1);
loop {
match listener.accept().await {
Ok((stream, _)) => {
let runtime = active_runtime.load_full();
if !*runtime.admission_rx.borrow() {
drop(stream);
continue;
}
let config = runtime.config();
let timeout_ms = config.server.accept_permit_timeout_ms;
let permit = if timeout_ms == 0 {
match runtime.max_connections.clone().acquire_owned().await {
Ok(permit) => permit,
Err(_) => {
error!("Connection limiter is closed");
break;
}
}
} else {
match tokio::time::timeout(
Duration::from_millis(timeout_ms),
runtime.max_connections.clone().acquire_owned(),
)
.await
{
Ok(Ok(permit)) => permit,
Ok(Err(_)) => {
error!("Connection limiter is closed");
break;
}
Err(_) => {
runtime.stats.increment_accept_permit_timeout_total();
debug!(
timeout_ms,
"Dropping accepted unix connection: permit wait timeout"
);
drop(stream);
continue;
}
}
};
let connection_id = connection_counter.fetch_add(1, Ordering::Relaxed);
let fake_peer =
SocketAddr::from(([127, 0, 0, 1], (connection_id % 65535) as u16));
let stats = runtime.stats.clone();
let upstream_manager = runtime.upstream_manager.clone();
let replay_checker = runtime.replay_checker.clone();
let buffer_pool = runtime.buffer_pool.clone();
let rng = runtime.rng.clone();
let me_pool = runtime.me_pool.clone();
let me_pool_runtime = runtime.me_pool_runtime.clone();
let route_runtime = runtime.route_runtime.clone();
let tls_cache = runtime.tls_cache.clone();
let ip_tracker = runtime.ip_tracker.clone();
let beobachten = runtime.beobachten.clone();
let shared = runtime.proxy_shared.clone();
let proxy_protocol_enabled = config.server.proxy_protocol;
let _ = runtime.spawn_session(async move {
let _permit = permit;
if let Err(error) =
crate::proxy::client::handle_client_stream_with_shared_and_pool_runtime(
stream,
fake_peer,
config,
stats,
upstream_manager,
replay_checker,
buffer_pool,
rng,
me_pool,
Some(me_pool_runtime),
route_runtime,
tls_cache,
ip_tracker,
beobachten,
shared,
proxy_protocol_enabled,
)
.await
{
debug!(error = %error, "Unix socket connection error");
}
});
}
Err(error) => {
error!(error = %error, "Unix socket accept error");
tokio::time::sleep(Duration::from_millis(100)).await;
}
}
}
});
}
+143 -158
View File
@@ -1,9 +1,11 @@
#![allow(clippy::too_many_arguments)] #![allow(clippy::too_many_arguments)]
use std::future::Future;
use std::sync::Arc; use std::sync::Arc;
use std::time::Duration; use std::time::Duration;
use tokio::sync::{RwLock, watch}; use tokio::sync::{RwLock, watch};
use tokio_util::task::AbortOnDropHandle;
use tracing::{error, info, warn}; use tracing::{error, info, warn};
use crate::config::ProxyConfig; use crate::config::ProxyConfig;
@@ -17,8 +19,61 @@ use crate::stats::Stats;
use crate::transport::UpstreamManager; use crate::transport::UpstreamManager;
use crate::transport::middle_proxy::MePool; use crate::transport::middle_proxy::MePool;
use super::generation::RuntimeTaskScope;
use super::helpers::load_startup_proxy_config_snapshot; use super::helpers::load_startup_proxy_config_snapshot;
async fn supervise_me_task<F, Fut>(task_name: &'static str, mut task: F)
where
F: FnMut() -> Fut,
Fut: Future<Output = ()> + Send + 'static,
{
loop {
let result = AbortOnDropHandle::new(tokio::spawn(task())).await;
match result {
Ok(()) => warn!(
task = task_name,
"Middle-End supervisor task exited unexpectedly, restarting"
),
Err(error) => {
error!(task = task_name, error = %error, "Middle-End supervisor task panicked, restarting in 1s");
tokio::time::sleep(Duration::from_secs(1)).await;
}
}
}
}
fn spawn_me_supervisors(
task_scope: RuntimeTaskScope,
pool: Arc<MePool>,
rng: Arc<SecureRandom>,
min_connections: usize,
) {
let health_pool = pool.clone();
let health_rng = rng;
task_scope.spawn(supervise_me_task("health_monitor", move || {
let pool = health_pool.clone();
let rng = health_rng.clone();
async move {
crate::transport::middle_proxy::me_health_monitor(pool, rng, min_connections).await;
}
}));
let drain_pool = pool.clone();
task_scope.spawn(supervise_me_task("drain_timeout_enforcer", move || {
let pool = drain_pool.clone();
async move {
crate::transport::middle_proxy::me_drain_timeout_enforcer(pool).await;
}
}));
task_scope.spawn(supervise_me_task("zombie_writer_watchdog", move || {
let pool = pool.clone();
async move {
crate::transport::middle_proxy::me_zombie_writer_watchdog(pool).await;
}
}));
}
pub(crate) async fn initialize_me_pool( pub(crate) async fn initialize_me_pool(
use_middle_proxy: bool, use_middle_proxy: bool,
config: &ProxyConfig, config: &ProxyConfig,
@@ -30,6 +85,7 @@ pub(crate) async fn initialize_me_pool(
stats: Arc<Stats>, stats: Arc<Stats>,
api_me_pool: Arc<RwLock<Option<Arc<MePool>>>>, api_me_pool: Arc<RwLock<Option<Arc<MePool>>>>,
me_ready_tx: watch::Sender<u64>, me_ready_tx: watch::Sender<u64>,
task_scope: RuntimeTaskScope,
) -> Option<Arc<MePool>> { ) -> Option<Arc<MePool>> {
if !use_middle_proxy { if !use_middle_proxy {
return None; return None;
@@ -52,15 +108,8 @@ pub(crate) async fn initialize_me_pool(
.as_ref() .as_ref()
.map(|tag| hex::decode(tag).expect("general.ad_tag must be validated before startup")); .map(|tag| hex::decode(tag).expect("general.ad_tag must be validated before startup"));
// ============================================================= // The Telegram proxy-secret authenticates ME RPC and is distinct from client secrets.
// CRITICAL: Download Telegram proxy-secret (NOT user secret!) // It corresponds to the C MTProxy --aes-pwd input and may be fetched from Telegram.
//
// C MTProxy uses TWO separate secrets:
// -S flag = 16-byte user secret for client obfuscation
// --aes-pwd = 32-512 byte binary file for ME RPC auth
//
// proxy-secret is from: https://core.telegram.org/getProxySecret
// =============================================================
let proxy_secret_path = config.general.proxy_secret_path.as_deref(); let proxy_secret_path = config.general.proxy_secret_path.as_deref();
let pool_size = config.general.middle_proxy_pool_size.max(1); let pool_size = config.general.middle_proxy_pool_size.max(1);
let proxy_secret = loop { let proxy_secret = loop {
@@ -279,6 +328,7 @@ pub(crate) async fn initialize_me_pool(
config.general.me_writer_pick_sample_size, config.general.me_writer_pick_sample_size,
config.general.me_socks_kdf_policy, config.general.me_socks_kdf_policy,
config.general.me_writer_cmd_channel_capacity, config.general.me_writer_cmd_channel_capacity,
config.general.me_writer_byte_budget_bytes,
config.general.me_route_channel_capacity, config.general.me_route_channel_capacity,
config.general.me_route_backpressure_enabled, config.general.me_route_backpressure_enabled,
config.general.me_route_fairshare_enabled, config.general.me_route_fairshare_enabled,
@@ -318,23 +368,13 @@ pub(crate) async fn initialize_me_pool(
let rng_bg = rng.clone(); let rng_bg = rng.clone();
let startup_tracker_bg = startup_tracker.clone(); let startup_tracker_bg = startup_tracker.clone();
let me_ready_tx_bg = me_ready_tx.clone(); let me_ready_tx_bg = me_ready_tx.clone();
let task_scope_bg = task_scope.clone();
let retry_limit = if me_init_retry_attempts == 0 { let retry_limit = if me_init_retry_attempts == 0 {
String::from("unlimited") String::from("unlimited")
} else { } else {
me_init_retry_attempts.to_string() me_init_retry_attempts.to_string()
}; };
std::thread::spawn(move || { task_scope.spawn(async move {
let runtime = match tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
{
Ok(runtime) => runtime,
Err(error) => {
error!(error = %error, "Failed to build background runtime for ME initialization");
return;
}
};
runtime.block_on(async move {
let mut init_attempt: u32 = 0; let mut init_attempt: u32 = 0;
loop { loop {
init_attempt = init_attempt.saturating_add(1); init_attempt = init_attempt.saturating_add(1);
@@ -358,80 +398,18 @@ pub(crate) async fn initialize_me_pool(
attempt = init_attempt, attempt = init_attempt,
"Middle-End pool initialized successfully" "Middle-End pool initialized successfully"
); );
spawn_me_supervisors(
// ── Supervised background tasks ────────────────── task_scope_bg,
// Each task runs inside a nested tokio::spawn so pool_bg.clone(),
// that a panic is caught via JoinHandle and the rng_bg.clone(),
// outer loop restarts the task automatically. pool_size,
let pool_health = pool_bg.clone(); );
let rng_health = rng_bg.clone(); break;
let min_conns = pool_size;
tokio::spawn(async move {
loop {
let p = pool_health.clone();
let r = rng_health.clone();
let res = tokio::spawn(async move {
crate::transport::middle_proxy::me_health_monitor(
p, r, min_conns,
)
.await;
})
.await;
match res {
Ok(()) => warn!("me_health_monitor exited unexpectedly, restarting"),
Err(e) => {
error!(error = %e, "me_health_monitor panicked, restarting in 1s");
tokio::time::sleep(Duration::from_secs(1)).await;
}
}
}
});
let pool_drain_enforcer = pool_bg.clone();
tokio::spawn(async move {
loop {
let p = pool_drain_enforcer.clone();
let res = tokio::spawn(async move {
crate::transport::middle_proxy::me_drain_timeout_enforcer(p).await;
})
.await;
match res {
Ok(()) => warn!("me_drain_timeout_enforcer exited unexpectedly, restarting"),
Err(e) => {
error!(error = %e, "me_drain_timeout_enforcer panicked, restarting in 1s");
tokio::time::sleep(Duration::from_secs(1)).await;
}
}
}
});
let pool_watchdog = pool_bg.clone();
tokio::spawn(async move {
loop {
let p = pool_watchdog.clone();
let res = tokio::spawn(async move {
crate::transport::middle_proxy::me_zombie_writer_watchdog(p).await;
})
.await;
match res {
Ok(()) => warn!("me_zombie_writer_watchdog exited unexpectedly, restarting"),
Err(e) => {
error!(error = %e, "me_zombie_writer_watchdog panicked, restarting in 1s");
tokio::time::sleep(Duration::from_secs(1)).await;
}
}
}
});
// CRITICAL: keep the current-thread runtime
// alive. Without this, block_on() returns,
// the Runtime is dropped, and ALL spawned
// background tasks (health monitor, drain
// enforcer, zombie watchdog) are silently
// cancelled — causing the draining-writer
// leak that brought us here.
std::future::pending::<()>().await;
unreachable!();
} }
Err(e) => { Err(e) => {
startup_tracker_bg.set_me_last_error(Some(e.to_string())).await; startup_tracker_bg
.set_me_last_error(Some(e.to_string()))
.await;
if init_attempt >= me_init_warn_after_attempts { if init_attempt >= me_init_warn_after_attempts {
warn!( warn!(
error = %e, error = %e,
@@ -455,7 +433,6 @@ pub(crate) async fn initialize_me_pool(
} }
} }
}); });
});
startup_tracker startup_tracker
.set_me_status(StartupMeStatus::Initializing, "background_init") .set_me_status(StartupMeStatus::Initializing, "background_init")
.await; .await;
@@ -489,70 +466,12 @@ pub(crate) async fn initialize_me_pool(
"Middle-End pool initialized successfully" "Middle-End pool initialized successfully"
); );
// ── Supervised background tasks ────────────────── spawn_me_supervisors(
let pool_clone = pool.clone(); task_scope.clone(),
let rng_clone = rng.clone(); pool.clone(),
let min_conns = pool_size; rng.clone(),
tokio::spawn(async move { pool_size,
loop { );
let p = pool_clone.clone();
let r = rng_clone.clone();
let res = tokio::spawn(async move {
crate::transport::middle_proxy::me_health_monitor(
p, r, min_conns,
)
.await;
})
.await;
match res {
Ok(()) => warn!(
"me_health_monitor exited unexpectedly, restarting"
),
Err(e) => {
error!(error = %e, "me_health_monitor panicked, restarting in 1s");
tokio::time::sleep(Duration::from_secs(1)).await;
}
}
}
});
let pool_drain_enforcer = pool.clone();
tokio::spawn(async move {
loop {
let p = pool_drain_enforcer.clone();
let res = tokio::spawn(async move {
crate::transport::middle_proxy::me_drain_timeout_enforcer(p).await;
})
.await;
match res {
Ok(()) => warn!(
"me_drain_timeout_enforcer exited unexpectedly, restarting"
),
Err(e) => {
error!(error = %e, "me_drain_timeout_enforcer panicked, restarting in 1s");
tokio::time::sleep(Duration::from_secs(1)).await;
}
}
}
});
let pool_watchdog = pool.clone();
tokio::spawn(async move {
loop {
let p = pool_watchdog.clone();
let res = tokio::spawn(async move {
crate::transport::middle_proxy::me_zombie_writer_watchdog(p).await;
})
.await;
match res {
Ok(()) => warn!(
"me_zombie_writer_watchdog exited unexpectedly, restarting"
),
Err(e) => {
error!(error = %e, "me_zombie_writer_watchdog panicked, restarting in 1s");
tokio::time::sleep(Duration::from_secs(1)).await;
}
}
}
});
break Some(pool); break Some(pool);
} }
@@ -665,3 +584,69 @@ pub(crate) async fn initialize_me_pool(
} }
} }
} }
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicUsize, Ordering};
use tokio::sync::Notify;
struct DropSignal(Arc<Notify>);
impl Drop for DropSignal {
fn drop(&mut self) {
self.0.notify_one();
}
}
#[tokio::test]
async fn scoped_supervisor_aborts_its_current_child() {
let scope = RuntimeTaskScope::new();
let dropped = Arc::new(Notify::new());
let dropped_for_task = dropped.clone();
scope.spawn(supervise_me_task("test", move || {
let dropped = dropped_for_task.clone();
async move {
let _signal = DropSignal(dropped);
std::future::pending::<()>().await;
}
}));
tokio::task::yield_now().await;
scope.stop().await;
tokio::time::timeout(Duration::from_secs(1), dropped.notified())
.await
.unwrap();
}
#[tokio::test]
async fn supervisor_restarts_exited_child_and_stops_with_runtime_scope() {
let scope = RuntimeTaskScope::new();
let starts = Arc::new(AtomicUsize::new(0));
let restarted = Arc::new(Notify::new());
let starts_task = starts.clone();
let restarted_task = restarted.clone();
scope.spawn(supervise_me_task("restart_test", move || {
let starts = starts_task.clone();
let restarted = restarted_task.clone();
async move {
if starts.fetch_add(1, Ordering::AcqRel) + 1 >= 3 {
restarted.notify_one();
}
}
}));
tokio::time::timeout(Duration::from_secs(1), restarted.notified())
.await
.unwrap();
scope.stop().await;
let stopped_at = starts.load(Ordering::Acquire);
for _ in 0..100 {
tokio::task::yield_now().await;
}
assert!(stopped_at >= 3);
assert_eq!(starts.load(Ordering::Acquire), stopped_at);
}
}
+121 -73
View File
@@ -13,19 +13,24 @@
// - shutdown: graceful shutdown sequence and uptime logging. // - shutdown: graceful shutdown sequence and uptime logging.
mod admission; mod admission;
mod connectivity; mod connectivity;
pub(crate) mod generation;
mod helpers; mod helpers;
mod listeners; mod listeners;
mod me_startup; mod me_startup;
pub(crate) mod reload;
mod reload_supervisor;
pub(crate) mod runtime_build;
mod runtime_tasks; mod runtime_tasks;
mod shutdown; mod shutdown;
mod tls_bootstrap; mod tls_bootstrap;
use arc_swap::ArcSwap;
use std::net::{IpAddr, SocketAddr}; use std::net::{IpAddr, SocketAddr};
use std::sync::Arc; use std::sync::Arc;
use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH}; use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH};
use tokio::sync::{RwLock, Semaphore, watch}; use tokio::sync::{RwLock, Semaphore, watch};
use tracing::{error, info, warn}; use tracing::{error, info, warn};
use tracing_subscriber::{EnvFilter, fmt, prelude::*, reload}; use tracing_subscriber::{EnvFilter, fmt, prelude::*, reload as tracing_reload};
use crate::api; use crate::api;
use crate::config::{LogLevel, ProxyConfig}; use crate::config::{LogLevel, ProxyConfig};
@@ -33,6 +38,9 @@ use crate::conntrack_control;
use crate::crypto::SecureRandom; use crate::crypto::SecureRandom;
use crate::ip_tracker::UserIpTracker; use crate::ip_tracker::UserIpTracker;
use crate::network::probe::{decide_network_capabilities, log_probe_result, run_probe}; use crate::network::probe::{decide_network_capabilities, log_probe_result, run_probe};
use crate::proxy::direct_buffer_budget::{
DirectBufferBudget, resolve_direct_buffer_hard_limit, run_direct_buffer_budget_controller,
};
use crate::proxy::route_mode::{RelayRouteMode, RouteRuntimeController}; use crate::proxy::route_mode::{RelayRouteMode, RouteRuntimeController};
use crate::proxy::shared_state::ProxySharedState; use crate::proxy::shared_state::ProxySharedState;
use crate::startup::{ use crate::startup::{
@@ -43,7 +51,7 @@ use crate::startup::{
}; };
use crate::stats::beobachten::BeobachtenStore; use crate::stats::beobachten::BeobachtenStore;
use crate::stats::telemetry::TelemetryPolicy; use crate::stats::telemetry::TelemetryPolicy;
use crate::stats::{ReplayChecker, Stats}; use crate::stats::{QuotaStore, ReplayChecker, Stats};
use crate::stream::BufferPool; use crate::stream::BufferPool;
use crate::synlimit_control; use crate::synlimit_control;
use crate::transport::UpstreamManager; use crate::transport::UpstreamManager;
@@ -340,7 +348,7 @@ async fn run_telemt_core(
} }
}; };
let (filter_layer, filter_handle) = let (filter_layer, filter_handle) =
reload::Layer::new(EnvFilter::new(initial_filter_spec.clone())); tracing_reload::Layer::new(EnvFilter::new(initial_filter_spec.clone()));
startup_tracker startup_tracker
.start_component( .start_component(
COMPONENT_TRACING_INIT, COMPONENT_TRACING_INIT,
@@ -384,6 +392,7 @@ async fn run_telemt_core(
_logging_guard = Some(guard); _logging_guard = Some(guard);
} }
} }
let runtime_log_filter = runtime_tasks::RuntimeLogFilter::new(filter_handle);
startup_tracker startup_tracker
.complete_component( .complete_component(
@@ -430,12 +439,15 @@ async fn run_telemt_core(
warn!("Using default tls_domain. Consider setting a custom domain."); warn!("Using default tls_domain. Consider setting a custom domain.");
} }
let stats = Arc::new(Stats::new()); let quota_store = Arc::new(QuotaStore::default());
let stats = Arc::new(Stats::with_quota_store(quota_store.clone()));
let runtime_task_scope = generation::RuntimeTaskScope::new();
stats.apply_telemetry_policy(TelemetryPolicy::from_config(&config.general.telemetry)); stats.apply_telemetry_policy(TelemetryPolicy::from_config(&config.general.telemetry));
let quota_state_path = config.general.quota_state_path.clone(); let quota_state_path = config.general.quota_state_path.clone();
crate::quota_state::load_quota_state(&quota_state_path, stats.as_ref()).await; crate::quota_state::load_quota_state(&quota_state_path, stats.as_ref()).await;
let upstream_manager = Arc::new(UpstreamManager::new( let upstream_manager = Arc::new(
UpstreamManager::new(
config.upstreams.clone(), config.upstreams.clone(),
config.general.upstream_connect_retry_attempts, config.general.upstream_connect_retry_attempts,
config.general.upstream_connect_retry_backoff_ms, config.general.upstream_connect_retry_backoff_ms,
@@ -444,7 +456,9 @@ async fn run_telemt_core(
config.general.upstream_unhealthy_fail_threshold, config.general.upstream_unhealthy_fail_threshold,
config.general.upstream_connect_failfast_hard_errors, config.general.upstream_connect_failfast_hard_errors,
stats.clone(), stats.clone(),
)); )
.with_dns_overrides(&config.network.dns_overrides)?,
);
let ip_tracker = Arc::new(UserIpTracker::new()); let ip_tracker = Arc::new(UserIpTracker::new());
ip_tracker ip_tracker
.load_limits( .load_limits(
@@ -473,18 +487,31 @@ async fn run_telemt_core(
config.network.dns_overrides.len() config.network.dns_overrides.len()
); );
} }
let shared_state = ProxySharedState::new(); let direct_buffer_hard_limit =
resolve_direct_buffer_hard_limit(config.general.direct_relay_buffer_budget_max_bytes).await;
let direct_buffer_budget = DirectBufferBudget::new(direct_buffer_hard_limit);
info!(
hard_limit_bytes = direct_buffer_hard_limit,
configured_override_bytes = config.general.direct_relay_buffer_budget_max_bytes,
"Direct relay buffer budget initialized"
);
let shared_state =
ProxySharedState::new_with_direct_buffer_budget(direct_buffer_budget.clone());
shared_state.apply_user_enabled_config(&config.access.user_enabled); shared_state.apply_user_enabled_config(&config.access.user_enabled);
shared_state.traffic_limiter.apply_policy( shared_state.traffic_limiter.apply_policy(
config.access.user_rate_limits.clone(), config.access.user_rate_limits.clone(),
config.access.cidr_rate_limits.clone(), config.access.cidr_rate_limits.clone(),
); );
let (api_config_tx, api_config_rx) = watch::channel(Arc::new(config.clone()));
let (detected_ips_tx, detected_ips_rx) = watch::channel((None::<IpAddr>, None::<IpAddr>)); let (detected_ips_tx, detected_ips_rx) = watch::channel((None::<IpAddr>, None::<IpAddr>));
let initial_direct_first = config.general.use_middle_proxy && config.general.me2dc_fallback; let initial_direct_first = config.general.use_middle_proxy && config.general.me2dc_fallback;
let initial_admission_open = !config.general.use_middle_proxy || initial_direct_first; let initial_admission_open = !config.general.use_middle_proxy || initial_direct_first;
let (admission_tx, admission_rx) = watch::channel(initial_admission_open); let (admission_tx, admission_rx) = watch::channel(initial_admission_open);
let (reload_control, reload_commands) = reload::ReloadControl::channel(1);
let (active_runtime_tx, active_runtime_rx) =
watch::channel(None::<Arc<ArcSwap<generation::RuntimeGeneration>>>);
let (runtime_watch_tx, runtime_watch_rx) =
watch::channel(None::<generation::RuntimeWatchState>);
let initial_route_mode = if !config.general.use_middle_proxy || initial_direct_first { let initial_route_mode = if !config.general.use_middle_proxy || initial_direct_first {
RelayRouteMode::Direct RelayRouteMode::Direct
} else { } else {
@@ -518,12 +545,13 @@ async fn run_telemt_core(
let upstream_manager_api = upstream_manager.clone(); let upstream_manager_api = upstream_manager.clone();
let route_runtime_api = route_runtime.clone(); let route_runtime_api = route_runtime.clone();
let proxy_shared_api = shared_state.clone(); let proxy_shared_api = shared_state.clone();
let config_rx_api = api_config_rx.clone();
let admission_rx_api = admission_rx.clone();
let config_path_api = config_path.clone(); let config_path_api = config_path.clone();
let quota_state_path_api = quota_state_path.clone(); let quota_state_path_api = quota_state_path.clone();
let startup_tracker_api = startup_tracker.clone(); let startup_tracker_api = startup_tracker.clone();
let detected_ips_rx_api = detected_ips_rx.clone(); let detected_ips_rx_api = detected_ips_rx.clone();
let reload_control_api = reload_control.clone();
let active_runtime_rx_api = active_runtime_rx.clone();
let runtime_watch_rx_api = runtime_watch_rx.clone();
tokio::spawn(async move { tokio::spawn(async move {
api::serve( api::serve(
listen, listen,
@@ -533,13 +561,14 @@ async fn run_telemt_core(
route_runtime_api, route_runtime_api,
proxy_shared_api, proxy_shared_api,
upstream_manager_api, upstream_manager_api,
config_rx_api,
admission_rx_api,
config_path_api, config_path_api,
quota_state_path_api, quota_state_path_api,
detected_ips_rx_api, detected_ips_rx_api,
process_started_at_epoch_secs, process_started_at_epoch_secs,
startup_tracker_api, startup_tracker_api,
reload_control_api,
active_runtime_rx_api,
runtime_watch_rx_api,
) )
.await; .await;
}); });
@@ -579,8 +608,10 @@ async fn run_telemt_core(
&tls_domains, &tls_domains,
upstream_manager.clone(), upstream_manager.clone(),
&startup_tracker, &startup_tracker,
runtime_task_scope.clone(),
tls_bootstrap::TlsBootstrapPolicy::BestEffort,
) )
.await; .await?;
startup_tracker startup_tracker
.start_component( .start_component(
@@ -706,6 +737,7 @@ async fn run_telemt_core(
stats.clone(), stats.clone(),
api_me_pool.clone(), api_me_pool.clone(),
me_ready_tx.clone(), me_ready_tx.clone(),
runtime_task_scope.clone(),
) )
.await .await
}; };
@@ -793,16 +825,22 @@ async fn run_telemt_core(
rng.clone(), rng.clone(),
ip_tracker.clone(), ip_tracker.clone(),
beobachten.clone(), beobachten.clone(),
api_config_tx.clone(),
me_pool.clone(), me_pool.clone(),
shared_state.clone(), shared_state.clone(),
me_ready_tx.clone(), me_ready_tx.clone(),
runtime_task_scope.clone(),
) )
.await; .await;
let config_rx = runtime_watches.config_rx; let config_rx = runtime_watches.config_rx;
let log_level_rx = runtime_watches.log_level_rx; let log_level_rx = runtime_watches.log_level_rx;
let detected_ip_v4 = runtime_watches.detected_ip_v4; let detected_ip_v4 = runtime_watches.detected_ip_v4;
let detected_ip_v6 = runtime_watches.detected_ip_v6; let detected_ip_v6 = runtime_watches.detected_ip_v6;
runtime_log_filter.start(
has_rust_log,
&effective_log_level,
log_level_rx,
runtime_task_scope.clone(),
);
if direct_first_startup { if direct_first_startup {
let config_bg = config.clone(); let config_bg = config.clone();
@@ -815,7 +853,8 @@ async fn run_telemt_core(
let api_me_pool_bg = api_me_pool.clone(); let api_me_pool_bg = api_me_pool.clone();
let me_ready_tx_bg = me_ready_tx.clone(); let me_ready_tx_bg = me_ready_tx.clone();
let config_rx_bg = config_rx.clone(); let config_rx_bg = config_rx.clone();
tokio::spawn(async move { let task_scope_bg = runtime_task_scope.clone();
runtime_task_scope.spawn(async move {
let mut bootstrap_attempt: u32 = 0; let mut bootstrap_attempt: u32 = 0;
loop { loop {
bootstrap_attempt = bootstrap_attempt.saturating_add(1); bootstrap_attempt = bootstrap_attempt.saturating_add(1);
@@ -830,6 +869,7 @@ async fn run_telemt_core(
stats_bg.clone(), stats_bg.clone(),
api_me_pool_bg.clone(), api_me_pool_bg.clone(),
me_ready_tx_bg.clone(), me_ready_tx_bg.clone(),
task_scope_bg.clone(),
) )
.await; .await;
if let Some(pool) = pool { if let Some(pool) = pool {
@@ -839,6 +879,7 @@ async fn run_telemt_core(
pool, pool,
rng_bg, rng_bg,
me_ready_tx_bg, me_ready_tx_bg,
task_scope_bg,
); );
break; break;
} }
@@ -852,7 +893,7 @@ async fn run_telemt_core(
let startup_tracker_ready = startup_tracker.clone(); let startup_tracker_ready = startup_tracker.clone();
let api_me_pool_ready = api_me_pool.clone(); let api_me_pool_ready = api_me_pool.clone();
let mut me_ready_rx_transport = me_ready_tx.subscribe(); let mut me_ready_rx_transport = me_ready_tx.subscribe();
tokio::spawn(async move { runtime_task_scope.spawn(async move {
if me_ready_rx_transport.changed().await.is_ok() { if me_ready_rx_transport.changed().await.is_ok() {
if let Some(pool) = api_me_pool_ready.read().await.as_ref() { if let Some(pool) = api_me_pool_ready.read().await.as_ref() {
pool.set_runtime_ready(true); pool.set_runtime_ready(true);
@@ -874,13 +915,54 @@ async fn run_telemt_core(
&admission_tx, &admission_tx,
config_rx.clone(), config_rx.clone(),
me_ready_rx, me_ready_rx,
runtime_task_scope.clone(),
) )
.await; .await;
let _admission_tx_hold = admission_tx; let _admission_tx_hold = admission_tx;
conntrack_control::spawn_conntrack_controller( let conntrack_scope = runtime_task_scope.clone();
runtime_task_scope.spawn(conntrack_control::run_conntrack_controller(
config_rx.clone(), config_rx.clone(),
stats.clone(), stats.clone(),
shared_state.clone(), shared_state.clone(),
conntrack_scope.cancellation_token(),
));
runtime_task_scope.spawn(run_direct_buffer_budget_controller(
direct_buffer_budget,
buffer_pool.clone(),
stats.clone(),
shared_state.clone(),
config.server.max_connections,
));
let runtime_generation = generation::RuntimeGeneration::new(
1,
config_rx.clone(),
admission_rx.clone(),
stats.clone(),
upstream_manager.clone(),
replay_checker.clone(),
buffer_pool.clone(),
rng.clone(),
me_pool.clone(),
api_me_pool.clone(),
route_runtime.clone(),
tls_cache.clone(),
ip_tracker.clone(),
beobachten.clone(),
shared_state.clone(),
max_connections.clone(),
runtime_task_scope.clone(),
);
let active_runtime = Arc::new(ArcSwap::from(runtime_generation));
let reload_supervisor = reload_supervisor::ReloadSupervisor::spawn(
active_runtime.clone(),
reload_control,
reload_commands,
config_path.clone(),
quota_store,
detected_ips_tx,
runtime_log_filter,
runtime_watch_tx.clone(),
); );
let bound = listeners::bind_listeners( let bound = listeners::bind_listeners(
@@ -890,25 +972,15 @@ async fn run_telemt_core(
detected_ip_v4, detected_ip_v4,
detected_ip_v6, detected_ip_v6,
&startup_tracker, &startup_tracker,
config_rx.clone(),
admission_rx.clone(),
stats.clone(),
upstream_manager.clone(),
replay_checker.clone(),
buffer_pool.clone(),
rng.clone(),
me_pool.clone(),
api_me_pool.clone(),
route_runtime.clone(),
tls_cache.clone(),
ip_tracker.clone(),
beobachten.clone(),
shared_state.clone(),
max_connections.clone(),
) )
.await?; .await?;
let listeners = bound.listeners; let listeners = bound.listeners;
let has_unix_listener = bound.has_unix_listener; #[cfg(unix)]
let unix_listener = bound.unix_listener;
#[cfg(unix)]
let has_unix_listener = unix_listener.is_some();
#[cfg(not(unix))]
let has_unix_listener = false;
if listeners.is_empty() && !has_unix_listener { if listeners.is_empty() && !has_unix_listener {
error!("No listeners. Exiting."); error!("No listeners. Exiting.");
@@ -918,54 +990,30 @@ async fn run_telemt_core(
// On Unix, caller supplies privilege drop after bind (may require root for port < 1024). // On Unix, caller supplies privilege drop after bind (may require root for port < 1024).
drop_after_bind(); drop_after_bind();
synlimit_control::reconcile_synlimit_rules(&config).await; let synlimit_controller = synlimit_control::spawn_synlimit_controller(runtime_watch_rx);
synlimit_control::spawn_synlimit_controller(config_rx.clone());
runtime_tasks::apply_runtime_log_filter( runtime_tasks::spawn_metrics_if_configured(&config, &startup_tracker, active_runtime.clone())
has_rust_log,
&effective_log_level,
filter_handle,
log_level_rx,
)
.await;
runtime_tasks::spawn_metrics_if_configured(
&config,
&startup_tracker,
stats.clone(),
beobachten.clone(),
shared_state.clone(),
ip_tracker.clone(),
tls_cache.clone(),
config_rx.clone(),
)
.await; .await;
runtime_watch_tx.send_replace(Some(active_runtime.load_full().watch_state()));
active_runtime_tx.send_replace(Some(active_runtime.clone()));
runtime_tasks::mark_runtime_ready(&startup_tracker).await; runtime_tasks::mark_runtime_ready(&startup_tracker).await;
// Spawn signal handlers for SIGUSR1/SIGUSR2 (non-shutdown signals) // Spawn signal handlers for SIGUSR1/SIGUSR2 (non-shutdown signals)
shutdown::spawn_signal_handlers(stats.clone(), process_started_at); shutdown::spawn_signal_handlers(active_runtime.clone(), process_started_at);
listeners::spawn_tcp_accept_loops( listeners::spawn_tcp_accept_loops(listeners, active_runtime.clone());
listeners, #[cfg(unix)]
config_rx.clone(), listeners::spawn_unix_accept_loop(unix_listener, active_runtime.clone());
admission_rx.clone(),
stats.clone(),
upstream_manager.clone(),
replay_checker.clone(),
buffer_pool.clone(),
rng.clone(),
me_pool.clone(),
api_me_pool.clone(),
route_runtime.clone(),
tls_cache.clone(),
ip_tracker.clone(),
beobachten.clone(),
shared_state,
max_connections.clone(),
);
shutdown::wait_for_shutdown(process_started_at, me_pool, stats, quota_state_path).await; shutdown::wait_for_shutdown(
process_started_at,
active_runtime,
quota_state_path,
synlimit_controller,
reload_supervisor,
)
.await;
Ok(()) Ok(())
} }
+449
View File
@@ -0,0 +1,449 @@
use std::collections::VecDeque;
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::{SystemTime, UNIX_EPOCH};
use serde::{Deserialize, Serialize};
use tokio::sync::{Mutex, mpsc};
use crate::config::ProxyConfig;
const RELOAD_HISTORY_CAPACITY: usize = 32;
const RELOAD_COMMAND_CAPACITY: usize = 1;
const MAX_DRAIN_TIMEOUT_SECS: u64 = 3_600;
/// Session handling policy for an in-process runtime reload.
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Deserialize, Serialize)]
#[serde(rename_all = "snake_case")]
pub(crate) enum ReloadMode {
#[default]
Instant,
Drain,
}
/// Failure policy applied during the activation barrier.
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Deserialize, Serialize)]
#[serde(rename_all = "snake_case")]
pub(crate) enum ReloadFailurePolicy {
#[default]
KeepNew,
Rollback,
}
/// Request body accepted by the maestro reload endpoint.
#[derive(Debug, Clone, Default, PartialEq, Eq, Deserialize, Serialize)]
#[serde(deny_unknown_fields)]
pub(crate) struct ReloadRequest {
#[serde(default)]
pub(crate) mode: ReloadMode,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub(crate) timeout_secs: Option<u64>,
#[serde(default)]
pub(crate) failure_policy: ReloadFailurePolicy,
}
impl ReloadRequest {
/// Validates mode-specific request parameters.
pub(crate) fn validate(&self) -> Result<(), &'static str> {
match (self.mode, self.timeout_secs) {
(ReloadMode::Instant, None) => Ok(()),
(ReloadMode::Instant, Some(_)) => Err("timeout_secs is only valid when mode is drain"),
(ReloadMode::Drain, Some(1..=MAX_DRAIN_TIMEOUT_SECS)) => Ok(()),
(ReloadMode::Drain, Some(_)) => Err("timeout_secs must be within 1..=3600"),
(ReloadMode::Drain, None) => Err("timeout_secs is required when mode is drain"),
}
}
/// Parses optional PATCH query parameters into a reload request.
pub(crate) fn from_query(query: Option<&str>) -> Result<Option<Self>, String> {
let Some(query) = query.filter(|query| !query.is_empty()) else {
return Ok(None);
};
let mut mode = None;
let mut timeout_secs = None;
let mut failure_policy = None;
for (key, value) in url::form_urlencoded::parse(query.as_bytes()) {
match key.as_ref() {
"reload" if mode.is_none() => {
mode = Some(match value.as_ref() {
"instant" => ReloadMode::Instant,
"drain" => ReloadMode::Drain,
_ => return Err("reload must be instant or drain".to_string()),
});
}
"timeout_secs" if timeout_secs.is_none() => {
timeout_secs = Some(
value
.parse::<u64>()
.map_err(|_| "timeout_secs must be an integer".to_string())?,
);
}
"failure_policy" if failure_policy.is_none() => {
failure_policy = Some(match value.as_ref() {
"keep_new" => ReloadFailurePolicy::KeepNew,
"rollback" => ReloadFailurePolicy::Rollback,
_ => {
return Err("failure_policy must be keep_new or rollback".to_string());
}
});
}
"reload" | "timeout_secs" | "failure_policy" => {
return Err(format!("duplicate query parameter: {}", key));
}
_ => return Err(format!("unknown query parameter: {}", key)),
}
}
let mode = mode.ok_or_else(|| "reload query parameter is required".to_string())?;
let request = Self {
mode,
timeout_secs,
failure_policy: failure_policy.unwrap_or_default(),
};
request.validate().map_err(str::to_string)?;
Ok(Some(request))
}
}
/// Observable phase of one reload operation.
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize)]
#[serde(rename_all = "snake_case")]
pub(crate) enum ReloadPhase {
Accepted,
Preparing,
Activating,
Draining,
Succeeded,
RolledBack,
Failed,
}
impl ReloadPhase {
fn is_terminal(self) -> bool {
matches!(
self,
ReloadPhase::Succeeded | ReloadPhase::RolledBack | ReloadPhase::Failed
)
}
}
/// Bounded public status for one reload operation.
#[derive(Debug, Clone, Serialize)]
pub(crate) struct ReloadStatus {
pub(crate) reload_id: u64,
pub(crate) target_generation: u64,
pub(crate) config_revision: String,
pub(crate) state: ReloadPhase,
pub(crate) mode: ReloadMode,
pub(crate) failure_policy: ReloadFailurePolicy,
pub(crate) requested_at_epoch_secs: u64,
#[serde(skip_serializing_if = "Option::is_none")]
pub(crate) started_at_epoch_secs: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub(crate) finished_at_epoch_secs: Option<u64>,
#[serde(
rename = "deferred_process_fields",
default,
skip_serializing_if = "Vec::is_empty"
)]
pub(crate) deferred_fields: Vec<String>,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub(crate) warnings: Vec<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub(crate) error: Option<String>,
}
/// Accepted operation metadata returned before asynchronous preparation starts.
#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
pub(crate) struct ReloadAccepted {
pub(crate) reload_id: u64,
pub(crate) target_generation: u64,
pub(crate) config_revision: String,
pub(crate) state: ReloadPhase,
pub(crate) mode: ReloadMode,
pub(crate) failure_policy: ReloadFailurePolicy,
}
pub(crate) struct ReloadCommand {
pub(crate) reload_id: u64,
pub(crate) target_generation: u64,
pub(crate) config: Arc<ProxyConfig>,
pub(crate) config_revision: String,
pub(crate) request: ReloadRequest,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum ReloadSubmitError {
InProgress(u64),
MaestroUnavailable,
}
#[derive(Clone)]
pub(crate) struct ReloadControl {
command_tx: mpsc::Sender<ReloadCommand>,
status_store: Arc<ReloadStatusStore>,
active_generation: Arc<AtomicU64>,
}
pub(crate) struct ReloadCommandReceiver {
command_rx: mpsc::Receiver<ReloadCommand>,
}
struct ReloadStatusState {
next_reload_id: u64,
active_reload_id: Option<u64>,
statuses: VecDeque<ReloadStatus>,
accepting_commands: bool,
}
impl Default for ReloadStatusState {
fn default() -> Self {
Self {
next_reload_id: 0,
active_reload_id: None,
statuses: VecDeque::new(),
accepting_commands: true,
}
}
}
#[derive(Default)]
struct ReloadStatusStore {
state: Mutex<ReloadStatusState>,
}
impl ReloadControl {
/// Creates the process-scoped coordinator channel and status store.
pub(crate) fn channel(initial_generation: u64) -> (Self, ReloadCommandReceiver) {
let (command_tx, command_rx) = mpsc::channel(RELOAD_COMMAND_CAPACITY);
(
Self {
command_tx,
status_store: Arc::new(ReloadStatusStore::default()),
active_generation: Arc::new(AtomicU64::new(initial_generation)),
},
ReloadCommandReceiver { command_rx },
)
}
/// Atomically reserves and enqueues one reload operation.
pub(crate) async fn submit(
&self,
config: Arc<ProxyConfig>,
config_revision: String,
request: ReloadRequest,
) -> Result<ReloadAccepted, ReloadSubmitError> {
let target_generation = self
.active_generation
.load(Ordering::Acquire)
.saturating_add(1);
let status = self
.status_store
.reserve(target_generation, config_revision, request.clone())
.await?;
let command = ReloadCommand {
reload_id: status.reload_id,
target_generation,
config,
config_revision: status.config_revision.clone(),
request,
};
if self.command_tx.try_send(command).is_err() {
self.status_store
.finish(
status.reload_id,
ReloadPhase::Failed,
Some("maestro command channel is closed".to_string()),
)
.await;
return Err(ReloadSubmitError::MaestroUnavailable);
}
Ok(ReloadAccepted {
reload_id: status.reload_id,
target_generation,
config_revision: status.config_revision,
state: ReloadPhase::Accepted,
mode: status.mode,
failure_policy: status.failure_policy,
})
}
/// Returns a retained reload status by identifier.
pub(crate) async fn status(&self, reload_id: u64) -> Option<ReloadStatus> {
self.status_store.get(reload_id).await
}
/// Returns the identifier of the currently active reload.
pub(crate) async fn in_progress(&self) -> Option<u64> {
self.status_store.state.lock().await.active_reload_id
}
/// Rejects new commands while preserving an already accepted operation.
pub(crate) async fn begin_shutdown(&self) {
self.status_store.state.lock().await.accepting_commands = false;
}
/// Records a non-terminal lifecycle phase.
pub(crate) async fn mark_phase(&self, reload_id: u64, phase: ReloadPhase) {
self.status_store.mark_phase(reload_id, phase).await;
}
/// Records process-owned fields deferred until the next process restart.
pub(crate) async fn set_deferred_fields(&self, reload_id: u64, fields: Vec<String>) {
self.status_store
.update(reload_id, |status| status.deferred_fields = fields)
.await;
}
/// Commits the active generation and completes the matching reload.
pub(crate) async fn succeed(&self, reload_id: u64, generation: u64) {
self.status_store
.finish_success(reload_id, generation, &self.active_generation)
.await;
}
/// Marks the matching reload as failed.
pub(crate) async fn fail(&self, reload_id: u64, error: impl Into<String>) {
self.status_store
.finish(reload_id, ReloadPhase::Failed, Some(error.into()))
.await;
}
/// Marks the matching reload as rolled back.
pub(crate) async fn rolled_back(&self, reload_id: u64, error: impl Into<String>) {
self.status_store
.finish(reload_id, ReloadPhase::RolledBack, Some(error.into()))
.await;
}
/// Appends a non-fatal warning to the matching reload status.
pub(crate) async fn add_warning(&self, reload_id: u64, warning: impl Into<String>) {
let warning = warning.into();
self.status_store
.update(reload_id, |status| status.warnings.push(warning))
.await;
}
}
impl ReloadCommandReceiver {
/// Receives the next accepted reload command.
pub(crate) async fn recv(&mut self) -> Option<ReloadCommand> {
self.command_rx.recv().await
}
}
impl ReloadStatusStore {
async fn reserve(
&self,
target_generation: u64,
config_revision: String,
request: ReloadRequest,
) -> Result<ReloadStatus, ReloadSubmitError> {
let mut state = self.state.lock().await;
if !state.accepting_commands {
return Err(ReloadSubmitError::MaestroUnavailable);
}
if let Some(reload_id) = state.active_reload_id {
return Err(ReloadSubmitError::InProgress(reload_id));
}
state.next_reload_id = state.next_reload_id.saturating_add(1).max(1);
let reload_id = state.next_reload_id;
let status = ReloadStatus {
reload_id,
target_generation,
config_revision,
state: ReloadPhase::Accepted,
mode: request.mode,
failure_policy: request.failure_policy,
requested_at_epoch_secs: now_epoch_secs(),
started_at_epoch_secs: None,
finished_at_epoch_secs: None,
deferred_fields: Vec::new(),
warnings: Vec::new(),
error: None,
};
state.active_reload_id = Some(reload_id);
state.statuses.push_back(status.clone());
while state.statuses.len() > RELOAD_HISTORY_CAPACITY {
state.statuses.pop_front();
}
Ok(status)
}
async fn get(&self, reload_id: u64) -> Option<ReloadStatus> {
self.state
.lock()
.await
.statuses
.iter()
.find(|status| status.reload_id == reload_id)
.cloned()
}
async fn mark_phase(&self, reload_id: u64, phase: ReloadPhase) {
self.update(reload_id, |status| {
status.state = phase;
if status.started_at_epoch_secs.is_none() && phase != ReloadPhase::Accepted {
status.started_at_epoch_secs = Some(now_epoch_secs());
}
})
.await;
}
async fn finish(&self, reload_id: u64, phase: ReloadPhase, error: Option<String>) {
debug_assert!(phase.is_terminal());
let mut state = self.state.lock().await;
if let Some(status) = state
.statuses
.iter_mut()
.find(|status| status.reload_id == reload_id)
{
status.state = phase;
status.error = error;
status.finished_at_epoch_secs = Some(now_epoch_secs());
}
if state.active_reload_id == Some(reload_id) {
state.active_reload_id = None;
}
}
async fn finish_success(&self, reload_id: u64, generation: u64, active_generation: &AtomicU64) {
let mut state = self.state.lock().await;
if state.active_reload_id != Some(reload_id) {
return;
}
let Some(status) = state
.statuses
.iter_mut()
.find(|status| status.reload_id == reload_id)
else {
return;
};
status.state = ReloadPhase::Succeeded;
status.error = None;
status.finished_at_epoch_secs = Some(now_epoch_secs());
active_generation.store(generation, Ordering::Release);
state.active_reload_id = None;
}
async fn update(&self, reload_id: u64, update: impl FnOnce(&mut ReloadStatus)) {
let mut state = self.state.lock().await;
if let Some(status) = state
.statuses
.iter_mut()
.find(|status| status.reload_id == reload_id)
{
update(status);
}
}
}
fn now_epoch_secs() -> u64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs()
}
#[cfg(test)]
#[path = "reload_tests.rs"]
mod tests;
+280
View File
@@ -0,0 +1,280 @@
use std::path::PathBuf;
use std::sync::Arc;
use std::time::Duration;
use arc_swap::ArcSwap;
use tokio::sync::watch;
use tokio_util::sync::CancellationToken;
use tracing::{info, warn};
use crate::stats::QuotaStore;
use super::generation::{RuntimeGeneration, RuntimeWatchState};
use super::reload::{
ReloadCommand, ReloadCommandReceiver, ReloadControl, ReloadFailurePolicy, ReloadMode,
ReloadPhase,
};
use super::runtime_build::{PreparedRuntime, deferred_process_fields, prepare_runtime};
use super::runtime_tasks::RuntimeLogFilter;
pub(crate) struct ReloadSupervisor {
active_runtime: Arc<ArcSwap<RuntimeGeneration>>,
control: ReloadControl,
commands: ReloadCommandReceiver,
config_path: PathBuf,
quota_store: Arc<QuotaStore>,
detected_ips_tx: watch::Sender<(Option<std::net::IpAddr>, Option<std::net::IpAddr>)>,
runtime_log_filter: RuntimeLogFilter,
runtime_watch_tx: watch::Sender<Option<RuntimeWatchState>>,
}
/// Process-owned handle that quiesces reloads before shutdown snapshots the runtime.
pub(crate) struct ReloadSupervisorHandle {
control: ReloadControl,
shutdown: CancellationToken,
join: tokio::task::JoinHandle<()>,
}
impl ReloadSupervisorHandle {
/// Stops new submissions and waits for the accepted reload to finish.
pub(crate) async fn quiesce(self) {
self.control.begin_shutdown().await;
self.shutdown.cancel();
if let Err(error) = self.join.await {
warn!(error = %error, "Reload supervisor failed while quiescing");
}
}
}
#[derive(Debug, PartialEq, Eq)]
enum RevisionGateAction {
Proceed,
Warn(String),
Rollback(String),
}
fn revision_gate_action(
accepted_revision: &str,
current_revision: Result<String, String>,
failure_policy: ReloadFailurePolicy,
) -> RevisionGateAction {
let warning = match current_revision {
Ok(current) if current == accepted_revision => return RevisionGateAction::Proceed,
Ok(current) => format!(
"config revision changed during preparation: accepted={} current={}",
accepted_revision, current
),
Err(error) => format!("config revision verification failed: {}", error),
};
match failure_policy {
ReloadFailurePolicy::KeepNew => RevisionGateAction::Warn(warning),
ReloadFailurePolicy::Rollback => RevisionGateAction::Rollback(warning),
}
}
async fn stop_background_and_middle_end(generation: &RuntimeGeneration) -> bool {
generation.stop_background_tasks().await;
let Some(pool) = generation.current_me_pool().await else {
return false;
};
tokio::time::timeout(Duration::from_secs(2), pool.shutdown_send_close_conn_all())
.await
.is_err()
}
async fn cleanup_candidate(generation: &RuntimeGeneration) -> bool {
generation.stop_sessions().await;
stop_background_and_middle_end(generation).await
}
impl ReloadSupervisor {
#[allow(clippy::too_many_arguments)]
/// Starts the process-scoped reload supervisor and returns its shutdown owner.
pub(crate) fn spawn(
active_runtime: Arc<ArcSwap<RuntimeGeneration>>,
control: ReloadControl,
commands: ReloadCommandReceiver,
config_path: PathBuf,
quota_store: Arc<QuotaStore>,
detected_ips_tx: watch::Sender<(Option<std::net::IpAddr>, Option<std::net::IpAddr>)>,
runtime_log_filter: RuntimeLogFilter,
runtime_watch_tx: watch::Sender<Option<RuntimeWatchState>>,
) -> ReloadSupervisorHandle {
let supervisor = Self {
active_runtime,
control,
commands,
config_path,
quota_store,
detected_ips_tx,
runtime_log_filter,
runtime_watch_tx,
};
let control = supervisor.control.clone();
let shutdown = CancellationToken::new();
let join = tokio::spawn(supervisor.run(shutdown.clone()));
ReloadSupervisorHandle {
control,
shutdown,
join,
}
}
async fn run(mut self, shutdown: CancellationToken) {
loop {
tokio::select! {
biased;
_ = shutdown.cancelled() => {
if self.control.in_progress().await.is_some()
&& let Some(command) = self.commands.recv().await
{
self.reload(command).await;
}
break;
}
command = self.commands.recv() => {
let Some(command) = command else {
break;
};
self.reload(command).await;
}
}
}
}
async fn reload(&self, command: ReloadCommand) {
self.control
.mark_phase(command.reload_id, ReloadPhase::Preparing)
.await;
let old_runtime = self.active_runtime.load_full();
let deferred = deferred_process_fields(&old_runtime.config(), &command.config);
self.control
.set_deferred_fields(command.reload_id, deferred)
.await;
let prepared = match prepare_runtime(
command.target_generation,
command.config.as_ref().clone(),
&self.config_path,
self.quota_store.clone(),
self.runtime_log_filter.clone(),
)
.await
{
Ok(prepared) => prepared,
Err(error) => {
self.control.fail(command.reload_id, error).await;
return;
}
};
let revision_action = revision_gate_action(
&command.config_revision,
crate::api::config_store::current_revision_for_maestro(&self.config_path).await,
command.request.failure_policy,
);
self.activate_prepared(command, old_runtime, prepared, revision_action, |entries| {
crate::network::dns_overrides::install_entries(entries)
.map_err(|error| error.to_string())
})
.await;
}
async fn activate_prepared<InstallDns>(
&self,
command: ReloadCommand,
old_runtime: Arc<RuntimeGeneration>,
prepared: PreparedRuntime,
revision_action: RevisionGateAction,
install_dns: InstallDns,
) where
InstallDns: FnOnce(&[String]) -> Result<(), String>,
{
match revision_action {
RevisionGateAction::Proceed => {}
RevisionGateAction::Warn(warning) => {
self.control.add_warning(command.reload_id, warning).await;
}
RevisionGateAction::Rollback(warning) => {
let _ = cleanup_candidate(&prepared.generation).await;
self.runtime_log_filter
.apply_reload(&old_runtime.config().general.log_level);
self.control.rolled_back(command.reload_id, warning).await;
return;
}
}
self.control
.mark_phase(command.reload_id, ReloadPhase::Activating)
.await;
let new_runtime = prepared.generation;
old_runtime.stop_accepting_sessions();
if let Err(error) = install_dns(&new_runtime.config().network.dns_overrides) {
let message = format!("runtime DNS activation failed: {}", error);
if command.request.failure_policy == ReloadFailurePolicy::Rollback {
old_runtime.resume_accepting_sessions();
let _ = cleanup_candidate(&new_runtime).await;
self.runtime_log_filter
.apply_reload(&old_runtime.config().general.log_level);
self.control.rolled_back(command.reload_id, message).await;
return;
}
self.control.add_warning(command.reload_id, message).await;
}
let replaced = self.active_runtime.swap(new_runtime.clone());
self.detected_ips_tx.send_replace(prepared.detected_ips);
self.runtime_log_filter
.apply_reload(&new_runtime.config().general.log_level);
self.runtime_watch_tx
.send_replace(Some(new_runtime.watch_state()));
info!(
reload_id = command.reload_id,
old_generation = replaced.id,
new_generation = new_runtime.id,
config_revision = %command.config_revision,
"Runtime generation activated"
);
match command.request.mode {
ReloadMode::Instant => {
replaced.stop_sessions().await;
}
ReloadMode::Drain => {
self.control
.mark_phase(command.reload_id, ReloadPhase::Draining)
.await;
let timeout = Duration::from_secs(
command
.request
.timeout_secs
.expect("validated drain request must carry timeout_secs"),
);
if !replaced.drain_sessions(timeout).await {
let warning = format!(
"generation {} exceeded drain timeout; remaining sessions were cancelled",
replaced.id
);
warn!(reload_id = command.reload_id, warning = %warning);
self.control.add_warning(command.reload_id, warning).await;
}
}
}
if stop_background_and_middle_end(&replaced).await {
let warning = format!(
"generation {} Middle-End close broadcast timed out",
replaced.id
);
warn!(reload_id = command.reload_id, warning = %warning);
self.control.add_warning(command.reload_id, warning).await;
}
self.control
.succeed(command.reload_id, new_runtime.id)
.await;
}
}
#[cfg(test)]
#[path = "reload_supervisor_tests.rs"]
mod tests;
+320
View File
@@ -0,0 +1,320 @@
use super::*;
use crate::config::ProxyConfig;
use crate::maestro::generation::test_runtime_generation;
use crate::maestro::reload::{ReloadRequest, ReloadSubmitError};
use crate::stats::QuotaStore;
use tokio::sync::Notify;
use tracing_subscriber::{EnvFilter, Registry};
struct ReloadFixture {
supervisor: Arc<ReloadSupervisor>,
control: ReloadControl,
command: ReloadCommand,
old_runtime: Arc<RuntimeGeneration>,
new_runtime: Arc<RuntimeGeneration>,
runtime_watch_rx: watch::Receiver<Option<RuntimeWatchState>>,
}
fn runtime_log_filter() -> RuntimeLogFilter {
let (_layer, handle) =
tracing_subscriber::reload::Layer::<EnvFilter, Registry>::new(EnvFilter::new("info"));
RuntimeLogFilter::new(handle)
}
async fn fixture(request: ReloadRequest) -> ReloadFixture {
let old_runtime = test_runtime_generation(1, ProxyConfig::default());
let new_config = Arc::new(ProxyConfig::default());
let new_runtime = test_runtime_generation(2, new_config.as_ref().clone());
let active_runtime = Arc::new(ArcSwap::from(old_runtime.clone()));
let (control, commands) = ReloadControl::channel(old_runtime.id);
let accepted = control
.submit(new_config.clone(), "revision".to_string(), request.clone())
.await
.unwrap();
let (detected_ips_tx, _detected_ips_rx) = watch::channel((None, None));
let (runtime_watch_tx, runtime_watch_rx) = watch::channel(Some(old_runtime.watch_state()));
let supervisor = Arc::new(ReloadSupervisor {
active_runtime,
control: control.clone(),
commands,
config_path: PathBuf::new(),
quota_store: Arc::new(QuotaStore::default()),
detected_ips_tx,
runtime_log_filter: runtime_log_filter(),
runtime_watch_tx,
});
let command = ReloadCommand {
reload_id: accepted.reload_id,
target_generation: accepted.target_generation,
config: new_config,
config_revision: accepted.config_revision,
request,
};
ReloadFixture {
supervisor,
control,
command,
old_runtime,
new_runtime,
runtime_watch_rx,
}
}
struct DropSignal(Arc<Notify>);
impl Drop for DropSignal {
fn drop(&mut self) {
self.0.notify_one();
}
}
#[test]
fn revision_gate_proceeds_only_on_verified_match() {
assert_eq!(
revision_gate_action(
"accepted",
Ok("accepted".to_string()),
ReloadFailurePolicy::Rollback,
),
RevisionGateAction::Proceed
);
}
#[test]
fn revision_gate_applies_failure_policy_to_mismatch_and_read_error() {
for result in [Ok("changed".to_string()), Err("read failed".to_string())] {
assert!(matches!(
revision_gate_action("accepted", result.clone(), ReloadFailurePolicy::KeepNew,),
RevisionGateAction::Warn(_)
));
assert!(matches!(
revision_gate_action("accepted", result, ReloadFailurePolicy::Rollback),
RevisionGateAction::Rollback(_)
));
}
}
#[tokio::test]
async fn revision_rollback_keeps_old_generation_and_cleans_candidate() {
let fixture = fixture(ReloadRequest {
failure_policy: ReloadFailurePolicy::Rollback,
..ReloadRequest::default()
})
.await;
let candidate_dropped = Arc::new(Notify::new());
let candidate_drop = candidate_dropped.clone();
assert!(fixture.new_runtime.spawn_session(async move {
let _drop_signal = DropSignal(candidate_drop);
std::future::pending::<()>().await;
}));
tokio::task::yield_now().await;
fixture
.supervisor
.activate_prepared(
fixture.command,
fixture.old_runtime.clone(),
PreparedRuntime {
generation: fixture.new_runtime,
detected_ips: (None, None),
},
RevisionGateAction::Rollback("revision changed".to_string()),
|_| -> Result<(), String> { panic!("DNS activation must not run on rollback") },
)
.await;
tokio::time::timeout(Duration::from_secs(1), candidate_dropped.notified())
.await
.unwrap();
assert_eq!(fixture.supervisor.active_runtime.load().id, 1);
assert_eq!(
fixture
.runtime_watch_rx
.borrow()
.as_ref()
.unwrap()
.generation_id,
1
);
assert!(fixture.old_runtime.spawn_session(async {}));
let status = fixture.control.status(1).await.unwrap();
assert_eq!(status.state, ReloadPhase::RolledBack);
fixture.old_runtime.stop_sessions().await;
}
#[tokio::test]
async fn dns_failure_policy_controls_rollback_or_keep_new() {
for policy in [ReloadFailurePolicy::Rollback, ReloadFailurePolicy::KeepNew] {
let fixture = fixture(ReloadRequest {
failure_policy: policy,
..ReloadRequest::default()
})
.await;
fixture
.supervisor
.activate_prepared(
fixture.command,
fixture.old_runtime.clone(),
PreparedRuntime {
generation: fixture.new_runtime.clone(),
detected_ips: (None, None),
},
RevisionGateAction::Proceed,
|_| Err("invalid DNS entry".to_string()),
)
.await;
let status = fixture.control.status(1).await.unwrap();
match policy {
ReloadFailurePolicy::Rollback => {
assert_eq!(fixture.supervisor.active_runtime.load().id, 1);
assert_eq!(status.state, ReloadPhase::RolledBack);
assert!(fixture.old_runtime.spawn_session(async {}));
fixture.old_runtime.stop_sessions().await;
}
ReloadFailurePolicy::KeepNew => {
assert_eq!(fixture.supervisor.active_runtime.load().id, 2);
assert_eq!(status.state, ReloadPhase::Succeeded);
assert_eq!(status.warnings.len(), 1);
assert!(!fixture.old_runtime.spawn_session(async {}));
fixture.new_runtime.stop_sessions().await;
}
}
}
}
#[tokio::test]
async fn drain_publishes_new_generation_before_old_sessions_finish() {
let mut fixture = fixture(ReloadRequest {
mode: ReloadMode::Drain,
timeout_secs: Some(30),
..ReloadRequest::default()
})
.await;
let old_started = Arc::new(Notify::new());
let old_release = Arc::new(Notify::new());
let started = old_started.clone();
let release = old_release.clone();
assert!(fixture.old_runtime.spawn_session(async move {
started.notify_one();
release.notified().await;
}));
old_started.notified().await;
let supervisor = fixture.supervisor.clone();
let old_runtime = fixture.old_runtime.clone();
let new_runtime = fixture.new_runtime.clone();
let activation = tokio::spawn(async move {
supervisor
.activate_prepared(
fixture.command,
old_runtime,
PreparedRuntime {
generation: new_runtime,
detected_ips: (None, None),
},
RevisionGateAction::Proceed,
|_| Ok(()),
)
.await;
});
fixture.runtime_watch_rx.changed().await.unwrap();
assert_eq!(
fixture
.runtime_watch_rx
.borrow()
.as_ref()
.unwrap()
.generation_id,
2
);
assert!(!activation.is_finished());
assert!(!fixture.old_runtime.spawn_session(async {}));
old_release.notify_one();
activation.await.unwrap();
assert_eq!(
fixture.control.status(1).await.unwrap().state,
ReloadPhase::Succeeded
);
fixture.new_runtime.stop_sessions().await;
}
#[tokio::test(start_paused = true)]
async fn drain_timeout_cancels_old_sessions_and_records_one_warning() {
let mut fixture = fixture(ReloadRequest {
mode: ReloadMode::Drain,
timeout_secs: Some(1),
..ReloadRequest::default()
})
.await;
let dropped = Arc::new(Notify::new());
let drop_signal = dropped.clone();
assert!(fixture.old_runtime.spawn_session(async move {
let _drop_signal = DropSignal(drop_signal);
std::future::pending::<()>().await;
}));
tokio::task::yield_now().await;
let supervisor = fixture.supervisor.clone();
let old_runtime = fixture.old_runtime.clone();
let new_runtime = fixture.new_runtime.clone();
let activation = tokio::spawn(async move {
supervisor
.activate_prepared(
fixture.command,
old_runtime,
PreparedRuntime {
generation: new_runtime,
detected_ips: (None, None),
},
RevisionGateAction::Proceed,
|_| Ok(()),
)
.await;
});
fixture.runtime_watch_rx.changed().await.unwrap();
tokio::task::yield_now().await;
tokio::time::advance(Duration::from_secs(1)).await;
activation.await.unwrap();
dropped.notified().await;
let status = fixture.control.status(1).await.unwrap();
assert_eq!(status.state, ReloadPhase::Succeeded);
assert_eq!(status.warnings.len(), 1);
assert!(status.warnings[0].contains("exceeded drain timeout"));
fixture.new_runtime.stop_sessions().await;
}
#[tokio::test]
async fn quiesce_joins_idle_supervisor_and_rejects_later_submissions() {
let runtime = test_runtime_generation(1, ProxyConfig::default());
let active_runtime = Arc::new(ArcSwap::from(runtime.clone()));
let (control, commands) = ReloadControl::channel(runtime.id);
let (detected_ips_tx, _detected_ips_rx) = watch::channel((None, None));
let (runtime_watch_tx, _runtime_watch_rx) = watch::channel(Some(runtime.watch_state()));
let handle = ReloadSupervisor::spawn(
active_runtime,
control.clone(),
commands,
PathBuf::new(),
Arc::new(QuotaStore::default()),
detected_ips_tx,
runtime_log_filter(),
runtime_watch_tx,
);
tokio::time::timeout(Duration::from_secs(1), handle.quiesce())
.await
.unwrap();
let result = control
.submit(
Arc::new(ProxyConfig::default()),
"revision".to_string(),
ReloadRequest::default(),
)
.await;
assert_eq!(result, Err(ReloadSubmitError::MaestroUnavailable));
runtime.stop_sessions().await;
}
+256
View File
@@ -0,0 +1,256 @@
use super::*;
#[test]
fn request_defaults_to_instant_keep_new() {
let request: ReloadRequest = serde_json::from_str("{}").unwrap();
assert_eq!(request, ReloadRequest::default());
assert_eq!(request.validate(), Ok(()));
}
#[test]
fn drain_requires_bounded_timeout() {
let missing = ReloadRequest {
mode: ReloadMode::Drain,
..ReloadRequest::default()
};
assert!(missing.validate().is_err());
let valid = ReloadRequest {
mode: ReloadMode::Drain,
timeout_secs: Some(30),
..ReloadRequest::default()
};
assert_eq!(valid.validate(), Ok(()));
}
#[test]
fn patch_query_parses_reload_policy() {
let request =
ReloadRequest::from_query(Some("reload=drain&timeout_secs=30&failure_policy=rollback"))
.unwrap()
.unwrap();
assert_eq!(request.mode, ReloadMode::Drain);
assert_eq!(request.timeout_secs, Some(30));
assert_eq!(request.failure_policy, ReloadFailurePolicy::Rollback);
assert!(ReloadRequest::from_query(Some("timeout_secs=30")).is_err());
}
#[test]
fn status_uses_documented_deferred_process_fields_key() {
let status = ReloadStatus {
reload_id: 1,
target_generation: 2,
config_revision: "revision".to_string(),
state: ReloadPhase::Succeeded,
mode: ReloadMode::Instant,
failure_policy: ReloadFailurePolicy::KeepNew,
requested_at_epoch_secs: 10,
started_at_epoch_secs: Some(11),
finished_at_epoch_secs: Some(12),
deferred_fields: vec!["server.listeners".to_string()],
warnings: Vec::new(),
error: None,
};
let value = serde_json::to_value(status).unwrap();
assert_eq!(
value["deferred_process_fields"],
serde_json::json!(["server.listeners"])
);
assert!(value.get("deferred_fields").is_none());
}
#[tokio::test]
async fn coordinator_rejects_concurrent_reload_and_releases_terminal_slot() {
let (control, mut receiver) = ReloadControl::channel(1);
let first = control
.submit(
Arc::new(ProxyConfig::default()),
"rev-1".to_string(),
ReloadRequest::default(),
)
.await
.unwrap();
let _command = receiver.recv().await.unwrap();
let second = control
.submit(
Arc::new(ProxyConfig::default()),
"rev-2".to_string(),
ReloadRequest::default(),
)
.await;
assert_eq!(second, Err(ReloadSubmitError::InProgress(first.reload_id)));
control
.succeed(first.reload_id, first.target_generation)
.await;
let third = control
.submit(
Arc::new(ProxyConfig::default()),
"rev-3".to_string(),
ReloadRequest::default(),
)
.await
.unwrap();
assert_eq!(third.reload_id, first.reload_id + 1);
}
#[tokio::test]
async fn terminal_outcomes_release_slot_and_only_success_advances_generation() {
let (control, mut receiver) = ReloadControl::channel(7);
let failed = control
.submit(
Arc::new(ProxyConfig::default()),
"rev-failed".to_string(),
ReloadRequest::default(),
)
.await
.unwrap();
let _command = receiver.recv().await.unwrap();
control
.mark_phase(failed.reload_id, ReloadPhase::Preparing)
.await;
control.fail(failed.reload_id, "prepare failed").await;
let failed_status = control.status(failed.reload_id).await.unwrap();
assert_eq!(failed_status.state, ReloadPhase::Failed);
assert_eq!(failed_status.error.as_deref(), Some("prepare failed"));
assert!(failed_status.started_at_epoch_secs.is_some());
assert!(failed_status.finished_at_epoch_secs.is_some());
let rolled_back = control
.submit(
Arc::new(ProxyConfig::default()),
"rev-rollback".to_string(),
ReloadRequest::default(),
)
.await
.unwrap();
let _command = receiver.recv().await.unwrap();
assert_eq!(rolled_back.target_generation, 8);
control
.rolled_back(rolled_back.reload_id, "revision changed")
.await;
let succeeded = control
.submit(
Arc::new(ProxyConfig::default()),
"rev-success".to_string(),
ReloadRequest::default(),
)
.await
.unwrap();
let _command = receiver.recv().await.unwrap();
assert_eq!(succeeded.target_generation, 8);
control
.succeed(succeeded.reload_id, succeeded.target_generation)
.await;
let next = control
.submit(
Arc::new(ProxyConfig::default()),
"rev-next".to_string(),
ReloadRequest::default(),
)
.await
.unwrap();
assert_eq!(next.target_generation, 9);
}
#[tokio::test]
async fn stale_success_cannot_advance_generation_or_release_active_reload() {
let (control, mut receiver) = ReloadControl::channel(3);
let active = control
.submit(
Arc::new(ProxyConfig::default()),
"rev-active".to_string(),
ReloadRequest::default(),
)
.await
.unwrap();
let _command = receiver.recv().await.unwrap();
control.succeed(active.reload_id + 100, 99).await;
assert_eq!(control.in_progress().await, Some(active.reload_id));
control.fail(active.reload_id, "expected failure").await;
let next = control
.submit(
Arc::new(ProxyConfig::default()),
"rev-next".to_string(),
ReloadRequest::default(),
)
.await
.unwrap();
assert_eq!(next.target_generation, 4);
}
#[tokio::test]
async fn status_history_retains_only_the_latest_entries() {
let (control, mut receiver) = ReloadControl::channel(1);
let mut reload_ids = Vec::new();
for index in 0..=RELOAD_HISTORY_CAPACITY {
let accepted = control
.submit(
Arc::new(ProxyConfig::default()),
format!("rev-{index}"),
ReloadRequest::default(),
)
.await
.unwrap();
let _command = receiver.recv().await.unwrap();
reload_ids.push(accepted.reload_id);
control.fail(accepted.reload_id, "expected failure").await;
}
assert!(control.status(reload_ids[0]).await.is_none());
assert!(control.status(reload_ids[1]).await.is_some());
assert!(control.status(*reload_ids.last().unwrap()).await.is_some());
}
#[tokio::test]
async fn closed_command_channel_marks_reload_failed_and_releases_slot() {
let (control, receiver) = ReloadControl::channel(1);
drop(receiver);
let result = control
.submit(
Arc::new(ProxyConfig::default()),
"rev-closed".to_string(),
ReloadRequest::default(),
)
.await;
assert_eq!(result, Err(ReloadSubmitError::MaestroUnavailable));
assert_eq!(control.in_progress().await, None);
let status = control.status(1).await.unwrap();
assert_eq!(status.state, ReloadPhase::Failed);
assert_eq!(
status.error.as_deref(),
Some("maestro command channel is closed")
);
}
#[tokio::test]
async fn shutdown_gate_rejects_new_commands_without_disturbing_active_status() {
let (control, mut receiver) = ReloadControl::channel(4);
let active = control
.submit(
Arc::new(ProxyConfig::default()),
"rev-active".to_string(),
ReloadRequest::default(),
)
.await
.unwrap();
let _command = receiver.recv().await.unwrap();
control.begin_shutdown().await;
let rejected = control
.submit(
Arc::new(ProxyConfig::default()),
"rev-rejected".to_string(),
ReloadRequest::default(),
)
.await;
assert_eq!(rejected, Err(ReloadSubmitError::MaestroUnavailable));
assert_eq!(control.in_progress().await, Some(active.reload_id));
control.fail(active.reload_id, "shutdown test").await;
}
+383
View File
@@ -0,0 +1,383 @@
use std::net::IpAddr;
use std::path::Path;
use std::sync::Arc;
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use tokio::sync::{RwLock, Semaphore, watch};
use crate::config::ProxyConfig;
use crate::crypto::SecureRandom;
use crate::ip_tracker::UserIpTracker;
use crate::network::probe::{decide_network_capabilities, run_probe};
use crate::proxy::direct_buffer_budget::{
DirectBufferBudget, resolve_direct_buffer_hard_limit, run_direct_buffer_budget_controller,
};
use crate::proxy::route_mode::{RelayRouteMode, RouteRuntimeController};
use crate::proxy::shared_state::ProxySharedState;
use crate::startup::StartupTracker;
use crate::stats::beobachten::BeobachtenStore;
use crate::stats::telemetry::TelemetryPolicy;
use crate::stats::{QuotaStore, ReplayChecker, Stats};
use crate::stream::BufferPool;
use crate::transport::UpstreamManager;
use crate::transport::middle_proxy::MePool;
use super::admission;
use super::generation::{RuntimeGeneration, RuntimeTaskScope};
use super::runtime_tasks::RuntimeLogFilter;
use super::{me_startup, runtime_tasks, tls_bootstrap};
pub(crate) struct PreparedRuntime {
pub(crate) generation: Arc<RuntimeGeneration>,
pub(crate) detected_ips: (Option<IpAddr>, Option<IpAddr>),
}
pub(crate) async fn prepare_runtime(
generation_id: u64,
config: ProxyConfig,
config_path: &Path,
quota_store: Arc<QuotaStore>,
runtime_log_filter: RuntimeLogFilter,
) -> Result<PreparedRuntime, String> {
let started_at_epoch_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
let startup_tracker = Arc::new(StartupTracker::new(started_at_epoch_secs));
let task_scope = RuntimeTaskScope::new();
let stats = Arc::new(Stats::with_quota_store(quota_store));
stats.apply_telemetry_policy(TelemetryPolicy::from_config(&config.general.telemetry));
let upstream_manager = Arc::new(
UpstreamManager::new(
config.upstreams.clone(),
config.general.upstream_connect_retry_attempts,
config.general.upstream_connect_retry_backoff_ms,
config.general.upstream_connect_budget_ms,
config.general.tg_connect,
config.general.upstream_unhealthy_fail_threshold,
config.general.upstream_connect_failfast_hard_errors,
stats.clone(),
)
.with_dns_overrides(&config.network.dns_overrides)
.map_err(|error| format!("DNS override preparation failed: {}", error))?,
);
let ip_tracker = Arc::new(UserIpTracker::new());
ip_tracker
.load_limits(
config.access.user_max_unique_ips_global_each,
&config.access.user_max_unique_ips,
)
.await;
ip_tracker
.set_limit_policy(
config.access.user_max_unique_ips_mode,
config.access.user_max_unique_ips_window_secs,
)
.await;
let hard_limit =
resolve_direct_buffer_hard_limit(config.general.direct_relay_buffer_budget_max_bytes).await;
let direct_buffer_budget = DirectBufferBudget::new(hard_limit);
let proxy_shared =
ProxySharedState::new_with_direct_buffer_budget(direct_buffer_budget.clone());
proxy_shared.apply_user_enabled_config(&config.access.user_enabled);
proxy_shared.traffic_limiter.apply_policy(
config.access.user_rate_limits.clone(),
config.access.cidr_rate_limits.clone(),
);
let probe = run_probe(
&config.network,
&config.upstreams,
config.general.middle_proxy_nat_probe,
config.general.stun_nat_probe_concurrency,
)
.await
.map_err(|error| format!("network probe failed: {}", error))?;
let decision =
decide_network_capabilities(&config.network, &probe, config.general.middle_proxy_nat_ip);
let prefer_ipv6 = decision.prefer_ipv6();
let mut tls_domains = Vec::with_capacity(1 + config.censorship.tls_domains.len());
tls_domains.push(config.censorship.tls_domain.clone());
for domain in &config.censorship.tls_domains {
if !tls_domains.contains(domain) {
tls_domains.push(domain.clone());
}
}
let tls_cache = tls_bootstrap::bootstrap_tls_front(
&config,
&tls_domains,
upstream_manager.clone(),
&startup_tracker,
task_scope.clone(),
tls_bootstrap::TlsBootstrapPolicy::RequireReady,
)
.await
.map_err(|error| error.to_string())?;
let beobachten = Arc::new(BeobachtenStore::new());
let rng = Arc::new(SecureRandom::new());
let route_mode = if !config.general.use_middle_proxy || config.general.me2dc_fallback {
RelayRouteMode::Direct
} else {
RelayRouteMode::Middle
};
let route_runtime = Arc::new(RouteRuntimeController::new(route_mode));
let me_pool_runtime = Arc::new(RwLock::new(None::<Arc<MePool>>));
let (me_ready_tx, me_ready_rx) = watch::channel(0_u64);
let direct_first_startup = config.general.use_middle_proxy && config.general.me2dc_fallback;
let me_pool = if direct_first_startup {
None
} else {
me_startup::initialize_me_pool(
config.general.use_middle_proxy,
&config,
&decision,
&probe,
&startup_tracker,
upstream_manager.clone(),
rng.clone(),
stats.clone(),
me_pool_runtime.clone(),
me_ready_tx.clone(),
task_scope.clone(),
)
.await
};
if strict_middle_proxy_unavailable(
config.general.use_middle_proxy,
direct_first_startup,
me_pool.is_some(),
) {
task_scope.stop().await;
return Err(
"Middle-End pool is required but did not become ready during reload preparation"
.to_string(),
);
}
let config = Arc::new(config);
let replay_checker = Arc::new(ReplayChecker::new(
config.access.replay_check_len,
Duration::from_secs(config.access.replay_window_secs),
));
let buffer_pool = Arc::new(BufferPool::with_config(64 * 1024, 4096));
let max_connections_limit = if config.server.max_connections == 0 {
Semaphore::MAX_PERMITS
} else {
config.server.max_connections as usize
};
let max_connections = Arc::new(Semaphore::new(max_connections_limit));
let watches = runtime_tasks::spawn_runtime_tasks(
&config,
config_path,
&probe,
prefer_ipv6,
decision.ipv4_dc,
decision.ipv6_dc,
&startup_tracker,
stats.clone(),
upstream_manager.clone(),
replay_checker.clone(),
me_pool.clone(),
rng.clone(),
ip_tracker.clone(),
beobachten.clone(),
me_pool.clone(),
proxy_shared.clone(),
me_ready_tx.clone(),
task_scope.clone(),
)
.await;
let config_rx = watches.config_rx;
runtime_log_filter.spawn_watcher(watches.log_level_rx, task_scope.clone());
let initial_admission_open = !config.general.use_middle_proxy || me_pool.is_some();
let (admission_tx, admission_rx) = watch::channel(initial_admission_open);
admission::configure_admission_gate(
&config,
me_pool.clone(),
me_pool_runtime.clone(),
route_runtime.clone(),
&admission_tx,
config_rx.clone(),
me_ready_rx,
task_scope.clone(),
)
.await;
if direct_first_startup {
let config_bg = config.clone();
let decision_bg = decision.clone();
let probe_bg = probe.clone();
let startup_tracker_bg = startup_tracker.clone();
let upstream_manager_bg = upstream_manager.clone();
let rng_bg = rng.clone();
let stats_bg = stats.clone();
let me_pool_runtime_bg = me_pool_runtime.clone();
let me_ready_tx_bg = me_ready_tx.clone();
let config_rx_bg = config_rx.clone();
let task_scope_bg = task_scope.clone();
let retry_limit = config.general.me_init_retry_attempts;
task_scope.spawn(async move {
let mut attempt = 0_u32;
loop {
attempt = attempt.saturating_add(1);
let pool = me_startup::initialize_me_pool(
true,
config_bg.as_ref(),
&decision_bg,
&probe_bg,
&startup_tracker_bg,
upstream_manager_bg.clone(),
rng_bg.clone(),
stats_bg.clone(),
me_pool_runtime_bg.clone(),
me_ready_tx_bg.clone(),
task_scope_bg.clone(),
)
.await;
if let Some(pool) = pool {
runtime_tasks::spawn_middle_proxy_runtime_tasks(
config_bg.as_ref(),
config_rx_bg,
pool,
rng_bg,
me_ready_tx_bg,
task_scope_bg,
);
break;
}
if retry_limit > 0 && attempt >= retry_limit {
break;
}
tokio::time::sleep(Duration::from_secs(2)).await;
}
});
}
let conntrack_scope = task_scope.clone();
task_scope.spawn(crate::conntrack_control::run_conntrack_controller(
config_rx.clone(),
stats.clone(),
proxy_shared.clone(),
conntrack_scope.cancellation_token(),
));
task_scope.spawn(run_direct_buffer_budget_controller(
direct_buffer_budget,
buffer_pool.clone(),
stats.clone(),
proxy_shared.clone(),
config.server.max_connections,
));
let generation = RuntimeGeneration::new(
generation_id,
config_rx,
admission_rx,
stats,
upstream_manager,
replay_checker,
buffer_pool,
rng,
me_pool,
me_pool_runtime,
route_runtime,
tls_cache,
ip_tracker,
beobachten,
proxy_shared,
max_connections,
task_scope,
);
drop(admission_tx);
Ok(PreparedRuntime {
generation,
detected_ips: (
probe.detected_ipv4.map(IpAddr::V4),
probe.detected_ipv6.map(IpAddr::V6),
),
})
}
fn strict_middle_proxy_unavailable(
use_middle_proxy: bool,
direct_first_startup: bool,
pool_available: bool,
) -> bool {
use_middle_proxy && !direct_first_startup && !pool_available
}
pub(crate) fn deferred_process_fields(old: &ProxyConfig, new: &ProxyConfig) -> Vec<String> {
let mut fields = Vec::new();
if old.server.port != new.server.port
|| old.server.proxy_protocol != new.server.proxy_protocol
|| old.server.listen_backlog != new.server.listen_backlog
|| serde_json::to_value(&old.server.listeners).ok()
!= serde_json::to_value(&new.server.listeners).ok()
{
fields.push("server.listeners".to_string());
}
if old.server.listen_unix_sock != new.server.listen_unix_sock
|| old.server.listen_unix_sock_perm != new.server.listen_unix_sock_perm
{
fields.push("server.listen_unix_sock".to_string());
}
if old.server.api.listen != new.server.api.listen
|| old.server.api.enabled != new.server.api.enabled
{
fields.push("server.api.listen".to_string());
}
if old.server.metrics_listen != new.server.metrics_listen
|| old.server.metrics_port != new.server.metrics_port
{
fields.push("server.metrics_listen".to_string());
}
if old.general.quota_state_path != new.general.quota_state_path {
fields.push("general.quota_state_path".to_string());
}
if old.general.disable_colors != new.general.disable_colors {
fields.push("general.disable_colors".to_string());
}
if old.general.data_path != new.general.data_path {
fields.push("general.data_path".to_string());
}
if serde_json::to_value(&old.logging).ok() != serde_json::to_value(&new.logging).ok() {
fields.push("logging".to_string());
}
fields
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn process_socket_and_logging_changes_are_deferred() {
let old = ProxyConfig::default();
let mut new = old.clone();
new.server.listen_backlog = new.server.listen_backlog.saturating_add(1);
new.general.disable_colors = !new.general.disable_colors;
let fields = deferred_process_fields(&old, &new);
assert!(fields.contains(&"server.listeners".to_string()));
assert!(fields.contains(&"general.disable_colors".to_string()));
}
#[test]
fn runtime_only_change_does_not_require_process_rebind() {
let old = ProxyConfig::default();
let mut new = old.clone();
new.censorship.tls_domain = "reload.example".to_string();
assert!(deferred_process_fields(&old, &new).is_empty());
}
#[test]
fn strict_middle_proxy_requires_a_prepared_pool() {
assert!(strict_middle_proxy_unavailable(true, false, false));
assert!(!strict_middle_proxy_unavailable(true, false, true));
assert!(!strict_middle_proxy_unavailable(true, true, false));
assert!(!strict_middle_proxy_unavailable(false, false, false));
}
}
+84 -78
View File
@@ -2,9 +2,11 @@ use std::net::IpAddr;
use std::path::Path; use std::path::Path;
use std::sync::Arc; use std::sync::Arc;
use arc_swap::ArcSwap;
use tokio::sync::{mpsc, watch}; use tokio::sync::{mpsc, watch};
use tracing::{debug, info, warn}; use tracing::{debug, info, warn};
use tracing_subscriber::EnvFilter; use tracing_subscriber::EnvFilter;
use tracing_subscriber::Registry;
use tracing_subscriber::reload; use tracing_subscriber::reload;
use crate::config::hot_reload::spawn_config_watcher; use crate::config::hot_reload::spawn_config_watcher;
@@ -21,10 +23,11 @@ use crate::startup::{
use crate::stats::beobachten::BeobachtenStore; use crate::stats::beobachten::BeobachtenStore;
use crate::stats::telemetry::TelemetryPolicy; use crate::stats::telemetry::TelemetryPolicy;
use crate::stats::{ReplayChecker, Stats}; use crate::stats::{ReplayChecker, Stats};
use crate::tls_front::TlsFrontCache;
use crate::transport::UpstreamManager; use crate::transport::UpstreamManager;
use crate::transport::middle_proxy::{MePool, MeReinitTrigger}; use crate::transport::middle_proxy::{MePool, MeReinitTrigger};
use super::generation::RuntimeGeneration;
use super::generation::RuntimeTaskScope;
use super::helpers::write_beobachten_snapshot; use super::helpers::write_beobachten_snapshot;
pub(crate) struct RuntimeWatches { pub(crate) struct RuntimeWatches {
@@ -34,6 +37,56 @@ pub(crate) struct RuntimeWatches {
pub(crate) detected_ip_v6: Option<IpAddr>, pub(crate) detected_ip_v6: Option<IpAddr>,
} }
#[derive(Clone)]
pub(crate) struct RuntimeLogFilter {
handle: reload::Handle<EnvFilter, Registry>,
}
impl RuntimeLogFilter {
pub(crate) fn new(handle: reload::Handle<EnvFilter, Registry>) -> Self {
Self { handle }
}
pub(crate) fn start(
&self,
has_rust_log: bool,
effective_log_level: &LogLevel,
log_level_rx: watch::Receiver<LogLevel>,
task_scope: RuntimeTaskScope,
) {
self.apply(effective_log_level, has_rust_log);
self.spawn_watcher(log_level_rx, task_scope);
}
pub(crate) fn apply_reload(&self, level: &LogLevel) {
self.apply(level, false);
}
pub(crate) fn spawn_watcher(
&self,
mut log_level_rx: watch::Receiver<LogLevel>,
task_scope: RuntimeTaskScope,
) {
let filter = self.clone();
task_scope.spawn(async move {
loop {
if log_level_rx.changed().await.is_err() {
break;
}
let level = log_level_rx.borrow_and_update().clone();
filter.apply_reload(&level);
}
});
}
fn apply(&self, level: &LogLevel, has_rust_log: bool) {
let runtime_filter = EnvFilter::new(log_filter_spec(has_rust_log, level));
if let Err(error) = self.handle.reload(runtime_filter) {
tracing::error!(error = %error, "Failed to update runtime log filter");
}
}
}
#[allow(clippy::too_many_arguments)] #[allow(clippy::too_many_arguments)]
pub(crate) async fn spawn_runtime_tasks( pub(crate) async fn spawn_runtime_tasks(
config: &Arc<ProxyConfig>, config: &Arc<ProxyConfig>,
@@ -50,14 +103,14 @@ pub(crate) async fn spawn_runtime_tasks(
rng: Arc<SecureRandom>, rng: Arc<SecureRandom>,
ip_tracker: Arc<UserIpTracker>, ip_tracker: Arc<UserIpTracker>,
beobachten: Arc<BeobachtenStore>, beobachten: Arc<BeobachtenStore>,
api_config_tx: watch::Sender<Arc<ProxyConfig>>,
me_pool_for_policy: Option<Arc<MePool>>, me_pool_for_policy: Option<Arc<MePool>>,
shared_state: Arc<ProxySharedState>, shared_state: Arc<ProxySharedState>,
me_ready_tx: watch::Sender<u64>, me_ready_tx: watch::Sender<u64>,
task_scope: RuntimeTaskScope,
) -> RuntimeWatches { ) -> RuntimeWatches {
let um_clone = upstream_manager.clone(); let um_clone = upstream_manager.clone();
let dc_overrides_for_health = config.dc_overrides.clone(); let dc_overrides_for_health = config.dc_overrides.clone();
tokio::spawn(async move { task_scope.spawn(async move {
um_clone um_clone
.run_health_checks( .run_health_checks(
prefer_ipv6, prefer_ipv6,
@@ -69,19 +122,19 @@ pub(crate) async fn spawn_runtime_tasks(
}); });
let rc_clone = replay_checker.clone(); let rc_clone = replay_checker.clone();
tokio::spawn(async move { task_scope.spawn(async move {
rc_clone.run_periodic_cleanup().await; rc_clone.run_periodic_cleanup().await;
}); });
let stats_maintenance = stats.clone(); let stats_maintenance = stats.clone();
tokio::spawn(async move { task_scope.spawn(async move {
stats_maintenance stats_maintenance
.run_periodic_user_stats_maintenance() .run_periodic_user_stats_maintenance()
.await; .await;
}); });
let ip_tracker_maintenance = ip_tracker.clone(); let ip_tracker_maintenance = ip_tracker.clone();
tokio::spawn(async move { task_scope.spawn(async move {
ip_tracker_maintenance.run_periodic_maintenance().await; ip_tracker_maintenance.run_periodic_maintenance().await;
}); });
@@ -104,6 +157,7 @@ pub(crate) async fn spawn_runtime_tasks(
config.clone(), config.clone(),
detected_ip_v4, detected_ip_v4,
detected_ip_v6, detected_ip_v6,
task_scope.cancellation_token(),
); );
startup_tracker startup_tracker
.complete_component( .complete_component(
@@ -111,21 +165,10 @@ pub(crate) async fn spawn_runtime_tasks(
Some("config hot-reload watcher started".to_string()), Some("config hot-reload watcher started".to_string()),
) )
.await; .await;
let mut config_rx_api_bridge = config_rx.clone();
let api_config_tx_bridge = api_config_tx.clone();
tokio::spawn(async move {
loop {
if config_rx_api_bridge.changed().await.is_err() {
break;
}
let cfg = config_rx_api_bridge.borrow_and_update().clone();
api_config_tx_bridge.send_replace(cfg);
}
});
let stats_policy = stats.clone(); let stats_policy = stats.clone();
let upstream_policy = upstream_manager.clone();
let mut config_rx_policy = config_rx.clone(); let mut config_rx_policy = config_rx.clone();
tokio::spawn(async move { task_scope.spawn(async move {
loop { loop {
if config_rx_policy.changed().await.is_err() { if config_rx_policy.changed().await.is_err() {
break; break;
@@ -133,6 +176,9 @@ pub(crate) async fn spawn_runtime_tasks(
let cfg = config_rx_policy.borrow_and_update().clone(); let cfg = config_rx_policy.borrow_and_update().clone();
stats_policy stats_policy
.apply_telemetry_policy(TelemetryPolicy::from_config(&cfg.general.telemetry)); .apply_telemetry_policy(TelemetryPolicy::from_config(&cfg.general.telemetry));
if let Err(error) = upstream_policy.update_dns_overrides(&cfg.network.dns_overrides) {
warn!(error = %error, "Failed to update generation DNS overrides");
}
if let Some(pool) = &me_pool_for_policy { if let Some(pool) = &me_pool_for_policy {
pool.update_runtime_transport_policy( pool.update_runtime_transport_policy(
cfg.general.me_socks_kdf_policy, cfg.general.me_socks_kdf_policy,
@@ -149,7 +195,7 @@ pub(crate) async fn spawn_runtime_tasks(
let ip_tracker_policy = ip_tracker.clone(); let ip_tracker_policy = ip_tracker.clone();
let mut config_rx_ip_limits = config_rx.clone(); let mut config_rx_ip_limits = config_rx.clone();
tokio::spawn(async move { task_scope.spawn(async move {
let mut prev_limits = config_rx_ip_limits let mut prev_limits = config_rx_ip_limits
.borrow() .borrow()
.access .access
@@ -205,7 +251,7 @@ pub(crate) async fn spawn_runtime_tasks(
config.access.cidr_rate_limits.clone(), config.access.cidr_rate_limits.clone(),
); );
let mut config_rx_rate_limits = config_rx.clone(); let mut config_rx_rate_limits = config_rx.clone();
tokio::spawn(async move { task_scope.spawn(async move {
let mut prev_user_limits = config_rx_rate_limits let mut prev_user_limits = config_rx_rate_limits
.borrow() .borrow()
.access .access
@@ -236,7 +282,7 @@ pub(crate) async fn spawn_runtime_tasks(
let shared_user_enabled = shared_state.clone(); let shared_user_enabled = shared_state.clone();
let mut config_rx_user_enabled = config_rx.clone(); let mut config_rx_user_enabled = config_rx.clone();
tokio::spawn(async move { task_scope.spawn(async move {
loop { loop {
if config_rx_user_enabled.changed().await.is_err() { if config_rx_user_enabled.changed().await.is_err() {
break; break;
@@ -257,7 +303,7 @@ pub(crate) async fn spawn_runtime_tasks(
let beobachten_writer = beobachten.clone(); let beobachten_writer = beobachten.clone();
let config_rx_beobachten = config_rx.clone(); let config_rx_beobachten = config_rx.clone();
tokio::spawn(async move { task_scope.spawn(async move {
loop { loop {
let cfg = config_rx_beobachten.borrow().clone(); let cfg = config_rx_beobachten.borrow().clone();
let sleep_secs = cfg.general.beobachten_flush_secs.max(1); let sleep_secs = cfg.general.beobachten_flush_secs.max(1);
@@ -278,7 +324,14 @@ pub(crate) async fn spawn_runtime_tasks(
}); });
if let Some(pool) = me_pool { if let Some(pool) = me_pool {
spawn_middle_proxy_runtime_tasks(config, config_rx.clone(), pool, rng, me_ready_tx); spawn_middle_proxy_runtime_tasks(
config,
config_rx.clone(),
pool,
rng,
me_ready_tx,
task_scope,
);
} }
RuntimeWatches { RuntimeWatches {
@@ -295,6 +348,7 @@ pub(crate) fn spawn_middle_proxy_runtime_tasks(
pool: Arc<MePool>, pool: Arc<MePool>,
rng: Arc<SecureRandom>, rng: Arc<SecureRandom>,
me_ready_tx: watch::Sender<u64>, me_ready_tx: watch::Sender<u64>,
task_scope: RuntimeTaskScope,
) { ) {
let reinit_trigger_capacity = config.general.me_reinit_trigger_channel.max(1); let reinit_trigger_capacity = config.general.me_reinit_trigger_channel.max(1);
let (reinit_tx, reinit_rx) = mpsc::channel::<MeReinitTrigger>(reinit_trigger_capacity); let (reinit_tx, reinit_rx) = mpsc::channel::<MeReinitTrigger>(reinit_trigger_capacity);
@@ -303,7 +357,7 @@ pub(crate) fn spawn_middle_proxy_runtime_tasks(
let rng_clone_sched = rng.clone(); let rng_clone_sched = rng.clone();
let config_rx_clone_sched = config_rx.clone(); let config_rx_clone_sched = config_rx.clone();
let me_ready_tx_sched = me_ready_tx.clone(); let me_ready_tx_sched = me_ready_tx.clone();
tokio::spawn(async move { task_scope.spawn(async move {
crate::transport::middle_proxy::me_reinit_scheduler( crate::transport::middle_proxy::me_reinit_scheduler(
pool_clone_sched, pool_clone_sched,
rng_clone_sched, rng_clone_sched,
@@ -317,7 +371,7 @@ pub(crate) fn spawn_middle_proxy_runtime_tasks(
let pool_clone = pool.clone(); let pool_clone = pool.clone();
let config_rx_clone = config_rx.clone(); let config_rx_clone = config_rx.clone();
let reinit_tx_updater = reinit_tx.clone(); let reinit_tx_updater = reinit_tx.clone();
tokio::spawn(async move { task_scope.spawn(async move {
crate::transport::middle_proxy::me_config_updater( crate::transport::middle_proxy::me_config_updater(
pool_clone, pool_clone,
config_rx_clone, config_rx_clone,
@@ -328,37 +382,12 @@ pub(crate) fn spawn_middle_proxy_runtime_tasks(
let config_rx_clone_rot = config_rx.clone(); let config_rx_clone_rot = config_rx.clone();
let reinit_tx_rotation = reinit_tx.clone(); let reinit_tx_rotation = reinit_tx.clone();
tokio::spawn(async move { task_scope.spawn(async move {
crate::transport::middle_proxy::me_rotation_task(config_rx_clone_rot, reinit_tx_rotation) crate::transport::middle_proxy::me_rotation_task(config_rx_clone_rot, reinit_tx_rotation)
.await; .await;
}); });
} }
pub(crate) async fn apply_runtime_log_filter(
has_rust_log: bool,
effective_log_level: &LogLevel,
filter_handle: reload::Handle<EnvFilter, tracing_subscriber::Registry>,
mut log_level_rx: watch::Receiver<LogLevel>,
) {
let runtime_filter = EnvFilter::new(log_filter_spec(has_rust_log, effective_log_level));
filter_handle
.reload(runtime_filter)
.expect("Failed to switch log filter");
tokio::spawn(async move {
loop {
if log_level_rx.changed().await.is_err() {
break;
}
let level = log_level_rx.borrow_and_update().clone();
let new_filter = tracing_subscriber::EnvFilter::new(log_filter_spec(false, &level));
if let Err(e) = filter_handle.reload(new_filter) {
tracing::error!("config reload: failed to update log filter: {}", e);
}
}
});
}
pub(crate) fn log_filter_spec(has_rust_log: bool, effective_log_level: &LogLevel) -> String { pub(crate) fn log_filter_spec(has_rust_log: bool, effective_log_level: &LogLevel) -> String {
if has_rust_log { if has_rust_log {
std::env::var("RUST_LOG") std::env::var("RUST_LOG")
@@ -373,12 +402,7 @@ pub(crate) fn log_filter_spec(has_rust_log: bool, effective_log_level: &LogLevel
pub(crate) async fn spawn_metrics_if_configured( pub(crate) async fn spawn_metrics_if_configured(
config: &Arc<ProxyConfig>, config: &Arc<ProxyConfig>,
startup_tracker: &Arc<StartupTracker>, startup_tracker: &Arc<StartupTracker>,
stats: Arc<Stats>, active_runtime: Arc<ArcSwap<RuntimeGeneration>>,
beobachten: Arc<BeobachtenStore>,
shared_state: Arc<ProxySharedState>,
ip_tracker: Arc<UserIpTracker>,
tls_cache: Option<Arc<TlsFrontCache>>,
config_rx: watch::Receiver<Arc<ProxyConfig>>,
) { ) {
// metrics_listen takes precedence; fall back to metrics_port for backward compat. // metrics_listen takes precedence; fall back to metrics_port for backward compat.
let metrics_target: Option<(u16, Option<String>)> = let metrics_target: Option<(u16, Option<String>)> =
@@ -408,28 +432,10 @@ pub(crate) async fn spawn_metrics_if_configured(
Some(format!("spawn metrics endpoint on {}", label)), Some(format!("spawn metrics endpoint on {}", label)),
) )
.await; .await;
let stats = stats.clone(); let active_runtime = active_runtime.clone();
let beobachten = beobachten.clone();
let shared_state = shared_state.clone();
let config_rx_metrics = config_rx.clone();
let ip_tracker_metrics = ip_tracker.clone();
let tls_cache_metrics = tls_cache.clone();
let whitelist = config.server.metrics_whitelist.clone();
let listen_backlog = config.server.listen_backlog; let listen_backlog = config.server.listen_backlog;
tokio::spawn(async move { tokio::spawn(async move {
metrics::serve( metrics::serve(port, listen, listen_backlog, active_runtime).await;
port,
listen,
listen_backlog,
stats,
beobachten,
shared_state,
ip_tracker_metrics,
tls_cache_metrics,
config_rx_metrics,
whitelist,
)
.await;
}); });
startup_tracker startup_tracker
.complete_component( .complete_component(
+35 -17
View File
@@ -12,17 +12,18 @@ use std::path::PathBuf;
use std::sync::Arc; use std::sync::Arc;
use std::time::{Duration, Instant}; use std::time::{Duration, Instant};
use arc_swap::ArcSwap;
#[cfg(not(unix))] #[cfg(not(unix))]
use tokio::signal; use tokio::signal;
#[cfg(unix)] #[cfg(unix)]
use tokio::signal::unix::{SignalKind, signal}; use tokio::signal::unix::{SignalKind, signal};
use tracing::{info, warn}; use tracing::{info, warn};
use super::generation::RuntimeGeneration;
use super::helpers::{format_uptime, unit_label};
use super::reload_supervisor::ReloadSupervisorHandle;
use crate::stats::Stats; use crate::stats::Stats;
use crate::synlimit_control; use crate::synlimit_control;
use crate::transport::middle_proxy::MePool;
use super::helpers::{format_uptime, unit_label};
/// Signal that triggered shutdown. /// Signal that triggered shutdown.
#[derive(Debug, Clone, Copy, PartialEq, Eq)] #[derive(Debug, Clone, Copy, PartialEq, Eq)]
@@ -48,17 +49,19 @@ impl std::fmt::Display for ShutdownSignal {
/// Waits for a shutdown signal and performs graceful shutdown. /// Waits for a shutdown signal and performs graceful shutdown.
pub(crate) async fn wait_for_shutdown( pub(crate) async fn wait_for_shutdown(
process_started_at: Instant, process_started_at: Instant,
me_pool: Option<Arc<MePool>>, active_runtime: Arc<ArcSwap<RuntimeGeneration>>,
stats: Arc<Stats>,
quota_state_path: PathBuf, quota_state_path: PathBuf,
synlimit_controller: synlimit_control::SynlimitController,
reload_supervisor: ReloadSupervisorHandle,
) { ) {
let signal = wait_for_shutdown_signal().await; let signal = wait_for_shutdown_signal().await;
perform_shutdown( perform_shutdown(
signal, signal,
process_started_at, process_started_at,
me_pool, active_runtime,
&stats,
quota_state_path, quota_state_path,
synlimit_controller,
reload_supervisor,
) )
.await; .await;
} }
@@ -87,13 +90,18 @@ async fn wait_for_shutdown_signal() -> ShutdownSignal {
async fn perform_shutdown( async fn perform_shutdown(
signal: ShutdownSignal, signal: ShutdownSignal,
process_started_at: Instant, process_started_at: Instant,
me_pool: Option<Arc<MePool>>, active_runtime: Arc<ArcSwap<RuntimeGeneration>>,
stats: &Stats,
quota_state_path: PathBuf, quota_state_path: PathBuf,
synlimit_controller: synlimit_control::SynlimitController,
reload_supervisor: ReloadSupervisorHandle,
) { ) {
let shutdown_started_at = Instant::now(); let shutdown_started_at = Instant::now();
info!(signal = %signal, "Received shutdown signal"); info!(signal = %signal, "Received shutdown signal");
reload_supervisor.quiesce().await;
let runtime = active_runtime.load_full();
let stats = runtime.stats.as_ref();
// Dump stats if SIGQUIT // Dump stats if SIGQUIT
if signal == ShutdownSignal::Quit { if signal == ShutdownSignal::Quit {
dump_stats(stats, process_started_at); dump_stats(stats, process_started_at);
@@ -103,12 +111,10 @@ async fn perform_shutdown(
let uptime_secs = process_started_at.elapsed().as_secs(); let uptime_secs = process_started_at.elapsed().as_secs();
info!("Uptime: {}", format_uptime(uptime_secs)); info!("Uptime: {}", format_uptime(uptime_secs));
if let Err(error) = synlimit_control::clear_synlimit_rules_all_backends().await {
warn!(error = %error, "Failed to clear SYN limiter rules during shutdown");
}
// Graceful ME pool shutdown // Graceful ME pool shutdown
if let Some(pool) = &me_pool { runtime.stop_sessions().await;
runtime.stop_background_tasks().await;
if let Some(pool) = runtime.current_me_pool().await {
match tokio::time::timeout(Duration::from_secs(2), pool.shutdown_send_close_conn_all()) match tokio::time::timeout(Duration::from_secs(2), pool.shutdown_send_close_conn_all())
.await .await
{ {
@@ -124,6 +130,11 @@ async fn perform_shutdown(
} }
} }
synlimit_controller.shutdown().await;
if let Err(error) = synlimit_control::clear_synlimit_rules_all_backends().await {
warn!(error = %error, "Failed to clear SYN limiter rules during shutdown");
}
match crate::quota_state::save_quota_state(&quota_state_path, stats).await { match crate::quota_state::save_quota_state(&quota_state_path, stats).await {
Ok(()) => { Ok(()) => {
info!( info!(
@@ -191,7 +202,10 @@ fn dump_stats(stats: &Stats, process_started_at: Instant) {
/// - SIGUSR1: Log rotation acknowledgment (for external log rotation tools) /// - SIGUSR1: Log rotation acknowledgment (for external log rotation tools)
/// - SIGUSR2: Dump runtime status to log /// - SIGUSR2: Dump runtime status to log
#[cfg(unix)] #[cfg(unix)]
pub(crate) fn spawn_signal_handlers(stats: Arc<Stats>, process_started_at: Instant) { pub(crate) fn spawn_signal_handlers(
active_runtime: Arc<ArcSwap<RuntimeGeneration>>,
process_started_at: Instant,
) {
tokio::spawn(async move { tokio::spawn(async move {
let mut sigusr1 = let mut sigusr1 =
signal(SignalKind::user_defined1()).expect("Failed to register SIGUSR1 handler"); signal(SignalKind::user_defined1()).expect("Failed to register SIGUSR1 handler");
@@ -204,7 +218,8 @@ pub(crate) fn spawn_signal_handlers(stats: Arc<Stats>, process_started_at: Insta
handle_sigusr1(); handle_sigusr1();
} }
_ = sigusr2.recv() => { _ = sigusr2.recv() => {
handle_sigusr2(&stats, process_started_at); let runtime = active_runtime.load_full();
handle_sigusr2(runtime.stats.as_ref(), process_started_at);
} }
} }
} }
@@ -213,7 +228,10 @@ pub(crate) fn spawn_signal_handlers(stats: Arc<Stats>, process_started_at: Insta
/// No-op on non-Unix platforms. /// No-op on non-Unix platforms.
#[cfg(not(unix))] #[cfg(not(unix))]
pub(crate) fn spawn_signal_handlers(_stats: Arc<Stats>, _process_started_at: Instant) { pub(crate) fn spawn_signal_handlers(
_active_runtime: Arc<ArcSwap<RuntimeGeneration>>,
_process_started_at: Instant,
) {
// No SIGUSR1/SIGUSR2 on non-Unix // No SIGUSR1/SIGUSR2 on non-Unix
} }
+287 -127
View File
@@ -5,11 +5,86 @@ use rand::RngExt;
use tracing::warn; use tracing::warn;
use crate::config::ProxyConfig; use crate::config::ProxyConfig;
use crate::error::{ProxyError, Result};
use crate::startup::{COMPONENT_TLS_FRONT_BOOTSTRAP, StartupTracker}; use crate::startup::{COMPONENT_TLS_FRONT_BOOTSTRAP, StartupTracker};
use crate::tls_front::TlsFrontCache; use crate::tls_front::TlsFrontCache;
use crate::tls_front::fetcher::TlsFetchStrategy; use crate::tls_front::fetcher::TlsFetchStrategy;
use crate::transport::UpstreamManager; use crate::transport::UpstreamManager;
use super::generation::RuntimeTaskScope;
/// Readiness requirement for TLS-front cache initialization.
#[derive(Clone, Copy)]
pub(crate) enum TlsBootstrapPolicy {
BestEffort,
RequireReady,
}
#[derive(Clone)]
struct TlsFetchContext {
cache: Arc<TlsFrontCache>,
domains: Vec<String>,
mask_host: String,
primary_domain: String,
mask_unix_sock: Option<String>,
tls_fetch_scope: Option<String>,
upstream_manager: Arc<UpstreamManager>,
strategy: TlsFetchStrategy,
port: u16,
proxy_protocol: u8,
}
impl TlsFetchContext {
async fn fetch_all(&self, failure_message: &'static str) {
let mut join = tokio::task::JoinSet::new();
for domain in self.domains.clone() {
let cache = self.cache.clone();
let host = tls_fetch_host_for_domain(&self.mask_host, &self.primary_domain, &domain);
let unix_sock = self.mask_unix_sock.clone();
let scope = self.tls_fetch_scope.clone();
let upstream = self.upstream_manager.clone();
let strategy = self.strategy.clone();
let port = self.port;
let proxy_protocol = self.proxy_protocol;
join.spawn(async move {
match crate::tls_front::fetcher::fetch_real_tls_with_strategy(
&host,
port,
&domain,
&strategy,
Some(upstream),
scope.as_deref(),
proxy_protocol,
unix_sock.as_deref(),
)
.await
{
Ok(result) => cache.update_from_fetch(&domain, result).await,
Err(error) => warn!(domain = %domain, error = %error, failure_message),
}
});
}
while let Some(result) = join.join_next().await {
if let Err(error) = result {
warn!(error = %error, "TLS emulation fetch task join failed");
}
}
}
async fn fetch_all_with_budget(&self, phase: &'static str) {
if tokio::time::timeout(self.strategy.total_budget, self.fetch_all(phase))
.await
.is_err()
{
warn!(
phase,
timeout_ms = self.strategy.total_budget.as_millis(),
"TLS emulation fetch budget exhausted"
);
}
}
}
fn tls_fetch_host_for_domain(mask_host: &str, primary_tls_domain: &str, domain: &str) -> String { fn tls_fetch_host_for_domain(mask_host: &str, primary_tls_domain: &str, domain: &str) -> String {
if mask_host.eq_ignore_ascii_case(primary_tls_domain) { if mask_host.eq_ignore_ascii_case(primary_tls_domain) {
domain.to_string() domain.to_string()
@@ -18,12 +93,24 @@ fn tls_fetch_host_for_domain(mask_host: &str, primary_tls_domain: &str, domain:
} }
} }
fn readiness_error(default_domains: &[String]) -> Option<String> {
(!default_domains.is_empty()).then(|| {
format!(
"TLS-front profiles are not ready for domains: {}",
default_domains.join(", ")
)
})
}
/// Initializes the TLS-front cache and generation-owned refresh tasks.
pub(crate) async fn bootstrap_tls_front( pub(crate) async fn bootstrap_tls_front(
config: &ProxyConfig, config: &ProxyConfig,
tls_domains: &[String], tls_domains: &[String],
upstream_manager: Arc<UpstreamManager>, upstream_manager: Arc<UpstreamManager>,
startup_tracker: &Arc<StartupTracker>, startup_tracker: &Arc<StartupTracker>,
) -> Option<Arc<TlsFrontCache>> { task_scope: RuntimeTaskScope,
policy: TlsBootstrapPolicy,
) -> Result<Option<Arc<TlsFrontCache>>> {
startup_tracker startup_tracker
.start_component( .start_component(
COMPONENT_TLS_FRONT_BOOTSTRAP, COMPONENT_TLS_FRONT_BOOTSTRAP,
@@ -31,7 +118,16 @@ pub(crate) async fn bootstrap_tls_front(
) )
.await; .await;
let tls_cache: Option<Arc<TlsFrontCache>> = if config.censorship.tls_emulation { if !config.censorship.tls_emulation {
startup_tracker
.skip_component(
COMPONENT_TLS_FRONT_BOOTSTRAP,
Some("censorship.tls_emulation is false".to_string()),
)
.await;
return Ok(None);
}
let cache = Arc::new(TlsFrontCache::new( let cache = Arc::new(TlsFrontCache::new(
tls_domains, tls_domains,
config.censorship.fake_cert_len, config.censorship.fake_cert_len,
@@ -39,18 +135,21 @@ pub(crate) async fn bootstrap_tls_front(
)); ));
cache.load_from_disk().await; cache.load_from_disk().await;
let port = config.censorship.mask_port; let tls_fetch = config.censorship.tls_fetch.clone();
let proxy_protocol = config.censorship.mask_proxy_protocol; let fetch_context = TlsFetchContext {
let mask_host = config cache: cache.clone(),
domains: tls_domains.to_vec(),
mask_host: config
.censorship .censorship
.mask_host .mask_host
.clone() .clone()
.unwrap_or_else(|| config.censorship.tls_domain.clone()); .unwrap_or_else(|| config.censorship.tls_domain.clone()),
let mask_unix_sock = config.censorship.mask_unix_sock.clone(); primary_domain: config.censorship.tls_domain.clone(),
let tls_fetch_scope = (!config.censorship.tls_fetch_scope.is_empty()) mask_unix_sock: config.censorship.mask_unix_sock.clone(),
.then(|| config.censorship.tls_fetch_scope.clone()); tls_fetch_scope: (!config.censorship.tls_fetch_scope.is_empty())
let tls_fetch = config.censorship.tls_fetch.clone(); .then(|| config.censorship.tls_fetch_scope.clone()),
let fetch_strategy = TlsFetchStrategy { upstream_manager,
strategy: TlsFetchStrategy {
profiles: tls_fetch.profiles, profiles: tls_fetch.profiles,
strict_route: tls_fetch.strict_route, strict_route: tls_fetch.strict_route,
attempt_timeout: Duration::from_millis(tls_fetch.attempt_timeout_ms.max(1)), attempt_timeout: Duration::from_millis(tls_fetch.attempt_timeout_ms.max(1)),
@@ -58,150 +157,110 @@ pub(crate) async fn bootstrap_tls_front(
grease_enabled: tls_fetch.grease_enabled, grease_enabled: tls_fetch.grease_enabled,
deterministic: tls_fetch.deterministic, deterministic: tls_fetch.deterministic,
profile_cache_ttl: Duration::from_secs(tls_fetch.profile_cache_ttl_secs), profile_cache_ttl: Duration::from_secs(tls_fetch.profile_cache_ttl_secs),
},
port: config.censorship.mask_port,
proxy_protocol: config.censorship.mask_proxy_protocol,
}; };
let fetch_timeout = fetch_strategy.total_budget;
let cache_initial = cache.clone(); match policy {
let domains_initial = tls_domains.to_vec(); TlsBootstrapPolicy::BestEffort => {
let host_initial = mask_host.clone(); let initial_fetch = fetch_context.clone();
let primary_initial = config.censorship.tls_domain.clone(); let fake_cert_len = config.censorship.fake_cert_len;
let unix_sock_initial = mask_unix_sock.clone(); task_scope.spawn(async move {
let scope_initial = tls_fetch_scope.clone(); initial_fetch
let upstream_initial = upstream_manager.clone(); .fetch_all_with_budget("TLS emulation initial fetch failed")
let strategy_initial = fetch_strategy.clone(); .await;
tokio::spawn(async move { for domain in initial_fetch
let mut join = tokio::task::JoinSet::new(); .cache
for domain in domains_initial { .default_profile_domains(&initial_fetch.domains)
let cache_domain = cache_initial.clone();
let host_domain =
tls_fetch_host_for_domain(&host_initial, &primary_initial, &domain);
let unix_sock_domain = unix_sock_initial.clone();
let scope_domain = scope_initial.clone();
let upstream_domain = upstream_initial.clone();
let strategy_domain = strategy_initial.clone();
join.spawn(async move {
match crate::tls_front::fetcher::fetch_real_tls_with_strategy(
&host_domain,
port,
&domain,
&strategy_domain,
Some(upstream_domain),
scope_domain.as_deref(),
proxy_protocol,
unix_sock_domain.as_deref(),
)
.await .await
{ {
Ok(res) => cache_domain.update_from_fetch(&domain, res).await,
Err(e) => {
warn!(domain = %domain, error = %e, "TLS emulation initial fetch failed")
}
}
});
}
while let Some(res) = join.join_next().await {
if let Err(e) = res {
warn!(error = %e, "TLS emulation initial fetch task join failed");
}
}
});
let cache_timeout = cache.clone();
let domains_timeout = tls_domains.to_vec();
let fake_cert_len = config.censorship.fake_cert_len;
tokio::spawn(async move {
tokio::time::sleep(fetch_timeout).await;
for domain in domains_timeout {
let cached = cache_timeout.get(&domain).await;
if cached.domain == "default" {
warn!( warn!(
domain = %domain, domain = %domain,
timeout_secs = fetch_timeout.as_secs(), timeout_ms = initial_fetch.strategy.total_budget.as_millis(),
fake_cert_len, fake_cert_len,
"TLS-front fetch not ready within timeout; using cache/default fake cert fallback" "TLS-front fetch not ready within timeout; using cache/default fake cert fallback"
); );
} }
}
}); });
}
TlsBootstrapPolicy::RequireReady => {
fetch_context
.fetch_all_with_budget("TLS emulation initial fetch failed")
.await;
let default_domains = cache.default_profile_domains(tls_domains).await;
if let Some(error) = readiness_error(&default_domains) {
startup_tracker
.fail_component(COMPONENT_TLS_FRONT_BOOTSTRAP, Some(error.clone()))
.await;
return Err(ProxyError::Proxy(error));
}
}
}
let cache_refresh = cache.clone(); let refresh_context = fetch_context;
let domains_refresh = tls_domains.to_vec(); task_scope.spawn(async move {
let host_refresh = mask_host.clone();
let primary_refresh = config.censorship.tls_domain.clone();
let unix_sock_refresh = mask_unix_sock.clone();
let scope_refresh = tls_fetch_scope.clone();
let upstream_refresh = upstream_manager.clone();
let strategy_refresh = fetch_strategy.clone();
tokio::spawn(async move {
loop { loop {
let base_secs = rand::rng().random_range(4 * 3600..=6 * 3600); let base_secs = rand::rng().random_range(4 * 3600..=6 * 3600);
let jitter_secs = rand::rng().random_range(0..=7200); let jitter_secs = rand::rng().random_range(0..=7200);
tokio::time::sleep(Duration::from_secs(base_secs + jitter_secs)).await; tokio::time::sleep(Duration::from_secs(base_secs + jitter_secs)).await;
refresh_context
let mut join = tokio::task::JoinSet::new(); .fetch_all_with_budget("TLS emulation refresh failed")
for domain in domains_refresh.clone() {
let cache_domain = cache_refresh.clone();
let host_domain =
tls_fetch_host_for_domain(&host_refresh, &primary_refresh, &domain);
let unix_sock_domain = unix_sock_refresh.clone();
let scope_domain = scope_refresh.clone();
let upstream_domain = upstream_refresh.clone();
let strategy_domain = strategy_refresh.clone();
join.spawn(async move {
match crate::tls_front::fetcher::fetch_real_tls_with_strategy(
&host_domain,
port,
&domain,
&strategy_domain,
Some(upstream_domain),
scope_domain.as_deref(),
proxy_protocol,
unix_sock_domain.as_deref(),
)
.await
{
Ok(res) => cache_domain.update_from_fetch(&domain, res).await,
Err(e) => {
warn!(domain = %domain, error = %e, "TLS emulation refresh failed")
}
}
});
}
while let Some(res) = join.join_next().await {
if let Err(e) = res {
warn!(error = %e, "TLS emulation refresh task join failed");
}
}
}
});
Some(cache)
} else {
startup_tracker
.skip_component(
COMPONENT_TLS_FRONT_BOOTSTRAP,
Some("censorship.tls_emulation is false".to_string()),
)
.await; .await;
None }
}; });
if tls_cache.is_some() {
startup_tracker startup_tracker
.complete_component( .complete_component(
COMPONENT_TLS_FRONT_BOOTSTRAP, COMPONENT_TLS_FRONT_BOOTSTRAP,
Some("tls front cache is initialized".to_string()), Some("tls front cache is initialized".to_string()),
) )
.await; .await;
} Ok(Some(cache))
tls_cache
} }
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::tls_fetch_host_for_domain; use super::*;
use crate::startup::StartupComponentStatus;
use crate::stats::Stats;
fn test_config(cache_dir: &std::path::Path) -> ProxyConfig {
let mut config = ProxyConfig::default();
config.censorship.tls_emulation = true;
config.censorship.tls_domain = "front.example".to_string();
config.censorship.mask_host = Some("127.0.0.1".to_string());
config.censorship.mask_port = 1;
config.censorship.tls_front_dir = cache_dir.display().to_string();
config.censorship.tls_fetch.profiles.truncate(1);
config.censorship.tls_fetch.attempt_timeout_ms = 10;
config.censorship.tls_fetch.total_budget_ms = 20;
config
}
fn upstream_manager(config: &ProxyConfig) -> Arc<UpstreamManager> {
Arc::new(UpstreamManager::new(
Vec::new(),
config.general.upstream_connect_retry_attempts,
config.general.upstream_connect_retry_backoff_ms,
config.general.upstream_connect_budget_ms,
config.general.tg_connect,
config.general.upstream_unhealthy_fail_threshold,
config.general.upstream_connect_failfast_hard_errors,
Arc::new(Stats::new()),
))
}
async fn tls_component_status(tracker: &StartupTracker) -> StartupComponentStatus {
tracker
.snapshot()
.await
.components
.into_iter()
.find(|component| component.id == COMPONENT_TLS_FRONT_BOOTSTRAP)
.unwrap()
.status
}
#[test] #[test]
fn tls_fetch_host_uses_each_domain_when_mask_host_is_primary_default() { fn tls_fetch_host_uses_each_domain_when_mask_host_is_primary_default() {
@@ -218,4 +277,105 @@ mod tests {
"origin.example" "origin.example"
); );
} }
#[test]
fn readiness_rejects_only_default_profiles() {
assert!(readiness_error(&[]).is_none());
assert_eq!(
readiness_error(&["front.example".to_string()]),
Some("TLS-front profiles are not ready for domains: front.example".to_string())
);
}
#[tokio::test]
async fn require_ready_rejects_default_cache_after_bounded_fetch_failure() {
let cache_dir = tempfile::tempdir().unwrap();
let config = test_config(cache_dir.path());
let domains = vec![config.censorship.tls_domain.clone()];
let tracker = Arc::new(StartupTracker::new(1));
let scope = RuntimeTaskScope::new();
let result = bootstrap_tls_front(
&config,
&domains,
upstream_manager(&config),
&tracker,
scope.clone(),
TlsBootstrapPolicy::RequireReady,
)
.await;
assert!(result.is_err());
assert_eq!(
tls_component_status(&tracker).await,
StartupComponentStatus::Failed
);
scope.stop().await;
}
#[tokio::test]
async fn require_ready_accepts_non_default_disk_cache_when_refresh_fails() {
let cache_dir = tempfile::tempdir().unwrap();
let config = test_config(cache_dir.path());
let domains = vec![config.censorship.tls_domain.clone()];
let seed = TlsFrontCache::new(&domains, config.censorship.fake_cert_len, cache_dir.path());
let mut cached = seed.default_entry().as_ref().clone();
cached.domain = domains[0].clone();
tokio::fs::write(
cache_dir.path().join("front.example.json"),
serde_json::to_vec(&cached).unwrap(),
)
.await
.unwrap();
let tracker = Arc::new(StartupTracker::new(1));
let scope = RuntimeTaskScope::new();
let cache = bootstrap_tls_front(
&config,
&domains,
upstream_manager(&config),
&tracker,
scope.clone(),
TlsBootstrapPolicy::RequireReady,
)
.await
.unwrap()
.unwrap();
assert!(cache.default_profile_domains(&domains).await.is_empty());
assert_eq!(
tls_component_status(&tracker).await,
StartupComponentStatus::Ready
);
scope.stop().await;
}
#[tokio::test]
async fn best_effort_returns_ready_and_refresh_tasks_are_scope_owned() {
let cache_dir = tempfile::tempdir().unwrap();
let config = test_config(cache_dir.path());
let domains = vec![config.censorship.tls_domain.clone()];
let tracker = Arc::new(StartupTracker::new(1));
let scope = RuntimeTaskScope::new();
let cache = bootstrap_tls_front(
&config,
&domains,
upstream_manager(&config),
&tracker,
scope.clone(),
TlsBootstrapPolicy::BestEffort,
)
.await
.unwrap();
assert!(cache.is_some());
assert_eq!(
tls_component_status(&tracker).await,
StartupComponentStatus::Ready
);
tokio::time::timeout(Duration::from_secs(1), scope.stop())
.await
.unwrap();
}
} }
+177 -84
View File
@@ -4,12 +4,12 @@ use std::net::SocketAddr;
use std::sync::Arc; use std::sync::Arc;
use std::time::Duration; use std::time::Duration;
use arc_swap::ArcSwap;
use http_body_util::Full; use http_body_util::Full;
use hyper::body::Bytes; use hyper::body::Bytes;
use hyper::server::conn::http1; use hyper::server::conn::http1;
use hyper::service::service_fn; use hyper::service::service_fn;
use hyper::{Request, Response, StatusCode}; use hyper::{Request, Response, StatusCode};
use ipnetwork::IpNetwork;
use tokio::net::TcpListener; use tokio::net::TcpListener;
use tokio::sync::Semaphore; use tokio::sync::Semaphore;
use tokio::time::timeout; use tokio::time::timeout;
@@ -17,6 +17,7 @@ use tracing::{debug, info, warn};
use crate::config::ProxyConfig; use crate::config::ProxyConfig;
use crate::ip_tracker::UserIpTracker; use crate::ip_tracker::UserIpTracker;
use crate::maestro::generation::RuntimeGeneration;
use crate::proxy::shared_state::ProxySharedState; use crate::proxy::shared_state::ProxySharedState;
use crate::stats::Stats; use crate::stats::Stats;
use crate::stats::beobachten::BeobachtenStore; use crate::stats::beobachten::BeobachtenStore;
@@ -36,16 +37,8 @@ pub async fn serve(
port: u16, port: u16,
listen: Option<String>, listen: Option<String>,
listen_backlog: u32, listen_backlog: u32,
stats: Arc<Stats>, active_runtime: Arc<ArcSwap<RuntimeGeneration>>,
beobachten: Arc<BeobachtenStore>,
shared_state: Arc<ProxySharedState>,
ip_tracker: Arc<UserIpTracker>,
tls_cache: Option<Arc<TlsFrontCache>>,
config_rx: tokio::sync::watch::Receiver<Arc<ProxyConfig>>,
whitelist: Vec<IpNetwork>,
) { ) {
let whitelist = Arc::new(whitelist);
// If `metrics_listen` is set, bind on that single address only. // If `metrics_listen` is set, bind on that single address only.
if let Some(ref listen_addr) = listen { if let Some(ref listen_addr) = listen {
let addr: SocketAddr = match listen_addr.parse() { let addr: SocketAddr = match listen_addr.parse() {
@@ -61,17 +54,7 @@ pub async fn serve(
match bind_metrics_listener(addr, ipv6_only, listen_backlog) { match bind_metrics_listener(addr, ipv6_only, listen_backlog) {
Ok(listener) => { Ok(listener) => {
info!("Metrics endpoint: http://{}/metrics and /beobachten", addr); info!("Metrics endpoint: http://{}/metrics and /beobachten", addr);
serve_listener( serve_listener(listener, active_runtime).await;
listener,
stats,
beobachten,
shared_state,
ip_tracker,
tls_cache,
config_rx,
whitelist,
)
.await;
} }
Err(e) => { Err(e) => {
warn!(error = %e, "Failed to bind metrics on {}", addr); warn!(error = %e, "Failed to bind metrics on {}", addr);
@@ -117,50 +100,14 @@ pub async fn serve(
warn!("Metrics listener is unavailable on both IPv4 and IPv6"); warn!("Metrics listener is unavailable on both IPv4 and IPv6");
} }
(Some(listener), None) | (None, Some(listener)) => { (Some(listener), None) | (None, Some(listener)) => {
serve_listener( serve_listener(listener, active_runtime).await;
listener,
stats,
beobachten,
shared_state,
ip_tracker,
tls_cache,
config_rx,
whitelist,
)
.await;
} }
(Some(listener4), Some(listener6)) => { (Some(listener4), Some(listener6)) => {
let stats_v6 = stats.clone(); let active_runtime_v6 = active_runtime.clone();
let beobachten_v6 = beobachten.clone();
let shared_state_v6 = shared_state.clone();
let ip_tracker_v6 = ip_tracker.clone();
let tls_cache_v6 = tls_cache.clone();
let config_rx_v6 = config_rx.clone();
let whitelist_v6 = whitelist.clone();
tokio::spawn(async move { tokio::spawn(async move {
serve_listener( serve_listener(listener6, active_runtime_v6).await;
listener6,
stats_v6,
beobachten_v6,
shared_state_v6,
ip_tracker_v6,
tls_cache_v6,
config_rx_v6,
whitelist_v6,
)
.await;
}); });
serve_listener( serve_listener(listener4, active_runtime).await;
listener4,
stats,
beobachten,
shared_state,
ip_tracker,
tls_cache,
config_rx,
whitelist,
)
.await;
} }
} }
} }
@@ -180,16 +127,7 @@ fn bind_metrics_listener(
TcpListener::from_std(socket.into()) TcpListener::from_std(socket.into())
} }
async fn serve_listener( async fn serve_listener(listener: TcpListener, active_runtime: Arc<ArcSwap<RuntimeGeneration>>) {
listener: TcpListener,
stats: Arc<Stats>,
beobachten: Arc<BeobachtenStore>,
shared_state: Arc<ProxySharedState>,
ip_tracker: Arc<UserIpTracker>,
tls_cache: Option<Arc<TlsFrontCache>>,
config_rx: tokio::sync::watch::Receiver<Arc<ProxyConfig>>,
whitelist: Arc<Vec<IpNetwork>>,
) {
let connection_permits = Arc::new(Semaphore::new(METRICS_MAX_CONTROL_CONNECTIONS)); let connection_permits = Arc::new(Semaphore::new(METRICS_MAX_CONTROL_CONNECTIONS));
loop { loop {
@@ -201,7 +139,15 @@ async fn serve_listener(
} }
}; };
if !whitelist.is_empty() && !whitelist.iter().any(|net| net.contains(peer.ip())) { let runtime = active_runtime.load_full();
let config = runtime.config();
if !config.server.metrics_whitelist.is_empty()
&& !config
.server
.metrics_whitelist
.iter()
.any(|net| net.contains(peer.ip()))
{
debug!(peer = %peer, "Metrics request denied by whitelist"); debug!(peer = %peer, "Metrics request denied by whitelist");
continue; continue;
} }
@@ -218,21 +164,17 @@ async fn serve_listener(
} }
}; };
let stats = stats.clone(); let active_runtime = active_runtime.clone();
let beobachten = beobachten.clone();
let shared_state = shared_state.clone();
let ip_tracker = ip_tracker.clone();
let tls_cache = tls_cache.clone();
let config_rx_conn = config_rx.clone();
tokio::spawn(async move { tokio::spawn(async move {
let _connection_permit = connection_permit; let _connection_permit = connection_permit;
let svc = service_fn(move |req| { let svc = service_fn(move |req| {
let stats = stats.clone(); let runtime = active_runtime.load_full();
let beobachten = beobachten.clone(); let stats = runtime.stats.clone();
let shared_state = shared_state.clone(); let beobachten = runtime.beobachten.clone();
let ip_tracker = ip_tracker.clone(); let shared_state = runtime.proxy_shared.clone();
let tls_cache = tls_cache.clone(); let ip_tracker = runtime.ip_tracker.clone();
let config = config_rx_conn.borrow().clone(); let tls_cache = runtime.tls_cache.clone();
let config = runtime.config();
async move { async move {
handle( handle(
req, req,
@@ -595,6 +537,81 @@ async fn render_metrics(
"telemt_buffer_pool_buffers_total{{kind=\"in_use\"}} {}", "telemt_buffer_pool_buffers_total{{kind=\"in_use\"}} {}",
stats.get_buffer_pool_in_use_gauge() stats.get_buffer_pool_in_use_gauge()
); );
let _ = writeln!(
out,
"# HELP telemt_buffer_pool_events_total Buffer-pool allocation lifecycle events"
);
let _ = writeln!(out, "# TYPE telemt_buffer_pool_events_total counter");
let _ = writeln!(
out,
"telemt_buffer_pool_events_total{{event=\"replaced_nonstandard\"}} {}",
stats.get_buffer_pool_replaced_nonstandard_total()
);
let direct_budget = shared_state.direct_buffer_budget.snapshot();
let _ = writeln!(
out,
"# HELP telemt_direct_relay_buffer_budget_bytes Direct relay copy-buffer budget and memory inputs"
);
let _ = writeln!(out, "# TYPE telemt_direct_relay_buffer_budget_bytes gauge");
for (kind, value) in [
("hard_limit", direct_budget.hard_limit_bytes),
("target", direct_budget.target_bytes),
("reserved", direct_budget.reserved_bytes),
("memory_total", direct_budget.memory_total_bytes),
("memory_available", direct_budget.memory_available_bytes),
("process_rss", direct_budget.process_rss_bytes),
] {
let _ = writeln!(
out,
"telemt_direct_relay_buffer_budget_bytes{{kind=\"{}\"}} {}",
kind, value
);
}
let _ = writeln!(
out,
"# HELP telemt_direct_relay_buffer_budget_events_total Direct relay buffer-budget lifecycle events"
);
let _ = writeln!(
out,
"# TYPE telemt_direct_relay_buffer_budget_events_total counter"
);
for (result, value) in [
("promotion", direct_budget.promotion_total),
("promotion_denied", direct_budget.promotion_denied_total),
("minimum_fallback", direct_budget.minimum_fallback_total),
("admission_rejected", direct_budget.admission_rejected_total),
("quiet_demotion", direct_budget.quiet_demotion_total),
(
"write_pressure_demotion",
direct_budget.write_pressure_demotion_total,
),
(
"global_pressure_demotion",
direct_budget.global_pressure_demotion_total,
),
] {
let _ = writeln!(
out,
"telemt_direct_relay_buffer_budget_events_total{{result=\"{}\"}} {}",
result, value
);
}
let _ = writeln!(
out,
"# HELP telemt_direct_relay_buffer_sessions Current Direct relay sessions by adaptive tier"
);
let _ = writeln!(out, "# TYPE telemt_direct_relay_buffer_sessions gauge");
for (tier, value) in ["base", "tier1", "tier2", "tier3"]
.into_iter()
.zip(direct_budget.tier_sessions)
{
let _ = writeln!(
out,
"telemt_direct_relay_buffer_sessions{{tier=\"{}\"}} {}",
tier, value
);
}
let _ = writeln!( let _ = writeln!(
out, out,
@@ -2435,6 +2452,82 @@ async fn render_metrics(
} }
); );
let _ = writeln!(
out,
"# HELP telemt_me_writer_byte_budget_limit_bytes Configured resident-memory budget per ME writer"
);
let _ = writeln!(out, "# TYPE telemt_me_writer_byte_budget_limit_bytes gauge");
let _ = writeln!(
out,
"telemt_me_writer_byte_budget_limit_bytes {}",
if me_allows_normal {
stats.get_me_writer_byte_budget_limit_bytes_gauge()
} else {
0
}
);
let _ = writeln!(
out,
"# HELP telemt_me_writer_byte_budget_reserved_bytes Aggregate ME writer memory reservations by lifecycle state"
);
let _ = writeln!(
out,
"# TYPE telemt_me_writer_byte_budget_reserved_bytes gauge"
);
let _ = writeln!(
out,
"telemt_me_writer_byte_budget_reserved_bytes{{state=\"queued\"}} {}",
if me_allows_normal {
stats.get_me_writer_byte_budget_queued_bytes_gauge()
} else {
0
}
);
let _ = writeln!(
out,
"telemt_me_writer_byte_budget_reserved_bytes{{state=\"inflight\"}} {}",
if me_allows_normal {
stats.get_me_writer_byte_budget_inflight_bytes_gauge()
} else {
0
}
);
let _ = writeln!(
out,
"# HELP telemt_me_writer_byte_budget_events_total ME writer byte-budget outcomes"
);
let _ = writeln!(
out,
"# TYPE telemt_me_writer_byte_budget_events_total counter"
);
let _ = writeln!(
out,
"telemt_me_writer_byte_budget_events_total{{result=\"wait\"}} {}",
if me_allows_normal {
stats.get_me_writer_byte_budget_wait_total()
} else {
0
}
);
let _ = writeln!(
out,
"telemt_me_writer_byte_budget_events_total{{result=\"timeout\"}} {}",
if me_allows_normal {
stats.get_me_writer_byte_budget_timeout_total()
} else {
0
}
);
let _ = writeln!(
out,
"telemt_me_writer_byte_budget_events_total{{result=\"oversize\"}} {}",
if me_allows_normal {
stats.get_me_writer_byte_budget_oversize_total()
} else {
0
}
);
let _ = writeln!( let _ = writeln!(
out, out,
"# HELP telemt_me_writer_pick_total ME writer-pick outcomes by mode and result" "# HELP telemt_me_writer_pick_total ME writer-pick outcomes by mode and result"
+43
View File
@@ -8,6 +8,33 @@ use crate::error::{ProxyError, Result};
type OverrideMap = HashMap<(String, u16), IpAddr>; type OverrideMap = HashMap<(String, u16), IpAddr>;
/// Immutable DNS override snapshot owned by one runtime generation.
#[derive(Debug, Clone, Default)]
pub struct DnsOverrides {
entries: std::sync::Arc<OverrideMap>,
}
impl DnsOverrides {
/// Parses a validated generation-local override snapshot.
pub fn from_entries(entries: &[String]) -> Result<Self> {
Ok(Self {
entries: std::sync::Arc::new(parse_entries(entries)?),
})
}
/// Resolves a generation-local hostname override.
pub fn resolve(&self, host: &str, port: u16) -> Option<IpAddr> {
self.entries
.get(&(host.to_ascii_lowercase(), port))
.copied()
}
/// Resolves a generation-local override as a socket address.
pub fn resolve_socket_addr(&self, host: &str, port: u16) -> Option<SocketAddr> {
self.resolve(host, port).map(|ip| SocketAddr::new(ip, port))
}
}
static DNS_OVERRIDES: OnceLock<RwLock<OverrideMap>> = OnceLock::new(); static DNS_OVERRIDES: OnceLock<RwLock<OverrideMap>> = OnceLock::new();
fn overrides_store() -> &'static RwLock<OverrideMap> { fn overrides_store() -> &'static RwLock<OverrideMap> {
@@ -180,6 +207,22 @@ mod tests {
assert_eq!(resolved, Some("127.0.0.1".parse().unwrap())); assert_eq!(resolved, Some("127.0.0.1".parse().unwrap()));
} }
#[test]
fn generation_snapshots_do_not_observe_each_other() {
let first = DnsOverrides::from_entries(&["example.com:443:127.0.0.1".to_string()]).unwrap();
let second =
DnsOverrides::from_entries(&["example.com:443:127.0.0.2".to_string()]).unwrap();
assert_eq!(
first.resolve("example.com", 443),
Some("127.0.0.1".parse().unwrap())
);
assert_eq!(
second.resolve("example.com", 443),
Some("127.0.0.2".parse().unwrap())
);
}
#[test] #[test]
fn split_host_port_parses_supported_shapes() { fn split_host_port_parses_supported_shapes() {
assert_eq!( assert_eq!(
+9 -3
View File
@@ -246,7 +246,7 @@ pub fn secure_payload_len_from_wire_len(wire_len: usize) -> Option<usize> {
} }
/// Generate padding length for Secure Intermediate protocol. /// Generate padding length for Secure Intermediate protocol.
/// Telegram Desktop uses a 4-bit random padding length for VersionD packets. /// Outbound padding is 1..=3 so a receiver can strip it by 4-byte alignment.
pub fn secure_padding_len(data_len: usize, rng: &SecureRandom) -> usize { pub fn secure_padding_len(data_len: usize, rng: &SecureRandom) -> usize {
debug_assert!( debug_assert!(
is_valid_secure_payload_len(data_len), is_valid_secure_payload_len(data_len),
@@ -425,15 +425,21 @@ mod tests {
} }
#[test] #[test]
fn secure_padding_matches_tdesktop_range() { fn secure_padding_never_produces_aligned_total() {
let rng = SecureRandom::new(); let rng = SecureRandom::new();
for data_len in (0..1000).step_by(4) { for data_len in (0..1000).step_by(4) {
for _ in 0..100 { for _ in 0..100 {
let padding = secure_padding_len(data_len, &rng); let padding = secure_padding_len(data_len, &rng);
assert!( assert!(
padding <= 15, (1..=3).contains(&padding),
"padding out of range: data_len={data_len}, padding={padding}" "padding out of range: data_len={data_len}, padding={padding}"
); );
assert_ne!(
(data_len + padding) % 4,
0,
"invariant violated: data_len={data_len}, padding={padding}, total={}",
data_len + padding
);
} }
} }
} }
+4 -4
View File
@@ -8,8 +8,8 @@ pub(crate) const INTERMEDIATE_QUICKACK_FLAG: u32 = 0x8000_0000;
/// Payload length mask used by Intermediate and Secure Intermediate headers. /// Payload length mask used by Intermediate and Secure Intermediate headers.
pub(crate) const INTERMEDIATE_WIRE_LEN_MASK: u32 = 0x7fff_ffff; pub(crate) const INTERMEDIATE_WIRE_LEN_MASK: u32 = 0x7fff_ffff;
/// Maximum random tail length used by Telegram Desktop VersionD packets. /// Maximum outbound Secure tail length that keeps wire lengths non-aligned.
pub(crate) const SECURE_VERSION_D_PADDING_MAX: usize = 15; pub(crate) const SECURE_VERSION_D_PADDING_MAX: usize = 3;
/// Parsed Intermediate/Secure Intermediate length header. /// Parsed Intermediate/Secure Intermediate length header.
#[derive(Clone, Copy, Debug, PartialEq, Eq)] #[derive(Clone, Copy, Debug, PartialEq, Eq)]
@@ -51,9 +51,9 @@ pub(crate) fn secure_version_d_body_len_from_wire_len(wire_len: usize) -> Option
Some(wire_len - (wire_len % 4)) Some(wire_len - (wire_len % 4))
} }
/// Generate Telegram Desktop-compatible VersionD random tail length. /// Generate outbound Secure tail length without ambiguous full-word padding.
pub(crate) fn secure_version_d_padding_len(rng: &SecureRandom) -> usize { pub(crate) fn secure_version_d_padding_len(rng: &SecureRandom) -> usize {
rng.range(SECURE_VERSION_D_PADDING_MAX + 1) rng.range(SECURE_VERSION_D_PADDING_MAX) + 1
} }
#[cfg(test)] #[cfg(test)]
+137 -41
View File
@@ -1,7 +1,4 @@
#![allow(dead_code)] // Adaptive buffer policy shared by active Direct relay sessions.
// Adaptive buffer policy is staged and retained for deterministic rollout.
// Keep definitions compiled for compatibility and security test scaffolding.
use dashmap::DashMap; use dashmap::DashMap;
use std::cmp::max; use std::cmp::max;
@@ -13,29 +10,42 @@ const PROFILE_TTL: Duration = Duration::from_secs(300);
const THROUGHPUT_UP_BPS: f64 = 8_000_000.0; const THROUGHPUT_UP_BPS: f64 = 8_000_000.0;
const THROUGHPUT_DOWN_BPS: f64 = 2_000_000.0; const THROUGHPUT_DOWN_BPS: f64 = 2_000_000.0;
const RATIO_CONFIRM_THRESHOLD: f64 = 1.12; const RATIO_CONFIRM_THRESHOLD: f64 = 1.12;
const TIER1_HOLD_TICKS: u32 = 8; const TIER1_HOLD: Duration = Duration::from_secs(2);
const TIER2_HOLD_TICKS: u32 = 4; const TIER2_HOLD: Duration = Duration::from_secs(1);
const QUIET_DEMOTE_TICKS: u32 = 480; const QUIET_DEMOTE: Duration = Duration::from_secs(120);
const HARD_COOLDOWN_TICKS: u32 = 20; const HARD_COOLDOWN: Duration = Duration::from_secs(5);
const SUSTAINED_PRESSURE_DEMOTE: Duration = Duration::from_secs(30);
const PRESSURE_DEMOTE_COOLDOWN: Duration = Duration::from_secs(60);
const HARD_PENDING_THRESHOLD: u32 = 3; const HARD_PENDING_THRESHOLD: u32 = 3;
const HARD_PARTIAL_RATIO_THRESHOLD: f64 = 0.25; const HARD_PARTIAL_RATIO_THRESHOLD: f64 = 0.25;
#[cfg(test)]
const DIRECT_C2S_CAP_BYTES: usize = 128 * 1024; const DIRECT_C2S_CAP_BYTES: usize = 128 * 1024;
#[cfg(test)]
const DIRECT_S2C_CAP_BYTES: usize = 512 * 1024; const DIRECT_S2C_CAP_BYTES: usize = 512 * 1024;
#[cfg(test)]
const ME_FRAMES_CAP: usize = 96; const ME_FRAMES_CAP: usize = 96;
#[cfg(test)]
const ME_BYTES_CAP: usize = 384 * 1024; const ME_BYTES_CAP: usize = 384 * 1024;
#[cfg(test)]
const ME_DELAY_MIN_US: u64 = 150; const ME_DELAY_MIN_US: u64 = 150;
const MAX_USER_PROFILES_ENTRIES: usize = 50_000; const MAX_USER_PROFILES_ENTRIES: usize = 50_000;
const MAX_USER_KEY_BYTES: usize = 512; const MAX_USER_KEY_BYTES: usize = 512;
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)] #[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
/// Per-session Direct copy-buffer capacity tier.
pub enum AdaptiveTier { pub enum AdaptiveTier {
/// Conservative baseline capacity.
Base = 0, Base = 0,
/// First throughput promotion.
Tier1 = 1, Tier1 = 1,
/// Sustained bidirectional pressure promotion.
Tier2 = 2, Tier2 = 2,
/// Configured per-direction ceilings.
Tier3 = 3, Tier3 = 3,
} }
impl AdaptiveTier { impl AdaptiveTier {
/// Returns the next larger tier, saturating at `Tier3`.
pub fn promote(self) -> Self { pub fn promote(self) -> Self {
match self { match self {
Self::Base => Self::Tier1, Self::Base => Self::Tier1,
@@ -45,6 +55,7 @@ impl AdaptiveTier {
} }
} }
/// Returns the next smaller tier, saturating at `Base`.
pub fn demote(self) -> Self { pub fn demote(self) -> Self {
match self { match self {
Self::Base => Self::Base, Self::Base => Self::Base,
@@ -54,6 +65,7 @@ impl AdaptiveTier {
} }
} }
#[cfg(test)]
fn ratio(self) -> (usize, usize) { fn ratio(self) -> (usize, usize) {
match self { match self {
Self::Base => (1, 1), Self::Base => (1, 1),
@@ -63,49 +75,70 @@ impl AdaptiveTier {
} }
} }
/// Returns the stable numeric tier used by bounded metrics.
pub fn as_u8(self) -> u8 { pub fn as_u8(self) -> u8 {
self as u8 self as u8
} }
} }
#[derive(Debug, Clone, Copy, PartialEq, Eq)] #[derive(Debug, Clone, Copy, PartialEq, Eq)]
/// Signal that caused an accepted adaptive tier transition.
pub enum TierTransitionReason { pub enum TierTransitionReason {
/// Sustained throughput and directional ratio confirmation.
SoftConfirmed, SoftConfirmed,
/// Short pending or partial-write pressure burst.
HardPressure, HardPressure,
/// Sustained low-throughput period.
QuietDemotion, QuietDemotion,
/// Sustained pending or partial-write pressure.
SustainedWritePressure,
} }
#[derive(Debug, Clone, Copy, PartialEq, Eq)] #[derive(Debug, Clone, Copy, PartialEq, Eq)]
/// Proposed transition emitted by the per-session controller.
pub struct TierTransition { pub struct TierTransition {
/// Tier active before the observation.
pub from: AdaptiveTier, pub from: AdaptiveTier,
/// Tier requested after the observation.
pub to: AdaptiveTier, pub to: AdaptiveTier,
/// Pressure or throughput condition that requested the transition.
pub reason: TierTransitionReason, pub reason: TierTransitionReason,
} }
#[derive(Debug, Clone, Copy, Default)] #[derive(Debug, Clone, Copy, Default)]
/// Directional byte and write-pressure deltas for one observation period.
pub struct RelaySignalSample { pub struct RelaySignalSample {
/// Client-to-DC bytes copied during the period.
pub c2s_bytes: u64, pub c2s_bytes: u64,
/// Bytes offered to DC-to-client writes during the period.
pub s2c_requested_bytes: u64, pub s2c_requested_bytes: u64,
/// Bytes accepted by DC-to-client writes during the period.
pub s2c_written_bytes: u64, pub s2c_written_bytes: u64,
/// Successful DC-to-client write operations during the period.
pub s2c_write_ops: u64, pub s2c_write_ops: u64,
/// Partial DC-to-client write operations during the period.
pub s2c_partial_writes: u64, pub s2c_partial_writes: u64,
/// Consecutive pending DC-to-client writes at sample time.
pub s2c_consecutive_pending_writes: u32, pub s2c_consecutive_pending_writes: u32,
} }
#[derive(Debug, Clone, Copy)] #[derive(Debug, Clone, Copy)]
/// Stateful hysteresis controller for one active Direct session.
pub struct SessionAdaptiveController { pub struct SessionAdaptiveController {
tier: AdaptiveTier, tier: AdaptiveTier,
max_tier_seen: AdaptiveTier, max_tier_seen: AdaptiveTier,
throughput_ema_bps: f64, throughput_ema_bps: f64,
incoming_ema_bps: f64, incoming_ema_bps: f64,
outgoing_ema_bps: f64, outgoing_ema_bps: f64,
tier1_hold_ticks: u32, tier1_hold: Duration,
tier2_hold_ticks: u32, tier2_hold: Duration,
quiet_ticks: u32, quiet: Duration,
hard_cooldown_ticks: u32, hard_cooldown: Duration,
sustained_pressure: Duration,
} }
impl SessionAdaptiveController { impl SessionAdaptiveController {
/// Creates a controller at the tier whose memory reservation was accepted.
pub fn new(initial_tier: AdaptiveTier) -> Self { pub fn new(initial_tier: AdaptiveTier) -> Self {
Self { Self {
tier: initial_tier, tier: initial_tier,
@@ -113,25 +146,33 @@ impl SessionAdaptiveController {
throughput_ema_bps: 0.0, throughput_ema_bps: 0.0,
incoming_ema_bps: 0.0, incoming_ema_bps: 0.0,
outgoing_ema_bps: 0.0, outgoing_ema_bps: 0.0,
tier1_hold_ticks: 0, tier1_hold: Duration::ZERO,
tier2_hold_ticks: 0, tier2_hold: Duration::ZERO,
quiet_ticks: 0, quiet: Duration::ZERO,
hard_cooldown_ticks: 0, hard_cooldown: Duration::ZERO,
sustained_pressure: Duration::ZERO,
} }
} }
/// Returns the highest tier proposed by controller observations.
#[allow(dead_code)]
pub fn max_tier_seen(&self) -> AdaptiveTier { pub fn max_tier_seen(&self) -> AdaptiveTier {
self.max_tier_seen self.max_tier_seen
} }
/// Returns the controller's current logical tier.
pub fn tier(&self) -> AdaptiveTier {
self.tier
}
/// Observes one period and returns at most one hysteresis-controlled transition.
pub fn observe(&mut self, sample: RelaySignalSample, tick_secs: f64) -> Option<TierTransition> { pub fn observe(&mut self, sample: RelaySignalSample, tick_secs: f64) -> Option<TierTransition> {
if tick_secs <= f64::EPSILON { if tick_secs <= f64::EPSILON {
return None; return None;
} }
if self.hard_cooldown_ticks > 0 { let tick = Duration::from_secs_f64(tick_secs);
self.hard_cooldown_ticks -= 1; self.hard_cooldown = self.hard_cooldown.saturating_sub(tick);
}
let c2s_bps = (sample.c2s_bytes as f64 * 8.0) / tick_secs; let c2s_bps = (sample.c2s_bytes as f64 * 8.0) / tick_secs;
let incoming_bps = (sample.s2c_requested_bytes as f64 * 8.0) / tick_secs; let incoming_bps = (sample.s2c_requested_bytes as f64 * 8.0) / tick_secs;
@@ -144,9 +185,9 @@ impl SessionAdaptiveController {
let tier1_now = self.throughput_ema_bps >= THROUGHPUT_UP_BPS; let tier1_now = self.throughput_ema_bps >= THROUGHPUT_UP_BPS;
if tier1_now { if tier1_now {
self.tier1_hold_ticks = self.tier1_hold_ticks.saturating_add(1); self.tier1_hold = self.tier1_hold.saturating_add(tick);
} else { } else {
self.tier1_hold_ticks = 0; self.tier1_hold = Duration::ZERO;
} }
let ratio = if self.outgoing_ema_bps <= f64::EPSILON { let ratio = if self.outgoing_ema_bps <= f64::EPSILON {
@@ -156,9 +197,9 @@ impl SessionAdaptiveController {
}; };
let tier2_now = ratio >= RATIO_CONFIRM_THRESHOLD; let tier2_now = ratio >= RATIO_CONFIRM_THRESHOLD;
if tier2_now { if tier2_now {
self.tier2_hold_ticks = self.tier2_hold_ticks.saturating_add(1); self.tier2_hold = self.tier2_hold.saturating_add(tick);
} else { } else {
self.tier2_hold_ticks = 0; self.tier2_hold = Duration::ZERO;
} }
let partial_ratio = if sample.s2c_write_ops == 0 { let partial_ratio = if sample.s2c_write_ops == 0 {
@@ -169,24 +210,37 @@ impl SessionAdaptiveController {
let hard_now = sample.s2c_consecutive_pending_writes >= HARD_PENDING_THRESHOLD let hard_now = sample.s2c_consecutive_pending_writes >= HARD_PENDING_THRESHOLD
|| partial_ratio >= HARD_PARTIAL_RATIO_THRESHOLD; || partial_ratio >= HARD_PARTIAL_RATIO_THRESHOLD;
if hard_now && self.hard_cooldown_ticks == 0 { if hard_now {
return self.promote(TierTransitionReason::HardPressure, HARD_COOLDOWN_TICKS); self.sustained_pressure = self.sustained_pressure.saturating_add(tick);
if self.sustained_pressure >= SUSTAINED_PRESSURE_DEMOTE {
self.sustained_pressure = Duration::ZERO;
return self.demote(
TierTransitionReason::SustainedWritePressure,
PRESSURE_DEMOTE_COOLDOWN,
);
}
} else {
self.sustained_pressure = Duration::ZERO;
} }
if self.tier1_hold_ticks >= TIER1_HOLD_TICKS && self.tier2_hold_ticks >= TIER2_HOLD_TICKS { if hard_now && self.hard_cooldown.is_zero() {
return self.promote(TierTransitionReason::SoftConfirmed, 0); return self.promote(TierTransitionReason::HardPressure, HARD_COOLDOWN);
}
if self.tier1_hold >= TIER1_HOLD && self.tier2_hold >= TIER2_HOLD {
return self.promote(TierTransitionReason::SoftConfirmed, Duration::ZERO);
} }
let demote_candidate = let demote_candidate =
self.throughput_ema_bps < THROUGHPUT_DOWN_BPS && !tier2_now && !hard_now; self.throughput_ema_bps < THROUGHPUT_DOWN_BPS && !tier2_now && !hard_now;
if demote_candidate { if demote_candidate {
self.quiet_ticks = self.quiet_ticks.saturating_add(1); self.quiet = self.quiet.saturating_add(tick);
if self.quiet_ticks >= QUIET_DEMOTE_TICKS { if self.quiet >= QUIET_DEMOTE {
self.quiet_ticks = 0; self.quiet = Duration::ZERO;
return self.demote(TierTransitionReason::QuietDemotion); return self.demote(TierTransitionReason::QuietDemotion, Duration::ZERO);
} }
} else { } else {
self.quiet_ticks = 0; self.quiet = Duration::ZERO;
} }
None None
@@ -195,7 +249,7 @@ impl SessionAdaptiveController {
fn promote( fn promote(
&mut self, &mut self,
reason: TierTransitionReason, reason: TierTransitionReason,
hard_cooldown_ticks: u32, hard_cooldown: Duration,
) -> Option<TierTransition> { ) -> Option<TierTransition> {
let from = self.tier; let from = self.tier;
let to = from.promote(); let to = from.promote();
@@ -204,22 +258,27 @@ impl SessionAdaptiveController {
} }
self.tier = to; self.tier = to;
self.max_tier_seen = max(self.max_tier_seen, to); self.max_tier_seen = max(self.max_tier_seen, to);
self.hard_cooldown_ticks = hard_cooldown_ticks; self.hard_cooldown = hard_cooldown;
self.tier1_hold_ticks = 0; self.tier1_hold = Duration::ZERO;
self.tier2_hold_ticks = 0; self.tier2_hold = Duration::ZERO;
self.quiet_ticks = 0; self.quiet = Duration::ZERO;
Some(TierTransition { from, to, reason }) Some(TierTransition { from, to, reason })
} }
fn demote(&mut self, reason: TierTransitionReason) -> Option<TierTransition> { fn demote(
&mut self,
reason: TierTransitionReason,
hard_cooldown: Duration,
) -> Option<TierTransition> {
let from = self.tier; let from = self.tier;
let to = from.demote(); let to = from.demote();
if from == to { if from == to {
return None; return None;
} }
self.tier = to; self.tier = to;
self.tier1_hold_ticks = 0; self.hard_cooldown = hard_cooldown;
self.tier2_hold_ticks = 0; self.tier1_hold = Duration::ZERO;
self.tier2_hold = Duration::ZERO;
Some(TierTransition { from, to, reason }) Some(TierTransition { from, to, reason })
} }
} }
@@ -235,6 +294,8 @@ fn profiles() -> &'static DashMap<String, UserAdaptiveProfile> {
USER_PROFILES.get_or_init(DashMap::new) USER_PROFILES.get_or_init(DashMap::new)
} }
/// Returns a fresh user's recent successful Direct tier, or `Base` when stale.
#[allow(dead_code)]
pub fn seed_tier_for_user(user: &str) -> AdaptiveTier { pub fn seed_tier_for_user(user: &str) -> AdaptiveTier {
if user.len() > MAX_USER_KEY_BYTES { if user.len() > MAX_USER_KEY_BYTES {
return AdaptiveTier::Base; return AdaptiveTier::Base;
@@ -253,6 +314,8 @@ pub fn seed_tier_for_user(user: &str) -> AdaptiveTier {
AdaptiveTier::Base AdaptiveTier::Base
} }
/// Records the highest successfully allocated tier for bounded session seeding.
#[allow(dead_code)]
pub fn record_user_tier(user: &str, tier: AdaptiveTier) { pub fn record_user_tier(user: &str, tier: AdaptiveTier) {
if user.len() > MAX_USER_KEY_BYTES { if user.len() > MAX_USER_KEY_BYTES {
return; return;
@@ -282,6 +345,8 @@ pub fn record_user_tier(user: &str, tier: AdaptiveTier) {
} }
} }
#[cfg(test)]
/// Returns the legacy staged scaling policy retained by security fixtures.
pub fn direct_copy_buffers_for_tier( pub fn direct_copy_buffers_for_tier(
tier: AdaptiveTier, tier: AdaptiveTier,
base_c2s: usize, base_c2s: usize,
@@ -294,6 +359,32 @@ pub fn direct_copy_buffers_for_tier(
) )
} }
/// Maps an adaptive tier to independent capacities within configured ceilings.
pub(crate) fn direct_copy_buffers_for_tier_with_ceilings(
tier: AdaptiveTier,
base_c2s: usize,
base_s2c: usize,
ceiling_c2s: usize,
ceiling_s2c: usize,
) -> (usize, usize) {
(
direct_direction_size(tier, base_c2s, ceiling_c2s),
direct_direction_size(tier, base_s2c, ceiling_s2c),
)
}
fn direct_direction_size(tier: AdaptiveTier, base: usize, ceiling: usize) -> usize {
let target = match tier {
AdaptiveTier::Base => base,
AdaptiveTier::Tier1 => ceiling / 4,
AdaptiveTier::Tier2 => ceiling / 2,
AdaptiveTier::Tier3 => ceiling,
};
target.max(base).min(ceiling.max(base)).max(1)
}
#[cfg(test)]
/// Returns the staged Middle-End flush policy retained by security fixtures.
pub fn me_flush_policy_for_tier( pub fn me_flush_policy_for_tier(
tier: AdaptiveTier, tier: AdaptiveTier,
base_frames: usize, base_frames: usize,
@@ -323,6 +414,7 @@ fn ema(prev: f64, value: f64) -> f64 {
} }
} }
#[cfg(test)]
fn scale(base: usize, numerator: usize, denominator: usize, cap: usize) -> usize { fn scale(base: usize, numerator: usize, denominator: usize, cap: usize) -> usize {
let scaled = base let scaled = base
.saturating_mul(numerator) .saturating_mul(numerator)
@@ -338,6 +430,10 @@ mod adaptive_buffers_security_tests;
#[path = "tests/adaptive_buffers_record_race_security_tests.rs"] #[path = "tests/adaptive_buffers_record_race_security_tests.rs"]
mod adaptive_buffers_record_race_security_tests; mod adaptive_buffers_record_race_security_tests;
#[cfg(test)]
#[path = "tests/adaptive_direct_budget_policy_tests.rs"]
mod adaptive_direct_budget_policy_tests;
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::*; use super::*;
@@ -396,7 +492,7 @@ mod tests {
fn test_quiet_demotion_is_slow_and_stepwise() { fn test_quiet_demotion_is_slow_and_stepwise() {
let mut ctrl = SessionAdaptiveController::new(AdaptiveTier::Tier2); let mut ctrl = SessionAdaptiveController::new(AdaptiveTier::Tier2);
let mut demotion = None; let mut demotion = None;
for _ in 0..QUIET_DEMOTE_TICKS { for _ in 0..480 {
demotion = ctrl.observe(sample(1, 1, 1, 1, 0, 0), 0.25); demotion = ctrl.observe(sample(1, 1, 1, 1, 0, 0), 0.25);
} }
+568
View File
@@ -0,0 +1,568 @@
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::Duration;
use tokio::sync::watch;
use crate::stats::Stats;
use crate::stream::BufferPool;
use super::shared_state::ProxySharedState;
/// Accounting granularity for process-wide Direct copy-buffer reservations.
pub(crate) const DIRECT_BUFFER_UNIT_BYTES: usize = 4 * 1024;
/// Minimum client-to-DC copy-buffer capacity for one Direct session.
pub(crate) const DIRECT_BASE_C2S_BYTES: usize = 4 * 1024;
/// Minimum DC-to-client copy-buffer capacity for one Direct session.
pub(crate) const DIRECT_BASE_S2C_BYTES: usize = 8 * 1024;
const AUTO_HARD_MIN_BYTES: usize = 64 * 1024 * 1024;
const AUTO_HARD_MAX_BYTES: usize = 2 * 1024 * 1024 * 1024;
const AUTO_HARD_FALLBACK_BYTES: usize = 512 * 1024 * 1024;
const TARGET_FLOOR_MIN_BYTES: usize = 16 * 1024 * 1024;
const CONTROL_INTERVAL: Duration = Duration::from_secs(1);
const HEALTHY_RECOVERY_SAMPLES: u8 = 30;
const BUFFER_POOL_TRIM_LOW_WATERMARK: usize = 64;
const BUFFER_POOL_TRIM_HIGH_WATERMARK: usize = 128;
#[derive(Debug, Clone, Copy, Default)]
/// Lock-free observability snapshot of the Direct copy-buffer envelope.
pub(crate) struct DirectBufferBudgetSnapshot {
/// Absolute process-wide copy-buffer ceiling.
pub(crate) hard_limit_bytes: u64,
/// Current pressure-adjusted promotion target.
pub(crate) target_bytes: u64,
/// Bytes currently covered by active session leases.
pub(crate) reserved_bytes: u64,
/// Effective host or cgroup memory limit.
pub(crate) memory_total_bytes: u64,
/// Effective host or cgroup memory headroom.
pub(crate) memory_available_bytes: u64,
/// Current process resident set size.
pub(crate) process_rss_bytes: u64,
/// Successful tier growth reservations.
pub(crate) promotion_total: u64,
/// Tier growth attempts rejected by the adaptive target.
pub(crate) promotion_denied_total: u64,
/// Sessions admitted at minimum size above the adaptive target.
pub(crate) minimum_fallback_total: u64,
/// Sessions rejected by the absolute ceiling.
pub(crate) admission_rejected_total: u64,
/// Quiet-period tier reductions.
pub(crate) quiet_demotion_total: u64,
/// Sustained write-pressure tier reductions.
pub(crate) write_pressure_demotion_total: u64,
/// Process-wide pressure tier reductions.
pub(crate) global_pressure_demotion_total: u64,
/// Current sessions for Base through Tier3.
pub(crate) tier_sessions: [u64; 4],
}
#[derive(Debug, Clone, Copy, Default)]
struct SystemMemorySample {
total_bytes: u64,
available_bytes: u64,
process_rss_bytes: u64,
}
/// Process-wide hard envelope and adaptive target for Direct copy buffers.
pub(crate) struct DirectBufferBudget {
hard_limit_bytes: u64,
target_bytes: AtomicU64,
reserved_bytes: AtomicU64,
pressure_generation: AtomicU64,
pressure_tx: watch::Sender<u64>,
memory_total_bytes: AtomicU64,
memory_available_bytes: AtomicU64,
process_rss_bytes: AtomicU64,
promotion_total: AtomicU64,
promotion_denied_total: AtomicU64,
minimum_fallback_total: AtomicU64,
admission_rejected_total: AtomicU64,
quiet_demotion_total: AtomicU64,
write_pressure_demotion_total: AtomicU64,
global_pressure_demotion_total: AtomicU64,
tier_sessions: [AtomicU64; 4],
}
impl DirectBufferBudget {
/// Creates an envelope with a fixed absolute ceiling.
pub(crate) fn new(hard_limit_bytes: usize) -> Arc<Self> {
let hard_limit_bytes = align_down(hard_limit_bytes.max(DIRECT_BUFFER_UNIT_BYTES)) as u64;
let (pressure_tx, _) = watch::channel(0);
Arc::new(Self {
hard_limit_bytes,
target_bytes: AtomicU64::new(hard_limit_bytes),
reserved_bytes: AtomicU64::new(0),
pressure_generation: AtomicU64::new(0),
pressure_tx,
memory_total_bytes: AtomicU64::new(0),
memory_available_bytes: AtomicU64::new(0),
process_rss_bytes: AtomicU64::new(0),
promotion_total: AtomicU64::new(0),
promotion_denied_total: AtomicU64::new(0),
minimum_fallback_total: AtomicU64::new(0),
admission_rejected_total: AtomicU64::new(0),
quiet_demotion_total: AtomicU64::new(0),
write_pressure_demotion_total: AtomicU64::new(0),
global_pressure_demotion_total: AtomicU64::new(0),
tier_sessions: std::array::from_fn(|_| AtomicU64::new(0)),
})
}
/// Returns the current pressure-adjusted reservation target.
pub(crate) fn target_bytes(&self) -> usize {
self.target_bytes.load(Ordering::Relaxed) as usize
}
/// Subscribes to target reductions that require prompt session demotion.
pub(crate) fn subscribe_pressure(&self) -> watch::Receiver<u64> {
self.pressure_tx.subscribe()
}
/// Reserves bytes against either the adaptive target or the absolute ceiling.
pub(crate) fn try_reserve(
self: &Arc<Self>,
bytes: usize,
allow_above_target: bool,
) -> Option<DirectBufferLease> {
let bytes = align_up(bytes) as u64;
let limit = if allow_above_target {
self.hard_limit_bytes
} else {
self.target_bytes
.load(Ordering::Relaxed)
.min(self.hard_limit_bytes)
};
if !self.try_add_reserved(bytes, limit) {
return None;
}
self.tier_sessions[0].fetch_add(1, Ordering::Relaxed);
Some(DirectBufferLease {
budget: Arc::clone(self),
reserved_bytes: bytes,
tier: 0,
})
}
fn try_add_reserved(&self, bytes: u64, limit: u64) -> bool {
let mut current = self.reserved_bytes.load(Ordering::Acquire);
loop {
if bytes > limit.saturating_sub(current) {
return false;
}
match self.reserved_bytes.compare_exchange_weak(
current,
current + bytes,
Ordering::AcqRel,
Ordering::Acquire,
) {
Ok(_) => return true,
Err(observed) => current = observed,
}
}
}
fn target_floor_bytes(&self) -> u64 {
(self.hard_limit_bytes / 8)
.max(TARGET_FLOOR_MIN_BYTES as u64)
.min(self.hard_limit_bytes)
}
fn set_target_bytes(&self, target: u64) {
let target =
align_down(target.clamp(self.target_floor_bytes(), self.hard_limit_bytes) as usize)
as u64;
let previous = self.target_bytes.swap(target, Ordering::AcqRel);
if target < previous {
let generation = self
.pressure_generation
.fetch_add(1, Ordering::AcqRel)
.wrapping_add(1);
self.pressure_tx.send_replace(generation);
}
}
fn update_system_sample(&self, sample: SystemMemorySample) {
self.memory_total_bytes
.store(sample.total_bytes, Ordering::Relaxed);
self.memory_available_bytes
.store(sample.available_bytes, Ordering::Relaxed);
self.process_rss_bytes
.store(sample.process_rss_bytes, Ordering::Relaxed);
}
/// Records a session that had to bypass the adaptive target at minimum size.
pub(crate) fn increment_minimum_fallback(&self) {
self.minimum_fallback_total.fetch_add(1, Ordering::Relaxed);
}
/// Records a session rejected because the absolute ceiling was exhausted.
pub(crate) fn increment_admission_rejected(&self) {
self.admission_rejected_total
.fetch_add(1, Ordering::Relaxed);
}
/// Records a tier reduction after sustained low throughput.
pub(crate) fn increment_quiet_demotion(&self) {
self.quiet_demotion_total.fetch_add(1, Ordering::Relaxed);
}
/// Records a tier reduction after sustained partial or pending writes.
pub(crate) fn increment_write_pressure_demotion(&self) {
self.write_pressure_demotion_total
.fetch_add(1, Ordering::Relaxed);
}
/// Records a tier reduction requested by the process-wide controller.
pub(crate) fn increment_global_pressure_demotion(&self) {
self.global_pressure_demotion_total
.fetch_add(1, Ordering::Relaxed);
}
/// Captures all bounded metrics without allocating or locking.
pub(crate) fn snapshot(&self) -> DirectBufferBudgetSnapshot {
DirectBufferBudgetSnapshot {
hard_limit_bytes: self.hard_limit_bytes,
target_bytes: self.target_bytes.load(Ordering::Relaxed),
reserved_bytes: self.reserved_bytes.load(Ordering::Relaxed),
memory_total_bytes: self.memory_total_bytes.load(Ordering::Relaxed),
memory_available_bytes: self.memory_available_bytes.load(Ordering::Relaxed),
process_rss_bytes: self.process_rss_bytes.load(Ordering::Relaxed),
promotion_total: self.promotion_total.load(Ordering::Relaxed),
promotion_denied_total: self.promotion_denied_total.load(Ordering::Relaxed),
minimum_fallback_total: self.minimum_fallback_total.load(Ordering::Relaxed),
admission_rejected_total: self.admission_rejected_total.load(Ordering::Relaxed),
quiet_demotion_total: self.quiet_demotion_total.load(Ordering::Relaxed),
write_pressure_demotion_total: self
.write_pressure_demotion_total
.load(Ordering::Relaxed),
global_pressure_demotion_total: self
.global_pressure_demotion_total
.load(Ordering::Relaxed),
tier_sessions: std::array::from_fn(|index| {
self.tier_sessions[index].load(Ordering::Relaxed)
}),
}
}
}
/// Returns the conservative ceiling used when memory discovery is unavailable.
pub(crate) fn fallback_direct_buffer_hard_limit() -> usize {
AUTO_HARD_FALLBACK_BYTES
}
/// RAII ownership of all copy-buffer bytes retained by one Direct session.
pub(crate) struct DirectBufferLease {
budget: Arc<DirectBufferBudget>,
reserved_bytes: u64,
tier: usize,
}
impl DirectBufferLease {
/// Returns the currently covered allocation rounded to accounting units.
pub(crate) fn reserved_bytes(&self) -> usize {
self.reserved_bytes as usize
}
/// Attempts to cover a larger tier before its buffers are resized.
pub(crate) fn try_grow_to(&mut self, bytes: usize) -> bool {
let bytes = align_up(bytes) as u64;
if bytes <= self.reserved_bytes {
return true;
}
let delta = bytes - self.reserved_bytes;
let limit = self
.budget
.target_bytes
.load(Ordering::Relaxed)
.min(self.budget.hard_limit_bytes);
if !self.budget.try_add_reserved(delta, limit) {
self.budget
.promotion_denied_total
.fetch_add(1, Ordering::Relaxed);
return false;
}
self.reserved_bytes = bytes;
self.budget.promotion_total.fetch_add(1, Ordering::Relaxed);
true
}
/// Releases bytes only after both directional buffers report smaller coverage.
pub(crate) fn shrink_to(&mut self, bytes: usize) {
let bytes = align_up(bytes) as u64;
if bytes >= self.reserved_bytes {
return;
}
let released = self.reserved_bytes - bytes;
self.reserved_bytes = bytes;
self.budget
.reserved_bytes
.fetch_sub(released, Ordering::AcqRel);
}
/// Updates bounded per-tier session gauges for an accepted transition.
pub(crate) fn set_tier(&mut self, tier: usize) {
let tier = tier.min(self.budget.tier_sessions.len() - 1);
if tier == self.tier {
return;
}
decrement_saturating(&self.budget.tier_sessions[self.tier]);
self.budget.tier_sessions[tier].fetch_add(1, Ordering::Relaxed);
self.tier = tier;
}
}
impl Drop for DirectBufferLease {
fn drop(&mut self) {
self.budget
.reserved_bytes
.fetch_sub(self.reserved_bytes, Ordering::AcqRel);
decrement_saturating(&self.budget.tier_sessions[self.tier]);
}
}
/// Resolves the startup hard ceiling from config, cgroup, and host memory.
pub(crate) async fn resolve_direct_buffer_hard_limit(configured: usize) -> usize {
if configured != 0 {
return align_down(configured);
}
let sample = read_system_memory_sample().await;
if sample.total_bytes == 0 {
return AUTO_HARD_FALLBACK_BYTES;
}
let derived = (sample.total_bytes / 4)
.clamp(AUTO_HARD_MIN_BYTES as u64, AUTO_HARD_MAX_BYTES as u64)
.min(sample.total_bytes);
align_down(derived as usize).max(DIRECT_BUFFER_UNIT_BYTES)
}
/// Runs the control-plane loop for Direct budget and shared pool pressure.
pub(crate) async fn run_direct_buffer_budget_controller(
budget: Arc<DirectBufferBudget>,
buffer_pool: Arc<BufferPool>,
stats: Arc<Stats>,
shared: Arc<ProxySharedState>,
max_connections: u32,
) {
let mut interval = tokio::time::interval(CONTROL_INTERVAL);
interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
let mut healthy_streak = 0u8;
let mut previous_denied = 0u64;
let mut previous_fallback = 0u64;
let mut previous_rejected = 0u64;
let pool_trim_low = buffer_pool
.max_buffers()
.min(BUFFER_POOL_TRIM_LOW_WATERMARK);
let pool_trim_high = buffer_pool
.max_buffers()
.min(BUFFER_POOL_TRIM_HIGH_WATERMARK);
let mut pool_trim_armed = true;
loop {
interval.tick().await;
let sample = read_system_memory_sample().await;
budget.update_system_sample(sample);
let snapshot = budget.snapshot();
let denied_delta = snapshot
.promotion_denied_total
.saturating_sub(previous_denied);
previous_denied = snapshot.promotion_denied_total;
let fallback_delta = snapshot
.minimum_fallback_total
.saturating_sub(previous_fallback);
previous_fallback = snapshot.minimum_fallback_total;
let rejected_delta = snapshot
.admission_rejected_total
.saturating_sub(previous_rejected);
previous_rejected = snapshot.admission_rejected_total;
let connection_pct = connection_fill_pct(stats.as_ref(), max_connections);
let memory_available_pct = percentage(sample.available_bytes, sample.total_bytes);
let target_utilization_pct = percentage(snapshot.reserved_bytes, snapshot.target_bytes);
let pressure = shared.conntrack_pressure_active()
|| connection_pct.is_some_and(|value| value >= 85)
|| memory_available_pct.is_some_and(|value| value <= 15)
|| target_utilization_pct.is_some_and(|value| value >= 90)
|| denied_delta > 0
|| fallback_delta > 0
|| rejected_delta > 0;
if !pressure {
pool_trim_armed = true;
} else if pool_trim_armed && buffer_pool.pooled() > pool_trim_high {
buffer_pool.trim_to(pool_trim_low);
pool_trim_armed = false;
}
let pool_snapshot = buffer_pool.stats();
stats.set_buffer_pool_gauges(
pool_snapshot.pooled,
pool_snapshot.allocated,
pool_snapshot.allocated.saturating_sub(pool_snapshot.pooled),
);
stats.set_buffer_pool_replaced_nonstandard_total(pool_snapshot.replaced_nonstandard);
let headroom_target = if sample.total_bytes == 0 {
snapshot.hard_limit_bytes
} else {
snapshot
.reserved_bytes
.saturating_add(sample.available_bytes / 4)
.min(snapshot.hard_limit_bytes)
};
if pressure {
healthy_streak = 0;
let reduced = snapshot.target_bytes.saturating_mul(3) / 4;
budget.set_target_bytes(reduced.min(headroom_target));
continue;
}
let healthy = memory_available_pct.is_none_or(|value| value >= 30)
&& connection_pct.is_none_or(|value| value <= 70);
if !healthy {
healthy_streak = 0;
if headroom_target < snapshot.target_bytes {
budget.set_target_bytes(headroom_target);
}
continue;
}
healthy_streak = healthy_streak.saturating_add(1);
if healthy_streak >= HEALTHY_RECOVERY_SAMPLES {
healthy_streak = 0;
let increment = (snapshot.target_bytes / 16).max(4 * 1024 * 1024);
budget.set_target_bytes(
snapshot
.target_bytes
.saturating_add(increment)
.min(headroom_target),
);
}
}
}
fn connection_fill_pct(stats: &Stats, max_connections: u32) -> Option<u8> {
if max_connections == 0 {
return None;
}
Some(
((stats.get_current_connections_total().saturating_mul(100)) / u64::from(max_connections))
.min(100) as u8,
)
}
fn percentage(value: u64, total: u64) -> Option<u8> {
if total == 0 {
return None;
}
Some(((value.saturating_mul(100)) / total).min(100) as u8)
}
async fn read_system_memory_sample() -> SystemMemorySample {
#[cfg(target_os = "linux")]
{
let meminfo = tokio::fs::read_to_string("/proc/meminfo")
.await
.unwrap_or_default();
let status = tokio::fs::read_to_string("/proc/self/status")
.await
.unwrap_or_default();
let host_total = parse_kib_field(&meminfo, "MemTotal:");
let host_available = parse_kib_field(&meminfo, "MemAvailable:");
let process_rss = parse_kib_field(&status, "VmRSS:");
let cgroup_v2_max = read_cgroup_limit("/sys/fs/cgroup/memory.max").await;
let cgroup_v2_current = read_u64_file("/sys/fs/cgroup/memory.current").await;
let cgroup_v1_max = read_cgroup_limit("/sys/fs/cgroup/memory/memory.limit_in_bytes").await;
let cgroup_v1_current = read_u64_file("/sys/fs/cgroup/memory/memory.usage_in_bytes").await;
let cgroup_max = cgroup_v2_max.or(cgroup_v1_max);
let cgroup_current = cgroup_v2_current.or(cgroup_v1_current);
let total = match (host_total, cgroup_max) {
(0, Some(limit)) => limit,
(host, Some(limit)) => host.min(limit),
(host, None) => host,
};
let cgroup_available = cgroup_max
.zip(cgroup_current)
.map(|(limit, current)| limit.saturating_sub(current));
let available = match (host_available, cgroup_available) {
(0, Some(value)) => value,
(host, Some(value)) => host.min(value),
(host, None) => host,
};
return SystemMemorySample {
total_bytes: total,
available_bytes: available,
process_rss_bytes: process_rss,
};
}
#[cfg(not(target_os = "linux"))]
{
SystemMemorySample::default()
}
}
#[cfg(target_os = "linux")]
async fn read_cgroup_limit(path: &str) -> Option<u64> {
let raw = tokio::fs::read_to_string(path).await.ok()?;
let raw = raw.trim();
if raw == "max" {
return None;
}
let value = raw.parse::<u64>().ok()?;
(value < (1u64 << 60)).then_some(value)
}
#[cfg(target_os = "linux")]
async fn read_u64_file(path: &str) -> Option<u64> {
tokio::fs::read_to_string(path)
.await
.ok()?
.trim()
.parse()
.ok()
}
#[cfg(target_os = "linux")]
fn parse_kib_field(raw: &str, key: &str) -> u64 {
raw.lines()
.find_map(|line| {
let value = line.strip_prefix(key)?.split_whitespace().next()?;
value.parse::<u64>().ok()
})
.unwrap_or(0)
.saturating_mul(1024)
}
fn align_up(bytes: usize) -> usize {
bytes
.div_ceil(DIRECT_BUFFER_UNIT_BYTES)
.saturating_mul(DIRECT_BUFFER_UNIT_BYTES)
}
fn align_down(bytes: usize) -> usize {
bytes / DIRECT_BUFFER_UNIT_BYTES * DIRECT_BUFFER_UNIT_BYTES
}
fn decrement_saturating(value: &AtomicU64) {
let mut current = value.load(Ordering::Relaxed);
while current != 0 {
match value.compare_exchange_weak(
current,
current - 1,
Ordering::Relaxed,
Ordering::Relaxed,
) {
Ok(_) => return,
Err(observed) => current = observed,
}
}
}
#[cfg(test)]
#[path = "tests/direct_buffer_budget_tests.rs"]
mod tests;
+3 -4
View File
@@ -345,21 +345,21 @@ where
} else { } else {
Duration::from_secs(1800) Duration::from_secs(1800)
}; };
let relay_result = let relay_result = crate::proxy::relay::relay_direct_adaptive(
crate::proxy::relay::relay_bidirectional_with_activity_timeout_lease_and_cancel(
client_reader, client_reader,
client_writer, client_writer,
tg_reader, tg_reader,
tg_writer, tg_writer,
config.general.direct_relay_copy_buf_c2s_bytes, config.general.direct_relay_copy_buf_c2s_bytes,
config.general.direct_relay_copy_buf_s2c_bytes, config.general.direct_relay_copy_buf_s2c_bytes,
config.server.max_connections,
user, user,
Arc::clone(&stats), Arc::clone(&stats),
config.access.user_data_quota.get(user).copied(), config.access.user_data_quota.get(user).copied(),
buffer_pool,
traffic_lease, traffic_lease,
relay_activity_timeout, relay_activity_timeout,
session_cancel.clone(), session_cancel.clone(),
Arc::clone(&shared.direct_buffer_budget),
); );
tokio::pin!(relay_result); tokio::pin!(relay_result);
let relay_result = loop { let relay_result = loop {
@@ -400,7 +400,6 @@ where
Err(e) => debug!(user = %user, error = %e, "Direct relay ended with error"), Err(e) => debug!(user = %user, error = %e, "Direct relay ended with error"),
} }
buffer_pool_trim.trim_to(buffer_pool_trim.max_buffers().min(64));
let pool_snapshot = buffer_pool_trim.stats(); let pool_snapshot = buffer_pool_trim.stats();
stats.set_buffer_pool_gauges( stats.set_buffer_pool_gauges(
pool_snapshot.pooled, pool_snapshot.pooled,
+1
View File
@@ -959,6 +959,7 @@ fn configure_mask_backend_socket(stream: &TcpStream) {
} }
/// Handles a bad client by forwarding it to the configured mask target. /// Handles a bad client by forwarding it to the configured mask target.
#[cfg(test)]
pub async fn handle_bad_client<R, W>( pub async fn handle_bad_client<R, W>(
reader: R, reader: R,
writer: W, writer: W,
+7 -1
View File
@@ -69,7 +69,9 @@ use self::quota::{
#[cfg(test)] #[cfg(test)]
use self::c2me::enqueue_c2me_command; use self::c2me::enqueue_c2me_command;
#[cfg(test)] #[cfg(test)]
use self::d2c::{compute_intermediate_secure_wire_len, process_me_writer_response}; use self::d2c::{
compute_intermediate_secure_wire_len, process_me_writer_response, write_client_payload,
};
#[cfg(test)] #[cfg(test)]
pub(crate) use self::desync::{ pub(crate) use self::desync::{
clear_desync_dedup_for_testing_in_shared, desync_dedup_get_for_testing, clear_desync_dedup_for_testing_in_shared, desync_dedup_get_for_testing,
@@ -166,3 +168,7 @@ mod middle_relay_atomic_quota_invariant_tests;
#[cfg(test)] #[cfg(test)]
#[path = "tests/middle_relay_baseline_invariant_tests.rs"] #[path = "tests/middle_relay_baseline_invariant_tests.rs"]
mod middle_relay_baseline_invariant_tests; mod middle_relay_baseline_invariant_tests;
#[cfg(test)]
#[path = "tests/middle_relay_d2c_flush_padding_security_tests.rs"]
mod middle_relay_d2c_flush_padding_security_tests;
+4 -4
View File
@@ -55,8 +55,8 @@ pub(super) fn classify_me_d2c_flush_reason(
MeD2cFlushReason::QueueDrain MeD2cFlushReason::QueueDrain
} }
pub(super) fn me_d2c_flush_reason_requires_client_flush(reason: MeD2cFlushReason) -> bool { pub(super) fn me_d2c_flush_reason_requires_client_flush(_reason: MeD2cFlushReason) -> bool {
!matches!(reason, MeD2cFlushReason::QueueDrain) true
} }
#[cfg(test)] #[cfg(test)]
@@ -64,8 +64,8 @@ mod tests {
use super::*; use super::*;
#[test] #[test]
fn queue_drain_is_not_a_physical_flush_trigger() { fn all_flush_reasons_trigger_physical_flush() {
assert!(!me_d2c_flush_reason_requires_client_flush( assert!(me_d2c_flush_reason_requires_client_flush(
MeD2cFlushReason::QueueDrain MeD2cFlushReason::QueueDrain
)); ));
assert!(me_d2c_flush_reason_requires_client_flush( assert!(me_d2c_flush_reason_requires_client_flush(
+1 -1
View File
@@ -149,6 +149,7 @@ where
peer, peer,
translated_local_addr, translated_local_addr,
payload, payload,
_permit,
flags, flags,
effective_tag_array, effective_tag_array,
) )
@@ -822,7 +823,6 @@ where
clear_relay_idle_candidate_in(shared.as_ref(), conn_id); clear_relay_idle_candidate_in(shared.as_ref(), conn_id);
me_pool.registry().unregister(conn_id).await; me_pool.registry().unregister(conn_id).await;
buffer_pool.trim_to(buffer_pool.max_buffers().min(64));
let pool_snapshot = buffer_pool.stats(); let pool_snapshot = buffer_pool.stats();
stats.set_buffer_pool_gauges( stats.set_buffer_pool_gauges(
pool_snapshot.pooled, pool_snapshot.pooled,
+2
View File
@@ -60,6 +60,8 @@
pub mod adaptive_buffers; pub mod adaptive_buffers;
pub mod client; pub mod client;
// Process-wide Direct relay copy-buffer ownership and pressure policy.
pub(crate) mod direct_buffer_budget;
pub mod direct_relay; pub mod direct_relay;
pub mod handshake; pub mod handshake;
pub mod masking; pub mod masking;
+4
View File
@@ -84,8 +84,11 @@ fn watchdog_delta(current: u64, previous: u64) -> u64 {
current.saturating_sub(previous) current.saturating_sub(previous)
} }
mod adaptive_copy;
mod io; mod io;
pub(crate) use self::adaptive_copy::relay_direct_adaptive;
use self::io::{CombinedStream, SharedCounters, StatsIo, is_quota_io_error}; use self::io::{CombinedStream, SharedCounters, StatsIo, is_quota_io_error};
#[cfg(test)] #[cfg(test)]
use self::io::{quota_adaptive_interval_bytes, should_immediate_quota_check}; use self::io::{quota_adaptive_interval_bytes, should_immediate_quota_check};
@@ -217,6 +220,7 @@ where
.await .await
} }
#[allow(dead_code)]
pub async fn relay_bidirectional_with_activity_timeout_lease_and_cancel<CR, CW, SR, SW>( pub async fn relay_bidirectional_with_activity_timeout_lease_and_cancel<CR, CW, SR, SW>(
client_reader: CR, client_reader: CR,
client_writer: CW, client_writer: CW,
+527
View File
@@ -0,0 +1,527 @@
use std::io;
use std::pin::Pin;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::task::{Context, Poll};
use std::time::Duration;
use tokio::io::{AsyncBufRead, AsyncRead, AsyncWrite, AsyncWriteExt, ReadBuf, copy_buf};
use tokio::time::Instant;
use tokio_util::sync::CancellationToken;
use tracing::{debug, warn};
use crate::error::{ProxyError, Result};
use crate::proxy::adaptive_buffers::{
AdaptiveTier, RelaySignalSample, SessionAdaptiveController, TierTransitionReason,
direct_copy_buffers_for_tier_with_ceilings,
};
use crate::proxy::direct_buffer_budget::{
DIRECT_BASE_C2S_BYTES, DIRECT_BASE_S2C_BYTES, DirectBufferBudget, DirectBufferLease,
};
use crate::proxy::traffic_limiter::TrafficLease;
use crate::stats::Stats;
use super::WATCHDOG_INTERVAL;
use super::io::{SharedCounters, StatsIo, is_quota_io_error};
use super::watchdog_delta;
mod write_pressure;
use self::write_pressure::WritePressureIo;
struct AdaptiveBufferState {
desired_bytes: AtomicUsize,
actual_bytes: AtomicUsize,
}
impl AdaptiveBufferState {
fn new(bytes: usize) -> Arc<Self> {
Arc::new(Self {
desired_bytes: AtomicUsize::new(bytes.max(1)),
actual_bytes: AtomicUsize::new(bytes.max(1)),
})
}
}
struct AdaptiveBufReader<R> {
inner: R,
buffer: Box<[u8]>,
pos: usize,
cap: usize,
state: Arc<AdaptiveBufferState>,
}
impl<R> AdaptiveBufReader<R> {
fn new(inner: R, state: Arc<AdaptiveBufferState>) -> Self {
let bytes = state.actual_bytes.load(Ordering::Relaxed).max(1);
Self {
inner,
buffer: vec![0; bytes].into_boxed_slice(),
pos: 0,
cap: 0,
state,
}
}
fn resize_if_drained(&mut self) {
if self.pos != self.cap {
return;
}
let desired = self.state.desired_bytes.load(Ordering::Acquire).max(1);
if desired == self.buffer.len() {
return;
}
self.buffer = vec![0; desired].into_boxed_slice();
self.pos = 0;
self.cap = 0;
self.state.actual_bytes.store(desired, Ordering::Release);
}
}
impl<R: AsyncRead + Unpin> AsyncRead for AdaptiveBufReader<R> {
fn poll_read(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
output: &mut ReadBuf<'_>,
) -> Poll<io::Result<()>> {
let this = self.get_mut();
if this.pos < this.cap {
let available = &this.buffer[this.pos..this.cap];
let copied = available.len().min(output.remaining());
output.put_slice(&available[..copied]);
this.pos += copied;
return Poll::Ready(Ok(()));
}
this.resize_if_drained();
Pin::new(&mut this.inner).poll_read(cx, output)
}
}
impl<R: AsyncRead + Unpin> AsyncBufRead for AdaptiveBufReader<R> {
fn poll_fill_buf(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<&[u8]>> {
let this = self.get_mut();
if this.pos < this.cap {
return Poll::Ready(Ok(&this.buffer[this.pos..this.cap]));
}
this.resize_if_drained();
let mut read_buf = ReadBuf::new(&mut this.buffer);
match Pin::new(&mut this.inner).poll_read(cx, &mut read_buf) {
Poll::Ready(Ok(())) => {
this.pos = 0;
this.cap = read_buf.filled().len();
Poll::Ready(Ok(&this.buffer[..this.cap]))
}
Poll::Ready(Err(error)) => Poll::Ready(Err(error)),
Poll::Pending => Poll::Pending,
}
}
fn consume(self: Pin<&mut Self>, amount: usize) {
let this = self.get_mut();
this.pos = this.pos.saturating_add(amount).min(this.cap);
}
}
enum AdaptiveRelayOutcome {
Copy(io::Result<(u64, u64)>),
ActivityTimeout,
UserDisabled,
}
#[allow(clippy::too_many_arguments)]
/// Relays one Direct session with independently resizable directional buffers.
pub(crate) async fn relay_direct_adaptive<CR, CW, SR, SW>(
client_reader: CR,
client_writer: CW,
server_reader: SR,
server_writer: SW,
ceiling_c2s_bytes: usize,
ceiling_s2c_bytes: usize,
max_connections: u32,
user: &str,
stats: Arc<Stats>,
quota_limit: Option<u64>,
traffic_lease: Option<Arc<TrafficLease>>,
activity_timeout: Duration,
session_cancel: CancellationToken,
budget: Arc<DirectBufferBudget>,
) -> Result<()>
where
CR: AsyncRead + Unpin + Send + 'static,
CW: AsyncWrite + Unpin + Send + 'static,
SR: AsyncRead + Unpin + Send + 'static,
SW: AsyncWrite + Unpin + Send + 'static,
{
let activity_timeout = activity_timeout.max(Duration::from_secs(1));
let epoch = Instant::now();
let counters = Arc::new(SharedCounters::new());
let quota_exceeded = Arc::new(AtomicBool::new(false));
let user_owned = user.to_string();
let (base_c2s, base_s2c) = initial_base_sizes(
ceiling_c2s_bytes,
ceiling_s2c_bytes,
max_connections,
budget.target_bytes(),
);
let base_total = base_c2s.saturating_add(base_s2c);
let mut lease = match budget.try_reserve(base_total, false) {
Some(lease) => lease,
None => {
let minimum_total = DIRECT_BASE_C2S_BYTES + DIRECT_BASE_S2C_BYTES;
match budget.try_reserve(minimum_total, true) {
Some(lease) => {
budget.increment_minimum_fallback();
lease
}
None => {
budget.increment_admission_rejected();
return Err(ProxyError::Proxy(
"Direct relay buffer pressure: budget exhausted".to_string(),
));
}
}
}
};
let effective_base = if lease.reserved_bytes() < base_total {
(DIRECT_BASE_C2S_BYTES, DIRECT_BASE_S2C_BYTES)
} else {
(base_c2s, base_s2c)
};
let c2s_state = AdaptiveBufferState::new(effective_base.0);
let s2c_state = AdaptiveBufferState::new(effective_base.1);
let mut controller = SessionAdaptiveController::new(AdaptiveTier::Base);
let c2s_client = StatsIo::new_with_traffic_lease(
client_reader,
Arc::clone(&counters),
Arc::clone(&stats),
user_owned.clone(),
traffic_lease.clone(),
quota_limit,
Arc::clone(&quota_exceeded),
epoch,
);
let client_writer = StatsIo::new_with_traffic_lease(
client_writer,
Arc::clone(&counters),
Arc::clone(&stats),
user_owned.clone(),
traffic_lease,
quota_limit,
Arc::clone(&quota_exceeded),
epoch,
);
let mut client_writer = WritePressureIo::new(client_writer, Arc::clone(&counters));
let mut c2s_reader = AdaptiveBufReader::new(c2s_client, Arc::clone(&c2s_state));
let mut s2c_reader = AdaptiveBufReader::new(server_reader, Arc::clone(&s2c_state));
let mut server_writer = server_writer;
let mut pressure_rx = budget.subscribe_pressure();
let relay_outcome = {
let copy = async {
let c2s = async {
let copied = copy_buf(&mut c2s_reader, &mut server_writer).await?;
server_writer.shutdown().await?;
Ok::<u64, io::Error>(copied)
};
let s2c = async {
let copied = copy_buf(&mut s2c_reader, &mut client_writer).await?;
client_writer.shutdown().await?;
Ok::<u64, io::Error>(copied)
};
tokio::try_join!(c2s, s2c)
};
tokio::pin!(copy);
let mut interval = tokio::time::interval(WATCHDOG_INTERVAL);
interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
interval.tick().await;
let mut previous = RelaySignalSample::default();
let mut previous_log_c2s = 0u64;
let mut previous_log_s2c = 0u64;
let mut previous_sample_at = epoch;
loop {
tokio::select! {
result = &mut copy => break AdaptiveRelayOutcome::Copy(result),
_ = session_cancel.cancelled() => break AdaptiveRelayOutcome::UserDisabled,
changed = pressure_rx.changed() => {
if changed.is_ok() {
apply_global_pressure_demotion(
&mut controller,
&mut lease,
&c2s_state,
&s2c_state,
effective_base,
(ceiling_c2s_bytes, ceiling_s2c_bytes),
budget.as_ref(),
);
reconcile_reservation(&mut lease, &c2s_state, &s2c_state);
}
}
_ = interval.tick() => {
let now = Instant::now();
let idle = counters.idle_duration(now, epoch);
if quota_exceeded.load(Ordering::Acquire) {
warn!(user = %user_owned, "User data quota reached, closing relay");
break AdaptiveRelayOutcome::ActivityTimeout;
}
if idle >= activity_timeout {
warn!(
user = %user_owned,
c2s_bytes = counters.c2s_bytes.load(Ordering::Relaxed),
s2c_bytes = counters.s2c_bytes.load(Ordering::Relaxed),
idle_secs = idle.as_secs(),
"Activity timeout"
);
break AdaptiveRelayOutcome::ActivityTimeout;
}
let sample = current_sample(counters.as_ref());
let c2s_delta = watchdog_delta(sample.c2s_bytes, previous_log_c2s);
let s2c_delta = watchdog_delta(sample.s2c_written_bytes, previous_log_s2c);
if c2s_delta > 0 || s2c_delta > 0 {
let secs = now.saturating_duration_since(previous_sample_at).as_secs_f64();
debug!(
user = %user_owned,
c2s_kbps = (c2s_delta as f64 / secs / 1024.0) as u64,
s2c_kbps = (s2c_delta as f64 / secs / 1024.0) as u64,
c2s_total = sample.c2s_bytes,
s2c_total = sample.s2c_written_bytes,
"Relay active"
);
}
let delta = sample_delta(sample, previous);
let tick_secs = now.saturating_duration_since(previous_sample_at).as_secs_f64();
if let Some(transition) = controller.observe(delta, tick_secs) {
apply_controller_transition(
transition,
&mut controller,
&mut lease,
&c2s_state,
&s2c_state,
effective_base,
(ceiling_c2s_bytes, ceiling_s2c_bytes),
budget.as_ref(),
);
}
reconcile_reservation(&mut lease, &c2s_state, &s2c_state);
previous = sample;
previous_log_c2s = sample.c2s_bytes;
previous_log_s2c = sample.s2c_written_bytes;
previous_sample_at = now;
}
}
}
};
let _ = client_writer.shutdown().await;
let _ = server_writer.shutdown().await;
let c2s_ops = counters.c2s_ops.load(Ordering::Relaxed);
let s2c_ops = counters.s2c_ops.load(Ordering::Relaxed);
let duration = epoch.elapsed();
match relay_outcome {
AdaptiveRelayOutcome::Copy(Ok((c2s, s2c))) => {
debug!(
user = %user_owned,
c2s_bytes = c2s,
s2c_bytes = s2c,
c2s_msgs = c2s_ops,
s2c_msgs = s2c_ops,
duration_secs = duration.as_secs(),
"Relay finished"
);
Ok(())
}
AdaptiveRelayOutcome::Copy(Err(error)) if is_quota_io_error(&error) => {
warn!(
user = %user_owned,
c2s_bytes = counters.c2s_bytes.load(Ordering::Relaxed),
s2c_bytes = counters.s2c_bytes.load(Ordering::Relaxed),
c2s_msgs = c2s_ops,
s2c_msgs = s2c_ops,
duration_secs = duration.as_secs(),
"Data quota reached, closing relay"
);
Err(ProxyError::DataQuotaExceeded { user: user_owned })
}
AdaptiveRelayOutcome::Copy(Err(error)) => {
debug!(
user = %user_owned,
c2s_bytes = counters.c2s_bytes.load(Ordering::Relaxed),
s2c_bytes = counters.s2c_bytes.load(Ordering::Relaxed),
c2s_msgs = c2s_ops,
s2c_msgs = s2c_ops,
duration_secs = duration.as_secs(),
error = %error,
"Relay error"
);
Err(error.into())
}
AdaptiveRelayOutcome::ActivityTimeout => {
debug!(
user = %user_owned,
c2s_bytes = counters.c2s_bytes.load(Ordering::Relaxed),
s2c_bytes = counters.s2c_bytes.load(Ordering::Relaxed),
c2s_msgs = c2s_ops,
s2c_msgs = s2c_ops,
duration_secs = duration.as_secs(),
"Relay finished (activity timeout)"
);
Ok(())
}
AdaptiveRelayOutcome::UserDisabled => {
debug!(
user = %user_owned,
c2s_bytes = counters.c2s_bytes.load(Ordering::Relaxed),
s2c_bytes = counters.s2c_bytes.load(Ordering::Relaxed),
c2s_msgs = c2s_ops,
s2c_msgs = s2c_ops,
duration_secs = duration.as_secs(),
"Relay finished (user disabled)"
);
Err(ProxyError::UserDisabled { user: user_owned })
}
}
}
fn initial_base_sizes(
ceiling_c2s: usize,
ceiling_s2c: usize,
max_connections: u32,
target_bytes: usize,
) -> (usize, usize) {
let configured_total = ceiling_c2s.saturating_add(ceiling_s2c);
let configured_worst_case = configured_total.saturating_mul(max_connections as usize);
if max_connections != 0 && configured_worst_case <= target_bytes {
return (ceiling_c2s, ceiling_s2c);
}
(
DIRECT_BASE_C2S_BYTES.min(ceiling_c2s),
DIRECT_BASE_S2C_BYTES.min(ceiling_s2c),
)
}
fn current_sample(counters: &SharedCounters) -> RelaySignalSample {
RelaySignalSample {
c2s_bytes: counters.c2s_bytes.load(Ordering::Relaxed),
s2c_requested_bytes: counters.s2c_requested_bytes.load(Ordering::Relaxed),
s2c_written_bytes: counters.s2c_bytes.load(Ordering::Relaxed),
s2c_write_ops: counters.s2c_ops.load(Ordering::Relaxed),
s2c_partial_writes: counters.s2c_partial_writes.load(Ordering::Relaxed),
s2c_consecutive_pending_writes: counters
.s2c_consecutive_pending_writes
.load(Ordering::Relaxed),
}
}
fn sample_delta(current: RelaySignalSample, previous: RelaySignalSample) -> RelaySignalSample {
RelaySignalSample {
c2s_bytes: current.c2s_bytes.saturating_sub(previous.c2s_bytes),
s2c_requested_bytes: current
.s2c_requested_bytes
.saturating_sub(previous.s2c_requested_bytes),
s2c_written_bytes: current
.s2c_written_bytes
.saturating_sub(previous.s2c_written_bytes),
s2c_write_ops: current.s2c_write_ops.saturating_sub(previous.s2c_write_ops),
s2c_partial_writes: current
.s2c_partial_writes
.saturating_sub(previous.s2c_partial_writes),
s2c_consecutive_pending_writes: current.s2c_consecutive_pending_writes,
}
}
#[allow(clippy::too_many_arguments)]
fn apply_controller_transition(
transition: crate::proxy::adaptive_buffers::TierTransition,
controller: &mut SessionAdaptiveController,
lease: &mut DirectBufferLease,
c2s_state: &AdaptiveBufferState,
s2c_state: &AdaptiveBufferState,
base: (usize, usize),
ceilings: (usize, usize),
budget: &DirectBufferBudget,
) {
let sizes = direct_copy_buffers_for_tier_with_ceilings(
transition.to,
base.0,
base.1,
ceilings.0,
ceilings.1,
);
if transition.to > transition.from {
if !lease.try_grow_to(sizes.0.saturating_add(sizes.1)) {
*controller = SessionAdaptiveController::new(transition.from);
return;
}
} else {
match transition.reason {
TierTransitionReason::QuietDemotion => budget.increment_quiet_demotion(),
TierTransitionReason::SustainedWritePressure => {
budget.increment_write_pressure_demotion();
}
TierTransitionReason::SoftConfirmed | TierTransitionReason::HardPressure => {}
}
}
set_desired_sizes(c2s_state, s2c_state, sizes);
lease.set_tier(transition.to.as_u8() as usize);
}
fn apply_global_pressure_demotion(
controller: &mut SessionAdaptiveController,
lease: &mut DirectBufferLease,
c2s_state: &AdaptiveBufferState,
s2c_state: &AdaptiveBufferState,
base: (usize, usize),
ceilings: (usize, usize),
budget: &DirectBufferBudget,
) {
let current = controller.tier();
let target = current.demote();
if target == current {
return;
}
*controller = SessionAdaptiveController::new(target);
let sizes =
direct_copy_buffers_for_tier_with_ceilings(target, base.0, base.1, ceilings.0, ceilings.1);
set_desired_sizes(c2s_state, s2c_state, sizes);
lease.set_tier(target.as_u8() as usize);
budget.increment_global_pressure_demotion();
}
fn set_desired_sizes(
c2s_state: &AdaptiveBufferState,
s2c_state: &AdaptiveBufferState,
sizes: (usize, usize),
) {
c2s_state
.desired_bytes
.store(sizes.0.max(1), Ordering::Release);
s2c_state
.desired_bytes
.store(sizes.1.max(1), Ordering::Release);
}
fn reconcile_reservation(
lease: &mut DirectBufferLease,
c2s_state: &AdaptiveBufferState,
s2c_state: &AdaptiveBufferState,
) {
// Promotion reserves the desired allocation before either reader grows.
// Demotion keeps the actual allocation covered until its buffered bytes drain.
let covered_c2s = c2s_state
.actual_bytes
.load(Ordering::Acquire)
.max(c2s_state.desired_bytes.load(Ordering::Acquire));
let covered_s2c = s2c_state
.actual_bytes
.load(Ordering::Acquire)
.max(s2c_state.desired_bytes.load(Ordering::Acquire));
lease.shrink_to(covered_c2s.saturating_add(covered_s2c));
}
@@ -0,0 +1,72 @@
use std::io;
use std::pin::Pin;
use std::sync::Arc;
use std::sync::atomic::Ordering;
use std::task::{Context, Poll};
use tokio::io::AsyncWrite;
use super::super::io::SharedCounters;
/// Direct-only writer wrapper that exposes bounded backpressure signals.
pub(super) struct WritePressureIo<W> {
inner: W,
counters: Arc<SharedCounters>,
}
impl<W> WritePressureIo<W> {
/// Wraps the client writer without changing its I/O or error contract.
pub(super) fn new(inner: W, counters: Arc<SharedCounters>) -> Self {
Self { inner, counters }
}
}
impl<W: AsyncWrite + Unpin> AsyncWrite for WritePressureIo<W> {
fn poll_write(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buffer: &[u8],
) -> Poll<io::Result<usize>> {
let this = self.get_mut();
if !buffer.is_empty() {
this.counters
.s2c_requested_bytes
.fetch_add(buffer.len() as u64, Ordering::Relaxed);
}
match Pin::new(&mut this.inner).poll_write(cx, buffer) {
Poll::Ready(Ok(written)) => {
this.counters
.s2c_consecutive_pending_writes
.store(0, Ordering::Relaxed);
if written < buffer.len() {
this.counters
.s2c_partial_writes
.fetch_add(1, Ordering::Relaxed);
}
Poll::Ready(Ok(written))
}
Poll::Ready(Err(error)) => {
this.counters
.s2c_consecutive_pending_writes
.store(0, Ordering::Relaxed);
Poll::Ready(Err(error))
}
Poll::Pending => {
let _ = this.counters.s2c_consecutive_pending_writes.fetch_update(
Ordering::Relaxed,
Ordering::Relaxed,
|current| Some(current.saturating_add(1)),
);
Poll::Pending
}
}
}
fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Pin::new(&mut self.get_mut().inner).poll_flush(cx)
}
fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Pin::new(&mut self.get_mut().inner).poll_shutdown(cx)
}
}
+10 -1
View File
@@ -1,4 +1,4 @@
use std::sync::atomic::{AtomicU64, Ordering}; use std::sync::atomic::{AtomicU32, AtomicU64, Ordering};
use std::time::Duration; use std::time::Duration;
use tokio::time::Instant; use tokio::time::Instant;
@@ -20,6 +20,12 @@ pub(in crate::proxy::relay) struct SharedCounters {
pub(in crate::proxy::relay) c2s_ops: AtomicU64, pub(in crate::proxy::relay) c2s_ops: AtomicU64,
/// Number of poll_write completions (≈ S→C chunks) /// Number of poll_write completions (≈ S→C chunks)
pub(in crate::proxy::relay) s2c_ops: AtomicU64, pub(in crate::proxy::relay) s2c_ops: AtomicU64,
/// Bytes presented to client writes, including retried pending writes.
pub(in crate::proxy::relay) s2c_requested_bytes: AtomicU64,
/// Successful client writes that consumed only part of the offered slice.
pub(in crate::proxy::relay) s2c_partial_writes: AtomicU64,
/// Consecutive pending client writes observed by the active copy loop.
pub(in crate::proxy::relay) s2c_consecutive_pending_writes: AtomicU32,
/// Milliseconds since relay epoch of last I/O activity /// Milliseconds since relay epoch of last I/O activity
last_activity_ms: AtomicU64, last_activity_ms: AtomicU64,
} }
@@ -31,6 +37,9 @@ impl SharedCounters {
s2c_bytes: AtomicU64::new(0), s2c_bytes: AtomicU64::new(0),
c2s_ops: AtomicU64::new(0), c2s_ops: AtomicU64::new(0),
s2c_ops: AtomicU64::new(0), s2c_ops: AtomicU64::new(0),
s2c_requested_bytes: AtomicU64::new(0),
s2c_partial_writes: AtomicU64::new(0),
s2c_consecutive_pending_writes: AtomicU32::new(0),
last_activity_ms: AtomicU64::new(0), last_activity_ms: AtomicU64::new(0),
} }
} }
+1 -14
View File
@@ -1,6 +1,5 @@
use crate::stats::UserStats; use crate::stats::UserStats;
use std::io; use std::io;
use std::sync::atomic::Ordering;
#[derive(Debug)] #[derive(Debug)]
struct QuotaIoSentinel; struct QuotaIoSentinel;
@@ -52,17 +51,5 @@ pub(super) fn refund_reserved_quota_bytes(user_stats: &UserStats, reserved_bytes
if reserved_bytes == 0 { if reserved_bytes == 0 {
return; return;
} }
let mut current = user_stats.quota_used.load(Ordering::Relaxed); user_stats.refund_quota(reserved_bytes);
loop {
let next = current.saturating_sub(reserved_bytes);
match user_stats.quota_used.compare_exchange_weak(
current,
next,
Ordering::Relaxed,
Ordering::Relaxed,
) {
Ok(_) => return,
Err(observed) => current = observed,
}
}
} }
+12
View File
@@ -9,6 +9,7 @@ use dashmap::DashMap;
use tokio::sync::{OwnedSemaphorePermit, Semaphore, mpsc}; use tokio::sync::{OwnedSemaphorePermit, Semaphore, mpsc};
use tokio_util::sync::CancellationToken; use tokio_util::sync::CancellationToken;
use crate::proxy::direct_buffer_budget::{DirectBufferBudget, fallback_direct_buffer_hard_limit};
use crate::proxy::handshake::{AuthProbeSaturationState, AuthProbeState}; use crate::proxy::handshake::{AuthProbeSaturationState, AuthProbeState};
use crate::proxy::middle_relay::{DesyncDedupRotationState, RelayIdleCandidateRegistry}; use crate::proxy::middle_relay::{DesyncDedupRotationState, RelayIdleCandidateRegistry};
use crate::proxy::traffic_limiter::TrafficLimiter; use crate::proxy::traffic_limiter::TrafficLimiter;
@@ -69,6 +70,7 @@ pub(crate) struct ProxySharedState {
pub(crate) handshake: HandshakeSharedState, pub(crate) handshake: HandshakeSharedState,
pub(crate) middle_relay: MiddleRelaySharedState, pub(crate) middle_relay: MiddleRelaySharedState,
pub(crate) traffic_limiter: Arc<TrafficLimiter>, pub(crate) traffic_limiter: Arc<TrafficLimiter>,
pub(crate) direct_buffer_budget: Arc<DirectBufferBudget>,
disabled_users: DashMap<String, ()>, disabled_users: DashMap<String, ()>,
active_user_sessions: DashMap<(String, u64), CancellationToken>, active_user_sessions: DashMap<(String, u64), CancellationToken>,
pub(crate) conntrack_pressure_active: AtomicBool, pub(crate) conntrack_pressure_active: AtomicBool,
@@ -101,6 +103,15 @@ impl Drop for UserSessionGuard {
impl ProxySharedState { impl ProxySharedState {
pub(crate) fn new() -> Arc<Self> { pub(crate) fn new() -> Arc<Self> {
Self::new_with_direct_buffer_budget(DirectBufferBudget::new(
fallback_direct_buffer_hard_limit(),
))
}
/// Creates process state with the startup-resolved Direct buffer envelope.
pub(crate) fn new_with_direct_buffer_budget(
direct_buffer_budget: Arc<DirectBufferBudget>,
) -> Arc<Self> {
Arc::new(Self { Arc::new(Self {
handshake: HandshakeSharedState { handshake: HandshakeSharedState {
auth_probe: DashMap::new(), auth_probe: DashMap::new(),
@@ -129,6 +140,7 @@ impl ProxySharedState {
relay_idle_mark_seq: AtomicU64::new(0), relay_idle_mark_seq: AtomicU64::new(0),
}, },
traffic_limiter: TrafficLimiter::new(), traffic_limiter: TrafficLimiter::new(),
direct_buffer_budget,
disabled_users: DashMap::new(), disabled_users: DashMap::new(),
active_user_sessions: DashMap::new(), active_user_sessions: DashMap::new(),
conntrack_pressure_active: AtomicBool::new(false), conntrack_pressure_active: AtomicBool::new(false),
@@ -0,0 +1,58 @@
use super::*;
#[test]
fn configured_direct_sizes_are_strict_tier_ceilings() {
let base = (4 * 1024, 8 * 1024);
let ceilings = (64 * 1024, 256 * 1024);
assert_eq!(
direct_copy_buffers_for_tier_with_ceilings(
AdaptiveTier::Base,
base.0,
base.1,
ceilings.0,
ceilings.1,
),
base
);
assert_eq!(
direct_copy_buffers_for_tier_with_ceilings(
AdaptiveTier::Tier3,
base.0,
base.1,
ceilings.0,
ceilings.1,
),
ceilings
);
}
#[test]
fn sustained_pending_pressure_demotes_after_transient_promotion() {
let mut controller = SessionAdaptiveController::new(AdaptiveTier::Tier1);
let pressure = RelaySignalSample {
c2s_bytes: 0,
s2c_requested_bytes: 1024,
s2c_written_bytes: 0,
s2c_write_ops: 0,
s2c_partial_writes: 0,
s2c_consecutive_pending_writes: 3,
};
let first = controller
.observe(pressure, 10.0)
.expect("transient pressure must retain the staged promotion");
assert_eq!(first.reason, TierTransitionReason::HardPressure);
let second = controller
.observe(pressure, 10.0)
.expect("bounded transient pressure may promote one additional tier");
assert_eq!(second.reason, TierTransitionReason::HardPressure);
let sustained = controller
.observe(pressure, 10.0)
.expect("sustained pressure must release one tier");
assert_eq!(
sustained.reason,
TierTransitionReason::SustainedWritePressure
);
assert_eq!(sustained.to, AdaptiveTier::Tier2);
}
@@ -0,0 +1,38 @@
use super::*;
#[test]
fn lease_drop_releases_the_complete_reservation() {
let budget = DirectBufferBudget::new(16 * 1024);
{
let lease = budget
.try_reserve(12 * 1024, false)
.expect("minimum reservation must fit");
assert_eq!(lease.reserved_bytes(), 12 * 1024);
assert_eq!(budget.snapshot().reserved_bytes, 12 * 1024);
}
assert_eq!(budget.snapshot().reserved_bytes, 0);
}
#[test]
fn absolute_ceiling_rejects_excess_minimum_reservations() {
let budget = DirectBufferBudget::new(16 * 1024);
let _lease = budget
.try_reserve(12 * 1024, true)
.expect("first minimum reservation must fit");
assert!(budget.try_reserve(8 * 1024, true).is_none());
assert_eq!(budget.snapshot().reserved_bytes, 12 * 1024);
}
#[test]
fn growth_and_shrink_keep_accounting_balanced() {
let budget = DirectBufferBudget::new(32 * 1024);
let mut lease = budget
.try_reserve(12 * 1024, false)
.expect("base reservation must fit");
assert!(lease.try_grow_to(24 * 1024));
assert_eq!(budget.snapshot().reserved_bytes, 24 * 1024);
lease.shrink_to(16 * 1024);
assert_eq!(budget.snapshot().reserved_bytes, 16 * 1024);
drop(lease);
assert_eq!(budget.snapshot().reserved_bytes, 0);
}
@@ -0,0 +1,150 @@
use std::io;
use std::pin::Pin;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::{Arc, Mutex};
use std::task::{Context, Poll};
use tokio::io::AsyncWrite;
use super::*;
use crate::crypto::AesCtr;
use crate::protocol::framing::INTERMEDIATE_WIRE_LEN_MASK;
#[derive(Clone, Default)]
struct RecordingWriter {
writes: Arc<Mutex<Vec<u8>>>,
flushes: Arc<AtomicUsize>,
}
impl RecordingWriter {
fn captured(&self) -> Vec<u8> {
self.writes
.lock()
.expect("test writer capture lock must not be poisoned")
.clone()
}
}
impl AsyncWrite for RecordingWriter {
fn poll_write(
self: Pin<&mut Self>,
_cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<io::Result<usize>> {
self.writes
.lock()
.expect("test writer capture lock must not be poisoned")
.extend_from_slice(buf);
Poll::Ready(Ok(buf.len()))
}
fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
self.flushes.fetch_add(1, Ordering::Relaxed);
Poll::Ready(Ok(()))
}
fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Poll::Ready(Ok(()))
}
}
fn crypto_writer(inner: RecordingWriter) -> CryptoWriter<RecordingWriter> {
let key = [0u8; 32];
CryptoWriter::new(inner, AesCtr::new(&key, 0), 8 * 1024 * 1024)
}
fn decrypt_capture(mut encrypted: Vec<u8>) -> Vec<u8> {
let key = [0u8; 32];
let mut cipher = AesCtr::new(&key, 0);
cipher.apply(&mut encrypted);
encrypted
}
fn secure_wire_len(cleartext: &[u8]) -> usize {
let header = cleartext
.get(..4)
.expect("secure frame must include an intermediate header");
(u32::from_le_bytes(
header
.try_into()
.expect("secure frame header must be four bytes"),
) & INTERMEDIATE_WIRE_LEN_MASK) as usize
}
async fn write_secure_payload(payload_len: usize) -> (MeD2cWriteMode, Vec<u8>) {
let inner = RecordingWriter::default();
let capture = inner.clone();
let mut writer = crypto_writer(inner);
let payload = vec![0xa5; payload_len];
let mut frame_buf = Vec::new();
let cancel = CancellationToken::new();
let rng = SecureRandom::new();
let mode = write_client_payload(
&mut writer,
ProtoTag::Secure,
0,
&payload,
&rng,
&mut frame_buf,
&cancel,
)
.await
.expect("secure payload write must succeed");
flush_client_or_cancel(&mut writer, &cancel)
.await
.expect("secure payload flush must succeed");
(mode, decrypt_capture(capture.captured()))
}
fn assert_secure_payload_with_tail_padding(cleartext: &[u8], payload_len: usize) {
let wire_len = secure_wire_len(cleartext);
assert_eq!(cleartext.len(), 4 + wire_len);
assert!(
cleartext[4..4 + payload_len]
.iter()
.all(|byte| *byte == 0xa5)
);
let padding_len = wire_len
.checked_sub(payload_len)
.expect("secure wire length must include payload bytes");
assert!((1..=3).contains(&padding_len));
assert_ne!(wire_len % 4, 0);
}
#[tokio::test]
async fn queue_drain_flush_reason_performs_physical_client_flush() {
let inner = RecordingWriter::default();
let flushes = inner.flushes.clone();
let mut writer = crypto_writer(inner);
let cancel = CancellationToken::new();
assert!(me_d2c_flush_reason_requires_client_flush(
MeD2cFlushReason::QueueDrain
));
flush_client_or_cancel(&mut writer, &cancel)
.await
.expect("client flush must succeed");
assert_eq!(flushes.load(Ordering::Relaxed), 1);
}
#[tokio::test]
async fn secure_payload_coalesced_path_keeps_tail_padding() {
let payload_len = 8;
let (mode, cleartext) = write_secure_payload(payload_len).await;
assert!(matches!(mode, MeD2cWriteMode::Coalesced));
assert_secure_payload_with_tail_padding(&cleartext, payload_len);
}
#[tokio::test]
async fn secure_payload_split_path_keeps_tail_padding() {
let payload_len = ME_D2C_SINGLE_WRITE_COALESCE_MAX_BYTES;
let (mode, cleartext) = write_secure_payload(payload_len).await;
assert!(matches!(mode, MeD2cWriteMode::Split));
assert_secure_payload_with_tail_padding(&cleartext, payload_len);
}
+153 -8
View File
@@ -10,7 +10,7 @@ use arc_swap::ArcSwap;
use dashmap::DashMap; use dashmap::DashMap;
use ipnetwork::IpNetwork; use ipnetwork::IpNetwork;
use crate::config::RateLimitBps; use crate::config::{CidrRateLimitKey, RateLimitBps};
const REGISTRY_SHARDS: usize = 64; const REGISTRY_SHARDS: usize = 64;
const FAIR_EPOCH_MS: u64 = 20; const FAIR_EPOCH_MS: u64 = 20;
@@ -413,16 +413,29 @@ struct CidrRule {
prefix_len: u8, prefix_len: u8,
} }
#[derive(Clone, Copy)]
struct CidrAutoRule {
prefix_len: u8,
limits: RateLimitBps,
}
enum CidrPolicyMatch<'a> {
Explicit(&'a CidrRule),
Auto { key: String, limits: RateLimitBps },
}
#[derive(Default)] #[derive(Default)]
struct PolicySnapshot { struct PolicySnapshot {
user_limits: HashMap<String, RateLimitBps>, user_limits: HashMap<String, RateLimitBps>,
cidr_rules_v4: Vec<CidrRule>, cidr_rules_v4: Vec<CidrRule>,
cidr_rules_v6: Vec<CidrRule>, cidr_rules_v6: Vec<CidrRule>,
cidr_auto_rules_v4: Vec<CidrAutoRule>,
cidr_auto_rules_v6: Vec<CidrAutoRule>,
cidr_rule_keys: HashSet<String>, cidr_rule_keys: HashSet<String>,
} }
impl PolicySnapshot { impl PolicySnapshot {
fn match_cidr(&self, ip: IpAddr) -> Option<&CidrRule> { fn match_cidr(&self, ip: IpAddr) -> Option<CidrPolicyMatch<'_>> {
match ip { match ip {
IpAddr::V4(_) => self IpAddr::V4(_) => self
.cidr_rules_v4 .cidr_rules_v4
@@ -433,6 +446,20 @@ impl PolicySnapshot {
.iter() .iter()
.find(|rule| rule.cidr.contains(ip)), .find(|rule| rule.cidr.contains(ip)),
} }
.map(CidrPolicyMatch::Explicit)
.or_else(|| self.match_auto_cidr(ip))
}
fn match_auto_cidr(&self, ip: IpAddr) -> Option<CidrPolicyMatch<'_>> {
let rule = match ip {
IpAddr::V4(_) => self.cidr_auto_rules_v4.first()?,
IpAddr::V6(_) => self.cidr_auto_rules_v6.first()?,
};
let key = auto_cidr_bucket_key(ip, rule.prefix_len)?;
Some(CidrPolicyMatch::Auto {
key,
limits: rule.limits,
})
} }
} }
@@ -633,7 +660,7 @@ impl TrafficLimiter {
pub fn apply_policy( pub fn apply_policy(
&self, &self,
user_limits: HashMap<String, RateLimitBps>, user_limits: HashMap<String, RateLimitBps>,
cidr_limits: HashMap<IpNetwork, RateLimitBps>, cidr_limits: HashMap<CidrRateLimitKey, RateLimitBps>,
) { ) {
let filtered_users = user_limits let filtered_users = user_limits
.into_iter() .into_iter()
@@ -642,11 +669,15 @@ impl TrafficLimiter {
let mut cidr_rules_v4 = Vec::new(); let mut cidr_rules_v4 = Vec::new();
let mut cidr_rules_v6 = Vec::new(); let mut cidr_rules_v6 = Vec::new();
let mut cidr_auto_rules_v4 = Vec::new();
let mut cidr_auto_rules_v6 = Vec::new();
let mut cidr_rule_keys = HashSet::new(); let mut cidr_rule_keys = HashSet::new();
for (cidr, limits) in cidr_limits { for (key, limits) in cidr_limits {
if limits.up_bps == 0 && limits.down_bps == 0 { if limits.up_bps == 0 && limits.down_bps == 0 {
continue; continue;
} }
match key {
CidrRateLimitKey::Network(cidr) => {
let key = cidr.to_string(); let key = cidr.to_string();
let rule = CidrRule { let rule = CidrRule {
key: key.clone(), key: key.clone(),
@@ -660,21 +691,42 @@ impl TrafficLimiter {
IpNetwork::V6(_) => cidr_rules_v6.push(rule), IpNetwork::V6(_) => cidr_rules_v6.push(rule),
} }
} }
CidrRateLimitKey::AutoV4(prefix_len) => {
cidr_auto_rules_v4.push(CidrAutoRule { prefix_len, limits });
}
CidrRateLimitKey::AutoV6(prefix_len) => {
cidr_auto_rules_v6.push(CidrAutoRule { prefix_len, limits });
}
CidrRateLimitKey::AutoDual(prefix_len) => {
cidr_auto_rules_v4.push(CidrAutoRule { prefix_len, limits });
cidr_auto_rules_v6.push(CidrAutoRule {
prefix_len: prefix_len.saturating_mul(4),
limits,
});
}
}
}
cidr_rules_v4.sort_by(|a, b| b.prefix_len.cmp(&a.prefix_len)); cidr_rules_v4.sort_by(|a, b| b.prefix_len.cmp(&a.prefix_len));
cidr_rules_v6.sort_by(|a, b| b.prefix_len.cmp(&a.prefix_len)); cidr_rules_v6.sort_by(|a, b| b.prefix_len.cmp(&a.prefix_len));
cidr_auto_rules_v4.sort_by(|a, b| b.prefix_len.cmp(&a.prefix_len));
cidr_auto_rules_v6.sort_by(|a, b| b.prefix_len.cmp(&a.prefix_len));
let cidr_policy_entries =
cidr_rule_keys.len() + cidr_auto_rules_v4.len() + cidr_auto_rules_v6.len();
self.user_scope self.user_scope
.policy_entries .policy_entries
.store(filtered_users.len() as u64, Ordering::Relaxed); .store(filtered_users.len() as u64, Ordering::Relaxed);
self.cidr_scope self.cidr_scope
.policy_entries .policy_entries
.store(cidr_rule_keys.len() as u64, Ordering::Relaxed); .store(cidr_policy_entries as u64, Ordering::Relaxed);
self.policy.store(Arc::new(PolicySnapshot { self.policy.store(Arc::new(PolicySnapshot {
user_limits: filtered_users, user_limits: filtered_users,
cidr_rules_v4, cidr_rules_v4,
cidr_rules_v6, cidr_rules_v6,
cidr_auto_rules_v4,
cidr_auto_rules_v6,
cidr_rule_keys, cidr_rule_keys,
})); }));
@@ -703,11 +755,15 @@ impl TrafficLimiter {
let mut cidr_bucket = None; let mut cidr_bucket = None;
let mut cidr_user_key = None; let mut cidr_user_key = None;
let mut cidr_user_share = None; let mut cidr_user_share = None;
if let Some(rule) = policy.match_cidr(client_ip) { if let Some(rule_match) = policy.match_cidr(client_ip) {
let (key, limits) = match &rule_match {
CidrPolicyMatch::Explicit(rule) => (rule.key.as_str(), rule.limits),
CidrPolicyMatch::Auto { key, limits } => (key.as_str(), *limits),
};
let bucket = self let bucket = self
.cidr_buckets .cidr_buckets
.get_or_insert_with(rule.key.as_str(), || CidrBucket::new(rule.limits)); .get_or_insert_with(key, || CidrBucket::new(limits));
bucket.set_rates(rule.limits); bucket.set_rates(limits);
bucket.active_leases.fetch_add(1, Ordering::Relaxed); bucket.active_leases.fetch_add(1, Ordering::Relaxed);
self.cidr_scope self.cidr_scope
.active_leases .active_leases
@@ -841,6 +897,16 @@ fn bytes_per_epoch(bps: u64) -> u64 {
bytes.max(1) bytes.max(1)
} }
fn auto_cidr_bucket_key(ip: IpAddr, prefix_len: u8) -> Option<String> {
let cidr = IpNetwork::new(ip, prefix_len).ok()?;
let network = IpNetwork::new(cidr.network(), prefix_len).ok()?;
let family = match network {
IpNetwork::V4(_) => "4",
IpNetwork::V6(_) => "6",
};
Some(format!("auto:{family}:{network}"))
}
fn current_epoch() -> u64 { fn current_epoch() -> u64 {
let start = limiter_epoch_start(); let start = limiter_epoch_start();
let elapsed_ms = start.elapsed().as_millis() as u64; let elapsed_ms = start.elapsed().as_millis() as u64;
@@ -851,3 +917,82 @@ fn limiter_epoch_start() -> &'static Instant {
static START: OnceLock<Instant> = OnceLock::new(); static START: OnceLock<Instant> = OnceLock::new();
START.get_or_init(Instant::now) START.get_or_init(Instant::now)
} }
#[cfg(test)]
mod tests {
use super::*;
fn rate(up_bps: u64, down_bps: u64) -> RateLimitBps {
RateLimitBps { up_bps, down_bps }
}
#[test]
fn explicit_cidr_rule_wins_over_auto_template() {
let limiter = TrafficLimiter::new();
let mut cidr_limits = HashMap::new();
cidr_limits.insert(CidrRateLimitKey::AutoV4(24), rate(1_000, 0));
cidr_limits.insert(
CidrRateLimitKey::Network("203.0.113.7/32".parse().unwrap()),
rate(2_000, 0),
);
limiter.apply_policy(HashMap::new(), cidr_limits);
let policy = limiter.policy.load_full();
let matched = policy.match_cidr("203.0.113.7".parse().unwrap()).unwrap();
match matched {
CidrPolicyMatch::Explicit(rule) => assert_eq!(rule.key.as_str(), "203.0.113.7/32"),
CidrPolicyMatch::Auto { .. } => panic!("explicit CIDR must have priority"),
}
}
#[test]
fn auto_template_uses_longest_prefix() {
let limiter = TrafficLimiter::new();
let mut cidr_limits = HashMap::new();
cidr_limits.insert(CidrRateLimitKey::AutoV4(24), rate(1_000, 0));
cidr_limits.insert(CidrRateLimitKey::AutoV4(32), rate(2_000, 0));
limiter.apply_policy(HashMap::new(), cidr_limits);
let policy = limiter.policy.load_full();
let matched = policy.match_cidr("203.0.113.129".parse().unwrap()).unwrap();
match matched {
CidrPolicyMatch::Auto { key, limits } => {
assert_eq!(key, "auto:4:203.0.113.129/32");
assert_eq!(limits.up_bps, 2_000);
}
CidrPolicyMatch::Explicit(_) => panic!("auto-template match expected"),
}
}
#[test]
fn dual_auto_template_maps_v6_prefix_by_four() {
let limiter = TrafficLimiter::new();
let mut cidr_limits = HashMap::new();
cidr_limits.insert(CidrRateLimitKey::AutoDual(32), rate(1_000, 0));
limiter.apply_policy(HashMap::new(), cidr_limits);
let policy = limiter.policy.load_full();
let matched = policy.match_cidr("2001:db8::1".parse().unwrap()).unwrap();
match matched {
CidrPolicyMatch::Auto { key, .. } => {
assert_eq!(key, "auto:6:2001:db8::1/128");
}
CidrPolicyMatch::Explicit(_) => panic!("auto-template match expected"),
}
}
#[test]
fn auto_cidr_bucket_key_canonicalizes_network_address() {
assert_eq!(
auto_cidr_bucket_key("203.0.113.129".parse().unwrap(), 24).unwrap(),
"auto:4:203.0.113.0/24"
);
assert_eq!(
auto_cidr_bucket_key("2001:db8::abcd".parse().unwrap(), 64).unwrap(),
"auto:6:2001:db8::/64"
);
}
}
+6 -5
View File
@@ -80,7 +80,11 @@ impl Stats {
return handle; return handle;
} }
let entry = self.user_stats.entry(user.to_string()).or_default(); let quota = self.quota_store.user(user);
let entry = self
.user_stats
.entry(user.to_string())
.or_insert_with(|| Arc::new(UserStats::with_quota(quota)));
if entry.last_seen_epoch_secs.load(Ordering::Relaxed) == 0 { if entry.last_seen_epoch_secs.load(Ordering::Relaxed) == 0 {
self.touch_user_stats(entry.value().as_ref()); self.touch_user_stats(entry.value().as_ref());
} }
@@ -166,10 +170,7 @@ impl Stats {
#[inline] #[inline]
pub(crate) fn quota_charge_post_write(&self, user_stats: &UserStats, bytes: u64) -> u64 { pub(crate) fn quota_charge_post_write(&self, user_stats: &UserStats, bytes: u64) -> u64 {
self.touch_user_stats(user_stats); self.touch_user_stats(user_stats);
user_stats user_stats.quota.charge(bytes)
.quota_used
.fetch_add(bytes, Ordering::Relaxed)
.saturating_add(bytes)
} }
pub(super) fn maybe_cleanup_user_stats(&self) { pub(super) fn maybe_cleanup_user_stats(&self) {
+36
View File
@@ -204,6 +204,12 @@ impl Stats {
self.buffer_pool_in_use_gauge.load(Ordering::Relaxed) self.buffer_pool_in_use_gauge.load(Ordering::Relaxed)
} }
/// Returns the count of non-standard buffers replaced before pooling.
pub fn get_buffer_pool_replaced_nonstandard_total(&self) -> u64 {
self.buffer_pool_replaced_nonstandard_total
.load(Ordering::Relaxed)
}
pub fn get_me_c2me_send_full_total(&self) -> u64 { pub fn get_me_c2me_send_full_total(&self) -> u64 {
self.me_c2me_send_full_total.load(Ordering::Relaxed) self.me_c2me_send_full_total.load(Ordering::Relaxed)
} }
@@ -269,6 +275,36 @@ impl Stats {
self.me_writer_pick_mode_switch_total self.me_writer_pick_mode_switch_total
.load(Ordering::Relaxed) .load(Ordering::Relaxed)
} }
/// Returns the configured resident-memory limit per ME writer.
pub fn get_me_writer_byte_budget_limit_bytes_gauge(&self) -> u64 {
self.me_writer_byte_budget_limit_bytes_gauge
.load(Ordering::Relaxed)
}
/// Returns aggregate queued or enqueueing writer memory reservations.
pub fn get_me_writer_byte_budget_queued_bytes_gauge(&self) -> u64 {
self.me_writer_byte_budget_queued_bytes_gauge
.load(Ordering::Relaxed)
}
/// Returns aggregate writer reservations currently owned by socket writes.
pub fn get_me_writer_byte_budget_inflight_bytes_gauge(&self) -> u64 {
self.me_writer_byte_budget_inflight_bytes_gauge
.load(Ordering::Relaxed)
}
/// Returns the count of blocking writer byte-budget waits.
pub fn get_me_writer_byte_budget_wait_total(&self) -> u64 {
self.me_writer_byte_budget_wait_total
.load(Ordering::Relaxed)
}
/// Returns the count of writer byte-budget wait timeouts.
pub fn get_me_writer_byte_budget_timeout_total(&self) -> u64 {
self.me_writer_byte_budget_timeout_total
.load(Ordering::Relaxed)
}
/// Returns the count of payloads that cannot fit the configured writer budget.
pub fn get_me_writer_byte_budget_oversize_total(&self) -> u64 {
self.me_writer_byte_budget_oversize_total
.load(Ordering::Relaxed)
}
pub fn get_me_socks_kdf_strict_reject(&self) -> u64 { pub fn get_me_socks_kdf_strict_reject(&self) -> u64 {
self.me_socks_kdf_strict_reject.load(Ordering::Relaxed) self.me_socks_kdf_strict_reject.load(Ordering::Relaxed)
} }
+33 -23
View File
@@ -8,6 +8,7 @@ mod core_getters;
mod helpers; mod helpers;
mod me_counters; mod me_counters;
mod me_getters; mod me_getters;
mod quota_store;
mod replay; mod replay;
pub mod telemetry; pub mod telemetry;
pub mod tls_fingerprints; pub mod tls_fingerprints;
@@ -20,6 +21,7 @@ use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicU8, AtomicU64, Ordering}; use std::sync::atomic::{AtomicBool, AtomicU8, AtomicU64, Ordering};
use std::time::Instant; use std::time::Instant;
pub(crate) use self::quota_store::QuotaStore;
#[allow(unused_imports)] #[allow(unused_imports)]
pub use self::replay::{ReplayChecker, ReplayStats}; pub use self::replay::{ReplayChecker, ReplayStats};
use self::telemetry::TelemetryPolicy; use self::telemetry::TelemetryPolicy;
@@ -274,6 +276,7 @@ pub struct Stats {
buffer_pool_pooled_gauge: AtomicU64, buffer_pool_pooled_gauge: AtomicU64,
buffer_pool_allocated_gauge: AtomicU64, buffer_pool_allocated_gauge: AtomicU64,
buffer_pool_in_use_gauge: AtomicU64, buffer_pool_in_use_gauge: AtomicU64,
buffer_pool_replaced_nonstandard_total: AtomicU64,
// C2ME enqueue observability // C2ME enqueue observability
me_c2me_send_full_total: AtomicU64, me_c2me_send_full_total: AtomicU64,
me_c2me_send_high_water_total: AtomicU64, me_c2me_send_high_water_total: AtomicU64,
@@ -292,6 +295,12 @@ pub struct Stats {
me_writer_pick_p2c_no_candidate_total: AtomicU64, me_writer_pick_p2c_no_candidate_total: AtomicU64,
me_writer_pick_blocking_fallback_total: AtomicU64, me_writer_pick_blocking_fallback_total: AtomicU64,
me_writer_pick_mode_switch_total: AtomicU64, me_writer_pick_mode_switch_total: AtomicU64,
me_writer_byte_budget_limit_bytes_gauge: AtomicU64,
me_writer_byte_budget_queued_bytes_gauge: AtomicU64,
me_writer_byte_budget_inflight_bytes_gauge: AtomicU64,
me_writer_byte_budget_wait_total: AtomicU64,
me_writer_byte_budget_timeout_total: AtomicU64,
me_writer_byte_budget_oversize_total: AtomicU64,
me_socks_kdf_strict_reject: AtomicU64, me_socks_kdf_strict_reject: AtomicU64,
me_socks_kdf_compat_fallback: AtomicU64, me_socks_kdf_compat_fallback: AtomicU64,
secure_padding_invalid: AtomicU64, secure_padding_invalid: AtomicU64,
@@ -337,6 +346,7 @@ pub struct Stats {
cached_epoch_secs: AtomicU64, cached_epoch_secs: AtomicU64,
tls_fingerprints: tls_fingerprints::TlsFingerprintCollector, tls_fingerprints: tls_fingerprints::TlsFingerprintCollector,
user_stats: DashMap<String, Arc<UserStats>>, user_stats: DashMap<String, Arc<UserStats>>,
quota_store: Arc<QuotaStore>,
user_stats_last_cleanup_epoch_secs: AtomicU64, user_stats_last_cleanup_epoch_secs: AtomicU64,
start_time: parking_lot::RwLock<Option<Instant>>, start_time: parking_lot::RwLock<Option<Instant>>,
} }
@@ -349,12 +359,7 @@ pub struct UserStats {
pub octets_to_client: AtomicU64, pub octets_to_client: AtomicU64,
pub msgs_from_client: AtomicU64, pub msgs_from_client: AtomicU64,
pub msgs_to_client: AtomicU64, pub msgs_to_client: AtomicU64,
/// Total bytes charged against per-user quota admission. quota: Arc<quota_store::UserQuotaCounters>,
///
/// This counter is the single source of truth for quota enforcement and
/// intentionally tracks attempted traffic, not guaranteed delivery.
pub quota_used: AtomicU64,
pub quota_last_reset_epoch_secs: AtomicU64,
pub last_seen_epoch_secs: AtomicU64, pub last_seen_epoch_secs: AtomicU64,
} }
@@ -371,9 +376,21 @@ pub enum QuotaReserveError {
} }
impl UserStats { impl UserStats {
fn with_quota(quota: Arc<quota_store::UserQuotaCounters>) -> Self {
Self {
quota,
..Self::default()
}
}
#[inline] #[inline]
pub fn quota_used(&self) -> u64 { pub fn quota_used(&self) -> u64 {
self.quota_used.load(Ordering::Relaxed) self.quota.used()
}
#[inline]
pub(crate) fn refund_quota(&self, bytes: u64) {
self.quota.refund(bytes);
} }
/// Attempts one CAS reservation step against the quota counter. /// Attempts one CAS reservation step against the quota counter.
@@ -383,27 +400,20 @@ impl UserStats {
/// with their own contention strategy. /// with their own contention strategy.
#[inline] #[inline]
pub fn quota_try_reserve(&self, bytes: u64, limit: u64) -> Result<u64, QuotaReserveError> { pub fn quota_try_reserve(&self, bytes: u64, limit: u64) -> Result<u64, QuotaReserveError> {
let current = self.quota_used.load(Ordering::Relaxed); self.quota.try_reserve(bytes, limit)
if bytes > limit.saturating_sub(current) {
return Err(QuotaReserveError::LimitExceeded);
}
let next = current.saturating_add(bytes);
match self.quota_used.compare_exchange_weak(
current,
next,
Ordering::Relaxed,
Ordering::Relaxed,
) {
Ok(_) => Ok(next),
Err(_) => Err(QuotaReserveError::Contended),
}
} }
} }
impl Stats { impl Stats {
pub fn new() -> Self { pub fn new() -> Self {
let stats = Self::default(); Self::with_quota_store(Arc::new(QuotaStore::default()))
}
pub(crate) fn with_quota_store(quota_store: Arc<QuotaStore>) -> Self {
let stats = Self {
quota_store,
..Self::default()
};
stats.apply_telemetry_policy(TelemetryPolicy::default()); stats.apply_telemetry_policy(TelemetryPolicy::default());
stats.refresh_cached_epoch_secs(); stats.refresh_cached_epoch_secs();
*stats.start_time.write() = Some(Instant::now()); *stats.start_time.write() = Some(Instant::now());
+147
View File
@@ -0,0 +1,147 @@
use std::collections::HashMap;
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use dashmap::DashMap;
use super::{QuotaReserveError, UserQuotaSnapshot};
/// Process-scoped per-user quota accounting shared by runtime generations.
#[derive(Default)]
pub struct QuotaStore {
users: DashMap<String, Arc<UserQuotaCounters>>,
}
/// Atomic quota state for one configured user.
#[derive(Default)]
pub(crate) struct UserQuotaCounters {
used_bytes: AtomicU64,
last_reset_epoch_secs: AtomicU64,
}
impl QuotaStore {
pub(crate) fn user(&self, user: &str) -> Arc<UserQuotaCounters> {
if let Some(existing) = self.users.get(user) {
return Arc::clone(existing.value());
}
Arc::clone(
self.users
.entry(user.to_string())
.or_insert_with(|| Arc::new(UserQuotaCounters::default()))
.value(),
)
}
pub(crate) fn used(&self, user: &str) -> u64 {
self.users.get(user).map(|state| state.used()).unwrap_or(0)
}
pub(crate) fn load(&self, user: &str, used_bytes: u64, last_reset_epoch_secs: u64) {
let state = self.user(user);
state.used_bytes.store(used_bytes, Ordering::Relaxed);
state
.last_reset_epoch_secs
.store(last_reset_epoch_secs, Ordering::Relaxed);
}
pub(crate) fn reset(&self, user: &str, now_epoch_secs: u64) -> UserQuotaSnapshot {
let state = self.user(user);
state.used_bytes.store(0, Ordering::Relaxed);
state
.last_reset_epoch_secs
.store(now_epoch_secs, Ordering::Relaxed);
UserQuotaSnapshot {
used_bytes: 0,
last_reset_epoch_secs: now_epoch_secs,
}
}
pub(crate) fn snapshot(&self) -> HashMap<String, UserQuotaSnapshot> {
let mut out = HashMap::new();
for entry in self.users.iter() {
let state = entry.value();
let used_bytes = state.used();
let last_reset_epoch_secs = state.last_reset_epoch_secs.load(Ordering::Relaxed);
if used_bytes == 0 && last_reset_epoch_secs == 0 {
continue;
}
out.insert(
entry.key().clone(),
UserQuotaSnapshot {
used_bytes,
last_reset_epoch_secs,
},
);
}
out
}
}
impl UserQuotaCounters {
#[inline]
pub(crate) fn used(&self) -> u64 {
self.used_bytes.load(Ordering::Relaxed)
}
#[inline]
pub(crate) fn charge(&self, bytes: u64) -> u64 {
self.used_bytes
.fetch_add(bytes, Ordering::Relaxed)
.saturating_add(bytes)
}
#[inline]
pub(crate) fn refund(&self, bytes: u64) {
let mut current = self.used_bytes.load(Ordering::Relaxed);
loop {
let next = current.saturating_sub(bytes);
match self.used_bytes.compare_exchange_weak(
current,
next,
Ordering::Relaxed,
Ordering::Relaxed,
) {
Ok(_) => return,
Err(observed) => current = observed,
}
}
}
#[inline]
pub(crate) fn try_reserve(&self, bytes: u64, limit: u64) -> Result<u64, QuotaReserveError> {
let current = self.used_bytes.load(Ordering::Relaxed);
if bytes > limit.saturating_sub(current) {
return Err(QuotaReserveError::LimitExceeded);
}
let next = current.saturating_add(bytes);
match self.used_bytes.compare_exchange_weak(
current,
next,
Ordering::Relaxed,
Ordering::Relaxed,
) {
Ok(_) => Ok(next),
Err(_) => Err(QuotaReserveError::Contended),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::stats::Stats;
#[test]
fn quota_counters_are_shared_across_stats_generations() {
let store = Arc::new(QuotaStore::default());
let first = Stats::with_quota_store(store.clone());
store.user("alice").charge(512);
assert_eq!(first.get_user_quota_used("alice"), 512);
let second = Stats::with_quota_store(store);
assert_eq!(second.get_user_quota_used("alice"), 512);
second.reset_user_quota("alice");
assert_eq!(first.get_user_quota_used("alice"), 0);
}
}
+5 -35
View File
@@ -119,51 +119,21 @@ impl Stats {
} }
pub fn get_user_quota_used(&self, user: &str) -> u64 { pub fn get_user_quota_used(&self, user: &str) -> u64 {
self.user_stats self.quota_store.used(user)
.get(user)
.map(|s| s.quota_used.load(Ordering::Relaxed))
.unwrap_or(0)
} }
pub fn load_user_quota_state(&self, user: &str, used_bytes: u64, last_reset_epoch_secs: u64) { pub fn load_user_quota_state(&self, user: &str, used_bytes: u64, last_reset_epoch_secs: u64) {
let stats = self.get_or_create_user_stats_handle(user); self.quota_store
stats.quota_used.store(used_bytes, Ordering::Relaxed); .load(user, used_bytes, last_reset_epoch_secs);
stats
.quota_last_reset_epoch_secs
.store(last_reset_epoch_secs, Ordering::Relaxed);
} }
pub fn reset_user_quota(&self, user: &str) -> UserQuotaSnapshot { pub fn reset_user_quota(&self, user: &str) -> UserQuotaSnapshot {
let stats = self.get_or_create_user_stats_handle(user);
let last_reset_epoch_secs = Self::now_epoch_secs(); let last_reset_epoch_secs = Self::now_epoch_secs();
stats.quota_used.store(0, Ordering::Relaxed); self.quota_store.reset(user, last_reset_epoch_secs)
stats
.quota_last_reset_epoch_secs
.store(last_reset_epoch_secs, Ordering::Relaxed);
UserQuotaSnapshot {
used_bytes: 0,
last_reset_epoch_secs,
}
} }
pub fn user_quota_snapshot(&self) -> HashMap<String, UserQuotaSnapshot> { pub fn user_quota_snapshot(&self) -> HashMap<String, UserQuotaSnapshot> {
let mut out = HashMap::new(); self.quota_store.snapshot()
for entry in self.user_stats.iter() {
let stats = entry.value();
let used_bytes = stats.quota_used.load(Ordering::Relaxed);
let last_reset_epoch_secs = stats.quota_last_reset_epoch_secs.load(Ordering::Relaxed);
if used_bytes == 0 && last_reset_epoch_secs == 0 {
continue;
}
out.insert(
entry.key().clone(),
UserQuotaSnapshot {
used_bytes,
last_reset_epoch_secs,
},
);
}
out
} }
pub fn get_handshake_timeouts(&self) -> u64 { pub fn get_handshake_timeouts(&self) -> u64 {
+58
View File
@@ -88,6 +88,56 @@ impl Stats {
.fetch_add(1, Ordering::Relaxed); .fetch_add(1, Ordering::Relaxed);
} }
} }
/// Publishes the configured resident-memory limit applied to each ME writer.
pub fn set_me_writer_byte_budget_limit_bytes(&self, bytes: usize) {
self.me_writer_byte_budget_limit_bytes_gauge
.store(bytes as u64, Ordering::Relaxed);
}
pub(crate) fn add_me_writer_byte_budget_queued_bytes(&self, bytes: u64) {
self.me_writer_byte_budget_queued_bytes_gauge
.fetch_add(bytes, Ordering::Relaxed);
}
pub(crate) fn move_me_writer_byte_budget_to_inflight(&self, bytes: u64) {
let _ = self.me_writer_byte_budget_queued_bytes_gauge.fetch_update(
Ordering::Relaxed,
Ordering::Relaxed,
|current| Some(current.saturating_sub(bytes)),
);
self.me_writer_byte_budget_inflight_bytes_gauge
.fetch_add(bytes, Ordering::Relaxed);
}
pub(crate) fn release_me_writer_byte_budget_queued_bytes(&self, bytes: u64) {
let _ = self.me_writer_byte_budget_queued_bytes_gauge.fetch_update(
Ordering::Relaxed,
Ordering::Relaxed,
|current| Some(current.saturating_sub(bytes)),
);
}
pub(crate) fn release_me_writer_byte_budget_inflight_bytes(&self, bytes: u64) {
let _ = self
.me_writer_byte_budget_inflight_bytes_gauge
.fetch_update(Ordering::Relaxed, Ordering::Relaxed, |current| {
Some(current.saturating_sub(bytes))
});
}
pub(crate) fn increment_me_writer_byte_budget_wait_total(&self) {
if self.telemetry_me_allows_normal() {
self.me_writer_byte_budget_wait_total
.fetch_add(1, Ordering::Relaxed);
}
}
pub(crate) fn increment_me_writer_byte_budget_timeout_total(&self) {
if self.telemetry_me_allows_normal() {
self.me_writer_byte_budget_timeout_total
.fetch_add(1, Ordering::Relaxed);
}
}
pub(crate) fn increment_me_writer_byte_budget_oversize_total(&self) {
if self.telemetry_me_allows_normal() {
self.me_writer_byte_budget_oversize_total
.fetch_add(1, Ordering::Relaxed);
}
}
pub fn increment_me_socks_kdf_strict_reject(&self) { pub fn increment_me_socks_kdf_strict_reject(&self) {
if self.telemetry_me_allows_normal() { if self.telemetry_me_allows_normal() {
self.me_socks_kdf_strict_reject self.me_socks_kdf_strict_reject
@@ -502,6 +552,14 @@ impl Stats {
} }
} }
/// Publishes the cumulative count of non-standard pool buffer replacements.
pub fn set_buffer_pool_replaced_nonstandard_total(&self, value: usize) {
if self.telemetry_me_allows_normal() {
self.buffer_pool_replaced_nonstandard_total
.store(value as u64, Ordering::Relaxed);
}
}
pub fn increment_me_c2me_send_full_total(&self) { pub fn increment_me_c2me_send_full_total(&self) {
if self.telemetry_me_allows_normal() { if self.telemetry_me_allows_normal() {
self.me_c2me_send_full_total.fetch_add(1, Ordering::Relaxed); self.me_c2me_send_full_total.fetch_add(1, Ordering::Relaxed);
+20 -18
View File
@@ -68,8 +68,8 @@ use crate::crypto::AesCtr;
/// Actual limit is supplied at runtime from configuration. /// Actual limit is supplied at runtime from configuration.
const DEFAULT_MAX_PENDING_WRITE: usize = 64 * 1024; const DEFAULT_MAX_PENDING_WRITE: usize = 64 * 1024;
/// Default read buffer capacity (reader mostly decrypts in-place into caller buffer). /// Maximum scratch capacity retained after a completed write.
const DEFAULT_READ_CAPACITY: usize = 16 * 1024; const MAX_RETAINED_SCRATCH_CAPACITY: usize = 32 * 1024;
// ============= CryptoReader State ============= // ============= CryptoReader State =============
@@ -110,10 +110,6 @@ pub struct CryptoReader<R> {
upstream: R, upstream: R,
decryptor: AesCtr, decryptor: AesCtr,
state: CryptoReaderState, state: CryptoReaderState,
/// Reserved for future coalescing optimizations.
#[allow(dead_code)]
read_buf: BytesMut,
} }
impl<R> CryptoReader<R> { impl<R> CryptoReader<R> {
@@ -122,7 +118,6 @@ impl<R> CryptoReader<R> {
upstream, upstream,
decryptor, decryptor,
state: CryptoReaderState::Idle, state: CryptoReaderState::Idle,
read_buf: BytesMut::with_capacity(DEFAULT_READ_CAPACITY),
} }
} }
@@ -321,7 +316,7 @@ struct PendingCiphertext {
impl PendingCiphertext { impl PendingCiphertext {
fn new(max_len: usize) -> Self { fn new(max_len: usize) -> Self {
Self { Self {
buf: BytesMut::with_capacity(16 * 1024), buf: BytesMut::new(),
pos: 0, pos: 0,
max_len, max_len,
} }
@@ -372,15 +367,12 @@ impl PendingCiphertext {
} }
} }
/// Replace the entire pending ciphertext by moving `src` in (swap, no copy). /// Replace the entire pending ciphertext by moving `src` in without copying.
fn replace_with(&mut self, mut src: BytesMut) { fn replace_with(&mut self, src: BytesMut) {
debug_assert!(src.len() <= self.max_len); debug_assert!(src.len() <= self.max_len);
self.buf.clear(); self.buf = src;
self.pos = 0; self.pos = 0;
// Swap: keep allocations hot and avoid copying bytes.
std::mem::swap(&mut self.buf, &mut src);
} }
/// Append plaintext and encrypt appended range in-place. /// Append plaintext and encrypt appended range in-place.
@@ -465,7 +457,7 @@ impl<W> CryptoWriter<W> {
upstream, upstream,
encryptor, encryptor,
state: CryptoWriterState::Idle, state: CryptoWriterState::Idle,
scratch: BytesMut::with_capacity(16 * 1024), scratch: BytesMut::new(),
max_pending_write: max_pending.max(4 * 1024), max_pending_write: max_pending.max(4 * 1024),
} }
} }
@@ -502,6 +494,7 @@ impl<W> CryptoWriter<W> {
} }
fn poison(&mut self, error: io::Error) { fn poison(&mut self, error: io::Error) {
self.scratch = BytesMut::new();
self.state = CryptoWriterState::Poisoned { error: Some(error) }; self.state = CryptoWriterState::Poisoned { error: Some(error) };
} }
@@ -552,6 +545,15 @@ impl<W> CryptoWriter<W> {
scratch.extend_from_slice(plaintext); scratch.extend_from_slice(plaintext);
encryptor.apply(&mut scratch[..]); encryptor.apply(&mut scratch[..]);
} }
/// Clear reusable scratch while releasing allocations inflated by large writes.
fn recycle_scratch(scratch: &mut BytesMut) {
if scratch.capacity() > MAX_RETAINED_SCRATCH_CAPACITY {
*scratch = BytesMut::new();
} else {
scratch.clear();
}
}
} }
impl<W: AsyncWrite + Unpin> CryptoWriter<W> { impl<W: AsyncWrite + Unpin> CryptoWriter<W> {
@@ -698,13 +700,13 @@ impl<W: AsyncWrite + Unpin> AsyncWrite for CryptoWriter<W> {
Poll::Ready(Ok(n)) => { Poll::Ready(Ok(n)) => {
if n == this.scratch.len() { if n == this.scratch.len() {
this.scratch.clear(); Self::recycle_scratch(&mut this.scratch);
return Poll::Ready(Ok(to_accept)); return Poll::Ready(Ok(to_accept));
} }
// Partial upstream write of ciphertext // Partial upstream write of ciphertext
let remainder = this.scratch.split_off(n); let mut remainder = std::mem::take(&mut this.scratch);
this.scratch.clear(); let _ = remainder.split_to(n);
let pending = Self::ensure_pending(&mut this.state, this.max_pending_write); let pending = Self::ensure_pending(&mut this.state, this.max_pending_write);
pending.replace_with(remainder); pending.replace_with(remainder);
+6 -11
View File
@@ -312,7 +312,7 @@ fn encode_secure(frame: &Frame, dst: &mut BytesMut, rng: &SecureRandom) -> io::R
)); ));
} }
// Telegram Desktop VersionD uses a 4-bit random padding length. // Outbound Secure padding avoids full-word tails that readers cannot strip.
let padding_len = secure_padding_len(data.len(), rng); let padding_len = secure_padding_len(data.len(), rng);
let total_len = data.len() + padding_len; let total_len = data.len() + padding_len;
@@ -521,13 +521,7 @@ mod tests {
use tokio_util::codec::{FramedRead, FramedWrite}; use tokio_util::codec::{FramedRead, FramedWrite};
fn assert_secure_decoded_payload(decoded: &[u8], original: &[u8]) { fn assert_secure_decoded_payload(decoded: &[u8], original: &[u8]) {
assert!(decoded.starts_with(original)); assert_eq!(decoded, original);
assert!(
(original.len()..=original.len() + 12).contains(&decoded.len()),
"Secure decoded payload may retain up to 12 bytes of full-word padding, got {}",
decoded.len()
);
assert_eq!(decoded.len() % 4, 0);
} }
#[tokio::test] #[tokio::test]
@@ -653,7 +647,7 @@ mod tests {
} }
#[test] #[test]
fn secure_codec_uses_tdesktop_padding_range_and_jitters_wire_length() { fn secure_codec_uses_non_aligned_padding_and_jitters_wire_length() {
let codec = SecureCodec::new(Arc::new(SecureRandom::new())); let codec = SecureCodec::new(Arc::new(SecureRandom::new()));
let payload = Bytes::from_static(&[1, 2, 3, 4, 5, 6, 7, 8]); let payload = Bytes::from_static(&[1, 2, 3, 4, 5, 6, 7, 8]);
let mut wire_lens = HashSet::new(); let mut wire_lens = HashSet::new();
@@ -666,9 +660,10 @@ mod tests {
let wire_len = u32::from_le_bytes([out[0], out[1], out[2], out[3]]) as usize; let wire_len = u32::from_le_bytes([out[0], out[1], out[2], out[3]]) as usize;
assert_eq!(out.len(), 4 + wire_len); assert_eq!(out.len(), 4 + wire_len);
assert!( assert!(
(payload.len()..=payload.len() + 15).contains(&wire_len), (payload.len() + 1..=payload.len() + 3).contains(&wire_len),
"Secure wire length must be payload+0..15, got {wire_len}" "Secure wire length must be payload+1..3, got {wire_len}"
); );
assert_ne!(wire_len % 4, 0);
wire_lens.insert(wire_len); wire_lens.insert(wire_len);
} }
+2 -8
View File
@@ -367,7 +367,7 @@ impl<W: AsyncWrite + Unpin> SecureIntermediateFrameWriter<W> {
)); ));
} }
// Telegram Desktop VersionD uses a 4-bit random padding length. // Outbound Secure padding avoids full-word tails that readers cannot strip.
let padding_len = secure_padding_len(data.len(), &self.rng); let padding_len = secure_padding_len(data.len(), &self.rng);
let padding = self.rng.bytes(padding_len); let padding = self.rng.bytes(padding_len);
@@ -633,13 +633,7 @@ mod tests {
use tokio::time::{Duration, timeout}; use tokio::time::{Duration, timeout};
fn assert_secure_decoded_payload(decoded: &[u8], original: &[u8]) { fn assert_secure_decoded_payload(decoded: &[u8], original: &[u8]) {
assert!(decoded.starts_with(original)); assert_eq!(decoded, original);
assert!(
(original.len()..=original.len() + 12).contains(&decoded.len()),
"Secure decoded payload may retain up to 12 bytes of full-word padding, got {}",
decoded.len()
);
assert_eq!(decoded.len() % 4, 0);
} }
#[tokio::test] #[tokio::test]
+66 -34
View File
@@ -40,7 +40,7 @@ use std::pin::Pin;
use std::task::{Context, Poll}; use std::task::{Context, Poll};
use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt, ReadBuf}; use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt, ReadBuf};
use super::state::{HeaderBuffer, StreamState, WriteBuffer, YieldBuffer}; use super::state::{HeaderBuffer, StreamState, YieldBuffer};
use crate::protocol::constants::{ use crate::protocol::constants::{
MAX_TLS_CIPHERTEXT_SIZE, MAX_TLS_PLAINTEXT_SIZE, TLS_RECORD_ALERT, TLS_RECORD_APPLICATION, MAX_TLS_CIPHERTEXT_SIZE, MAX_TLS_PLAINTEXT_SIZE, TLS_RECORD_ALERT, TLS_RECORD_APPLICATION,
TLS_RECORD_CHANGE_CIPHER, TLS_RECORD_HANDSHAKE, TLS_VERSION, TLS_RECORD_CHANGE_CIPHER, TLS_RECORD_HANDSHAKE, TLS_VERSION,
@@ -59,6 +59,9 @@ const MAX_TLS_PAYLOAD: usize = MAX_TLS_CIPHERTEXT_SIZE;
/// Note: we never queue unlimited amount of data here; state holds at most one record. /// Note: we never queue unlimited amount of data here; state holds at most one record.
const MAX_PENDING_WRITE: usize = 64 * 1024; const MAX_PENDING_WRITE: usize = 64 * 1024;
/// Maximum record buffer capacity retained between writes.
const MAX_RETAINED_RECORD_CAPACITY: usize = 2 * (TLS_HEADER_SIZE + MAX_TLS_PAYLOAD);
// ============= TLS Record Types ============= // ============= TLS Record Types =============
/// Parsed TLS record header (5 bytes) /// Parsed TLS record header (5 bytes)
@@ -639,7 +642,7 @@ enum TlsWriterState {
/// Writing a complete TLS record (header + body), possibly partially /// Writing a complete TLS record (header + body), possibly partially
WritingRecord { WritingRecord {
record: WriteBuffer, position: usize,
payload_size: usize, payload_size: usize,
}, },
@@ -676,6 +679,7 @@ impl StreamState for TlsWriterState {
pub struct FakeTlsWriter<W> { pub struct FakeTlsWriter<W> {
upstream: W, upstream: W,
state: TlsWriterState, state: TlsWriterState,
record_buffer: BytesMut,
} }
impl<W> FakeTlsWriter<W> { impl<W> FakeTlsWriter<W> {
@@ -683,6 +687,7 @@ impl<W> FakeTlsWriter<W> {
Self { Self {
upstream, upstream,
state: TlsWriterState::Idle, state: TlsWriterState::Idle,
record_buffer: BytesMut::new(),
} }
} }
@@ -704,10 +709,15 @@ impl<W> FakeTlsWriter<W> {
} }
pub fn has_pending(&self) -> bool { pub fn has_pending(&self) -> bool {
matches!(&self.state, TlsWriterState::WritingRecord { record, .. } if !record.is_empty()) matches!(
&self.state,
TlsWriterState::WritingRecord { position, .. }
if *position < self.record_buffer.len()
)
} }
fn poison(&mut self, error: io::Error) { fn poison(&mut self, error: io::Error) {
self.record_buffer = BytesMut::new();
self.state = TlsWriterState::Poisoned { error: Some(error) }; self.state = TlsWriterState::Poisoned { error: Some(error) };
} }
@@ -720,17 +730,26 @@ impl<W> FakeTlsWriter<W> {
} }
} }
fn build_record(data: &[u8]) -> BytesMut { fn build_record(&mut self, data: &[u8]) {
let header = TlsRecordHeader { let header = TlsRecordHeader {
record_type: TLS_RECORD_APPLICATION, record_type: TLS_RECORD_APPLICATION,
version: TLS_VERSION, version: TLS_VERSION,
length: data.len() as u16, length: data.len() as u16,
}; };
let mut record = BytesMut::with_capacity(TLS_HEADER_SIZE + data.len()); self.record_buffer.clear();
record.extend_from_slice(&header.to_bytes()); self.record_buffer.reserve(TLS_HEADER_SIZE + data.len());
record.extend_from_slice(data); self.record_buffer.extend_from_slice(&header.to_bytes());
record self.record_buffer.extend_from_slice(data);
debug_assert!(self.record_buffer.len() <= MAX_PENDING_WRITE);
}
fn recycle_record_buffer(&mut self) {
if self.record_buffer.capacity() > MAX_RETAINED_RECORD_CAPACITY {
self.record_buffer = BytesMut::new();
} else {
self.record_buffer.clear();
}
} }
} }
@@ -744,10 +763,11 @@ impl<W: AsyncWrite + Unpin> FakeTlsWriter<W> {
fn poll_flush_record_inner( fn poll_flush_record_inner(
upstream: &mut W, upstream: &mut W,
cx: &mut Context<'_>, cx: &mut Context<'_>,
record: &mut WriteBuffer, record: &[u8],
position: &mut usize,
) -> FlushResult { ) -> FlushResult {
while !record.is_empty() { while *position < record.len() {
let data = record.pending(); let data = &record[*position..];
match Pin::new(&mut *upstream).poll_write(cx, data) { match Pin::new(&mut *upstream).poll_write(cx, data) {
Poll::Pending => return FlushResult::Pending, Poll::Pending => return FlushResult::Pending,
Poll::Ready(Err(e)) => return FlushResult::Error(e), Poll::Ready(Err(e)) => return FlushResult::Error(e),
@@ -757,7 +777,7 @@ impl<W: AsyncWrite + Unpin> FakeTlsWriter<W> {
"upstream returned 0 bytes written", "upstream returned 0 bytes written",
)); ));
} }
Poll::Ready(Ok(n)) => record.advance(n), Poll::Ready(Ok(n)) => *position += n,
} }
} }
@@ -780,14 +800,19 @@ impl<W: AsyncWrite + Unpin> AsyncWrite for FakeTlsWriter<W> {
} }
TlsWriterState::WritingRecord { TlsWriterState::WritingRecord {
mut record, mut position,
payload_size, payload_size,
} => { } => {
// Finish writing previous record before accepting new bytes. // Finish writing previous record before accepting new bytes.
match Self::poll_flush_record_inner(&mut this.upstream, cx, &mut record) { match Self::poll_flush_record_inner(
&mut this.upstream,
cx,
&this.record_buffer,
&mut position,
) {
FlushResult::Pending => { FlushResult::Pending => {
this.state = TlsWriterState::WritingRecord { this.state = TlsWriterState::WritingRecord {
record, position,
payload_size, payload_size,
}; };
return Poll::Pending; return Poll::Pending;
@@ -797,6 +822,7 @@ impl<W: AsyncWrite + Unpin> AsyncWrite for FakeTlsWriter<W> {
return Poll::Ready(Err(e)); return Poll::Ready(Err(e));
} }
FlushResult::Complete(_) => { FlushResult::Complete(_) => {
this.recycle_record_buffer();
this.state = TlsWriterState::Idle; this.state = TlsWriterState::Idle;
// continue to accept new buf below // continue to accept new buf below
} }
@@ -818,19 +844,17 @@ impl<W: AsyncWrite + Unpin> AsyncWrite for FakeTlsWriter<W> {
let chunk = &buf[..chunk_size]; let chunk = &buf[..chunk_size];
// Build the complete record (header + payload) // Build the complete record (header + payload)
let record_data = Self::build_record(chunk); this.build_record(chunk);
match Pin::new(&mut this.upstream).poll_write(cx, &record_data) { match Pin::new(&mut this.upstream).poll_write(cx, &this.record_buffer) {
Poll::Ready(Ok(n)) if n == record_data.len() => Poll::Ready(Ok(chunk_size)), Poll::Ready(Ok(n)) if n == this.record_buffer.len() => {
this.recycle_record_buffer();
Poll::Ready(Ok(chunk_size))
}
Poll::Ready(Ok(n)) => { Poll::Ready(Ok(n)) => {
// Partial write of the record: store remainder.
let mut write_buffer = WriteBuffer::with_max_size(MAX_PENDING_WRITE);
// record_data length is <= 16389, fits MAX_PENDING_WRITE
let _ = write_buffer.extend(&record_data[n..]);
this.state = TlsWriterState::WritingRecord { this.state = TlsWriterState::WritingRecord {
record: write_buffer, position: n,
payload_size: chunk_size, payload_size: chunk_size,
}; };
@@ -844,12 +868,8 @@ impl<W: AsyncWrite + Unpin> AsyncWrite for FakeTlsWriter<W> {
} }
Poll::Pending => { Poll::Pending => {
// Buffer entire record and report success for this chunk.
let mut write_buffer = WriteBuffer::with_max_size(MAX_PENDING_WRITE);
let _ = write_buffer.extend(&record_data);
this.state = TlsWriterState::WritingRecord { this.state = TlsWriterState::WritingRecord {
record: write_buffer, position: 0,
payload_size: chunk_size, payload_size: chunk_size,
}; };
@@ -871,12 +891,17 @@ impl<W: AsyncWrite + Unpin> AsyncWrite for FakeTlsWriter<W> {
} }
TlsWriterState::WritingRecord { TlsWriterState::WritingRecord {
mut record, mut position,
payload_size, payload_size,
} => match Self::poll_flush_record_inner(&mut this.upstream, cx, &mut record) { } => match Self::poll_flush_record_inner(
&mut this.upstream,
cx,
&this.record_buffer,
&mut position,
) {
FlushResult::Pending => { FlushResult::Pending => {
this.state = TlsWriterState::WritingRecord { this.state = TlsWriterState::WritingRecord {
record, position,
payload_size, payload_size,
}; };
return Poll::Pending; return Poll::Pending;
@@ -886,6 +911,7 @@ impl<W: AsyncWrite + Unpin> AsyncWrite for FakeTlsWriter<W> {
return Poll::Ready(Err(e)); return Poll::Ready(Err(e));
} }
FlushResult::Complete(_) => { FlushResult::Complete(_) => {
this.recycle_record_buffer();
this.state = TlsWriterState::Idle; this.state = TlsWriterState::Idle;
} }
}, },
@@ -905,11 +931,17 @@ impl<W: AsyncWrite + Unpin> AsyncWrite for FakeTlsWriter<W> {
match state { match state {
TlsWriterState::WritingRecord { TlsWriterState::WritingRecord {
mut record, mut position,
payload_size: _, payload_size: _,
} => { } => {
// Best-effort flush (do not block shutdown forever). // Best-effort flush (do not block shutdown forever).
let _ = Self::poll_flush_record_inner(&mut this.upstream, cx, &mut record); let _ = Self::poll_flush_record_inner(
&mut this.upstream,
cx,
&this.record_buffer,
&mut position,
);
this.recycle_record_buffer();
this.state = TlsWriterState::Idle; this.state = TlsWriterState::Idle;
} }
_ => { _ => {
-605
View File
@@ -1,605 +0,0 @@
use std::collections::BTreeSet;
use std::net::IpAddr;
use std::path::PathBuf;
use std::sync::Arc;
use tokio::io::AsyncWriteExt;
use tokio::process::Command;
use tokio::sync::watch;
use tracing::warn;
use crate::config::{ProxyConfig, SynLimitMode};
const IPTABLES_CHAIN: &str = "TELEMT_SYNLIMIT";
const IPTABLES_HASHLIMIT_NAME: &str = "TELEMT-BUMPER";
const NFT_TABLE: &str = "telemt_synlimit";
const NFT_CHAIN: &str = "input";
type SynLimitTarget = (Option<IpAddr>, u16, u32, u32, u32);
#[derive(Default)]
struct SynLimitTargets {
iptables_v4: Vec<SynLimitTarget>,
iptables_v6: Vec<SynLimitTarget>,
nft_v4: Vec<SynLimitTarget>,
nft_v6: Vec<SynLimitTarget>,
}
#[derive(Clone, Copy)]
struct NftTableFamilies {
inet: bool,
ip: bool,
ip6: bool,
}
#[derive(Clone, Copy)]
enum NftFamily {
Inet,
Ip,
Ip6,
}
struct NftApplyPlan<'a> {
family: NftFamily,
v4_targets: &'a [SynLimitTarget],
v6_targets: &'a [SynLimitTarget],
}
impl SynLimitTargets {
fn is_empty(&self) -> bool {
self.iptables_v4.is_empty()
&& self.iptables_v6.is_empty()
&& self.nft_v4.is_empty()
&& self.nft_v6.is_empty()
}
fn has_iptables_targets(&self) -> bool {
!self.iptables_v4.is_empty() || !self.iptables_v6.is_empty()
}
fn has_nft_targets(&self) -> bool {
!self.nft_v4.is_empty() || !self.nft_v6.is_empty()
}
}
impl NftFamily {
fn as_str(self) -> &'static str {
match self {
Self::Inet => "inet",
Self::Ip => "ip",
Self::Ip6 => "ip6",
}
}
}
pub(crate) fn spawn_synlimit_controller(config_rx: watch::Receiver<Arc<ProxyConfig>>) {
if !cfg!(target_os = "linux") {
if has_synlimit_config(&config_rx.borrow()) {
warn!("SYN limiter is configured but unsupported on this OS; skipping netfilter rules");
}
return;
}
tokio::spawn(async move {
wait_for_config_channel_close_and_reconcile(config_rx).await;
if let Err(error) = clear_synlimit_rules_all_backends().await {
warn!(error = %error, "Failed to clear SYN limiter rules after config channel close");
}
});
}
async fn wait_for_config_channel_close_and_reconcile(
mut config_rx: watch::Receiver<Arc<ProxyConfig>>,
) {
while config_rx.changed().await.is_ok() {
let cfg = config_rx.borrow_and_update().clone();
reconcile_synlimit_rules(&cfg).await;
}
}
pub(crate) async fn reconcile_synlimit_rules(cfg: &ProxyConfig) {
if let Err(error) = clear_synlimit_rules_all_backends().await {
warn!(error = %error, "Failed to clear existing SYN limiter rules before reconcile");
}
let targets = synlimit_targets(cfg);
if targets.is_empty() {
return;
}
if !has_cap_net_admin() {
warn!(
"SYN limiter configured but CAP_NET_ADMIN is not available; netfilter rules not applied"
);
return;
}
if targets.has_iptables_targets()
&& let Err(error) = apply_iptables_synlimit_rules(&targets).await
{
warn!(error = %error, "Failed to apply iptables SYN limiter rules");
}
if targets.has_nft_targets()
&& let Err(error) = apply_nft_synlimit_rules(&targets).await
{
warn!(error = %error, "Failed to apply nftables SYN limiter rules");
}
}
pub(crate) async fn clear_synlimit_rules_all_backends() -> Result<(), String> {
if !has_cap_net_admin() {
return Ok(());
}
let mut errors = Vec::new();
if let Err(error) = clear_nft_synlimit_rules_all_families().await {
errors.push(error);
}
if let Err(error) = clear_iptables_synlimit_rules_for_binary("iptables").await {
errors.push(error);
}
if let Err(error) = clear_iptables_synlimit_rules_for_binary("ip6tables").await {
errors.push(error);
}
if errors.is_empty() {
Ok(())
} else {
Err(errors.join("; "))
}
}
fn has_synlimit_config(cfg: &ProxyConfig) -> bool {
cfg.server
.listeners
.iter()
.any(|listener| !matches!(listener.synlimit, SynLimitMode::Off))
}
fn synlimit_targets(cfg: &ProxyConfig) -> SynLimitTargets {
let mut iptables_v4 = BTreeSet::new();
let mut iptables_v6 = BTreeSet::new();
let mut nft_v4 = BTreeSet::new();
let mut nft_v6 = BTreeSet::new();
for listener in &cfg.server.listeners {
let backend = listener.synlimit;
if matches!(backend, SynLimitMode::Off) {
continue;
}
let port = listener.port.unwrap_or(cfg.server.port);
let ip = (!listener.ip.is_unspecified()).then_some(listener.ip);
let seconds = listener.synlimit_seconds;
let hitcount = listener.synlimit_hitcount;
let burst = listener.synlimit_burst;
match (backend, listener.ip.is_ipv4()) {
(SynLimitMode::Iptables, true) => {
iptables_v4.insert((ip, port, seconds, hitcount, burst));
}
(SynLimitMode::Iptables, false) => {
iptables_v6.insert((ip, port, seconds, hitcount, burst));
}
(SynLimitMode::Nftables, true) => {
nft_v4.insert((ip, port, seconds, hitcount, burst));
}
(SynLimitMode::Nftables, false) => {
nft_v6.insert((ip, port, seconds, hitcount, burst));
}
(SynLimitMode::Off, _) => {}
}
}
SynLimitTargets {
iptables_v4: iptables_v4.into_iter().collect(),
iptables_v6: iptables_v6.into_iter().collect(),
nft_v4: nft_v4.into_iter().collect(),
nft_v6: nft_v6.into_iter().collect(),
}
}
async fn apply_iptables_synlimit_rules(targets: &SynLimitTargets) -> Result<(), String> {
apply_iptables_synlimit_rules_for_binary("iptables", &targets.iptables_v4).await?;
apply_iptables_synlimit_rules_for_binary("ip6tables", &targets.iptables_v6).await
}
async fn apply_iptables_synlimit_rules_for_binary(
binary: &str,
targets: &[SynLimitTarget],
) -> Result<(), String> {
if targets.is_empty() {
return Ok(());
}
let _ = run_command(binary, &["-t", "filter", "-N", IPTABLES_CHAIN], None).await;
run_command(binary, &["-t", "filter", "-F", IPTABLES_CHAIN], None).await?;
if run_command(
binary,
&["-t", "filter", "-C", "INPUT", "-j", IPTABLES_CHAIN],
None,
)
.await
.is_err()
{
run_command(
binary,
&["-t", "filter", "-A", "INPUT", "-j", IPTABLES_CHAIN],
None,
)
.await?;
}
for (idx, (ip, port, seconds, hitcount, burst)) in targets.iter().enumerate() {
let hashlimit_name = format!("{IPTABLES_HASHLIMIT_NAME}-{idx}");
let accept_args = iptables_hashlimit_accept_rule_args(
ip,
*port,
*seconds,
*hitcount,
*burst,
&hashlimit_name,
);
let drop_args = iptables_synlimit_drop_rule_args(ip, *port);
let drop_refs: Vec<&str> = drop_args.iter().map(String::as_str).collect();
let accept_refs: Vec<&str> = accept_args.iter().map(String::as_str).collect();
run_command(binary, &accept_refs, None).await?;
run_command(binary, &drop_refs, None).await?;
}
run_command(
binary,
&["-t", "filter", "-A", IPTABLES_CHAIN, "-j", "RETURN"],
None,
)
.await?;
Ok(())
}
fn iptables_hashlimit_accept_rule_args(
ip: &Option<IpAddr>,
port: u16,
seconds: u32,
hitcount: u32,
burst: u32,
hashlimit_name: &str,
) -> Vec<String> {
let mut args = vec![
"-t".to_string(),
"filter".to_string(),
"-A".to_string(),
IPTABLES_CHAIN.to_string(),
"-p".to_string(),
"tcp".to_string(),
"--syn".to_string(),
];
if let Some(ip) = ip {
args.push("-d".to_string());
args.push(ip.to_string());
}
let rate = synlimit_rate_arg(seconds, hitcount);
args.extend([
"--dport".to_string(),
port.to_string(),
"-m".to_string(),
"hashlimit".to_string(),
"--hashlimit-name".to_string(),
hashlimit_name.to_string(),
"--hashlimit-mode".to_string(),
"srcip".to_string(),
"--hashlimit-upto".to_string(),
rate,
"--hashlimit-burst".to_string(),
burst.to_string(),
"--hashlimit-htable-expire".to_string(),
"15000".to_string(),
"-j".to_string(),
"ACCEPT".to_string(),
]);
args
}
fn iptables_synlimit_drop_rule_args(ip: &Option<IpAddr>, port: u16) -> Vec<String> {
let mut args = vec![
"-t".to_string(),
"filter".to_string(),
"-A".to_string(),
IPTABLES_CHAIN.to_string(),
"-p".to_string(),
"tcp".to_string(),
"--syn".to_string(),
];
if let Some(ip) = ip {
args.push("-d".to_string());
args.push(ip.to_string());
}
args.extend([
"--dport".to_string(),
port.to_string(),
"-j".to_string(),
"DROP".to_string(),
]);
args
}
fn synlimit_rate_arg(seconds: u32, hitcount: u32) -> String {
let seconds = u64::from(seconds.max(1));
let hitcount = u64::from(hitcount.max(1));
for (unit_seconds, unit_name) in [
(1_u64, "second"),
(60_u64, "minute"),
(3_600_u64, "hour"),
(86_400_u64, "day"),
] {
let amount = hitcount.saturating_mul(unit_seconds);
if amount >= seconds && amount % seconds == 0 {
return format!("{}/{}", amount / seconds, unit_name);
}
}
let amount = hitcount.saturating_mul(86_400).saturating_add(seconds - 1) / seconds;
format!("{}/day", amount.max(1))
}
async fn clear_iptables_synlimit_rules_for_binary(binary: &str) -> Result<(), String> {
let mut errors = Vec::new();
for _ in 0..8 {
match run_command(
binary,
&["-t", "filter", "-D", "INPUT", "-j", IPTABLES_CHAIN],
None,
)
.await
{
Ok(()) => {}
Err(error) if is_missing_command_or_iptables_rule(&error) => break,
Err(error) => {
errors.push(format!("{binary} delete INPUT jump failed: {error}"));
break;
}
}
}
if let Err(error) = run_command(binary, &["-t", "filter", "-F", IPTABLES_CHAIN], None).await
&& !is_missing_command_or_iptables_rule(&error)
{
errors.push(format!("{binary} flush chain failed: {error}"));
}
if let Err(error) = run_command(binary, &["-t", "filter", "-X", IPTABLES_CHAIN], None).await
&& !is_missing_command_or_iptables_rule(&error)
{
errors.push(format!("{binary} delete chain failed: {error}"));
}
if errors.is_empty() {
Ok(())
} else {
Err(errors.join(", "))
}
}
async fn apply_nft_synlimit_rules(targets: &SynLimitTargets) -> Result<(), String> {
let families = detect_nft_table_families().await;
for plan in nft_apply_plan(families, &targets.nft_v4, &targets.nft_v6) {
let script = nft_synlimit_script(plan);
run_command("nft", &["-f", "-"], Some(script)).await?;
}
Ok(())
}
async fn detect_nft_table_families() -> NftTableFamilies {
let Ok(output) = run_command_stdout("nft", &["list", "tables"]).await else {
return NftTableFamilies {
inet: false,
ip: false,
ip6: false,
};
};
let mut families = NftTableFamilies {
inet: false,
ip: false,
ip6: false,
};
for line in output.lines() {
let mut fields = line.split_whitespace();
if fields.next() != Some("table") {
continue;
}
match fields.next() {
Some("inet") => families.inet = true,
Some("ip") => families.ip = true,
Some("ip6") => families.ip6 = true,
_ => {}
}
}
families
}
fn nft_apply_plan<'a>(
families: NftTableFamilies,
v4_targets: &'a [SynLimitTarget],
v6_targets: &'a [SynLimitTarget],
) -> Vec<NftApplyPlan<'a>> {
if !v4_targets.is_empty() && !v6_targets.is_empty() {
return vec![NftApplyPlan {
family: NftFamily::Inet,
v4_targets,
v6_targets,
}];
}
if !v4_targets.is_empty() {
return vec![NftApplyPlan {
family: if families.inet || !families.ip {
NftFamily::Inet
} else {
NftFamily::Ip
},
v4_targets,
v6_targets: &[],
}];
}
if !v6_targets.is_empty() {
return vec![NftApplyPlan {
family: if families.inet || !families.ip6 {
NftFamily::Inet
} else {
NftFamily::Ip6
},
v4_targets: &[],
v6_targets,
}];
}
Vec::new()
}
fn nft_synlimit_script(plan: NftApplyPlan<'_>) -> String {
let mut script = String::new();
script.push_str(&format!("table {} {NFT_TABLE} {{\n", plan.family.as_str()));
script.push_str(&format!(" chain {NFT_CHAIN} {{\n"));
script.push_str(" type filter hook input priority filter; policy accept;\n");
for (idx, (ip, port, seconds, hitcount, burst)) in plan.v4_targets.iter().enumerate() {
let daddr = ip
.map(|ip| format!(" ip daddr {ip}"))
.unwrap_or_else(String::new);
let rate = synlimit_rate_arg(*seconds, *hitcount);
script.push_str(&format!(
" tcp flags & (fin|syn|rst|ack) == syn{daddr} tcp dport {port} meter telemt_synlimit_v4_{idx} {{ ip saddr limit rate over {rate} burst {burst} packets }} drop\n"
));
script.push_str(&format!(
" tcp flags & (fin|syn|rst|ack) == syn{daddr} tcp dport {port} accept\n"
));
}
for (idx, (ip, port, seconds, hitcount, burst)) in plan.v6_targets.iter().enumerate() {
let daddr = ip
.map(|ip| format!(" ip6 daddr {ip}"))
.unwrap_or_else(String::new);
let rate = synlimit_rate_arg(*seconds, *hitcount);
script.push_str(&format!(
" tcp flags & (fin|syn|rst|ack) == syn{daddr} tcp dport {port} meter telemt_synlimit_v6_{idx} {{ ip6 saddr limit rate over {rate} burst {burst} packets }} drop\n"
));
script.push_str(&format!(
" tcp flags & (fin|syn|rst|ack) == syn{daddr} tcp dport {port} accept\n"
));
}
script.push_str(" }\n");
script.push_str("}\n");
script
}
async fn clear_nft_synlimit_rules_all_families() -> Result<(), String> {
let mut errors = Vec::new();
for family in [NftFamily::Inet, NftFamily::Ip, NftFamily::Ip6] {
if let Err(error) = run_command(
"nft",
&["delete", "table", family.as_str(), NFT_TABLE],
None,
)
.await
&& !is_missing_command_or_nft_table(&error)
{
errors.push(format!(
"nft delete table {} {NFT_TABLE} failed: {error}",
family.as_str()
));
}
}
if errors.is_empty() {
Ok(())
} else {
Err(errors.join(", "))
}
}
fn is_missing_command_or_iptables_rule(error: &str) -> bool {
error.contains("is not available")
|| error.contains("No chain/target/match by that name")
|| error.contains("does not exist")
}
fn is_missing_command_or_nft_table(error: &str) -> bool {
error.contains("is not available") || error.contains("No such file or directory")
}
async fn run_command(binary: &str, args: &[&str], stdin: Option<String>) -> Result<(), String> {
let Some(command_path) = resolve_command(binary) else {
return Err(format!("{binary} is not available"));
};
let mut command = Command::new(command_path);
command.args(args);
if stdin.is_some() {
command.stdin(std::process::Stdio::piped());
}
command.stdout(std::process::Stdio::null());
command.stderr(std::process::Stdio::piped());
let mut child = command
.spawn()
.map_err(|e| format!("spawn {binary} failed: {e}"))?;
if let Some(blob) = stdin
&& let Some(mut writer) = child.stdin.take()
{
writer
.write_all(blob.as_bytes())
.await
.map_err(|e| format!("stdin write {binary} failed: {e}"))?;
}
let output = child
.wait_with_output()
.await
.map_err(|e| format!("wait {binary} failed: {e}"))?;
if output.status.success() {
return Ok(());
}
let stderr = String::from_utf8_lossy(&output.stderr).trim().to_string();
Err(if stderr.is_empty() {
format!("{binary} exited with status {}", output.status)
} else {
stderr
})
}
async fn run_command_stdout(binary: &str, args: &[&str]) -> Result<String, String> {
let Some(command_path) = resolve_command(binary) else {
return Err(format!("{binary} is not available"));
};
let output = Command::new(command_path)
.args(args)
.output()
.await
.map_err(|e| format!("wait {binary} failed: {e}"))?;
if output.status.success() {
return Ok(String::from_utf8_lossy(&output.stdout).to_string());
}
let stderr = String::from_utf8_lossy(&output.stderr).trim().to_string();
Err(if stderr.is_empty() {
format!("{binary} exited with status {}", output.status)
} else {
stderr
})
}
fn resolve_command(binary: &str) -> Option<PathBuf> {
let mut dirs = std::env::var_os("PATH")
.map(|path| std::env::split_paths(&path).collect::<Vec<_>>())
.unwrap_or_default();
dirs.extend(["/usr/sbin", "/sbin", "/usr/bin", "/bin"].map(PathBuf::from));
dirs.into_iter()
.map(|dir| dir.join(binary))
.find(|candidate| candidate.exists() && candidate.is_file())
}
fn has_cap_net_admin() -> bool {
#[cfg(target_os = "linux")]
{
let Ok(status) = std::fs::read_to_string("/proc/self/status") else {
return false;
};
for line in status.lines() {
if let Some(raw) = line.strip_prefix("CapEff:") {
let caps = raw.trim();
if let Ok(bits) = u64::from_str_radix(caps, 16) {
const CAP_NET_ADMIN_BIT: u64 = 12;
return (bits & (1u64 << CAP_NET_ADMIN_BIT)) != 0;
}
}
}
false
}
#[cfg(not(target_os = "linux"))]
{
false
}
}
+98
View File
@@ -0,0 +1,98 @@
use std::path::PathBuf;
use tokio::io::AsyncWriteExt;
use tokio::process::Command;
pub(super) async fn run_command(
binary: &str,
args: &[&str],
stdin: Option<String>,
) -> Result<(), String> {
let Some(command_path) = resolve_command(binary) else {
return Err(format!("{binary} is not available"));
};
let mut command = Command::new(command_path);
command.args(args);
if stdin.is_some() {
command.stdin(std::process::Stdio::piped());
}
command.stdout(std::process::Stdio::null());
command.stderr(std::process::Stdio::piped());
let mut child = command
.spawn()
.map_err(|e| format!("spawn {binary} failed: {e}"))?;
if let Some(blob) = stdin
&& let Some(mut writer) = child.stdin.take()
{
writer
.write_all(blob.as_bytes())
.await
.map_err(|e| format!("stdin write {binary} failed: {e}"))?;
}
let output = child
.wait_with_output()
.await
.map_err(|e| format!("wait {binary} failed: {e}"))?;
if output.status.success() {
return Ok(());
}
let stderr = String::from_utf8_lossy(&output.stderr).trim().to_string();
Err(if stderr.is_empty() {
format!("{binary} exited with status {}", output.status)
} else {
stderr
})
}
pub(super) async fn run_command_stdout(binary: &str, args: &[&str]) -> Result<String, String> {
let Some(command_path) = resolve_command(binary) else {
return Err(format!("{binary} is not available"));
};
let output = Command::new(command_path)
.args(args)
.output()
.await
.map_err(|e| format!("wait {binary} failed: {e}"))?;
if output.status.success() {
return Ok(String::from_utf8_lossy(&output.stdout).to_string());
}
let stderr = String::from_utf8_lossy(&output.stderr).trim().to_string();
Err(if stderr.is_empty() {
format!("{binary} exited with status {}", output.status)
} else {
stderr
})
}
fn resolve_command(binary: &str) -> Option<PathBuf> {
let mut dirs = std::env::var_os("PATH")
.map(|path| std::env::split_paths(&path).collect::<Vec<_>>())
.unwrap_or_default();
dirs.extend(["/usr/sbin", "/sbin", "/usr/bin", "/bin"].map(PathBuf::from));
dirs.into_iter()
.map(|dir| dir.join(binary))
.find(|candidate| candidate.exists() && candidate.is_file())
}
pub(super) fn has_cap_net_admin() -> bool {
#[cfg(target_os = "linux")]
{
let Ok(status) = std::fs::read_to_string("/proc/self/status") else {
return false;
};
for line in status.lines() {
if let Some(raw) = line.strip_prefix("CapEff:") {
let caps = raw.trim();
if let Ok(bits) = u64::from_str_radix(caps, 16) {
const CAP_NET_ADMIN_BIT: u64 = 12;
return (bits & (1u64 << CAP_NET_ADMIN_BIT)) != 0;
}
}
}
false
}
#[cfg(not(target_os = "linux"))]
{
false
}
}
+389
View File
@@ -0,0 +1,389 @@
use std::net::IpAddr;
use super::command::run_command;
use super::model::{SynLimitNamespace, SynLimitRule, SynLimitTargets, synlimit_rate_arg};
const IPV4_IOS_PACKET_LENGTH: u16 = 64;
const IPV6_IOS_PACKET_LENGTH: u16 = 84;
const IOS_TTL_LIMIT: u8 = 65;
#[derive(Clone, Copy)]
enum IpTablesFamily {
V4,
V6,
}
impl IpTablesFamily {
fn ios_packet_length(self) -> u16 {
match self {
Self::V4 => IPV4_IOS_PACKET_LENGTH,
Self::V6 => IPV6_IOS_PACKET_LENGTH,
}
}
fn ttl_match(self) -> [&'static str; 3] {
match self {
Self::V4 => ["-m", "ttl", "--ttl-lt"],
Self::V6 => ["-m", "hl", "--hl-lt"],
}
}
fn hashlimit_tag(self) -> &'static str {
match self {
Self::V4 => "4",
Self::V6 => "6",
}
}
}
pub(super) async fn apply_synlimit_rules(
targets: &SynLimitTargets,
namespace: &SynLimitNamespace,
) -> Result<(), String> {
apply_rules_for_binary(
"iptables",
&targets.iptables_v4,
IpTablesFamily::V4,
namespace,
)
.await?;
apply_rules_for_binary(
"ip6tables",
&targets.iptables_v6,
IpTablesFamily::V6,
namespace,
)
.await
}
async fn apply_rules_for_binary(
binary: &str,
targets: &[SynLimitRule],
family: IpTablesFamily,
namespace: &SynLimitNamespace,
) -> Result<(), String> {
if targets.is_empty() {
return Ok(());
}
let chain = namespace.iptables_chain.as_str();
let _ = run_command(binary, &["-t", "filter", "-N", chain], None).await;
run_command(binary, &["-t", "filter", "-F", chain], None).await?;
if run_command(binary, &["-t", "filter", "-C", "INPUT", "-j", chain], None)
.await
.is_err()
{
run_command(
binary,
&["-t", "filter", "-I", "INPUT", "1", "-j", chain],
None,
)
.await?;
}
for (idx, target) in targets.iter().enumerate() {
for rule in iptables_synfix_rule_args(target, idx, family, namespace) {
let refs: Vec<&str> = rule.iter().map(String::as_str).collect();
run_command(binary, &refs, None).await?;
}
}
run_command(binary, &["-t", "filter", "-A", chain, "-j", "RETURN"], None).await?;
Ok(())
}
fn iptables_synfix_rule_args(
target: &SynLimitRule,
idx: usize,
family: IpTablesFamily,
namespace: &SynLimitNamespace,
) -> Vec<Vec<String>> {
vec![
iptables_ios_accept_rule_args(target, idx, family, namespace),
iptables_ios_reject_rule_args(target, family, namespace),
iptables_generic_accept_rule_args(target, idx, family, namespace),
iptables_generic_reject_rule_args(target, namespace),
]
}
fn iptables_ios_accept_rule_args(
target: &SynLimitRule,
idx: usize,
family: IpTablesFamily,
namespace: &SynLimitNamespace,
) -> Vec<String> {
let hashlimit_name = format!(
"{}-I{}-{idx}",
namespace.iptables_hashlimit_prefix,
family.hashlimit_tag()
);
let mut args =
iptables_base_rule_args(namespace.iptables_chain.as_str(), target.ip, target.port);
args.extend(iptables_ios_match_args(family));
args.extend(iptables_hashlimit_args(
&hashlimit_name,
target.ios_seconds,
target.ios_hitcount,
target.ios_burst,
target.hashlimit_expire_ms,
target.hashlimit_size,
));
args.extend(["-j".to_string(), "ACCEPT".to_string()]);
args
}
fn iptables_ios_reject_rule_args(
target: &SynLimitRule,
family: IpTablesFamily,
namespace: &SynLimitNamespace,
) -> Vec<String> {
let mut args =
iptables_base_rule_args(namespace.iptables_chain.as_str(), target.ip, target.port);
args.extend(iptables_ios_match_args(family));
args.extend(iptables_reject_args());
args
}
fn iptables_generic_accept_rule_args(
target: &SynLimitRule,
idx: usize,
family: IpTablesFamily,
namespace: &SynLimitNamespace,
) -> Vec<String> {
let hashlimit_name = format!(
"{}-G{}-{idx}",
namespace.iptables_hashlimit_prefix,
family.hashlimit_tag()
);
let mut args =
iptables_base_rule_args(namespace.iptables_chain.as_str(), target.ip, target.port);
args.extend(iptables_hashlimit_args(
&hashlimit_name,
target.generic_seconds,
target.generic_hitcount,
target.generic_burst,
target.hashlimit_expire_ms,
target.hashlimit_size,
));
args.extend(["-j".to_string(), "ACCEPT".to_string()]);
args
}
fn iptables_generic_reject_rule_args(
target: &SynLimitRule,
namespace: &SynLimitNamespace,
) -> Vec<String> {
let mut args =
iptables_base_rule_args(namespace.iptables_chain.as_str(), target.ip, target.port);
args.extend(iptables_reject_args());
args
}
fn iptables_base_rule_args(chain: &str, ip: Option<IpAddr>, port: u16) -> Vec<String> {
let mut args = vec![
"-t".to_string(),
"filter".to_string(),
"-A".to_string(),
chain.to_string(),
"-p".to_string(),
"tcp".to_string(),
"--syn".to_string(),
"-m".to_string(),
"tcp".to_string(),
"--tcp-flags".to_string(),
"SYN".to_string(),
"SYN".to_string(),
];
if let Some(ip) = ip {
args.push("-d".to_string());
args.push(ip.to_string());
}
args.extend(["--dport".to_string(), port.to_string()]);
args
}
fn iptables_ios_match_args(family: IpTablesFamily) -> Vec<String> {
let mut args = vec![
"-m".to_string(),
"length".to_string(),
"--length".to_string(),
family.ios_packet_length().to_string(),
];
args.extend(family.ttl_match().map(str::to_string));
args.push(IOS_TTL_LIMIT.to_string());
args
}
fn iptables_hashlimit_args(
name: &str,
seconds: u32,
hitcount: u32,
burst: u32,
expire_ms: u32,
size: u32,
) -> Vec<String> {
vec![
"-m".to_string(),
"hashlimit".to_string(),
"--hashlimit-name".to_string(),
name.to_string(),
"--hashlimit-mode".to_string(),
"srcip".to_string(),
"--hashlimit-upto".to_string(),
synlimit_rate_arg(seconds, hitcount),
"--hashlimit-burst".to_string(),
burst.to_string(),
"--hashlimit-htable-expire".to_string(),
expire_ms.to_string(),
"--hashlimit-htable-size".to_string(),
size.to_string(),
]
}
fn iptables_reject_args() -> Vec<String> {
vec![
"-j".to_string(),
"REJECT".to_string(),
"--reject-with".to_string(),
"tcp-reset".to_string(),
]
}
pub(super) async fn clear_rules_for_binary(
binary: &str,
namespace: &SynLimitNamespace,
) -> Result<bool, String> {
let mut errors = Vec::new();
let mut removed = false;
let chain = namespace.iptables_chain.as_str();
for _ in 0..8 {
match run_command(binary, &["-t", "filter", "-D", "INPUT", "-j", chain], None).await {
Ok(()) => {
removed = true;
}
Err(error) if is_missing_command_or_iptables_rule(&error) => break,
Err(error) => {
errors.push(format!("{binary} delete INPUT jump failed: {error}"));
break;
}
}
}
match run_command(binary, &["-t", "filter", "-F", chain], None).await {
Ok(()) => {
removed = true;
}
Err(error) if is_missing_command_or_iptables_rule(&error) => {}
Err(error) => {
errors.push(format!("{binary} flush chain failed: {error}"));
}
}
match run_command(binary, &["-t", "filter", "-X", chain], None).await {
Ok(()) => {
removed = true;
}
Err(error) if is_missing_command_or_iptables_rule(&error) => {}
Err(error) => {
errors.push(format!("{binary} delete chain failed: {error}"));
}
}
if errors.is_empty() {
Ok(removed)
} else {
Err(errors.join(", "))
}
}
fn is_missing_command_or_iptables_rule(error: &str) -> bool {
error.contains("is not available")
|| error.contains("No chain/target/match by that name")
|| error.contains("does not exist")
|| error.contains("Couldn't load target")
}
#[cfg(test)]
mod tests {
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr};
use super::*;
use crate::synlimit_control::model::test_rule;
fn has_pair(args: &[String], key: &str, value: &str) -> bool {
args.windows(2)
.any(|pair| pair[0].as_str() == key && pair[1].as_str() == value)
}
fn has_key(args: &[String], key: &str) -> bool {
args.iter().any(|arg| arg == key)
}
fn test_namespace() -> SynLimitNamespace {
SynLimitNamespace {
nft_table: "telemt_synlimit_test".to_string(),
iptables_chain: "TMT_SYN_TEST".to_string(),
iptables_hashlimit_prefix: "TMTTEST".to_string(),
}
}
#[test]
fn iptables_rules_use_synfix_order_and_rejects() {
let target = test_rule(Some(IpAddr::V4(Ipv4Addr::new(203, 0, 113, 7))), 443);
let namespace = test_namespace();
let rules = iptables_synfix_rule_args(&target, 0, IpTablesFamily::V4, &namespace);
assert_eq!(rules.len(), 4);
assert!(has_pair(&rules[0], "-A", "TMT_SYN_TEST"));
assert!(has_pair(&rules[0], "--length", "64"));
assert!(has_pair(&rules[0], "--ttl-lt", "65"));
assert!(has_pair(&rules[0], "--hashlimit-upto", "12/second"));
assert!(has_pair(&rules[0], "--hashlimit-burst", "24"));
assert!(has_pair(&rules[0], "--hashlimit-htable-expire", "60000"));
assert!(has_pair(&rules[0], "--hashlimit-htable-size", "32768"));
assert!(has_pair(&rules[0], "-j", "ACCEPT"));
assert!(has_pair(&rules[1], "-j", "REJECT"));
assert!(has_pair(&rules[1], "--reject-with", "tcp-reset"));
assert!(has_pair(&rules[2], "--hashlimit-upto", "48/minute"));
assert!(has_pair(&rules[3], "--reject-with", "tcp-reset"));
}
#[test]
fn ip6tables_rules_use_ipv6_hoplimit_classifier() {
let target = test_rule(Some(IpAddr::V6(Ipv6Addr::LOCALHOST)), 443);
let namespace = test_namespace();
let rules = iptables_synfix_rule_args(&target, 0, IpTablesFamily::V6, &namespace);
assert!(has_pair(&rules[0], "--length", "84"));
assert!(has_pair(&rules[0], "--hl-lt", "65"));
assert!(has_pair(&rules[0], "-d", "::1"));
}
#[test]
fn iptables_missing_rule_errors_are_cleanup_benign() {
assert!(is_missing_command_or_iptables_rule(
"iptables is not available"
));
assert!(is_missing_command_or_iptables_rule(
"iptables: No chain/target/match by that name."
));
assert!(is_missing_command_or_iptables_rule(
"iptables: Chain TELEMT_SYNLIMIT does not exist."
));
assert!(is_missing_command_or_iptables_rule(
"Couldn't load target `TELEMT_SYNLIMIT': No such file or directory"
));
assert!(!is_missing_command_or_iptables_rule(
"iptables: Permission denied"
));
}
#[test]
fn iptables_wildcard_rule_omits_destination_match() {
let target = test_rule(None, 443);
let namespace = test_namespace();
let rules = iptables_synfix_rule_args(&target, 0, IpTablesFamily::V4, &namespace);
for rule in rules {
assert!(!has_key(&rule, "-d"));
assert!(has_pair(&rule, "--dport", "443"));
}
}
}
+418
View File
@@ -0,0 +1,418 @@
use std::sync::{Arc, Mutex};
use tokio::sync::watch;
use tokio_util::sync::CancellationToken;
use tracing::warn;
use crate::config::{ProxyConfig, SynLimitMode};
use crate::maestro::generation::RuntimeWatchState;
mod command;
mod iptables;
mod model;
mod nftables;
use self::command::has_cap_net_admin;
use self::model::{SynLimitNamespace, synlimit_namespace, synlimit_targets};
static ACTIVE_SYNLIMIT_NAMESPACE: Mutex<Option<SynLimitNamespace>> = Mutex::new(None);
/// Process-owned lifecycle handle for the SYN limiter reconciler.
pub(crate) struct SynlimitController {
shutdown: CancellationToken,
join: tokio::task::JoinHandle<()>,
}
impl SynlimitController {
/// Stops config observation after any in-flight reconcile completes.
pub(crate) async fn shutdown(self) {
self.shutdown.cancel();
let _ = self.join.await;
}
}
/// Spawns the process-scoped SYN limiter reconciler for active generations.
pub(crate) fn spawn_synlimit_controller(
runtime_watch_rx: watch::Receiver<Option<RuntimeWatchState>>,
) -> SynlimitController {
let shutdown = CancellationToken::new();
let join = if !cfg!(target_os = "linux") {
tokio::spawn(watch_active_runtime_configs(
runtime_watch_rx,
shutdown.clone(),
|_generation_id, cfg| async move {
if has_synlimit_config(&cfg) {
warn!(
"SYN limiter is configured but unsupported on this OS; skipping netfilter rules"
);
}
},
))
} else {
tokio::spawn(watch_active_runtime_configs(
runtime_watch_rx,
shutdown.clone(),
|_generation_id, cfg| async move {
reconcile_synlimit_rules(&cfg).await;
},
))
};
SynlimitController { shutdown, join }
}
async fn watch_active_runtime_configs<F, Fut>(
mut runtime_watch_rx: watch::Receiver<Option<RuntimeWatchState>>,
shutdown: CancellationToken,
mut on_config: F,
) where
F: FnMut(u64, Arc<ProxyConfig>) -> Fut,
Fut: std::future::Future<Output = ()>,
{
let mut current = loop {
if let Some(state) = runtime_watch_rx.borrow().clone() {
break state;
}
tokio::select! {
biased;
_ = shutdown.cancelled() => return,
changed = runtime_watch_rx.changed() => {
if changed.is_err() {
return;
}
}
}
};
if shutdown.is_cancelled() {
return;
}
let initial_config = current.config_rx.borrow().clone();
on_config(current.generation_id, initial_config).await;
loop {
tokio::select! {
biased;
_ = shutdown.cancelled() => break,
changed = runtime_watch_rx.changed() => {
if changed.is_err() {
break;
}
let Some(next) = runtime_watch_rx.borrow().clone() else {
continue;
};
if next.generation_id != current.generation_id {
current = next;
let config = current.config_rx.borrow().clone();
on_config(current.generation_id, config).await;
}
}
changed = current.config_rx.changed() => {
if changed.is_err() {
let Some(next) = wait_for_new_runtime(
&mut runtime_watch_rx,
current.generation_id,
&shutdown,
).await else {
break;
};
current = next;
let config = current.config_rx.borrow().clone();
on_config(current.generation_id, config).await;
continue;
}
let active_generation_id = runtime_watch_rx
.borrow()
.as_ref()
.map(|state| state.generation_id);
if active_generation_id == Some(current.generation_id) {
let cfg = current.config_rx.borrow_and_update().clone();
on_config(current.generation_id, cfg).await;
}
}
}
}
}
async fn wait_for_new_runtime(
runtime_watch_rx: &mut watch::Receiver<Option<RuntimeWatchState>>,
previous_generation_id: u64,
shutdown: &CancellationToken,
) -> Option<RuntimeWatchState> {
loop {
if let Some(state) = runtime_watch_rx.borrow().clone()
&& state.generation_id != previous_generation_id
{
return Some(state);
}
tokio::select! {
biased;
_ = shutdown.cancelled() => return None,
changed = runtime_watch_rx.changed() => {
if changed.is_err() {
return None;
}
}
}
}
}
pub(crate) async fn reconcile_synlimit_rules(cfg: &ProxyConfig) {
let targets = synlimit_targets(cfg);
let namespace = synlimit_namespace(&targets);
if let Some(previous_namespace) = set_active_synlimit_namespace(namespace.clone()) {
match clear_synlimit_rules_for_namespace(&previous_namespace).await {
Ok(true) => {
warn!("Removed previous SYN limiter namespace before reconcile");
}
Ok(false) => {}
Err(error) => {
warn!(error = %error, "Failed to clear previous SYN limiter namespace before reconcile");
}
}
}
if targets.is_empty() {
return;
}
let Some(namespace) = namespace else {
return;
};
if !has_cap_net_admin() {
warn!(
"SYN limiter configured but CAP_NET_ADMIN is not available; netfilter rules not applied"
);
return;
}
match clear_synlimit_rules_for_namespace(&namespace).await {
Ok(true) => {
warn!("Removed stale SYN limiter rules left by a previous run before reconcile");
}
Ok(false) => {}
Err(error) => {
warn!(error = %error, "Failed to clear stale SYN limiter rules before reconcile");
}
}
if targets.has_iptables_targets()
&& let Err(error) = iptables::apply_synlimit_rules(&targets, &namespace).await
{
warn!(error = %error, "Failed to apply iptables SYN limiter rules");
}
if targets.has_nft_targets()
&& let Err(error) = nftables::apply_synlimit_rules(&targets, &namespace).await
{
warn!(error = %error, "Failed to apply nftables SYN limiter rules");
}
}
pub(crate) async fn clear_synlimit_rules_all_backends() -> Result<bool, String> {
let Some(namespace) = take_active_synlimit_namespace() else {
return Ok(false);
};
clear_synlimit_rules_for_namespace(&namespace).await
}
async fn clear_synlimit_rules_for_namespace(namespace: &SynLimitNamespace) -> Result<bool, String> {
if !has_cap_net_admin() {
return Ok(false);
}
let mut errors = Vec::new();
let mut removed = false;
match nftables::clear_rules_all_families(namespace).await {
Ok(value) => {
removed |= value;
}
Err(error) => {
errors.push(error);
}
}
match iptables::clear_rules_for_binary("iptables", namespace).await {
Ok(value) => {
removed |= value;
}
Err(error) => {
errors.push(error);
}
}
match iptables::clear_rules_for_binary("ip6tables", namespace).await {
Ok(value) => {
removed |= value;
}
Err(error) => {
errors.push(error);
}
}
if errors.is_empty() {
Ok(removed)
} else {
Err(errors.join("; "))
}
}
fn set_active_synlimit_namespace(next: Option<SynLimitNamespace>) -> Option<SynLimitNamespace> {
match ACTIVE_SYNLIMIT_NAMESPACE.lock() {
Ok(mut active) => {
if *active == next {
None
} else {
std::mem::replace(&mut *active, next)
}
}
Err(error) => {
warn!(error = %error, "Failed to update active SYN limiter namespace");
None
}
}
}
fn take_active_synlimit_namespace() -> Option<SynLimitNamespace> {
match ACTIVE_SYNLIMIT_NAMESPACE.lock() {
Ok(mut active) => active.take(),
Err(error) => {
warn!(error = %error, "Failed to read active SYN limiter namespace");
None
}
}
}
fn has_synlimit_config(cfg: &ProxyConfig) -> bool {
cfg.server
.listeners
.iter()
.any(|listener| !matches!(listener.synlimit, SynLimitMode::Off))
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::time::Duration;
use tokio::sync::{Notify, mpsc};
fn runtime_state(
generation_id: u64,
max_connections: u32,
) -> (
RuntimeWatchState,
watch::Sender<Arc<ProxyConfig>>,
watch::Sender<bool>,
) {
let mut config = ProxyConfig::default();
config.server.max_connections = max_connections;
let (config_tx, config_rx) = watch::channel(Arc::new(config));
let (admission_tx, admission_rx) = watch::channel(true);
(
RuntimeWatchState {
generation_id,
config_rx,
admission_rx,
},
config_tx,
admission_tx,
)
}
#[tokio::test]
async fn config_watcher_ignores_retired_generation_updates() {
let (initial, initial_config_tx, _initial_admission_tx) = runtime_state(1, 10);
let (runtime_tx, runtime_rx) = watch::channel(Some(initial));
let (observed_tx, mut observed_rx) = mpsc::unbounded_channel();
let watcher = tokio::spawn(watch_active_runtime_configs(
runtime_rx,
CancellationToken::new(),
move |generation_id, cfg| {
let observed_tx = observed_tx.clone();
async move {
let _ = observed_tx.send((generation_id, cfg.server.max_connections));
}
},
));
assert_eq!(observed_rx.recv().await, Some((1, 10)));
let (next, next_config_tx, _next_admission_tx) = runtime_state(2, 20);
runtime_tx.send_replace(Some(next));
assert_eq!(observed_rx.recv().await, Some((2, 20)));
let mut stale = ProxyConfig::default();
stale.server.max_connections = 30;
initial_config_tx.send_replace(Arc::new(stale));
assert!(
tokio::time::timeout(Duration::from_millis(50), observed_rx.recv())
.await
.is_err()
);
let mut active = ProxyConfig::default();
active.server.max_connections = 40;
next_config_tx.send_replace(Arc::new(active));
assert_eq!(observed_rx.recv().await, Some((2, 40)));
watcher.abort();
}
#[tokio::test]
async fn shutdown_waits_for_inflight_reconcile_and_stops_future_updates() {
let (initial, config_tx, _admission_tx) = runtime_state(1, 10);
let (_runtime_tx, runtime_rx) = watch::channel(Some(initial));
let shutdown = CancellationToken::new();
let started = Arc::new(Notify::new());
let release = Arc::new(Notify::new());
let calls = Arc::new(AtomicUsize::new(0));
let started_callback = started.clone();
let release_callback = release.clone();
let calls_callback = calls.clone();
let watcher_shutdown = shutdown.clone();
let watcher = tokio::spawn(watch_active_runtime_configs(
runtime_rx,
watcher_shutdown,
move |_generation_id, _cfg| {
let started = started_callback.clone();
let release = release_callback.clone();
let calls = calls_callback.clone();
async move {
calls.fetch_add(1, Ordering::AcqRel);
started.notify_one();
release.notified().await;
}
},
));
started.notified().await;
shutdown.cancel();
tokio::task::yield_now().await;
assert!(!watcher.is_finished());
release.notify_one();
tokio::time::timeout(Duration::from_secs(1), watcher)
.await
.unwrap()
.unwrap();
drop(_runtime_tx);
let mut updated = ProxyConfig::default();
updated.server.max_connections = 20;
assert!(config_tx.send(Arc::new(updated)).is_err());
assert_eq!(calls.load(Ordering::Acquire), 1);
}
#[tokio::test]
async fn shutdown_before_start_skips_initial_reconcile() {
let (initial, _config_tx, _admission_tx) = runtime_state(1, 10);
let (_runtime_tx, runtime_rx) = watch::channel(Some(initial));
let shutdown = CancellationToken::new();
shutdown.cancel();
let calls = Arc::new(AtomicUsize::new(0));
let calls_callback = calls.clone();
watch_active_runtime_configs(runtime_rx, shutdown, move |_generation_id, _cfg| {
let calls = calls_callback.clone();
async move {
calls.fetch_add(1, Ordering::AcqRel);
}
})
.await;
assert_eq!(calls.load(Ordering::Acquire), 0);
}
}
+365
View File
@@ -0,0 +1,365 @@
use std::collections::BTreeSet;
use std::net::IpAddr;
use crate::config::{ProxyConfig, SynLimitMode};
#[derive(Clone, Copy, Debug, Eq, Ord, PartialEq, PartialOrd)]
pub(super) struct SynLimitRule {
pub(super) ip: Option<IpAddr>,
pub(super) port: u16,
pub(super) generic_seconds: u32,
pub(super) generic_hitcount: u32,
pub(super) generic_burst: u32,
pub(super) ios_seconds: u32,
pub(super) ios_hitcount: u32,
pub(super) ios_burst: u32,
pub(super) hashlimit_expire_ms: u32,
pub(super) hashlimit_size: u32,
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub(super) struct SynLimitNamespace {
pub(super) nft_table: String,
pub(super) iptables_chain: String,
pub(super) iptables_hashlimit_prefix: String,
}
#[derive(Default)]
pub(super) struct SynLimitTargets {
pub(super) iptables_v4: Vec<SynLimitRule>,
pub(super) iptables_v6: Vec<SynLimitRule>,
pub(super) nft_v4: Vec<SynLimitRule>,
pub(super) nft_v6: Vec<SynLimitRule>,
}
impl SynLimitTargets {
pub(super) fn is_empty(&self) -> bool {
self.iptables_v4.is_empty()
&& self.iptables_v6.is_empty()
&& self.nft_v4.is_empty()
&& self.nft_v6.is_empty()
}
pub(super) fn has_iptables_targets(&self) -> bool {
!self.iptables_v4.is_empty() || !self.iptables_v6.is_empty()
}
pub(super) fn has_nft_targets(&self) -> bool {
!self.nft_v4.is_empty() || !self.nft_v6.is_empty()
}
}
struct SynLimitNamespaceHasher {
value: u64,
}
impl SynLimitNamespaceHasher {
const OFFSET: u64 = 0xcbf2_9ce4_8422_2325;
const PRIME: u64 = 0x0000_0100_0000_01b3;
fn new() -> Self {
Self {
value: Self::OFFSET,
}
}
fn write(&mut self, bytes: &[u8]) {
for byte in bytes {
self.value ^= u64::from(*byte);
self.value = self.value.wrapping_mul(Self::PRIME);
}
}
fn write_u8(&mut self, value: u8) {
self.write(&[value]);
}
fn write_u16(&mut self, value: u16) {
self.write(&value.to_le_bytes());
}
fn write_u32(&mut self, value: u32) {
self.write(&value.to_le_bytes());
}
fn finish(self) -> u64 {
self.value
}
}
pub(super) fn synlimit_targets(cfg: &ProxyConfig) -> SynLimitTargets {
let mut iptables_v4 = BTreeSet::new();
let mut iptables_v6 = BTreeSet::new();
let mut nft_v4 = BTreeSet::new();
let mut nft_v6 = BTreeSet::new();
for listener in &cfg.server.listeners {
let backend = listener.synlimit;
if matches!(backend, SynLimitMode::Off) {
continue;
}
let target = SynLimitRule {
ip: (!listener.ip.is_unspecified()).then_some(listener.ip),
port: listener.port.unwrap_or(cfg.server.port),
generic_seconds: listener.synlimit_seconds,
generic_hitcount: listener.synlimit_hitcount,
generic_burst: listener.synlimit_burst,
ios_seconds: listener.synlimit_ios_seconds,
ios_hitcount: listener.synlimit_ios_hitcount,
ios_burst: listener.synlimit_ios_burst,
hashlimit_expire_ms: listener.synlimit_hashlimit_expire_ms,
hashlimit_size: listener.synlimit_hashlimit_size,
};
match (backend, listener.ip.is_ipv4()) {
(SynLimitMode::Iptables, true) => {
iptables_v4.insert(target);
}
(SynLimitMode::Iptables, false) => {
iptables_v6.insert(target);
}
(SynLimitMode::Nftables, true) => {
nft_v4.insert(target);
}
(SynLimitMode::Nftables, false) => {
nft_v6.insert(target);
}
(SynLimitMode::Off, _) => {}
}
}
SynLimitTargets {
iptables_v4: iptables_v4.into_iter().collect(),
iptables_v6: iptables_v6.into_iter().collect(),
nft_v4: nft_v4.into_iter().collect(),
nft_v6: nft_v6.into_iter().collect(),
}
}
pub(super) fn synlimit_namespace(targets: &SynLimitTargets) -> Option<SynLimitNamespace> {
if targets.is_empty() {
return None;
}
let mut hasher = SynLimitNamespaceHasher::new();
write_namespace_rule_group(&mut hasher, b"iptables-v4", &targets.iptables_v4);
write_namespace_rule_group(&mut hasher, b"iptables-v6", &targets.iptables_v6);
write_namespace_rule_group(&mut hasher, b"nft-v4", &targets.nft_v4);
write_namespace_rule_group(&mut hasher, b"nft-v6", &targets.nft_v6);
let suffix = format!("{:016x}", hasher.finish());
let iptables_suffix = &suffix[..12];
let hashlimit_suffix = &suffix[..10];
Some(SynLimitNamespace {
nft_table: format!("telemt_synlimit_{suffix}"),
iptables_chain: format!("TMT_SYN_{iptables_suffix}"),
iptables_hashlimit_prefix: format!("TMT{hashlimit_suffix}"),
})
}
fn write_namespace_rule_group(
hasher: &mut SynLimitNamespaceHasher,
group: &[u8],
rules: &[SynLimitRule],
) {
hasher.write(group);
hasher.write_u32(rules.len() as u32);
for rule in rules {
write_namespace_rule(hasher, rule);
}
}
fn write_namespace_rule(hasher: &mut SynLimitNamespaceHasher, rule: &SynLimitRule) {
match rule.ip {
Some(IpAddr::V4(ip)) => {
hasher.write_u8(4);
hasher.write(&ip.octets());
}
Some(IpAddr::V6(ip)) => {
hasher.write_u8(6);
hasher.write(&ip.octets());
}
None => {
hasher.write_u8(0);
}
}
hasher.write_u16(rule.port);
hasher.write_u32(rule.generic_seconds);
hasher.write_u32(rule.generic_hitcount);
hasher.write_u32(rule.generic_burst);
hasher.write_u32(rule.ios_seconds);
hasher.write_u32(rule.ios_hitcount);
hasher.write_u32(rule.ios_burst);
hasher.write_u32(rule.hashlimit_expire_ms);
hasher.write_u32(rule.hashlimit_size);
}
pub(super) fn synlimit_rate_arg(seconds: u32, hitcount: u32) -> String {
let seconds = u64::from(seconds.max(1));
let hitcount = u64::from(hitcount.max(1));
for (unit_seconds, unit_name) in [
(1_u64, "second"),
(60_u64, "minute"),
(3_600_u64, "hour"),
(86_400_u64, "day"),
] {
let amount = hitcount.saturating_mul(unit_seconds);
if amount >= seconds && amount % seconds == 0 {
return format!("{}/{}", amount / seconds, unit_name);
}
}
let amount = hitcount.saturating_mul(86_400).saturating_add(seconds - 1) / seconds;
format!("{}/day", amount.max(1))
}
#[cfg(test)]
pub(super) fn test_rule(ip: Option<IpAddr>, port: u16) -> SynLimitRule {
SynLimitRule {
ip,
port,
generic_seconds: 60,
generic_hitcount: 48,
generic_burst: 1,
ios_seconds: 1,
ios_hitcount: 12,
ios_burst: 24,
hashlimit_expire_ms: 60_000,
hashlimit_size: 32_768,
}
}
#[cfg(test)]
mod tests {
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr};
use super::*;
use crate::config::ListenerConfig;
fn listener(ip: IpAddr, port: Option<u16>, synlimit: SynLimitMode) -> ListenerConfig {
ListenerConfig {
ip,
port,
client_mss: None,
synlimit,
synlimit_seconds: 60,
synlimit_hitcount: 48,
synlimit_burst: 1,
synlimit_ios_seconds: 1,
synlimit_ios_hitcount: 12,
synlimit_ios_burst: 24,
synlimit_hashlimit_expire_ms: 60_000,
synlimit_hashlimit_size: 32_768,
announce: None,
announce_ip: None,
proxy_protocol: None,
reuse_allow: false,
}
}
#[test]
fn synlimit_targets_deduplicate_and_use_legacy_port_fallback() {
let mut cfg = ProxyConfig::default();
cfg.server.port = 9443;
cfg.server.listeners = vec![
listener(
IpAddr::V4(Ipv4Addr::UNSPECIFIED),
None,
SynLimitMode::Iptables,
),
listener(
IpAddr::V4(Ipv4Addr::UNSPECIFIED),
None,
SynLimitMode::Iptables,
),
];
let targets = synlimit_targets(&cfg);
assert_eq!(targets.iptables_v4.len(), 1);
assert_eq!(targets.iptables_v4[0].ip, None);
assert_eq!(targets.iptables_v4[0].port, 9443);
assert!(targets.iptables_v6.is_empty());
assert!(targets.nft_v4.is_empty());
assert!(targets.nft_v6.is_empty());
}
#[test]
fn synlimit_targets_separate_backends_and_ip_families() {
let mut cfg = ProxyConfig::default();
cfg.server.listeners = vec![
listener(
IpAddr::V4(Ipv4Addr::new(203, 0, 113, 1)),
Some(443),
SynLimitMode::Iptables,
),
listener(
IpAddr::V6(Ipv6Addr::LOCALHOST),
Some(443),
SynLimitMode::Iptables,
),
listener(
IpAddr::V4(Ipv4Addr::new(203, 0, 113, 2)),
Some(444),
SynLimitMode::Nftables,
),
listener(
IpAddr::V6(Ipv6Addr::UNSPECIFIED),
Some(444),
SynLimitMode::Nftables,
),
];
let targets = synlimit_targets(&cfg);
assert_eq!(targets.iptables_v4.len(), 1);
assert_eq!(targets.iptables_v6.len(), 1);
assert_eq!(targets.nft_v4.len(), 1);
assert_eq!(targets.nft_v6.len(), 1);
assert_eq!(
targets.iptables_v4[0].ip,
Some(IpAddr::V4(Ipv4Addr::new(203, 0, 113, 1)))
);
assert_eq!(
targets.iptables_v6[0].ip,
Some(IpAddr::V6(Ipv6Addr::LOCALHOST))
);
assert_eq!(
targets.nft_v4[0].ip,
Some(IpAddr::V4(Ipv4Addr::new(203, 0, 113, 2)))
);
assert_eq!(targets.nft_v6[0].ip, None);
}
#[test]
fn synlimit_namespace_is_stable_and_changes_by_targets() {
let mut cfg = ProxyConfig::default();
cfg.server.listeners = vec![listener(
IpAddr::V4(Ipv4Addr::new(203, 0, 113, 1)),
Some(443),
SynLimitMode::Nftables,
)];
let first = synlimit_namespace(&synlimit_targets(&cfg))
.expect("configured targets must have a namespace");
let second = synlimit_namespace(&synlimit_targets(&cfg))
.expect("configured targets must have a namespace");
cfg.server.listeners[0].port = Some(444);
let changed = synlimit_namespace(&synlimit_targets(&cfg))
.expect("configured targets must have a namespace");
assert_eq!(first, second);
assert_ne!(first, changed);
assert!(first.nft_table.starts_with("telemt_synlimit_"));
assert!(first.iptables_chain.starts_with("TMT_SYN_"));
assert!(first.iptables_chain.len() <= 28);
assert!(first.iptables_hashlimit_prefix.starts_with("TMT"));
}
#[test]
fn synlimit_rate_arg_uses_native_units_without_fractional_rates() {
assert_eq!(synlimit_rate_arg(1, 12), "12/second");
assert_eq!(synlimit_rate_arg(60, 48), "48/minute");
assert_eq!(synlimit_rate_arg(3600, 121), "121/hour");
assert_eq!(synlimit_rate_arg(86400, 241), "241/day");
}
}
+316
View File
@@ -0,0 +1,316 @@
use super::command::{run_command, run_command_stdout};
use super::model::{SynLimitNamespace, SynLimitRule, SynLimitTargets, synlimit_rate_arg};
const NFT_CHAIN: &str = "input";
const NFT_INPUT_PRIORITY: i16 = -5;
const IPV4_IOS_PACKET_LENGTH: u16 = 64;
const IPV6_IOS_PACKET_LENGTH: u16 = 84;
const IOS_TTL_LIMIT: u8 = 65;
#[derive(Clone, Copy)]
struct NftTableFamilies {
inet: bool,
ip: bool,
ip6: bool,
}
#[derive(Clone, Copy)]
enum NftFamily {
Inet,
Ip,
Ip6,
}
struct NftApplyPlan<'a> {
family: NftFamily,
v4_targets: &'a [SynLimitRule],
v6_targets: &'a [SynLimitRule],
}
impl NftFamily {
fn as_str(self) -> &'static str {
match self {
Self::Inet => "inet",
Self::Ip => "ip",
Self::Ip6 => "ip6",
}
}
}
pub(super) async fn apply_synlimit_rules(
targets: &SynLimitTargets,
namespace: &SynLimitNamespace,
) -> Result<(), String> {
let families = detect_nft_table_families().await;
for plan in nft_apply_plan(families, &targets.nft_v4, &targets.nft_v6) {
let script = nft_synlimit_script(plan, namespace);
run_command("nft", &["-f", "-"], Some(script)).await?;
}
Ok(())
}
async fn detect_nft_table_families() -> NftTableFamilies {
let Ok(output) = run_command_stdout("nft", &["list", "tables"]).await else {
return NftTableFamilies {
inet: false,
ip: false,
ip6: false,
};
};
let mut families = NftTableFamilies {
inet: false,
ip: false,
ip6: false,
};
for line in output.lines() {
let mut fields = line.split_whitespace();
if fields.next() != Some("table") {
continue;
}
match fields.next() {
Some("inet") => families.inet = true,
Some("ip") => families.ip = true,
Some("ip6") => families.ip6 = true,
_ => {}
}
}
families
}
fn nft_apply_plan<'a>(
families: NftTableFamilies,
v4_targets: &'a [SynLimitRule],
v6_targets: &'a [SynLimitRule],
) -> Vec<NftApplyPlan<'a>> {
if !v4_targets.is_empty() && !v6_targets.is_empty() {
return vec![NftApplyPlan {
family: NftFamily::Inet,
v4_targets,
v6_targets,
}];
}
if !v4_targets.is_empty() {
return vec![NftApplyPlan {
family: if families.inet || !families.ip {
NftFamily::Inet
} else {
NftFamily::Ip
},
v4_targets,
v6_targets: &[],
}];
}
if !v6_targets.is_empty() {
return vec![NftApplyPlan {
family: if families.inet || !families.ip6 {
NftFamily::Inet
} else {
NftFamily::Ip6
},
v4_targets: &[],
v6_targets,
}];
}
Vec::new()
}
fn nft_synlimit_script(plan: NftApplyPlan<'_>, namespace: &SynLimitNamespace) -> String {
let mut script = String::new();
script.push_str(&format!(
"table {} {} {{\n",
plan.family.as_str(),
namespace.nft_table
));
script.push_str(&format!(" chain {NFT_CHAIN} {{\n"));
script.push_str(&format!(
" type filter hook input priority {NFT_INPUT_PRIORITY}; policy accept;\n"
));
for (idx, target) in plan.v4_targets.iter().enumerate() {
push_nft_v4_rules(&mut script, target, idx);
}
for (idx, target) in plan.v6_targets.iter().enumerate() {
push_nft_v6_rules(&mut script, target, idx);
}
script.push_str(" }\n");
script.push_str("}\n");
script
}
fn push_nft_v4_rules(script: &mut String, target: &SynLimitRule, idx: usize) {
let daddr = target
.ip
.map(|ip| format!(" ip daddr {ip}"))
.unwrap_or_default();
let ios_rate = synlimit_rate_arg(target.ios_seconds, target.ios_hitcount);
let generic_rate = synlimit_rate_arg(target.generic_seconds, target.generic_hitcount);
script.push_str(&format!(
" tcp flags & (fin|syn|rst|ack) == syn{daddr} meta length {IPV4_IOS_PACKET_LENGTH} ip ttl < {IOS_TTL_LIMIT} tcp dport {port} meter telemt_synfix_ios_v4_{idx} {{ ip saddr limit rate over {ios_rate} burst {ios_burst} packets }} reject with tcp reset\n",
port = target.port,
ios_burst = target.ios_burst,
));
script.push_str(&format!(
" tcp flags & (fin|syn|rst|ack) == syn{daddr} meta length {IPV4_IOS_PACKET_LENGTH} ip ttl < {IOS_TTL_LIMIT} tcp dport {port} accept\n",
port = target.port,
));
script.push_str(&format!(
" tcp flags & (fin|syn|rst|ack) == syn{daddr} tcp dport {port} meter telemt_synfix_v4_{idx} {{ ip saddr limit rate over {generic_rate} burst {generic_burst} packets }} reject with tcp reset\n",
port = target.port,
generic_burst = target.generic_burst,
));
script.push_str(&format!(
" tcp flags & (fin|syn|rst|ack) == syn{daddr} tcp dport {port} accept\n",
port = target.port,
));
}
fn push_nft_v6_rules(script: &mut String, target: &SynLimitRule, idx: usize) {
let daddr = target
.ip
.map(|ip| format!(" ip6 daddr {ip}"))
.unwrap_or_default();
let ios_rate = synlimit_rate_arg(target.ios_seconds, target.ios_hitcount);
let generic_rate = synlimit_rate_arg(target.generic_seconds, target.generic_hitcount);
script.push_str(&format!(
" tcp flags & (fin|syn|rst|ack) == syn{daddr} meta length {IPV6_IOS_PACKET_LENGTH} ip6 hoplimit < {IOS_TTL_LIMIT} tcp dport {port} meter telemt_synfix_ios_v6_{idx} {{ ip6 saddr limit rate over {ios_rate} burst {ios_burst} packets }} reject with tcp reset\n",
port = target.port,
ios_burst = target.ios_burst,
));
script.push_str(&format!(
" tcp flags & (fin|syn|rst|ack) == syn{daddr} meta length {IPV6_IOS_PACKET_LENGTH} ip6 hoplimit < {IOS_TTL_LIMIT} tcp dport {port} accept\n",
port = target.port,
));
script.push_str(&format!(
" tcp flags & (fin|syn|rst|ack) == syn{daddr} tcp dport {port} meter telemt_synfix_v6_{idx} {{ ip6 saddr limit rate over {generic_rate} burst {generic_burst} packets }} reject with tcp reset\n",
port = target.port,
generic_burst = target.generic_burst,
));
script.push_str(&format!(
" tcp flags & (fin|syn|rst|ack) == syn{daddr} tcp dport {port} accept\n",
port = target.port,
));
}
pub(super) async fn clear_rules_all_families(
namespace: &SynLimitNamespace,
) -> Result<bool, String> {
let mut errors = Vec::new();
let mut removed = false;
let table = namespace.nft_table.as_str();
for family in [NftFamily::Inet, NftFamily::Ip, NftFamily::Ip6] {
match run_command("nft", &["delete", "table", family.as_str(), table], None).await {
Ok(()) => {
removed = true;
}
Err(error) if is_missing_command_or_nft_table(&error) => {}
Err(error) => {
errors.push(format!(
"nft delete table {} {table} failed: {error}",
family.as_str(),
));
}
}
}
if errors.is_empty() {
Ok(removed)
} else {
Err(errors.join(", "))
}
}
fn is_missing_command_or_nft_table(error: &str) -> bool {
error.contains("is not available") || error.contains("No such file or directory")
}
#[cfg(test)]
mod tests {
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr};
use super::*;
use crate::synlimit_control::model::test_rule;
fn test_namespace(table: &str) -> SynLimitNamespace {
SynLimitNamespace {
nft_table: table.to_string(),
iptables_chain: "TMT_SYN_TEST".to_string(),
iptables_hashlimit_prefix: "TMTTEST".to_string(),
}
}
#[test]
fn nft_script_uses_synfix_v4_rules_and_early_priority() {
let rule = test_rule(Some(IpAddr::V4(Ipv4Addr::new(203, 0, 113, 7))), 443);
let namespace = test_namespace("telemt_synlimit_test_a");
let script = nft_synlimit_script(
NftApplyPlan {
family: NftFamily::Inet,
v4_targets: &[rule],
v6_targets: &[],
},
&namespace,
);
assert!(script.contains("table inet telemt_synlimit_test_a"));
assert!(script.contains("type filter hook input priority -5; policy accept;"));
assert!(script.contains("ip daddr 203.0.113.7"));
assert!(script.contains("meta length 64 ip ttl < 65"));
assert!(script.contains("limit rate over 12/second burst 24 packets"));
assert!(script.contains("limit rate over 48/minute burst 1 packets"));
assert!(script.contains("reject with tcp reset"));
}
#[test]
fn nft_script_uses_ipv6_hoplimit_classifier() {
let rule = test_rule(Some(IpAddr::V6(Ipv6Addr::LOCALHOST)), 443);
let namespace = test_namespace("telemt_synlimit_test_b");
let script = nft_synlimit_script(
NftApplyPlan {
family: NftFamily::Inet,
v4_targets: &[],
v6_targets: &[rule],
},
&namespace,
);
assert!(script.contains("table inet telemt_synlimit_test_b"));
assert!(script.contains("ip6 daddr ::1"));
assert!(script.contains("meta length 84 ip6 hoplimit < 65"));
assert!(script.contains("ip6 saddr limit rate over 12/second burst 24 packets"));
assert!(script.contains("ip6 saddr limit rate over 48/minute burst 1 packets"));
}
#[test]
fn nft_missing_table_errors_are_cleanup_benign() {
assert!(is_missing_command_or_nft_table("nft is not available"));
assert!(is_missing_command_or_nft_table(
"Error: No such file or directory"
));
assert!(!is_missing_command_or_nft_table(
"Error: Operation not permitted"
));
}
#[test]
fn nft_apply_plan_keeps_dual_stack_rules_in_inet_table() {
let v4_rule = test_rule(Some(IpAddr::V4(Ipv4Addr::new(203, 0, 113, 7))), 443);
let v6_rule = test_rule(Some(IpAddr::V6(Ipv6Addr::LOCALHOST)), 443);
let v4_rules = [v4_rule];
let v6_rules = [v6_rule];
let plans = nft_apply_plan(
NftTableFamilies {
inet: false,
ip: false,
ip6: false,
},
&v4_rules,
&v6_rules,
);
assert_eq!(plans.len(), 1);
assert_eq!(plans[0].family.as_str(), "inet");
assert_eq!(plans[0].v4_targets, v4_rules.as_slice());
assert_eq!(plans[0].v6_targets, v6_rules.as_slice());
}
}
+29
View File
@@ -190,6 +190,18 @@ impl TlsFrontCache {
(snapshot, suppressed) (snapshot, suppressed)
} }
/// Returns configured domains that still resolve to the synthetic default profile.
pub(crate) async fn default_profile_domains(&self, domains: &[String]) -> Vec<String> {
let guard = self.memory.read().await;
domains
.iter()
.filter(|domain| {
guard.get(domain.as_str()).unwrap_or(&self.default).domain == "default"
})
.cloned()
.collect()
}
fn full_cert_sent_shard_index(client_ip: IpAddr) -> usize { fn full_cert_sent_shard_index(client_ip: IpAddr) -> usize {
let mut hasher = DefaultHasher::new(); let mut hasher = DefaultHasher::new();
client_ip.hash(&mut hasher); client_ip.hash(&mut hasher);
@@ -546,6 +558,23 @@ mod tests {
assert!(!cert_info_matches_domain(&cached)); assert!(!cert_info_matches_domain(&cached));
} }
#[tokio::test]
async fn default_profile_domains_reports_only_unprepared_entries() {
let domains = vec!["ready.example".to_string(), "pending.example".to_string()];
let cache = TlsFrontCache::new(&domains, 1024, "tlsfront-test-cache");
cache
.set(
"ready.example",
cached_with_cert_info("ready.example", None, Vec::new()),
)
.await;
assert_eq!(
cache.default_profile_domains(&domains).await,
vec!["pending.example".to_string()]
);
}
#[tokio::test] #[tokio::test]
async fn test_take_full_cert_budget_for_ip_uses_ttl() { async fn test_take_full_cert_budget_for_ip_uses_ttl() {
let cache = TlsFrontCache::new(&["example.com".to_string()], 1024, "tlsfront-test-cache"); let cache = TlsFrontCache::new(&["example.com".to_string()], 1024, "tlsfront-test-cache");
+5 -6
View File
@@ -916,11 +916,8 @@ async fn connect_tcp_with_upstream(
strict_route: bool, strict_route: bool,
) -> Result<UpstreamStream> { ) -> Result<UpstreamStream> {
if let Some(manager) = upstream { if let Some(manager) = upstream {
let resolved = if let Some(addr) = resolve_socket_addr(host, port) { let resolved = match manager.resolve_hostname(host, port).await {
Some(addr) Ok(addr) => Some(addr),
} else {
match tokio::net::lookup_host((host, port)).await {
Ok(mut addrs) => addrs.find(|a| a.is_ipv4()),
Err(e) => { Err(e) => {
if strict_route { if strict_route {
return Err(anyhow!( return Err(anyhow!(
@@ -936,7 +933,6 @@ async fn connect_tcp_with_upstream(
); );
None None
} }
}
}; };
if let Some(addr) = resolved { if let Some(addr) = resolved {
@@ -955,6 +951,9 @@ async fn connect_tcp_with_upstream(
error = %e, error = %e,
"Upstream connect failed, using direct connect" "Upstream connect failed, using direct connect"
); );
return Ok(UpstreamStream::Tcp(
timeout(connect_timeout, TcpStream::connect(addr)).await??,
));
} }
} }
} else if strict_route { } else if strict_route {
+100 -3
View File
@@ -1,9 +1,12 @@
use bytes::Bytes; use bytes::Bytes;
use std::sync::Arc;
use tokio::io::{AsyncReadExt, AsyncWriteExt}; use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::sync::OwnedSemaphorePermit;
use crate::crypto::{AesCbc, crc32, crc32c}; use crate::crypto::{AesCbc, crc32, crc32c};
use crate::error::{ProxyError, Result}; use crate::error::{ProxyError, Result};
use crate::protocol::constants::*; use crate::protocol::constants::*;
use crate::stats::Stats;
use crate::stream::PooledBuffer; use crate::stream::PooledBuffer;
use super::wire::{append_proxy_req_payload_into, proxy_req_payload_len}; use super::wire::{append_proxy_req_payload_into, proxy_req_payload_len};
@@ -11,9 +14,66 @@ use super::wire::{append_proxy_req_payload_into, proxy_req_payload_len};
const RPC_WRITER_FRAME_BUF_SHRINK_THRESHOLD: usize = 256 * 1024; const RPC_WRITER_FRAME_BUF_SHRINK_THRESHOLD: usize = 256 * 1024;
const RPC_WRITER_FRAME_BUF_RETAIN: usize = 64 * 1024; const RPC_WRITER_FRAME_BUF_RETAIN: usize = 64 * 1024;
enum WriterBytePermitState {
Queued,
Inflight,
}
/// Holds one writer memory reservation through queueing and socket write completion.
pub(crate) struct WriterBytePermit {
_permit: OwnedSemaphorePermit,
reserved_bytes: u64,
stats: Arc<Stats>,
state: WriterBytePermitState,
}
impl WriterBytePermit {
/// Creates a queued reservation and publishes its rounded resident-byte cost.
pub(crate) fn new(
permit: OwnedSemaphorePermit,
reserved_bytes: usize,
stats: Arc<Stats>,
) -> Self {
let reserved_bytes = reserved_bytes as u64;
stats.add_me_writer_byte_budget_queued_bytes(reserved_bytes);
Self {
_permit: permit,
reserved_bytes,
stats,
state: WriterBytePermitState::Queued,
}
}
/// Moves the reservation from queued to in-flight accounting exactly once.
pub(crate) fn mark_inflight(&mut self) {
if matches!(self.state, WriterBytePermitState::Queued) {
self.stats
.move_me_writer_byte_budget_to_inflight(self.reserved_bytes);
self.state = WriterBytePermitState::Inflight;
}
}
}
impl Drop for WriterBytePermit {
fn drop(&mut self) {
match self.state {
WriterBytePermitState::Queued => self
.stats
.release_me_writer_byte_budget_queued_bytes(self.reserved_bytes),
WriterBytePermitState::Inflight => self
.stats
.release_me_writer_byte_budget_inflight_bytes(self.reserved_bytes),
}
}
}
/// Commands sent to dedicated writer tasks to avoid mutex contention on TCP writes. /// Commands sent to dedicated writer tasks to avoid mutex contention on TCP writes.
pub(crate) enum WriterCommand { pub(crate) enum WriterCommand {
Data(Bytes), Data {
payload: Bytes,
_permit: Option<OwnedSemaphorePermit>,
writer_permit: WriterBytePermit,
},
DataAndFlush(Bytes), DataAndFlush(Bytes),
ProxyReq(ProxyReqCommand), ProxyReq(ProxyReqCommand),
ControlAndFlush([u8; 12]), ControlAndFlush([u8; 12]),
@@ -28,6 +88,8 @@ pub(crate) struct ProxyReqCommand {
pub(crate) proto_flags: u32, pub(crate) proto_flags: u32,
pub(crate) proxy_tag: Option<[u8; 16]>, pub(crate) proxy_tag: Option<[u8; 16]>,
pub(crate) payload: PooledBuffer, pub(crate) payload: PooledBuffer,
pub(crate) _permit: OwnedSemaphorePermit,
pub(crate) writer_permit: WriterBytePermit,
} }
#[derive(Debug, Clone, Copy, PartialEq, Eq)] #[derive(Debug, Clone, Copy, PartialEq, Eq)]
@@ -82,7 +144,7 @@ fn build_rpc_frame_into(
) { ) {
let total_len = 4 + 4 + payload.len() + 4; let total_len = 4 + 4 + payload.len() + 4;
frame.clear(); frame.clear();
frame.reserve(total_len + 15); frame.reserve_exact(total_len + 15);
let total_len = total_len as u32; let total_len = total_len as u32;
frame.extend_from_slice(&total_len.to_le_bytes()); frame.extend_from_slice(&total_len.to_le_bytes());
frame.extend_from_slice(&seq_no.to_le_bytes()); frame.extend_from_slice(&seq_no.to_le_bytes());
@@ -274,7 +336,7 @@ impl RpcWriter {
); );
let total_len = 4 + 4 + payload_len + 4; let total_len = 4 + 4 + payload_len + 4;
self.frame_buf.clear(); self.frame_buf.clear();
self.frame_buf.reserve(total_len + 15); self.frame_buf.reserve_exact(total_len + 15);
self.frame_buf self.frame_buf
.extend_from_slice(&(total_len as u32).to_le_bytes()); .extend_from_slice(&(total_len as u32).to_le_bytes());
self.frame_buf.extend_from_slice(&self.seq_no.to_le_bytes()); self.frame_buf.extend_from_slice(&self.seq_no.to_le_bytes());
@@ -322,3 +384,38 @@ impl RpcWriter {
self.writer.flush().await.map_err(ProxyError::Io) self.writer.flush().await.map_err(ProxyError::Io)
} }
} }
#[cfg(test)]
mod tests {
use super::*;
use tokio::sync::Semaphore;
#[test]
fn writer_byte_permit_tracks_queued_and_inflight_lifecycle() {
let stats = Arc::new(Stats::default());
let semaphore = Arc::new(Semaphore::new(2));
let permit = semaphore
.clone()
.try_acquire_many_owned(2)
.expect("writer byte permits must be available");
let mut writer_permit = WriterBytePermit::new(permit, 32 * 1024, stats.clone());
assert_eq!(
stats.get_me_writer_byte_budget_queued_bytes_gauge(),
32 * 1024
);
assert_eq!(stats.get_me_writer_byte_budget_inflight_bytes_gauge(), 0);
writer_permit.mark_inflight();
assert_eq!(stats.get_me_writer_byte_budget_queued_bytes_gauge(), 0);
assert_eq!(
stats.get_me_writer_byte_budget_inflight_bytes_gauge(),
32 * 1024
);
drop(writer_permit);
assert_eq!(stats.get_me_writer_byte_budget_queued_bytes_gauge(), 0);
assert_eq!(stats.get_me_writer_byte_budget_inflight_bytes_gauge(), 0);
assert_eq!(semaphore.available_permits(), 2);
}
}
+114
View File
@@ -675,8 +675,122 @@ fn hex_dump(data: &[u8]) -> String {
mod tests { mod tests {
use super::*; use super::*;
use std::io::ErrorKind; use std::io::ErrorKind;
use std::net::{Ipv4Addr, Ipv6Addr};
use tokio::net::{TcpListener, TcpStream}; use tokio::net::{TcpListener, TcpStream};
fn upstream_egress(
route_kind: UpstreamRouteKind,
socks_bound_addr: Option<SocketAddr>,
) -> UpstreamEgressInfo {
UpstreamEgressInfo {
upstream_id: 7,
route_kind,
local_addr: None,
direct_bind_ip: None,
socks_bound_addr,
socks_proxy_addr: None,
}
}
#[test]
fn socks_bound_addr_is_used_only_for_public_same_family_tuple() {
let v4_bound = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(93, 184, 216, 34)), 443);
let v6_bound = SocketAddr::new(
IpAddr::V6(
"2606:4700:4700::1111"
.parse::<Ipv6Addr>()
.expect("test IPv6 address must parse"),
),
443,
);
assert_eq!(
MePool::select_socks_bound_addr(
IpFamily::V4,
Some(upstream_egress(UpstreamRouteKind::Socks5, Some(v4_bound)))
),
Some(v4_bound)
);
assert_eq!(
MePool::select_socks_bound_addr(
IpFamily::V6,
Some(upstream_egress(UpstreamRouteKind::Socks5, Some(v6_bound)))
),
Some(v6_bound)
);
}
#[test]
fn socks_bound_addr_rejects_bogon_unspecified_wrong_family_and_non_socks_routes() {
let bogon_bound = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 0, 0, 1)), 443);
let unspecified_bound = SocketAddr::new(IpAddr::V4(Ipv4Addr::UNSPECIFIED), 443);
let public_v4_bound = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(93, 184, 216, 34)), 443);
assert_eq!(
MePool::select_socks_bound_addr(
IpFamily::V4,
Some(upstream_egress(
UpstreamRouteKind::Socks5,
Some(bogon_bound)
))
),
None
);
assert_eq!(
MePool::select_socks_bound_addr(
IpFamily::V4,
Some(upstream_egress(
UpstreamRouteKind::Socks5,
Some(unspecified_bound)
))
),
None
);
assert_eq!(
MePool::select_socks_bound_addr(
IpFamily::V6,
Some(upstream_egress(
UpstreamRouteKind::Socks5,
Some(public_v4_bound)
))
),
None
);
assert_eq!(
MePool::select_socks_bound_addr(
IpFamily::V4,
Some(upstream_egress(
UpstreamRouteKind::Direct,
Some(public_v4_bound)
))
),
None
);
}
#[test]
fn kdf_client_port_source_tracks_only_valid_socks_bound_port() {
let bound = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(93, 184, 216, 34)), 443);
let zero_port_bound = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(93, 184, 216, 34)), 0);
assert_eq!(
KdfClientPortSource::from_socks_bound_port(Some(bound.port())),
KdfClientPortSource::SocksBound
);
assert_eq!(
KdfClientPortSource::from_socks_bound_port(
Some(zero_port_bound)
.filter(|addr| addr.port() != 0)
.map(|addr| addr.port())
),
KdfClientPortSource::LocalSocket
);
assert_eq!(
KdfClientPortSource::from_socks_bound_port(None),
KdfClientPortSource::LocalSocket
);
}
#[tokio::test] #[tokio::test]
async fn test_configure_keepalive_loopback() { async fn test_configure_keepalive_loopback() {
let listener = match TcpListener::bind("127.0.0.1:0").await { let listener = match TcpListener::bind("127.0.0.1:0").await {
+11 -2
View File
@@ -1795,6 +1795,7 @@ mod tests {
general.me_writer_pick_sample_size, general.me_writer_pick_sample_size,
MeSocksKdfPolicy::default(), MeSocksKdfPolicy::default(),
general.me_writer_cmd_channel_capacity, general.me_writer_cmd_channel_capacity,
general.me_writer_byte_budget_bytes,
general.me_route_channel_capacity, general.me_route_channel_capacity,
general.me_route_backpressure_enabled, general.me_route_backpressure_enabled,
general.me_route_fairshare_enabled, general.me_route_fairshare_enabled,
@@ -1821,6 +1822,7 @@ mod tests {
) -> u64 { ) -> u64 {
let (conn_id, _rx) = pool.registry.register().await; let (conn_id, _rx) = pool.registry.register().await;
let (tx, _writer_rx) = mpsc::channel::<WriterCommand>(8); let (tx, _writer_rx) = mpsc::channel::<WriterCommand>(8);
let byte_budget = pool.new_writer_byte_budget();
let writer = MeWriter { let writer = MeWriter {
id: writer_id, id: writer_id,
addr: SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 4000 + writer_id as u16), addr: SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 4000 + writer_id as u16),
@@ -1830,6 +1832,7 @@ mod tests {
contour: Arc::new(AtomicU8::new(WriterContour::Draining.as_u8())), contour: Arc::new(AtomicU8::new(WriterContour::Draining.as_u8())),
created_at: Instant::now() - Duration::from_secs(writer_id), created_at: Instant::now() - Duration::from_secs(writer_id),
tx: tx.clone(), tx: tx.clone(),
byte_budget: byte_budget.clone(),
cancel: CancellationToken::new(), cancel: CancellationToken::new(),
degraded: Arc::new(AtomicBool::new(false)), degraded: Arc::new(AtomicBool::new(false)),
rtt_ema_ms_x10: Arc::new(AtomicU32::new(0)), rtt_ema_ms_x10: Arc::new(AtomicU32::new(0)),
@@ -1839,7 +1842,9 @@ mod tests {
allow_drain_fallback: Arc::new(AtomicBool::new(false)), allow_drain_fallback: Arc::new(AtomicBool::new(false)),
}; };
pool.writers.write().await.push(writer); pool.writers.write().await.push(writer);
pool.registry.register_writer(writer_id, tx).await; pool.registry
.register_writer(writer_id, tx, byte_budget)
.await;
pool.conn_count.fetch_add(1, Ordering::Relaxed); pool.conn_count.fetch_add(1, Ordering::Relaxed);
assert!( assert!(
pool.registry pool.registry
@@ -1860,6 +1865,7 @@ mod tests {
async fn insert_live_writer(pool: &Arc<MePool>, writer_id: u64, writer_dc: i32) { async fn insert_live_writer(pool: &Arc<MePool>, writer_id: u64, writer_dc: i32) {
let (tx, _writer_rx) = mpsc::channel::<WriterCommand>(8); let (tx, _writer_rx) = mpsc::channel::<WriterCommand>(8);
let byte_budget = pool.new_writer_byte_budget();
let writer = MeWriter { let writer = MeWriter {
id: writer_id, id: writer_id,
addr: SocketAddr::new( addr: SocketAddr::new(
@@ -1877,6 +1883,7 @@ mod tests {
contour: Arc::new(AtomicU8::new(WriterContour::Active.as_u8())), contour: Arc::new(AtomicU8::new(WriterContour::Active.as_u8())),
created_at: Instant::now(), created_at: Instant::now(),
tx: tx.clone(), tx: tx.clone(),
byte_budget: byte_budget.clone(),
cancel: CancellationToken::new(), cancel: CancellationToken::new(),
degraded: Arc::new(AtomicBool::new(false)), degraded: Arc::new(AtomicBool::new(false)),
rtt_ema_ms_x10: Arc::new(AtomicU32::new(0)), rtt_ema_ms_x10: Arc::new(AtomicU32::new(0)),
@@ -1886,7 +1893,9 @@ mod tests {
allow_drain_fallback: Arc::new(AtomicBool::new(false)), allow_drain_fallback: Arc::new(AtomicBool::new(false)),
}; };
pool.writers.write().await.push(writer); pool.writers.write().await.push(writer);
pool.registry.register_writer(writer_id, tx).await; pool.registry
.register_writer(writer_id, tx, byte_budget)
.await;
pool.conn_count.fetch_add(1, Ordering::Relaxed); pool.conn_count.fetch_add(1, Ordering::Relaxed);
} }
+1 -21
View File
@@ -66,33 +66,13 @@ fn extract_host_port_path(url: &str) -> Result<(String, u16, String)> {
Ok((host, port, path)) Ok((host, port, path))
} }
async fn resolve_target_addr(host: &str, port: u16) -> Result<std::net::SocketAddr> {
if let Some(addr) = resolve_socket_addr(host, port) {
return Ok(addr);
}
let addrs: Vec<std::net::SocketAddr> = tokio::net::lookup_host((host, port))
.await
.map_err(|e| ProxyError::Proxy(format!("DNS resolve failed for {host}:{port}: {e}")))?
.collect();
if let Some(addr) = addrs.iter().copied().find(|addr| addr.is_ipv4()) {
return Ok(addr);
}
addrs
.first()
.copied()
.ok_or_else(|| ProxyError::Proxy(format!("DNS returned no addresses for {host}:{port}")))
}
async fn connect_https_transport( async fn connect_https_transport(
host: &str, host: &str,
port: u16, port: u16,
upstream: Option<Arc<UpstreamManager>>, upstream: Option<Arc<UpstreamManager>>,
) -> Result<UpstreamStream> { ) -> Result<UpstreamStream> {
if let Some(manager) = upstream { if let Some(manager) = upstream {
let target = resolve_target_addr(host, port).await?; let target = manager.resolve_hostname(host, port).await?;
return timeout(HTTP_CONNECT_TIMEOUT, manager.connect(target, None, None)) return timeout(HTTP_CONNECT_TIMEOUT, manager.connect(target, None, None))
.await .await
.map_err(|_| ProxyError::Proxy(format!("upstream connect timeout for {host}:{port}")))? .map_err(|_| ProxyError::Proxy(format!("upstream connect timeout for {host}:{port}")))?
+16 -1
View File
@@ -10,7 +10,7 @@ use std::sync::atomic::{
use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH}; use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH};
use arc_swap::ArcSwap; use arc_swap::ArcSwap;
use tokio::sync::{Mutex, RwLock, mpsc, watch}; use tokio::sync::{Mutex, RwLock, Semaphore, mpsc, watch};
use tokio_util::sync::CancellationToken; use tokio_util::sync::CancellationToken;
use crate::config::{ use crate::config::{
@@ -48,6 +48,8 @@ pub struct MeWriter {
pub contour: Arc<AtomicU8>, pub contour: Arc<AtomicU8>,
pub created_at: Instant, pub created_at: Instant,
pub tx: mpsc::Sender<WriterCommand>, pub tx: mpsc::Sender<WriterCommand>,
/// Aggregate resident-memory budget shared by all data commands for this writer.
pub byte_budget: Arc<Semaphore>,
pub cancel: CancellationToken, pub cancel: CancellationToken,
pub degraded: Arc<AtomicBool>, pub degraded: Arc<AtomicBool>,
pub rtt_ema_ms_x10: Arc<AtomicU32>, pub rtt_ema_ms_x10: Arc<AtomicU32>,
@@ -277,6 +279,7 @@ pub(super) struct WriterLifecycleCore {
pub(super) me_keepalive_payload_random: bool, pub(super) me_keepalive_payload_random: bool,
pub(super) rpc_proxy_req_every_secs: AtomicU64, pub(super) rpc_proxy_req_every_secs: AtomicU64,
pub(super) writer_cmd_channel_capacity: usize, pub(super) writer_cmd_channel_capacity: usize,
pub(super) writer_byte_budget_permits: usize,
} }
pub(super) struct RouteRuntimeCore { pub(super) struct RouteRuntimeCore {
@@ -553,6 +556,7 @@ impl MePool {
me_writer_pick_sample_size: u8, me_writer_pick_sample_size: u8,
me_socks_kdf_policy: MeSocksKdfPolicy, me_socks_kdf_policy: MeSocksKdfPolicy,
me_writer_cmd_channel_capacity: usize, me_writer_cmd_channel_capacity: usize,
me_writer_byte_budget_bytes: usize,
me_route_channel_capacity: usize, me_route_channel_capacity: usize,
me_route_backpressure_enabled: bool, me_route_backpressure_enabled: bool,
me_route_fairshare_enabled: bool, me_route_fairshare_enabled: bool,
@@ -583,6 +587,7 @@ impl MePool {
); );
let (writer_epoch, _) = watch::channel(0u64); let (writer_epoch, _) = watch::channel(0u64);
let now_epoch_secs = Self::now_epoch_secs(); let now_epoch_secs = Self::now_epoch_secs();
stats.set_me_writer_byte_budget_limit_bytes(me_writer_byte_budget_bytes);
Arc::new(Self { Arc::new(Self {
routing: Arc::new(RoutingCore { routing: Arc::new(RoutingCore {
registry, registry,
@@ -615,6 +620,9 @@ impl MePool {
me_keepalive_payload_random, me_keepalive_payload_random,
rpc_proxy_req_every_secs: AtomicU64::new(rpc_proxy_req_every_secs), rpc_proxy_req_every_secs: AtomicU64::new(rpc_proxy_req_every_secs),
writer_cmd_channel_capacity: me_writer_cmd_channel_capacity.max(1), writer_cmd_channel_capacity: me_writer_cmd_channel_capacity.max(1),
writer_byte_budget_permits: me_writer_byte_budget_bytes
.div_ceil(crate::config::defaults::ME_WRITER_BYTE_PERMIT_UNIT_BYTES)
.max(1),
}), }),
route_runtime: Arc::new(RouteRuntimeCore { route_runtime: Arc::new(RouteRuntimeCore {
me_route_no_writer_mode: AtomicU8::new(me_route_no_writer_mode.as_u8()), me_route_no_writer_mode: AtomicU8::new(me_route_no_writer_mode.as_u8()),
@@ -833,6 +841,13 @@ impl MePool {
}) })
} }
/// Creates the immutable byte semaphore assigned to one ME writer generation.
pub(crate) fn new_writer_byte_budget(&self) -> Arc<Semaphore> {
Arc::new(Semaphore::new(
self.writer_lifecycle.writer_byte_budget_permits,
))
}
pub fn current_generation(&self) -> u64 { pub fn current_generation(&self) -> u64 {
self.reinit.active_generation.load(Ordering::Relaxed) self.reinit.active_generation.load(Ordering::Relaxed)
} }
+13 -3
View File
@@ -63,13 +63,19 @@ async fn writer_command_loop(
_ = cancel.cancelled() => return Ok(()), _ = cancel.cancelled() => return Ok(()),
cmd = rx.recv() => { cmd = rx.recv() => {
match cmd { match cmd {
Some(WriterCommand::Data(payload)) => { Some(WriterCommand::Data {
payload,
_permit,
mut writer_permit,
}) => {
writer_permit.mark_inflight();
rpc_writer.send(&payload).await?; rpc_writer.send(&payload).await?;
} }
Some(WriterCommand::DataAndFlush(payload)) => { Some(WriterCommand::DataAndFlush(payload)) => {
rpc_writer.send_and_flush(&payload).await?; rpc_writer.send_and_flush(&payload).await?;
} }
Some(WriterCommand::ProxyReq(command)) => { Some(WriterCommand::ProxyReq(mut command)) => {
command.writer_permit.mark_inflight();
rpc_writer.send_proxy_req(&command).await?; rpc_writer.send_proxy_req(&command).await?;
} }
Some(WriterCommand::ControlAndFlush(payload)) => { Some(WriterCommand::ControlAndFlush(payload)) => {
@@ -419,6 +425,7 @@ impl MePool {
let draining_started_at_epoch_secs = Arc::new(AtomicU64::new(0)); let draining_started_at_epoch_secs = Arc::new(AtomicU64::new(0));
let drain_deadline_epoch_secs = Arc::new(AtomicU64::new(0)); let drain_deadline_epoch_secs = Arc::new(AtomicU64::new(0));
let allow_drain_fallback = Arc::new(AtomicBool::new(false)); let allow_drain_fallback = Arc::new(AtomicBool::new(false));
let byte_budget = self.new_writer_byte_budget();
let (tx, rx) = let (tx, rx) =
mpsc::channel::<WriterCommand>(self.writer_lifecycle.writer_cmd_channel_capacity); mpsc::channel::<WriterCommand>(self.writer_lifecycle.writer_cmd_channel_capacity);
let rpc_writer = RpcWriter { let rpc_writer = RpcWriter {
@@ -438,6 +445,7 @@ impl MePool {
contour: contour.clone(), contour: contour.clone(),
created_at: Instant::now(), created_at: Instant::now(),
tx: tx.clone(), tx: tx.clone(),
byte_budget: byte_budget.clone(),
cancel: cancel.clone(), cancel: cancel.clone(),
degraded: degraded.clone(), degraded: degraded.clone(),
rtt_ema_ms_x10: rtt_ema_ms_x10.clone(), rtt_ema_ms_x10: rtt_ema_ms_x10.clone(),
@@ -449,7 +457,9 @@ impl MePool {
self.writers self.writers
.update(|writers| writers.push(writer.clone())) .update(|writers| writers.push(writer.clone()))
.await; .await;
self.registry.register_writer(writer_id, tx.clone()).await; self.registry
.register_writer(writer_id, tx.clone(), byte_budget)
.await;
self.registry.mark_writer_idle(writer_id).await; self.registry.mark_writer_idle(writer_id).await;
self.conn_count.fetch_add(1, Ordering::Relaxed); self.conn_count.fetch_add(1, Ordering::Relaxed);
self.notify_writer_epoch(); self.notify_writer_epoch();
+9 -1
View File
@@ -48,6 +48,8 @@ pub struct BoundConn {
pub struct ConnWriter { pub struct ConnWriter {
pub writer_id: u64, pub writer_id: u64,
pub tx: mpsc::Sender<WriterCommand>, pub tx: mpsc::Sender<WriterCommand>,
/// Writer-local memory budget used by the hot bound-client route.
pub byte_budget: Arc<Semaphore>,
} }
#[derive(Clone, Debug, Default)] #[derive(Clone, Debug, Default)]
@@ -62,7 +64,13 @@ struct RoutingTable {
} }
struct WriterTable { struct WriterTable {
map: DashMap<u64, mpsc::Sender<WriterCommand>>, map: DashMap<u64, WriterRoute>,
}
#[derive(Clone)]
struct WriterRoute {
tx: mpsc::Sender<WriterCommand>,
byte_budget: Arc<Semaphore>,
} }
#[derive(Clone)] #[derive(Clone)]
+30 -8
View File
@@ -1,10 +1,16 @@
use std::net::{IpAddr, Ipv4Addr, SocketAddr}; use std::net::{IpAddr, Ipv4Addr, SocketAddr};
use std::sync::Arc;
use bytes::Bytes; use bytes::Bytes;
use tokio::sync::Semaphore;
use super::{ConnMeta, ConnRegistry, RouteResult}; use super::{ConnMeta, ConnRegistry, RouteResult};
use crate::transport::middle_proxy::MeResponse; use crate::transport::middle_proxy::MeResponse;
fn writer_byte_budget() -> Arc<Semaphore> {
Arc::new(Semaphore::new(2049))
}
#[tokio::test] #[tokio::test]
async fn writer_activity_snapshot_tracks_writer_and_dc_load() { async fn writer_activity_snapshot_tracks_writer_and_dc_load() {
let registry = ConnRegistry::new(); let registry = ConnRegistry::new();
@@ -14,8 +20,12 @@ async fn writer_activity_snapshot_tracks_writer_and_dc_load() {
let (conn_c, _rx_c) = registry.register().await; let (conn_c, _rx_c) = registry.register().await;
let (writer_tx_a, _writer_rx_a) = tokio::sync::mpsc::channel(8); let (writer_tx_a, _writer_rx_a) = tokio::sync::mpsc::channel(8);
let (writer_tx_b, _writer_rx_b) = tokio::sync::mpsc::channel(8); let (writer_tx_b, _writer_rx_b) = tokio::sync::mpsc::channel(8);
registry.register_writer(10, writer_tx_a.clone()).await; registry
registry.register_writer(20, writer_tx_b.clone()).await; .register_writer(10, writer_tx_a.clone(), writer_byte_budget())
.await;
registry
.register_writer(20, writer_tx_b.clone(), writer_byte_budget())
.await;
let addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 443); let addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 443);
assert!( assert!(
@@ -124,8 +134,12 @@ async fn bind_writer_rebinds_conn_atomically() {
let (conn_id, _rx) = registry.register().await; let (conn_id, _rx) = registry.register().await;
let (writer_tx_a, _writer_rx_a) = tokio::sync::mpsc::channel(8); let (writer_tx_a, _writer_rx_a) = tokio::sync::mpsc::channel(8);
let (writer_tx_b, _writer_rx_b) = tokio::sync::mpsc::channel(8); let (writer_tx_b, _writer_rx_b) = tokio::sync::mpsc::channel(8);
registry.register_writer(10, writer_tx_a).await; registry
registry.register_writer(20, writer_tx_b).await; .register_writer(10, writer_tx_a, writer_byte_budget())
.await;
registry
.register_writer(20, writer_tx_b, writer_byte_budget())
.await;
let client_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 443); let client_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 443);
let first_our_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(1, 1, 1, 1)), 443); let first_our_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(1, 1, 1, 1)), 443);
@@ -184,8 +198,12 @@ async fn writer_lost_does_not_drop_rebound_conn() {
let (conn_id, _rx) = registry.register().await; let (conn_id, _rx) = registry.register().await;
let (writer_tx_a, _writer_rx_a) = tokio::sync::mpsc::channel(8); let (writer_tx_a, _writer_rx_a) = tokio::sync::mpsc::channel(8);
let (writer_tx_b, _writer_rx_b) = tokio::sync::mpsc::channel(8); let (writer_tx_b, _writer_rx_b) = tokio::sync::mpsc::channel(8);
registry.register_writer(10, writer_tx_a).await; registry
registry.register_writer(20, writer_tx_b).await; .register_writer(10, writer_tx_a, writer_byte_budget())
.await;
registry
.register_writer(20, writer_tx_b, writer_byte_budget())
.await;
let addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 443); let addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 443);
assert!( assert!(
@@ -262,8 +280,12 @@ async fn non_empty_writer_ids_returns_only_writers_with_bound_clients() {
let (conn_id, _rx) = registry.register().await; let (conn_id, _rx) = registry.register().await;
let (writer_tx_a, _writer_rx_a) = tokio::sync::mpsc::channel(8); let (writer_tx_a, _writer_rx_a) = tokio::sync::mpsc::channel(8);
let (writer_tx_b, _writer_rx_b) = tokio::sync::mpsc::channel(8); let (writer_tx_b, _writer_rx_b) = tokio::sync::mpsc::channel(8);
registry.register_writer(10, writer_tx_a).await; registry
registry.register_writer(20, writer_tx_b).await; .register_writer(10, writer_tx_a, writer_byte_budget())
.await;
registry
.register_writer(20, writer_tx_b, writer_byte_budget())
.await;
let addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 443); let addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 443);
assert!( assert!(
+15 -4
View File
@@ -1,4 +1,5 @@
use std::collections::{HashMap, HashSet}; use std::collections::{HashMap, HashSet};
use std::sync::Arc;
use std::sync::atomic::Ordering; use std::sync::atomic::Ordering;
use std::time::Duration; use std::time::Duration;
@@ -56,7 +57,13 @@ impl ConnRegistry {
} }
} }
pub async fn register_writer(&self, writer_id: u64, tx: mpsc::Sender<WriterCommand>) { /// Registers one writer command route and its matching memory budget atomically.
pub async fn register_writer(
&self,
writer_id: u64,
tx: mpsc::Sender<WriterCommand>,
byte_budget: Arc<tokio::sync::Semaphore>,
) {
let mut binding = self.binding.inner.lock().await; let mut binding = self.binding.inner.lock().await;
binding binding
.conns_for_writer .conns_for_writer
@@ -70,7 +77,9 @@ impl ConnRegistry {
.writer_idle_since_epoch_secs .writer_idle_since_epoch_secs
.entry(writer_id) .entry(writer_id)
.or_insert_with(Self::now_epoch_secs); .or_insert_with(Self::now_epoch_secs);
self.writers.map.insert(writer_id, tx); self.writers
.map
.insert(writer_id, super::WriterRoute { tx, byte_budget });
} }
/// Unregister connection, returning associated writer_id if any. /// Unregister connection, returning associated writer_id if any.
@@ -417,7 +426,8 @@ impl ConnRegistry {
.map(|entry| entry.value().clone())?; .map(|entry| entry.value().clone())?;
Some(ConnWriter { Some(ConnWriter {
writer_id, writer_id,
tx: writer, tx: writer.tx,
byte_budget: writer.byte_budget,
}) })
} }
@@ -438,7 +448,8 @@ impl ConnRegistry {
Some(( Some((
ConnWriter { ConnWriter {
writer_id, writer_id,
tx: writer, tx: writer.tx,
byte_budget: writer.byte_budget,
}, },
meta, meta,
)) ))
+310 -46
View File
@@ -6,16 +6,18 @@ use std::sync::Arc;
use std::sync::atomic::Ordering; use std::sync::atomic::Ordering;
use std::time::{Duration, Instant}; use std::time::{Duration, Instant};
use tokio::sync::mpsc;
use tokio::sync::mpsc::error::TrySendError; use tokio::sync::mpsc::error::TrySendError;
use tokio::sync::{OwnedSemaphorePermit, Semaphore, TryAcquireError, mpsc};
use tracing::{debug, warn}; use tracing::{debug, warn};
use super::MePool; use super::MePool;
use super::codec::{ProxyReqCommand, WriterCommand}; use super::codec::{ProxyReqCommand, WriterBytePermit, WriterCommand};
use super::registry::ConnMeta; use super::registry::ConnMeta;
use super::wire::build_proxy_req_payload; use super::wire::{build_proxy_req_payload, proxy_req_payload_len};
use crate::config::defaults::ME_WRITER_BYTE_PERMIT_UNIT_BYTES;
use crate::config::{MeRouteNoWriterMode, MeWriterPickMode}; use crate::config::{MeRouteNoWriterMode, MeWriterPickMode};
use crate::error::{ProxyError, Result}; use crate::error::{ProxyError, Result};
use crate::stats::Stats;
use crate::stream::PooledBuffer; use crate::stream::PooledBuffer;
use rand::seq::SliceRandom; use rand::seq::SliceRandom;
@@ -29,6 +31,8 @@ const PICK_PENALTY_WARM: u64 = 200;
const PICK_PENALTY_DRAINING: u64 = 600; const PICK_PENALTY_DRAINING: u64 = 600;
const PICK_PENALTY_STALE: u64 = 300; const PICK_PENALTY_STALE: u64 = 300;
const PICK_PENALTY_DEGRADED: u64 = 250; const PICK_PENALTY_DEGRADED: u64 = 250;
const RPC_WRITER_FRAME_CAPACITY_OVERHEAD_BYTES: usize = 27;
const LEGACY_PROXY_REQ_SOURCE_CAPACITY_OVERHEAD_BYTES: usize = 128;
mod close; mod close;
mod recovery; mod recovery;
@@ -39,34 +43,133 @@ enum WriterCommandReserveError {
TimedOut, TimedOut,
} }
enum WriterByteReserveError {
Closed,
TimedOut,
}
fn proxy_tag_array(tag: Option<&[u8]>) -> Option<[u8; 16]> { fn proxy_tag_array(tag: Option<&[u8]>) -> Option<[u8; 16]> {
tag.and_then(|tag| <[u8; 16]>::try_from(tag).ok()) tag.and_then(|tag| <[u8; 16]>::try_from(tag).ok())
} }
fn proxy_req_payload_from_command(cmd: WriterCommand) -> Option<PooledBuffer> { fn proxy_req_payload_from_command(
cmd: WriterCommand,
) -> Option<(PooledBuffer, OwnedSemaphorePermit)> {
match cmd { match cmd {
WriterCommand::ProxyReq(command) => Some(command.payload), WriterCommand::ProxyReq(command) => Some((command.payload, command._permit)),
_ => None,
}
}
fn payload_permit_from_data_command(cmd: WriterCommand) -> Option<OwnedSemaphorePermit> {
match cmd {
WriterCommand::Data { _permit, .. } => _permit,
_ => None, _ => None,
} }
} }
async fn reserve_writer_command_slot( async fn reserve_writer_command_slot(
tx: &mpsc::Sender<WriterCommand>, tx: &mpsc::Sender<WriterCommand>,
wait: Option<Duration>, deadline: Option<Instant>,
) -> std::result::Result<mpsc::OwnedPermit<WriterCommand>, WriterCommandReserveError> { ) -> std::result::Result<mpsc::OwnedPermit<WriterCommand>, WriterCommandReserveError> {
let reserve = tx.clone().reserve_owned(); let reserve = tx.clone().reserve_owned();
match wait { match deadline {
Some(wait) => match tokio::time::timeout(wait, reserve).await { Some(deadline) => {
match tokio::time::timeout(deadline.saturating_duration_since(Instant::now()), reserve)
.await
{
Ok(Ok(permit)) => Ok(permit), Ok(Ok(permit)) => Ok(permit),
Ok(Err(_)) => Err(WriterCommandReserveError::Closed), Ok(Err(_)) => Err(WriterCommandReserveError::Closed),
Err(_) => Err(WriterCommandReserveError::TimedOut), Err(_) => Err(WriterCommandReserveError::TimedOut),
}, }
}
None => reserve.await.map_err(|_| WriterCommandReserveError::Closed), None => reserve.await.map_err(|_| WriterCommandReserveError::Closed),
} }
} }
fn writer_send_deadline(wait: Option<Duration>) -> Option<Instant> {
wait.map(|wait| Instant::now() + wait)
}
fn writer_resident_permits(
source_capacity: usize,
encoded_payload_len: usize,
) -> Option<(u32, usize)> {
let resident_bytes = source_capacity
.checked_add(encoded_payload_len)?
.checked_add(RPC_WRITER_FRAME_CAPACITY_OVERHEAD_BYTES)?;
let permits = resident_bytes.div_ceil(ME_WRITER_BYTE_PERMIT_UNIT_BYTES);
let permits = u32::try_from(permits).ok()?;
let reserved_bytes = (permits as usize).checked_mul(ME_WRITER_BYTE_PERMIT_UNIT_BYTES)?;
Some((
permits.max(1),
reserved_bytes.max(ME_WRITER_BYTE_PERMIT_UNIT_BYTES),
))
}
fn proxy_req_resident_permits(
source_capacity: usize,
data_len: usize,
proxy_tag: Option<&[u8]>,
proto_flags: u32,
) -> Option<(u32, usize)> {
writer_resident_permits(
source_capacity,
proxy_req_payload_len(data_len, proxy_tag, proto_flags),
)
}
fn try_reserve_writer_bytes(
byte_budget: &Arc<Semaphore>,
permits: u32,
reserved_bytes: usize,
stats: &Arc<Stats>,
) -> std::result::Result<WriterBytePermit, TryAcquireError> {
byte_budget
.clone()
.try_acquire_many_owned(permits)
.map(|permit| WriterBytePermit::new(permit, reserved_bytes, stats.clone()))
}
async fn reserve_writer_bytes(
byte_budget: &Arc<Semaphore>,
permits: u32,
reserved_bytes: usize,
deadline: Option<Instant>,
stats: &Arc<Stats>,
) -> std::result::Result<WriterBytePermit, WriterByteReserveError> {
match try_reserve_writer_bytes(byte_budget, permits, reserved_bytes, stats) {
Ok(permit) => return Ok(permit),
Err(TryAcquireError::Closed) => return Err(WriterByteReserveError::Closed),
Err(TryAcquireError::NoPermits) => {
stats.increment_me_writer_byte_budget_wait_total();
}
}
let acquire = byte_budget.clone().acquire_many_owned(permits);
match deadline {
Some(deadline) => {
match tokio::time::timeout(deadline.saturating_duration_since(Instant::now()), acquire)
.await
{
Ok(Ok(permit)) => Ok(WriterBytePermit::new(permit, reserved_bytes, stats.clone())),
Ok(Err(_)) => Err(WriterByteReserveError::Closed),
Err(_) => {
stats.increment_me_writer_byte_budget_timeout_total();
Err(WriterByteReserveError::TimedOut)
}
}
}
None => acquire
.await
.map(|permit| WriterBytePermit::new(permit, reserved_bytes, stats.clone()))
.map_err(|_| WriterByteReserveError::Closed),
}
}
impl MePool { impl MePool {
/// Send RPC_PROXY_REQ. `tag_override`: per-user ad_tag (from access.user_ad_tags); if None, uses pool default. /// Send RPC_PROXY_REQ. `tag_override`: per-user ad_tag (from access.user_ad_tags); if None, uses pool default.
/// `payload_permit` keeps optional client byte accounting alive until the writer consumes the command.
pub async fn send_proxy_req( pub async fn send_proxy_req(
self: &Arc<Self>, self: &Arc<Self>,
conn_id: u64, conn_id: u64,
@@ -76,8 +179,32 @@ impl MePool {
data: &[u8], data: &[u8],
proto_flags: u32, proto_flags: u32,
tag_override: Option<&[u8]>, tag_override: Option<&[u8]>,
mut payload_permit: Option<OwnedSemaphorePermit>,
) -> Result<()> { ) -> Result<()> {
let tag = tag_override.or(self.proxy_tag.as_deref()); let tag = tag_override.or(self.proxy_tag.as_deref());
let Some(source_capacity) = data
.len()
.checked_add(LEGACY_PROXY_REQ_SOURCE_CAPACITY_OVERHEAD_BYTES)
else {
self.stats.increment_me_writer_byte_budget_oversize_total();
return Err(ProxyError::Proxy(
"ME writer payload residency calculation overflow".into(),
));
};
let Some((writer_byte_permits, writer_reserved_bytes)) =
proxy_req_resident_permits(source_capacity, data.len(), tag, proto_flags)
else {
self.stats.increment_me_writer_byte_budget_oversize_total();
return Err(ProxyError::Proxy(
"ME writer payload residency calculation overflow".into(),
));
};
if writer_byte_permits as usize > self.writer_lifecycle.writer_byte_budget_permits {
self.stats.increment_me_writer_byte_budget_oversize_total();
return Err(ProxyError::Proxy(
"ME writer payload exceeds configured byte budget".into(),
));
}
let build_routed_payload = |effective_our_addr: SocketAddr| { let build_routed_payload = |effective_our_addr: SocketAddr| {
( (
build_proxy_req_payload( build_proxy_req_payload(
@@ -118,19 +245,48 @@ impl MePool {
loop { loop {
if let Some((current, current_meta)) = self.registry.get_writer_with_meta(conn_id).await if let Some((current, current_meta)) = self.registry.get_writer_with_meta(conn_id).await
{ {
let deadline =
writer_send_deadline(self.route_runtime.me_route_blocking_send_timeout);
let writer_permit = match reserve_writer_bytes(
&current.byte_budget,
writer_byte_permits,
writer_reserved_bytes,
deadline,
&self.stats,
)
.await
{
Ok(permit) => permit,
Err(WriterByteReserveError::TimedOut) => {
self.stats
.increment_me_writer_pick_full_total(self.writer_pick_mode());
return Err(ProxyError::Proxy(
"ME writer byte budget full within blocking send timeout".into(),
));
}
Err(WriterByteReserveError::Closed) => {
warn!(
writer_id = current.writer_id,
"ME writer byte budget closed"
);
self.remove_writer_and_close_clients(current.writer_id)
.await;
continue;
}
};
let (current_payload, _) = build_routed_payload(current_meta.our_addr); let (current_payload, _) = build_routed_payload(current_meta.our_addr);
match current.tx.try_send(WriterCommand::Data(current_payload)) { let command = WriterCommand::Data {
payload: current_payload,
_permit: payload_permit.take(),
writer_permit,
};
match current.tx.try_send(command) {
Ok(()) => { Ok(()) => {
self.note_hybrid_route_success(); self.note_hybrid_route_success();
return Ok(()); return Ok(());
} }
Err(TrySendError::Full(cmd)) => { Err(TrySendError::Full(cmd)) => {
match reserve_writer_command_slot( match reserve_writer_command_slot(&current.tx, deadline).await {
&current.tx,
self.route_runtime.me_route_blocking_send_timeout,
)
.await
{
Ok(permit) => { Ok(permit) => {
permit.send(cmd); permit.send(cmd);
self.note_hybrid_route_success(); self.note_hybrid_route_success();
@@ -143,14 +299,17 @@ impl MePool {
"ME writer channel full within blocking send timeout".into(), "ME writer channel full within blocking send timeout".into(),
)); ));
} }
Err(WriterCommandReserveError::Closed) => {} Err(WriterCommandReserveError::Closed) => {
payload_permit = payload_permit_from_data_command(cmd);
}
} }
warn!(writer_id = current.writer_id, "ME writer channel closed"); warn!(writer_id = current.writer_id, "ME writer channel closed");
self.remove_writer_and_close_clients(current.writer_id) self.remove_writer_and_close_clients(current.writer_id)
.await; .await;
continue; continue;
} }
Err(TrySendError::Closed(_)) => { Err(TrySendError::Closed(cmd)) => {
payload_permit = payload_permit_from_data_command(cmd);
warn!(writer_id = current.writer_id, "ME writer channel closed"); warn!(writer_id = current.writer_id, "ME writer channel closed");
self.remove_writer_and_close_clients(current.writer_id) self.remove_writer_and_close_clients(current.writer_id)
.await; .await;
@@ -464,9 +623,31 @@ impl MePool {
if !self.writer_accepts_new_binding(w) { if !self.writer_accepts_new_binding(w) {
continue; continue;
} }
let (payload, meta) = build_routed_payload(our_addr); let writer_permit = match try_reserve_writer_bytes(
&w.byte_budget,
writer_byte_permits,
writer_reserved_bytes,
&self.stats,
) {
Ok(permit) => permit,
Err(TryAcquireError::NoPermits) => {
if fallback_blocking_idx.is_none() {
fallback_blocking_idx = Some(idx);
}
continue;
}
Err(TryAcquireError::Closed) => {
self.stats.increment_me_writer_pick_closed_total(pick_mode);
warn!(writer_id = w.id, "ME writer byte budget closed");
self.remove_writer_and_close_clients(w.id).await;
continue;
}
};
match w.tx.clone().try_reserve_owned() { match w.tx.clone().try_reserve_owned() {
Ok(permit) => { Ok(permit) => {
// Keep the advertised proxy IP aligned with the selected ME writer source.
let effective_our_addr = SocketAddr::new(w.source_ip, our_addr.port());
let (payload, meta) = build_routed_payload(effective_our_addr);
if !self.registry.bind_writer(conn_id, w.id, meta).await { if !self.registry.bind_writer(conn_id, w.id, meta).await {
debug!( debug!(
conn_id, conn_id,
@@ -477,7 +658,11 @@ impl MePool {
self.remove_writer_and_close_clients(w.id).await; self.remove_writer_and_close_clients(w.id).await;
continue; continue;
} }
permit.send(WriterCommand::Data(payload)); permit.send(WriterCommand::Data {
payload,
_permit: payload_permit.take(),
writer_permit,
});
self.stats self.stats
.increment_me_writer_pick_success_try_total(pick_mode); .increment_me_writer_pick_success_try_total(pick_mode);
if w.generation < self.current_generation() { if w.generation < self.current_generation() {
@@ -519,21 +704,44 @@ impl MePool {
} }
self.stats self.stats
.increment_me_writer_pick_blocking_fallback_total(); .increment_me_writer_pick_blocking_fallback_total();
let (payload, meta) = build_routed_payload(our_addr); let deadline = writer_send_deadline(self.route_runtime.me_route_blocking_send_timeout);
let reserve_result = let writer_permit = match reserve_writer_bytes(
if let Some(timeout) = self.route_runtime.me_route_blocking_send_timeout { &w.byte_budget,
match tokio::time::timeout(timeout, w.tx.clone().reserve_owned()).await { writer_byte_permits,
Ok(result) => result, writer_reserved_bytes,
Err(_) => { deadline,
&self.stats,
)
.await
{
Ok(permit) => permit,
Err(WriterByteReserveError::TimedOut) => {
self.stats.increment_me_writer_pick_full_total(pick_mode); self.stats.increment_me_writer_pick_full_total(pick_mode);
continue; continue;
} }
Err(WriterByteReserveError::Closed) => {
self.stats.increment_me_writer_pick_closed_total(pick_mode);
warn!(writer_id = w.id, "ME writer byte budget closed (blocking)");
self.remove_writer_and_close_clients(w.id).await;
continue;
} }
} else {
w.tx.clone().reserve_owned().await
}; };
match reserve_result { let permit = match reserve_writer_command_slot(&w.tx, deadline).await {
Ok(permit) => { Ok(permit) => permit,
Err(WriterCommandReserveError::TimedOut) => {
self.stats.increment_me_writer_pick_full_total(pick_mode);
continue;
}
Err(WriterCommandReserveError::Closed) => {
self.stats.increment_me_writer_pick_closed_total(pick_mode);
warn!(writer_id = w.id, "ME writer channel closed (blocking)");
self.remove_writer_and_close_clients(w.id).await;
continue;
}
};
// Keep the advertised proxy IP aligned with the selected ME writer source.
let effective_our_addr = SocketAddr::new(w.source_ip, our_addr.port());
let (payload, meta) = build_routed_payload(effective_our_addr);
if !self.registry.bind_writer(conn_id, w.id, meta).await { if !self.registry.bind_writer(conn_id, w.id, meta).await {
debug!( debug!(
conn_id, conn_id,
@@ -544,7 +752,11 @@ impl MePool {
self.remove_writer_and_close_clients(w.id).await; self.remove_writer_and_close_clients(w.id).await;
continue; continue;
} }
permit.send(WriterCommand::Data(payload)); permit.send(WriterCommand::Data {
payload,
_permit: payload_permit.take(),
writer_permit,
});
self.stats self.stats
.increment_me_writer_pick_success_fallback_total(pick_mode); .increment_me_writer_pick_success_fallback_total(pick_mode);
if w.generation < self.current_generation() { if w.generation < self.current_generation() {
@@ -553,16 +765,10 @@ impl MePool {
self.note_hybrid_route_success(); self.note_hybrid_route_success();
return Ok(()); return Ok(());
} }
Err(_) => {
self.stats.increment_me_writer_pick_closed_total(pick_mode);
warn!(writer_id = w.id, "ME writer channel closed (blocking)");
self.remove_writer_and_close_clients(w.id).await;
}
}
}
} }
/// Send RPC_PROXY_REQ while keeping the first bound-writer path allocation-light. /// Send RPC_PROXY_REQ while keeping the first bound-writer path allocation-light.
/// The client byte permit follows the payload until writer completion or command drop.
pub async fn send_proxy_req_pooled( pub async fn send_proxy_req_pooled(
self: &Arc<Self>, self: &Arc<Self>,
conn_id: u64, conn_id: u64,
@@ -570,12 +776,69 @@ impl MePool {
client_addr: SocketAddr, client_addr: SocketAddr,
our_addr: SocketAddr, our_addr: SocketAddr,
payload: PooledBuffer, payload: PooledBuffer,
_permit: OwnedSemaphorePermit,
proto_flags: u32, proto_flags: u32,
tag_override: Option<[u8; 16]>, tag_override: Option<[u8; 16]>,
) -> Result<()> { ) -> Result<()> {
let tag = tag_override.or_else(|| proxy_tag_array(self.proxy_tag.as_deref())); let tag = tag_override.or_else(|| proxy_tag_array(self.proxy_tag.as_deref()));
let Some((writer_byte_permits, writer_reserved_bytes)) = proxy_req_resident_permits(
payload.capacity(),
payload.len(),
tag.as_ref().map(|tag| tag.as_slice()),
proto_flags,
) else {
self.stats.increment_me_writer_byte_budget_oversize_total();
return Err(ProxyError::Proxy(
"ME writer payload residency calculation overflow".into(),
));
};
if writer_byte_permits as usize > self.writer_lifecycle.writer_byte_budget_permits {
self.stats.increment_me_writer_byte_budget_oversize_total();
return Err(ProxyError::Proxy(
"ME writer payload exceeds configured byte budget".into(),
));
}
if let Some((current, current_meta)) = self.registry.get_writer_with_meta(conn_id).await { if let Some((current, current_meta)) = self.registry.get_writer_with_meta(conn_id).await {
let deadline = writer_send_deadline(self.route_runtime.me_route_blocking_send_timeout);
let writer_permit = match reserve_writer_bytes(
&current.byte_budget,
writer_byte_permits,
writer_reserved_bytes,
deadline,
&self.stats,
)
.await
{
Ok(permit) => permit,
Err(WriterByteReserveError::TimedOut) => {
self.stats
.increment_me_writer_pick_full_total(self.writer_pick_mode());
return Err(ProxyError::Proxy(
"ME writer byte budget full within blocking send timeout".into(),
));
}
Err(WriterByteReserveError::Closed) => {
warn!(
writer_id = current.writer_id,
"ME writer byte budget closed"
);
self.remove_writer_and_close_clients(current.writer_id)
.await;
return self
.send_proxy_req(
conn_id,
target_dc,
client_addr,
our_addr,
payload.as_ref(),
proto_flags,
tag.as_ref().map(|tag| tag.as_slice()),
Some(_permit),
)
.await;
}
};
let command = WriterCommand::ProxyReq(ProxyReqCommand { let command = WriterCommand::ProxyReq(ProxyReqCommand {
conn_id, conn_id,
client_addr, client_addr,
@@ -583,6 +846,8 @@ impl MePool {
proto_flags, proto_flags,
proxy_tag: tag, proxy_tag: tag,
payload, payload,
_permit,
writer_permit,
}); });
match current.tx.try_send(command) { match current.tx.try_send(command) {
Ok(()) => { Ok(()) => {
@@ -590,12 +855,7 @@ impl MePool {
return Ok(()); return Ok(());
} }
Err(TrySendError::Full(cmd)) => { Err(TrySendError::Full(cmd)) => {
match reserve_writer_command_slot( match reserve_writer_command_slot(&current.tx, deadline).await {
&current.tx,
self.route_runtime.me_route_blocking_send_timeout,
)
.await
{
Ok(permit) => { Ok(permit) => {
permit.send(cmd); permit.send(cmd);
self.note_hybrid_route_success(); self.note_hybrid_route_success();
@@ -609,7 +869,8 @@ impl MePool {
)); ));
} }
Err(WriterCommandReserveError::Closed) => { Err(WriterCommandReserveError::Closed) => {
let Some(payload) = proxy_req_payload_from_command(cmd) else { let Some((payload, _permit)) = proxy_req_payload_from_command(cmd)
else {
return Err(ProxyError::Proxy( return Err(ProxyError::Proxy(
"ME writer rejected unexpected command type".into(), "ME writer rejected unexpected command type".into(),
)); ));
@@ -626,13 +887,14 @@ impl MePool {
payload.as_ref(), payload.as_ref(),
proto_flags, proto_flags,
tag.as_ref().map(|tag| tag.as_slice()), tag.as_ref().map(|tag| tag.as_slice()),
Some(_permit),
) )
.await; .await;
} }
} }
} }
Err(TrySendError::Closed(cmd)) => { Err(TrySendError::Closed(cmd)) => {
let Some(payload) = proxy_req_payload_from_command(cmd) else { let Some((payload, _permit)) = proxy_req_payload_from_command(cmd) else {
return Err(ProxyError::Proxy( return Err(ProxyError::Proxy(
"ME writer rejected unexpected command type".into(), "ME writer rejected unexpected command type".into(),
)); ));
@@ -649,6 +911,7 @@ impl MePool {
payload.as_ref(), payload.as_ref(),
proto_flags, proto_flags,
tag.as_ref().map(|tag| tag.as_slice()), tag.as_ref().map(|tag| tag.as_slice()),
Some(_permit),
) )
.await; .await;
} }
@@ -663,6 +926,7 @@ impl MePool {
payload.as_ref(), payload.as_ref(),
proto_flags, proto_flags,
tag.as_ref().map(|tag| tag.as_slice()), tag.as_ref().map(|tag| tag.as_slice()),
Some(_permit),
) )
.await .await
} }
+9 -3
View File
@@ -10,7 +10,7 @@ use crate::protocol::constants::{RPC_CLOSE_CONN_U32, RPC_CLOSE_EXT_U32};
use super::super::MePool; use super::super::MePool;
use super::super::codec::{WriterCommand, build_control_payload}; use super::super::codec::{WriterCommand, build_control_payload};
use super::{WriterCommandReserveError, reserve_writer_command_slot}; use super::{WriterCommandReserveError, reserve_writer_command_slot, writer_send_deadline};
const ME_CLOSE_SIGNAL_SEND_TIMEOUT: Duration = Duration::from_millis(50); const ME_CLOSE_SIGNAL_SEND_TIMEOUT: Duration = Duration::from_millis(50);
@@ -22,7 +22,10 @@ impl MePool {
match w.tx.try_send(WriterCommand::ControlAndFlush(payload)) { match w.tx.try_send(WriterCommand::ControlAndFlush(payload)) {
Ok(()) => {} Ok(()) => {}
Err(TrySendError::Full(cmd)) => { Err(TrySendError::Full(cmd)) => {
match reserve_writer_command_slot(&w.tx, Some(ME_CLOSE_SIGNAL_SEND_TIMEOUT)) match reserve_writer_command_slot(
&w.tx,
writer_send_deadline(Some(ME_CLOSE_SIGNAL_SEND_TIMEOUT)),
)
.await .await
{ {
Ok(permit) => { Ok(permit) => {
@@ -63,7 +66,10 @@ impl MePool {
match w.tx.try_send(WriterCommand::ControlAndFlush(payload)) { match w.tx.try_send(WriterCommand::ControlAndFlush(payload)) {
Ok(()) => {} Ok(()) => {}
Err(TrySendError::Full(cmd)) => { Err(TrySendError::Full(cmd)) => {
let _ = reserve_writer_command_slot(&w.tx, Some(ME_CLOSE_SIGNAL_SEND_TIMEOUT)) let _ = reserve_writer_command_slot(
&w.tx,
writer_send_deadline(Some(ME_CLOSE_SIGNAL_SEND_TIMEOUT)),
)
.await .await
.map(|permit| permit.send(cmd)); .map(|permit| permit.send(cmd));
} }
@@ -105,6 +105,7 @@ async fn make_pool(
general.me_writer_pick_sample_size, general.me_writer_pick_sample_size,
MeSocksKdfPolicy::default(), MeSocksKdfPolicy::default(),
general.me_writer_cmd_channel_capacity, general.me_writer_cmd_channel_capacity,
general.me_writer_byte_budget_bytes,
general.me_route_channel_capacity, general.me_route_channel_capacity,
general.me_route_backpressure_enabled, general.me_route_backpressure_enabled,
general.me_route_fairshare_enabled, general.me_route_fairshare_enabled,
@@ -134,6 +135,7 @@ async fn insert_draining_writer(
drain_deadline_epoch_secs: u64, drain_deadline_epoch_secs: u64,
) { ) {
let (tx, _writer_rx) = mpsc::channel::<WriterCommand>(8); let (tx, _writer_rx) = mpsc::channel::<WriterCommand>(8);
let byte_budget = pool.new_writer_byte_budget();
let writer = MeWriter { let writer = MeWriter {
id: writer_id, id: writer_id,
addr: SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 6000 + writer_id as u16), addr: SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 6000 + writer_id as u16),
@@ -143,6 +145,7 @@ async fn insert_draining_writer(
contour: Arc::new(AtomicU8::new(WriterContour::Draining.as_u8())), contour: Arc::new(AtomicU8::new(WriterContour::Draining.as_u8())),
created_at: Instant::now() - Duration::from_secs(writer_id), created_at: Instant::now() - Duration::from_secs(writer_id),
tx: tx.clone(), tx: tx.clone(),
byte_budget: byte_budget.clone(),
cancel: CancellationToken::new(), cancel: CancellationToken::new(),
degraded: Arc::new(AtomicBool::new(false)), degraded: Arc::new(AtomicBool::new(false)),
rtt_ema_ms_x10: Arc::new(AtomicU32::new(0)), rtt_ema_ms_x10: Arc::new(AtomicU32::new(0)),
@@ -153,7 +156,9 @@ async fn insert_draining_writer(
}; };
pool.writers.write().await.push(writer); pool.writers.write().await.push(writer);
pool.registry.register_writer(writer_id, tx).await; pool.registry
.register_writer(writer_id, tx, byte_budget)
.await;
pool.conn_count.fetch_add(1, Ordering::Relaxed); pool.conn_count.fetch_add(1, Ordering::Relaxed);
for idx in 0..bound_clients { for idx in 0..bound_clients {
@@ -103,6 +103,7 @@ async fn make_pool(
general.me_writer_pick_sample_size, general.me_writer_pick_sample_size,
MeSocksKdfPolicy::default(), MeSocksKdfPolicy::default(),
general.me_writer_cmd_channel_capacity, general.me_writer_cmd_channel_capacity,
general.me_writer_byte_budget_bytes,
general.me_route_channel_capacity, general.me_route_channel_capacity,
general.me_route_backpressure_enabled, general.me_route_backpressure_enabled,
general.me_route_fairshare_enabled, general.me_route_fairshare_enabled,
@@ -131,6 +132,7 @@ async fn insert_draining_writer(
drain_deadline_epoch_secs: u64, drain_deadline_epoch_secs: u64,
) { ) {
let (tx, _writer_rx) = mpsc::channel::<WriterCommand>(8); let (tx, _writer_rx) = mpsc::channel::<WriterCommand>(8);
let byte_budget = pool.new_writer_byte_budget();
let writer = MeWriter { let writer = MeWriter {
id: writer_id, id: writer_id,
addr: SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 5500 + writer_id as u16), addr: SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 5500 + writer_id as u16),
@@ -140,6 +142,7 @@ async fn insert_draining_writer(
contour: Arc::new(AtomicU8::new(WriterContour::Draining.as_u8())), contour: Arc::new(AtomicU8::new(WriterContour::Draining.as_u8())),
created_at: Instant::now() - Duration::from_secs(writer_id), created_at: Instant::now() - Duration::from_secs(writer_id),
tx: tx.clone(), tx: tx.clone(),
byte_budget: byte_budget.clone(),
cancel: CancellationToken::new(), cancel: CancellationToken::new(),
degraded: Arc::new(AtomicBool::new(false)), degraded: Arc::new(AtomicBool::new(false)),
rtt_ema_ms_x10: Arc::new(AtomicU32::new(0)), rtt_ema_ms_x10: Arc::new(AtomicU32::new(0)),
@@ -149,7 +152,9 @@ async fn insert_draining_writer(
allow_drain_fallback: Arc::new(AtomicBool::new(false)), allow_drain_fallback: Arc::new(AtomicBool::new(false)),
}; };
pool.writers.write().await.push(writer); pool.writers.write().await.push(writer);
pool.registry.register_writer(writer_id, tx).await; pool.registry
.register_writer(writer_id, tx, byte_budget)
.await;
pool.conn_count.fetch_add(1, Ordering::Relaxed); pool.conn_count.fetch_add(1, Ordering::Relaxed);
for idx in 0..bound_clients { for idx in 0..bound_clients {
let (conn_id, _rx) = pool.registry.register().await; let (conn_id, _rx) = pool.registry.register().await;
@@ -98,6 +98,7 @@ async fn make_pool(me_pool_drain_threshold: u64) -> Arc<MePool> {
general.me_writer_pick_sample_size, general.me_writer_pick_sample_size,
MeSocksKdfPolicy::default(), MeSocksKdfPolicy::default(),
general.me_writer_cmd_channel_capacity, general.me_writer_cmd_channel_capacity,
general.me_writer_byte_budget_bytes,
general.me_route_channel_capacity, general.me_route_channel_capacity,
general.me_route_backpressure_enabled, general.me_route_backpressure_enabled,
general.me_route_fairshare_enabled, general.me_route_fairshare_enabled,
@@ -126,6 +127,7 @@ async fn insert_draining_writer(
) -> Vec<u64> { ) -> Vec<u64> {
let mut conn_ids = Vec::with_capacity(bound_clients); let mut conn_ids = Vec::with_capacity(bound_clients);
let (tx, _writer_rx) = mpsc::channel::<WriterCommand>(8); let (tx, _writer_rx) = mpsc::channel::<WriterCommand>(8);
let byte_budget = pool.new_writer_byte_budget();
let writer = MeWriter { let writer = MeWriter {
id: writer_id, id: writer_id,
addr: SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 4500 + writer_id as u16), addr: SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 4500 + writer_id as u16),
@@ -135,6 +137,7 @@ async fn insert_draining_writer(
contour: Arc::new(AtomicU8::new(WriterContour::Draining.as_u8())), contour: Arc::new(AtomicU8::new(WriterContour::Draining.as_u8())),
created_at: Instant::now() - Duration::from_secs(writer_id), created_at: Instant::now() - Duration::from_secs(writer_id),
tx: tx.clone(), tx: tx.clone(),
byte_budget: byte_budget.clone(),
cancel: CancellationToken::new(), cancel: CancellationToken::new(),
degraded: Arc::new(AtomicBool::new(false)), degraded: Arc::new(AtomicBool::new(false)),
rtt_ema_ms_x10: Arc::new(AtomicU32::new(0)), rtt_ema_ms_x10: Arc::new(AtomicU32::new(0)),
@@ -144,7 +147,9 @@ async fn insert_draining_writer(
allow_drain_fallback: Arc::new(AtomicBool::new(false)), allow_drain_fallback: Arc::new(AtomicBool::new(false)),
}; };
pool.writers.write().await.push(writer); pool.writers.write().await.push(writer);
pool.registry.register_writer(writer_id, tx).await; pool.registry
.register_writer(writer_id, tx, byte_budget)
.await;
pool.conn_count.fetch_add(1, Ordering::Relaxed); pool.conn_count.fetch_add(1, Ordering::Relaxed);
for idx in 0..bound_clients { for idx in 0..bound_clients {
let (conn_id, _rx) = pool.registry.register().await; let (conn_id, _rx) = pool.registry.register().await;
@@ -87,6 +87,7 @@ async fn make_pool() -> Arc<MePool> {
general.me_writer_pick_sample_size, general.me_writer_pick_sample_size,
MeSocksKdfPolicy::default(), MeSocksKdfPolicy::default(),
general.me_writer_cmd_channel_capacity, general.me_writer_cmd_channel_capacity,
general.me_writer_byte_budget_bytes,
general.me_route_channel_capacity, general.me_route_channel_capacity,
general.me_route_backpressure_enabled, general.me_route_backpressure_enabled,
general.me_route_fairshare_enabled, general.me_route_fairshare_enabled,

Some files were not shown because too many files have changed in this diff Show More