diff --git a/.github/workflows/docker-publish.yml b/.github/workflows/docker-publish.yml index b5d2ffd..3b22465 100644 --- a/.github/workflows/docker-publish.yml +++ b/.github/workflows/docker-publish.yml @@ -36,6 +36,20 @@ jobs: - name: Set up Docker Buildx uses: docker/setup-buildx-action@v3 + - name: Verify release tag matches Cargo version + shell: bash + run: | + set -euo pipefail + package_version="$(awk ' + /^\[workspace.package\]$/ { in_workspace = 1; next } + in_workspace && /^version = / { + gsub(/"/, "", $3) + print $3 + exit + } + ' Cargo.toml)" + test "${GITHUB_REF_NAME}" = "v${package_version}" + - name: Log in to Docker Hub uses: docker/login-action@v3 with: @@ -47,7 +61,7 @@ jobs: uses: houseabsolute/actions-rust-cross@v1 with: target: ${{ matrix.target.arch }}-unknown-linux-musl - toolchain: 1.88.0 + toolchain: 1.98.0 args: "--locked --release --bin pb-mapper" strip: true diff --git a/.github/workflows/release-ui.yml b/.github/workflows/release-ui.yml index a108439..6cefaf3 100644 --- a/.github/workflows/release-ui.yml +++ b/.github/workflows/release-ui.yml @@ -80,7 +80,7 @@ jobs: - name: Install Rust uses: dtolnay/rust-toolchain@stable with: - toolchain: 1.88.0 + toolchain: 1.98.0 - name: Build latest Windows FFI run: | make build-pb-mapper-ffi-windows @@ -128,7 +128,7 @@ jobs: - name: Install Rust uses: dtolnay/rust-toolchain@stable with: - toolchain: 1.88.0 + toolchain: 1.98.0 - name: Install dependencies run: | sudo apt-get update -y @@ -236,7 +236,7 @@ jobs: - name: Install Rust uses: dtolnay/rust-toolchain@stable with: - toolchain: 1.88.0 + toolchain: 1.98.0 - name: Set up Android NDK uses: nttld/setup-ndk@v1 with: @@ -399,7 +399,7 @@ jobs: - name: Install Rust uses: dtolnay/rust-toolchain@stable with: - toolchain: 1.88.0 + toolchain: 1.98.0 - name: Install appdmg run: | npm install -g appdmg @@ -480,7 +480,7 @@ jobs: - name: Install Rust uses: dtolnay/rust-toolchain@stable with: - toolchain: 1.88.0 + toolchain: 1.98.0 - name: Build latest iOS FFI run: | make build-pb-mapper-ffi-ios diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml index 518b522..dabf432 100644 --- a/.github/workflows/release.yml +++ b/.github/workflows/release.yml @@ -61,7 +61,7 @@ jobs: uses: houseabsolute/actions-rust-cross@v1 with: target: ${{ matrix.platform.target }} - toolchain: 1.88.0 + toolchain: 1.98.0 args: "--locked --release --bin pb-mapper" strip: true @@ -77,3 +77,7 @@ jobs: LICENSE README.md README.zh-CN.md + docs/authentication-v2.md + docs/authentication-v2.zh-CN.md + docs/user-guide.md + docs/user-guide.zh-CN.md diff --git a/.github/workflows/syntax-check.yml b/.github/workflows/syntax-check.yml index 03e5e8d..e2804bc 100644 --- a/.github/workflows/syntax-check.yml +++ b/.github/workflows/syntax-check.yml @@ -62,8 +62,11 @@ jobs: # belongs to the Rust side and must be matched before ui/*. ui/native/*) rust=true ;; ui/*) flutter=true ;; - src/*|tests/*|examples/*) rust=true ;; + crates/*|src/*|tests/*|examples/*) rust=true ;; Cargo.toml|Cargo.lock|rust-toolchain.toml|rustfmt.toml) rust=true ;; + # Per-crate manifests. `case` patterns match the whole path, so + # the unanchored entry above only ever catches the root manifest. + */Cargo.toml|*/Cargo.lock) rust=true ;; # A change to this workflow has to prove itself on both. .github/workflows/syntax-check.yml) rust=true; flutter=true ;; esac @@ -88,7 +91,7 @@ jobs: - name: Install Rust uses: dtolnay/rust-toolchain@stable with: - toolchain: 1.88.0 + toolchain: 1.98.0 components: clippy, rustfmt - name: Cache Rust dependencies diff --git a/AGENTS.md b/AGENTS.md index 5c96352..e21de91 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -1,19 +1,24 @@ # Repository Guidelines ## Architecture Overview -- One `pb-mapper` binary in `src/bin/` with four role commands: +- One `pb-mapper` binary in `crates/pb-mapper-cli/src/bin/` with five role commands: - `server`: central router (default port 7666) - `register`: registers local TCP/UDP services with the router - `connect`: connects to a registered service and exposes a local port - `status`: queries router IDs and registered keys -- Core crates: `src/pb_server`, `src/local/{server,client}`, `src/common` (protocol, streams, listeners), `src/utils`. + - `admin`: issues, lists, and revokes credentials; rotates the administrator key +- Crates, bottom-up: `pb-mapper-core` (credentials, checksum, config, addressing) + → `pb-mapper-auth` (credential lifecycle and persistence) → `pb-mapper-protocol` + (framing and secure sessions) → `pb-mapper-server` and `pb-mapper-client`, which + are peers → `pb-mapper-cli`. `ui/native/pb_mapper_ffi` is the C ABI cdylib. ## Project Structure & Modules -- `src/`: Rust backend and CLI - - `src/bin/pb-mapper.rs`: unified CLI entry point - - `src/pb_server`, `src/local`, `src/common`, `src/utils` +- `crates/`: the Rust workspace; the root `Cargo.toml` is a virtual manifest + - `crates/pb-mapper-cli/src/bin/pb-mapper.rs`: unified CLI entry point + - `crates/pb-mapper-{core,auth,protocol,server,client,cli}` + - `crates/pb-mapper-cli/tests/`: integration tests; loads env from `tests/.env` + - `crates/pb-mapper-cli/examples/`: runnable examples - `ui/`: Flutter UI; Rust bridge under `ui/native/*` -- `tests/`: integration tests; loads env from `tests/.env` - `docker/`, `services/`, `scripts/`: container, systemd, build/release ## Build, Test, and Development Commands @@ -28,7 +33,9 @@ Notes: CI builds release artifacts on tags `vX.Y.Z` (see `.github/workflows/release.yml`). ## Coding Style & Naming Conventions -- Rust 2021; toolchain pinned via `rust-toolchain.toml` (CI uses 1.88.0) +- Edition is set once in `[workspace.package]`; the toolchain is pinned in + `rust-toolchain.toml`, which CI installs. Both are deliberately not repeated + here — a version in prose goes stale on the next upgrade. - Format: `cargo fmt --all` (4 spaces; import grouping per `rustfmt.toml`) - Lint: `cargo clippy --all-targets -- -D warnings` - Naming: modules/functions `snake_case`, types/traits `PascalCase`, consts `SCREAMING_SNAKE_CASE` diff --git a/CHANGELOG.md b/CHANGELOG.md index 14aa463..9d6e546 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -2,6 +2,17 @@ All notable changes to this project will be documented in this file. +## [0.4.0] - 2026-08-18 +- Added a sole administrator credential plus renewable, expiring, and immediately revocable `pbmt1_` temporary credentials with fixed-slot O(1) lookup and isolated per-key service namespaces. +- Added single-flight protocol-v2 authentication with directional AES-256-GCM keys, monotonic frame counters, authenticated routing metadata, durable first-flight replay protection, and optional legacy framing during migration. +- Added encrypted snapshot/WAL authentication state, exclusive `auth.lock`, lifecycle audit records, hierarchical timing-wheel expiry, hard closure of revoked live connections, recoverable root-key rotation, and explicit auth-state reset. +- Extended the unified CLI with temporary-key lifecycle, service/connection inventory, auth status, protocol policy, root rotation, namespace targeting, and human/JSON/NDJSON output. +- Replaced insecure default-key fallback with first-start random administrator-key generation, retained machine-derived keys only for explicit compatibility, and updated Flutter, installers, systemd, Docker, release metadata, and bilingual documentation. +- Recovery keys must decrypt existing snapshot or WAL state before they are persisted. Interrupted rotation and reset recover from staged `admin.key.next` and `server-instance-id.next`. +- First-flight salts are unique for nonce 0: admission is atomic under one lock, torn replay records fail closed, and a nonce-0 error frame is sent only after that salt is reserved. +- Pinned UI and local tunnels to the credential and relay address captured at start, and bound tunneled-frame checksums to each hop's authenticated session key. +- Aborted pooled registration workers and accepted connection tasks on shutdown; the relay reaps connection tasks with a `JoinSet`. + ## [0.3.0] - 2026-08-18 - Replaced the three role-specific executables with one `pb-mapper` CLI and explicit `server`, `register`, `connect`, and `status` commands. - Consolidated release archives into one cross-platform binary artifact per target and updated Docker, installers, systemd templates, build scripts, deployment skills, and documentation to use it. diff --git a/CLAUDE.md b/CLAUDE.md index 37d6151..8cb35e4 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -6,7 +6,7 @@ This file provides guidance to Claude Code (claude.ai/code) when working with co This is a Rust-based network tunneling/proxy system called `pb-mapper` that allows exposing local services to clients over a public network. The project enables users to access their home services (like file transfer servers) from anywhere by creating secure tunnels through a public server. -The system uses one **pb-mapper** binary (`src/bin/pb-mapper.rs`) with explicit role commands: +The system uses one **pb-mapper** binary (`crates/pb-mapper-cli/src/bin/pb-mapper.rs`) with explicit role commands: 1. **`pb-mapper server`**: Central server that manages connections between local services and clients - Runs on port 7666 by default @@ -25,7 +25,11 @@ The system uses one **pb-mapper** binary (`src/bin/pb-mapper.rs`) with explicit 4. **`pb-mapper status`**: Queries remote IDs and registered service keys -5. **UI Module** (`ui/`): Flutter graphical interface +5. **`pb-mapper admin`**: Administrator operations against a running server — + issuing, listing, and revoking temporary credentials, rotating the + administrator key, and listing services and connections + +6. **UI Module** (`ui/`): Flutter graphical interface - Replaces all CLI functionality with a user-friendly GUI - Calls into Rust through raw `dart:ffi` against the `pb-mapper-ffi` crate - Provides comprehensive service management interface @@ -36,62 +40,101 @@ The system works by creating a bridge between local services and remote clients ### Project Structure +The root `Cargo.toml` is a virtual manifest; every crate lives under `crates/`, +except the FFI cdylib, which sits next to the Flutter code that loads it. + ``` pb-mapper/ -├── src/ # Main Rust codebase -│ ├── bin/ # Unified pb-mapper CLI entry point -│ ├── pb_server/ # Central server implementation -│ ├── local/ # Local service handlers (server/client) -│ ├── common/ # Shared utilities and protocols -│ └── utils/ # Helper functions +├── crates/ +│ ├── pb-mapper-core/ # Bottom layer: checksum, config, conn_id, error, +│ │ # addr, codec, timeout, durable_file, DataLenType +│ ├── pb-mapper-auth/ # Credential lifecycle, persistence, timing wheel +│ ├── pb-mapper-protocol/ # Message framing, v2 secure sessions, forwarding +│ ├── pb-mapper-server/ # Central relay server, plus the task manager +│ ├── pb-mapper-client/ # Both tunnel ends: `register` and `connect` +│ └── pb-mapper-cli/ # The `pb-mapper` binary, integration tests, examples ├── ui/ # Flutter UI, talking to Rust over dart:ffi │ ├── lib/ # Flutter application code │ │ ├── l10n/ # ARB sources and generated AppLocalizations │ │ └── src/ffi/ # The Dart side of the FFI boundary │ ├── native/pb_mapper_ffi/ # C ABI crate (a workspace member) │ └── test/ # Widget tests -├── examples/ # Example implementations -├── tests/ # Integration tests ├── docker/ # Docker deployment configuration └── services/ # Systemd service files ``` +The dependency graph is a DAG, and the layering is what the crate split +encodes: + +``` +pb-mapper-cli pb-mapper-ffi + │ │ + └────┬─────────────────┤ + ▼ ▼ + pb-mapper-server pb-mapper-client (peers: no reference either way) + └──────┬──────────┘ + ▼ + pb-mapper-protocol + ▼ + pb-mapper-auth + ▼ + pb-mapper-core +``` + +Note that the binary is still named `pb-mapper`, discovered from +`src/bin/pb-mapper.rs` inside `pb-mapper-cli`. The release workflows, both +Dockerfiles, and the install scripts hardcode that name, and `cargo build --bin +pb-mapper` resolves it from the workspace root regardless of the crate name. +Likewise `pb-mapper-ffi` keeps its package name, because it determines the +`libpb_mapper_ffi.{so,dylib,a}` / `pb_mapper_ffi.dll` filenames that the Dart +loader, two CMakeLists, four xcconfigs, and the release-ui hash checks expect. + ### Core Modules -#### Rust Backend (`src/`) -- **`src/pb_server/`**: Central server implementation - - `server.rs`: Main server logic with connection management - - `client.rs`: Client connection handling - - `status.rs`: Server status reporting - - `mod.rs`: Server manager with ManagerTask and ConnTask enums - -- **`src/local/server/`**: Local service registration (`register` functionality) - - `stream.rs`: Stream handling for service registration - - `mod.rs`: Registration logic and server-side CLI implementation - - `error.rs`: Server-specific error handling - -- **`src/local/client/`**: Client connection handling (`connect` functionality) - - `stream.rs`: Stream management for client connections - - `status.rs`: Status checking and reporting - - `mod.rs`: Client-side CLI implementation - - `error.rs`: Client-specific error handling - -- **`src/common/`**: Shared utilities and protocols - - `message/`: Protocol definitions (command.rs, forward.rs) - - `config.rs`: Configuration management and environment variables - - `stream.rs`: Stream abstractions (TcpStreamProvider, UdpStreamProvider) - - `listener.rs`: Listener abstractions (TcpListenerProvider, UdpListenerProvider) - - `manager.rs`: Connection management utilities - - `buffer.rs`: Buffer management for data streaming - - `checksum.rs`: Data integrity verification - - `conn_id.rs`: Connection ID management - - `error.rs`: Common error definitions - -- **`src/utils/`**: Helper functions - - `addr.rs`: Address resolution with OneOrMore enum for multiple addresses - - `codec.rs`: Encryption/decryption utilities - - `timeout.rs`: Timeout handling mechanisms - - `udp.rs`: UDP-specific utilities +#### Rust Backend (`crates/`) +- **`pb-mapper-core/`**: The bottom layer; depends on no other crate here + - `checksum.rs`: The process credential, and the framing checksum over `datalen` + - `config.rs`: Environment configuration and address resolution entry points + - `conn_id.rs`: Connection ID types + - `error.rs`: The shared error type, plus the `snafu_error_*` macros + - `addr.rs`: Address resolution; custom DNS servers on the async path + - `codec.rs`: AES-256-GCM encrypt/decrypt + - `timeout.rs`: `RetryBackoff` + - `durable_file.rs`: Atomic replace and parent-directory fsync + - `test_support.rs`: `PROCESS_CREDENTIAL_TEST_LOCK`, shared across crates' tests + - `lib.rs`: `DataLenType`, which lives here so `checksum` and `error` can name it + +- **`pb-mapper-auth/`**: The credential subsystem, and the largest one + - `lib.rs`: `AuthRuntime`, `AuthContext`, `AuthFailure`, `KeyId` + - `runtime.rs`: Key derivation and authentication of a presented key + - `actor/`: The lifecycle actor — `epoch.rs` for root rotation + - `persistence/`: `snapshot.rs`, `wal.rs`, `blob.rs`, `admin_key.rs`, `fs.rs` + - `timing_wheel.rs`: Hierarchical wheel driving credential expiry + - `leases.rs`, `keys.rs`, `ids.rs`, `config.rs`: Leases, key material, platform dirs + +- **`pb-mapper-protocol/`**: Framing and the authenticated session + - `lib.rs`: The checksum + length framing, and the reader/writer traits + - `command.rs`: Request/response types (`PbConnRequest`, `LocalServer`, `AdminRequest`, …) + - `secure.rs`: Protocol-v2 single-flight sessions, client and server + - `secure/`: `frame.rs`, `first_flight.rs`, `replay.rs`, `limiter.rs` + - `forward.rs`: Stream and datagram forwarding + - `buffer.rs`: Read buffers for the framing + +- **`pb-mapper-server/`**: The central relay + - `lib.rs`: `ManagerTask` / `ConnTask`, and the routing domain model + - `runtime.rs`: Serialises the global routing maps and quotas (the largest file) + - `connection.rs`: Per-socket authentication and dispatch + - `server.rs`, `client.rs`: The service-side and subscriber-side loops + - `admin.rs`: Administrator request handling + - `status.rs`, `error.rs`, `manager.rs`: Status replies, errors, the task manager + +- **`pb-mapper-client/`**: Both ends of a tunnel + - `server/`: `register` — publishes a local service (`mod.rs`, `stream.rs`, `error.rs`) + - `client/`: `connect` — subscribes and listens locally, plus `status.rs` + +- **`pb-mapper-cli/`**: The binary, integration tests, and examples + - `src/bin/pb-mapper.rs`: Argument parsing and the role commands + - `src/bin/pb-mapper/admin.rs`: The `admin` subcommand #### Flutter UI (`ui/`) - **`lib/src/views/`**: One file per zone the shell can show @@ -126,12 +169,17 @@ pb-mapper/ ### Key Components -1. **Message Protocol** (`src/common/message/`): +1. **Message Protocol** (`crates/pb-mapper-protocol/`): - **Command Protocol** (`command.rs`): Defines request/response types: - `PbConnStatusReq`/`PbConnStatusResp`: Status checking - `PbConnRequest`/`PbConnResponse`: Connection management - `PbServerRequest`: Server operation requests - - `LocalService`: Service type definitions (TCP/UDP) + - `LocalServer`: Service type definitions (TCP/UDP) + - `AdminRequest`/`AdminResponse`: Administrator operations + - **Secure sessions** (`secure.rs`): Protocol-v2 first flight — the initial + frame carries a clear-text routing prefix plus an authenticated encrypted + request, adding no extra round trip, and later frames on the connection use + directional keys with monotonic counters - **Forward Protocol** (`forward.rs`): Data forwarding mechanisms - Uses JSON serialization with custom framing (checksum + length header) - Supports encryption/decryption for secure communication via ring crate @@ -142,15 +190,17 @@ pb-mapper/ - Implements keep-alive and timeout mechanisms - Uses actor model for concurrent connection handling -3. **Stream Abstractions**: - - `StreamProvider` trait for TCP/UDP stream handling - - `ListenerProvider` trait for TCP/UDP listener management - - Unified interface for different transport protocols +3. **Stream Abstractions**: `StreamProvider` and `ListenerProvider` give TCP and + UDP one interface. These live in the external `uni-stream` crate, not in this + repository. + +4. **Authentication** (`crates/pb-mapper-auth/`): An administrator key plus + derived temporary credentials, persisted through a write-ahead log and + snapshots, with expiry driven by a hierarchical timing wheel. See + `docs/authentication-v2.md`. -4. **Configuration System**: - - Environment variable support: - - `PB_MAPPER_SERVER`: Remote server address - - `PB_MAPPER_KEEP_ALIVE`: TCP keep-alive setting +5. **Configuration System**: + - Environment variables (see Environment Variables below) - Command-line argument parsing with clap - Workspace-based dependency management @@ -184,16 +234,20 @@ and it is what lets a widget test substitute `FakePbMapperApi` ### Current UI Implementation Status -The UI is fully implemented with the following structure: +Every view under `ui/lib/src/views/`: - **Main App** (`ui/lib/main.dart`): Entry point with navigation and theme management - **Landing Page** (`main_landing_view.dart`): Central navigation hub -- **Server Management** (`server_management_page.dart`, `server_management_view.dart`): Complete server control -- **Service Registration** (`service_registration_page.dart`, `service_registration_view.dart`): Service registration interface -- **Client Connection** (`client_connection_page.dart`, `client_connection_view.dart`): Client connection management +- **Setup Wizard** (`setup_wizard_view.dart`): First-run guided setup +- **Service Registration** (`service_registration_view.dart`): The register workspace +- **Registered Services** (`registered_services_view.dart`): What this process has registered +- **Client Connection** (`client_connection_view.dart`): The connect workspace - **Status Monitoring** (`status_monitoring_view.dart`): Real-time status dashboard - **Configuration** (`configuration_view.dart`): Environment and settings management -- **Logging** (`log_display_widget.dart`, `log_manager.dart`): Comprehensive log viewing +- **Logging** (`log_view_page.dart`, `src/common/log_manager.dart` (under `ui/lib/`)): The log stream + +There is no separate server-management view: starting and stopping the relay is +part of the landing page and the setup wizard. ### UI Features Implemented @@ -265,29 +319,34 @@ The UI is fully implemented with the following structure: - **FFI Integration**: Direct `dart:ffi` calls into the `pb-mapper-ffi` crate - **Real-time Updates**: Live status monitoring and log streaming - **Configuration Management**: Persistent settings and environment variable management -- **Multi-platform**: Desktop, mobile, and web support +- **Multi-platform**: Desktop and mobile. There is no web/wasm target — the UI + loads a native library over `dart:ffi`, which the web cannot do. ## Development Notes ### Project Structure & Dependencies -- **Workspace Configuration**: Multi-crate workspace with shared dependencies in root `Cargo.toml` +- **Workspace Configuration**: Virtual manifest at the root; versions are pinned + once in `[workspace.dependencies]` and crates take them with `.workspace = true` - **Memory Optimization**: Uses mimalloc-rust for improved memory allocation performance -- **Error Handling**: Comprehensive error handling with snafu crate across all modules +- **Error Handling**: snafu, with each crate owning its own error type and wrapping + the layer below as a `source` rather than sharing one workspace-wide enum - **Async Runtime**: Built on Tokio with full async/await support - **Serialization**: serde and serde_json for message serialization -- **Networking**: socket2 for low-level socket operations, trust-dns-resolver for DNS +- **Networking**: uni-stream for the stream/listener abstractions, hickory-resolver + for DNS (custom resolvers on the async path only — the sync path uses `std`, + since hickory has no blocking resolver) - **Cryptography**: ring crate for encryption/decryption functionality ### Code Quality & Standards -- **Linting**: Strict clippy rules in UI native hub (deny unwrap_used, expect_used, wildcard_imports) +- **Linting**: `unwrap_used` and `expect_used` are denied for the whole workspace + via `[workspace.lints]`; `clippy.toml` exempts test code, and `tests/` and + `examples/` targets carry a file-level allow. A production `unwrap` needs a + reason recorded at the site. - **Formatting**: rustfmt.toml configuration for consistent code style - **Toolchain**: rust-toolchain.toml for reproducible builds -- **Testing**: Comprehensive test suite in `tests/` directory - -### Build Profiles -- **wasm-dev**: Optimized for WebAssembly builds -- **server-dev**: Development profile for server components -- **android-dev**: Android-specific build optimizations +- **Testing**: Unit tests live beside the code; integration tests are in + `crates/pb-mapper-cli/tests/`, which is the crate that depends on every layer + they exercise ### UI Development Guidelines - **Framework**: Flutter 3.44.9, Material 3. CI pins the same version. @@ -299,9 +358,23 @@ The UI is fully implemented with the following structure: - **Responsive Design**: Adaptive layouts for different screen sizes ### Environment Variables + +The commonly used ones: + - **`PB_MAPPER_SERVER`**: Default remote server address for CLI tools -- **`PB_MAPPER_KEEP_ALIVE`**: Global TCP keep-alive setting ("ON" to enable) +- **`PB_MAPPER_KEEP_ALIVE`**: TCP keep-alive ("ON", "1", "true", "yes" to enable). + Read on every call, not cached — the UI's per-service toggle depends on that. +- **`MSG_HEADER_KEY`**: The process credential, administrator or temporary. + Required; there is no insecure default. - **`RUST_LOG`**: Tracing level configuration (supports env-filter) +- **`PB_MAPPER_LOG_FORMAT`**: Log output format + +Timeouts, intervals, and pool sizes are also configurable, and there are more +than a dozen: the authoritative list is the `pub const PB_MAPPER_*` declarations +at the top of `crates/pb-mapper-core/src/config.rs`, each read by the accessor +named after it. Beyond those, `PB_MAPPER_AUTH_STATE_DIR`, +`PB_MAPPER_LEGACY_PROTOCOL`, and `PB_MAPPER_NEW_STREAMS_PER_SECOND` are read by +name where they are used; the first two are also settable as `server` flags. ## Development Workflow diff --git a/Cargo.lock b/Cargo.lock index d0d1077..3e44e6f 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -81,9 +81,21 @@ checksum = "9035ad2d096bed7955a320ee7e2230574d28fd3c3a0f186cbea1ff3c7eed5dbb" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.114", ] +[[package]] +name = "autocfg" +version = "1.5.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f2032f911046de80f0a198e0901378627c33f59ea0ac00e363d481118bd70a53" + +[[package]] +name = "base64" +version = "0.23.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ac07cdecf99051d9a5238b80f35af32cdeba5b336e55d957b318b50137e18da5" + [[package]] name = "better_mimalloc_rs" version = "0.1.2" @@ -119,6 +131,12 @@ dependencies = [ "rustc_version", ] +[[package]] +name = "bumpalo" +version = "3.20.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "72f5acc6cb2ba439de613abc23857ec3d78374d8ed5ac84e9d11336e87da8649" + [[package]] name = "bytes" version = "1.11.0" @@ -189,7 +207,7 @@ dependencies = [ "heck", "proc-macro2", "quote", - "syn", + "syn 2.0.114", ] [[package]] @@ -204,6 +222,32 @@ version = "1.0.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b05b61dc5112cbb17e4b6cd61790d9845d13888356391624cbe7e41efeac1e75" +[[package]] +name = "combine" +version = "4.6.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ba5a308b75df32fe02788e748662718f03fde005016435c444eea572398219fd" +dependencies = [ + "bytes", + "memchr", +] + +[[package]] +name = "core-foundation" +version = "0.9.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "91e195e091a93c46f7102ec7818a2aa394e1e1771c3ab4825963fa03e45afb8f" +dependencies = [ + "core-foundation-sys", + "libc", +] + +[[package]] +name = "core-foundation-sys" +version = "0.8.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "773648b94d0e5d620f64f280777445740e61fe701025087ec8b57f45c791888b" + [[package]] name = "cpufeatures" version = "0.3.0" @@ -213,6 +257,36 @@ dependencies = [ "libc", ] +[[package]] +name = "critical-section" +version = "1.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "790eea4361631c5e7d22598ecd5723ff611904e3344ce8720784c93e3d83d40b" + +[[package]] +name = "crossbeam-channel" +version = "0.5.16" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d85363c37faeca707aef026efa9f3b34d077bce547e48f770770625c6013679e" +dependencies = [ + "crossbeam-utils", +] + +[[package]] +name = "crossbeam-epoch" +version = "0.9.20" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2d6914041f254d6e9176c01941b21115dcfb7089e55135a35411081bd106ef3f" +dependencies = [ + "crossbeam-utils", +] + +[[package]] +name = "crossbeam-utils" +version = "0.8.22" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "61803da095bee82a81bb1a452ecc25d3b2f1416d1897eb86430c6159ef717c17" + [[package]] name = "cty" version = "0.2.2" @@ -227,23 +301,23 @@ checksum = "d7a1e2f27636f116493b8b860f5546edb47c8d8f8ea73e1d2a20be88e28d1fea" [[package]] name = "dirs" -version = "5.0.1" +version = "6.0.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "44c45a9d03d6676652bcb5e724c7e988de1acad23a711b5217ab9cbecbec2225" +checksum = "c3e8aa94d75141228480295a7d0e7feb620b1a5ad9f12bc40be62411e38cce4e" dependencies = [ "dirs-sys", ] [[package]] name = "dirs-sys" -version = "0.4.1" +version = "0.5.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "520f05a5cbd335fae5a99ff7a6ab8627577660ee5cfd6a94a6a929b52ff0321c" +checksum = "e01a3366d27ee9890022452ee61b2b63a67e6f13f58900b651ff5665f0bb1fab" dependencies = [ "libc", "option-ext", "redox_users", - "windows-sys 0.48.0", + "windows-sys 0.61.2", ] [[package]] @@ -254,7 +328,7 @@ checksum = "97369cbbc041bc366949bc74d34658d6cda5621039731c6310521892a3a20ae0" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.114", ] [[package]] @@ -263,6 +337,12 @@ version = "0.15.7" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1aaf95b3e5c8f23aa320147307562d361db0ae0d51242340f558153b4eb2439b" +[[package]] +name = "either" +version = "1.18.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "252afb9ae5eaa683babdc6a068b3f5726eb19e05070c731f9b2a23a7c3e8ed34" + [[package]] name = "enum-as-inner" version = "0.6.1" @@ -272,7 +352,7 @@ dependencies = [ "heck", "proc-macro2", "quote", - "syn", + "syn 2.0.114", ] [[package]] @@ -374,7 +454,7 @@ checksum = "162ee34ebcb7c64a8abebc059ce0fee27c2262618d7b60ed8faf72fef13c3650" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.114", ] [[package]] @@ -452,12 +532,93 @@ dependencies = [ "foldhash 0.2.0", ] +[[package]] +name = "hashbrown" +version = "0.17.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ed5909b6e89a2db4456e54cd5f673791d7eca6732202bbf2a9cc504fe2f9b84a" +dependencies = [ + "allocator-api2", + "equivalent", + "foldhash 0.2.0", +] + [[package]] name = "heck" version = "0.5.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "2304e00983f87ffb38b55b444b5e3b60a884b5d30c0fca7d82fe33449bbe55ea" +[[package]] +name = "hickory-net" +version = "0.26.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e2295ed2f9c31e471e1428a8f88a3f0e1f4b27c15049592138d1eebe9c35b183" +dependencies = [ + "async-trait", + "cfg-if", + "data-encoding", + "futures-channel", + "futures-io", + "futures-util", + "hickory-proto", + "idna 1.1.0", + "ipnet", + "jni", + "rand 0.10.0", + "thiserror 2.0.20", + "tinyvec", + "tokio", + "tracing", + "url", +] + +[[package]] +name = "hickory-proto" +version = "0.26.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0bab31817bfb44672a252e97fe81cd0c18d1b2cf892108922f6818820df8c643" +dependencies = [ + "data-encoding", + "idna 1.1.0", + "ipnet", + "jni", + "once_cell", + "prefix-trie", + "rand 0.10.0", + "ring", + "thiserror 2.0.20", + "tinyvec", + "tracing", + "url", +] + +[[package]] +name = "hickory-resolver" +version = "0.26.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f0d58d28879ceecde6607729660c2667a081ccdc082e082675042793960f178c" +dependencies = [ + "cfg-if", + "futures-util", + "hickory-net", + "hickory-proto", + "ipconfig", + "ipnet", + "jni", + "moka", + "ndk-context", + "once_cell", + "parking_lot", + "rand 0.10.0", + "resolv-conf", + "smallvec", + "system-configuration", + "thiserror 2.0.20", + "tokio", + "tracing", +] + [[package]] name = "icu_collections" version = "2.1.1" @@ -605,6 +766,9 @@ name = "ipnet" version = "2.11.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "469fb0b9cefa57e3ef31275ee7cacb78f2fdca44e4765491884a2b119d4eb130" +dependencies = [ + "serde", +] [[package]] name = "is_terminal_polyfill" @@ -618,6 +782,66 @@ version = "1.0.17" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "92ecc6618181def0457392ccd0ee51198e065e016d1d527a7ac1b6dc7c1f09d2" +[[package]] +name = "jni" +version = "0.22.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5efd9a482cf3a427f00d6b35f14332adc7902ce91efb778580e180ff90fa3498" +dependencies = [ + "cfg-if", + "combine", + "jni-macros", + "jni-sys", + "log", + "simd_cesu8", + "thiserror 2.0.20", + "walkdir", + "windows-link", +] + +[[package]] +name = "jni-macros" +version = "0.22.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a00109accc170f0bdb141fed3e393c565b6f5e072365c3bd58f5b062591560a3" +dependencies = [ + "proc-macro2", + "quote", + "rustc_version", + "simd_cesu8", + "syn 2.0.114", +] + +[[package]] +name = "jni-sys" +version = "0.4.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c6377a88cb3910bee9b0fa88d4f42e1d2da8e79915598f65fb0c7ee14c878af2" +dependencies = [ + "jni-sys-macros", +] + +[[package]] +name = "jni-sys-macros" +version = "0.4.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "38c0b942f458fe50cdac086d2f946512305e5631e720728f2a61aabcd47a6264" +dependencies = [ + "quote", + "syn 2.0.114", +] + +[[package]] +name = "js-sys" +version = "0.3.104" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0e0c1080212aad755ea003d18543e8768dd432c48819efd73a7bf1e39b7a5a3a" +dependencies = [ + "cfg-if", + "futures-util", + "wasm-bindgen", +] + [[package]] name = "kanal" version = "0.2.0-beta2" @@ -719,6 +943,29 @@ dependencies = [ "windows-sys 0.61.2", ] +[[package]] +name = "moka" +version = "0.12.16" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4293f18e7567a1caf3c584855554377025c65e0aa445344d04171f5ad63d19b9" +dependencies = [ + "crossbeam-channel", + "crossbeam-epoch", + "crossbeam-utils", + "equivalent", + "parking_lot", + "portable-atomic", + "smallvec", + "tagptr", + "uuid", +] + +[[package]] +name = "ndk-context" +version = "0.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "27b02d87554356db9e9a873add8782d4ea6e3e58ea071a9adb9a2e8ddb884a8b" + [[package]] name = "nu-ansi-term" version = "0.50.3" @@ -728,11 +975,24 @@ dependencies = [ "windows-sys 0.61.2", ] +[[package]] +name = "num-traits" +version = "0.2.19" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "071dfc062690e90b734c0b2273ce72ad0ffa95f0c74596bc250dcfd960262841" +dependencies = [ + "autocfg", +] + [[package]] name = "once_cell" version = "1.21.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "42f5e15c9953c5e4ccceeb2e7382a716482c34515315f7b03532b8b4e8393d2d" +dependencies = [ + "critical-section", + "portable-atomic", +] [[package]] name = "once_cell_polyfill" @@ -770,29 +1030,69 @@ dependencies = [ ] [[package]] -name = "pb-mapper" -version = "0.3.0" +name = "pb-mapper-auth" +version = "0.4.0" +dependencies = [ + "parking_lot", + "pb-mapper-core", + "rand 0.10.0", + "ring", + "serde", + "serde_json", + "subtle", + "tokio", + "tokio-util", + "tracing", +] + +[[package]] +name = "pb-mapper-cli" +version = "0.4.0" dependencies = [ "better_mimalloc_rs", - "bytes", "clap", "dotenvy", - "futures", - "hashbrown 0.16.1", - "kanal", - "once_cell", + "pb-mapper-auth", + "pb-mapper-client", + "pb-mapper-core", + "pb-mapper-protocol", + "pb-mapper-server", + "rand 0.10.0", + "serde_json", + "tokio", + "tokio-util", + "tracing", + "uni-stream", +] + +[[package]] +name = "pb-mapper-client" +version = "0.4.0" +dependencies = [ + "pb-mapper-core", + "pb-mapper-protocol", + "serde_json", + "snafu", + "tokio", + "tracing", + "uni-stream", +] + +[[package]] +name = "pb-mapper-core" +version = "0.4.0" +dependencies = [ + "base64", + "clap", + "hickory-resolver", + "parking_lot", "rand 0.10.0", "ring", - "serde", "serde_json", "snafu", - "socket2 0.6.1", "tokio", - "tokio-util", "tracing", "tracing-subscriber", - "trust-dns-resolver", - "uni-stream", ] [[package]] @@ -802,7 +1102,12 @@ dependencies = [ "better_mimalloc_rs", "clap", "dirs", - "pb-mapper", + "parking_lot", + "pb-mapper-auth", + "pb-mapper-client", + "pb-mapper-core", + "pb-mapper-protocol", + "pb-mapper-server", "serde", "serde_json", "tokio", @@ -812,6 +1117,41 @@ dependencies = [ "uni-stream", ] +[[package]] +name = "pb-mapper-protocol" +version = "0.4.0" +dependencies = [ + "bytes", + "parking_lot", + "pb-mapper-auth", + "pb-mapper-core", + "rand 0.10.0", + "ring", + "serde", + "serde_json", + "snafu", + "tokio", + "tracing", + "uni-stream", +] + +[[package]] +name = "pb-mapper-server" +version = "0.4.0" +dependencies = [ + "hashbrown 0.17.1", + "kanal", + "pb-mapper-auth", + "pb-mapper-core", + "pb-mapper-protocol", + "rand 0.10.0", + "snafu", + "tokio", + "tokio-util", + "tracing", + "uni-stream", +] + [[package]] name = "percent-encoding" version = "2.3.2" @@ -830,6 +1170,12 @@ version = "0.1.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8b870d8c151b6f2fb93e84a13146138f05d02ed11c7e7c54f8826aaaf7c9f184" +[[package]] +name = "portable-atomic" +version = "1.15.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "05c8b63e8d9609db387f0324918f81d68fe27748f084ef092fb35954d0539a85" + [[package]] name = "potential_utf" version = "0.1.4" @@ -848,6 +1194,17 @@ dependencies = [ "zerocopy", ] +[[package]] +name = "prefix-trie" +version = "0.8.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4cf6e3177f0684016a5c209b00882e15f8bdd3f3bb48f0491df10cd102d0c6e7" +dependencies = [ + "either", + "ipnet", + "num-traits", +] + [[package]] name = "prettyplease" version = "0.2.37" @@ -855,7 +1212,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "479ca8adacdd7ce8f1fb39ce9ecccbfe93a3f1344b3d0d97f20bc0196208f62b" dependencies = [ "proc-macro2", - "syn", + "syn 2.0.114", ] [[package]] @@ -940,13 +1297,13 @@ dependencies = [ [[package]] name = "redox_users" -version = "0.4.6" +version = "0.5.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "ba009ff324d1fc1b900bd1fdb31564febe58a8ccc8a6fdbb93b543d33b13ca43" +checksum = "a4e608c6638b9c18977b00b475ac1f28d14e84b27d8d42f70e0bf1e3dec127ac" dependencies = [ "getrandom 0.2.17", "libredox", - "thiserror", + "thiserror 2.0.20", ] [[package]] @@ -995,6 +1352,21 @@ dependencies = [ "semver", ] +[[package]] +name = "rustversion" +version = "1.0.23" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cf54715a573b99ac80df0bc206da022bcd442c974952c7b9720069370852e21f" + +[[package]] +name = "same-file" +version = "1.0.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "93fc1dc3aaa9bfed95e02e6eadabb4baf7e3078b0bd1b4d7b6b0b68378900502" +dependencies = [ + "winapi-util", +] + [[package]] name = "scopeguard" version = "1.2.0" @@ -1034,7 +1406,7 @@ checksum = "d540f220d3187173da220f885ab66608367b6574e925011a9353e4badda91d79" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.114", ] [[package]] @@ -1075,6 +1447,22 @@ dependencies = [ "libc", ] +[[package]] +name = "simd_cesu8" +version = "1.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "11031e251abf8611c80f460e19dbdeb54a66db918e49c65a7065b46ac7aec520" +dependencies = [ + "rustc_version", + "simdutf8", +] + +[[package]] +name = "simdutf8" +version = "0.1.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e3a9fe34e3e7a50316060351f37187a3f546bce95496156754b601a5fa71b76e" + [[package]] name = "slab" version = "0.4.11" @@ -1089,23 +1477,23 @@ checksum = "67b1b7a3b5fe4f1376887184045fcf45c69e92af734b7aaddc05fb777b6fbd03" [[package]] name = "snafu" -version = "0.8.9" +version = "0.9.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "6e84b3f4eacbf3a1ce05eac6763b4d629d60cbc94d632e4092c54ade71f1e1a2" +checksum = "e45cb604038abb7b926b679887b3226d8d0f23874b66623625a0454be425a4b7" dependencies = [ "snafu-derive", ] [[package]] name = "snafu-derive" -version = "0.8.9" +version = "0.9.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c1c97747dbf44bb1ca44a561ece23508e99cb592e862f22222dcf42f51d1e451" +checksum = "287f59010008f0d7cf5e3b03196d666c1acc46c8d3e9cf34c28a1a7157601e72" dependencies = [ "heck", "proc-macro2", "quote", - "syn", + "syn 2.0.114", ] [[package]] @@ -1140,6 +1528,12 @@ version = "0.11.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7da8b5736845d9f2fcb837ea5d9e2628564b3b043a70948a3f0b778838c5fb4f" +[[package]] +name = "subtle" +version = "2.6.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "13c2bddecc57b384dee18652358fb23172facb8a2c51ccc10d74c157bdea3292" + [[package]] name = "syn" version = "2.0.114" @@ -1151,6 +1545,17 @@ dependencies = [ "unicode-ident", ] +[[package]] +name = "syn" +version = "3.0.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "53e9bae58849f64dfa4f5d5ae372c8341f7305f82a3868709269343628b659a3" +dependencies = [ + "proc-macro2", + "quote", + "unicode-ident", +] + [[package]] name = "synstructure" version = "0.13.2" @@ -1159,16 +1564,52 @@ checksum = "728a70f3dbaf5bab7f0c4b1ac8d7ae5ea60a4b5549c8a5914361c99147a709d2" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.114", +] + +[[package]] +name = "system-configuration" +version = "0.7.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a13f3d0daba03132c0aa9767f98351b3488edc2c100cda2d2ec2b04f3d8d3c8b" +dependencies = [ + "bitflags", + "core-foundation", + "system-configuration-sys", +] + +[[package]] +name = "system-configuration-sys" +version = "0.6.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8e1d1b10ced5ca923a1fcb8d03e96b8d3268065d724548c0211415ff6ac6bac4" +dependencies = [ + "core-foundation-sys", + "libc", ] +[[package]] +name = "tagptr" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7b2093cf4c8eb1e67749a6762251bc9cd836b6fc171623bd0a9d324d37af2417" + [[package]] name = "thiserror" version = "1.0.69" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b6aaf5339b578ea85b50e080feb250a3e8ae8cfcdff9a461c9ec2904bc923f52" dependencies = [ - "thiserror-impl", + "thiserror-impl 1.0.69", +] + +[[package]] +name = "thiserror" +version = "2.0.20" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ec86235f5fcc2a73650310756d2ac5b138a5780bbbdfae3eeccec992c435ba4f" +dependencies = [ + "thiserror-impl 2.0.20", ] [[package]] @@ -1179,7 +1620,18 @@ checksum = "4fee6c4efc90059e10f81e6d42c60a18f76588c3d74cb83a0b242a2b6c7504c1" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.114", +] + +[[package]] +name = "thiserror-impl" +version = "2.0.20" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bc04cd3e1236dd4a98afca4569f2deb3f120e5422a4023be2cb683f8486292af" +dependencies = [ + "proc-macro2", + "quote", + "syn 3.0.3", ] [[package]] @@ -1241,7 +1693,7 @@ checksum = "af407857209536a95c8e56f8231ef2c2e2aff839b22e07a1ffcbc617e9db9fa5" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.114", ] [[package]] @@ -1276,7 +1728,7 @@ checksum = "7490cfa5ec963746568740651ac6781f701c9c5ea257c58e057f3ba8cf69e8da" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.114", ] [[package]] @@ -1349,7 +1801,7 @@ dependencies = [ "once_cell", "rand 0.8.5", "smallvec", - "thiserror", + "thiserror 1.0.69", "tinyvec", "tokio", "tracing", @@ -1371,7 +1823,7 @@ dependencies = [ "rand 0.8.5", "resolv-conf", "smallvec", - "thiserror", + "thiserror 1.0.69", "tokio", "tracing", "trust-dns-proto", @@ -1450,12 +1902,33 @@ version = "0.2.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "06abde3611657adf66d383f00b093d7faecc7fa57071cce2578660c9f1010821" +[[package]] +name = "uuid" +version = "1.24.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2cefc03fd367c0c6d4305de1b312cf00248c4114f4a0418ce6a6af769e3b0bd9" +dependencies = [ + "getrandom 0.4.1", + "js-sys", + "wasm-bindgen", +] + [[package]] name = "valuable" version = "0.1.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ba73ea9cf16a25df0c8caa16c51acb937d5712a8429db78a3ee29d5dcacd3a65" +[[package]] +name = "walkdir" +version = "2.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "29790946404f91d9c5d06f9874efddea1dc06c5efe94541a7d6863108e3a5e4b" +dependencies = [ + "same-file", + "winapi-util", +] + [[package]] name = "wasi" version = "0.11.1+wasi-snapshot-preview1" @@ -1480,6 +1953,51 @@ dependencies = [ "wit-bindgen", ] +[[package]] +name = "wasm-bindgen" +version = "0.2.127" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1b70935747edd64d89de3efa29d73789b806c15798f8e7dca4d8ac356b50ce70" +dependencies = [ + "cfg-if", + "once_cell", + "rustversion", + "wasm-bindgen-macro", + "wasm-bindgen-shared", +] + +[[package]] +name = "wasm-bindgen-macro" +version = "0.2.127" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "77775f8f3f7217702089053b94958f8f54061a3f663417df76e19cbdcca29bc1" +dependencies = [ + "quote", + "wasm-bindgen-macro-support", +] + +[[package]] +name = "wasm-bindgen-macro-support" +version = "0.2.127" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e11d33f857dc2fb11b8bc75aee111aa9cbeb12cd9f25efd3d4c2a3dd4e235284" +dependencies = [ + "bumpalo", + "proc-macro2", + "quote", + "syn 2.0.114", + "wasm-bindgen-shared", +] + +[[package]] +name = "wasm-bindgen-shared" +version = "0.2.127" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7ef64dbcc55df09c7e5a46182d181c2cfa3e925f3da937ea764728b4bbb9dcbf" +dependencies = [ + "unicode-ident", +] + [[package]] name = "wasm-encoder" version = "0.244.0" @@ -1520,6 +2038,15 @@ version = "1.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "72069c3113ab32ab29e5584db3c6ec55d416895e60715417b5b883a357c3e471" +[[package]] +name = "winapi-util" +version = "0.1.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c2a7b1c03c876122aa43f3020e6c3c3ee5c05081c9a00739faf7503aeba10d22" +dependencies = [ + "windows-sys 0.61.2", +] + [[package]] name = "windows-link" version = "0.2.1" @@ -1788,7 +2315,7 @@ dependencies = [ "heck", "indexmap", "prettyplease", - "syn", + "syn 2.0.114", "wasm-metadata", "wit-bindgen-core", "wit-component", @@ -1804,7 +2331,7 @@ dependencies = [ "prettyplease", "proc-macro2", "quote", - "syn", + "syn 2.0.114", "wit-bindgen-core", "wit-bindgen-rust", ] @@ -1871,7 +2398,7 @@ checksum = "b659052874eb698efe5b9e8cf382204678a0086ebf46982b79d6ca3182927e5d" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.114", "synstructure", ] @@ -1892,7 +2419,7 @@ checksum = "2c7962b26b0a8685668b671ee4b54d007a67d4eaf05fda79ac0ecf41e32270f1" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.114", ] [[package]] @@ -1912,7 +2439,7 @@ checksum = "d71e5d6e06ab090c67b5e44993ec16b72dcbaabc526db883a360057678b48502" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.114", "synstructure", ] @@ -1946,7 +2473,7 @@ checksum = "eadce39539ca5cb3985590102671f2567e659fca9666581ad3411d59207951f3" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.114", ] [[package]] diff --git a/Cargo.toml b/Cargo.toml index 68584db..4848d5f 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,80 +1,52 @@ -[package] -name = "pb-mapper" -version.workspace = true -edition.workspace = true -authors.workspace = true - -# See more keys and their definitions at https://doc.rust-lang.org/cargo/reference/manifest.html - -[dependencies] -rand.workspace = true -socket2.workspace = true -tokio.workspace = true -tokio-util.workspace = true -snafu.workspace = true -serde.workspace = true -serde_json.workspace = true -tracing.workspace = true -tracing-subscriber.workspace = true -hashbrown.workspace = true -clap.workspace = true -futures.workspace = true -better_mimalloc_rs.workspace = true -bytes.workspace = true -trust-dns-resolver.workspace = true -ring.workspace = true -once_cell.workspace = true -uni-stream.workspace = true -kanal.workspace = true - -[dev-dependencies] -dotenvy = "0.15.7" - -[features] -udp-timeout = ["uni-stream/udp-timeout"] - [workspace] -members = ["ui/native/pb_mapper_ffi"] +# Every library and the CLI live under `crates/`. The FFI cdylib stays next to +# the Flutter code that loads it. +members = ["crates/*", "ui/native/pb_mapper_ffi"] exclude = ["deps/uni-stream", "deps/kanal"] +# Spelled out: a virtual manifest does not infer the resolver from the edition, +# and without this it silently falls back to resolver 1. +resolver = "3" [workspace.package] -version = "0.3.0" +# `version` must stay the first key here: `release.yml` and `docker-publish.yml` +# both parse it positionally with awk to check the tag against it. +version = "0.4.0" authors = ["L_B__"] -edition = "2021" +edition = "2024" + +[workspace.lints.clippy] +unwrap_used = "deny" +expect_used = "deny" [workspace.dependencies] +pb-mapper-auth = { path = "crates/pb-mapper-auth" } +pb-mapper-client = { path = "crates/pb-mapper-client" } +pb-mapper-core = { path = "crates/pb-mapper-core" } +pb-mapper-protocol = { path = "crates/pb-mapper-protocol" } +pb-mapper-server = { path = "crates/pb-mapper-server" } + +base64 = "0.23.1" +better_mimalloc_rs = { version = "0.1.2", features = ["config"] } +bytes = "1.11" +clap = { version = "4.5", features = ["derive"] } +dirs = "6.0.0" +dotenvy = "0.15.7" +hashbrown = { version = "0.17.1" } +hickory-resolver = { version = "0.26.1" } +kanal = { git = "https://github.com/acking-you/kanal.git", branch = "dev/pb-mapper" } +parking_lot = "0.12" rand = "0.10" -socket2 = "0.6" -tokio = { version = "1", features = ["full"] } -tokio-util = "0.7" -snafu = "0.8.7" +ring = "0.17.14" serde = { version = "1.0", features = ["derive"] } serde_json = { version = "1.0", default-features = false, features = ["alloc"] } +snafu = "0.9.2" +subtle = "2.6.1" +tokio = { version = "1", features = ["full"] } +tokio-util = "0.7" tracing = "0.1.40" tracing-subscriber = { version = "0.3.18", features = [ "env-filter", "fmt", "json", ], default-features = true } -hashbrown = { version = "0.16" } -clap = { version = "4.5", features = ["derive"] } -futures = "0.3.31" -better_mimalloc_rs = { version = "0.1.2", features = ["config"] } -bytes = "1.11" -trust-dns-resolver = { version = "0.23.2" } -ring = "0.17.14" -once_cell = "1.20.2" uni-stream = { git = "https://github.com/acking-you/uni-stream.git", branch = "master" } -kanal = { git = "https://github.com/acking-you/kanal.git", branch = "dev/pb-mapper" } - -[profile] - -[profile.wasm-dev] -inherits = "dev" -opt-level = 1 - -[profile.server-dev] -inherits = "dev" - -[profile.android-dev] -inherits = "dev" diff --git a/DOCKER_README.md b/DOCKER_README.md index 7cc7bdb..da26103 100644 --- a/DOCKER_README.md +++ b/DOCKER_README.md @@ -13,8 +13,8 @@ docker run -d \ --name pb-mapper \ -p 7666:7666 \ -e PB_MAPPER_PORT=7666 \ - -e USE_MACHINE_MSG_HEADER_KEY=true \ -e RUST_LOG=error \ + -v pb-mapper-auth:/var/lib/pb-mapper/auth \ ackingliu/pb-mapper:latest-x86_64_musl ``` @@ -28,11 +28,16 @@ services: environment: - PB_MAPPER_PORT=7666 - USE_IPV6=false - - USE_MACHINE_MSG_HEADER_KEY=true + - USE_MACHINE_MSG_HEADER_KEY=false - RUST_LOG=error + volumes: + - pb-mapper-auth:/var/lib/pb-mapper/auth ports: - "7666:7666" restart: unless-stopped + +volumes: + pb-mapper-auth: ``` Save as `docker-compose.yml` and run: @@ -46,10 +51,12 @@ docker-compose up -d |----------|---------|-------------| | `PB_MAPPER_PORT` | `7666` | **Required** - Port for the pb-mapper server to listen on | | `USE_IPV6` | `false` | Enable IPv6 support (`true`/`false`) | -| `USE_MACHINE_MSG_HEADER_KEY` | `true` | Derive `MSG_HEADER_KEY` from hostname + MAC and persist to `/var/lib/pb-mapper-server/msg_header_key` | +| `MSG_HEADER_KEY` | unset | Optional 32-character administrator key used only to initialize a new persistent auth volume | +| `USE_MACHINE_MSG_HEADER_KEY` | `false` | Legacy compatibility: derive the administrator key from hostname + MAC | +| `PB_MAPPER_AUTH_STATE_DIR` | `/var/lib/pb-mapper/auth` | Persistent encrypted authentication state | | `RUST_LOG` | `error` | Logging level (`error`, `warn`, `info`, `debug`, `trace`) | -⚠️ **Important**: `PB_MAPPER_PORT` must be set or the container will exit with an error. +⚠️ **Important**: `PB_MAPPER_PORT` must be set and `/var/lib/pb-mapper/auth` must be persistent. The first start creates a random administrator key at `admin.key`; losing the volume changes the root credential and loses temporary-key state. ## 📋 Ubuntu Deployment Guide @@ -85,11 +92,16 @@ services: environment: PB_MAPPER_PORT: 7666 USE_IPV6: false - USE_MACHINE_MSG_HEADER_KEY: true + USE_MACHINE_MSG_HEADER_KEY: false RUST_LOG: error + volumes: + - pb-mapper-auth:/var/lib/pb-mapper/auth ports: - "7666:7666" restart: unless-stopped + +volumes: + pb-mapper-auth: EOF ``` @@ -112,6 +124,9 @@ docker-compose ps # View logs docker-compose logs -f pb-mapper + +# Read the administrator key on the Docker host +docker exec pb-mapper cat /var/lib/pb-mapper/auth/admin.key ``` ### Step 5: Verify Installation @@ -154,18 +169,18 @@ For other architectures, you can build the image yourself using the provided Doc |-----|-------------| | `latest-x86_64_musl` | Latest stable x86_64 build (recommended) | | `latest-aarch64_musl` | Latest stable ARM64 build | -| `v0.3.0-x86_64_musl` | Tagged-release x86_64 build | -| `v0.3.0-aarch64_musl` | Tagged-release ARM64 build | -| `0.3.0-x86_64_musl` | Semver x86_64 alias | -| `0.3.0-aarch64_musl` | Semver ARM64 alias | +| `v0.4.0-x86_64_musl` | Tagged-release x86_64 build | +| `v0.4.0-aarch64_musl` | Tagged-release ARM64 build | +| `0.4.0-x86_64_musl` | Semver x86_64 alias | +| `0.4.0-aarch64_musl` | Semver ARM64 alias | **Recommendation**: Use `latest-x86_64_musl` for x86_64 systems or `latest-aarch64_musl` for ARM64 systems for best compatibility. ## 🛡️ Security Considerations - **Firewall**: Only expose port 7666 to trusted networks -- **Encryption**: Use the encryption features in client/server tools -- **Access Control**: Implement service key management strategy +- **Authentication**: Keep `admin.key` on the relay and distribute expiring `pbmt1_` temporary credentials to workloads +- **Forwarded payload encryption**: Use `register --codec` when the inner application protocol is plaintext - **Updates**: Regularly update to the latest version for security patches ## 📊 Monitoring and Logs diff --git a/README.md b/README.md index 3745f7c..e94280c 100644 --- a/README.md +++ b/README.md @@ -3,7 +3,7 @@ pb-mapper

- Rust 2021 + Rust 2024 Tokio Flutter License: MIT @@ -27,6 +27,8 @@ ## Highlights - **One binary, one public port** — `pb-mapper` provides every runtime role, while a service-key registry replaces per-service port planning. +- **Scoped temporary credentials** — the administrator key can issue renewable, expiring `pbmt1_` credentials; each credential gets an isolated service namespace and can only inspect, register, and connect inside it. +- **Authenticated protocol v2** — directional AES-256-GCM control frames authenticate in the first request without adding a handshake round trip. New clients use v2; the server can temporarily allow legacy clients during migration. - **Optional encryption** — AES-256-GCM (via `ring`) on forwarded traffic, enabled with `--codec` at registration. - **Proven in production** — on real workloads (e.g. a Palworld UDP server), latency matches frp with a directly exposed port. @@ -41,18 +43,21 @@ With an AI coding agent (Claude Code, Cursor, Kiro), the built-in skills handle ### Alternative — one-liner install script -If the remote host can reach GitHub directly, this installs the unified `pb-mapper` binary and runs its `server` command as a systemd service on Linux (x86_64, musl) — port `7666`, `--use-machine-msg-header-key` on, key stored at `/var/lib/pb-mapper-server/msg_header_key`. +If the remote host can reach GitHub directly, this installs the unified `pb-mapper` binary and runs its `server` command as a systemd service on Linux (x86_64, musl). The relay listens on port `7666` and creates a random administrator key at `/var/lib/pb-mapper/auth/admin.key` on first start. ```bash curl -fsSL https://raw.githubusercontent.com/acking-you/pb-mapper/master/scripts/install-server-github.sh | bash ``` -After install, load the same key before running `pb-mapper register` or `pb-mapper connect`: +Use the administrator key only for management and issue a temporary credential for a workload: ```bash -export MSG_HEADER_KEY="$(cat /var/lib/pb-mapper-server/msg_header_key)" +export MSG_HEADER_KEY="$(sudo cat /var/lib/pb-mapper/auth/admin.key)" +pb-mapper admin --server :7666 key issue --ttl 24h --label home-web ``` +Copy the printed `pbmt1_...` credential to the register and connect machines as their `MSG_HEADER_KEY`. They may use the same service name without colliding with another temporary credential's namespace. + ## Architecture ![pb-mapper architecture](docs/assets/architecture-flow.svg) @@ -80,10 +85,13 @@ Your web server runs on `localhost:8080` at home. # 1. on the public server — start the central router pb-mapper server --port 7666 -# 2. at home — register the web server under key 'web' +# 2. issue a temporary credential, then export it on both endpoint machines +export MSG_HEADER_KEY='' + +# 3. at home — register the web server under key 'web' pb-mapper register tcp --server :7666 --key web --addr 127.0.0.1:8080 -# 3. at the coffee shop — subscribe and expose it locally +# 4. at the coffee shop — subscribe and expose it locally pb-mapper connect tcp --server :7666 --key web --addr 127.0.0.1:3000 ``` @@ -97,25 +105,27 @@ Open `http://localhost:3000` in the coffee-shop browser — traffic flows throug | `pb-mapper register tcp\|udp` | Registers a local TCP/UDP service with the server | | `pb-mapper connect tcp\|udp` | Subscribes to a registered service and exposes a local port | | `pb-mapper status keys\|remote-id` | Queries the central router | +| `pb-mapper admin ...` | Issues/renews/revokes credentials and inspects auth, services, and connections | | **Flutter UI** (`ui/`) | GUI for server, register, connect, and status workflows | ## Developer view -- **Rust core** — the unified entry point is `src/bin/pb-mapper.rs`; shared protocol and networking live in `src/common` and `src/utils`; server / register / connect internals live in `src/pb_server`, `src/local/server`, and `src/local/client`. +- **Rust core** — a workspace under `crates/`, layered bottom-up: `pb-mapper-core` (credentials, checksum, config, addressing), `pb-mapper-auth` (credential lifecycle and persistence), `pb-mapper-protocol` (framing and secure sessions), then `pb-mapper-server` and `pb-mapper-client` as peers, with the `pb-mapper` binary in `pb-mapper-cli`. - **Flutter UI** — views in `ui/lib/src/views`, FFI layers in `ui/lib/src/ffi`, Rust bridge in `ui/native/pb_mapper_ffi`. FFI calls run on a background isolate, and Rust returns JSON (`{success, message, data}`) to keep the C ABI stable. ## Documentation - User guide (build / run / use): [`docs/user-guide.md`](docs/user-guide.md) +- Authentication and protocol v2: [`docs/authentication-v2.md`](docs/authentication-v2.md) - Docker server guide: [`DOCKER_README.md`](DOCKER_README.md) - 中文文档: [`README.zh-CN.md`](README.zh-CN.md), [`docs/user-guide.zh-CN.md`](docs/user-guide.zh-CN.md) ## Repository layout -- `src/` — Rust backend +- `crates/` — the Rust workspace (six crates; the root manifest is virtual) - `ui/` — Flutter UI + native bridge - `docs/` — documentation and assets -- `docker/`, `services/`, `scripts/`, `tests/` — deployment and tooling +- `docker/`, `services/`, `scripts/` — deployment and tooling - `skills/` — AI coding agent deployment skills (server and connect tunnel) ## License diff --git a/README.zh-CN.md b/README.zh-CN.md index 58fe928..562351e 100644 --- a/README.zh-CN.md +++ b/README.zh-CN.md @@ -3,7 +3,7 @@ pb-mapper

- Rust 2021 + Rust 2024 Tokio Flutter License: MIT @@ -27,6 +27,8 @@ ## 亮点 - **单二进制、单公网端口**:统一的 `pb-mapper` 命令覆盖所有运行角色,服务 key 注册表取代逐个服务规划端口。 +- **临时凭据与命名空间隔离**:管理员密钥可签发可续期、自动过期的 `pbmt1_` 凭据;每把临时凭据只能查看、注册和连接自己的命名空间。 +- **V2 首帧鉴权**:控制帧使用按方向派生的 AES-256-GCM 密钥,在第一个请求内完成鉴权,不增加额外握手往返;新客户端固定使用 V2,服务端可在迁移期兼容旧协议。 - **可选加密**:转发流量可启用 AES-256-GCM(基于 `ring`),注册服务时用 `--codec` 开启。 - **生产可用**:真实负载下(例如 Palworld UDP 服务器),延迟与 frp 直暴端口相当。 @@ -41,18 +43,21 @@ ### 备选方式:一键安装脚本 -远程主机能直连 GitHub 时,一条命令即可在 Linux(x86_64,musl)上安装统一的 `pb-mapper` 二进制,并以 `server` 子命令启动 systemd 服务:端口 `7666`,启用 `--use-machine-msg-header-key`,key 落盘在 `/var/lib/pb-mapper-server/msg_header_key`。 +远程主机能直连 GitHub 时,一条命令即可在 Linux(x86_64,musl)上安装统一的 `pb-mapper` 二进制,并以 `server` 子命令启动 systemd 服务。中继监听 `7666`,首次启动时会在 `/var/lib/pb-mapper/auth/admin.key` 创建随机管理员密钥。 ```bash curl -fsSL https://raw.githubusercontent.com/acking-you/pb-mapper/master/scripts/install-server-github.sh | bash ``` -安装完成后,在运行 `pb-mapper register` 或 `pb-mapper connect` 前加载同一把 key: +管理员密钥只用于管理;先为一项业务签发临时凭据: ```bash -export MSG_HEADER_KEY="$(cat /var/lib/pb-mapper-server/msg_header_key)" +export MSG_HEADER_KEY="$(sudo cat /var/lib/pb-mapper/auth/admin.key)" +pb-mapper admin --server :7666 key issue --ttl 24h --label home-web ``` +把输出的 `pbmt1_...` 凭据作为 register 与 connect 机器上的 `MSG_HEADER_KEY`。不同临时凭据即使使用相同的 service name,也不会相互冲突。 + ## 架构 ![pb-mapper architecture](docs/assets/architecture-flow.svg) @@ -80,10 +85,13 @@ register 与 connect 工作流也可以通过 Flutter UI 操作。 # 1. 公网服务器:启动中心路由 pb-mapper server --port 7666 -# 2. 家中机器:以 key 'web' 注册服务 +# 2. 签发临时凭据,并在两端机器导入 +export MSG_HEADER_KEY='' + +# 3. 家中机器:以 key 'web' 注册服务 pb-mapper register tcp --server :7666 --key web --addr 127.0.0.1:8080 -# 3. 咖啡店机器:订阅并在本地暴露 +# 4. 咖啡店机器:订阅并在本地暴露 pb-mapper connect tcp --server :7666 --key web --addr 127.0.0.1:3000 ``` @@ -97,25 +105,27 @@ pb-mapper connect tcp --server :7666 --key web --addr 127.0.0.1:3000 | `pb-mapper register tcp\|udp` | 将本地 TCP/UDP 服务注册到服务器 | | `pb-mapper connect tcp\|udp` | 订阅已注册的服务并在本地暴露端口 | | `pb-mapper status keys\|remote-id` | 查询中心路由状态 | +| `pb-mapper admin ...` | 签发/续期/吊销临时凭据并查看认证、服务与连接状态 | | **Flutter UI**(`ui/`) | server、register、connect、status 的图形化界面 | ## 开发者视角 -- **Rust 核心**:统一入口为 `src/bin/pb-mapper.rs`;协议与网络通用逻辑在 `src/common`、`src/utils`;server/register/connect 实现在 `src/pb_server`、`src/local/server`、`src/local/client`。 +- **Rust 核心**:`crates/` 下的 workspace,自底向上分层:`pb-mapper-core`(凭据、校验和、配置、地址解析)→ `pb-mapper-auth`(凭据生命周期与持久化)→ `pb-mapper-protocol`(帧格式与安全会话)→ `pb-mapper-server` 与 `pb-mapper-client`(二者平级,互不引用)→ `pb-mapper-cli`(`pb-mapper` 二进制所在)。 - **Flutter UI**:界面在 `ui/lib/src/views`,FFI 各层在 `ui/lib/src/ffi`,Rust 桥接在 `ui/native/pb_mapper_ffi`。FFI 调用跑在后台 isolate,Rust 统一返回 JSON(`{success, message, data}`)以保持 C ABI 稳定。 ## 文档 - 使用手册(编译/运行/使用):[`docs/user-guide.zh-CN.md`](docs/user-guide.zh-CN.md) +- 认证与 V2 协议:[`docs/authentication-v2.zh-CN.md`](docs/authentication-v2.zh-CN.md) - Docker 服务器指南:[`DOCKER_README.md`](DOCKER_README.md) - English docs: [`README.md`](README.md)、[`docs/user-guide.md`](docs/user-guide.md) ## 仓库结构 -- `src/` — Rust 后端 +- `crates/` — Rust workspace(六个 crate,根清单为虚拟清单) - `ui/` — Flutter UI + 原生桥接 - `docs/` — 文档与素材 -- `docker/`、`services/`、`scripts/`、`tests/` — 部署与工具 +- `docker/`、`services/`、`scripts/` — 部署与工具 - `skills/` — AI 编程助手部署 skill(服务端、客户端隧道) ## 许可证 diff --git a/clippy.toml b/clippy.toml new file mode 100644 index 0000000..96ba758 --- /dev/null +++ b/clippy.toml @@ -0,0 +1,5 @@ +# `unwrap_used` and `expect_used` are denied workspace-wide (see the root +# `[workspace.lints.clippy]`). A panic in a test is a failing test, which is the +# point, so exempt test code rather than annotating every assertion. +allow-unwrap-in-tests = true +allow-expect-in-tests = true diff --git a/crates/pb-mapper-auth/Cargo.toml b/crates/pb-mapper-auth/Cargo.toml new file mode 100644 index 0000000..a76d7a3 --- /dev/null +++ b/crates/pb-mapper-auth/Cargo.toml @@ -0,0 +1,21 @@ +[package] +name = "pb-mapper-auth" +version.workspace = true +edition.workspace = true +authors.workspace = true + +[dependencies] +pb-mapper-core.workspace = true + +parking_lot.workspace = true +rand.workspace = true +ring.workspace = true +serde.workspace = true +serde_json.workspace = true +subtle.workspace = true +tokio.workspace = true +tokio-util.workspace = true +tracing.workspace = true + +[lints] +workspace = true diff --git a/crates/pb-mapper-auth/src/actor/epoch.rs b/crates/pb-mapper-auth/src/actor/epoch.rs new file mode 100644 index 0000000..c69ba8a --- /dev/null +++ b/crates/pb-mapper-auth/src/actor/epoch.rs @@ -0,0 +1,142 @@ +//! Root rotation, auth-state reset, and live temporary-key wipe. +use super::super::*; +use super::{audit, ensure_store_available}; + +fn remember_previous_root(inner: &AuthStateInner) { + *inner.previous_root.write() = Some(PreviousRoot { + admin_key: inner.admin_key(), + instance_id: inner.instance_id(), + }); +} + +pub(super) fn actor_reset( + inner: &Arc, + config: &AuthConfig, + leases: &mut Leases, + admin_replays: &VecDeque, + action: &str, +) -> Result<(), AuthFailure> { + let new_instance_id = random_instance_id(); + inner.root_epoch.fetch_add(1, Ordering::AcqRel); + let reset_audit = audit(action, None, None); + let mut snapshot = empty_snapshot(inner, new_instance_id, admin_replays); + push_persisted_audit(&mut snapshot.audit_records, reset_audit.clone()); + let admin_key = inner.admin_key(); + let next_instance_path = config.state_dir.join("server-instance-id.next"); + if let Err(error) = atomic_write(&next_instance_path, &new_instance_id, 0o600) + .and_then(|()| write_snapshot_and_truncate_wal(config, &admin_key, &snapshot)) + .and_then(|()| { + atomic_write( + &config.state_dir.join("server-instance-id"), + &new_instance_id, + 0o600, + ) + }) + { + if !reset_already_installed(&config.state_dir, &admin_key, &new_instance_id) { + inner.safe_mode.store(true, Ordering::Release); + cancel_all_temporary_leases(inner); + return Err(error); + } + tracing::warn!( + event = "auth_state_reset_finalized_after_sync_error", + error = %error, + "server-instance-id replacement reported an error, but the live id and snapshot already match the new instance; finishing in-memory reset" + ); + } + let _ = std::fs::remove_file(&next_instance_path); + push_audit_record(inner, reset_audit); + remember_previous_root(inner); + leases.wipe(unix_seconds()); + *inner.instance_id.write() = new_instance_id; + inner.safe_mode.store(false, Ordering::Release); + Ok(()) +} + +pub(super) fn actor_rotate_root( + inner: &Arc, + config: &AuthConfig, + leases: &mut Leases, + admin_lease: &mut Arc, + new_key: AesKeyType, +) -> Result<(), AuthFailure> { + if new_key == inner.admin_key() { + return Err(AuthFailure::new( + "administrator_key_unchanged", + "new administrator key must differ from the current key", + false, + )); + } + if !is_env_safe_admin_key(&new_key) { + return Err(AuthFailure::new( + "administrator_key_invalid", + env_safe_admin_key_error(), + false, + )); + } + // Unreachable: `is_env_safe_admin_key` above accepts only printable ASCII. + // Reported rather than panicked, since this already returns `Result`. + let new_key_string = String::from_utf8(new_key.to_vec()).map_err(|_| { + AuthFailure::new( + "administrator_key_invalid", + env_safe_admin_key_error(), + false, + ) + })?; + + inner.root_epoch.fetch_add(1, Ordering::AcqRel); + let rotate_audit = audit("administrator_key_rotate", None, None); + let mut snapshot = empty_snapshot(inner, inner.instance_id(), &VecDeque::new()); + push_persisted_audit(&mut snapshot.audit_records, rotate_audit.clone()); + let next_key_path = config.state_dir.join("admin.key.next"); + if let Err(error) = write_admin_key_file(&next_key_path, &new_key_string, true) + .and_then(|()| write_snapshot_and_truncate_wal(config, &new_key, &snapshot)) + .and_then(|()| write_admin_key(&config.state_dir, &new_key_string)) + { + if !rotation_already_installed(&config.state_dir, &new_key_string) { + inner.safe_mode.store(true, Ordering::Release); + cancel_all_temporary_leases(inner); + return Err(error); + } + tracing::warn!( + event = "administrator_key_rotate_finalized_after_sync_error", + error = %error, + "admin.key replacement reported an error, but the new snapshot already decrypts with the new key; finishing in-memory rotation" + ); + } + let _ = std::fs::remove_file(&next_key_path); + push_audit_record(inner, rotate_audit); + remember_previous_root(inner); + leases.wipe(unix_seconds()); + let old_admin_lease = admin_lease.clone(); + let new_admin_lease = Arc::new(AuthLease::new(ADMIN_KEY_ID, u64::MAX)); + *inner.admin.write() = AdminState { + key: new_key, + lease: Arc::downgrade(&new_admin_lease), + }; + if inner.sync_process_credential { + set_process_msg_header_key(Some(&new_key_string)).map_err(AuthFailure::internal)?; + } + inner.safe_mode.store(false, Ordering::Release); + old_admin_lease.cancel_rotated(); + *admin_lease = new_admin_lease; + Ok(()) +} + +pub(super) fn actor_set_legacy_protocol( + inner: &Arc, + config: &AuthConfig, + policy: LegacyProtocolPolicy, +) -> Result<(), AuthFailure> { + ensure_store_available(inner)?; + append_mutation( + config, + inner, + StateMutation::LegacyProtocol(policy), + audit("legacy_protocol_update", None, Some(format!("{policy:?}"))), + )?; + inner + .legacy_protocol_allowed + .store(policy.is_allowed(), Ordering::Release); + Ok(()) +} diff --git a/crates/pb-mapper-auth/src/actor/lifecycle.rs b/crates/pb-mapper-auth/src/actor/lifecycle.rs new file mode 100644 index 0000000..ee3fa89 --- /dev/null +++ b/crates/pb-mapper-auth/src/actor/lifecycle.rs @@ -0,0 +1,432 @@ +//! Issue, inspect, renew, revoke, and collect temporary keys. +use super::super::*; +use super::{ + audit, ensure_store_available, key_not_active, key_not_found, key_not_renewable, + slot_state_name, validate_slot_identity, +}; + +fn validate_ttl(config: &AuthConfig, ttl: Duration) -> Result { + if ttl < MIN_TEMP_KEY_TTL { + return Err(AuthFailure::new( + "temporary_key_ttl_too_short", + format!( + "temporary key TTL must be at least {} seconds", + MIN_TEMP_KEY_TTL.as_secs() + ), + false, + )); + } + if ttl > config.max_temporary_key_ttl { + return Err(AuthFailure::new( + "temporary_key_ttl_too_long", + format!( + "temporary key TTL exceeds the configured maximum of {} seconds", + config.max_temporary_key_ttl.as_secs() + ), + false, + )); + } + Ok(unix_seconds().saturating_add(ttl.as_secs())) +} + +fn validate_label(label: Option) -> Result, AuthFailure> { + let label = label + .map(|label| label.trim().to_string()) + .filter(|label| !label.is_empty()); + if label.as_ref().is_some_and(|label| label.len() > 64) { + return Err(AuthFailure::new( + "temporary_key_label_too_long", + "temporary key label must not exceed 64 UTF-8 bytes", + false, + )); + } + Ok(label) +} + +/// One key's lifecycle, read from wherever that key lives. +struct KeyState { + state: SlotState, + expires_at: u64, + issued_at: u64, + label: Option, +} + +/// Reads a key's lifecycle from the slot table, falling back to the entries +/// retained for slots the configured capacity no longer covers. Every operation +/// that accepts any live key id needs both paths; see +/// `AuthStateInner::high_slot_generations`. +fn key_state(inner: &AuthStateInner, key_id: KeyId) -> Result { + let slots = inner.slots(); + if let Some(slot) = slots.get(key_id.slot().as_index()) { + validate_slot_identity(slot, key_id)?; + let metadata = inner + .cold() + .get(&key_id) + .cloned() + .ok_or_else(|| key_not_found(key_id))?; + return Ok(KeyState { + state: slot.state, + expires_at: slot.expires_at, + issued_at: metadata.issued_at, + label: metadata.label.clone(), + }); + } + drop(slots); + let high = inner.high(); + let entry = high_slot_entry(&high, key_id)?; + Ok(KeyState { + state: entry.state, + expires_at: entry.expires_at, + issued_at: entry.issued_at, + label: entry.label.clone(), + }) +} + +pub(super) fn actor_issue( + inner: &Arc, + config: &AuthConfig, + leases: &mut Leases, + ttl: Duration, + label: Option, +) -> Result { + ensure_store_available(inner)?; + let expires_at = validate_ttl(config, ttl)?; + let label = validate_label(label)?; + let issued_at = unix_seconds(); + let (index, generation, key_id, entry) = { + let slots = inner.slots(); + // A row whose generation cannot advance is skipped rather than reused: it + // has no unused identity left to hand out. + let Some((slot_index, generation)) = slots.iter().enumerate().find_map(|(index, slot)| { + (slot.state == SlotState::Free) + .then(|| slot.generation.next()) + .flatten() + .map(|generation| (SlotIndex::from_index(index), generation)) + }) else { + return Err(AuthFailure::new( + "temporary_key_capacity_exhausted", + "temporary key slot table is full", + true, + )); + }; + let key_id = KeyId::new(generation, slot_index); + ( + slot_index.as_index(), + generation, + key_id, + PersistedEntry { + key_id, + state: SlotState::Active, + issued_at, + expires_at, + label: label.clone(), + tombstoned_at: None, + }, + ) + }; + // Persist before taking the slot write lock. A fail-closed WAL error + // cancels leases via slots.read() and must not nest under slots.write(). + append_mutation( + config, + inner, + StateMutation::Issue(entry), + audit("temporary_key_issue", Some(key_id), label.clone()), + )?; + let mut slots = inner.slots_mut(); + let slot = slots + .get_mut(index) + .ok_or_else(|| AuthFailure::internal("issued slot disappeared"))?; + let lease = Arc::new(AuthLease::new(key_id, expires_at)); + slot.generation = generation; + slot.state = SlotState::Active; + slot.expires_at = expires_at; + slot.lease = Arc::downgrade(&lease); + drop(slots); + leases.issue(&lease, issued_at, label); + metadata_with_credential(inner, key_id, true) +} + +pub(super) fn actor_list( + inner: &Arc, + page: u32, + page_size: u16, +) -> Result { + let page_size = page_size.clamp(1, 1000) as usize; + let start = (page as usize).saturating_mul(page_size); + let slots = inner.slots(); + let cold = inner.cold(); + let mut all = slots + .iter() + .enumerate() + .filter_map(|(index, slot)| { + if slot.state == SlotState::Free { + return None; + } + let key_id = KeyId::new(slot.generation, SlotIndex::from_index(index)); + let cold = cold.get(&key_id)?; + Some(TemporaryKeyMetadata { + key_id, + state: slot_state_name(slot.state).to_string(), + issued_at: cold.issued_at, + expires_at: slot.expires_at, + label: cold.label.clone(), + }) + }) + .collect::>(); + all.extend( + inner + .high() + .iter() + .filter(|entry| entry.state != SlotState::Free) + .map(high_slot_metadata), + ); + all.sort_by_key(|item| std::cmp::Reverse(item.issued_at)); + let items = all.iter().skip(start).take(page_size).cloned().collect(); + let next_page = (start.saturating_add(page_size) < all.len()).then_some(page.saturating_add(1)); + Ok(KeyPage { + schema_version: 1, + items, + next_page, + }) +} + +pub(super) fn actor_show( + inner: &Arc, + config: &AuthConfig, + key_id: KeyId, + reveal: bool, +) -> Result { + let result = metadata_with_credential(inner, key_id, reveal)?; + append_audit( + config, + inner, + audit( + if reveal { + "temporary_key_reveal" + } else { + "temporary_key_show" + }, + Some(key_id), + result.metadata.label.clone(), + ), + )?; + Ok(result) +} + +pub(super) fn actor_renew( + inner: &Arc, + config: &AuthConfig, + leases: &mut Leases, + key_id: KeyId, + ttl: Duration, +) -> Result { + ensure_store_available(inner)?; + let expires_at = validate_ttl(config, ttl)?; + let index = key_id.slot().as_index(); + let current = key_state(inner, key_id)?; + if current.state != SlotState::Active || current.expires_at <= unix_seconds() { + return Err(key_not_renewable()); + } + let label = current.label; + append_mutation( + config, + inner, + StateMutation::Renew { key_id, expires_at }, + audit("temporary_key_renew", Some(key_id), label.clone()), + )?; + let mut slots = inner.slots_mut(); + if let Some(slot) = slots.get_mut(index) { + validate_slot_identity(slot, key_id)?; + if slot.state != SlotState::Active { + return Err(AuthFailure::new( + "temporary_key_inactive", + "temporary key lease is no longer active", + true, + )); + } + slot.expires_at = expires_at; + match slot.lease.upgrade() { + Some(lease) if !lease.cancellation_token().is_cancelled() => { + lease.expires_at.store(expires_at, Ordering::Release); + drop(slots); + leases.renew(key_id, expires_at); + } + // A cancelled lease cannot be revived, so the renewal installs a + // replacement; that drops the handle on the lease it succeeds. + _ => { + let lease = Arc::new(AuthLease::new(key_id, expires_at)); + slot.lease = Arc::downgrade(&lease); + drop(slots); + leases.adopt(&lease); + } + } + return metadata_with_credential(inner, key_id, true); + } + drop(slots); + { + let mut high = inner.high_mut(); + let entry = high_slot_entry_mut(&mut high, key_id)?; + if entry.state != SlotState::Active { + return Err(AuthFailure::new( + "temporary_key_inactive", + "temporary key lease is no longer active", + true, + )); + } + entry.expires_at = expires_at; + entry.tombstoned_at = None; + } + metadata_with_credential(inner, key_id, true) +} + +pub(super) fn actor_revoke( + inner: &Arc, + config: &AuthConfig, + leases: &mut Leases, + key_id: KeyId, +) -> Result { + ensure_store_available(inner)?; + let now = unix_seconds(); + let index = key_id.slot().as_index(); + let current = key_state(inner, key_id)?; + if current.state != SlotState::Active { + return Err(key_not_active()); + } + let KeyState { + label, + issued_at, + expires_at, + .. + } = current; + append_mutation( + config, + inner, + StateMutation::Revoke { key_id, at: now }, + audit("temporary_key_revoke", Some(key_id), label.clone()), + )?; + let mut slots = inner.slots_mut(); + if let Some(slot) = slots.get_mut(index) { + validate_slot_identity(slot, key_id)?; + slot.state = SlotState::Revoked; + if let Some(lease) = slot.lease.upgrade() { + lease.cancel_revoked(); + } + let state = slot_state_name(slot.state).to_string(); + let expires_at = slot.expires_at; + drop(slots); + // Retire only: the row stays until its retention elapses, because the + // slot table holds a `Weak` and a later request has to be able to read + // the revoked reason rather than find a recycled row. + leases.retire_now(key_id); + let metadata = inner + .cold() + .get(&key_id) + .cloned() + .ok_or_else(|| key_not_found(key_id))?; + return Ok(TemporaryKeyMetadata { + key_id, + state, + issued_at: metadata.issued_at, + expires_at, + label: metadata.label.clone(), + }); + } + drop(slots); + let mut high = inner.high_mut(); + let entry = high_slot_entry_mut(&mut high, key_id)?; + if entry.state != SlotState::Active { + return Err(key_not_active()); + } + entry.state = SlotState::Revoked; + entry.tombstoned_at = Some(now); + let state = slot_state_name(entry.state).to_string(); + drop(high); + leases.retire_now(key_id); + Ok(TemporaryKeyMetadata { + key_id, + state, + issued_at, + expires_at, + label, + }) +} + +pub(super) fn actor_gc( + inner: &Arc, + config: &AuthConfig, + leases: &mut Leases, + admin_replays: &VecDeque, +) -> Result { + ensure_store_available(inner)?; + let removed = leases.collect_garbage(unix_seconds()); + let gc_audit = audit("temporary_key_gc", None, Some(format!("removed={removed}"))); + let mut snapshot = build_snapshot(inner, admin_replays); + push_persisted_audit(&mut snapshot.audit_records, gc_audit.clone()); + let admin_key = inner.admin_key(); + if let Err(error) = write_snapshot_and_truncate_wal(config, &admin_key, &snapshot) { + inner.safe_mode.store(true, Ordering::Release); + cancel_all_temporary_leases(inner); + return Err(error); + } + push_audit_record(inner, gc_audit); + Ok(removed) +} + +fn high_slot_entry(high: &[PersistedEntry], key_id: KeyId) -> Result<&PersistedEntry, AuthFailure> { + high.iter() + .find(|entry| entry.key_id == key_id) + .ok_or_else(|| key_not_found(key_id)) +} + +fn high_slot_entry_mut( + high: &mut [PersistedEntry], + key_id: KeyId, +) -> Result<&mut PersistedEntry, AuthFailure> { + high.iter_mut() + .find(|entry| entry.key_id == key_id) + .ok_or_else(|| key_not_found(key_id)) +} + +fn high_slot_metadata(entry: &PersistedEntry) -> TemporaryKeyMetadata { + TemporaryKeyMetadata { + key_id: entry.key_id, + state: slot_state_name(entry.state).to_string(), + issued_at: entry.issued_at, + expires_at: entry.expires_at, + label: entry.label.clone(), + } +} + +fn metadata_with_credential( + inner: &Arc, + key_id: KeyId, + reveal: bool, +) -> Result { + let slots = inner.slots(); + let credential = if reveal { + let key = derive_temporary_key(&inner.admin_key(), &inner.instance_id(), key_id)?; + encode_temporary_credential(key_id.as_u64(), &key) + } else { + String::new() + }; + if let Some(slot) = slots.get(key_id.slot().as_index()) { + validate_slot_identity(slot, key_id)?; + let cold = inner.cold(); + let cold = cold.get(&key_id).ok_or_else(|| key_not_found(key_id))?; + return Ok(IssuedTemporaryKey { + metadata: TemporaryKeyMetadata { + key_id, + state: slot_state_name(slot.state).to_string(), + issued_at: cold.issued_at, + expires_at: slot.expires_at, + label: cold.label.clone(), + }, + credential, + }); + } + let high = inner.high(); + Ok(IssuedTemporaryKey { + metadata: high_slot_metadata(high_slot_entry(&high, key_id)?), + credential, + }) +} diff --git a/crates/pb-mapper-auth/src/actor/mod.rs b/crates/pb-mapper-auth/src/actor/mod.rs new file mode 100644 index 0000000..2945614 --- /dev/null +++ b/crates/pb-mapper-auth/src/actor/mod.rs @@ -0,0 +1,408 @@ +//! Serialized owner of mutable authentication lifecycle state. +//! +//! ```text +//! authenticated admin command +//! | +//! v +//! validate current admin lease +//! | +//! v +//! append encrypted WAL -> mutate slots / leases / timing wheel +//! | +//! +-> periodic snapshot + bounded replay/audit retention +//! ``` +//! +//! Keeping authorization revalidation and mutations in one actor prevents a request +//! authenticated before root rotation from executing against the new administrator +//! state. The actor is also the sole strong owner of temporary-key leases. + +use super::*; + +mod epoch; +mod lifecycle; +use epoch::*; +use lifecycle::*; + +pub(super) struct AuthActorState { + leases: Leases, + admin_replays: HashSet<[u8; 32]>, + admin_replay_order: VecDeque, +} + +impl AuthActorState { + pub(super) fn new( + leases: Leases, + admin_replays: HashSet<[u8; 32]>, + admin_replay_order: VecDeque, + ) -> Self { + Self { + leases, + admin_replays, + admin_replay_order, + } + } +} + +pub(super) async fn run_auth_actor( + inner: Arc, + mut admin_lease: Arc, + mut command_rx: mpsc::Receiver, + config: AuthConfig, + state: AuthActorState, + _state_lock: Arc, +) { + let AuthActorState { + mut leases, + mut admin_replays, + mut admin_replay_order, + } = state; + let mut last_snapshot_at = unix_seconds(); + let mut tick = tokio::time::interval(Duration::from_secs(1)); + tick.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip); + loop { + tokio::select! { + _ = tick.tick() => { + let now = unix_seconds(); + leases.tick(now); + prune_expired_admin_replays( + now, + &mut admin_replays, + &mut admin_replay_order, + ); + // WHY: A failed load starts safe mode with empty in-memory + // generations. Compacting that reconstruction would replace the + // damaged snapshot, truncate the WAL, and let the next start + // exit safe mode without rotating the instance id. + if compaction_is_allowed(inner.safe_mode.load(Ordering::Acquire)) + && now.saturating_sub(last_snapshot_at) + >= SNAPSHOT_COMPACTION_INTERVAL.as_secs() + { + let snapshot = build_snapshot(&inner, &admin_replay_order); + if let Err(error) = write_snapshot_and_truncate_wal( + &config, + &inner.admin_key(), + &snapshot, + ) { + inner.safe_mode.store(true, Ordering::Release); + cancel_all_temporary_leases(&inner); + tracing::error!( + event = "auth_state_safe_mode", + auth_stage = "snapshot_compaction", + reason = %error.code, + error = %error, + "authentication state compaction failed closed" + ); + } else { + last_snapshot_at = now; + } + } + } + command = command_rx.recv() => { + let Some(command) = command else { + admin_lease.cancel_rotated(); + cancel_all_temporary_leases(&inner); + break; + }; + match command { + AuthCommand::ClaimAdminMutation { + authority, + fingerprint, + client_timestamp, + response, + } => { + let result = validate_admin_authority(&inner, &authority).and_then(|()| { + actor_claim_admin_mutation( + &inner, + &config, + &mut admin_replays, + &mut admin_replay_order, + fingerprint, + client_timestamp, + ) + }); + let _ = response.send(result); + } + AuthCommand::Issue { authority, ttl, label, response } => { + let result = validate_admin_authority(&inner, &authority) + .and_then(|()| actor_issue(&inner, &config, &mut leases, ttl, label)); + let _ = response.send(result); + } + AuthCommand::List { authority, page, page_size, response } => { + let result = validate_admin_authority(&inner, &authority) + .and_then(|()| actor_list(&inner, page, page_size)); + let _ = response.send(result); + } + AuthCommand::Show { authority, key_id, reveal, response } => { + let result = validate_admin_authority(&inner, &authority) + .and_then(|()| actor_show(&inner, &config, key_id, reveal)); + let _ = response.send(result); + } + AuthCommand::Renew { authority, key_id, ttl, response } => { + let result = validate_admin_authority(&inner, &authority) + .and_then(|()| actor_renew(&inner, &config, &mut leases, key_id, ttl)); + let _ = response.send(result); + } + AuthCommand::Revoke { authority, key_id, response } => { + let result = validate_admin_authority(&inner, &authority) + .and_then(|()| actor_revoke(&inner, &config, &mut leases, key_id)); + let _ = response.send(result); + } + AuthCommand::Gc { authority, response } => { + let result = validate_admin_authority(&inner, &authority) + .and_then(|()| actor_gc(&inner, &config, &mut leases, &admin_replay_order)); + let _ = response.send(result); + } + AuthCommand::Reset { authority, response } => { + let result = validate_admin_authority(&inner, &authority) + .and_then(|()| actor_reset(&inner, &config, &mut leases, &admin_replay_order, "auth_state_reset")); + let _ = response.send(result); + } + AuthCommand::RotateRoot { authority, new_key, response } => { + let result = validate_admin_authority(&inner, &authority) + .and_then(|()| actor_rotate_root(&inner, &config, &mut leases, &mut admin_lease, new_key)); + if result.is_ok() { + admin_replays.clear(); + admin_replay_order.clear(); + } + let _ = response.send(result); + } + AuthCommand::SetLegacyProtocol { authority, policy, response } => { + let result = validate_admin_authority(&inner, &authority) + .and_then(|()| actor_set_legacy_protocol(&inner, &config, policy)); + let _ = response.send(result); + } + AuthCommand::Status { authority, response } => { + let result = validate_admin_authority(&inner, &authority) + .map(|()| actor_status(&inner)); + let _ = response.send(result); + } + AuthCommand::Audit { authority, action, key_id, detail, response } => { + let result = validate_admin_authority(&inner, &authority).and_then(|()| { + append_audit( + &config, + &inner, + audit(&action, key_id, detail), + ) + }); + let _ = response.send(result); + } + AuthCommand::Shutdown { response } => { + admin_lease.cancel_rotated(); + cancel_all_temporary_leases(&inner); + let _ = response.send(()); + break; + } + } + } + } + } +} + +fn actor_claim_admin_mutation( + inner: &AuthStateInner, + config: &AuthConfig, + admin_replays: &mut HashSet<[u8; 32]>, + admin_replay_order: &mut VecDeque, + fingerprint: [u8; 32], + client_timestamp: u64, +) -> Result<(), AuthFailure> { + let now = unix_seconds(); + prune_expired_admin_replays(now, admin_replays, admin_replay_order); + if admin_replays.contains(&fingerprint) { + return Err(AuthFailure::new( + "admin_request_replayed", + "administrator mutation was already admitted", + false, + )); + } + if admin_replays.len() >= ADMIN_REPLAY_CAPACITY { + return Err(AuthFailure::new( + "admin_replay_capacity_exhausted", + "administrator mutation replay window is full; retry after older claims expire", + true, + )); + } + if now.abs_diff(client_timestamp) > ADMIN_REPLAY_RETENTION.as_secs() / 2 { + return Err(AuthFailure::new( + "admin_request_timestamp_invalid", + "administrator mutation timestamp is outside the accepted window", + false, + )); + } + let record = AdminReplayRecord { + fingerprint, + client_timestamp, + accepted_at: now, + }; + fail_closed_on_uncertain_wal( + inner, + append_wal( + config, + &inner.admin_key(), + &WalRecord::AdminReplay(record.clone()), + ), + )?; + admin_replays.insert(fingerprint); + admin_replay_order.push_back(record); + Ok(()) +} + +pub(super) fn prune_expired_admin_replays( + now: u64, + admin_replays: &mut HashSet<[u8; 32]>, + admin_replay_order: &mut VecDeque, +) { + admin_replay_order.retain(|record| { + let keep = record.within_retention(now); + if !keep { + admin_replays.remove(&record.fingerprint); + } + keep + }); +} + +fn actor_status(inner: &Arc) -> AuthStatus { + let slots = inner.slots(); + let high = inner.high(); + let active_keys = slots + .iter() + .filter(|slot| slot.state == SlotState::Active) + .count() + + high + .iter() + .filter(|entry| entry.state == SlotState::Active) + .count(); + let expired_keys = slots + .iter() + .filter(|slot| slot.state == SlotState::Expired) + .count() + + high + .iter() + .filter(|entry| entry.state == SlotState::Expired) + .count(); + let revoked_keys = slots + .iter() + .filter(|slot| slot.state == SlotState::Revoked) + .count() + + high + .iter() + .filter(|entry| entry.state == SlotState::Revoked) + .count(); + let last_legacy_connection_at = inner.last_legacy_connection_at.load(Ordering::Acquire); + AuthStatus { + schema_version: 1, + safe_mode: inner.safe_mode.load(Ordering::Acquire), + capacity: slots.len(), + active_keys, + expired_keys, + revoked_keys, + legacy_protocol: if inner.legacy_protocol_allowed.load(Ordering::Acquire) { + LegacyProtocolPolicy::Allow + } else { + LegacyProtocolPolicy::Deny + }, + active_legacy_connections: inner.active_legacy_connections.load(Ordering::Acquire), + last_legacy_connection_at: (last_legacy_connection_at != 0) + .then_some(last_legacy_connection_at), + auth_successes: inner.auth_successes.load(Ordering::Relaxed), + auth_failures: inner.auth_failures.load(Ordering::Relaxed), + server_instance_id: hex(&inner.instance_id()), + } +} + +fn validate_admin_authority( + inner: &AuthStateInner, + authority: &Weak, +) -> Result<(), AuthFailure> { + let presented = authority.upgrade().ok_or_else(|| { + AuthFailure::new( + "administrator_key_rotated", + "administrator credential lease is no longer active", + false, + ) + })?; + if presented.cancellation.is_cancelled() { + return Err(AuthFailure::new( + "administrator_key_rotated", + "administrator credential lease has been cancelled", + false, + )); + } + let current = inner.admin.read().lease.upgrade().ok_or_else(|| { + AuthFailure::new( + "administrator_key_rotated", + "active administrator credential lease is unavailable", + false, + ) + })?; + if !Arc::ptr_eq(&presented, ¤t) { + return Err(AuthFailure::new( + "administrator_key_rotated", + "administrator request was authenticated before the latest root-key rotation", + false, + )); + } + Ok(()) +} + +fn ensure_store_available(inner: &AuthStateInner) -> Result<(), AuthFailure> { + if inner.safe_mode.load(Ordering::Acquire) { + Err(AuthFailure::new( + "temporary_key_store_unavailable", + "temporary key store is in administrator safe mode", + false, + )) + } else { + Ok(()) + } +} + +fn validate_slot_identity(slot: &SlotHot, key_id: KeyId) -> Result<(), AuthFailure> { + if slot.generation != key_id.generation() || slot.state == SlotState::Free { + Err(key_not_found(key_id)) + } else { + Ok(()) + } +} + +fn key_not_found(key_id: KeyId) -> AuthFailure { + AuthFailure::new( + "temporary_key_not_found", + format!("temporary key {key_id} does not exist"), + false, + ) +} + +fn key_not_renewable() -> AuthFailure { + AuthFailure::new( + "temporary_key_not_renewable", + "only an active, unexpired temporary key can be renewed", + false, + ) +} + +fn key_not_active() -> AuthFailure { + AuthFailure::new( + "temporary_key_not_active", + "temporary key is not active", + false, + ) +} + +fn slot_state_name(state: SlotState) -> &'static str { + match state { + SlotState::Free => "free", + SlotState::Active => "active", + SlotState::Expired => "expired", + SlotState::Revoked => "revoked", + } +} + +fn audit(action: &str, key_id: Option, label: Option) -> AuditRecord { + AuditRecord { + at: unix_seconds(), + action: action.to_string(), + key_id, + label, + } +} diff --git a/crates/pb-mapper-auth/src/config.rs b/crates/pb-mapper-auth/src/config.rs new file mode 100644 index 0000000..e320419 --- /dev/null +++ b/crates/pb-mapper-auth/src/config.rs @@ -0,0 +1,189 @@ +//! Authentication configuration and platform state-directory defaults. +use super::*; + +pub fn default_auth_state_dir() -> PathBuf { + std::env::var_os("PB_MAPPER_AUTH_STATE_DIR") + .map(PathBuf::from) + .unwrap_or_else(platform_default_auth_state_dir) +} + +/// Linux systemd/Docker keep `/var/lib/pb-mapper/auth` when that path is usable +/// (root, or an already-writable service directory). Unprivileged Linux, +/// macOS, and Windows binaries need an application data directory instead. +pub(crate) fn platform_default_auth_state_dir() -> PathBuf { + #[cfg(windows)] + { + let base = std::env::var_os("LOCALAPPDATA") + .or_else(|| std::env::var_os("APPDATA")) + .map(PathBuf::from) + .unwrap_or_else(|| PathBuf::from(r"C:\ProgramData")); + base.join("pb-mapper").join("auth") + } + #[cfg(target_os = "macos")] + { + match std::env::var_os("HOME") { + Some(home) => PathBuf::from(home) + .join("Library") + .join("Application Support") + .join("pb-mapper") + .join("auth"), + None => PathBuf::from("/Library/Application Support/pb-mapper/auth"), + } + } + #[cfg(not(any(windows, target_os = "macos")))] + { + linux_default_auth_state_dir( + unix_effective_uid(), + linux_system_auth_dir_usable(), + std::env::var_os("XDG_DATA_HOME").as_deref(), + std::env::var_os("HOME").as_deref(), + ) + } +} + +#[cfg(not(any(windows, target_os = "macos")))] +pub(crate) fn linux_default_auth_state_dir( + euid: u32, + system_dir_usable: bool, + xdg_data_home: Option<&std::ffi::OsStr>, + home: Option<&std::ffi::OsStr>, +) -> PathBuf { + if euid == 0 || system_dir_usable { + return PathBuf::from(DEFAULT_AUTH_STATE_DIR); + } + if let Some(xdg) = xdg_data_home + && !xdg.is_empty() + { + return PathBuf::from(xdg).join("pb-mapper").join("auth"); + } + if let Some(home) = home + && !home.is_empty() + { + return PathBuf::from(home) + .join(".local") + .join("share") + .join("pb-mapper") + .join("auth"); + } + PathBuf::from(DEFAULT_AUTH_STATE_DIR) +} + +#[cfg(not(any(windows, target_os = "macos")))] +pub(super) fn unix_effective_uid() -> u32 { + unsafe extern "C" { + fn geteuid() -> u32; + } + unsafe { geteuid() } +} + +#[cfg(not(any(windows, target_os = "macos")))] +pub(super) fn linux_system_auth_dir_usable() -> bool { + let path = Path::new(DEFAULT_AUTH_STATE_DIR); + path.is_dir() && unix_path_is_writable(path) +} + +#[cfg(not(any(windows, target_os = "macos")))] +fn unix_path_is_writable(path: &Path) -> bool { + use std::os::unix::ffi::OsStrExt; + let Ok(c_path) = std::ffi::CString::new(path.as_os_str().as_bytes()) else { + return false; + }; + unsafe extern "C" { + fn access(pathname: *const std::os::raw::c_char, mode: i32) -> i32; + } + const W_OK: i32 = 2; + unsafe { access(c_path.as_ptr(), W_OK) == 0 } +} + +impl Default for AuthConfig { + fn default() -> Self { + Self { + state_dir: default_auth_state_dir(), + max_temporary_keys: env_usize( + "PB_MAPPER_AUTH_MAX_TEMP_KEYS", + DEFAULT_TEMP_KEY_CAPACITY, + 1, + MAX_TEMP_KEY_CAPACITY, + ), + max_temporary_key_ttl: Duration::from_secs(env_u64( + "PB_MAPPER_AUTH_MAX_TEMP_TTL_SECS", + DEFAULT_MAX_TEMP_KEY_TTL.as_secs(), + MIN_TEMP_KEY_TTL.as_secs(), + MAX_TEMP_KEY_TTL.as_secs(), + )), + legacy_protocol: legacy_protocol_from_env(), + } + } +} + +fn legacy_protocol_from_env() -> LegacyProtocolPolicy { + match std::env::var("PB_MAPPER_LEGACY_PROTOCOL") { + Err(std::env::VarError::NotPresent) => LegacyProtocolPolicy::Allow, + Err(std::env::VarError::NotUnicode(_)) => { + tracing::error!( + event = "legacy_protocol_config_invalid", + "PB_MAPPER_LEGACY_PROTOCOL is not UTF-8; denying legacy framing" + ); + LegacyProtocolPolicy::Deny + } + Ok(value) => parse_legacy_protocol_policy(&value).unwrap_or_else(|| { + tracing::error!( + event = "legacy_protocol_config_invalid", + value, + "PB_MAPPER_LEGACY_PROTOCOL must be `allow` or `deny`; denying legacy framing" + ); + LegacyProtocolPolicy::Deny + }), + } +} + +pub(super) fn parse_legacy_protocol_policy(value: &str) -> Option { + match value.trim().to_ascii_lowercase().as_str() { + "allow" => Some(LegacyProtocolPolicy::Allow), + "deny" => Some(LegacyProtocolPolicy::Deny), + _ => None, + } +} + +fn env_usize(name: &str, default: usize, min: usize, max: usize) -> usize { + env_bounded(name, default, min, max) +} + +fn env_u64(name: &str, default: u64, min: u64, max: u64) -> u64 { + env_bounded(name, default, min, max) +} + +fn env_bounded(name: &str, default: T, min: T, max: T) -> T +where + T: std::str::FromStr + PartialOrd + Copy + fmt::Display, +{ + match std::env::var(name) { + Err(std::env::VarError::NotPresent) => default, + Ok(raw) => match raw.parse::() { + Ok(value) if value >= min && value <= max => value, + _ => { + tracing::warn!( + event = "auth_config_value_invalid", + variable = name, + value = raw, + min = %min, + max = %max, + fallback = %default, + "invalid authentication configuration value; using the default" + ); + default + } + }, + Err(std::env::VarError::NotUnicode(_)) => { + tracing::warn!( + event = "auth_config_value_invalid", + variable = name, + min = %min, + max = %max, + fallback = %default, + "authentication configuration value is not UTF-8; using the default" + ); + default + } + } +} diff --git a/crates/pb-mapper-auth/src/ids.rs b/crates/pb-mapper-auth/src/ids.rs new file mode 100644 index 0000000..fdd5a0e --- /dev/null +++ b/crates/pb-mapper-auth/src/ids.rs @@ -0,0 +1,129 @@ +//! Identity types for temporary credentials. +//! +//! ```text +//! KeyId (u64) — what a client presents +//! ┌──────────────────────────┬──────────────────────────┐ +//! │ Generation (high 32) │ SlotIndex (low 32) │ +//! └──────────────────────────┴──────────────────────────┘ +//! which tenant of the row which row of the table +//! ``` +//! +//! These were all bare integers, which made `make_key_id(generation, slot)` +//! accept its arguments in either order and let a slot index be compared against +//! a generation without complaint. Separate types make both a compile error, and +//! keep a `KeyId` from being used as an array index by mistake — the only way to +//! get one is [`KeyId::slot`], which is also the only place the truncation to a +//! row number is expressed. +//! +//! All three are `#[serde(transparent)]`, so persisted snapshots and the admin +//! wire protocol keep the plain-integer encoding they already had. + +use std::fmt; + +use serde::{Deserialize, Serialize}; + +/// The identity a client presents: a [`SlotIndex`] paired with the +/// [`Generation`] of the row it was issued from. +#[derive(Clone, Copy, Debug, Eq, PartialEq, Ord, PartialOrd, Hash, Serialize, Deserialize)] +#[serde(transparent)] +pub struct KeyId(u64); + +/// Which tenant of a slot a credential belongs to. Bumped every time the row is +/// reissued, and never reset, so a retired credential can never match the row +/// that replaced it. +#[derive(Clone, Copy, Debug, Eq, PartialEq, Ord, PartialOrd, Hash, Serialize, Deserialize)] +#[serde(transparent)] +pub struct Generation(u32); + +/// Which row of the slot table a credential lives in. +#[derive(Clone, Copy, Debug, Eq, PartialEq, Ord, PartialOrd, Hash, Serialize, Deserialize)] +#[serde(transparent)] +pub struct SlotIndex(u32); + +/// The administrator, which owns no slot and never expires. +pub const ADMIN_KEY_ID: KeyId = KeyId(0); + +impl KeyId { + pub const fn new(generation: Generation, slot: SlotIndex) -> Self { + Self(((generation.0 as u64) << 32) | slot.0 as u64) + } + + pub const fn generation(self) -> Generation { + Generation((self.0 >> 32) as u32) + } + + pub const fn slot(self) -> SlotIndex { + SlotIndex(self.0 as u32) + } + + pub const fn is_admin(self) -> bool { + self.0 == ADMIN_KEY_ID.0 + } + + /// The bytes mixed into the credential's key derivation. + pub const fn to_be_bytes(self) -> [u8; 8] { + self.0.to_be_bytes() + } + + pub const fn as_u64(self) -> u64 { + self.0 + } + + pub const fn from_u64(raw: u64) -> Self { + Self(raw) + } +} + +impl Generation { + pub const FIRST: Self = Self(0); + + /// The generation for a reissue of this row, or `None` once the row has been + /// cycled `u32::MAX` times and can no longer produce a fresh identity. + pub fn next(self) -> Option { + self.0.checked_add(1).map(Self) + } + + pub const fn as_u32(self) -> u32 { + self.0 + } + + pub const fn from_u32(raw: u32) -> Self { + Self(raw) + } +} + +impl SlotIndex { + pub const fn as_index(self) -> usize { + self.0 as usize + } + + /// # Panics + /// + /// If `index` exceeds `u32::MAX`. `MAX_TEMP_KEY_CAPACITY` caps the table far + /// below that, so a real index cannot reach it; panicking keeps a future + /// capacity change from silently wrapping into another row's identity. + pub fn from_index(index: usize) -> Self { + match u32::try_from(index) { + Ok(index) => Self(index), + Err(_) => panic!("slot index exceeds the addressable slot table"), + } + } +} + +impl fmt::Display for KeyId { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + self.0.fmt(formatter) + } +} + +impl fmt::Display for Generation { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + self.0.fmt(formatter) + } +} + +impl fmt::Display for SlotIndex { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + self.0.fmt(formatter) + } +} diff --git a/crates/pb-mapper-auth/src/keys.rs b/crates/pb-mapper-auth/src/keys.rs new file mode 100644 index 0000000..5df9e99 --- /dev/null +++ b/crates/pb-mapper-auth/src/keys.rs @@ -0,0 +1,231 @@ +//! Administrator credential load, recovery, and temporary-key derivation. +use super::*; + +fn read_admin_key(path: &Path) -> Result, AuthFailure> { + if !path.exists() { + return Ok(None); + } + #[cfg(unix)] + { + let metadata = std::fs::metadata(path).map_err(|error| { + AuthFailure::new( + "administrator_key_required", + format!( + "administrator key file `{}` metadata could not be read: {error}", + path.display() + ), + false, + ) + })?; + if metadata.permissions().mode() & 0o077 != 0 { + std::fs::set_permissions(path, std::fs::Permissions::from_mode(0o600)).map_err( + |error| { + AuthFailure::new( + "administrator_key_required", + format!( + "administrator key file `{}` permissions could not be secured: {error}", + path.display() + ), + false, + ) + }, + )?; + tracing::warn!( + event = "administrator_key_permissions_repaired", + path = %path.display(), + "restricted administrator key file permissions to 0600" + ); + } + } + std::fs::read_to_string(path).map(Some).map_err(|error| { + AuthFailure::new( + "administrator_key_required", + format!( + "administrator key file `{}` could not be read: {error}", + path.display() + ), + false, + ) + }) +} + +fn persist_recovery_admin_key( + state_dir: &Path, + key: &str, + mismatch_message: &'static str, +) -> Result<(), AuthFailure> { + if encrypted_auth_state_exists(state_dir) && !key_matches_existing_state(Some(state_dir), key) { + return Err(AuthFailure::new( + "administrator_key_invalid", + mismatch_message, + false, + )); + } + write_admin_key(state_dir, key) +} + +fn validate_admin_credential(raw: &str) -> Result { + let credential = parse_credential(raw.trim()) + .map_err(|error| AuthFailure::new("administrator_key_invalid", error, false))?; + if !credential.is_admin() { + return Err(AuthFailure::new( + "administrator_key_required", + "the server key file contains a temporary credential", + false, + )); + } + Ok(credential) +} + +pub(super) fn recover_admin_key_after_rotation( + state_dir: &Path, + current: &str, +) -> Result { + let snapshot_path = auth_snapshot_path(state_dir); + if !snapshot_path.exists() { + return Ok(current.to_string()); + } + let bytes = std::fs::read(&snapshot_path).map_err(|error| { + AuthFailure::new( + "temporary_key_store_unavailable", + format!("failed to read `{}`: {error}", snapshot_path.display()), + false, + ) + })?; + if let Ok(Credential::Admin(current_key)) = parse_credential(current.trim()) + && open_blob(¤t_key, &bytes).is_ok() + { + return Ok(current.to_string()); + } + let Some(next) = read_admin_key(&state_dir.join("admin.key.next"))? else { + return Ok(current.to_string()); + }; + let Ok(Credential::Admin(next_key)) = parse_credential(next.trim()) else { + return Ok(current.to_string()); + }; + if open_blob(&next_key, &bytes).is_err() { + return Ok(current.to_string()); + } + // The rotation snapshot is complete under the staged key. Leftover WAL + // records are still encrypted with the previous key. + truncate_auth_wal(state_dir)?; + write_admin_key(state_dir, next.trim())?; + let _ = std::fs::remove_file(state_dir.join("admin.key.next")); + Ok(next) +} + +pub(super) fn load_server_admin_credential(state_dir: &Path) -> Result { + let path = state_dir.join("admin.key"); + let raw = if let Some(raw) = read_admin_key(&path)? { + raw + } else if std::env::var_os(ENV_MSG_HEADER_KEY).is_some() { + let credential = get_process_credential() + .map_err(|error| AuthFailure::new("administrator_key_invalid", error, false))?; + let Credential::Admin(key) = credential else { + return Err(AuthFailure::new( + "administrator_key_required", + "the relay server cannot start with a temporary credential", + false, + )); + }; + let key = String::from_utf8(key.to_vec()).map_err(|_| { + AuthFailure::new( + "administrator_key_invalid", + "the relay administrator key must be printable UTF-8 so it can be persisted", + false, + ) + })?; + persist_recovery_admin_key( + state_dir, + &key, + "MSG_HEADER_KEY does not decrypt the existing authentication state; refusing to write admin.key", + )?; + key + } else if Path::new(MACHINE_MSG_HEADER_KEY_PATH).is_file() { + let key = std::fs::read_to_string(MACHINE_MSG_HEADER_KEY_PATH).map_err(|error| { + AuthFailure::new( + "administrator_key_required", + format!( + "legacy administrator key file `{MACHINE_MSG_HEADER_KEY_PATH}` could not be read: {error}" + ), + false, + ) + })?; + validate_admin_credential(&key)?; + persist_recovery_admin_key( + state_dir, + key.trim(), + "legacy administrator key does not decrypt the existing authentication state; refusing to write admin.key", + )?; + tracing::warn!( + event = "administrator_key_migrated", + source = MACHINE_MSG_HEADER_KEY_PATH, + destination = %path.display(), + "migrated the legacy administrator key into the v0.4 authentication state directory" + ); + key + } else { + let key = initialize_admin_key(&path, false)?; + tracing::warn!( + event = "administrator_key_initialized", + path = %path.display(), + "no administrator credential was configured; generated a random key file" + ); + key + }; + let raw = recover_admin_key_after_rotation(state_dir, &raw)?; + let credential = validate_admin_credential(&raw)?; + set_process_msg_header_key(Some(raw.trim())).map_err(AuthFailure::internal)?; + Ok(credential) +} + +/// Load or create an app-local relay root without reading or mutating the process credential. +/// +/// The Flutter process uses its configured process credential for the remote relay, while its +/// optional embedded relay owns an independent administrator key under the app data directory. +pub(super) fn load_isolated_server_admin_credential( + state_dir: &Path, +) -> Result { + let path = state_dir.join("admin.key"); + let raw = match read_admin_key(&path)? { + Some(raw) => raw, + None => { + let key = initialize_admin_key(&path, false)?; + tracing::warn!( + event = "isolated_administrator_key_initialized", + path = %path.display(), + "generated an administrator key for an embedded relay" + ); + key + } + }; + let raw = recover_admin_key_after_rotation(state_dir, &raw)?; + validate_admin_credential(&raw) +} + +pub fn derive_temporary_key( + admin_key: &AesKeyType, + instance_id: &[u8; INSTANCE_ID_LEN], + key_id: KeyId, +) -> Result { + let salt = Salt::new(HKDF_SHA256, instance_id); + let pseudo_random_key = salt.extract(admin_key); + let key_id_bytes = key_id.to_be_bytes(); + let info = [b"pb-mapper-temp-key-v1".as_slice(), key_id_bytes.as_slice()]; + let output = pseudo_random_key + .expand(&info, HkdfLen(32)) + .map_err(|_| AuthFailure::internal("failed to expand temporary key"))?; + let mut key = [0_u8; 32]; + output + .fill(&mut key) + .map_err(|_| AuthFailure::internal("failed to fill temporary key"))?; + Ok(key) +} + +struct HkdfLen(usize); + +impl ring::hkdf::KeyType for HkdfLen { + fn len(&self) -> usize { + self.0 + } +} diff --git a/crates/pb-mapper-auth/src/leases.rs b/crates/pb-mapper-auth/src/leases.rs new file mode 100644 index 0000000..8fb741c --- /dev/null +++ b/crates/pb-mapper-auth/src/leases.rs @@ -0,0 +1,356 @@ +//! Temporary-key lifetimes, scheduled on the timing wheel. +//! +//! ```text +//! issue schedule(expires_at) ── slots[i] Active, lease live +//! | +//! v deadline arrives, or a revoke fires the timer early +//! retire lease cancelled, slots[i] Expired, tombstoned_at recorded, +//! a reap timer scheduled for +TOMBSTONE_RETENTION +//! | <- a client presenting the dead credential is told +//! | "expired", not the "unknown key" it would get from +//! v an already-recycled row +//! reap slots[i] Free (generation kept), cold metadata and any high-slot +//! row removed +//! ``` +//! +//! Both stages are timers, so nothing sweeps and no queue has to stay in step +//! with the wheel. Every way a key can end runs the same callback: a deadline +//! arriving runs it on schedule, [`Timer::fire`] runs it early for a revoke or a +//! GC, and dropping the wheel runs it for a rotation or shutdown. +//! +//! `timers` maps each key to a `Weak` handle on its current timer, which is what +//! keeps key identity out of the wheel. Renewing upgrades the handle and +//! schedules the same timer at the later deadline: the earlier placement still +//! drains, but it is no longer the last reference, so nothing fires. Because the +//! map holds only `Weak` references, an entry whose timer has fired costs nothing +//! but a stale key, cleared by the callback itself. +//! +//! The callbacks hold a `Weak`, so they neither keep the state +//! alive nor touch it after a runtime has shut down. + +use super::*; + +/// A key's two scheduled stages. Both are `Weak`, so a stage that has already +/// run costs nothing but a stale map key. +#[derive(Default)] +struct Stages { + retire: Weak, + reap: Weak, +} + +pub(super) struct Leases { + inner: Weak, + wheel: TimingWheel, + /// The wall-clock second the wheel's current tick corresponds to. The wheel + /// itself only counts ticks, so this is where absolute deadlines are turned + /// into the relative delays it takes. + now: u64, + /// Each key's stages, so a renew or an early end can reach them without the + /// wheel knowing what a key is. + stages: HashMap, +} + +impl Leases { + /// Rebuilds a loaded state's schedule: live keys wait for their expiry, and + /// keys that were already dead wait out the rest of their retention. + pub(super) fn restored(inner: &Arc, now: u64) -> Self { + let mut leases = Self { + inner: Arc::downgrade(inner), + wheel: new_wheel(), + now, + stages: HashMap::new(), + }; + let mut live = Vec::new(); + let mut dead = Vec::new(); + for (index, slot) in inner.slots().iter().enumerate() { + let key_id = KeyId::new(slot.generation, SlotIndex::from_index(index)); + match slot.state { + SlotState::Active => live.extend(slot.lease.upgrade().map(|l| (key_id, l))), + SlotState::Expired | SlotState::Revoked => dead.push(key_id), + SlotState::Free => {} + } + } + dead.extend( + inner + .high() + .iter() + .filter(|entry| entry.state != SlotState::Active) + .map(|entry| entry.key_id), + ); + for (key_id, lease) in live { + leases.watch(key_id, lease); + } + for key_id in dead { + let tombstoned_at = inner + .cold() + .get(&key_id) + .map(|cold| cold.tombstoned_at) + .filter(|at| *at != 0) + .unwrap_or(now); + leases.schedule_reap(key_id, retention_ends(tombstoned_at)); + } + leases + } + + /// Takes over a newly issued key: records its description and schedules both + /// stages of its teardown. + pub(super) fn issue(&mut self, lease: &Arc, issued_at: u64, label: Option) { + let Some(inner) = self.inner.upgrade() else { + return; + }; + inner.cold_mut().insert( + lease.key_id(), + ColdMetadata { + issued_at, + label, + tombstoned_at: 0, + }, + ); + self.watch(lease.key_id(), lease.clone()); + } + + /// Hands a key's remaining life to a replacement lease, for a renewal whose + /// original lease had already been cancelled. + pub(super) fn adopt(&mut self, lease: &Arc) { + self.watch(lease.key_id(), lease.clone()); + } + + /// Moves a renewed key to its new expiry, and its reap along with it. Returns + /// `false` for a key with no live stages, as for a high slot. + /// + /// Each timer is scheduled a second time rather than moved: its earlier + /// placement drains on the old deadline but is no longer the last reference, + /// so it fires nothing. + pub(super) fn renew(&mut self, key_id: KeyId, expires_at: u64) -> bool { + let Some(stages) = self.stages.get(&key_id) else { + return false; + }; + let (Some(retire), Some(reap)) = (stages.retire.upgrade(), stages.reap.upgrade()) else { + return false; + }; + self.schedule_at(expires_at, retire); + self.schedule_at(retention_ends(expires_at), reap); + true + } + + /// Retires a key now rather than at its expiry, leaving its reap on schedule. + /// This is what a revoke needs: the credential stops working immediately, but + /// the row is held long enough to report *why*. + pub(super) fn retire_now(&mut self, key_id: KeyId) { + if let Some(retire) = self.stage(key_id, |stages| &stages.retire) { + retire.fire(); + } + } + + /// Ends a key outright, running both stages. Skips the retention wait, so it + /// is for a caller that wants the row back now. + pub(super) fn end(&mut self, key_id: KeyId) { + self.retire_now(key_id); + if let Some(reap) = self.stage(key_id, |stages| &stages.reap) { + reap.fire(); + } + self.stages.remove(&key_id); + } + + /// Runs every callback whose deadline has passed. + pub(super) fn tick(&mut self, now: u64) { + // A jump longer than anything a key can be scheduled for means every + // timer is due, so the schedule is dropped wholesale instead of ticked up + // to. That keeps a corrected hardware clock from spinning for hours. + // + // A clock stepping backwards is ignored: buckets are indexed relative to + // the wheel's `now`, so re-filing against an earlier one would place + // entries in slots it has already drained. + if now.saturating_sub(self.now) > self.wheel.max_delay() { + self.drop_schedule(now); + return; + } + while self.now < now { + self.now += 1; + self.wheel.tick(); + } + } + + /// Ends every key at once, for a root rotation or state reset. Dropping the + /// wheel releases the last reference to every timer, so no row, lease, or + /// metadata entry survives it. + pub(super) fn wipe(&mut self, now: u64) { + // Rotation is the one reason a callback cannot infer, so it is recorded + // before the drop; `record_cancel` keeps the first reason. + if let Some(inner) = self.inner.upgrade() { + for lease in inner.slots().iter().filter_map(|slot| slot.lease.upgrade()) { + lease.cancel_rotated(); + } + } + self.drop_schedule(now); + } + + /// Ends every key that is dead or past its deadline, skipping the retention + /// wait. Returns how many keys were ended. + pub(super) fn collect_garbage(&mut self, now: u64) -> u64 { + let Some(inner) = self.inner.upgrade() else { + return 0; + }; + let mut due = inner + .slots() + .iter() + .enumerate() + .filter(|(_, slot)| slot.is_collectable(now)) + .map(|(index, slot)| KeyId::new(slot.generation, SlotIndex::from_index(index))) + .collect::>(); + due.extend( + inner + .high() + .iter() + .filter(|entry| entry.state != SlotState::Active || entry.expires_at <= now) + .map(|entry| entry.key_id), + ); + for key_id in &due { + self.end(*key_id); + } + due.len() as u64 + } + + /// Replaces the whole schedule, running every callback the old one held. + fn drop_schedule(&mut self, now: u64) { + self.stages.clear(); + self.now = now; + self.wheel = new_wheel(); + } + + /// Schedules `timer` for an absolute second, as the delay from now the wheel + /// works in. A deadline already past releases the timer at once. + fn schedule_at(&mut self, deadline: u64, timer: Arc) { + self.wheel + .schedule(deadline.saturating_sub(self.now), timer); + } + + /// Upgrades one of a key's stages, forgetting the key once both have run. + fn stage( + &mut self, + key_id: KeyId, + which: impl Fn(&Stages) -> &Weak, + ) -> Option> { + let stages = self.stages.get(&key_id)?; + let timer = which(stages).upgrade(); + if stages.retire.strong_count() == 0 && stages.reap.strong_count() == 0 { + self.stages.remove(&key_id); + } + timer + } + + /// Schedules both stages of a live key: retirement at its lease's expiry, and + /// the reap a retention window later. + /// + /// WHY the reap timer owns the lease rather than the retire timer: the slot + /// table holds only a `Weak`, so this is the reference that lets a request + /// during the retention window read *why* the key died instead of finding a + /// vanished lease. Firing a timer consumes its callback, so an `Arc` held by + /// the retire stage would be released the moment that stage ran. + fn watch(&mut self, key_id: KeyId, lease: Arc) { + let inner = self.inner.clone(); + let expires_at = lease.expires_at(); + let retire_lease = lease.clone(); + let retire = Timer::new(move || { + if let Some(inner) = inner.upgrade() { + retire(&inner, key_id, &retire_lease); + } + }); + let reap = self.reap_timer(key_id, Some(lease)); + self.stages.insert( + key_id, + Stages { + retire: Arc::downgrade(&retire), + reap: Arc::downgrade(&reap), + }, + ); + self.schedule_at(expires_at, retire); + self.schedule_at(retention_ends(expires_at), reap); + } + + /// Schedules only the reap, for a key that is already dead. + fn schedule_reap(&mut self, key_id: KeyId, deadline: u64) { + let reap = self.reap_timer(key_id, None); + self.stages.insert( + key_id, + Stages { + reap: Arc::downgrade(&reap), + ..Stages::default() + }, + ); + self.schedule_at(deadline, reap); + } + + /// Builds the reap stage. `lease` is the key's live lease when there is one, + /// kept alive by this timer until the row is recycled. + fn reap_timer(&self, key_id: KeyId, lease: Option>) -> Arc { + let inner = self.inner.clone(); + Timer::new(move || { + drop(lease); + if let Some(inner) = inner.upgrade() { + reap(&inner, key_id); + } + }) + } +} + +fn retention_ends(tombstoned_at: u64) -> u64 { + tombstoned_at.saturating_add(TOMBSTONE_RETENTION.as_secs()) +} + +/// Ends a key's active stage: cancels the lease, marks the row dead, records when +/// its retention starts, and schedules the reap that frees the row. +fn retire(inner: &Arc, key_id: KeyId, lease: &Arc) { + // WHY expiry is the fallback reason: a key ended for any other reason was + // already cancelled by the code that knew that reason, and `record_cancel` + // keeps the first one, so this cannot mislabel it. + lease.cancel_expired(); + let mut slots = inner.slots_mut(); + let Some(slot) = slots.get_mut(key_id.slot().as_index()) else { + return; + }; + // Already dead, or the row moved on: a revoke marked it and recorded its + // tombstone time, and the reap is already scheduled either way. + if !slot.holds(key_id) || slot.state != SlotState::Active { + return; + } + slot.state = SlotState::Expired; + let tombstoned_at = slot.expires_at; + drop(slots); + inner + .cold_mut() + .entry(key_id) + .and_modify(|cold| cold.tombstoned_at = tombstoned_at); + tracing::info!( + event = "temporary_key_expired", + auth_stage = "expiry", + key_id = key_id.as_u64(), + expires_at = tombstoned_at, + "temporary key expired and active work was cancelled" + ); +} + +/// Frees a dead key's row and forgets it. +fn reap(inner: &Arc, key_id: KeyId) { + let mut slots = inner.slots_mut(); + match slots.get_mut(key_id.slot().as_index()) { + Some(slot) if slot.holds(key_id) => { + slot.retire(); + drop(slots); + } + Some(_) => return, + // Above the addressable table: the retained row is dropped outright, + // since only its generation has to survive. + None => { + drop(slots); + inner.high_mut().retain(|entry| entry.key_id != key_id); + } + } + inner.cold_mut().remove(&key_id); +} + +/// The wheel every schedule uses: wide enough for the longest lifetime a key can +/// have, at 64 buckets per level. +fn new_wheel() -> TimingWheel { + TimingWheel::new(MAX_SCHEDULABLE_DELAY.as_secs(), 64) +} diff --git a/crates/pb-mapper-auth/src/lib.rs b/crates/pb-mapper-auth/src/lib.rs new file mode 100644 index 0000000..ec1772b --- /dev/null +++ b/crates/pb-mapper-auth/src/lib.rs @@ -0,0 +1,784 @@ +//! Authentication state for protocol-v2 connections and administrator operations. +//! +//! # How a temporary credential works +//! +//! Nothing secret is stored per key. A temporary credential is *derived* from the +//! root key, the server instance id, and the key id, so the server can verify a +//! credential it holds no copy of, and a key id is all the state a key needs: +//! +//! ```text +//! issue: root key + instance id + key id --HKDF--> credential handed to the client +//! verify: root key + instance id + key id --HKDF--> compare against what was presented +//! ``` +//! +//! Because the material is derived, invalidating every key at once is a matter of +//! changing an input: a root rotation replaces the root key, a state reset replaces +//! the instance id. Neither has to touch individual keys. +//! +//! # Where the state lives +//! +//! ```text +//! key_id = generation:slot +//! | +//! request ──> derive & compare ──> slots[slot] ── lifecycle: Free/Active/ +//! │ Expired/Revoked, expires_at +//! │ +//! Weak lease ──> Arc lease, owned by the actor's +//! ^ timing wheel — the single place +//! │ a lease's lifetime ends +//! AuthContext (also Weak) ─────────────┘ +//! ``` +//! +//! The slot table is a preallocated array indexed straight off the key id, so +//! verification costs an array index and churn does not grow memory. The +//! `SlotState` docs below cover the table's layout, why generations exist, and +//! why dead rows linger. `Leases` (`leases.rs`) owns the three structures a +//! key's lifetime spans; `timing_wheel.rs` schedules the expiries. +//! +//! # Where mutations happen +//! +//! ```text +//! AuthRuntime (facade) ──channel──> one actor ──> encrypted snapshot + WAL +//! ``` +//! +//! Every mutation is serialized through a single actor, so a request authorized +//! before a root rotation cannot execute against the state that replaced it. +//! +//! The facade and model types stay in this root module; runtime checks, actor +//! mutations, persistence, expiry scheduling, and tests live in the children. + +use std::collections::{HashMap, HashSet, VecDeque}; +use std::fmt; +use std::fs::{File, OpenOptions}; +use std::io::{Read, Write}; +#[cfg(unix)] +use std::os::unix::fs::PermissionsExt; +use std::path::{Path, PathBuf}; +use std::sync::atomic::{AtomicBool, AtomicU8, AtomicU64, Ordering}; +use std::sync::{Arc, Weak}; +use std::time::{Duration, SystemTime, UNIX_EPOCH}; + +use parking_lot::{Mutex, RwLock}; +use rand::RngExt; +use ring::aead::{AES_256_GCM, Aad, LessSafeKey, Nonce, UnboundKey}; +use ring::hkdf::{HKDF_SHA256, Salt}; +use serde::{Deserialize, Serialize}; +use subtle::ConstantTimeEq; +use tokio::sync::{mpsc, oneshot}; +use tokio_util::sync::CancellationToken; + +use pb_mapper_core::checksum::{ + AesKeyType, Credential, ENV_MSG_HEADER_KEY, MACHINE_MSG_HEADER_KEY_PATH, + encode_temporary_credential, env_safe_admin_key_error, get_process_credential, + is_env_safe_admin_key, parse_credential, set_process_msg_header_key, +}; + +/// The namespace administrator connections operate in. Tenant namespaces are the +/// key id that owns them, so this mirrors [`ADMIN_KEY_ID`]. +pub const ADMIN_NAMESPACE: u64 = ADMIN_KEY_ID.as_u64(); +pub const DEFAULT_AUTH_STATE_DIR: &str = "/var/lib/pb-mapper/auth"; +pub const DEFAULT_TEMP_KEY_CAPACITY: usize = 65_536; +pub const MAX_TEMP_KEY_CAPACITY: usize = 1_048_576; +pub const DEFAULT_MAX_TEMP_KEY_TTL: Duration = Duration::from_secs(30 * 24 * 60 * 60); +pub const MIN_TEMP_KEY_TTL: Duration = Duration::from_secs(10); +pub const MAX_TEMP_KEY_TTL: Duration = Duration::from_secs(365 * 24 * 60 * 60); +const TOMBSTONE_RETENTION: Duration = Duration::from_secs(60); +/// Longest delay any scheduled cleanup can ask for, so the timing wheel can tell +/// a plausible wait from a clock correction. +const MAX_SCHEDULABLE_DELAY: Duration = + Duration::from_secs(MAX_TEMP_KEY_TTL.as_secs() + TOMBSTONE_RETENTION.as_secs()); +const SNAPSHOT_COMPACTION_INTERVAL: Duration = Duration::from_secs(5 * 60); +const SNAPSHOT_SCHEMA_VERSION: u16 = 1; +const STATE_BLOB_MAGIC: &[u8; 5] = b"PBAS1"; +const STATE_AAD: &[u8] = b"pb-mapper-auth-state-v1"; +const INSTANCE_ID_LEN: usize = 16; +const ADMIN_REPLAY_RETENTION: Duration = Duration::from_secs(10 * 60); +const ADMIN_REPLAY_CAPACITY: usize = 65_536; +const AUDIT_RECORD_CAPACITY: usize = 4096; + +#[derive(Clone, Copy, Debug, Eq, PartialEq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum LegacyProtocolPolicy { + Allow, + Deny, +} + +impl LegacyProtocolPolicy { + pub fn is_allowed(self) -> bool { + matches!(self, Self::Allow) + } +} + +#[derive(Clone, Debug)] +pub struct AuthConfig { + pub state_dir: PathBuf, + pub max_temporary_keys: usize, + pub max_temporary_key_ttl: Duration, + pub legacy_protocol: LegacyProtocolPolicy, +} + +#[derive(Clone, Debug, Serialize, Deserialize)] +pub struct AuthFailure { + pub code: String, + pub message: String, + pub retryable: bool, +} + +impl AuthFailure { + pub fn new(code: impl Into, message: impl Into, retryable: bool) -> Self { + Self { + code: code.into(), + message: message.into(), + retryable, + } + } + + pub fn internal(message: impl Into) -> Self { + Self::new("auth_internal_error", message, false) + } +} + +impl fmt::Display for AuthFailure { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + write!(formatter, "{}: {}", self.code, self.message) + } +} + +impl std::error::Error for AuthFailure {} + +const LEASE_CANCEL_NONE: u8 = 0; +const LEASE_CANCEL_EXPIRED: u8 = 1; +const LEASE_CANCEL_REVOKED: u8 = 2; +const LEASE_CANCEL_ROTATED: u8 = 3; + +#[derive(Debug)] +pub struct AuthLease { + key_id: KeyId, + expires_at: AtomicU64, + cancellation: CancellationToken, + cancel_reason: AtomicU8, +} + +impl AuthLease { + fn new(key_id: KeyId, expires_at: u64) -> Self { + Self { + key_id, + expires_at: AtomicU64::new(expires_at), + cancellation: CancellationToken::new(), + cancel_reason: AtomicU8::new(LEASE_CANCEL_NONE), + } + } + + pub fn key_id(&self) -> KeyId { + self.key_id + } + + pub fn expires_at(&self) -> u64 { + self.expires_at.load(Ordering::Acquire) + } + + pub fn cancellation_token(&self) -> CancellationToken { + self.cancellation.clone() + } + + fn record_cancel(&self, reason: u8) { + let _ = self.cancel_reason.compare_exchange( + LEASE_CANCEL_NONE, + reason, + Ordering::AcqRel, + Ordering::Acquire, + ); + self.cancellation.cancel(); + } + + pub(crate) fn cancel_expired(&self) { + self.record_cancel(LEASE_CANCEL_EXPIRED); + } + + pub(crate) fn cancel_revoked(&self) { + self.record_cancel(LEASE_CANCEL_REVOKED); + } + + pub(crate) fn cancel_rotated(&self) { + self.record_cancel(LEASE_CANCEL_ROTATED); + } + + #[cfg(test)] + pub(crate) fn expire_now(&self) { + self.expires_at.store(0, Ordering::Release); + } +} + +#[derive(Clone, Debug)] +pub struct AuthContext { + pub key_id: KeyId, + pub namespace: u64, + pub is_admin: bool, + lease: Weak, +} + +impl AuthContext { + fn from_lease(key_id: KeyId, is_admin: bool, lease: &Arc) -> Self { + Self { + key_id, + namespace: if is_admin { + ADMIN_NAMESPACE + } else { + key_id.as_u64() + }, + is_admin, + lease: Arc::downgrade(lease), + } + } + + pub fn ensure_active(&self) -> Result, AuthFailure> { + let lease = self.lease.upgrade().ok_or_else(|| { + AuthFailure::new( + if self.is_admin { + "administrator_key_rotated" + } else { + "temporary_key_inactive" + }, + "credential lease is no longer active", + false, + ) + })?; + if lease.cancellation.is_cancelled() { + return Err(cancelled_lease_failure(self.is_admin, &lease)); + } + if !self.is_admin && lease.expires_at() <= unix_seconds() { + lease.cancel_expired(); + return Err(AuthFailure::new( + "temporary_key_expired", + "temporary key has expired", + false, + )); + } + Ok(lease) + } + + pub fn cancellation_token(&self) -> Result { + Ok(self.ensure_active()?.cancellation_token()) + } + + pub fn admin_cancellation_token(&self) -> Result { + self.require_admin()?; + self.cancellation_token() + } + + fn admin_authority(&self) -> Result, AuthFailure> { + self.require_admin()?; + self.ensure_active()?; + Ok(self.lease.clone()) + } + + fn require_admin(&self) -> Result<(), AuthFailure> { + if self.is_admin { + Ok(()) + } else { + Err(AuthFailure::new( + "admin_permission_required", + "administrator credential is required for this operation", + false, + )) + } + } +} + +fn cancelled_lease_failure(is_admin: bool, lease: &AuthLease) -> AuthFailure { + if is_admin { + return AuthFailure::new( + "administrator_key_rotated", + "credential lease has been cancelled", + false, + ); + } + match lease.cancel_reason.load(Ordering::Acquire) { + LEASE_CANCEL_EXPIRED => { + AuthFailure::new("temporary_key_expired", "temporary key has expired", false) + } + LEASE_CANCEL_ROTATED => AuthFailure::new( + "temporary_key_rotated", + "temporary credential was invalidated by administrator root rotation or auth-state reset", + false, + ), + LEASE_CANCEL_REVOKED => { + AuthFailure::new("temporary_key_revoked", "temporary key was revoked", false) + } + _ => AuthFailure::new( + "temporary_key_inactive", + "credential lease has been cancelled", + false, + ), + } +} + +/// # The slot table +/// +/// A temporary key is never stored. It is *derived* on demand from +/// `(root key, instance id, key id)`, so the server can verify a credential it +/// has no copy of. That makes the key id the whole identity of a key, and a +/// key id is a slot index plus a generation counter: +/// +/// ```text +/// key_id: u64 +/// ┌───────────────────────────┬───────────────────────────┐ +/// │ generation (high 32) │ slot index (low 32) │ +/// └───────────────────────────┴───────────────────────────┘ +/// ^ bumped on reuse ^ where the row lives +/// ``` +/// +/// The slot index is a direct offset into `AuthStateInner::slots`, a +/// preallocated `Box<[SlotHot]>`. So verifying a credential is an array index, +/// not a map lookup or a scan, and the table's memory does not grow with churn: +/// +/// ```text +/// slots: [ SlotHot; max_temporary_keys ] +/// idx 0 gen 7 Active expires_at=… lease─┐ +/// idx 1 gen 0 Free │ Weak, so the actor's +/// idx 2 gen 3 Expired (tombstoned) │ timing wheel is the +/// idx 3 gen 9 Active expires_at=… lease─┴─ only strong owner +/// ``` +/// +/// ## Why the generation counter +/// +/// A freed slot is reused, so the index alone would let a *retired* credential +/// authenticate against the *new* tenant of that row. The generation bump makes +/// the old key id refer to a row that no longer exists: +/// +/// ```text +/// issue -> idx 2, gen 3 => key_id 0x0000_0003_0000_0002 +/// expire -> idx 2 retired, generation kept at 3 +/// reissue -> idx 2, gen 4 => key_id 0x0000_0004_0000_0002 +/// the old key id still names gen 3, which nothing matches +/// ``` +/// +/// This is why [`SlotHot::retire`] clears the row but preserves `generation`, +/// and why a generation is never reset — not by expiry, GC, root rotation, or a +/// full state reset. +/// +/// ## The lifecycle +/// +/// ```text +/// issue deadline passes / revoke +/// Free ─────────> Active ──────────────────────────────> Expired +/// ^ │ Revoked +/// │ └── renew: same row, later expires_at │ +/// │ │ +/// └──────────── retire, after TOMBSTONE_RETENTION ─────────┘ +/// ``` +/// +/// `Expired`/`Revoked` are tombstones, not garbage. A row lingers in that state +/// for `TOMBSTONE_RETENTION` so a client that presents a dead credential is told +/// *why* ("expired", "revoked") instead of receiving the indistinguishable +/// "unknown key" it would get from an already-recycled row. `Leases` owns that +/// delay; see `leases.rs`. +/// +/// ## Slots above capacity +/// +/// `max_temporary_keys` is configurable, so a restart can shrink the table below +/// what the persisted state used. Those rows cannot be indexed any more, but +/// their generations still have to be honoured — otherwise growing the table +/// again would reissue a key id that was already handed out. They are retained +/// out-of-line in `high_slot_generations` / `high_slot_entries`, which is why so +/// many operations check the array first and fall back to a scan of that vector. +#[derive(Clone, Copy, Debug, Eq, PartialEq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +enum SlotState { + /// Never used, or retired and past its tombstone. + Free, + /// A live credential. `expires_at` is authoritative. + Active, + /// Dead. Retained for `TOMBSTONE_RETENTION` so the reason survives. + Expired, + Revoked, +} + +/// One row of the slot table. `generation` outlives every other field. +#[derive(Debug)] +struct SlotHot { + generation: Generation, + state: SlotState, + expires_at: u64, + /// `Weak`, because the actor's timing wheel holds the strong reference and is + /// the single place a lease's lifetime ends. See `timing_wheel.rs`. + lease: Weak, +} + +impl SlotHot { + /// Whether this row still belongs to `key_id`'s generation. A row that has + /// been reissued belongs to a newer tenant and must not be touched on the + /// old one's behalf. + fn holds(&self, key_id: KeyId) -> bool { + self.generation == key_id.generation() && self.state != SlotState::Free + } + + /// Whether a garbage collection should free this row: it is already dead, or + /// it is active but past its deadline. + fn is_collectable(&self, now: u64) -> bool { + match self.state { + SlotState::Expired | SlotState::Revoked => true, + SlotState::Active => self.expires_at <= now, + SlotState::Free => false, + } + } + + /// Frees the slot for reuse while keeping its generation, so a key id that + /// has been handed out is never issued a second time. + fn retire(&mut self) { + *self = Self { + generation: self.generation, + ..Self::default() + }; + } +} + +impl Default for SlotHot { + fn default() -> Self { + Self { + generation: Generation::FIRST, + state: SlotState::Free, + expires_at: 0, + lease: Weak::new(), + } + } +} + +#[derive(Debug)] +struct AdminState { + key: AesKeyType, + lease: Weak, +} + +#[derive(Clone, Debug)] +struct PreviousRoot { + admin_key: AesKeyType, + instance_id: [u8; INSTANCE_ID_LEN], +} + +#[derive(Debug)] +struct AuthStateInner { + admin: RwLock, + sync_process_credential: bool, + instance_id: RwLock<[u8; INSTANCE_ID_LEN]>, + /// Preallocated, indexed directly by `key_slot(key_id)`. Documented on + /// [`SlotState`]. + slots: RwLock>, + /// Rows the configured capacity no longer covers, because a restart shrank + /// the table below what the persisted state used: + /// + /// ```text + /// slots: [ 0 1 2 3 ] <- indexable + /// high: [ 4 5 ] <- generations still honoured, out-of-line + /// ``` + /// + /// Their generations must be kept so growing the table again cannot reissue + /// a key id that was already handed out, and their entries so a still-live + /// credential in that range keeps working. This is the fallback path that + /// operations take after missing in `slots`. + high_slot_generations: RwLock>, + high_slot_entries: RwLock>, + /// Per-key description that no authentication check needs, kept out of the + /// hot slot row. Lives here rather than inside the actor so a key's handle + /// can drop it without the actor being involved; see `leases.rs`. + cold: RwLock>, + safe_mode: AtomicBool, + legacy_protocol_allowed: AtomicBool, + active_legacy_connections: AtomicU64, + last_legacy_connection_at: AtomicU64, + auth_successes: AtomicU64, + auth_failures: AtomicU64, + root_epoch: AtomicU64, + previous_root: RwLock>, + audit_records: RwLock>, +} + +impl AuthStateInner { + /// Rows the configured capacity no longer covers. See the field's docs; the + /// fallback is a scan because the range is small and rarely touched. + fn high(&self) -> parking_lot::RwLockReadGuard<'_, Vec> { + self.high_slot_entries.read() + } + + fn high_mut(&self) -> parking_lot::RwLockWriteGuard<'_, Vec> { + self.high_slot_entries.write() + } + + fn slots(&self) -> parking_lot::RwLockReadGuard<'_, Box<[SlotHot]>> { + self.slots.read() + } + + fn slots_mut(&self) -> parking_lot::RwLockWriteGuard<'_, Box<[SlotHot]>> { + self.slots.write() + } + + fn cold(&self) -> parking_lot::RwLockReadGuard<'_, HashMap> { + self.cold.read() + } + + fn cold_mut(&self) -> parking_lot::RwLockWriteGuard<'_, HashMap> { + self.cold.write() + } + + fn admin_key(&self) -> AesKeyType { + self.admin.read().key + } + + fn instance_id(&self) -> [u8; INSTANCE_ID_LEN] { + *self.instance_id.read() + } +} + +#[derive(Clone)] +pub struct AuthRuntime { + inner: Weak, + command_tx: mpsc::Sender, + config: AuthConfig, + _state_lock: Arc, + actor: Arc>>>, + actor_abort: tokio::task::AbortHandle, +} + +#[derive(Clone, Debug, Serialize, Deserialize)] +pub struct TemporaryKeyMetadata { + pub key_id: KeyId, + pub state: String, + pub issued_at: u64, + pub expires_at: u64, + pub label: Option, +} + +#[derive(Clone, Debug, Serialize, Deserialize)] +pub struct IssuedTemporaryKey { + #[serde(flatten)] + pub metadata: TemporaryKeyMetadata, + pub credential: String, +} + +#[derive(Clone, Debug, Serialize, Deserialize)] +pub struct KeyPage { + pub schema_version: u16, + pub items: Vec, + pub next_page: Option, +} + +#[derive(Clone, Debug, Serialize, Deserialize)] +pub struct AuthStatus { + pub schema_version: u16, + pub safe_mode: bool, + pub capacity: usize, + pub active_keys: usize, + pub expired_keys: usize, + pub revoked_keys: usize, + pub legacy_protocol: LegacyProtocolPolicy, + pub active_legacy_connections: u64, + pub last_legacy_connection_at: Option, + pub auth_successes: u64, + pub auth_failures: u64, + pub server_instance_id: String, +} + +#[derive(Clone, Debug)] +struct ColdMetadata { + issued_at: u64, + label: Option, + tombstoned_at: u64, +} + +enum AuthCommand { + ClaimAdminMutation { + authority: Weak, + fingerprint: [u8; 32], + client_timestamp: u64, + response: oneshot::Sender>, + }, + Issue { + authority: Weak, + ttl: Duration, + label: Option, + response: oneshot::Sender>, + }, + List { + authority: Weak, + page: u32, + page_size: u16, + response: oneshot::Sender>, + }, + Show { + authority: Weak, + key_id: KeyId, + reveal: bool, + response: oneshot::Sender>, + }, + Renew { + authority: Weak, + key_id: KeyId, + ttl: Duration, + response: oneshot::Sender>, + }, + Revoke { + authority: Weak, + key_id: KeyId, + response: oneshot::Sender>, + }, + Gc { + authority: Weak, + response: oneshot::Sender>, + }, + Reset { + authority: Weak, + response: oneshot::Sender>, + }, + RotateRoot { + authority: Weak, + new_key: AesKeyType, + response: oneshot::Sender>, + }, + SetLegacyProtocol { + authority: Weak, + policy: LegacyProtocolPolicy, + response: oneshot::Sender>, + }, + Status { + authority: Weak, + response: oneshot::Sender>, + }, + Audit { + authority: Weak, + action: String, + key_id: Option, + detail: Option, + response: oneshot::Sender>, + }, + Shutdown { + response: oneshot::Sender<()>, + }, +} + +mod config; +pub use config::default_auth_state_dir; +#[cfg(all(test, not(any(windows, target_os = "macos"))))] +pub(crate) use config::linux_default_auth_state_dir; +#[cfg(test)] +pub(crate) use config::parse_legacy_protocol_policy; +#[cfg(test)] +pub(crate) use config::platform_default_auth_state_dir; +#[cfg(all(test, not(any(windows, target_os = "macos"))))] +pub(crate) use config::{linux_system_auth_dir_usable, unix_effective_uid}; +mod keys; +pub use keys::derive_temporary_key; +#[cfg(test)] +pub(crate) use keys::recover_admin_key_after_rotation; +pub(crate) use keys::{load_isolated_server_admin_credential, load_server_admin_credential}; +mod runtime; + +pub struct LegacyConnectionGuard { + inner: Weak, +} + +impl Drop for LegacyConnectionGuard { + fn drop(&mut self) { + if let Some(inner) = self.inner.upgrade() { + inner + .active_legacy_connections + .fetch_sub(1, Ordering::AcqRel); + } + } +} + +#[derive(Clone, Debug, Serialize, Deserialize)] +struct PersistedEntry { + key_id: KeyId, + state: SlotState, + issued_at: u64, + expires_at: u64, + label: Option, + #[serde(default)] + tombstoned_at: Option, +} + +#[derive(Clone, Debug, Serialize, Deserialize)] +struct PersistedSnapshot { + schema_version: u16, + instance_id: [u8; INSTANCE_ID_LEN], + generations: Vec, + entries: Vec, + legacy_protocol: LegacyProtocolPolicy, + #[serde(default)] + admin_replays: Vec, + #[serde(default)] + audit_records: VecDeque, + #[serde(default)] + root_epoch: u64, +} + +#[derive(Clone, Debug, Serialize, Deserialize)] +struct AdminReplayRecord { + fingerprint: [u8; 32], + client_timestamp: u64, + /// Server receipt time used for retention. Older snapshots omit this field + /// (`0` after serde default) and fall back to `client_timestamp`. + #[serde(default)] + accepted_at: u64, +} + +impl AdminReplayRecord { + fn within_retention(&self, now: u64) -> bool { + let anchor = if self.accepted_at == 0 { + self.client_timestamp + } else { + self.accepted_at + }; + now.saturating_sub(anchor) <= ADMIN_REPLAY_RETENTION.as_secs() + } +} + +#[derive(Clone, Debug, Serialize, Deserialize)] +enum StateMutation { + Issue(PersistedEntry), + Renew { key_id: KeyId, expires_at: u64 }, + Revoke { key_id: KeyId, at: u64 }, + LegacyProtocol(LegacyProtocolPolicy), +} + +#[derive(Clone, Debug, Serialize, Deserialize)] +struct AuditRecord { + at: u64, + action: String, + key_id: Option, + label: Option, +} + +#[derive(Clone, Debug, Serialize, Deserialize)] +enum WalRecord { + Mutation { + mutation: StateMutation, + audit: AuditRecord, + }, + Audit(AuditRecord), + AdminReplay(AdminReplayRecord), +} + +mod actor; +use actor::{AuthActorState, run_auth_actor}; +mod persistence; +pub use persistence::*; +pub(crate) use persistence::{ + append_audit, append_mutation, append_wal, atomic_write, auth_snapshot_path, build_snapshot, + cancel_all_temporary_leases, compaction_is_allowed, empty_snapshot, + fail_closed_on_uncertain_wal, hex, key_matches_existing_state, load_or_create_instance_id, + load_persisted_state, normalize_tombstone_times, open_blob, prepare_state_dir_and_lock, + push_audit_record, push_persisted_audit, random_instance_id, recover_instance_id_after_reset, + reset_already_installed, rotation_already_installed, split_high_slot_state, truncate_auth_wal, + unix_seconds, write_admin_key, write_snapshot_and_truncate_wal, +}; +#[cfg(test)] +pub(crate) use persistence::{prepare_state_dir, read_instance_id_file, try_load_persisted_state}; +mod ids; +pub use ids::{ADMIN_KEY_ID, Generation, KeyId, SlotIndex}; +mod leases; +use leases::Leases; +mod timing_wheel; +use timing_wheel::{Timer, TimingWheel}; +#[cfg(test)] +mod tests; diff --git a/crates/pb-mapper-auth/src/persistence/admin_key.rs b/crates/pb-mapper-auth/src/persistence/admin_key.rs new file mode 100644 index 0000000..f523e42 --- /dev/null +++ b/crates/pb-mapper-auth/src/persistence/admin_key.rs @@ -0,0 +1,275 @@ +//! Administrator key files, instance id, and recovery-key identity checks. +use super::super::*; +use super::{ + atomic_write, auth_snapshot_path, auth_wal_path, encrypted_auth_state_exists, open_blob, + truncate_auth_wal, +}; + +pub(crate) fn load_or_create_instance_id( + path: &Path, +) -> Result<[u8; INSTANCE_ID_LEN], AuthFailure> { + let instance_path = path.join("server-instance-id"); + if let Some(instance_id) = read_instance_id_file(&instance_path)? { + return Ok(instance_id); + } + let instance_id = random_instance_id(); + atomic_write(&instance_path, &instance_id, 0o600)?; + Ok(instance_id) +} + +pub(crate) fn read_instance_id_file( + path: &Path, +) -> Result, AuthFailure> { + if !path.exists() { + return Ok(None); + } + let bytes = std::fs::read(path).map_err(|error| { + AuthFailure::new( + "auth_state_unavailable", + format!("failed to read `{}`: {error}", path.display()), + false, + ) + })?; + bytes.try_into().map(Some).map_err(|_| { + AuthFailure::new( + "auth_state_unavailable", + "server instance id must be exactly 16 bytes", + false, + ) + }) +} + +/// Promote `server-instance-id.next` when the snapshot already belongs to it. +/// +/// Reset writes that staged file, then the empty snapshot, then the live +/// instance-id file. A crash after the snapshot lands would otherwise fail +/// closed on the next start because the live file still has the old id. +pub(crate) fn recover_instance_id_after_reset( + state_dir: &Path, + admin_key: &AesKeyType, + current: [u8; INSTANCE_ID_LEN], +) -> Result<[u8; INSTANCE_ID_LEN], AuthFailure> { + let next_path = state_dir.join("server-instance-id.next"); + let Some(next) = read_instance_id_file(&next_path)? else { + return Ok(current); + }; + let snapshot_path = auth_snapshot_path(state_dir); + if !snapshot_path.exists() { + let _ = std::fs::remove_file(&next_path); + return Ok(current); + } + let bytes = std::fs::read(&snapshot_path).map_err(|error| { + AuthFailure::new( + "temporary_key_store_unavailable", + format!("failed to read `{}`: {error}", snapshot_path.display()), + false, + ) + })?; + let Ok(plain) = open_blob(admin_key, &bytes) else { + return Ok(current); + }; + let Ok(snapshot) = serde_json::from_slice::(&plain) else { + return Ok(current); + }; + if snapshot.instance_id == current { + let _ = std::fs::remove_file(&next_path); + return Ok(current); + } + if snapshot.instance_id != next { + return Ok(current); + } + // The reset snapshot is complete. Any leftover WAL still belongs to the + // previous instance and must not be replayed onto the new derivation id. + truncate_auth_wal(state_dir)?; + atomic_write(&state_dir.join("server-instance-id"), &next, 0o600)?; + let _ = std::fs::remove_file(&next_path); + Ok(next) +} + +pub(crate) fn random_instance_id() -> [u8; INSTANCE_ID_LEN] { + let mut instance_id = [0_u8; INSTANCE_ID_LEN]; + let mut rng = rand::rng(); + for byte in &mut instance_id { + *byte = rng.random(); + } + instance_id +} + +pub(crate) fn write_admin_key(state_dir: &Path, key: &str) -> Result<(), AuthFailure> { + atomic_write( + &state_dir.join("admin.key"), + format!("{key}\n").as_bytes(), + 0o600, + ) +} + +pub fn generate_admin_key() -> String { + const CHARSET: &[u8] = b"0123456789abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ"; + let mut rng = rand::rng(); + (0..32) + .map(|_| CHARSET[rng.random_range(0..CHARSET.len())] as char) + .collect() +} + +pub fn initialize_admin_key(path: &Path, force: bool) -> Result { + if path.exists() && !force { + return Err(AuthFailure::new( + "administrator_key_exists", + format!("administrator key file `{}` already exists", path.display()), + false, + )); + } + refuse_write_if_encrypted_state(path, force)?; + let key = generate_admin_key(); + atomic_write(path, format!("{key}\n").as_bytes(), 0o600)?; + Ok(key) +} + +pub fn write_admin_key_file(path: &Path, key: &str, force: bool) -> Result<(), AuthFailure> { + let Credential::Admin(_) = parse_credential(key) + .map_err(|error| AuthFailure::new("administrator_key_invalid", error, false))? + else { + return Err(AuthFailure::new( + "administrator_key_invalid", + "administrator key file requires a 32-byte administrator key", + false, + )); + }; + if path.exists() && !force { + return Err(AuthFailure::new( + "administrator_key_exists", + format!( + "administrator key file `{}` already exists; pass --force to replace it", + path.display() + ), + false, + )); + } + if path.file_name() == Some(std::ffi::OsStr::new("admin.key")) + && !key_matches_existing_state(path.parent(), key) + { + refuse_write_if_encrypted_state(path, force)?; + } + atomic_write(path, format!("{key}\n").as_bytes(), 0o600) +} + +pub(crate) fn reset_already_installed( + state_dir: &Path, + admin_key: &AesKeyType, + new_instance_id: &[u8; INSTANCE_ID_LEN], +) -> bool { + let Ok(Some(live)) = read_instance_id_file(&state_dir.join("server-instance-id")) else { + return false; + }; + if live != *new_instance_id { + return false; + } + let Ok(bytes) = std::fs::read(auth_snapshot_path(state_dir)) else { + return false; + }; + let Ok(plain) = open_blob(admin_key, &bytes) else { + return false; + }; + let Ok(snapshot) = serde_json::from_slice::(&plain) else { + return false; + }; + snapshot.instance_id == *new_instance_id +} + +pub(crate) fn rotation_already_installed(state_dir: &Path, new_key: &str) -> bool { + key_matches_existing_snapshot(Some(state_dir), new_key) + && live_admin_key_matches(state_dir, new_key) +} + +fn live_admin_key_matches(state_dir: &Path, new_key: &str) -> bool { + let Ok(raw) = std::fs::read(state_dir.join("admin.key")) else { + return false; + }; + let Ok(text) = std::str::from_utf8(&raw) else { + return false; + }; + text.trim().as_bytes() == new_key.trim().as_bytes() +} + +pub(crate) fn key_matches_existing_snapshot(state_dir: Option<&Path>, key: &str) -> bool { + let Some(state_dir) = state_dir else { + return false; + }; + let snapshot_path = auth_snapshot_path(state_dir); + if !snapshot_path.exists() { + return false; + } + let Ok(Credential::Admin(admin_key)) = parse_credential(key) else { + return false; + }; + let Ok(bytes) = std::fs::read(&snapshot_path) else { + return false; + }; + open_blob(&admin_key, &bytes).is_ok() +} + +pub(crate) fn key_matches_existing_state(state_dir: Option<&Path>, key: &str) -> bool { + if key_matches_existing_snapshot(state_dir, key) { + return true; + } + let Some(state_dir) = state_dir else { + return false; + }; + if auth_snapshot_path(state_dir).exists() { + return false; + } + let wal_path = auth_wal_path(state_dir); + if !wal_path.exists() { + return false; + } + let Ok(Credential::Admin(admin_key)) = parse_credential(key) else { + return false; + }; + wal_decrypts_with_key(&wal_path, &admin_key) +} + +fn wal_decrypts_with_key(path: &Path, admin_key: &AesKeyType) -> bool { + let Ok(mut file) = File::open(path) else { + return false; + }; + let Ok(metadata) = file.metadata() else { + return false; + }; + if metadata.len() == 0 { + return true; + } + let mut length = [0_u8; 4]; + if file.read_exact(&mut length).is_err() { + return false; + } + let length = u32::from_be_bytes(length) as usize; + if length == 0 || length > 1024 * 1024 { + return false; + } + let mut sealed = vec![0_u8; length]; + if file.read_exact(&mut sealed).is_err() { + return false; + } + open_blob(admin_key, &sealed).is_ok() +} + +fn refuse_write_if_encrypted_state(path: &Path, force: bool) -> Result<(), AuthFailure> { + // Creating or replacing the live root while snapshot/WAL remain leaves + // those files encrypted under the previous key. Staging `admin.key.next` + // is the rotate path and must stay allowed. + let Some(state_dir) = path.parent() else { + return Ok(()); + }; + if !encrypted_auth_state_exists(state_dir) { + return Ok(()); + } + Err(AuthFailure::new( + "administrator_key_state_exists", + format!( + "refusing to {} `{}` while encrypted auth state exists; use `pb-mapper admin root-key rotate` or `pb-mapper admin auth-state reset --confirm`", + if force { "replace" } else { "create" }, + path.display() + ), + false, + )) +} diff --git a/crates/pb-mapper-auth/src/persistence/blob.rs b/crates/pb-mapper-auth/src/persistence/blob.rs new file mode 100644 index 0000000..dbebd7f --- /dev/null +++ b/crates/pb-mapper-auth/src/persistence/blob.rs @@ -0,0 +1,74 @@ +//! AEAD wrap/unwrap for snapshot and WAL payloads. +use super::super::*; + +pub(crate) fn seal_blob(admin_key: &AesKeyType, plain: &[u8]) -> Result, AuthFailure> { + let key = LessSafeKey::new( + UnboundKey::new(&AES_256_GCM, admin_key) + .map_err(|_| AuthFailure::internal("failed to initialize state encryption key"))?, + ); + let mut nonce_bytes = [0_u8; 12]; + let mut rng = rand::rng(); + for byte in &mut nonce_bytes { + *byte = rng.random(); + } + let mut output = plain.to_vec(); + key.seal_in_place_append_tag( + Nonce::assume_unique_for_key(nonce_bytes), + Aad::from(STATE_AAD), + &mut output, + ) + .map_err(|_| AuthFailure::internal("failed to encrypt authentication state"))?; + let mut sealed = Vec::with_capacity(STATE_BLOB_MAGIC.len() + nonce_bytes.len() + output.len()); + sealed.extend_from_slice(STATE_BLOB_MAGIC); + sealed.extend_from_slice(&nonce_bytes); + sealed.extend_from_slice(&output); + Ok(sealed) +} + +pub(crate) fn open_blob(admin_key: &AesKeyType, sealed: &[u8]) -> Result, AuthFailure> { + if sealed.len() < STATE_BLOB_MAGIC.len() + 12 + AES_256_GCM.tag_len() + || &sealed[..STATE_BLOB_MAGIC.len()] != STATE_BLOB_MAGIC + { + return Err(AuthFailure::new( + "temporary_key_store_unavailable", + "authentication state blob has an invalid header", + false, + )); + } + let nonce_start = STATE_BLOB_MAGIC.len(); + let nonce_end = nonce_start + 12; + // Unreachable: the length check above guarantees these 12 bytes exist. This + // parses a file that may have been truncated or corrupted, so it reports + // rather than panics. + let nonce_bytes: [u8; 12] = sealed[nonce_start..nonce_end].try_into().map_err(|_| { + AuthFailure::new( + "temporary_key_store_unavailable", + "authentication state blob has an invalid nonce", + false, + ) + })?; + let mut plain = sealed[nonce_end..].to_vec(); + let key = LessSafeKey::new(UnboundKey::new(&AES_256_GCM, admin_key).map_err(|_| { + AuthFailure::new( + "temporary_key_store_unavailable", + "failed to initialize state decryption key", + false, + ) + })?); + let opened = key + .open_in_place( + Nonce::assume_unique_for_key(nonce_bytes), + Aad::from(STATE_AAD), + &mut plain, + ) + .map_err(|_| { + AuthFailure::new( + "temporary_key_store_unavailable", + "authentication state integrity check failed", + false, + ) + })?; + let len = opened.len(); + plain.truncate(len); + Ok(plain) +} diff --git a/crates/pb-mapper-auth/src/persistence/fs.rs b/crates/pb-mapper-auth/src/persistence/fs.rs new file mode 100644 index 0000000..c83100a --- /dev/null +++ b/crates/pb-mapper-auth/src/persistence/fs.rs @@ -0,0 +1,210 @@ +//! Directory lock, atomic replace, and parent-directory durability. +use super::super::*; +use super::hex; + +/// Create the state directory and take `auth.lock` before any credential or +/// snapshot file is read or written. +pub(crate) fn prepare_state_dir_and_lock(state_dir: &Path) -> Result, AuthFailure> { + prepare_state_dir(state_dir)?; + Ok(Arc::new(acquire_state_dir_lock(state_dir)?)) +} + +pub fn acquire_state_dir_lock(state_dir: &Path) -> Result { + let path = state_dir.join("auth.lock"); + let file = OpenOptions::new() + .create(true) + .read(true) + .write(true) + .truncate(false) + .open(&path) + .map_err(|error| { + AuthFailure::new( + "auth_state_unavailable", + format!("failed to open `{}`: {error}", path.display()), + false, + ) + })?; + lock_exclusive_nonblock(&file).map_err(|error| { + AuthFailure::new( + "auth_state_locked", + format!( + "authentication state directory `{}` is already in use: {error}", + state_dir.display() + ), + false, + ) + })?; + Ok(file) +} + +fn lock_exclusive_nonblock(file: &File) -> std::io::Result<()> { + #[cfg(unix)] + { + unsafe extern "C" { + fn flock(fd: i32, operation: i32) -> i32; + } + const LOCK_EX: i32 = 2; + const LOCK_NB: i32 = 4; + use std::os::unix::io::AsRawFd; + if unsafe { flock(file.as_raw_fd(), LOCK_EX | LOCK_NB) } != 0 { + return Err(std::io::Error::last_os_error()); + } + Ok(()) + } + #[cfg(windows)] + { + use std::os::windows::io::AsRawHandle; + const LOCKFILE_FAIL_IMMEDIATELY: u32 = 0x1; + const LOCKFILE_EXCLUSIVE_LOCK: u32 = 0x2; + #[repr(C)] + struct Overlapped { + internal: usize, + internal_high: usize, + offset: u32, + offset_high: u32, + event: *mut core::ffi::c_void, + } + extern "system" { + fn LockFileEx( + file: *mut core::ffi::c_void, + flags: u32, + reserved: u32, + bytes_low: u32, + bytes_high: u32, + overlapped: *mut Overlapped, + ) -> i32; + } + let mut overlapped = Overlapped { + internal: 0, + internal_high: 0, + offset: 0, + offset_high: 0, + event: core::ptr::null_mut(), + }; + let ok = unsafe { + LockFileEx( + file.as_raw_handle(), + LOCKFILE_FAIL_IMMEDIATELY | LOCKFILE_EXCLUSIVE_LOCK, + 0, + 1, + 0, + &mut overlapped, + ) + }; + if ok == 0 { + Err(std::io::Error::last_os_error()) + } else { + Ok(()) + } + } + #[cfg(not(any(unix, windows)))] + { + let _ = file; + Ok(()) + } +} + +/// `core`'s durability primitive, reported as an `AuthFailure`. +pub(crate) fn sync_parent_directory(path: &Path) -> Result<(), AuthFailure> { + pb_mapper_core::durable_file::sync_parent_directory(path).map_err(|error| { + AuthFailure::new( + "temporary_key_store_unavailable", + format!("failed to sync `{}`: {error}", path.display()), + false, + ) + }) +} + +pub(crate) fn prepare_state_dir(path: &Path) -> Result<(), AuthFailure> { + std::fs::create_dir_all(path).map_err(|error| { + AuthFailure::new( + "auth_state_unavailable", + format!( + "failed to create auth state directory `{}`: {error}", + path.display() + ), + false, + ) + })?; + #[cfg(unix)] + std::fs::set_permissions(path, std::fs::Permissions::from_mode(0o700)).map_err(|error| { + AuthFailure::new( + "auth_state_unavailable", + format!( + "failed to secure auth state directory `{}`: {error}", + path.display() + ), + false, + ) + })?; + Ok(()) +} + +pub(crate) fn atomic_write(path: &Path, data: &[u8], mode: u32) -> Result<(), AuthFailure> { + if let Some(parent) = path.parent() { + std::fs::create_dir_all(parent).map_err(|error| { + AuthFailure::new( + "auth_state_unavailable", + format!("failed to create `{}`: {error}", parent.display()), + false, + ) + })?; + } + let file_name = path + .file_name() + .and_then(|name| name.to_str()) + .unwrap_or("auth-state"); + let mut random_suffix = [0_u8; 8]; + let mut rng = rand::rng(); + for byte in &mut random_suffix { + *byte = rng.random(); + } + let temporary = path.with_file_name(format!( + ".{file_name}.tmp-{}-{}", + std::process::id(), + hex(&random_suffix) + )); + let mut file = OpenOptions::new() + .create_new(true) + .write(true) + .open(&temporary) + .map_err(|error| { + AuthFailure::new( + "auth_state_unavailable", + format!("failed to open `{}`: {error}", temporary.display()), + false, + ) + })?; + let result = (|| { + #[cfg(unix)] + file.set_permissions(std::fs::Permissions::from_mode(mode)) + .map_err(|error| { + AuthFailure::internal(format!("failed to set key permissions: {error}")) + })?; + #[cfg(not(unix))] + let _ = mode; + file.write_all(data) + .and_then(|()| file.sync_all()) + .map_err(|error| { + AuthFailure::new( + "auth_state_unavailable", + format!("failed to write `{}`: {error}", temporary.display()), + false, + ) + })?; + drop(file); + pb_mapper_core::durable_file::replace_file(&temporary, path).map_err(|error| { + AuthFailure::new( + "auth_state_unavailable", + format!("failed to replace `{}`: {error}", path.display()), + false, + ) + })?; + sync_parent_directory(path)?; + Ok(()) + })(); + if result.is_err() { + let _ = std::fs::remove_file(&temporary); + } + result +} diff --git a/crates/pb-mapper-auth/src/persistence/mod.rs b/crates/pb-mapper-auth/src/persistence/mod.rs new file mode 100644 index 0000000..144190d --- /dev/null +++ b/crates/pb-mapper-auth/src/persistence/mod.rs @@ -0,0 +1,77 @@ +//! Durable, encrypted authentication state and audit/replay retention. +//! +//! ```text +//! startup: lock -> admin.key -> recover instance id -> decrypt snapshot -> replay WAL +//! mutation: command -> fsync encrypted WAL -> publish hot-state change +//! compact: hot state + audit + replay set -> snapshot -> truncate WAL +//! ``` +//! +//! Snapshot replacement and administrator-key files use atomic rename. Bounded audit +//! and replay collections are carried through compaction so security history does not +//! disappear when the WAL is truncated. + +use super::*; + +mod admin_key; +mod blob; +mod fs; +mod snapshot; +mod wal; + +#[cfg(test)] +pub(crate) use admin_key::read_instance_id_file; +pub use admin_key::{generate_admin_key, initialize_admin_key, write_admin_key_file}; +pub(crate) use admin_key::{ + key_matches_existing_state, load_or_create_instance_id, random_instance_id, + recover_instance_id_after_reset, reset_already_installed, rotation_already_installed, + write_admin_key, +}; +pub(crate) use blob::{open_blob, seal_blob}; +pub use fs::acquire_state_dir_lock; +#[cfg(test)] +pub(crate) use fs::prepare_state_dir; +pub(crate) use fs::sync_parent_directory; +pub(crate) use fs::{atomic_write, prepare_state_dir_and_lock}; +#[cfg(test)] +pub(crate) use snapshot::try_load_persisted_state; +pub(crate) use snapshot::{ + build_snapshot, cancel_all_temporary_leases, compaction_is_allowed, empty_snapshot, + load_persisted_state, normalize_tombstone_times, push_audit_record, push_persisted_audit, + split_high_slot_state, +}; +pub(crate) use wal::{ + append_audit, append_mutation, append_wal, fail_closed_on_uncertain_wal, read_wal, + truncate_auth_wal, write_snapshot_and_truncate_wal, +}; + +pub(crate) const AUTH_SNAPSHOT_FILE: &str = "auth.snapshot"; +pub(crate) const AUTH_WAL_FILE: &str = "auth.wal"; + +pub(crate) fn auth_snapshot_path(state_dir: &Path) -> PathBuf { + state_dir.join(AUTH_SNAPSHOT_FILE) +} + +pub(crate) fn auth_wal_path(state_dir: &Path) -> PathBuf { + state_dir.join(AUTH_WAL_FILE) +} + +pub fn encrypted_auth_state_exists(state_dir: &Path) -> bool { + auth_snapshot_path(state_dir).exists() || auth_wal_path(state_dir).exists() +} + +pub(crate) fn unix_seconds() -> u64 { + SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_secs() +} + +pub(crate) fn hex(bytes: &[u8]) -> String { + const DIGITS: &[u8; 16] = b"0123456789abcdef"; + let mut output = String::with_capacity(bytes.len() * 2); + for byte in bytes { + output.push(DIGITS[(byte >> 4) as usize] as char); + output.push(DIGITS[(byte & 0x0f) as usize] as char); + } + output +} diff --git a/crates/pb-mapper-auth/src/persistence/snapshot.rs b/crates/pb-mapper-auth/src/persistence/snapshot.rs new file mode 100644 index 0000000..18d72f0 --- /dev/null +++ b/crates/pb-mapper-auth/src/persistence/snapshot.rs @@ -0,0 +1,284 @@ +//! Snapshot construction, load, and mutation replay onto persisted entries. +use super::super::*; +use super::{auth_snapshot_path, auth_wal_path, open_blob, read_wal}; + +pub(crate) fn compaction_is_allowed(safe_mode: bool) -> bool { + !safe_mode +} + +pub(crate) fn push_audit_record(inner: &AuthStateInner, record: AuditRecord) { + let mut records = inner.audit_records.write(); + while records.len() >= AUDIT_RECORD_CAPACITY { + records.pop_front(); + } + records.push_back(record); +} + +pub(crate) fn cancel_all_temporary_leases(inner: &AuthStateInner) { + let slots = inner.slots(); + for lease in slots.iter().filter_map(|slot| slot.lease.upgrade()) { + lease.cancel_rotated(); + } +} + +fn snapshot_generations(inner: &AuthStateInner) -> Vec { + let slots = inner.slots(); + let extra = inner.high_slot_generations.read(); + let mut generations = slots.iter().map(|slot| slot.generation).collect::>(); + generations.extend_from_slice(&extra); + generations +} + +pub(crate) fn split_high_slot_state( + snapshot: &PersistedSnapshot, + capacity: usize, +) -> (Vec, Vec) { + let high_generations = snapshot.generations.get(capacity..).unwrap_or(&[]).to_vec(); + let high_entries = snapshot + .entries + .iter() + .filter(|entry| entry.key_id.slot().as_index() >= capacity) + .cloned() + .collect(); + (high_generations, high_entries) +} + +pub(crate) fn build_snapshot( + inner: &AuthStateInner, + admin_replays: &VecDeque, +) -> PersistedSnapshot { + let slots = inner.slots(); + let cold = inner.cold(); + let generations = snapshot_generations(inner); + let mut entries = slots + .iter() + .enumerate() + .filter_map(|(index, slot)| { + if slot.state == SlotState::Free { + return None; + } + let key_id = KeyId::new(slot.generation, SlotIndex::from_index(index)); + let cold = cold.get(&key_id)?; + Some(PersistedEntry { + key_id, + state: slot.state, + issued_at: cold.issued_at, + expires_at: slot.expires_at, + label: cold.label.clone(), + tombstoned_at: (cold.tombstoned_at != 0).then_some(cold.tombstoned_at), + }) + }) + .collect::>(); + entries.extend(inner.high().iter().cloned()); + snapshot_with( + inner, + inner.instance_id(), + generations, + entries, + admin_replays, + ) +} + +pub(crate) fn normalize_tombstone_times(snapshot: &mut PersistedSnapshot, now: u64) -> bool { + let mut changed = false; + for entry in &mut snapshot.entries { + if entry.tombstoned_at.is_some() { + continue; + } + let tombstoned_at = match entry.state { + SlotState::Expired => Some(entry.expires_at), + SlotState::Revoked => snapshot + .audit_records + .iter() + .rev() + .find(|record| { + record.action == "temporary_key_revoke" && record.key_id == Some(entry.key_id) + }) + .map(|record| record.at) + .or(Some(now)), + SlotState::Free | SlotState::Active => None, + }; + if tombstoned_at.is_some() { + entry.tombstoned_at = tombstoned_at; + changed = true; + } + } + changed +} + +pub(crate) fn empty_snapshot( + inner: &AuthStateInner, + instance_id: [u8; INSTANCE_ID_LEN], + admin_replays: &VecDeque, +) -> PersistedSnapshot { + snapshot_with( + inner, + instance_id, + snapshot_generations(inner), + Vec::new(), + admin_replays, + ) +} + +fn snapshot_with( + inner: &AuthStateInner, + instance_id: [u8; INSTANCE_ID_LEN], + generations: Vec, + entries: Vec, + admin_replays: &VecDeque, +) -> PersistedSnapshot { + PersistedSnapshot { + schema_version: SNAPSHOT_SCHEMA_VERSION, + instance_id, + generations, + entries, + legacy_protocol: if inner.legacy_protocol_allowed.load(Ordering::Acquire) { + LegacyProtocolPolicy::Allow + } else { + LegacyProtocolPolicy::Deny + }, + admin_replays: admin_replays.iter().cloned().collect(), + audit_records: inner.audit_records.read().clone(), + root_epoch: inner.root_epoch.load(Ordering::Acquire), + } +} + +pub(crate) fn load_persisted_state( + config: &AuthConfig, + admin_key: &AesKeyType, + instance_id: [u8; INSTANCE_ID_LEN], +) -> (Option, bool) { + match try_load_persisted_state(config, admin_key, instance_id) { + Ok(state) => (Some(state), false), + Err(error) => { + tracing::error!( + event = "auth_state_safe_mode", + auth_stage = "state_load", + reason = %error.code, + error = %error, + "temporary key store failed closed in administrator safe mode" + ); + (None, true) + } + } +} + +pub(crate) fn try_load_persisted_state( + config: &AuthConfig, + admin_key: &AesKeyType, + instance_id: [u8; INSTANCE_ID_LEN], +) -> Result { + let snapshot_path = auth_snapshot_path(&config.state_dir); + let mut snapshot = if snapshot_path.exists() { + let bytes = std::fs::read(&snapshot_path).map_err(|error| { + AuthFailure::new( + "temporary_key_store_unavailable", + format!("failed to read `{}`: {error}", snapshot_path.display()), + false, + ) + })?; + let plain = open_blob(admin_key, &bytes)?; + serde_json::from_slice::(&plain).map_err(|error| { + AuthFailure::new( + "temporary_key_store_unavailable", + format!("failed to decode auth snapshot: {error}"), + false, + ) + })? + } else { + PersistedSnapshot { + schema_version: SNAPSHOT_SCHEMA_VERSION, + instance_id, + generations: vec![Generation::FIRST; config.max_temporary_keys], + entries: Vec::new(), + legacy_protocol: config.legacy_protocol, + admin_replays: Vec::new(), + audit_records: VecDeque::new(), + root_epoch: 0, + } + }; + if snapshot.schema_version != SNAPSHOT_SCHEMA_VERSION || snapshot.instance_id != instance_id { + return Err(AuthFailure::new( + "temporary_key_store_unavailable", + "auth snapshot schema or server instance id does not match", + false, + )); + } + if snapshot.generations.len() < config.max_temporary_keys { + snapshot + .generations + .resize(config.max_temporary_keys, Generation::FIRST); + } + + let wal_path = auth_wal_path(&config.state_dir); + if wal_path.exists() { + for record in read_wal(&wal_path, admin_key)? { + match record { + WalRecord::Mutation { mutation, audit } => { + apply_persisted_mutation(&mut snapshot, mutation)?; + push_persisted_audit(&mut snapshot.audit_records, audit); + } + WalRecord::AdminReplay(record) => snapshot.admin_replays.push(record), + WalRecord::Audit(audit) => push_persisted_audit(&mut snapshot.audit_records, audit), + } + } + } + Ok(snapshot) +} + +pub(crate) fn apply_persisted_mutation( + snapshot: &mut PersistedSnapshot, + mutation: StateMutation, +) -> Result<(), AuthFailure> { + match mutation { + StateMutation::Issue(entry) => { + let index = entry.key_id.slot().as_index(); + if snapshot.generations.len() <= index { + snapshot.generations.resize(index + 1, Generation::FIRST); + } + snapshot.generations[index] = entry.key_id.generation(); + snapshot + .entries + .retain(|current| current.key_id.slot().as_index() != index); + snapshot.entries.push(entry); + } + StateMutation::Renew { key_id, expires_at } => { + let entry = snapshot_entry_mut(snapshot, key_id, "renew")?; + entry.expires_at = expires_at; + entry.state = SlotState::Active; + entry.tombstoned_at = None; + } + StateMutation::Revoke { key_id, at } => { + let entry = snapshot_entry_mut(snapshot, key_id, "revoke")?; + entry.state = SlotState::Revoked; + entry.tombstoned_at = Some(at); + } + StateMutation::LegacyProtocol(policy) => snapshot.legacy_protocol = policy, + } + Ok(()) +} + +fn snapshot_entry_mut<'a>( + snapshot: &'a mut PersistedSnapshot, + key_id: KeyId, + operation: &str, +) -> Result<&'a mut PersistedEntry, AuthFailure> { + snapshot + .entries + .iter_mut() + .find(|entry| entry.key_id == key_id) + .ok_or_else(|| { + AuthFailure::new( + "temporary_key_store_unavailable", + format!("WAL {operation} record references an unknown key"), + false, + ) + }) +} + +pub(crate) fn push_persisted_audit(records: &mut VecDeque, record: AuditRecord) { + while records.len() >= AUDIT_RECORD_CAPACITY { + records.pop_front(); + } + records.push_back(record); +} diff --git a/crates/pb-mapper-auth/src/persistence/wal.rs b/crates/pb-mapper-auth/src/persistence/wal.rs new file mode 100644 index 0000000..a049ab8 --- /dev/null +++ b/crates/pb-mapper-auth/src/persistence/wal.rs @@ -0,0 +1,226 @@ +//! Encrypted WAL append, replay, and snapshot compaction. +use super::super::*; +use super::{ + atomic_write, auth_snapshot_path, auth_wal_path, cancel_all_temporary_leases, open_blob, + push_audit_record, seal_blob, sync_parent_directory, +}; + +pub(crate) fn fail_closed_on_uncertain_wal( + inner: &AuthStateInner, + result: Result<(), AuthFailure>, +) -> Result<(), AuthFailure> { + if let Err(error) = &result + && !error.retryable + { + inner.safe_mode.store(true, Ordering::Release); + cancel_all_temporary_leases(inner); + } + result +} + +pub(crate) fn append_mutation( + config: &AuthConfig, + inner: &AuthStateInner, + mutation: StateMutation, + audit: AuditRecord, +) -> Result<(), AuthFailure> { + fail_closed_on_uncertain_wal( + inner, + append_wal( + config, + &inner.admin_key(), + &WalRecord::Mutation { + mutation, + audit: audit.clone(), + }, + ), + )?; + push_audit_record(inner, audit); + Ok(()) +} + +pub(crate) fn append_audit( + config: &AuthConfig, + inner: &AuthStateInner, + audit: AuditRecord, +) -> Result<(), AuthFailure> { + fail_closed_on_uncertain_wal( + inner, + append_wal(config, &inner.admin_key(), &WalRecord::Audit(audit.clone())), + )?; + push_audit_record(inner, audit); + Ok(()) +} + +pub(crate) fn append_wal( + config: &AuthConfig, + admin_key: &AesKeyType, + record: &WalRecord, +) -> Result<(), AuthFailure> { + let plain = serde_json::to_vec(record).map_err(|error| { + AuthFailure::internal(format!("failed to encode auth WAL record: {error}")) + })?; + let sealed = seal_blob(admin_key, &plain)?; + let length = u32::try_from(sealed.len()) + .map_err(|_| AuthFailure::internal("auth WAL record is too large"))?; + let path = auth_wal_path(&config.state_dir); + let created = !path.exists(); + let mut file = OpenOptions::new() + .create(true) + .append(true) + .open(&path) + .map_err(|error| { + AuthFailure::new( + "temporary_key_store_unavailable", + format!("failed to open `{}`: {error}", path.display()), + true, + ) + })?; + #[cfg(unix)] + file.set_permissions(std::fs::Permissions::from_mode(0o600)) + .map_err(|error| { + AuthFailure::new( + "temporary_key_store_unavailable", + format!("failed to secure `{}`: {error}", path.display()), + false, + ) + })?; + let start_len = file + .metadata() + .map(|metadata| metadata.len()) + .map_err(|error| { + AuthFailure::new( + "temporary_key_store_unavailable", + format!("failed to inspect `{}`: {error}", path.display()), + true, + ) + })?; + if let Err(error) = file + .write_all(&length.to_be_bytes()) + .and_then(|()| file.write_all(&sealed)) + .and_then(|()| file.sync_data()) + { + // retryable == rolled_back. A later append can then start at a known + // good offset. If truncation fails, the next record would be unreadable. + let rolled_back = file + .set_len(start_len) + .and_then(|()| file.sync_data()) + .is_ok(); + return Err(AuthFailure::new( + "temporary_key_store_unavailable", + if rolled_back { + format!("failed to durably append `{}`: {error}", path.display()) + } else { + format!( + "failed to durably append `{}` and could not restore the previous WAL length: {error}", + path.display() + ) + }, + rolled_back, + )); + } + if created { + sync_parent_directory(&path)?; + } + Ok(()) +} + +pub(crate) fn read_wal(path: &Path, admin_key: &AesKeyType) -> Result, AuthFailure> { + let mut file = File::open(path).map_err(|error| { + AuthFailure::new( + "temporary_key_store_unavailable", + format!("failed to open `{}`: {error}", path.display()), + false, + ) + })?; + let mut records = Vec::new(); + loop { + let mut length = [0_u8; 4]; + match file.read(&mut length[..1]) { + Ok(0) => break, + Ok(1) => {} + Ok(_) => unreachable!("single-byte WAL prefix read"), + Err(error) => { + return Err(AuthFailure::new( + "temporary_key_store_unavailable", + format!("failed to read auth WAL length: {error}"), + false, + )); + } + } + file.read_exact(&mut length[1..]).map_err(|error| { + AuthFailure::new( + "temporary_key_store_unavailable", + format!("truncated auth WAL length: {error}"), + false, + ) + })?; + let length = u32::from_be_bytes(length) as usize; + if length > 1024 * 1024 { + return Err(AuthFailure::new( + "temporary_key_store_unavailable", + "auth WAL record exceeds 1 MiB", + false, + )); + } + let mut sealed = vec![0_u8; length]; + file.read_exact(&mut sealed).map_err(|error| { + AuthFailure::new( + "temporary_key_store_unavailable", + format!("truncated auth WAL record: {error}"), + false, + ) + })?; + let plain = open_blob(admin_key, &sealed)?; + records.push(serde_json::from_slice(&plain).map_err(|error| { + AuthFailure::new( + "temporary_key_store_unavailable", + format!("failed to decode auth WAL record: {error}"), + false, + ) + })?); + } + Ok(records) +} + +pub(crate) fn write_snapshot_and_truncate_wal( + config: &AuthConfig, + admin_key: &AesKeyType, + snapshot: &PersistedSnapshot, +) -> Result<(), AuthFailure> { + let plain = serde_json::to_vec(snapshot).map_err(|error| { + AuthFailure::internal(format!("failed to encode auth snapshot: {error}")) + })?; + let sealed = seal_blob(admin_key, &plain)?; + let snapshot_path = auth_snapshot_path(&config.state_dir); + atomic_write(&snapshot_path, &sealed, 0o600)?; + truncate_auth_wal(&config.state_dir) +} + +pub(crate) fn truncate_auth_wal(state_dir: &Path) -> Result<(), AuthFailure> { + let wal_path = auth_wal_path(state_dir); + let created = !wal_path.exists(); + let wal = OpenOptions::new() + .create(true) + .write(true) + .truncate(true) + .open(&wal_path) + .map_err(|error| { + AuthFailure::new( + "temporary_key_store_unavailable", + format!("failed to truncate `{}`: {error}", wal_path.display()), + true, + ) + })?; + wal.sync_all().map_err(|error| { + AuthFailure::new( + "temporary_key_store_unavailable", + format!("failed to sync `{}`: {error}", wal_path.display()), + true, + ) + })?; + if created { + sync_parent_directory(&wal_path)?; + } + Ok(()) +} diff --git a/crates/pb-mapper-auth/src/runtime.rs b/crates/pb-mapper-auth/src/runtime.rs new file mode 100644 index 0000000..0fb3266 --- /dev/null +++ b/crates/pb-mapper-auth/src/runtime.rs @@ -0,0 +1,638 @@ +//! Public authentication runtime facade and hot-path credential checks. +//! +//! ```text +//! process credential + persisted state +//! | +//! v +//! hot slot table (Weak leases) <---- request authentication +//! | +//! +----> lifecycle actor (strong leases + time wheel) +//! ``` +//! +//! Read-only authentication stays synchronous and allocation-light. Every administrator +//! API captures a weak authority lease and sends it to the actor, where it is compared +//! with the current lease immediately before the operation executes. + +use super::*; + +impl AuthRuntime { + pub async fn from_process(config: AuthConfig) -> Result { + let state_lock = prepare_state_dir_and_lock(&config.state_dir)?; + let credential = load_server_admin_credential(&config.state_dir)?; + let Credential::Admin(admin_key) = credential else { + return Err(AuthFailure::new( + "administrator_key_required", + "the relay server must start with the administrator credential", + false, + )); + }; + Self::start_locked(admin_key, config, true, state_lock).await + } + + /// Start an embedded relay with an administrator key owned only by its state directory. + /// + /// This deliberately leaves the process credential untouched because the containing UI uses + /// that credential for its outbound register, connect, status, and stream connections. + pub async fn from_isolated_state(config: AuthConfig) -> Result { + let state_lock = prepare_state_dir_and_lock(&config.state_dir)?; + let credential = load_isolated_server_admin_credential(&config.state_dir)?; + let Credential::Admin(admin_key) = credential else { + return Err(AuthFailure::new( + "administrator_key_required", + "the embedded relay must start with an administrator credential", + false, + )); + }; + Self::start_locked(admin_key, config, false, state_lock).await + } + + pub async fn start(admin_key: AesKeyType, config: AuthConfig) -> Result { + let state_lock = prepare_state_dir_and_lock(&config.state_dir)?; + Self::start_locked(admin_key, config, true, state_lock).await + } + + async fn start_locked( + admin_key: AesKeyType, + config: AuthConfig, + sync_process_credential: bool, + state_lock: Arc, + ) -> Result { + let instance_id = load_or_create_instance_id(&config.state_dir)?; + let instance_id = + recover_instance_id_after_reset(&config.state_dir, &admin_key, instance_id)?; + let (mut loaded, safe_mode) = load_persisted_state(&config, &admin_key, instance_id); + let now = unix_seconds(); + if let Some(state) = loaded.as_mut() + && normalize_tombstone_times(state, now) + { + write_snapshot_and_truncate_wal(&config, &admin_key, state)?; + } + let mut slots = (0..config.max_temporary_keys) + .map(|_| SlotHot::default()) + .collect::>() + .into_boxed_slice(); + let mut cold = HashMap::new(); + let mut restored_leases = Vec::new(); + + let admin_lease = Arc::new(AuthLease::new(ADMIN_KEY_ID, u64::MAX)); + if let Some(state) = loaded.as_ref() { + for (index, generation) in state.generations.iter().copied().enumerate() { + if let Some(slot) = slots.get_mut(index) { + slot.generation = generation; + } + } + for entry in &state.entries { + let index = entry.key_id.slot().as_index(); + let Some(slot) = slots.get_mut(index) else { + continue; + }; + if slot.generation != entry.key_id.generation() { + continue; + } + let state = if entry.state == SlotState::Active && entry.expires_at <= now { + SlotState::Expired + } else { + entry.state + }; + slot.state = state; + slot.expires_at = entry.expires_at; + cold.insert( + entry.key_id, + ColdMetadata { + issued_at: entry.issued_at, + label: entry.label.clone(), + tombstoned_at: match state { + SlotState::Expired => entry.tombstoned_at.unwrap_or(entry.expires_at), + SlotState::Revoked => entry.tombstoned_at.unwrap_or(now), + SlotState::Free | SlotState::Active => 0, + }, + }, + ); + if state == SlotState::Active { + let lease = Arc::new(AuthLease::new(entry.key_id, entry.expires_at)); + slot.lease = Arc::downgrade(&lease); + // Held only until the schedule below adopts them; the wheel + // is the lasting owner. + restored_leases.push(lease); + } + } + } + + let legacy_protocol = if safe_mode { + LegacyProtocolPolicy::Deny + } else { + loaded + .as_ref() + .map(|state| state.legacy_protocol) + .unwrap_or(config.legacy_protocol) + }; + let mut admin_replay_order = loaded + .as_ref() + .map(|state| { + state + .admin_replays + .iter() + .filter(|record| record.within_retention(now)) + .cloned() + .collect::>() + }) + .unwrap_or_default(); + while admin_replay_order.len() > ADMIN_REPLAY_CAPACITY { + admin_replay_order.pop_front(); + } + let admin_replays = admin_replay_order + .iter() + .map(|record| record.fingerprint) + .collect::>(); + let mut audit_records: VecDeque = loaded + .as_ref() + .map(|state| state.audit_records.iter().cloned().collect()) + .unwrap_or_default(); + while audit_records.len() > AUDIT_RECORD_CAPACITY { + audit_records.pop_front(); + } + let (high_slot_generations, mut high_slot_entries) = loaded + .as_ref() + .map(|state| split_high_slot_state(state, config.max_temporary_keys)) + .unwrap_or_default(); + for entry in &mut high_slot_entries { + if entry.state == SlotState::Active && entry.expires_at <= now { + entry.state = SlotState::Expired; + entry.tombstoned_at = Some(entry.tombstoned_at.unwrap_or(entry.expires_at)); + } + } + let inner = Arc::new(AuthStateInner { + admin: RwLock::new(AdminState { + key: admin_key, + lease: Arc::downgrade(&admin_lease), + }), + sync_process_credential, + instance_id: RwLock::new(instance_id), + slots: RwLock::new(slots), + high_slot_generations: RwLock::new(high_slot_generations), + high_slot_entries: RwLock::new(high_slot_entries), + safe_mode: AtomicBool::new(safe_mode), + legacy_protocol_allowed: AtomicBool::new(legacy_protocol.is_allowed()), + active_legacy_connections: AtomicU64::new(0), + last_legacy_connection_at: AtomicU64::new(0), + auth_successes: AtomicU64::new(0), + auth_failures: AtomicU64::new(0), + root_epoch: AtomicU64::new(loaded.as_ref().map(|state| state.root_epoch).unwrap_or(0)), + previous_root: RwLock::new(None), + audit_records: RwLock::new(audit_records), + cold: RwLock::new(cold), + }); + let (command_tx, command_rx) = mpsc::channel(256); + let actor = tokio::spawn(run_auth_actor( + inner.clone(), + admin_lease, + command_rx, + config.clone(), + AuthActorState::new( + Leases::restored(&inner, now), + admin_replays, + admin_replay_order, + ), + state_lock.clone(), + )); + let actor_abort = actor.abort_handle(); + let runtime = Self { + inner: Arc::downgrade(&inner), + command_tx, + config: config.clone(), + _state_lock: state_lock.clone(), + actor: Arc::new(Mutex::new(Some(actor))), + actor_abort, + }; + Ok(runtime) + } + + pub async fn shutdown_actor(&self) { + let (response, receiver) = oneshot::channel(); + let _ = self + .command_tx + .send(AuthCommand::Shutdown { response }) + .await; + let _ = receiver.await; + let handle = self.actor.lock().take(); + if let Some(handle) = handle { + let _ = handle.await; + } + } + + pub async fn abort_actor(&self) -> Result<(), AuthFailure> { + self.actor_abort.abort(); + let handle = self.actor.lock().take(); + if let Some(handle) = handle { + let _ = handle.await; + } + tokio::time::timeout(Duration::from_secs(5), async { + while self.inner.upgrade().is_some() { + tokio::time::sleep(Duration::from_millis(10)).await; + } + }) + .await + .map_err(|_| { + AuthFailure::new( + "auth_state_unavailable", + "authentication actor did not drop after abort", + true, + ) + }) + } + + pub fn config(&self) -> &AuthConfig { + &self.config + } + + fn inner(&self) -> Result, AuthFailure> { + self.inner.upgrade().ok_or_else(|| { + AuthFailure::new( + "auth_state_unavailable", + "authentication state manager is not running", + true, + ) + }) + } + + pub fn admin_key(&self) -> Result { + Ok(self.inner()?.admin_key()) + } + + pub fn derive_key(&self, key_id: KeyId) -> Result { + let inner = self.inner()?; + if key_id.is_admin() { + return Ok(inner.admin_key()); + } + derive_temporary_key(&inner.admin_key(), &inner.instance_id(), key_id) + } + + #[cfg(test)] + pub(crate) fn high_slot_entry_count(&self) -> usize { + self.inner().map(|inner| inner.high().len()).unwrap_or(0) + } + + pub fn derive_previous_key(&self, key_id: KeyId) -> Option { + let inner = self.inner().ok()?; + let previous = inner.previous_root.read().clone()?; + if key_id.is_admin() { + Some(previous.admin_key) + } else { + derive_temporary_key(&previous.admin_key, &previous.instance_id, key_id).ok() + } + } + + pub fn authenticate_presented( + &self, + key_id: KeyId, + presented_key: &AesKeyType, + ) -> Result { + let inner = self.inner()?; + if key_id.is_admin() { + let admin = inner.admin.read(); + if !bool::from(presented_key.ct_eq(&admin.key)) { + inner.auth_failures.fetch_add(1, Ordering::Relaxed); + return Err(AuthFailure::new( + "administrator_key_invalid", + "administrator credential does not match the active root key", + false, + )); + } + let lease = admin.lease.upgrade().ok_or_else(|| { + AuthFailure::new( + "administrator_key_rotated", + "administrator credential was rotated", + false, + ) + })?; + inner.auth_successes.fetch_add(1, Ordering::Relaxed); + return Ok(AuthContext::from_lease(ADMIN_KEY_ID, true, &lease)); + } + if inner.safe_mode.load(Ordering::Acquire) { + inner.auth_failures.fetch_add(1, Ordering::Relaxed); + return Err(AuthFailure::new( + "temporary_key_store_unavailable", + "temporary key state is unavailable; administrator reset is required", + false, + )); + } + + let expected_key = derive_temporary_key(&inner.admin_key(), &inner.instance_id(), key_id)?; + if !bool::from(presented_key.ct_eq(&expected_key)) { + inner.auth_failures.fetch_add(1, Ordering::Relaxed); + return Err(temporary_key_material_mismatch(&inner, key_id)); + } + + let index = key_id.slot().as_index(); + let generation = key_id.generation(); + let slots = inner.slots(); + let Some(slot) = slots.get(index) else { + inner.auth_failures.fetch_add(1, Ordering::Relaxed); + return Err(AuthFailure::new( + "temporary_key_not_found", + "temporary key id is outside the configured slot table", + false, + )); + }; + if slot.generation != generation { + inner.auth_failures.fetch_add(1, Ordering::Relaxed); + return Err(AuthFailure::new( + "temporary_key_generation_mismatch", + "temporary key generation does not match the current slot", + false, + )); + } + let failure = match slot.state { + SlotState::Free => Some(AuthFailure::new( + "temporary_key_not_found", + "temporary key does not exist", + false, + )), + SlotState::Expired => Some(AuthFailure::new( + "temporary_key_expired", + "temporary key has expired", + false, + )), + SlotState::Revoked => Some(AuthFailure::new( + "temporary_key_revoked", + "temporary key was revoked", + false, + )), + SlotState::Active if slot.expires_at <= unix_seconds() => { + if let Some(lease) = slot.lease.upgrade() { + lease.cancel_expired(); + } + Some(AuthFailure::new( + "temporary_key_expired", + "temporary key has expired", + false, + )) + } + SlotState::Active => None, + }; + if let Some(failure) = failure { + inner.auth_failures.fetch_add(1, Ordering::Relaxed); + return Err(failure); + } + let lease = slot.lease.upgrade().ok_or_else(|| { + inner.auth_failures.fetch_add(1, Ordering::Relaxed); + AuthFailure::new( + "temporary_key_inactive", + "temporary key lease is no longer active", + true, + ) + })?; + inner.auth_successes.fetch_add(1, Ordering::Relaxed); + Ok(AuthContext::from_lease(key_id, false, &lease)) + } + + pub fn legacy_protocol_allowed(&self) -> Result { + Ok(self + .inner()? + .legacy_protocol_allowed + .load(Ordering::Acquire)) + } + + pub fn record_legacy_connection(&self) -> Result { + let inner = self.inner()?; + inner + .active_legacy_connections + .fetch_add(1, Ordering::AcqRel); + inner + .last_legacy_connection_at + .store(unix_seconds(), Ordering::Release); + Ok(LegacyConnectionGuard { + inner: Arc::downgrade(&inner), + }) + } + + async fn request( + &self, + build: impl FnOnce(oneshot::Sender>) -> AuthCommand, + ) -> Result { + let (response, receiver) = oneshot::channel(); + self.command_tx.send(build(response)).await.map_err(|_| { + AuthFailure::new( + "auth_state_unavailable", + "authentication state manager is not running", + true, + ) + })?; + receiver.await.map_err(|_| { + AuthFailure::new( + "auth_state_unavailable", + "authentication state manager dropped the response", + true, + ) + })? + } + + pub async fn claim_admin_mutation( + &self, + authorization: &AuthContext, + fingerprint: [u8; 32], + client_timestamp: u64, + ) -> Result<(), AuthFailure> { + let authority = authorization.admin_authority()?; + self.request(|response| AuthCommand::ClaimAdminMutation { + authority, + fingerprint, + client_timestamp, + response, + }) + .await + } + + pub async fn issue( + &self, + authorization: &AuthContext, + ttl: Duration, + label: Option, + ) -> Result { + let authority = authorization.admin_authority()?; + self.request(|response| AuthCommand::Issue { + authority, + ttl, + label, + response, + }) + .await + } + + pub async fn list( + &self, + authorization: &AuthContext, + page: u32, + page_size: u16, + ) -> Result { + let authority = authorization.admin_authority()?; + self.request(|response| AuthCommand::List { + authority, + page, + page_size, + response, + }) + .await + } + + pub async fn show( + &self, + authorization: &AuthContext, + key_id: KeyId, + reveal: bool, + ) -> Result { + let authority = authorization.admin_authority()?; + self.request(|response| AuthCommand::Show { + authority, + key_id, + reveal, + response, + }) + .await + } + + pub async fn renew( + &self, + authorization: &AuthContext, + key_id: KeyId, + ttl: Duration, + ) -> Result { + let authority = authorization.admin_authority()?; + self.request(|response| AuthCommand::Renew { + authority, + key_id, + ttl, + response, + }) + .await + } + + pub async fn revoke( + &self, + authorization: &AuthContext, + key_id: KeyId, + ) -> Result { + let authority = authorization.admin_authority()?; + self.request(|response| AuthCommand::Revoke { + authority, + key_id, + response, + }) + .await + } + + pub async fn gc(&self, authorization: &AuthContext) -> Result { + let authority = authorization.admin_authority()?; + self.request(|response| AuthCommand::Gc { + authority, + response, + }) + .await + } + + pub async fn reset(&self, authorization: &AuthContext) -> Result<(), AuthFailure> { + let authority = authorization.admin_authority()?; + self.request(|response| AuthCommand::Reset { + authority, + response, + }) + .await + } + + pub async fn rotate_root( + &self, + authorization: &AuthContext, + new_key: AesKeyType, + ) -> Result<(), AuthFailure> { + let authority = authorization.admin_authority()?; + self.request(|response| AuthCommand::RotateRoot { + authority, + new_key, + response, + }) + .await + } + + pub async fn set_legacy_protocol( + &self, + authorization: &AuthContext, + policy: LegacyProtocolPolicy, + ) -> Result<(), AuthFailure> { + let authority = authorization.admin_authority()?; + self.request(|response| AuthCommand::SetLegacyProtocol { + authority, + policy, + response, + }) + .await + } + + pub async fn status(&self, authorization: &AuthContext) -> Result { + let authority = authorization.admin_authority()?; + self.request(|response| AuthCommand::Status { + authority, + response, + }) + .await + } + + pub async fn audit_admin( + &self, + authorization: &AuthContext, + action: impl Into, + key_id: Option, + detail: Option, + ) -> Result<(), AuthFailure> { + let authority = authorization.admin_authority()?; + let action = action.into(); + self.request(|response| AuthCommand::Audit { + authority, + action, + key_id, + detail, + response, + }) + .await + } +} + +fn temporary_key_material_mismatch(inner: &AuthStateInner, key_id: KeyId) -> AuthFailure { + let index = key_id.slot().as_index(); + let generation = key_id.generation(); + let slots = inner.slots(); + let current_generation = match slots.get(index) { + Some(slot) => Some(slot.generation), + None => { + let high = inner.high_slot_generations.read(); + index + .checked_sub(slots.len()) + .and_then(|offset| high.get(offset).copied()) + } + }; + let slot_is_active = slots + .get(index) + .is_some_and(|slot| slot.state == SlotState::Active && slot.generation == generation); + if slot_is_active { + return AuthFailure::new( + "temporary_key_invalid", + "temporary credential does not match the active relay key material", + false, + ); + } + let current_epoch = inner.root_epoch.load(Ordering::Acquire); + if current_epoch > 0 + && generation > Generation::FIRST + && current_generation.is_some_and(|issued| generation <= issued) + { + return AuthFailure::new( + "temporary_key_rotated", + "temporary credential was invalidated by administrator root rotation or auth-state reset", + false, + ); + } + AuthFailure::new( + "temporary_key_invalid", + "temporary credential does not match the active relay key material", + false, + ) +} diff --git a/crates/pb-mapper-auth/src/tests.rs b/crates/pb-mapper-auth/src/tests.rs new file mode 100644 index 0000000..be2755b --- /dev/null +++ b/crates/pb-mapper-auth/src/tests.rs @@ -0,0 +1,1571 @@ +//! Authentication invariants exercised at the state-machine boundary. +//! +//! ```text +//! issue -> renew -> expire/revoke -> persist/restart +//! | | +//! +-> lease cancellation +-> encrypted recovery +//! root rotate -> reject old key + reject already-authenticated old context +//! ``` +//! +//! Protocol framing has its own tests under `common::message::secure::tests`; this +//! module focuses on lifecycle, persistence, audit, replay, and timing-wheel behavior. + +use pb_mapper_core::test_support::PROCESS_CREDENTIAL_TEST_LOCK; + +use super::*; + +fn temp_state_dir(name: &str) -> PathBuf { + let mut suffix = [0_u8; 8]; + let mut rng = rand::rng(); + for byte in &mut suffix { + *byte = rng.random(); + } + std::env::temp_dir().join(format!("pb-mapper-{name}-{}", hex(&suffix))) +} + +fn authenticate_for_test(runtime: &AuthRuntime, key_id: KeyId) -> Result { + let key = runtime.derive_key(key_id)?; + runtime.authenticate_presented(key_id, &key) +} + +#[test] +fn initialize_admin_key_refuses_to_replace_a_key_when_encrypted_state_exists() { + let state_dir = temp_state_dir("force-init-state"); + std::fs::create_dir_all(&state_dir).unwrap(); + let key_path = state_dir.join("admin.key"); + std::fs::write(&key_path, b"0123456789abcdefghijklmnopqrstuv\n").unwrap(); + std::fs::write(state_dir.join("auth.snapshot"), b"encrypted").unwrap(); + let error = initialize_admin_key(&key_path, true).unwrap_err(); + assert_eq!(error.code, "administrator_key_state_exists"); + let missing = state_dir.join("missing-admin.key"); + let error = initialize_admin_key(&missing, false).unwrap_err(); + assert_eq!(error.code, "administrator_key_state_exists"); + let error = + write_admin_key_file(&key_path, "abcdefghijklmnopqrstuvwxyz012345", true).unwrap_err(); + assert_eq!(error.code, "administrator_key_state_exists"); + write_admin_key_file( + &state_dir.join("admin.key.next"), + "abcdefghijklmnopqrstuvwxyz012345", + true, + ) + .unwrap(); + let _ = std::fs::remove_dir_all(state_dir); +} + +#[tokio::test] +async fn shrinking_then_expanding_capacity_does_not_reuse_old_key_ids() { + let state_dir = temp_state_dir("capacity-shrink"); + let admin_key = *b"0123456789abcdefghijklmnopqrstuv"; + let config_two = AuthConfig { + state_dir: state_dir.clone(), + max_temporary_keys: 2, + max_temporary_key_ttl: Duration::from_secs(3600), + legacy_protocol: LegacyProtocolPolicy::Allow, + }; + let runtime = AuthRuntime::start(admin_key, config_two.clone()) + .await + .unwrap(); + let admin = authenticate_for_test(&runtime, ADMIN_KEY_ID).unwrap(); + let first = runtime + .issue(&admin, Duration::from_secs(60), Some("first".to_string())) + .await + .unwrap(); + let second = runtime + .issue(&admin, Duration::from_secs(60), Some("second".to_string())) + .await + .unwrap(); + let Credential::Temporary { + key_id: first_id, .. + } = parse_credential(&first.credential).unwrap() + else { + panic!("expected temporary credential"); + }; + let Credential::Temporary { + key_id: second_id, .. + } = parse_credential(&second.credential).unwrap() + else { + panic!("expected temporary credential"); + }; + runtime + .revoke(&admin, KeyId::from_u64(first_id)) + .await + .unwrap(); + runtime + .revoke(&admin, KeyId::from_u64(second_id)) + .await + .unwrap(); + runtime.gc(&admin).await.unwrap(); + drop(runtime); + tokio::time::sleep(Duration::from_millis(20)).await; + + let config_one = AuthConfig { + max_temporary_keys: 1, + ..config_two.clone() + }; + let runtime = AuthRuntime::start(admin_key, config_one).await.unwrap(); + let admin = authenticate_for_test(&runtime, ADMIN_KEY_ID).unwrap(); + assert!(!runtime.status(&admin).await.unwrap().safe_mode); + let _third = runtime + .issue(&admin, Duration::from_secs(60), Some("third".to_string())) + .await + .unwrap(); + drop(runtime); + tokio::time::sleep(Duration::from_millis(20)).await; + + let runtime = AuthRuntime::start(admin_key, config_two).await.unwrap(); + let admin = authenticate_for_test(&runtime, ADMIN_KEY_ID).unwrap(); + assert!(!runtime.status(&admin).await.unwrap().safe_mode); + let fourth = runtime + .issue(&admin, Duration::from_secs(60), Some("fourth".to_string())) + .await + .unwrap(); + let Credential::Temporary { + key_id: fourth_id, .. + } = parse_credential(&fourth.credential).unwrap() + else { + panic!("expected temporary credential"); + }; + assert_ne!(fourth_id, second_id); + drop(runtime); + tokio::time::sleep(Duration::from_millis(20)).await; + let _ = std::fs::remove_dir_all(state_dir); +} + +#[tokio::test] +async fn gc_removes_inactive_high_slot_entries_and_keeps_their_generations() { + let state_dir = temp_state_dir("gc-high-slots"); + let admin_key = *b"0123456789abcdefghijklmnopqrstuv"; + let config_two = AuthConfig { + state_dir: state_dir.clone(), + max_temporary_keys: 2, + max_temporary_key_ttl: Duration::from_secs(3600), + legacy_protocol: LegacyProtocolPolicy::Allow, + }; + let runtime = AuthRuntime::start(admin_key, config_two.clone()) + .await + .unwrap(); + let admin = authenticate_for_test(&runtime, ADMIN_KEY_ID).unwrap(); + let first = runtime + .issue(&admin, Duration::from_secs(60), Some("first".to_string())) + .await + .unwrap(); + let second = runtime + .issue(&admin, Duration::from_secs(60), Some("second".to_string())) + .await + .unwrap(); + let Credential::Temporary { + key_id: first_id, .. + } = parse_credential(&first.credential).unwrap() + else { + panic!("expected temporary credential"); + }; + let Credential::Temporary { + key_id: second_id, .. + } = parse_credential(&second.credential).unwrap() + else { + panic!("expected temporary credential"); + }; + runtime + .revoke(&admin, KeyId::from_u64(first_id)) + .await + .unwrap(); + runtime + .revoke(&admin, KeyId::from_u64(second_id)) + .await + .unwrap(); + drop(runtime); + tokio::time::sleep(Duration::from_millis(20)).await; + + let config_one = AuthConfig { + max_temporary_keys: 1, + ..config_two.clone() + }; + let runtime = AuthRuntime::start(admin_key, config_one).await.unwrap(); + let admin = authenticate_for_test(&runtime, ADMIN_KEY_ID).unwrap(); + assert_eq!(runtime.high_slot_entry_count(), 1); + let removed = runtime.gc(&admin).await.unwrap(); + assert!(removed >= 1); + assert_eq!(runtime.high_slot_entry_count(), 0); + drop(runtime); + tokio::time::sleep(Duration::from_millis(20)).await; + + let runtime = AuthRuntime::start(admin_key, config_two).await.unwrap(); + let admin = authenticate_for_test(&runtime, ADMIN_KEY_ID).unwrap(); + let _low = runtime + .issue(&admin, Duration::from_secs(60), Some("low".to_string())) + .await + .unwrap(); + let high = runtime + .issue(&admin, Duration::from_secs(60), Some("high".to_string())) + .await + .unwrap(); + let Credential::Temporary { + key_id: reused_id, .. + } = parse_credential(&high.credential).unwrap() + else { + panic!("expected temporary credential"); + }; + assert_ne!(reused_id, second_id); + drop(runtime); + tokio::time::sleep(Duration::from_millis(20)).await; + let _ = std::fs::remove_dir_all(state_dir); +} + +#[tokio::test] +async fn admin_lifecycle_covers_high_slot_keys_after_capacity_shrink() { + let state_dir = temp_state_dir("high-slot-admin"); + let admin_key = *b"0123456789abcdefghijklmnopqrstuv"; + let config_two = AuthConfig { + state_dir: state_dir.clone(), + max_temporary_keys: 2, + max_temporary_key_ttl: Duration::from_secs(3600), + legacy_protocol: LegacyProtocolPolicy::Allow, + }; + let runtime = AuthRuntime::start(admin_key, config_two.clone()) + .await + .unwrap(); + let admin = authenticate_for_test(&runtime, ADMIN_KEY_ID).unwrap(); + let first = runtime + .issue(&admin, Duration::from_secs(60), Some("first".to_string())) + .await + .unwrap(); + let second = runtime + .issue(&admin, Duration::from_secs(60), Some("second".to_string())) + .await + .unwrap(); + drop(runtime); + tokio::time::sleep(Duration::from_millis(20)).await; + + let config_one = AuthConfig { + max_temporary_keys: 1, + ..config_two.clone() + }; + let runtime = AuthRuntime::start(admin_key, config_one).await.unwrap(); + let admin = authenticate_for_test(&runtime, ADMIN_KEY_ID).unwrap(); + assert_eq!(runtime.high_slot_entry_count(), 1); + let high_id = [first.metadata.key_id, second.metadata.key_id] + .into_iter() + .find(|key_id| key_id.slot().as_index() >= 1) + .expect("one issued key should land above the shrunken table"); + let page = runtime.list(&admin, 0, 100).await.unwrap(); + assert_eq!(page.items.len(), 2); + assert!(page.items.iter().any(|item| item.key_id == high_id)); + let shown = runtime.show(&admin, high_id, false).await.unwrap(); + assert_eq!(shown.metadata.key_id, high_id); + assert_eq!(shown.metadata.state, "active"); + assert_eq!( + authenticate_for_test(&runtime, high_id).unwrap_err().code, + "temporary_key_not_found" + ); + let status = runtime.status(&admin).await.unwrap(); + assert_eq!(status.active_keys, 2); + let renewed = runtime + .renew(&admin, high_id, Duration::from_secs(120)) + .await + .unwrap(); + assert!(renewed.metadata.expires_at > shown.metadata.expires_at); + runtime.revoke(&admin, high_id).await.unwrap(); + let revoked = runtime.show(&admin, high_id, false).await.unwrap(); + assert_eq!(revoked.metadata.state, "revoked"); + drop(runtime); + tokio::time::sleep(Duration::from_millis(20)).await; + + let runtime = AuthRuntime::start(admin_key, config_two).await.unwrap(); + let admin = authenticate_for_test(&runtime, ADMIN_KEY_ID).unwrap(); + let restored = runtime.show(&admin, high_id, false).await.unwrap(); + assert_eq!(restored.metadata.state, "revoked"); + drop(runtime); + tokio::time::sleep(Duration::from_millis(20)).await; + let _ = std::fs::remove_dir_all(state_dir); +} + +#[tokio::test] +async fn safe_mode_denies_legacy_protocol_instead_of_restoring_the_default() { + let state_dir = temp_state_dir("safe-mode-legacy"); + let admin_key = *b"0123456789abcdefghijklmnopqrstuv"; + let config = AuthConfig { + state_dir: state_dir.clone(), + max_temporary_keys: 4, + max_temporary_key_ttl: Duration::from_secs(3600), + legacy_protocol: LegacyProtocolPolicy::Allow, + }; + let runtime = AuthRuntime::start(admin_key, config.clone()).await.unwrap(); + let admin = authenticate_for_test(&runtime, ADMIN_KEY_ID).unwrap(); + runtime + .set_legacy_protocol(&admin, LegacyProtocolPolicy::Deny) + .await + .unwrap(); + drop(runtime); + tokio::time::sleep(Duration::from_millis(20)).await; + std::fs::write(state_dir.join("auth.wal"), b"broken-wal").unwrap(); + + let runtime = AuthRuntime::start(admin_key, config).await.unwrap(); + let admin = authenticate_for_test(&runtime, ADMIN_KEY_ID).unwrap(); + let status = runtime.status(&admin).await.unwrap(); + assert!(status.safe_mode); + assert_eq!(status.legacy_protocol, LegacyProtocolPolicy::Deny); + drop(runtime); + tokio::time::sleep(Duration::from_millis(20)).await; + let _ = std::fs::remove_dir_all(state_dir); +} + +#[tokio::test] +async fn overlapping_runtimes_cannot_share_an_auth_state_directory() { + let state_dir = temp_state_dir("auth-dir-lock"); + let admin_key = *b"0123456789abcdefghijklmnopqrstuv"; + let config = AuthConfig { + state_dir: state_dir.clone(), + max_temporary_keys: 4, + max_temporary_key_ttl: Duration::from_secs(3600), + legacy_protocol: LegacyProtocolPolicy::Allow, + }; + let first = AuthRuntime::start(admin_key, config.clone()).await.unwrap(); + let error = match AuthRuntime::start(admin_key, config.clone()).await { + Ok(_) => panic!("second runtime should not share the auth directory"), + Err(error) => error, + }; + assert_eq!(error.code, "auth_state_locked"); + drop(first); + tokio::time::sleep(Duration::from_millis(20)).await; + let recovered = AuthRuntime::start(admin_key, config).await.unwrap(); + drop(recovered); + tokio::time::sleep(Duration::from_millis(20)).await; + let _ = std::fs::remove_dir_all(state_dir); +} + +#[tokio::test] +async fn env_recovery_key_is_not_written_when_it_cannot_decrypt_existing_state() { + let _process_credential_guard = PROCESS_CREDENTIAL_TEST_LOCK.lock().await; + let state_dir = temp_state_dir("env-key-must-match-snapshot"); + prepare_state_dir(&state_dir).unwrap(); + let good = *b"0123456789abcdefghijklmnopqrstuv"; + let bad = *b"abcdefghijklmnopqrstuvwxyz012345"; + let config = AuthConfig { + state_dir: state_dir.clone(), + max_temporary_keys: 1, + max_temporary_key_ttl: Duration::from_secs(3600), + legacy_protocol: LegacyProtocolPolicy::Allow, + }; + let snapshot = PersistedSnapshot { + schema_version: SNAPSHOT_SCHEMA_VERSION, + instance_id: [4_u8; INSTANCE_ID_LEN], + generations: vec![Generation::FIRST; 1], + entries: Vec::new(), + legacy_protocol: LegacyProtocolPolicy::Allow, + admin_replays: Vec::new(), + audit_records: VecDeque::new(), + root_epoch: 0, + }; + write_snapshot_and_truncate_wal(&config, &good, &snapshot).unwrap(); + set_process_msg_header_key(Some(std::str::from_utf8(&bad).unwrap())).unwrap(); + // SAFETY: this test holds `PROCESS_CREDENTIAL_TEST_LOCK`, which + // serialises every test that touches the process credential. + unsafe { + std::env::set_var(ENV_MSG_HEADER_KEY, std::str::from_utf8(&bad).unwrap()); + }; + let error = match AuthRuntime::from_process(config).await { + Ok(_) => panic!("a mismatched recovery key must not start the runtime"), + Err(error) => error, + }; + // SAFETY: this test holds `PROCESS_CREDENTIAL_TEST_LOCK`, which + // serialises every test that touches the process credential. + unsafe { + std::env::remove_var(ENV_MSG_HEADER_KEY); + }; + set_process_msg_header_key(None).unwrap(); + assert_eq!(error.code, "administrator_key_invalid"); + assert!( + !state_dir.join("admin.key").exists(), + "a mismatched MSG_HEADER_KEY must not become the live administrator key" + ); + let _ = std::fs::remove_dir_all(state_dir); +} + +#[tokio::test] +async fn env_recovery_key_is_accepted_for_wal_only_state() { + let _process_credential_guard = PROCESS_CREDENTIAL_TEST_LOCK.lock().await; + let state_dir = temp_state_dir("env-key-matches-wal"); + prepare_state_dir(&state_dir).unwrap(); + let good = *b"0123456789abcdefghijklmnopqrstuv"; + let config = AuthConfig { + state_dir: state_dir.clone(), + max_temporary_keys: 1, + max_temporary_key_ttl: Duration::from_secs(3600), + legacy_protocol: LegacyProtocolPolicy::Allow, + }; + append_wal( + &config, + &good, + &WalRecord::Audit(AuditRecord { + at: 1, + action: "issue".to_string(), + key_id: None, + label: None, + }), + ) + .unwrap(); + set_process_msg_header_key(Some(std::str::from_utf8(&good).unwrap())).unwrap(); + // SAFETY: this test holds `PROCESS_CREDENTIAL_TEST_LOCK`, which + // serialises every test that touches the process credential. + unsafe { + std::env::set_var(ENV_MSG_HEADER_KEY, std::str::from_utf8(&good).unwrap()); + }; + let started = AuthRuntime::from_process(config).await; + // SAFETY: this test holds `PROCESS_CREDENTIAL_TEST_LOCK`, which + // serialises every test that touches the process credential. + unsafe { + std::env::remove_var(ENV_MSG_HEADER_KEY); + }; + set_process_msg_header_key(None).unwrap(); + started.expect("a matching recovery key must start from WAL-only state"); + assert!( + state_dir.join("admin.key").exists(), + "a matching MSG_HEADER_KEY should become the live administrator key" + ); + let _ = std::fs::remove_dir_all(state_dir); +} + +#[tokio::test] +async fn env_recovery_key_is_not_written_when_wal_only_state_does_not_match() { + let _process_credential_guard = PROCESS_CREDENTIAL_TEST_LOCK.lock().await; + let state_dir = temp_state_dir("env-key-must-match-wal"); + prepare_state_dir(&state_dir).unwrap(); + let good = *b"0123456789abcdefghijklmnopqrstuv"; + let bad = *b"abcdefghijklmnopqrstuvwxyz012345"; + let config = AuthConfig { + state_dir: state_dir.clone(), + max_temporary_keys: 1, + max_temporary_key_ttl: Duration::from_secs(3600), + legacy_protocol: LegacyProtocolPolicy::Allow, + }; + append_wal( + &config, + &good, + &WalRecord::Audit(AuditRecord { + at: 1, + action: "issue".to_string(), + key_id: None, + label: None, + }), + ) + .unwrap(); + set_process_msg_header_key(Some(std::str::from_utf8(&bad).unwrap())).unwrap(); + // SAFETY: this test holds `PROCESS_CREDENTIAL_TEST_LOCK`, which + // serialises every test that touches the process credential. + unsafe { + std::env::set_var(ENV_MSG_HEADER_KEY, std::str::from_utf8(&bad).unwrap()); + }; + let error = match AuthRuntime::from_process(config).await { + Ok(_) => panic!("a mismatched recovery key must not start from WAL-only state"), + Err(error) => error, + }; + // SAFETY: this test holds `PROCESS_CREDENTIAL_TEST_LOCK`, which + // serialises every test that touches the process credential. + unsafe { + std::env::remove_var(ENV_MSG_HEADER_KEY); + }; + set_process_msg_header_key(None).unwrap(); + assert_eq!(error.code, "administrator_key_invalid"); + assert!( + !state_dir.join("admin.key").exists(), + "a mismatched MSG_HEADER_KEY must not become the live administrator key" + ); + let _ = std::fs::remove_dir_all(state_dir); +} + +#[tokio::test] +async fn from_isolated_state_takes_the_state_lock_before_creating_admin_key() { + let state_dir = temp_state_dir("lock-before-key"); + prepare_state_dir(&state_dir).unwrap(); + let _lock = acquire_state_dir_lock(&state_dir).unwrap(); + let config = AuthConfig { + state_dir: state_dir.clone(), + max_temporary_keys: 4, + max_temporary_key_ttl: Duration::from_secs(3600), + legacy_protocol: LegacyProtocolPolicy::Allow, + }; + let error = match AuthRuntime::from_isolated_state(config).await { + Ok(_) => panic!("a locked start should not create a second runtime"), + Err(error) => error, + }; + assert_eq!(error.code, "auth_state_locked"); + assert!( + !state_dir.join("admin.key").exists(), + "a locked start must not create a competing administrator key" + ); + drop(_lock); + tokio::time::sleep(Duration::from_millis(20)).await; + let _ = std::fs::remove_dir_all(state_dir); +} + +#[test] +fn safe_mode_startup_does_not_allow_compaction() { + assert!(!compaction_is_allowed(true)); + assert!(compaction_is_allowed(false)); +} + +#[tokio::test] +async fn reset_clears_retained_high_slot_entries() { + let state_dir = temp_state_dir("reset-high-slots"); + let admin_key = *b"0123456789abcdefghijklmnopqrstuv"; + let config_two = AuthConfig { + state_dir: state_dir.clone(), + max_temporary_keys: 2, + max_temporary_key_ttl: Duration::from_secs(3600), + legacy_protocol: LegacyProtocolPolicy::Allow, + }; + let runtime = AuthRuntime::start(admin_key, config_two.clone()) + .await + .unwrap(); + let admin = authenticate_for_test(&runtime, ADMIN_KEY_ID).unwrap(); + let first = runtime + .issue(&admin, Duration::from_secs(60), Some("first".to_string())) + .await + .unwrap(); + let second = runtime + .issue(&admin, Duration::from_secs(60), Some("second".to_string())) + .await + .unwrap(); + drop(runtime); + tokio::time::sleep(Duration::from_millis(20)).await; + + let config_one = AuthConfig { + max_temporary_keys: 1, + ..config_two.clone() + }; + let runtime = AuthRuntime::start(admin_key, config_one).await.unwrap(); + let admin = authenticate_for_test(&runtime, ADMIN_KEY_ID).unwrap(); + runtime.reset(&admin).await.unwrap(); + drop(runtime); + tokio::time::sleep(Duration::from_millis(20)).await; + + let runtime = AuthRuntime::start(admin_key, config_two).await.unwrap(); + let admin = authenticate_for_test(&runtime, ADMIN_KEY_ID).unwrap(); + let page = runtime.list(&admin, 0, 100).await.unwrap(); + assert!(page.items.is_empty()); + assert!(authenticate_for_test(&runtime, first.metadata.key_id).is_err()); + assert!(authenticate_for_test(&runtime, second.metadata.key_id).is_err()); + drop(runtime); + tokio::time::sleep(Duration::from_millis(20)).await; + let _ = std::fs::remove_dir_all(state_dir); +} + +#[tokio::test] +async fn rotate_root_rejects_a_nul_containing_key() { + let state_dir = temp_state_dir("rotate-nul"); + let admin_key = *b"0123456789abcdefghijklmnopqrstuv"; + let config = AuthConfig { + state_dir: state_dir.clone(), + max_temporary_keys: 4, + max_temporary_key_ttl: Duration::from_secs(3600), + legacy_protocol: LegacyProtocolPolicy::Allow, + }; + let runtime = AuthRuntime::start(admin_key, config).await.unwrap(); + let admin = authenticate_for_test(&runtime, ADMIN_KEY_ID).unwrap(); + let mut bad = *b"0123456789abcdefghijklmnopqrstuv"; + bad[4] = 0; + let error = runtime.rotate_root(&admin, bad).await.unwrap_err(); + assert_eq!(error.code, "administrator_key_invalid"); + drop(runtime); + tokio::time::sleep(Duration::from_millis(20)).await; + let _ = std::fs::remove_dir_all(state_dir); +} + +#[test] +fn platform_default_auth_state_dir_is_writable_outside_linux_system_paths() { + let dir = platform_default_auth_state_dir(); + #[cfg(windows)] + { + assert!( + dir.ends_with(std::path::Path::new("pb-mapper").join("auth")), + "windows default auth dir should be under a user-writable pb-mapper path: {}", + dir.display() + ); + assert_ne!(dir, PathBuf::from(r"\var\lib\pb-mapper\auth")); + } + #[cfg(target_os = "macos")] + { + assert!( + dir.ends_with("Library/Application Support/pb-mapper/auth") + || dir == PathBuf::from("/Library/Application Support/pb-mapper/auth"), + "macos default auth dir should be under Application Support: {}", + dir.display() + ); + } + #[cfg(not(any(windows, target_os = "macos")))] + { + let expected = linux_default_auth_state_dir( + unix_effective_uid(), + linux_system_auth_dir_usable(), + std::env::var_os("XDG_DATA_HOME").as_deref(), + std::env::var_os("HOME").as_deref(), + ); + assert_eq!(dir, expected); + if unix_effective_uid() != 0 && !linux_system_auth_dir_usable() { + assert_ne!( + dir, + PathBuf::from(DEFAULT_AUTH_STATE_DIR), + "unprivileged Linux should not default to the system auth directory: {}", + dir.display() + ); + } + } +} + +#[test] +fn sync_parent_directory_succeeds_for_a_local_file() { + let state_dir = temp_state_dir("dirsync"); + prepare_state_dir(&state_dir).unwrap(); + let path = state_dir.join("probe"); + std::fs::write(&path, b"x").unwrap(); + sync_parent_directory(&path).unwrap(); + let _ = std::fs::remove_dir_all(state_dir); +} + +#[cfg(not(any(windows, target_os = "macos")))] +#[test] +fn linux_default_auth_state_dir_prefers_user_data_when_system_dir_is_unusable() { + assert_eq!( + linux_default_auth_state_dir(0, false, None, Some(std::ffi::OsStr::new("/home/op"))), + PathBuf::from(DEFAULT_AUTH_STATE_DIR) + ); + assert_eq!( + linux_default_auth_state_dir(1000, true, None, Some(std::ffi::OsStr::new("/home/op"))), + PathBuf::from(DEFAULT_AUTH_STATE_DIR) + ); + assert_eq!( + linux_default_auth_state_dir( + 1000, + false, + Some(std::ffi::OsStr::new("/xdg")), + Some(std::ffi::OsStr::new("/home/op")) + ), + PathBuf::from("/xdg/pb-mapper/auth") + ); + assert_eq!( + linux_default_auth_state_dir(1000, false, None, Some(std::ffi::OsStr::new("/home/op"))), + PathBuf::from("/home/op/.local/share/pb-mapper/auth") + ); + assert_eq!( + linux_default_auth_state_dir(1000, false, None, None), + PathBuf::from(DEFAULT_AUTH_STATE_DIR) + ); +} + +#[test] +fn legacy_protocol_policy_trims_valid_values_and_rejects_unknown_values() { + assert_eq!( + parse_legacy_protocol_policy(" allow\n"), + Some(LegacyProtocolPolicy::Allow) + ); + assert_eq!( + parse_legacy_protocol_policy(" DENY "), + Some(LegacyProtocolPolicy::Deny) + ); + assert_eq!(parse_legacy_protocol_policy("enabled"), None); + assert_eq!(parse_legacy_protocol_policy(""), None); +} + +#[test] +fn key_id_serializes_as_a_plain_integer() { + let key_id = KeyId::new(Generation::from_u32(3), SlotIndex::from_index(2)); + assert_eq!(serde_json::to_string(&key_id).unwrap(), "12884901890"); + assert_eq!( + serde_json::from_str::("12884901890").unwrap(), + key_id + ); + assert_eq!( + serde_json::to_string(&Generation::from_u32(7)).unwrap(), + "7" + ); +} + +#[test] +fn key_id_round_trip() { + let key_id = KeyId::new(Generation::from_u32(42), SlotIndex::from_index(65_535)); + assert_eq!(key_id.generation(), Generation::from_u32(42)); + assert_eq!(key_id.slot(), SlotIndex::from_index(65_535)); +} + +#[test] +fn derived_key_is_bound_to_instance_and_key_id() { + let admin_key = *b"0123456789abcdefghijklmnopqrstuv"; + let instance_a = [1_u8; INSTANCE_ID_LEN]; + let instance_b = [2_u8; INSTANCE_ID_LEN]; + let key = derive_temporary_key( + &admin_key, + &instance_a, + KeyId::new(Generation::from_u32(1), SlotIndex::from_index(7)), + ) + .unwrap(); + assert_eq!( + key, + derive_temporary_key( + &admin_key, + &instance_a, + KeyId::new(Generation::from_u32(1), SlotIndex::from_index(7)) + ) + .unwrap() + ); + assert_ne!( + key, + derive_temporary_key( + &admin_key, + &instance_b, + KeyId::new(Generation::from_u32(1), SlotIndex::from_index(7)) + ) + .unwrap() + ); + assert_ne!( + key, + derive_temporary_key( + &admin_key, + &instance_a, + KeyId::new(Generation::from_u32(2), SlotIndex::from_index(7)) + ) + .unwrap() + ); +} + +#[tokio::test] +async fn isolated_runtime_preserves_remote_temporary_process_credential() { + let _process_credential_guard = PROCESS_CREDENTIAL_TEST_LOCK.lock().await; + let state_dir = temp_state_dir("isolated-relay"); + let temporary_key_id = KeyId::new(Generation::from_u32(1), SlotIndex::from_index(0)); + let temporary_key = *b"temporary-remote-key-0123456789a"; + let temporary_credential = + encode_temporary_credential(temporary_key_id.as_u64(), &temporary_key); + set_process_msg_header_key(Some(&temporary_credential)).unwrap(); + let config = AuthConfig { + state_dir: state_dir.clone(), + max_temporary_keys: 4, + max_temporary_key_ttl: Duration::from_secs(3600), + legacy_protocol: LegacyProtocolPolicy::Allow, + }; + + let runtime = AuthRuntime::from_isolated_state(config).await.unwrap(); + assert_eq!( + get_process_credential().unwrap(), + Credential::Temporary { + key_id: temporary_key_id.as_u64(), + key: temporary_key, + } + ); + + let local_admin_raw = std::fs::read_to_string(state_dir.join("admin.key")).unwrap(); + let Credential::Admin(local_admin_key) = parse_credential(local_admin_raw.trim()).unwrap() + else { + panic!("isolated relay key should be an administrator credential"); + }; + let local_admin = runtime + .authenticate_presented(ADMIN_KEY_ID, &local_admin_key) + .unwrap(); + runtime + .rotate_root(&local_admin, *b"isolated-new-admin-key-012345678") + .await + .unwrap(); + assert_eq!( + get_process_credential().unwrap(), + Credential::Temporary { + key_id: temporary_key_id.as_u64(), + key: temporary_key, + } + ); + + set_process_msg_header_key(None).unwrap(); + drop(runtime); + tokio::time::sleep(Duration::from_millis(20)).await; + let _ = std::fs::remove_dir_all(state_dir); +} + +#[tokio::test] +async fn issue_renew_revoke_and_persist() { + let state_dir = temp_state_dir("auth-lifecycle"); + let admin_key = *b"0123456789abcdefghijklmnopqrstuv"; + let config = AuthConfig { + state_dir: state_dir.clone(), + max_temporary_keys: 8, + max_temporary_key_ttl: Duration::from_secs(3600), + legacy_protocol: LegacyProtocolPolicy::Allow, + }; + let runtime = AuthRuntime::start(admin_key, config.clone()).await.unwrap(); + let admin = authenticate_for_test(&runtime, ADMIN_KEY_ID).unwrap(); + let issued = runtime + .issue(&admin, Duration::from_secs(60), Some("demo".to_string())) + .await + .unwrap(); + assert!(issued.credential.starts_with("pbmt1_")); + let context = authenticate_for_test(&runtime, issued.metadata.key_id).unwrap(); + assert!(!context.is_admin); + let cancellation = context.cancellation_token().unwrap(); + let renewed = runtime + .renew(&admin, issued.metadata.key_id, Duration::from_secs(120)) + .await + .unwrap(); + assert_eq!(renewed.metadata.key_id, issued.metadata.key_id); + assert_eq!(renewed.credential, issued.credential); + assert!(renewed.metadata.expires_at > issued.metadata.expires_at); + let presented = runtime.derive_key(issued.metadata.key_id).unwrap(); + runtime + .revoke(&admin, issued.metadata.key_id) + .await + .unwrap(); + let mut mistyped = presented; + mistyped[0] ^= 0x01; + assert_eq!( + runtime + .authenticate_presented(issued.metadata.key_id, &mistyped) + .unwrap_err() + .code, + "temporary_key_invalid" + ); + assert!(cancellation.is_cancelled()); + assert_eq!( + context.ensure_active().unwrap_err().code, + "temporary_key_revoked" + ); + assert_eq!( + authenticate_for_test(&runtime, issued.metadata.key_id) + .unwrap_err() + .code, + "temporary_key_revoked" + ); + let instance_id = load_or_create_instance_id(&state_dir).unwrap(); + let persisted = try_load_persisted_state(&config, &admin_key, instance_id).unwrap(); + let revoked = persisted + .entries + .iter() + .find(|entry| entry.key_id == issued.metadata.key_id) + .unwrap(); + assert_eq!(revoked.state, SlotState::Revoked); + assert!(revoked.tombstoned_at.is_some()); + drop(runtime); + + tokio::time::sleep(Duration::from_millis(20)).await; + let restored = AuthRuntime::start(admin_key, config).await.unwrap(); + assert_eq!( + authenticate_for_test(&restored, issued.metadata.key_id) + .unwrap_err() + .code, + "temporary_key_revoked" + ); + let _ = std::fs::remove_dir_all(state_dir); +} + +#[tokio::test] +async fn ensure_active_keeps_expiry_after_the_lease_is_cancelled() { + let state_dir = temp_state_dir("lease-expiry-reason"); + let admin_key = *b"0123456789abcdefghijklmnopqrstuv"; + let config = AuthConfig { + state_dir: state_dir.clone(), + max_temporary_keys: 4, + max_temporary_key_ttl: Duration::from_secs(3600), + legacy_protocol: LegacyProtocolPolicy::Allow, + }; + let runtime = AuthRuntime::start(admin_key, config).await.unwrap(); + let admin = authenticate_for_test(&runtime, ADMIN_KEY_ID).unwrap(); + let issued = runtime + .issue(&admin, Duration::from_secs(60), Some("exp".to_string())) + .await + .unwrap(); + let context = authenticate_for_test(&runtime, issued.metadata.key_id).unwrap(); + let lease = context.ensure_active().unwrap(); + lease.expire_now(); + assert_eq!( + context.ensure_active().unwrap_err().code, + "temporary_key_expired" + ); + assert_eq!( + context.ensure_active().unwrap_err().code, + "temporary_key_expired" + ); + drop(runtime); + tokio::time::sleep(Duration::from_millis(20)).await; + let _ = std::fs::remove_dir_all(state_dir); +} + +#[tokio::test] +async fn renew_replaces_a_lease_canceled_during_persistence() { + let state_dir = temp_state_dir("renew-canceled-lease"); + let admin_key = *b"0123456789abcdefghijklmnopqrstuv"; + let config = AuthConfig { + state_dir: state_dir.clone(), + max_temporary_keys: 4, + max_temporary_key_ttl: Duration::from_secs(3600), + legacy_protocol: LegacyProtocolPolicy::Allow, + }; + let runtime = AuthRuntime::start(admin_key, config).await.unwrap(); + let admin = authenticate_for_test(&runtime, ADMIN_KEY_ID).unwrap(); + let issued = runtime + .issue(&admin, Duration::from_secs(60), Some("renew".to_string())) + .await + .unwrap(); + let context = authenticate_for_test(&runtime, issued.metadata.key_id).unwrap(); + let canceled = context.cancellation_token().unwrap(); + canceled.cancel(); + assert!(canceled.is_cancelled()); + + let renewed = runtime + .renew(&admin, issued.metadata.key_id, Duration::from_secs(120)) + .await + .unwrap(); + assert_eq!(renewed.metadata.key_id, issued.metadata.key_id); + let restored = authenticate_for_test(&runtime, issued.metadata.key_id).unwrap(); + assert!(!restored.cancellation_token().unwrap().is_cancelled()); + assert!(canceled.is_cancelled()); + + drop(runtime); + tokio::time::sleep(Duration::from_millis(20)).await; + let _ = std::fs::remove_dir_all(state_dir); +} + +#[tokio::test] +async fn reset_rotates_instance_and_prevents_old_key_id_reuse() { + let state_dir = temp_state_dir("auth-reset"); + let admin_key = *b"0123456789abcdefghijklmnopqrstuv"; + let config = AuthConfig { + state_dir: state_dir.clone(), + max_temporary_keys: 1, + max_temporary_key_ttl: Duration::from_secs(3600), + legacy_protocol: LegacyProtocolPolicy::Allow, + }; + let runtime = AuthRuntime::start(admin_key, config).await.unwrap(); + let admin = authenticate_for_test(&runtime, ADMIN_KEY_ID).unwrap(); + let before = runtime.status(&admin).await.unwrap().server_instance_id; + let old = runtime + .issue( + &admin, + Duration::from_secs(60), + Some("before-reset".to_string()), + ) + .await + .unwrap(); + let old_context = authenticate_for_test(&runtime, old.metadata.key_id).unwrap(); + let old_cancellation = old_context.cancellation_token().unwrap(); + let old_presented = runtime.derive_key(old.metadata.key_id).unwrap(); + + runtime.reset(&admin).await.unwrap(); + + let after = runtime.status(&admin).await.unwrap().server_instance_id; + assert_ne!(after, before); + assert!(old_cancellation.is_cancelled()); + assert_eq!( + runtime + .authenticate_presented(old.metadata.key_id, &old_presented) + .unwrap_err() + .code, + "temporary_key_rotated" + ); + let replacement = runtime + .issue( + &admin, + Duration::from_secs(60), + Some("after-reset".to_string()), + ) + .await + .unwrap(); + assert_ne!(replacement.metadata.key_id, old.metadata.key_id); + assert_ne!(replacement.credential, old.credential); + + drop(runtime); + tokio::time::sleep(Duration::from_millis(20)).await; + let _ = std::fs::remove_dir_all(state_dir); +} + +#[test] +fn recover_instance_id_promotes_next_when_snapshot_matches() { + let state_dir = temp_state_dir("instance-next-promote"); + prepare_state_dir(&state_dir).unwrap(); + let admin_key = *b"0123456789abcdefghijklmnopqrstuv"; + let current = [1_u8; INSTANCE_ID_LEN]; + let next = [2_u8; INSTANCE_ID_LEN]; + atomic_write(&state_dir.join("server-instance-id"), ¤t, 0o600).unwrap(); + atomic_write(&state_dir.join("server-instance-id.next"), &next, 0o600).unwrap(); + let config = AuthConfig { + state_dir: state_dir.clone(), + max_temporary_keys: 1, + max_temporary_key_ttl: Duration::from_secs(3600), + legacy_protocol: LegacyProtocolPolicy::Allow, + }; + let snapshot = PersistedSnapshot { + schema_version: SNAPSHOT_SCHEMA_VERSION, + instance_id: next, + generations: vec![Generation::FIRST; 1], + entries: Vec::new(), + legacy_protocol: LegacyProtocolPolicy::Allow, + admin_replays: Vec::new(), + audit_records: VecDeque::new(), + root_epoch: 0, + }; + write_snapshot_and_truncate_wal(&config, &admin_key, &snapshot).unwrap(); + std::fs::write(state_dir.join("auth.wal"), b"old-instance-wal").unwrap(); + + let recovered = recover_instance_id_after_reset(&state_dir, &admin_key, current).unwrap(); + assert_eq!(recovered, next); + assert_eq!( + read_instance_id_file(&state_dir.join("server-instance-id")).unwrap(), + Some(next) + ); + assert!(!state_dir.join("server-instance-id.next").exists()); + assert_eq!(std::fs::read(state_dir.join("auth.wal")).unwrap(), b""); + let _ = std::fs::remove_dir_all(state_dir); +} + +#[test] +fn reset_already_installed_accepts_matching_live_id_and_snapshot() { + let state_dir = temp_state_dir("reset-already-installed"); + prepare_state_dir(&state_dir).unwrap(); + let admin_key = *b"0123456789abcdefghijklmnopqrstuv"; + let new_id = [9_u8; INSTANCE_ID_LEN]; + atomic_write(&state_dir.join("server-instance-id"), &new_id, 0o600).unwrap(); + let config = AuthConfig { + state_dir: state_dir.clone(), + max_temporary_keys: 1, + max_temporary_key_ttl: Duration::from_secs(3600), + legacy_protocol: LegacyProtocolPolicy::Allow, + }; + let snapshot = PersistedSnapshot { + schema_version: SNAPSHOT_SCHEMA_VERSION, + instance_id: new_id, + generations: vec![Generation::FIRST; 1], + entries: Vec::new(), + legacy_protocol: LegacyProtocolPolicy::Allow, + admin_replays: Vec::new(), + audit_records: VecDeque::new(), + root_epoch: 1, + }; + write_snapshot_and_truncate_wal(&config, &admin_key, &snapshot).unwrap(); + assert!(reset_already_installed(&state_dir, &admin_key, &new_id)); + assert!(!reset_already_installed( + &state_dir, + &admin_key, + &[8_u8; INSTANCE_ID_LEN] + )); + let _ = std::fs::remove_dir_all(state_dir); +} + +#[test] +fn recover_instance_id_discards_stale_next_when_snapshot_still_matches_current() { + let state_dir = temp_state_dir("instance-next-stale"); + prepare_state_dir(&state_dir).unwrap(); + let admin_key = *b"0123456789abcdefghijklmnopqrstuv"; + let current = [3_u8; INSTANCE_ID_LEN]; + let next = [4_u8; INSTANCE_ID_LEN]; + atomic_write(&state_dir.join("server-instance-id"), ¤t, 0o600).unwrap(); + atomic_write(&state_dir.join("server-instance-id.next"), &next, 0o600).unwrap(); + let config = AuthConfig { + state_dir: state_dir.clone(), + max_temporary_keys: 1, + max_temporary_key_ttl: Duration::from_secs(3600), + legacy_protocol: LegacyProtocolPolicy::Allow, + }; + let snapshot = PersistedSnapshot { + schema_version: SNAPSHOT_SCHEMA_VERSION, + instance_id: current, + generations: vec![Generation::FIRST; 1], + entries: Vec::new(), + legacy_protocol: LegacyProtocolPolicy::Allow, + admin_replays: Vec::new(), + audit_records: VecDeque::new(), + root_epoch: 0, + }; + write_snapshot_and_truncate_wal(&config, &admin_key, &snapshot).unwrap(); + + let recovered = recover_instance_id_after_reset(&state_dir, &admin_key, current).unwrap(); + assert_eq!(recovered, current); + assert!(!state_dir.join("server-instance-id.next").exists()); + let _ = std::fs::remove_dir_all(state_dir); +} + +#[test] +fn recover_admin_key_discards_leftover_wal_from_the_old_key() { + let state_dir = temp_state_dir("admin-next-wal"); + prepare_state_dir(&state_dir).unwrap(); + let old_key = *b"0123456789abcdefghijklmnopqrstuv"; + let new_key = *b"abcdefghijklmnopqrstuvwxyz012345"; + let old_key_str = std::str::from_utf8(&old_key).unwrap(); + let new_key_str = std::str::from_utf8(&new_key).unwrap(); + write_admin_key(&state_dir, old_key_str).unwrap(); + write_admin_key_file(&state_dir.join("admin.key.next"), new_key_str, true).unwrap(); + let config = AuthConfig { + state_dir: state_dir.clone(), + max_temporary_keys: 1, + max_temporary_key_ttl: Duration::from_secs(3600), + legacy_protocol: LegacyProtocolPolicy::Allow, + }; + let snapshot = PersistedSnapshot { + schema_version: SNAPSHOT_SCHEMA_VERSION, + instance_id: [9_u8; INSTANCE_ID_LEN], + generations: vec![Generation::FIRST; 1], + entries: Vec::new(), + legacy_protocol: LegacyProtocolPolicy::Allow, + admin_replays: Vec::new(), + audit_records: VecDeque::new(), + root_epoch: 0, + }; + write_snapshot_and_truncate_wal(&config, &new_key, &snapshot).unwrap(); + std::fs::write(state_dir.join("auth.wal"), b"old-key-wal").unwrap(); + + let recovered = recover_admin_key_after_rotation(&state_dir, old_key_str).unwrap(); + assert_eq!(recovered.trim(), new_key_str); + assert_eq!(std::fs::read(state_dir.join("auth.wal")).unwrap(), b""); + assert!(!state_dir.join("admin.key.next").exists()); + let _ = std::fs::remove_dir_all(state_dir); +} + +#[test] +fn rotation_finalize_requires_the_live_admin_key() { + let state_dir = temp_state_dir("rotate-requires-live-key"); + prepare_state_dir(&state_dir).unwrap(); + let old_key = *b"0123456789abcdefghijklmnopqrstuv"; + let new_key = *b"abcdefghijklmnopqrstuvwxyz012345"; + let old_key_str = std::str::from_utf8(&old_key).unwrap(); + let new_key_str = std::str::from_utf8(&new_key).unwrap(); + write_admin_key(&state_dir, old_key_str).unwrap(); + let config = AuthConfig { + state_dir: state_dir.clone(), + max_temporary_keys: 1, + max_temporary_key_ttl: Duration::from_secs(3600), + legacy_protocol: LegacyProtocolPolicy::Allow, + }; + let snapshot = PersistedSnapshot { + schema_version: SNAPSHOT_SCHEMA_VERSION, + instance_id: [3_u8; INSTANCE_ID_LEN], + generations: vec![Generation::FIRST; 1], + entries: Vec::new(), + legacy_protocol: LegacyProtocolPolicy::Allow, + admin_replays: Vec::new(), + audit_records: VecDeque::new(), + root_epoch: 1, + }; + write_snapshot_and_truncate_wal(&config, &new_key, &snapshot).unwrap(); + assert!(!rotation_already_installed(&state_dir, new_key_str)); + write_admin_key(&state_dir, new_key_str).unwrap(); + assert!(rotation_already_installed(&state_dir, new_key_str)); + let _ = std::fs::remove_dir_all(state_dir); +} + +#[tokio::test] +async fn interrupted_reset_recovers_the_staged_instance_id_on_restart() { + let state_dir = temp_state_dir("reset-recover"); + let admin_key = *b"0123456789abcdefghijklmnopqrstuv"; + let config = AuthConfig { + state_dir: state_dir.clone(), + max_temporary_keys: 4, + max_temporary_key_ttl: Duration::from_secs(3600), + legacy_protocol: LegacyProtocolPolicy::Allow, + }; + let runtime = AuthRuntime::start(admin_key, config.clone()).await.unwrap(); + let admin = authenticate_for_test(&runtime, ADMIN_KEY_ID).unwrap(); + let issued = runtime + .issue( + &admin, + Duration::from_secs(60), + Some("before-interrupted-reset".to_string()), + ) + .await + .unwrap(); + drop(runtime); + tokio::time::sleep(Duration::from_millis(20)).await; + + let old_instance_id = load_or_create_instance_id(&state_dir).unwrap(); + let next = random_instance_id(); + atomic_write(&state_dir.join("server-instance-id.next"), &next, 0o600).unwrap(); + let snapshot = PersistedSnapshot { + schema_version: SNAPSHOT_SCHEMA_VERSION, + instance_id: next, + generations: vec![Generation::FIRST; 4], + entries: Vec::new(), + legacy_protocol: LegacyProtocolPolicy::Allow, + admin_replays: Vec::new(), + audit_records: VecDeque::new(), + root_epoch: 0, + }; + write_snapshot_and_truncate_wal(&config, &admin_key, &snapshot).unwrap(); + std::fs::write(state_dir.join("auth.wal"), b"old-instance-wal").unwrap(); + atomic_write( + &state_dir.join("server-instance-id"), + &old_instance_id, + 0o600, + ) + .unwrap(); + + let restored = AuthRuntime::start(admin_key, config).await.unwrap(); + let restored_admin = authenticate_for_test(&restored, ADMIN_KEY_ID).unwrap(); + let status = restored.status(&restored_admin).await.unwrap(); + assert!(!status.safe_mode); + assert_eq!(status.server_instance_id, hex(&next)); + assert!(authenticate_for_test(&restored, issued.metadata.key_id).is_err()); + assert!(!state_dir.join("server-instance-id.next").exists()); + drop(restored); + tokio::time::sleep(Duration::from_millis(20)).await; + let _ = std::fs::remove_dir_all(state_dir); +} + +#[tokio::test] +async fn corrupt_wal_fails_temporary_keys_closed_until_admin_reset() { + let state_dir = temp_state_dir("auth-safe-mode"); + let admin_key = *b"0123456789abcdefghijklmnopqrstuv"; + let config = AuthConfig { + state_dir: state_dir.clone(), + max_temporary_keys: 4, + max_temporary_key_ttl: Duration::from_secs(3600), + legacy_protocol: LegacyProtocolPolicy::Allow, + }; + let runtime = AuthRuntime::start(admin_key, config.clone()).await.unwrap(); + let admin = authenticate_for_test(&runtime, ADMIN_KEY_ID).unwrap(); + let issued = runtime + .issue( + &admin, + Duration::from_secs(60), + Some("corrupt-me".to_string()), + ) + .await + .unwrap(); + drop(runtime); + tokio::time::sleep(Duration::from_millis(20)).await; + std::fs::write(state_dir.join("auth.wal"), b"broken-wal").unwrap(); + + let recovered = AuthRuntime::start(admin_key, config).await.unwrap(); + let recovered_admin = authenticate_for_test(&recovered, ADMIN_KEY_ID).unwrap(); + assert!(recovered.status(&recovered_admin).await.unwrap().safe_mode); + assert_eq!( + authenticate_for_test(&recovered, issued.metadata.key_id) + .unwrap_err() + .code, + "temporary_key_store_unavailable" + ); + recovered.reset(&recovered_admin).await.unwrap(); + assert!(!recovered.status(&recovered_admin).await.unwrap().safe_mode); + + drop(recovered); + tokio::time::sleep(Duration::from_millis(20)).await; + let _ = std::fs::remove_dir_all(state_dir); +} + +#[tokio::test] +async fn root_rotation_rejects_old_key_and_in_flight_admin_context() { + let _process_credential_guard = PROCESS_CREDENTIAL_TEST_LOCK.lock().await; + let state_dir = temp_state_dir("auth-root-rotation"); + let old_key = *b"0123456789abcdefghijklmnopqrstuv"; + let new_key = *b"abcdefghijklmnopqrstuvwxyz012345"; + let config = AuthConfig { + state_dir: state_dir.clone(), + max_temporary_keys: 4, + max_temporary_key_ttl: Duration::from_secs(3600), + legacy_protocol: LegacyProtocolPolicy::Allow, + }; + let runtime = AuthRuntime::start(old_key, config).await.unwrap(); + let old_admin = runtime + .authenticate_presented(ADMIN_KEY_ID, &old_key) + .unwrap(); + let issued = runtime + .issue( + &old_admin, + Duration::from_secs(60), + Some("before-rotate".to_string()), + ) + .await + .unwrap(); + let old_temporary = runtime.derive_key(issued.metadata.key_id).unwrap(); + let mut mistyped_temporary = old_temporary; + mistyped_temporary[0] ^= 0x01; + assert_eq!( + runtime + .authenticate_presented(issued.metadata.key_id, &mistyped_temporary) + .unwrap_err() + .code, + "temporary_key_invalid" + ); + let mistyped_key = *b"1123456789abcdefghijklmnopqrstuv"; + assert_eq!( + runtime + .authenticate_presented(ADMIN_KEY_ID, &mistyped_key) + .unwrap_err() + .code, + "administrator_key_invalid" + ); + + runtime + .rotate_root(&old_admin, new_key) + .await + .expect("root rotation should succeed"); + assert_eq!( + runtime + .authenticate_presented(issued.metadata.key_id, &old_temporary) + .unwrap_err() + .code, + "temporary_key_rotated" + ); + let new_admin = runtime + .authenticate_presented(ADMIN_KEY_ID, &new_key) + .unwrap(); + let _replacement = runtime + .issue( + &new_admin, + Duration::from_secs(60), + Some("after-rotate".to_string()), + ) + .await + .unwrap(); + assert_eq!( + runtime + .authenticate_presented(issued.metadata.key_id, &old_temporary) + .unwrap_err() + .code, + "temporary_key_rotated" + ); + + assert_eq!( + runtime + .authenticate_presented(ADMIN_KEY_ID, &old_key) + .unwrap_err() + .code, + "administrator_key_invalid" + ); + assert_eq!( + runtime + .issue(&old_admin, Duration::from_secs(60), None) + .await + .unwrap_err() + .code, + "administrator_key_rotated" + ); + assert!(runtime.status(&new_admin).await.is_ok()); + + drop(runtime); + tokio::time::sleep(Duration::from_millis(20)).await; + let _ = std::fs::remove_dir_all(state_dir); +} + +#[tokio::test] +async fn admitted_admin_mutation_replay_survives_restart() { + let state_dir = temp_state_dir("admin-replay-restart"); + let admin_key = *b"0123456789abcdefghijklmnopqrstuv"; + let config = AuthConfig { + state_dir: state_dir.clone(), + max_temporary_keys: 4, + max_temporary_key_ttl: Duration::from_secs(3600), + legacy_protocol: LegacyProtocolPolicy::Allow, + }; + let fingerprint = [0x5a; 32]; + let timestamp = unix_seconds(); + let runtime = AuthRuntime::start(admin_key, config.clone()).await.unwrap(); + let admin = authenticate_for_test(&runtime, ADMIN_KEY_ID).unwrap(); + runtime + .claim_admin_mutation(&admin, fingerprint, timestamp) + .await + .unwrap(); + drop(runtime); + tokio::time::sleep(Duration::from_millis(20)).await; + + let restored = AuthRuntime::start(admin_key, config).await.unwrap(); + let restored_admin = authenticate_for_test(&restored, ADMIN_KEY_ID).unwrap(); + assert_eq!( + restored + .claim_admin_mutation(&restored_admin, fingerprint, timestamp) + .await + .unwrap_err() + .code, + "admin_request_replayed" + ); + + drop(restored); + tokio::time::sleep(Duration::from_millis(20)).await; + let _ = std::fs::remove_dir_all(state_dir); +} + +#[tokio::test] +async fn snapshot_compaction_preserves_audit_records() { + let state_dir = temp_state_dir("audit-compaction"); + let admin_key = *b"0123456789abcdefghijklmnopqrstuv"; + let config = AuthConfig { + state_dir: state_dir.clone(), + max_temporary_keys: 4, + max_temporary_key_ttl: Duration::from_secs(3600), + legacy_protocol: LegacyProtocolPolicy::Allow, + }; + let runtime = AuthRuntime::start(admin_key, config.clone()).await.unwrap(); + let admin = authenticate_for_test(&runtime, ADMIN_KEY_ID).unwrap(); + runtime + .issue(&admin, Duration::from_secs(60), Some("audited".to_string())) + .await + .unwrap(); + runtime.gc(&admin).await.unwrap(); + + let instance_id = load_or_create_instance_id(&state_dir).unwrap(); + let persisted = try_load_persisted_state(&config, &admin_key, instance_id).unwrap(); + let actions = persisted + .audit_records + .iter() + .map(|record| record.action.as_str()) + .collect::>(); + assert!(actions.contains(&"temporary_key_issue")); + assert!(actions.contains(&"temporary_key_gc")); + + drop(runtime); + tokio::time::sleep(Duration::from_millis(20)).await; + let _ = std::fs::remove_dir_all(state_dir); +} + +#[tokio::test] +async fn revoking_keeps_the_row_until_its_retention_elapses() { + let state_dir = temp_state_dir("revoke-retention"); + let admin_key = *b"0123456789abcdefghijklmnopqrstuv"; + let config = AuthConfig { + state_dir: state_dir.clone(), + max_temporary_keys: 4, + max_temporary_key_ttl: Duration::from_secs(3600), + legacy_protocol: LegacyProtocolPolicy::Allow, + }; + let runtime = AuthRuntime::start(admin_key, config).await.unwrap(); + let admin = authenticate_for_test(&runtime, ADMIN_KEY_ID).unwrap(); + let issued = runtime + .issue(&admin, Duration::from_secs(60), Some("revoked".to_string())) + .await + .unwrap(); + let key_id = issued.metadata.key_id; + let presented = runtime.derive_key(key_id).unwrap(); + + runtime.revoke(&admin, key_id).await.unwrap(); + + // The credential stops working at once, but the row survives so the reason + // is still reportable rather than degrading to "unknown key". + assert_eq!( + runtime + .authenticate_presented(key_id, &presented) + .unwrap_err() + .code, + "temporary_key_revoked" + ); + assert!( + runtime + .list(&admin, 0, 100) + .await + .unwrap() + .items + .iter() + .any(|item| item.key_id == key_id && item.state == "revoked") + ); + + drop(runtime); + tokio::time::sleep(Duration::from_millis(20)).await; + let _ = std::fs::remove_dir_all(state_dir); +} + +#[test] +fn replay_pruning_removes_only_records_outside_the_retention_window() { + let now = 10_000; + let expired = AdminReplayRecord { + fingerprint: [1; 32], + client_timestamp: now - ADMIN_REPLAY_RETENTION.as_secs() - 1, + accepted_at: now - ADMIN_REPLAY_RETENTION.as_secs() - 1, + }; + let current = AdminReplayRecord { + fingerprint: [2; 32], + client_timestamp: now, + accepted_at: now, + }; + let mut replay_set = HashSet::from([expired.fingerprint, current.fingerprint]); + let mut replay_order = VecDeque::from([expired, current.clone()]); + + super::actor::prune_expired_admin_replays(now, &mut replay_set, &mut replay_order); + + assert_eq!(replay_set, HashSet::from([current.fingerprint])); + assert_eq!(replay_order.len(), 1); + assert_eq!(replay_order[0].fingerprint, current.fingerprint); +} + +#[test] +fn replay_pruning_uses_server_acceptance_not_client_timestamp() { + let now = 10_000; + let retention = ADMIN_REPLAY_RETENTION.as_secs(); + let backdated = AdminReplayRecord { + fingerprint: [3; 32], + client_timestamp: now - retention - 1, + accepted_at: now - 1, + }; + let future_dated_but_expired = AdminReplayRecord { + fingerprint: [4; 32], + client_timestamp: now + retention / 2, + accepted_at: now - retention - 1, + }; + let mut replay_set = + HashSet::from([backdated.fingerprint, future_dated_but_expired.fingerprint]); + let mut replay_order = VecDeque::from([backdated.clone(), future_dated_but_expired]); + + super::actor::prune_expired_admin_replays(now, &mut replay_set, &mut replay_order); + + assert_eq!(replay_set, HashSet::from([backdated.fingerprint])); + assert_eq!(replay_order.len(), 1); + assert_eq!(replay_order[0].fingerprint, backdated.fingerprint); +} + +#[test] +fn replay_pruning_falls_back_to_client_timestamp_for_legacy_records() { + let now = 10_000; + let legacy_expired = AdminReplayRecord { + fingerprint: [5; 32], + client_timestamp: now - ADMIN_REPLAY_RETENTION.as_secs() - 1, + accepted_at: 0, + }; + let legacy_current = AdminReplayRecord { + fingerprint: [6; 32], + client_timestamp: now, + accepted_at: 0, + }; + let mut replay_set = HashSet::from([legacy_expired.fingerprint, legacy_current.fingerprint]); + let mut replay_order = VecDeque::from([legacy_expired, legacy_current.clone()]); + + super::actor::prune_expired_admin_replays(now, &mut replay_set, &mut replay_order); + + assert_eq!(replay_set, HashSet::from([legacy_current.fingerprint])); + assert_eq!(replay_order.len(), 1); + assert_eq!(replay_order[0].fingerprint, legacy_current.fingerprint); +} + +#[test] +fn tombstone_migration_prefers_audit_time_and_persists_fail_closed_fallback() { + let now = 10_000; + let revoked_with_audit = PersistedEntry { + key_id: KeyId::new(Generation::from_u32(1), SlotIndex::from_index(0)), + state: SlotState::Revoked, + issued_at: 100, + expires_at: 20_000, + label: None, + tombstoned_at: None, + }; + let revoked_without_audit = PersistedEntry { + key_id: KeyId::new(Generation::from_u32(1), SlotIndex::from_index(1)), + state: SlotState::Revoked, + issued_at: 100, + expires_at: 20_000, + label: None, + tombstoned_at: None, + }; + let audit_at = now - 30; + let mut snapshot = PersistedSnapshot { + schema_version: SNAPSHOT_SCHEMA_VERSION, + instance_id: [1; INSTANCE_ID_LEN], + generations: vec![Generation::from_u32(1), Generation::from_u32(1)], + entries: vec![revoked_with_audit, revoked_without_audit], + legacy_protocol: LegacyProtocolPolicy::Deny, + admin_replays: Vec::new(), + audit_records: VecDeque::from([AuditRecord { + at: audit_at, + action: "temporary_key_revoke".to_string(), + key_id: Some(KeyId::new( + Generation::from_u32(1), + SlotIndex::from_index(0), + )), + label: None, + }]), + root_epoch: 0, + }; + + assert!(normalize_tombstone_times(&mut snapshot, now)); + assert_eq!(snapshot.entries[0].tombstoned_at, Some(audit_at)); + assert_eq!(snapshot.entries[1].tombstoned_at, Some(now)); + assert!(!normalize_tombstone_times(&mut snapshot, now + 1)); + assert_eq!(snapshot.entries[1].tombstoned_at, Some(now)); +} diff --git a/crates/pb-mapper-auth/src/timing_wheel.rs b/crates/pb-mapper-auth/src/timing_wheel.rs new file mode 100644 index 0000000..aefcfb4 --- /dev/null +++ b/crates/pb-mapper-auth/src/timing_wheel.rs @@ -0,0 +1,444 @@ +//! Hierarchical timer wheel: rotating bucket queues indexed by relative delay. +//! +//! Levels are digit positions in base `radix`, and how many exist is derived from +//! the longest delay the wheel must support. Scheduling decomposes the delay into +//! those digits and builds one nested link per digit, coarsest outermost: +//! +//! ```text +//! radix = 64, schedule(delay = 1*64² + 5*64 + 3) +//! +//! level 2 [ ][A][ ]… A pops after 1 rotation of 64² ticks; dropping it +//! level 1 [ ]…[B][ ]… files B, which pops 5 rotations of 64 ticks later +//! level 0 [ ][ ][C]… and files C, which pops 3 ticks later and fires +//! ``` +//! +//! `A` holds `B` holds `C` holds the timer, so the chain *is* the route: no per +//! timer list of future placements, and nothing to look up or recompute. A bucket +//! is just `Vec`, and dropping a link is what files the next one. +//! +//! ```text +//! tick() -> ticks += 1 +//! -> level 0 always rotates; level i rotates when ticks % radix^i == 0 +//! -> rotate = pop_front, push_back an empty bucket; dropping the popped +//! bucket files each link's successor, or fires the timer if the link +//! was the innermost +//! ``` +//! +//! So a tick moves one bucket per level that turns over and performs no +//! arithmetic per entry: the queues rotate, which keeps a bucket's index equal to +//! its distance from now. +//! +//! Only the wheel holds strong references to a timer, so it runs when the last +//! chain holding it is dropped. The wheel never looks a timer up, compares +//! identities, or has to be told one was superseded: to move a deadline, schedule +//! the same timer again — the earlier chain still drains, but it is no longer the +//! last reference, so it fires nothing. + +use super::*; + +/// A callback that runs once: when its delay elapses, or when it is cancelled, +/// whichever comes first. +pub(super) struct Timer { + /// `None` once the callback has run, so any route still holding this timer is + /// inert and a cancelled timer cannot fire twice. + /// + /// WHY a lock for state a single task owns: running a `FnOnce` moves it out, + /// which needs `&mut`, but a timer is reached through a shared handle so that + /// two routes can hold one. `Arc: Send` — which `tokio::spawn` requires of + /// the actor this runs in — implies `T: Sync`, and shared mutability that is + /// `Sync` needs a lock; `Cell` would be cheaper but is not `Sync`. It is never + /// contended, and the path that fires almost every timer skips it: `Drop` has + /// `&mut self`, so it reaches the callback directly. + callback: Mutex>>, +} + +impl Timer { + pub(super) fn new(callback: impl FnOnce() + Send + 'static) -> Arc { + Arc::new(Self { + callback: Mutex::new(Some(Box::new(callback))), + }) + } + + /// Runs the callback unless it has run already, for a caller cancelling ahead + /// of the deadline. + pub(super) fn fire(&self) { + let callback = self.callback.lock().take(); + run(callback); + } +} + +impl Drop for Timer { + /// Releasing the last reference is what fires a timer, so dropping the wheel + /// runs everything it was holding. Owning `&mut self` here is what lets the + /// usual path take the callback without locking. + fn drop(&mut self) { + let callback = self.callback.get_mut().take(); + run(callback); + } +} + +fn run(callback: Option>) { + if let Some(callback) = callback { + callback(); + } +} + +/// One leg of a timer's route through the levels. +/// +/// A delay spanning several digits cannot be filed in one bucket, so the route is +/// a chain: each [`Link::Relay`] waits in one bucket and, once that bucket comes +/// off the front, hands the leg nested inside it to the wheel, which files it in +/// the next, finer bucket. Only the outermost leg is ever in a bucket, and only +/// [`Link::Deliver`] holds the timer, so the chain unwinding one bucket at a time +/// *is* the timer descending the levels. That is what leaves a tick with nothing +/// to compute. +/// +/// ```text +/// delay = 1*64² + 5*64 + 3 +/// Relay{L2,slot 1} -> Relay{L1,slot 5} -> Relay{L0,slot 3} -> Deliver(timer) +/// ^ filed now ^ filed when the ^ …and so on ^ dropping this +/// one before it fires the timer +/// comes off +/// ``` +enum Link { + /// The end of a route. Never read: holding the reference *is* this leg's job, + /// and releasing it is what fires the timer. + Deliver(#[allow(dead_code)] Arc), + /// Files `next` into `level`'s `slot` when this leg comes off the front. + /// + /// `Box`, not `Arc`: exactly one bucket owns a route at a time, so a leg needs + /// no reference count of its own — only the timer at the end is shared. + Relay { + level: u8, + slot: u16, + next: Box, + }, +} + +/// A bucket's worth of routes. Dropping one without draining it releases the +/// timers at the end of every route it holds, which is how dropping the wheel +/// fires everything. +type Bucket = Vec; + +pub(super) struct TimingWheel { + /// Ticks elapsed since construction. Buckets are indexed relative to it, so + /// advancing re-indexes nothing. + ticks: u64, + radix: u64, + levels: Vec>, +} + +impl TimingWheel { + /// Builds the smallest wheel that can place `max_delay` ticks, adding a level + /// at a time until the levels together span it. + pub(super) fn new(max_delay: u64, radix: u64) -> Self { + assert!(radix > 1, "a level needs at least two buckets"); + let mut levels = 1_usize; + let mut span = radix; + while span < max_delay { + levels += 1; + span = span.saturating_mul(radix); + } + Self { + ticks: 0, + radix, + levels: (0..levels) + .map(|_| { + std::iter::repeat_with(Bucket::new) + .take(radix as usize) + .collect() + }) + .collect(), + } + } + + /// Longest delay this wheel can place exactly. + pub(super) fn max_delay(&self) -> u64 { + self.period(self.levels.len()) + } + + /// Holds `timer` for `delay` ticks. A delay of zero, or one past + /// [`Self::max_delay`], releases the timer at once rather than misplacing it. + /// + /// Scheduling a timer the wheel already holds builds a second route rather + /// than replacing the first, which is how a caller moves a deadline without + /// the wheel having to find the old one. + pub(super) fn schedule(&mut self, delay: u64, timer: Arc) { + if delay == 0 || delay > self.max_delay() { + // Dropping `timer` here fires it if this was the last reference. + return; + } + let deliver = Link::Deliver(timer); + // The coarsest reachable level absorbs however far the current tick sits + // into its rotation, so its bucket comes off on a rotation boundary. Every + // finer level is at zero offset there, which makes the delay still + // remaining a plain base-`radix` decomposition from that point down. + // The guard above bounds `delay`, so some level always takes it. + // Unreachable, and treated like the out-of-range case: dropping the + // timer fires it, which is the safe direction for a credential deadline. + let Some((level, slot, remaining)) = (0..self.levels.len()) + .rev() + .find_map(|level| self.entry_leg(level, self.ticks + delay)) + else { + return; + }; + let route = match remaining { + 0 => deliver, + remaining => self.route(remaining, deliver), + }; + self.file(level, slot, route); + } + + /// Advances one tick, rotating every level that turns over. + pub(super) fn tick(&mut self) { + self.ticks += 1; + for level in 0..self.levels.len() { + // A level turns over every `radix^level` ticks. Once one does not, no + // coarser one can either, since its period divides theirs. + if !self.ticks.is_multiple_of(self.period(level)) { + break; + } + let bucket = self.rotate(level); + for link in bucket { + match link { + // Out of legs: dropping it releases the timer, firing it if + // this was the last route holding it. + Link::Deliver(timer) => drop(timer), + Link::Relay { level, slot, next } => { + self.file(level as usize, slot as usize, *next) + } + } + } + } + } + + /// Takes `level`'s front bucket off and puts an empty one on the back, so + /// bucket indices stay relative to the current tick. + fn rotate(&mut self, level: usize) -> Bucket { + let queue = &mut self.levels[level]; + let bucket = queue.pop_front().unwrap_or_default(); + queue.push_back(Bucket::new()); + bucket + } + + fn file(&mut self, level: usize, slot: usize, link: Link) { + if let Some(bucket) = self.levels[level].get_mut(slot) { + bucket.push(link); + } + } + + /// The route to file for a timer due `remaining` ticks after a rotation + /// boundary, built by recursing into the finer levels so each leg owns the + /// part of the route it hands on. `remaining` must be non-zero. + fn route(&self, remaining: u64, inner: Link) -> Link { + let (level, slot, rest) = self.next_leg(remaining); + let next = match rest { + 0 => inner, + rest => self.route(rest, inner), + }; + Link::Relay { + level: level as u8, + slot: slot as u16, + next: Box::new(next), + } + } + + /// The route's first leg if it starts at `level`: which bucket holds a timer + /// due at tick `target`, and how much delay that leaves for the legs after it. + /// `None` when this level's next rotation already overshoots `target`, or when + /// `target` is more than one revolution away. + fn entry_leg(&self, level: usize, target: u64) -> Option<(usize, usize, u64)> { + let period = self.period(level); + let next_rotation = self.ticks - self.ticks % period + period; + let ahead = target.checked_sub(next_rotation)?; + let slot = ahead / period; + (slot < self.radix).then_some((level, slot as usize, ahead % period)) + } + + /// The next leg for a timer due `remaining` ticks after a rotation boundary: + /// the coarsest level whose rotation still fits. From a boundary that level's + /// front bucket comes off one period out, so bucket `j` comes off after + /// `j + 1` of them. + fn next_leg(&self, remaining: u64) -> (usize, usize, u64) { + let level = (0..self.levels.len()) + .rev() + .find(|level| self.period(*level) <= remaining) + .unwrap_or(0); + let period = self.period(level); + (level, (remaining / period - 1) as usize, remaining % period) + } + + /// Ticks spanned by `level` and every level below it: `radix^level`. + fn period(&self, level: usize) -> u64 { + self.radix.saturating_pow(level as u32) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + /// A wheel plus the tick each scheduled timer actually fired on. + struct Harness { + wheel: TimingWheel, + fired: Arc>>, + ticks: Arc, + } + + impl Harness { + fn new(max_delay: u64, radix: u64) -> Self { + Self { + wheel: TimingWheel::new(max_delay, radix), + fired: Arc::new(Mutex::new(Vec::new())), + ticks: Arc::new(AtomicU64::new(0)), + } + } + + fn timer(&self, id: u32) -> Arc { + let fired = self.fired.clone(); + let ticks = self.ticks.clone(); + Timer::new(move || { + fired.lock().push((id, ticks.load(Ordering::Acquire))); + }) + } + + fn schedule(&mut self, id: u32, delay: u64) { + let timer = self.timer(id); + self.wheel.schedule(delay, timer); + } + + fn tick(&mut self) { + self.ticks.fetch_add(1, Ordering::AcqRel); + self.wheel.tick(); + } + + fn fired_at(&self, id: u32) -> Option { + self.fired + .lock() + .iter() + .find(|(fired, _)| *fired == id) + .map(|(_, at)| *at) + } + } + + /// Every delay, from every starting offset, must fire on exactly the tick it + /// asked for. This is the wheel's whole contract, and a radix decomposition is + /// easy to get wrong by one bucket, so it is checked exhaustively rather than + /// sampled. + #[test] + fn every_delay_fires_on_its_exact_tick() { + let radix = 4; + let max_delay = radix * radix * radix; + for offset in 0..2 * radix * radix { + let mut harness = Harness::new(max_delay, radix); + for _ in 0..offset { + harness.tick(); + } + for delay in 1..=max_delay { + harness.schedule(delay as u32, delay); + } + for _ in 0..max_delay { + harness.tick(); + } + for delay in 1..=max_delay { + assert_eq!( + harness.fired_at(delay as u32), + Some(offset + delay), + "radix {radix}, offset {offset}, delay {delay}" + ); + } + } + } + + /// The same contract at the shape the wheel actually runs with. + #[test] + fn every_short_delay_fires_on_its_exact_tick_at_radix_64() { + let radix = 64; + let mut harness = Harness::new(radix * radix * radix * radix, radix); + for _ in 0..100 { + harness.tick(); + } + let delays = (1..=200).chain([radix - 1, radix, radix + 1, radix * radix, 4095, 4096]); + for delay in delays.clone() { + harness.schedule(delay as u32, delay); + } + for _ in 0..5000 { + harness.tick(); + } + for delay in delays { + assert_eq!( + harness.fired_at(delay as u32), + Some(100 + delay), + "delay {delay}" + ); + } + } + + #[test] + fn level_count_covers_the_requested_delay() { + assert_eq!(TimingWheel::new(64, 64).max_delay(), 64); + assert_eq!(TimingWheel::new(65, 64).max_delay(), 4096); + assert_eq!(TimingWheel::new(4096, 64).max_delay(), 4096); + assert_eq!(TimingWheel::new(4097, 64).max_delay(), 262_144); + } + + #[test] + fn a_delay_the_wheel_cannot_place_fires_at_once() { + let mut harness = Harness::new(64, 64); + harness.schedule(1, 65); + assert_eq!(harness.fired_at(1), Some(0)); + harness.schedule(2, 0); + assert_eq!(harness.fired_at(2), Some(0)); + } + + #[test] + fn rescheduling_the_same_timer_defers_it_to_the_later_route() { + let mut harness = Harness::new(4096, 64); + let timer = harness.timer(1); + harness.wheel.schedule(5, timer.clone()); + harness.wheel.schedule(20, timer); + + for _ in 0..5 { + harness.tick(); + } + assert_eq!( + harness.fired_at(1), + None, + "the earlier route must not fire the timer" + ); + for _ in 5..20 { + harness.tick(); + } + assert_eq!(harness.fired_at(1), Some(20)); + } + + #[test] + fn firing_early_makes_the_scheduled_route_inert() { + let mut harness = Harness::new(4096, 64); + let timer = harness.timer(1); + harness.wheel.schedule(10, timer.clone()); + + timer.fire(); + assert_eq!(harness.fired_at(1), Some(0)); + for _ in 0..10 { + harness.tick(); + } + assert_eq!(harness.fired.lock().len(), 1); + } + + #[test] + fn dropping_the_wheel_fires_everything_it_holds() { + let mut harness = Harness::new(262_144, 64); + harness.schedule(1, 5); + harness.schedule(2, 200_000); + + let fired = harness.fired.clone(); + drop(harness); + let ids = fired + .lock() + .iter() + .map(|(id, _)| *id) + .collect::>(); + assert_eq!(ids, HashSet::from([1, 2])); + } +} diff --git a/crates/pb-mapper-cli/Cargo.toml b/crates/pb-mapper-cli/Cargo.toml new file mode 100644 index 0000000..f45b297 --- /dev/null +++ b/crates/pb-mapper-cli/Cargo.toml @@ -0,0 +1,40 @@ +[package] +name = "pb-mapper-cli" +version.workspace = true +edition.workspace = true +authors.workspace = true + +# The binary keeps the name `pb-mapper`, discovered from +# `src/bin/pb-mapper.rs`. Release workflows, both Dockerfiles, and the install +# scripts all hardcode it, and `cargo build --bin pb-mapper` resolves it from +# the workspace root regardless of which crate holds it. + +[dependencies] +pb-mapper-auth.workspace = true +pb-mapper-client.workspace = true +pb-mapper-core.workspace = true +pb-mapper-protocol.workspace = true +pb-mapper-server.workspace = true + +better_mimalloc_rs.workspace = true +clap.workspace = true +serde_json.workspace = true +tokio.workspace = true +tokio-util.workspace = true +tracing.workspace = true +uni-stream.workspace = true + +[dev-dependencies] +dotenvy.workspace = true +rand.workspace = true + +[features] +udp-timeout = [ + "uni-stream/udp-timeout", + "pb-mapper-protocol/udp-timeout", + "pb-mapper-client/udp-timeout", + "pb-mapper-server/udp-timeout", +] + +[lints] +workspace = true diff --git a/examples/echo_tcp_client.rs b/crates/pb-mapper-cli/examples/echo_tcp_client.rs similarity index 100% rename from examples/echo_tcp_client.rs rename to crates/pb-mapper-cli/examples/echo_tcp_client.rs diff --git a/examples/echo_tcp_server.rs b/crates/pb-mapper-cli/examples/echo_tcp_server.rs similarity index 100% rename from examples/echo_tcp_server.rs rename to crates/pb-mapper-cli/examples/echo_tcp_server.rs diff --git a/examples/echo_udp_client.rs b/crates/pb-mapper-cli/examples/echo_udp_client.rs similarity index 93% rename from examples/echo_udp_client.rs rename to crates/pb-mapper-cli/examples/echo_udp_client.rs index 1a69b4e..9db1a3a 100644 --- a/examples/echo_udp_client.rs +++ b/crates/pb-mapper-cli/examples/echo_udp_client.rs @@ -1,6 +1,6 @@ use std::error::Error; -use pb_mapper::common::config::init_tracing; +use pb_mapper_core::config::init_tracing; use tokio::io::{AsyncReadExt, AsyncWriteExt}; use uni_stream::udp::UdpStream; diff --git a/examples/echo_udp_server.rs b/crates/pb-mapper-cli/examples/echo_udp_server.rs similarity index 88% rename from examples/echo_udp_server.rs rename to crates/pb-mapper-cli/examples/echo_udp_server.rs index 135aae1..6dc17c5 100644 --- a/examples/echo_udp_server.rs +++ b/crates/pb-mapper-cli/examples/echo_udp_server.rs @@ -1,8 +1,11 @@ +// An example: panicking on a failed bind is the clearest thing it can do. +#![allow(clippy::unwrap_used, clippy::expect_used)] + use std::error::Error; use std::net::SocketAddr; use std::str::FromStr; -use pb_mapper::common::config::init_tracing; +use pb_mapper_core::config::init_tracing; use tokio::io::{AsyncReadExt, AsyncWriteExt}; use uni_stream::udp::UdpListener; diff --git a/examples/pb_local_client.rs b/crates/pb-mapper-cli/examples/pb_local_client.rs similarity index 72% rename from examples/pb_local_client.rs rename to crates/pb-mapper-cli/examples/pb_local_client.rs index 3e7d1bc..e42731b 100644 --- a/examples/pb_local_client.rs +++ b/crates/pb-mapper-cli/examples/pb_local_client.rs @@ -1,5 +1,5 @@ -use pb_mapper::common::config::init_tracing; -use pb_mapper::local::client::run_client_side_cli; +use pb_mapper_client::client::run_client_side_cli; +use pb_mapper_core::config::init_tracing; use uni_stream::stream::TcpListenerProvider; #[tokio::main] diff --git a/examples/pb_local_server.rs b/crates/pb-mapper-cli/examples/pb_local_server.rs similarity index 67% rename from examples/pb_local_server.rs rename to crates/pb-mapper-cli/examples/pb_local_server.rs index 1ae9e97..584824d 100644 --- a/examples/pb_local_server.rs +++ b/crates/pb-mapper-cli/examples/pb_local_server.rs @@ -1,5 +1,5 @@ -use pb_mapper::common::config::init_tracing; -use pb_mapper::local::server::{run_server_side_cli, ServerTunnelOptions}; +use pb_mapper_client::server::{ServerTunnelOptions, run_server_side_cli}; +use pb_mapper_core::config::init_tracing; use uni_stream::stream::TcpStreamProvider; #[tokio::main] @@ -13,6 +13,8 @@ async fn main() { need_codec: false, is_datagram: false, keep_alive: false, + namespace: None, + force_namespace: false, }, ) .await; diff --git a/examples/pb_server.rs b/crates/pb-mapper-cli/examples/pb_server.rs similarity index 57% rename from examples/pb_server.rs rename to crates/pb-mapper-cli/examples/pb_server.rs index 1c09ff0..b8e3ba5 100644 --- a/examples/pb_server.rs +++ b/crates/pb-mapper-cli/examples/pb_server.rs @@ -1,5 +1,5 @@ -use pb_mapper::common::config::init_tracing; -use pb_mapper::pb_server::run_server; +use pb_mapper_core::config::init_tracing; +use pb_mapper_server::run_server; #[tokio::main] async fn main() -> std::io::Result<()> { diff --git a/crates/pb-mapper-cli/src/bin/pb-mapper.rs b/crates/pb-mapper-cli/src/bin/pb-mapper.rs new file mode 100644 index 0000000..14d6be5 --- /dev/null +++ b/crates/pb-mapper-cli/src/bin/pb-mapper.rs @@ -0,0 +1,641 @@ +//! Unified command-line entry point for every pb-mapper role. +//! +//! ```text +//! +-> server (relay) +//! process args -> clap -+-> register (publish a local service) +//! +-> connect (open a local listener) +//! +-> status (namespace-scoped inspection) +//! +-> admin (credential/control plane) +//! ``` +//! +//! Role-specific execution stays below this dispatch layer. Administrator parsing, +//! pagination, wire requests, and output rendering live in the `admin` module. + +use std::error::Error; +use std::net::{IpAddr, Ipv4Addr, Ipv6Addr}; +use std::path::PathBuf; +use std::time::Duration; + +use better_mimalloc_rs::MiMalloc; +use clap::{Args, Parser, Subcommand, ValueEnum}; +use pb_mapper_auth::{ + AuthConfig, KeyPage, LegacyProtocolPolicy, MAX_TEMP_KEY_CAPACITY, MAX_TEMP_KEY_TTL, + MIN_TEMP_KEY_TTL, acquire_state_dir_lock, generate_admin_key, initialize_admin_key, + write_admin_key_file, +}; +use pb_mapper_client::client::{ + handle_status_cli_scoped, run_client_side_cli_with_callback_scoped, +}; +use pb_mapper_client::server::{ServerTunnelOptions, run_server_side_cli_with_pinned_credential}; +use pb_mapper_core::checksum::set_process_msg_header_key; +use pb_mapper_core::checksum::{MACHINE_MSG_HEADER_KEY_PATH, setup_machine_msg_header_key}; +use pb_mapper_core::config::{ + StatusOp, control_io_timeout, get_pb_mapper_server_async, get_sockaddr_async, init_tracing, + keep_alive_from_env, +}; +use pb_mapper_protocol::MessageReader; +use pb_mapper_protocol::command::{ + AdminConnectionPage, AdminRequest, AdminResponse, AdminServicePage, MessageSerializer, + PbConnRequest, PbConnResponse, +}; +use pb_mapper_protocol::forward::StreamForward; +use pb_mapper_protocol::secure::ClientHeaderSession; +use pb_mapper_server::run_server_with_shutdown; +use tokio::net::TcpStream; +use tokio_util::sync::CancellationToken; +use uni_stream::stream::{ + StreamProvider, TcpListenerProvider, TcpStreamProvider, UdpListenerProvider, UdpStreamProvider, +}; + +#[global_allocator] +static GLOBAL_MIMALLOC: MiMalloc = MiMalloc; + +#[derive(Debug, Parser)] +#[command( + author = "L_B__", + version, + about = "Expose and consume keyed TCP/UDP services through a pb-mapper relay", + subcommand_required = true, + arg_required_else_help = true +)] +struct Cli { + #[command(subcommand)] + command: Command, +} + +#[derive(Debug, Subcommand)] +enum Command { + /// Run the public relay server. + Server(ServerArgs), + /// Register a local service with a relay. + Register(RegisterArgs), + /// Expose a registered service on a local listening address. + Connect(ConnectArgs), + /// Query relay status. + Status(StatusArgs), + /// Manage temporary credentials and inspect relay authentication state. + Admin(AdminArgs), +} + +#[derive(Debug, Args)] +struct ServerArgs { + /// Port exposed to registering services and connecting clients. + #[arg(short, long, visible_alias = "pb-mapper-port", default_value_t = 7666)] + port: u16, + /// Listen on IPv6 (::) instead of IPv4 (0.0.0.0). + #[arg(long, visible_alias = "use-ipv6", default_value_t = false)] + ipv6: bool, + /// Enable TCP keep-alive. PB_MAPPER_KEEP_ALIVE=ON is also supported. + #[arg(long, default_value_t = false)] + keep_alive: bool, + /// Derive MSG_HEADER_KEY from this machine and persist it for other roles. + #[arg(long, default_value_t = false)] + use_machine_msg_header_key: bool, + /// Directory containing encrypted authentication state and the administrator key file. + /// Defaults to /var/lib/pb-mapper/auth for Linux services or a writable system + /// directory; otherwise a user-writable application directory. + #[arg(long)] + auth_state_dir: Option, + /// Create a random administrator key before starting the relay. + #[arg( + long, + conflicts_with = "use_machine_msg_header_key", + default_value_t = false + )] + init_admin_key: bool, + /// Replace an existing administrator key when used with --init-admin-key. + #[arg(long, requires = "init_admin_key", default_value_t = false)] + force_init_admin_key: bool, + /// Maximum temporary-key slots allocated by the relay. + #[arg(long)] + max_temporary_keys: Option, + /// Maximum accepted temporary-key TTL. + #[arg(long, value_parser = parse_duration)] + max_temporary_key_ttl: Option, + /// Allow or deny the legacy encrypted framing protocol. + #[arg(long, value_enum)] + legacy_protocol: Option, +} + +#[path = "pb-mapper/admin.rs"] +mod admin; +use admin::AdminArgs; +#[derive(Debug, Args)] +struct RegisterArgs { + /// Transport used by the local service. + #[arg(value_enum)] + transport: Transport, + /// Service key registered with the relay. + #[arg(short, long)] + key: String, + /// Local service address to forward to. + #[arg(short, long, visible_alias = "local")] + addr: String, + #[command(flatten)] + relay: RelayArgs, + /// Encrypt forwarded traffic with the configured MSG_HEADER_KEY. + #[arg(short, long, default_value_t = false)] + codec: bool, + /// Administrator-only target namespace. Temporary credentials always use their own key id. + #[arg(long)] + namespace: Option, + /// Required when an administrator registers a service inside a temporary-key namespace. + #[arg(long, requires = "namespace", default_value_t = false)] + force: bool, +} + +#[derive(Debug, Args)] +struct ConnectArgs { + /// Transport exposed by the local listener. + #[arg(value_enum)] + transport: Transport, + /// Registered service key to subscribe to. + #[arg(short, long)] + key: String, + /// Local address on which downstream clients connect. + #[arg(short, long, visible_alias = "local")] + addr: String, + #[command(flatten)] + relay: RelayArgs, + /// Administrator-only target namespace. Temporary credentials always use their own key id. + #[arg(long)] + namespace: Option, +} + +#[derive(Debug, Args)] +struct StatusArgs { + /// Status query to execute. + #[arg(value_enum)] + op: StatusOp, + /// Relay address. Falls back to PB_MAPPER_SERVER. + #[arg(short, long, visible_alias = "pb-mapper-server", value_name = "ADDR")] + server: Option, + /// Administrator-only namespace to inspect. + #[arg(long)] + namespace: Option, +} + +#[derive(Debug, Args)] +struct RelayArgs { + /// Relay address. Falls back to PB_MAPPER_SERVER. + #[arg(short, long, visible_alias = "pb-mapper-server", value_name = "ADDR")] + server: Option, + /// Enable TCP keep-alive. PB_MAPPER_KEEP_ALIVE=ON is also supported. + #[arg(long, default_value_t = false)] + keep_alive: bool, +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq, ValueEnum)] +enum Transport { + Tcp, + Udp, +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq, ValueEnum)] +enum LegacyProtocolArg { + Allow, + Deny, +} + +impl From for LegacyProtocolPolicy { + fn from(value: LegacyProtocolArg) -> Self { + match value { + LegacyProtocolArg::Allow => Self::Allow, + LegacyProtocolArg::Deny => Self::Deny, + } + } +} + +#[tokio::main] +async fn main() { + MiMalloc::init(); + let cli = Cli::parse(); + init_tracing(); + + if let Err(error) = run(cli).await { + tracing::error!(%error, "pb-mapper command failed"); + std::process::exit(1); + } +} + +async fn run(cli: Cli) -> Result<(), Box> { + match cli.command { + Command::Server(args) => run_server(args).await?, + Command::Register(args) => run_register(args).await?, + Command::Connect(args) => run_connect(args).await?, + Command::Status(args) => run_status(args).await?, + Command::Admin(args) => admin::run_admin(args).await?, + } + Ok(()) +} + +/// Publishes CLI flags as environment variables, which is how the auth +/// subsystem reads its configuration. +/// +/// # Safety note +/// +/// Mutating the environment is unsafe in edition 2024 because it races +/// concurrent readers. Every call here happens during argument handling on the +/// main thread, before any runtime task or thread is spawned, so there is no +/// concurrent reader to race. +fn apply_server_auth_overrides(args: &ServerArgs) -> Result<(), Box> { + if let Some(auth_state_dir) = &args.auth_state_dir { + unsafe { std::env::set_var("PB_MAPPER_AUTH_STATE_DIR", auth_state_dir) }; + } + if let Some(max_temporary_keys) = args.max_temporary_keys { + if !(1..=MAX_TEMP_KEY_CAPACITY).contains(&max_temporary_keys) { + return Err(format!( + "`--max-temporary-keys` must be between 1 and {MAX_TEMP_KEY_CAPACITY}" + ) + .into()); + } + unsafe { + std::env::set_var( + "PB_MAPPER_AUTH_MAX_TEMP_KEYS", + max_temporary_keys.to_string(), + ) + }; + } + if let Some(max_temporary_key_ttl) = args.max_temporary_key_ttl { + if max_temporary_key_ttl < MIN_TEMP_KEY_TTL || max_temporary_key_ttl > MAX_TEMP_KEY_TTL { + return Err(format!( + "`--max-temporary-key-ttl` must be between {}s and {}d", + MIN_TEMP_KEY_TTL.as_secs(), + MAX_TEMP_KEY_TTL.as_secs() / 86_400 + ) + .into()); + } + unsafe { + std::env::set_var( + "PB_MAPPER_AUTH_MAX_TEMP_TTL_SECS", + max_temporary_key_ttl.as_secs().to_string(), + ) + }; + } + Ok(()) +} + +async fn run_server(args: ServerArgs) -> Result<(), Box> { + apply_server_auth_overrides(&args)?; + if let Some(legacy_protocol) = args.legacy_protocol { + // SAFETY: as in `apply_server_auth_overrides` — this runs before the + // server spawns anything that reads the environment. + unsafe { + std::env::set_var( + "PB_MAPPER_LEGACY_PROTOCOL", + match legacy_protocol { + LegacyProtocolArg::Allow => "allow", + LegacyProtocolArg::Deny => "deny", + }, + ) + }; + } + let auth_config = AuthConfig::default(); + if args.init_admin_key { + std::fs::create_dir_all(&auth_config.state_dir)?; + let _lock = acquire_state_dir_lock(&auth_config.state_dir)?; + let key_path = auth_config.state_dir.join("admin.key"); + let key = initialize_admin_key(&key_path, args.force_init_admin_key)?; + drop(_lock); + set_process_msg_header_key(Some(&key))?; + eprintln!("administrator key initialized at {}", key_path.display()); + } else if args.use_machine_msg_header_key { + let admin_key_path = auth_config.state_dir.join("admin.key"); + if admin_key_path.exists() { + return Err(format!( + "--use-machine-msg-header-key cannot replace `{}`; use `pb-mapper admin root-key rotate` to change the root key", + admin_key_path.display() + ) + .into()); + } + tracing::warn!( + "--use-machine-msg-header-key is a legacy compatibility option; prefer a random administrator key" + ); + setup_machine_msg_header_key()?; + tracing::info!( + path = MACHINE_MSG_HEADER_KEY_PATH, + "derived and persisted machine MSG_HEADER_KEY" + ); + } + + let ip_addr = if args.ipv6 { + IpAddr::V6(Ipv6Addr::UNSPECIFIED) + } else { + IpAddr::V4(Ipv4Addr::UNSPECIFIED) + }; + run_server_with_shutdown( + (ip_addr, args.port), + CancellationToken::new(), + None, + args.keep_alive || keep_alive_from_env(), + ) + .await?; + Ok(()) +} + +async fn run_register(args: RegisterArgs) -> Result<(), Box> { + let credential = pb_mapper_core::checksum::get_process_credential().map_err(|error| { + std::io::Error::other(format!("registration credential is required: {error}")) + })?; + let local_addr = get_sockaddr_async(&args.addr).await?; + let remote_addr = get_pb_mapper_server_async(args.relay.server.as_deref()).await?; + let options = ServerTunnelOptions { + need_codec: args.codec, + is_datagram: args.transport == Transport::Udp, + keep_alive: args.relay.keep_alive || keep_alive_from_env(), + namespace: args.namespace, + force_namespace: args.force, + }; + + match args.transport { + Transport::Tcp => { + register::(local_addr, remote_addr, args.key, options, credential) + .await + } + Transport::Udp => { + register::(local_addr, remote_addr, args.key, options, credential) + .await + } + } + Ok(()) +} + +async fn register( + local_addr: std::net::SocketAddr, + remote_addr: std::net::SocketAddr, + key: String, + options: ServerTunnelOptions, + credential: pb_mapper_core::checksum::Credential, +) where + LocalStream::Item: StreamForward, +{ + run_server_side_cli_with_pinned_credential::( + local_addr, + remote_addr, + key.into(), + options, + None, + credential, + ) + .await; +} + +async fn run_connect(args: ConnectArgs) -> Result<(), Box> { + let credential = pb_mapper_core::checksum::get_process_credential().map_err(|error| { + std::io::Error::other(format!("client credential is required: {error}")) + })?; + let local_addr = get_sockaddr_async(&args.addr).await?; + let remote_addr = get_pb_mapper_server_async(args.relay.server.as_deref()).await?; + let key = args.key.into(); + let keep_alive = args.relay.keep_alive || keep_alive_from_env(); + + match args.transport { + Transport::Tcp => { + run_client_side_cli_with_callback_scoped::( + local_addr, + remote_addr, + key, + keep_alive, + args.namespace, + None, + Some(credential), + ) + .await; + } + Transport::Udp => { + run_client_side_cli_with_callback_scoped::( + local_addr, + remote_addr, + key, + keep_alive, + args.namespace, + None, + Some(credential), + ) + .await; + } + } + Ok(()) +} + +async fn run_status(args: StatusArgs) -> Result<(), Box> { + let remote_addr = get_pb_mapper_server_async(args.server.as_deref()).await?; + handle_status_cli_scoped(args.op, remote_addr, args.namespace).await +} + +fn parse_duration(raw: &str) -> Result { + let raw = raw.trim(); + if raw.is_empty() { + return Err("duration must not be empty".to_string()); + } + let split = raw + .find(|character: char| !character.is_ascii_digit()) + .unwrap_or(raw.len()); + let (number, unit) = raw.split_at(split); + let value = number + .parse::() + .map_err(|_| format!("invalid duration `{raw}`"))?; + let multiplier = match unit { + "" | "s" => 1, + "m" => 60, + "h" => 60 * 60, + "d" => 24 * 60 * 60, + _ => { + return Err(format!( + "unsupported duration unit `{unit}`; use s, m, h, or d" + )); + } + }; + value + .checked_mul(multiplier) + .map(Duration::from_secs) + .ok_or_else(|| "duration is too large".to_string()) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn parses_each_runtime_role() { + let cases = [ + vec!["pb-mapper", "server", "--port", "7666", "--ipv6"], + vec![ + "pb-mapper", + "register", + "tcp", + "--key", + "web", + "--addr", + "127.0.0.1:8080", + "--server", + "relay:7666", + "--codec", + ], + vec![ + "pb-mapper", + "connect", + "udp", + "--key", + "game", + "--addr", + "127.0.0.1:8211", + "--server", + "relay:7666", + ], + vec!["pb-mapper", "status", "keys", "--server", "relay:7666"], + vec![ + "pb-mapper", + "admin", + "--server", + "relay:7666", + "key", + "issue", + "--ttl", + "30d", + "--label", + "build-agent", + ], + vec![ + "pb-mapper", + "admin", + "--output", + "ndjson", + "connection", + "list", + "--page-size", + "1000", + "--all", + ], + vec![ + "pb-mapper", + "admin", + "root-key", + "rotate", + "--key-file", + "/tmp/pb-mapper-admin.key", + ], + ]; + + for args in cases { + Cli::try_parse_from(args).expect("unified command should parse"); + } + } + + #[test] + fn accepts_documented_option_aliases() { + Cli::try_parse_from([ + "pb-mapper", + "server", + "--pb-mapper-port", + "7666", + "--use-ipv6", + ]) + .expect("server aliases should parse"); + Cli::try_parse_from([ + "pb-mapper", + "register", + "tcp", + "--key", + "web", + "--local", + "127.0.0.1:8080", + "--pb-mapper-server", + "relay:7666", + ]) + .expect("relay and local aliases should parse"); + Cli::try_parse_from([ + "pb-mapper", + "register", + "tcp", + "--key", + "web", + "--addr", + "127.0.0.1:8080", + "--namespace", + "4294967296", + "--force", + ]) + .expect("administrator namespace registration flags should parse"); + } + + #[test] + fn rejects_invalid_admin_paging_and_duration() { + assert!( + Cli::try_parse_from(["pb-mapper", "admin", "key", "list", "--page-size", "1001",]) + .is_err() + ); + assert!( + Cli::try_parse_from(["pb-mapper", "admin", "key", "issue", "--ttl", "1fortnight",]) + .is_err() + ); + assert!( + Cli::try_parse_from([ + "pb-mapper", + "server", + "--init-admin-key", + "--use-machine-msg-header-key", + ]) + .is_err() + ); + } + + #[test] + fn server_auth_options_only_override_environment_when_explicit() { + let cli = + Cli::try_parse_from(["pb-mapper", "server"]).expect("server defaults should parse"); + let Command::Server(defaults) = cli.command else { + panic!("expected server command"); + }; + assert_eq!(defaults.auth_state_dir, None); + assert_eq!(defaults.max_temporary_keys, None); + assert_eq!(defaults.max_temporary_key_ttl, None); + assert_eq!(defaults.legacy_protocol, None); + + let cli = Cli::try_parse_from([ + "pb-mapper", + "server", + "--auth-state-dir", + "/tmp/pb-mapper-auth", + "--max-temporary-keys", + "1024", + "--max-temporary-key-ttl", + "2h", + "--legacy-protocol", + "deny", + ]) + .expect("explicit server authentication options should parse"); + let Command::Server(explicit) = cli.command else { + panic!("expected server command"); + }; + assert_eq!( + explicit.auth_state_dir, + Some(PathBuf::from("/tmp/pb-mapper-auth")) + ); + assert_eq!(explicit.max_temporary_keys, Some(1024)); + assert_eq!( + explicit.max_temporary_key_ttl, + Some(Duration::from_secs(2 * 60 * 60)) + ); + assert_eq!(explicit.legacy_protocol, Some(LegacyProtocolArg::Deny)); + } + + #[test] + fn explicit_out_of_range_server_auth_flags_are_rejected() { + let cli = Cli::try_parse_from(["pb-mapper", "server", "--max-temporary-keys", "0"]) + .expect("clap should accept the token before bounds checking"); + let Command::Server(args) = cli.command else { + panic!("expected server command"); + }; + let error = apply_server_auth_overrides(&args).unwrap_err(); + assert!(error.to_string().contains("--max-temporary-keys")); + + let cli = Cli::try_parse_from(["pb-mapper", "server", "--max-temporary-key-ttl", "5s"]) + .expect("clap should accept the token before bounds checking"); + let Command::Server(args) = cli.command else { + panic!("expected server command"); + }; + let error = apply_server_auth_overrides(&args).unwrap_err(); + assert!(error.to_string().contains("--max-temporary-key-ttl")); + } +} diff --git a/crates/pb-mapper-cli/src/bin/pb-mapper/admin.rs b/crates/pb-mapper-cli/src/bin/pb-mapper/admin.rs new file mode 100644 index 0000000..1f9d620 --- /dev/null +++ b/crates/pb-mapper-cli/src/bin/pb-mapper/admin.rs @@ -0,0 +1,685 @@ +//! Administrator CLI: command parsing, one-shot V2 requests, pagination, and rendering. +//! +//! ```text +//! admin args -> AdminRequest -> authenticated V2 connection -> relay +//! ^ | +//! +--- human / JSON / NDJSON <- AdminResponse <--------+ +//! ``` +//! +//! `--all` keeps the selected output contract: human and JSON aggregate pages, +//! while NDJSON deliberately streams one item at a time. Root-key rotation stages +//! a recovery copy before contacting the relay, then verifies the new credential. + +use super::*; + +#[derive(Debug, Args)] +pub(super) struct AdminArgs { + /// Relay address. Falls back to PB_MAPPER_SERVER. + #[arg(short, long, visible_alias = "pb-mapper-server", value_name = "ADDR")] + server: Option, + /// Machine-readable output mode. + #[arg(long, value_enum, default_value_t = OutputFormat::Human)] + output: OutputFormat, + #[command(subcommand)] + command: AdminCommand, +} + +#[derive(Debug, Subcommand)] +enum AdminCommand { + /// Issue, inspect, renew, reveal, revoke, or collect temporary keys. + Key(AdminKeyArgs), + /// List relay connections across namespaces. + Connection(AdminConnectionArgs), + /// List registered services across namespaces. + Service(AdminServiceArgs), + /// Show authentication state and protocol counters. + Status, + /// Repair or reset encrypted temporary-key state. + AuthState(AdminAuthStateArgs), + /// Rotate the sole administrator key and invalidate every existing credential. + RootKey(AdminRootKeyArgs), + /// Change legacy protocol acceptance at runtime. + LegacyProtocol(AdminLegacyProtocolArgs), +} + +#[derive(Debug, Args)] +struct AdminKeyArgs { + #[command(subcommand)] + command: AdminKeyCommand, +} + +#[derive(Debug, Subcommand)] +enum AdminKeyCommand { + Issue { + #[arg(long, value_parser = parse_duration)] + ttl: Duration, + #[arg(long)] + label: Option, + }, + List { + #[arg(long, default_value_t = 0)] + page: u32, + #[arg(long, default_value_t = 100, value_parser = clap::value_parser!(u16).range(1..=1000))] + page_size: u16, + #[arg(long, default_value_t = false)] + all: bool, + }, + Show { + key_id: u64, + }, + Reveal { + key_id: u64, + }, + Renew { + key_id: u64, + #[arg(long, value_parser = parse_duration)] + ttl: Duration, + }, + Revoke { + key_id: u64, + }, + Gc, +} + +#[derive(Debug, Args)] +struct AdminConnectionArgs { + #[command(subcommand)] + command: AdminListCommand, +} + +#[derive(Debug, Args)] +struct AdminServiceArgs { + #[command(subcommand)] + command: AdminListCommand, +} + +#[derive(Debug, Clone, Subcommand)] +enum AdminListCommand { + List { + #[arg(long)] + key_id: Option, + #[arg(long, default_value_t = 0)] + page: u32, + #[arg(long, default_value_t = 100, value_parser = clap::value_parser!(u16).range(1..=1000))] + page_size: u16, + #[arg(long, default_value_t = false)] + all: bool, + }, +} + +#[derive(Debug, Args)] +struct AdminAuthStateArgs { + #[command(subcommand)] + command: AdminAuthStateCommand, +} + +#[derive(Debug, Subcommand)] +enum AdminAuthStateCommand { + Reset { + #[arg(long, default_value_t = false)] + confirm: bool, + }, +} + +#[derive(Debug, Args)] +struct AdminRootKeyArgs { + #[command(subcommand)] + command: AdminRootKeyCommand, +} + +#[derive(Debug, Subcommand)] +enum AdminRootKeyCommand { + Rotate { + /// New 32-byte administrator key. A cryptographically random printable key is generated when omitted. + #[arg(long)] + new_key: Option, + /// Save the new key here before asking the relay to rotate. + #[arg(long)] + key_file: Option, + }, +} + +#[derive(Debug, Args)] +struct AdminLegacyProtocolArgs { + #[command(subcommand)] + command: AdminLegacyProtocolCommand, +} + +#[derive(Debug, Subcommand)] +enum AdminLegacyProtocolCommand { + Set { policy: LegacyProtocolArg }, +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq, ValueEnum)] +enum OutputFormat { + Human, + Json, + Ndjson, +} + +pub(super) async fn run_admin(args: AdminArgs) -> Result<(), Box> { + let remote_addr = get_pb_mapper_server_async(args.server.as_deref()).await?; + match args.command { + AdminCommand::Key(AdminKeyArgs { command }) => match command { + AdminKeyCommand::Issue { ttl, label } => { + let response = send_admin_request( + remote_addr, + AdminRequest::KeyIssue { + ttl_seconds: ttl.as_secs(), + label, + }, + ) + .await?; + print_admin_response(args.output, &response)?; + } + AdminKeyCommand::List { + page, + page_size, + all, + } => { + stream_key_pages(remote_addr, args.output, page, page_size, all).await?; + } + AdminKeyCommand::Show { key_id } => { + let response = + send_admin_request(remote_addr, AdminRequest::KeyShow { key_id }).await?; + print_admin_response(args.output, &response)?; + } + AdminKeyCommand::Reveal { key_id } => { + let response = + send_admin_request(remote_addr, AdminRequest::KeyReveal { key_id }).await?; + print_admin_response(args.output, &response)?; + } + AdminKeyCommand::Renew { key_id, ttl } => { + let response = send_admin_request( + remote_addr, + AdminRequest::KeyRenew { + key_id, + ttl_seconds: ttl.as_secs(), + }, + ) + .await?; + print_admin_response(args.output, &response)?; + } + AdminKeyCommand::Revoke { key_id } => { + let response = + send_admin_request(remote_addr, AdminRequest::KeyRevoke { key_id }).await?; + print_admin_response(args.output, &response)?; + } + AdminKeyCommand::Gc => { + let response = send_admin_request(remote_addr, AdminRequest::KeyGc).await?; + print_admin_response(args.output, &response)?; + } + }, + AdminCommand::Connection(AdminConnectionArgs { command }) => { + let AdminListCommand::List { + key_id, + page, + page_size, + all, + } = command; + stream_connection_pages(remote_addr, args.output, key_id, page, page_size, all).await?; + } + AdminCommand::Service(AdminServiceArgs { command }) => { + let AdminListCommand::List { + key_id, + page, + page_size, + all, + } = command; + stream_service_pages(remote_addr, args.output, key_id, page, page_size, all).await?; + } + AdminCommand::Status => { + let response = send_admin_request(remote_addr, AdminRequest::AuthStatus).await?; + print_admin_response(args.output, &response)?; + } + AdminCommand::AuthState(AdminAuthStateArgs { + command: AdminAuthStateCommand::Reset { confirm }, + }) => { + let response = + send_admin_request(remote_addr, AdminRequest::AuthStateReset { confirm }).await?; + print_admin_response(args.output, &response)?; + } + AdminCommand::RootKey(AdminRootKeyArgs { + command: AdminRootKeyCommand::Rotate { new_key, key_file }, + }) => { + let key_file = key_file.unwrap_or_else(default_admin_recovery_key_file); + let new_key = new_key.unwrap_or_else(generate_admin_key); + let staged_key_file = key_file.with_file_name(format!( + ".{}.next", + key_file + .file_name() + .and_then(|name| name.to_str()) + .unwrap_or("admin.key") + )); + write_admin_key_file(&staged_key_file, &new_key, true)?; + let response = send_admin_request( + remote_addr, + AdminRequest::RootKeyRotate { + new_admin_key: new_key.clone(), + }, + ) + .await + .map_err(|error| { + std::io::Error::other(format!( + "root rotation request failed; the candidate key remains at `{}`: {error}", + staged_key_file.display() + )) + })?; + set_process_msg_header_key(Some(&new_key))?; + let verification = send_admin_request(remote_addr, AdminRequest::AuthStatus).await?; + if !matches!(verification, AdminResponse::AuthStatus(_)) { + return Err(std::io::Error::other( + "new administrator key did not pass the post-rotation status check", + ) + .into()); + } + write_admin_key_file(&key_file, &new_key, true).map_err(|error| { + std::io::Error::other(format!( + "administrator key rotated and verified, but `{}` could not be updated; recover the key from `{}`: {error}", + key_file.display(), + staged_key_file.display() + )) + })?; + if let Err(error) = std::fs::remove_file(&staged_key_file) { + tracing::warn!( + path = %staged_key_file.display(), + %error, + "administrator key was rotated, but the staged key file could not be removed" + ); + } + if args.output == OutputFormat::Human { + println!("administrator key rotated and verified"); + println!("all temporary credentials are now invalid (temporary_key_rotated)"); + println!("key file: {}", key_file.display()); + } else { + print_admin_response(args.output, &response)?; + } + } + AdminCommand::LegacyProtocol(AdminLegacyProtocolArgs { + command: AdminLegacyProtocolCommand::Set { policy }, + }) => { + let response = send_admin_request( + remote_addr, + AdminRequest::LegacyProtocolSet { + policy: policy.into(), + }, + ) + .await?; + print_admin_response(args.output, &response)?; + } + } + Ok(()) +} + +async fn send_admin_request( + remote_addr: std::net::SocketAddr, + request: AdminRequest, +) -> Result> { + send_admin_request_with_timeout(remote_addr, request, control_io_timeout()).await +} + +async fn send_admin_request_with_timeout( + remote_addr: std::net::SocketAddr, + request: AdminRequest, + io_timeout: Duration, +) -> Result> { + let encoded = PbConnRequest::Admin(request).encode()?; + for attempt in 0..2 { + let sent = std::sync::atomic::AtomicBool::new(false); + let attempt_result = tokio::time::timeout(io_timeout, async { + let mut stream = TcpStream::connect(remote_addr) + .await + .map_err(|error| -> Box { Box::new(error) })?; + let session = ClientHeaderSession::from_process()?; + session.write_initial(&mut stream, &encoded).await?; + sent.store(true, std::sync::atomic::Ordering::Release); + let mut reader = session.response_reader(&mut stream)?; + let message = reader.read_msg().await?; + Ok::<_, Box>(PbConnResponse::decode(message)?) + }) + .await + .map_err(|_| { + std::io::Error::new( + std::io::ErrorKind::TimedOut, + format!( + "administrator request attempt timed out after {} ms", + io_timeout.as_millis() + ), + ) + }); + let pre_send = !sent.load(std::sync::atomic::Ordering::Acquire); + let response = match attempt_result { + Ok(Ok(response)) => response, + Ok(Err(_)) if attempt == 0 && pre_send => continue, + Ok(Err(error)) => return Err(error), + Err(_) if attempt == 0 && pre_send => continue, + Err(error) => return Err(error.into()), + }; + match response { + PbConnResponse::Admin(response) => return Ok(response), + PbConnResponse::Error(error) + if error.code == "connection_salt_replayed" && error.retryable => + { + if attempt == 0 { + continue; + } + } + PbConnResponse::Error(error) => { + return Err(std::io::Error::other(format!( + "{}: {} (retryable={})", + error.code, error.message, error.retryable + )) + .into()); + } + response => { + return Err(std::io::Error::other(format!( + "unexpected administrator response: {response:?}" + )) + .into()); + } + } + } + Err(std::io::Error::other("connection salt replay retry was exhausted").into()) +} + +async fn stream_key_pages( + remote_addr: std::net::SocketAddr, + output: OutputFormat, + mut page: u32, + page_size: u16, + all: bool, +) -> Result<(), Box> { + let mut combined: Option = None; + loop { + let response = + send_admin_request(remote_addr, AdminRequest::KeyList { page, page_size }).await?; + let AdminResponse::KeyList(key_page) = &response else { + return Err(std::io::Error::other("unexpected key-list response").into()); + }; + if all { + if output == OutputFormat::Ndjson { + for item in &key_page.items { + println!("{}", serde_json::to_string(item)?); + } + } else { + let page = combined.get_or_insert_with(|| { + let mut page = key_page.clone(); + page.items.clear(); + page.next_page = None; + page + }); + page.items.extend(key_page.items.iter().cloned()); + } + } else { + print_admin_response(output, &response)?; + } + let Some(next_page) = key_page.next_page else { + break; + }; + if !all { + break; + } + page = next_page; + } + if let Some(page) = combined { + print_admin_response(output, &AdminResponse::KeyList(page))?; + } + Ok(()) +} + +async fn stream_service_pages( + remote_addr: std::net::SocketAddr, + output: OutputFormat, + key_id: Option, + mut page: u32, + page_size: u16, + all: bool, +) -> Result<(), Box> { + let mut combined: Option = None; + loop { + let response = send_admin_request( + remote_addr, + AdminRequest::ServiceList { + key_id, + page, + page_size, + }, + ) + .await?; + let AdminResponse::Services(service_page) = &response else { + return Err(std::io::Error::other("unexpected service-list response").into()); + }; + if all { + if output == OutputFormat::Ndjson { + for item in &service_page.items { + println!("{}", serde_json::to_string(item)?); + } + } else { + let page = combined.get_or_insert_with(|| { + let mut page = service_page.clone(); + page.items.clear(); + page.next_page = None; + page + }); + page.items.extend(service_page.items.iter().cloned()); + } + } else { + print_admin_response(output, &response)?; + } + let Some(next_page) = service_page.next_page else { + break; + }; + if !all { + break; + } + page = next_page; + } + if let Some(page) = combined { + print_admin_response(output, &AdminResponse::Services(page))?; + } + Ok(()) +} + +async fn stream_connection_pages( + remote_addr: std::net::SocketAddr, + output: OutputFormat, + key_id: Option, + mut page: u32, + page_size: u16, + all: bool, +) -> Result<(), Box> { + let mut combined: Option = None; + loop { + let response = send_admin_request( + remote_addr, + AdminRequest::ConnectionList { + key_id, + page, + page_size, + }, + ) + .await?; + let AdminResponse::Connections(connection_page) = &response else { + return Err(std::io::Error::other("unexpected connection-list response").into()); + }; + if all { + if output == OutputFormat::Ndjson { + for item in &connection_page.items { + println!("{}", serde_json::to_string(item)?); + } + } else { + let page = combined.get_or_insert_with(|| { + let mut page = connection_page.clone(); + page.items.clear(); + page.next_page = None; + page + }); + page.items.extend(connection_page.items.iter().cloned()); + } + } else { + print_admin_response(output, &response)?; + } + let Some(next_page) = connection_page.next_page else { + break; + }; + if !all { + break; + } + page = next_page; + } + if let Some(page) = combined { + print_admin_response(output, &AdminResponse::Connections(page))?; + } + Ok(()) +} + +fn print_admin_response( + output: OutputFormat, + response: &AdminResponse, +) -> Result<(), Box> { + match output { + OutputFormat::Json => println!( + "{}", + serde_json::to_string_pretty(&serde_json::json!({ + "schema_version": 1, + "data": response, + }))? + ), + OutputFormat::Ndjson => println!("{}", serde_json::to_string(response)?), + OutputFormat::Human => print_human_admin_response(response), + } + Ok(()) +} + +fn print_human_admin_response(response: &AdminResponse) { + match response { + AdminResponse::KeyIssued(key) + | AdminResponse::KeyShown(key) + | AdminResponse::KeyRenewed(key) => { + println!("key id: {}", key.metadata.key_id); + println!("state: {}", key.metadata.state); + println!("expires at: {}", key.metadata.expires_at); + if let Some(label) = &key.metadata.label { + println!("label: {label}"); + } + if !key.credential.is_empty() { + println!("credential: {}", key.credential); + } + } + AdminResponse::KeyRevoked(key) => { + println!("key {}: {}", key.key_id, key.state); + } + AdminResponse::KeyList(page) => { + println!("KEY ID\tSTATE\tEXPIRES\tLABEL"); + for key in &page.items { + println!( + "{}\t{}\t{}\t{}", + key.key_id, + key.state, + key.expires_at, + key.label.as_deref().unwrap_or("") + ); + } + if let Some(next) = page.next_page { + println!("next page: {next}"); + } + } + AdminResponse::KeyGc { removed } => println!("removed {removed} inactive keys"), + AdminResponse::AuthStatus(status) => { + println!("safe mode: {}", status.safe_mode); + println!( + "keys: {} active / {} expired / {} revoked / {} capacity", + status.active_keys, status.expired_keys, status.revoked_keys, status.capacity + ); + println!("legacy protocol: {:?}", status.legacy_protocol); + println!( + "active legacy connections: {}", + status.active_legacy_connections + ); + println!( + "last legacy connection: {}", + status + .last_legacy_connection_at + .map(|value| value.to_string()) + .unwrap_or_else(|| "never".to_string()) + ); + println!( + "authentication: {} succeeded / {} failed", + status.auth_successes, status.auth_failures + ); + println!("server instance: {}", status.server_instance_id); + } + AdminResponse::Services(page) => { + println!("KEY ID\tSERVICE\tTRANSPORT\tCODEC\tCONNECTIONS"); + for service in &page.items { + println!( + "{}\t{}\t{}\t{}\t{}", + service.key_id, + service.service_name, + service.transport, + service.codec_enabled, + service.connection_count + ); + } + } + AdminResponse::Connections(page) => { + println!("KEY ID\tSERVICE\tCONN ID\tHEALTHY\tTRANSPORT\tCODEC"); + for connection in &page.items { + println!( + "{}\t{}\t{}\t{}\t{}\t{}", + connection.key_id, + connection.service_name, + connection.conn_id, + connection.healthy, + connection.transport, + connection.codec_enabled + ); + } + } + AdminResponse::Ok { action } => println!("ok: {action}"), + } +} + +fn default_admin_recovery_key_file() -> PathBuf { + std::env::var_os("XDG_CONFIG_HOME") + .map(PathBuf::from) + .or_else(|| std::env::var_os("HOME").map(|home| PathBuf::from(home).join(".config"))) + .unwrap_or_else(|| PathBuf::from(".")) + .join("pb-mapper") + .join("admin.key") +} + +#[cfg(test)] +mod tests { + use super::*; + + #[tokio::test] + async fn administrator_request_times_out_when_peer_stalls() { + set_process_msg_header_key(Some("0123456789abcdefghijklmnopqrstuv")) + .expect("test administrator credential should be valid"); + let listener = tokio::net::TcpListener::bind((std::net::Ipv4Addr::LOCALHOST, 0)) + .await + .expect("test listener should bind"); + let remote_addr = listener + .local_addr() + .expect("listener should have an address"); + let stalled_peer = tokio::spawn(async move { + let (_stream, _) = listener.accept().await.expect("test peer should connect"); + std::future::pending::<()>().await; + }); + + let error = send_admin_request_with_timeout( + remote_addr, + AdminRequest::AuthStatus, + Duration::from_millis(50), + ) + .await + .expect_err("a stalled administrator request should time out"); + + let io_error = error + .downcast_ref::() + .expect("timeout should be reported as an I/O error"); + assert_eq!(io_error.kind(), std::io::ErrorKind::TimedOut); + stalled_peer.abort(); + } +} diff --git a/tests/.env b/crates/pb-mapper-cli/tests/.env similarity index 100% rename from tests/.env rename to crates/pb-mapper-cli/tests/.env diff --git a/tests/regression.rs b/crates/pb-mapper-cli/tests/regression.rs similarity index 60% rename from tests/regression.rs rename to crates/pb-mapper-cli/tests/regression.rs index 145b3ce..5b3e7b2 100644 --- a/tests/regression.rs +++ b/crates/pb-mapper-cli/tests/regression.rs @@ -1,19 +1,30 @@ +// An integration test: a failed `unwrap` is a failed test, which is the report +// this file exists to produce. `allow-unwrap-in-tests` covers `#[cfg(test)]` +// modules but not a `tests/` target, whose whole body is test code. +#![allow(clippy::unwrap_used, clippy::expect_used)] + use std::net::SocketAddr; use std::sync::atomic::{AtomicUsize, Ordering}; use std::sync::{Arc, Mutex}; use std::time::Duration; -use pb_mapper::common::message::command::{ - LocalServer, MessageSerializer, PbConnRequest, PbConnResponse, PbConnStatusReq, - PbConnStatusResp, PbServerRequest, PbServiceConnStatus, +use pb_mapper_auth::{ + ADMIN_KEY_ID, AuthConfig, AuthRuntime, LegacyProtocolPolicy, write_admin_key_file, +}; +use pb_mapper_client::client::run_client_side_cli_with_callback; +use pb_mapper_client::server::{ServerTunnelOptions, run_server_side_cli_with_callback}; +use pb_mapper_core::checksum::{Credential, parse_credential, set_process_msg_header_key}; +use pb_mapper_protocol::command::{ + AdminRequest, AdminResponse, LocalServer, MessageSerializer, PbConnRequest, PbConnResponse, + PbConnStatusReq, PbConnStatusResp, PbServerRequest, PbServiceConnStatus, }; -use pb_mapper::common::message::{ - get_header_msg_reader, get_header_msg_writer, MessageReader, MessageWriter, +use pb_mapper_protocol::secure::{ClientHeaderSession, ServerHeaderSession, ServerSecurity}; +use pb_mapper_protocol::{ + MessageReader, MessageWriter, get_header_msg_reader, get_header_msg_writer, }; -use pb_mapper::local::client::run_client_side_cli_with_callback; -use pb_mapper::local::server::{run_server_side_cli_with_callback, ServerTunnelOptions}; -use pb_mapper::pb_server::{get_init_request, run_server_with_shutdown}; +use pb_mapper_server::run_server_with_auth_config; use tokio::io::AsyncReadExt; +use tokio::io::AsyncWriteExt; use tokio::net::{TcpListener, TcpStream}; use tokio::time::timeout; use tokio_util::sync::CancellationToken; @@ -25,19 +36,27 @@ struct EnvVarGuard { } impl EnvVarGuard { + /// # Safety note + /// + /// Mutating the environment is unsafe in edition 2024 because it races + /// concurrent readers. The tests that use this guard set a variable no + /// other test reads, and the guard restores the previous value on drop. fn set(key: &'static str, value: &'static str) -> Self { let old_value = std::env::var(key).ok(); - std::env::set_var(key, value); + unsafe { std::env::set_var(key, value) }; Self { key, old_value } } } impl Drop for EnvVarGuard { fn drop(&mut self) { - if let Some(value) = self.old_value.take() { - std::env::set_var(self.key, value); - } else { - std::env::remove_var(self.key); + // SAFETY: as in `set`. + unsafe { + if let Some(value) = self.old_value.take() { + std::env::set_var(self.key, value); + } else { + std::env::remove_var(self.key); + } } } } @@ -55,6 +74,140 @@ async fn wait_for_server(server_addr: SocketAddr) -> TcpStream { .expect("server did not start") } +#[tokio::test] +async fn explicit_invalid_msg_header_key_fails_server_startup() { + let state_dir = + std::env::temp_dir().join(format!("pb-mapper-invalid-env-{}", std::process::id())); + let _ = std::fs::remove_dir_all(&state_dir); + let mut command = tokio::process::Command::new(env!("CARGO_BIN_EXE_pb-mapper")); + command + .arg("server") + .arg("--port") + .arg("0") + .arg("--auth-state-dir") + .arg(&state_dir) + .env("MSG_HEADER_KEY", "invalid") + .kill_on_drop(true); + + let output = timeout(Duration::from_secs(3), command.output()) + .await + .expect("invalid explicit key must fail instead of starting the server") + .unwrap(); + assert!(!output.status.success()); + let logs = format!( + "{}{}", + String::from_utf8_lossy(&output.stdout), + String::from_utf8_lossy(&output.stderr) + ); + assert!(logs.contains("administrator_key_invalid"), "logs: {logs}"); + assert!(!state_dir.join("admin.key").exists()); + let _ = std::fs::remove_dir_all(state_dir); +} + +#[tokio::test] +async fn admin_all_preserves_json_output_mode() { + let probe_listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let server_addr = probe_listener.local_addr().unwrap(); + drop(probe_listener); + let config = auth_config(server_addr); + let _ = std::fs::remove_dir_all(&config.state_dir); + write_admin_key_file(&config.state_dir.join("admin.key"), TEST_ADMIN_KEY, true).unwrap(); + let runtime = AuthRuntime::start( + *TEST_ADMIN_KEY.as_bytes().first_chunk::<32>().unwrap(), + config.clone(), + ) + .await + .unwrap(); + let admin = runtime + .authenticate_presented( + ADMIN_KEY_ID, + TEST_ADMIN_KEY.as_bytes().first_chunk::<32>().unwrap(), + ) + .unwrap(); + runtime + .issue( + &admin, + Duration::from_secs(120), + Some("json-output".to_string()), + ) + .await + .unwrap(); + drop(runtime); + tokio::time::sleep(Duration::from_millis(20)).await; + + let shutdown = CancellationToken::new(); + let server_shutdown = shutdown.clone(); + let server_config = config.clone(); + let server = tokio::spawn(async move { + run_server_with_auth_config(server_addr, server_shutdown, None, false, server_config) + .await + .unwrap(); + }); + drop(wait_for_server(server_addr).await); + + let mut command = tokio::process::Command::new(env!("CARGO_BIN_EXE_pb-mapper")); + command + .arg("admin") + .arg("--server") + .arg(server_addr.to_string()) + .arg("--output") + .arg("json") + .arg("key") + .arg("list") + .arg("--all") + .env("MSG_HEADER_KEY", TEST_ADMIN_KEY) + .env("RUST_LOG", "off") + .kill_on_drop(true); + let output = timeout(Duration::from_secs(3), command.output()) + .await + .expect("admin JSON request timed out") + .unwrap(); + assert!(output.status.success()); + let document: serde_json::Value = serde_json::from_slice(&output.stdout).unwrap(); + assert_eq!(document["schema_version"], 1); + assert_eq!( + document["data"]["KeyList"]["items"] + .as_array() + .unwrap() + .len(), + 1 + ); + + shutdown.cancel(); + server.await.unwrap(); + let _ = std::fs::remove_dir_all(config.state_dir); +} + +async fn read_secure_request( + security: &ServerSecurity, + stream: &mut TcpStream, +) -> (PbConnRequest, ServerHeaderSession) { + let initial = security.read_initial(stream).await.unwrap(); + ( + PbConnRequest::decode(&initial.payload).unwrap(), + initial.session, + ) +} + +fn auth_config(server_addr: SocketAddr) -> AuthConfig { + static CONFIG_SEQUENCE: AtomicUsize = AtomicUsize::new(0); + + set_process_msg_header_key(Some(TEST_ADMIN_KEY)).unwrap(); + let sequence = CONFIG_SEQUENCE.fetch_add(1, Ordering::Relaxed); + AuthConfig { + state_dir: std::env::temp_dir().join(format!( + "pb-mapper-regression-{}-{}-{sequence}", + std::process::id(), + server_addr.port() + )), + max_temporary_keys: 64, + max_temporary_key_ttl: Duration::from_secs(3600), + legacy_protocol: LegacyProtocolPolicy::Allow, + } +} + +const TEST_ADMIN_KEY: &str = "0123456789abcdefghijklmnopqrstuv"; + async fn register_control_conn_parts( reader: &mut impl MessageReader, writer: &mut impl MessageWriter, @@ -139,20 +292,386 @@ async fn read_status_keys(server_addr: SocketAddr) -> Vec { keys } +async fn send_v2_request( + server_addr: SocketAddr, + credential: &Credential, + request: PbConnRequest, +) -> (TcpStream, ClientHeaderSession, PbConnResponse) { + let mut stream = wait_for_server(server_addr).await; + let session = ClientHeaderSession::new_v2(credential).unwrap(); + session + .write_initial(&mut stream, &request.encode().unwrap()) + .await + .unwrap(); + let response = { + let mut reader = session.response_reader(&mut stream).unwrap(); + let message = timeout(Duration::from_secs(1), reader.read_msg()) + .await + .expect("v2 response timed out") + .unwrap(); + PbConnResponse::decode(message).unwrap() + }; + (stream, session, response) +} + #[tokio::test] -async fn status_service_reports_registered_v2_control_connection() { +async fn temporary_credentials_are_isolated_denied_admin_and_revoked_live() { let probe_listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); let server_addr = probe_listener.local_addr().unwrap(); drop(probe_listener); + let config = auth_config(server_addr); + let _ = std::fs::remove_dir_all(&config.state_dir); + write_admin_key_file(&config.state_dir.join("admin.key"), TEST_ADMIN_KEY, true).unwrap(); + let runtime = AuthRuntime::start( + *TEST_ADMIN_KEY.as_bytes().first_chunk::<32>().unwrap(), + config.clone(), + ) + .await + .unwrap(); + let admin = runtime + .authenticate_presented( + ADMIN_KEY_ID, + TEST_ADMIN_KEY.as_bytes().first_chunk::<32>().unwrap(), + ) + .unwrap(); + let first = runtime + .issue(&admin, Duration::from_secs(120), Some("first".to_string())) + .await + .unwrap(); + let second = runtime + .issue(&admin, Duration::from_secs(120), Some("second".to_string())) + .await + .unwrap(); + drop(runtime); + tokio::time::sleep(Duration::from_millis(20)).await; + let shutdown_token = CancellationToken::new(); let server_shutdown = shutdown_token.clone(); + let server_config = config.clone(); let server = tokio::spawn(async move { - run_server_with_shutdown(server_addr, server_shutdown, None, false) + run_server_with_auth_config(server_addr, server_shutdown, None, false, server_config) .await .unwrap(); }); + let first_credential = parse_credential(&first.credential).unwrap(); + let second_credential = parse_credential(&second.credential).unwrap(); + let register_request = |service: &str| PbConnRequest::Register { + need_codec: false, + is_datagram: false, + key: service.to_string(), + protocol_version: Some(2), + client_instance_id: Some("temporary-auth-regression".to_string()), + heartbeat_interval_ms: Some(50), + heartbeat_tolerance_ms: Some(150), + }; + let (mut first_control, _, first_register) = send_v2_request( + server_addr, + &first_credential, + register_request("same-name"), + ) + .await; + let (_second_control, _, second_register) = send_v2_request( + server_addr, + &second_credential, + register_request("same-name"), + ) + .await; + let first_conn_id = match first_register { + PbConnResponse::RegisterV2 { conn_id, .. } => conn_id, + response => panic!("unexpected first register response: {response:?}"), + }; + let second_conn_id = match second_register { + PbConnResponse::RegisterV2 { conn_id, .. } => conn_id, + response => panic!("unexpected second register response: {response:?}"), + }; + assert_ne!(first_conn_id, second_conn_id); + + let status_request = PbConnRequest::Status(PbConnStatusReq::Service { + key: "same-name".to_string(), + }); + for (credential, expected_conn_id) in [ + (&first_credential, first_conn_id), + (&second_credential, second_conn_id), + ] { + let (_, _, response) = + send_v2_request(server_addr, credential, status_request.clone()).await; + let PbConnResponse::Status(PbConnStatusResp::Service { connections, .. }) = response else { + panic!("unexpected scoped status response: {response:?}"); + }; + assert_eq!(connections.len(), 1); + assert_eq!(connections[0].conn_id, expected_conn_id); + } + + for admin_request in [ + AdminRequest::AuthStatus, + AdminRequest::ServiceList { + key_id: None, + page: 0, + page_size: 100, + }, + AdminRequest::ConnectionList { + key_id: None, + page: 0, + page_size: 100, + }, + ] { + let (_, _, denied) = send_v2_request( + server_addr, + &first_credential, + PbConnRequest::Admin(admin_request), + ) + .await; + let PbConnResponse::Error(denied) = denied else { + panic!("temporary credential unexpectedly received an admin response"); + }; + assert_eq!(denied.code, "admin_permission_required"); + } + + let admin_credential = + Credential::Admin(*TEST_ADMIN_KEY.as_bytes().first_chunk::<32>().unwrap()); + let (_, _, revoked) = send_v2_request( + server_addr, + &admin_credential, + PbConnRequest::Admin(AdminRequest::KeyRevoke { + key_id: first.metadata.key_id.as_u64(), + }), + ) + .await; + assert!(matches!( + revoked, + PbConnResponse::Admin(AdminResponse::KeyRevoked(_)) + )); + let mut byte = [0_u8; 1]; + let closed = timeout(Duration::from_secs(1), first_control.read(&mut byte)) + .await + .expect("revoked control connection was not closed") + .unwrap(); + assert_eq!(closed, 0); + + let (_, _, second_still_active) = + send_v2_request(server_addr, &second_credential, status_request).await; + assert!(matches!( + second_still_active, + PbConnResponse::Status(PbConnStatusResp::Service { .. }) + )); + + shutdown_token.cancel(); + server.await.unwrap(); + let _ = std::fs::remove_dir_all(config.state_dir); +} + +#[tokio::test] +async fn revoking_subscriber_credential_closes_cross_credential_data_stream() { + let probe_listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let server_addr = probe_listener.local_addr().unwrap(); + drop(probe_listener); + + let config = auth_config(server_addr); + let _ = std::fs::remove_dir_all(&config.state_dir); + write_admin_key_file(&config.state_dir.join("admin.key"), TEST_ADMIN_KEY, true).unwrap(); + let runtime = AuthRuntime::start( + *TEST_ADMIN_KEY.as_bytes().first_chunk::<32>().unwrap(), + config.clone(), + ) + .await + .unwrap(); + let admin = runtime + .authenticate_presented( + ADMIN_KEY_ID, + TEST_ADMIN_KEY.as_bytes().first_chunk::<32>().unwrap(), + ) + .unwrap(); + let issued = runtime + .issue( + &admin, + Duration::from_secs(120), + Some("active-stream".to_string()), + ) + .await + .unwrap(); + drop(runtime); + tokio::time::sleep(Duration::from_millis(20)).await; + + let shutdown_token = CancellationToken::new(); + let server_shutdown = shutdown_token.clone(); + let server_config = config.clone(); + let server = tokio::spawn(async move { + run_server_with_auth_config(server_addr, server_shutdown, None, false, server_config) + .await + .unwrap(); + }); + + let credential = parse_credential(&issued.credential).unwrap(); + let admin_credential = + Credential::Admin(*TEST_ADMIN_KEY.as_bytes().first_chunk::<32>().unwrap()); + let service = "revoked-stream"; + let mut control = wait_for_server(server_addr).await; + let control_session = ClientHeaderSession::new_v2(&admin_credential).unwrap(); + let register = PbConnRequest::RegisterScoped { + need_codec: false, + is_datagram: false, + key: service.to_string(), + namespace: issued.metadata.key_id.as_u64(), + force_namespace: true, + protocol_version: Some(2), + client_instance_id: Some("active-stream-test".to_string()), + heartbeat_interval_ms: Some(5_000), + heartbeat_tolerance_ms: Some(15_000), + }; + control_session + .write_initial(&mut control, ®ister.encode().unwrap()) + .await + .unwrap(); + let (mut control_read, mut control_write) = control.into_split(); + let mut control_reader = control_session.response_reader(&mut control_read).unwrap(); + let register_response = timeout(Duration::from_secs(1), control_reader.read_msg()) + .await + .expect("register response timed out") + .unwrap(); + assert!(matches!( + PbConnResponse::decode(register_response).unwrap(), + PbConnResponse::RegisterV2 { .. } + )); + + let mut subscriber = wait_for_server(server_addr).await; + let subscriber_session = ClientHeaderSession::new_v2(&credential).unwrap(); + subscriber_session + .write_initial( + &mut subscriber, + &PbConnRequest::Subcribe { + key: service.to_string(), + } + .encode() + .unwrap(), + ) + .await + .unwrap(); + + let stream_request = timeout(Duration::from_secs(1), control_reader.read_msg()) + .await + .expect("stream request timed out") + .unwrap(); + let LocalServer::Stream { + client_id, + server_generation, + } = LocalServer::decode(stream_request).unwrap() + else { + panic!("unexpected local server stream request"); + }; + let mut control_writer = control_session + .continuation_writer(&mut control_write) + .unwrap(); + control_writer + .write_msg( + &PbServerRequest::StreamAck { + client_id, + server_generation, + } + .encode() + .unwrap(), + ) + .await + .unwrap(); + + let mut provider = wait_for_server(server_addr).await; + let provider_session = ClientHeaderSession::new_v2(&admin_credential).unwrap(); + provider_session + .write_initial( + &mut provider, + &PbConnRequest::StreamScoped { + key: service.to_string(), + namespace: issued.metadata.key_id.as_u64(), + dst_id: client_id, + server_generation, + } + .encode() + .unwrap(), + ) + .await + .unwrap(); + { + let mut subscriber_reader = subscriber_session.response_reader(&mut subscriber).unwrap(); + let response = timeout(Duration::from_secs(1), subscriber_reader.read_msg()) + .await + .expect("subscribe response timed out") + .unwrap(); + assert!(matches!( + PbConnResponse::decode(response).unwrap(), + PbConnResponse::Subcribe { .. } + )); + } + { + let mut provider_reader = provider_session.response_reader(&mut provider).unwrap(); + let response = timeout(Duration::from_secs(1), provider_reader.read_msg()) + .await + .expect("provider stream response timed out") + .unwrap(); + assert!(matches!( + PbConnResponse::decode(response).unwrap(), + PbConnResponse::Stream { .. } + )); + } + + subscriber.write_all(b"ready").await.unwrap(); + let mut ready = [0_u8; 5]; + timeout(Duration::from_secs(1), provider.read_exact(&mut ready)) + .await + .expect("active data stream did not forward") + .unwrap(); + assert_eq!(&ready, b"ready"); + + let (_, _, revoked) = send_v2_request( + server_addr, + &admin_credential, + PbConnRequest::Admin(AdminRequest::KeyRevoke { + key_id: issued.metadata.key_id.as_u64(), + }), + ) + .await; + assert!(matches!( + revoked, + PbConnResponse::Admin(AdminResponse::KeyRevoked(_)) + )); + + let mut byte = [0_u8; 1]; + let read = timeout(Duration::from_secs(1), subscriber.read(&mut byte)) + .await + .expect("revoked subscriber data stream was not closed") + .unwrap(); + assert_eq!(read, 0, "revoked subscriber data stream remained open"); + + match timeout(Duration::from_millis(200), control_reader.read_msg()).await { + Err(_) | Ok(Ok(_)) => {} + Ok(Err(error)) => panic!("administrator registration was cancelled too: {error}"), + } + + shutdown_token.cancel(); + server.await.unwrap(); + let _ = std::fs::remove_dir_all(config.state_dir); +} + +#[tokio::test] +async fn status_service_reports_registered_v2_control_connection() { + let probe_listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let server_addr = probe_listener.local_addr().unwrap(); + drop(probe_listener); + + let shutdown_token = CancellationToken::new(); + let server_shutdown = shutdown_token.clone(); + let server = tokio::spawn(async move { + run_server_with_auth_config( + server_addr, + server_shutdown, + None, + false, + auth_config(server_addr), + ) + .await + .unwrap(); + }); + let key = "sf-backend"; let control = wait_for_server(server_addr).await; let (mut reader_stream, mut writer_stream) = control.into_split(); @@ -210,15 +729,19 @@ async fn local_server_reconnects_when_registered_conn_is_missing_from_remote_sta let fake_register_count = register_count.clone(); let fake_second_register_tx = second_register_tx.clone(); + let fake_security = ServerSecurity::new( + AuthRuntime::from_process(auth_config(remote_addr)) + .await + .unwrap(), + ); let fake_server = tokio::spawn(async move { loop { let (mut stream, _) = remote_listener.accept().await.unwrap(); let register_count = fake_register_count.clone(); let second_register_tx = fake_second_register_tx.clone(); + let security = fake_security.clone(); tokio::spawn(async move { - let Ok(request) = get_init_request(&mut stream, 0.into()).await else { - return; - }; + let (request, session) = read_secure_request(&security, &mut stream).await; match request { PbConnRequest::Register { key, .. } => { let count = register_count.fetch_add(1, Ordering::SeqCst) + 1; @@ -229,12 +752,12 @@ async fn local_server_reconnects_when_registered_conn_is_missing_from_remote_sta } .encode() .unwrap(); - let mut writer = get_header_msg_writer(&mut stream).unwrap(); + let mut writer = session.response_writer(&mut stream).unwrap(); writer.write_msg(&response).await.unwrap(); - if count == 2 { - if let Some(tx) = second_register_tx.lock().await.take() { - tx.send(()).unwrap(); - } + if count == 2 + && let Some(tx) = second_register_tx.lock().await.take() + { + tx.send(()).unwrap(); } tracing::debug!(key, count, "fake server accepted register"); std::future::pending::<()>().await; @@ -246,14 +769,14 @@ async fn local_server_reconnects_when_registered_conn_is_missing_from_remote_sta }) .encode() .unwrap(); - let mut writer = get_header_msg_writer(&mut stream).unwrap(); + let mut writer = session.response_writer(&mut stream).unwrap(); writer.write_msg(&response).await.unwrap(); } PbConnRequest::Status(PbConnStatusReq::Keys) => { let response = PbConnResponse::Status(PbConnStatusResp::Keys(Vec::new())) .encode() .unwrap(); - let mut writer = get_header_msg_writer(&mut stream).unwrap(); + let mut writer = session.response_writer(&mut stream).unwrap(); writer.write_msg(&response).await.unwrap(); } _ => {} @@ -271,6 +794,8 @@ async fn local_server_reconnects_when_registered_conn_is_missing_from_remote_sta need_codec: false, is_datagram: false, keep_alive: false, + namespace: None, + force_namespace: false, }, None, )); @@ -288,10 +813,15 @@ async fn local_server_reconnects_when_registered_conn_is_missing_from_remote_sta async fn client_closes_initial_status_probe_after_key_check() { let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); let remote_addr = listener.local_addr().unwrap(); + let security = ServerSecurity::new( + AuthRuntime::from_process(auth_config(remote_addr)) + .await + .unwrap(), + ); let fake_server = tokio::spawn(async move { let (mut stream, _) = listener.accept().await.unwrap(); - let request = get_init_request(&mut stream, 0.into()).await.unwrap(); + let (request, session) = read_secure_request(&security, &mut stream).await; let PbConnRequest::Status(PbConnStatusReq::Service { key }) = request else { panic!("client did not use service status for initial key check"); }; @@ -309,7 +839,7 @@ async fn client_closes_initial_status_probe_after_key_check() { .encode() .unwrap(); { - let mut writer = get_header_msg_writer(&mut stream).unwrap(); + let mut writer = session.response_writer(&mut stream).unwrap(); writer.write_msg(&response).await.unwrap(); } @@ -351,15 +881,19 @@ async fn client_tolerates_one_failed_health_check_while_listener_is_active() { let fake_failed_status_responses = failed_status_responses.clone(); let fake_status_count = status_count.clone(); + let fake_security = ServerSecurity::new( + AuthRuntime::from_process(auth_config(remote_addr)) + .await + .unwrap(), + ); let fake_server = tokio::spawn(async move { loop { let (mut stream, _) = remote_listener.accept().await.unwrap(); let failed_status_responses = fake_failed_status_responses.clone(); let status_count = fake_status_count.clone(); + let security = fake_security.clone(); tokio::spawn(async move { - let Ok(request) = get_init_request(&mut stream, 0.into()).await else { - return; - }; + let (request, session) = read_secure_request(&security, &mut stream).await; match request { PbConnRequest::Status(PbConnStatusReq::Service { key }) => { status_count.fetch_add(1, Ordering::SeqCst); @@ -383,7 +917,7 @@ async fn client_tolerates_one_failed_health_check_while_listener_is_active() { PbConnResponse::Status(PbConnStatusResp::Service { key, connections }) .encode() .unwrap(); - let mut writer = get_header_msg_writer(&mut stream).unwrap(); + let mut writer = session.response_writer(&mut stream).unwrap(); writer.write_msg(&response).await.unwrap(); } PbConnRequest::Status(PbConnStatusReq::Keys) => { @@ -401,7 +935,7 @@ async fn client_tolerates_one_failed_health_check_while_listener_is_active() { let response = PbConnResponse::Status(PbConnStatusResp::Keys(keys)) .encode() .unwrap(); - let mut writer = get_header_msg_writer(&mut stream).unwrap(); + let mut writer = session.response_writer(&mut stream).unwrap(); writer.write_msg(&response).await.unwrap(); } PbConnRequest::Subcribe { .. } => std::future::pending().await, @@ -484,9 +1018,15 @@ async fn subscribe_retires_unacked_control_connection() { let shutdown_token = CancellationToken::new(); let server_shutdown = shutdown_token.clone(); let server = tokio::spawn(async move { - run_server_with_shutdown(server_addr, server_shutdown, None, false) - .await - .unwrap(); + run_server_with_auth_config( + server_addr, + server_shutdown, + None, + false, + auth_config(server_addr), + ) + .await + .unwrap(); }); let key = "sf-backend"; @@ -548,9 +1088,15 @@ async fn subscribe_waits_for_replacement_after_retiring_stale_control_connection let shutdown_token = CancellationToken::new(); let server_shutdown = shutdown_token.clone(); let server = tokio::spawn(async move { - run_server_with_shutdown(server_addr, server_shutdown, None, false) - .await - .unwrap(); + run_server_with_auth_config( + server_addr, + server_shutdown, + None, + false, + auth_config(server_addr), + ) + .await + .unwrap(); }); let key = "sf-backend"; @@ -667,9 +1213,15 @@ async fn subscribe_missing_key_closes_without_hanging() { let shutdown_token = CancellationToken::new(); let server_shutdown = shutdown_token.clone(); let server = tokio::spawn(async move { - run_server_with_shutdown(server_addr, server_shutdown, None, false) - .await - .unwrap(); + run_server_with_auth_config( + server_addr, + server_shutdown, + None, + false, + auth_config(server_addr), + ) + .await + .unwrap(); }); let mut stream = timeout(Duration::from_secs(2), async { @@ -697,7 +1249,11 @@ async fn subscribe_missing_key_closes_without_hanging() { let result = timeout(Duration::from_millis(200), reader.read_msg()) .await .expect("missing-key subscribe hung instead of closing"); - assert!(result.is_err()); + let response = PbConnResponse::decode(result.unwrap()).unwrap(); + let PbConnResponse::Error(error) = response else { + panic!("expected structured missing-service error"); + }; + assert_eq!(error.code, "service_not_available"); shutdown_token.cancel(); server.await.unwrap(); @@ -712,9 +1268,15 @@ async fn subscribe_bypasses_unacked_stale_control_connection() { let shutdown_token = CancellationToken::new(); let server_shutdown = shutdown_token.clone(); let server = tokio::spawn(async move { - run_server_with_shutdown(server_addr, server_shutdown, None, false) - .await - .unwrap(); + run_server_with_auth_config( + server_addr, + server_shutdown, + None, + false, + auth_config(server_addr), + ) + .await + .unwrap(); }); let key = "sf-backend"; @@ -817,9 +1379,15 @@ async fn subscribe_bypasses_acked_control_connection_without_stream() { let shutdown_token = CancellationToken::new(); let server_shutdown = shutdown_token.clone(); let server = tokio::spawn(async move { - run_server_with_shutdown(server_addr, server_shutdown, None, false) - .await - .unwrap(); + run_server_with_auth_config( + server_addr, + server_shutdown, + None, + false, + auth_config(server_addr), + ) + .await + .unwrap(); }); let key = "sf-backend"; diff --git a/tests/test_delay.rs b/crates/pb-mapper-cli/tests/test_delay.rs similarity index 90% rename from tests/test_delay.rs rename to crates/pb-mapper-cli/tests/test_delay.rs index 533f330..78c747b 100644 --- a/tests/test_delay.rs +++ b/crates/pb-mapper-cli/tests/test_delay.rs @@ -1,18 +1,21 @@ +// See the note in `regression.rs`: the whole file is test code. +#![allow(clippy::unwrap_used, clippy::expect_used)] + use std::env; use std::sync::LazyLock; use std::time::Duration; -use pb_mapper::common::config::init_tracing; -use pb_mapper::common::message::{ - MessageReader, MessageWriter, NormalMessageReader, NormalMessageWriter, -}; -use pb_mapper::local::client::run_client_side_cli; -use pb_mapper::local::server::{run_server_side_cli, ServerTunnelOptions}; -use pb_mapper::pb_server::run_server; +use pb_mapper_auth::{AuthConfig, LegacyProtocolPolicy}; +use pb_mapper_client::client::run_client_side_cli; +use pb_mapper_client::server::{ServerTunnelOptions, run_server_side_cli}; +use pb_mapper_core::config::init_tracing; +use pb_mapper_protocol::{MessageReader, MessageWriter, NormalMessageReader, NormalMessageWriter}; +use pb_mapper_server::run_server_with_auth_config; use rand::RngExt; use tokio::io::{AsyncReadExt, AsyncWriteExt}; use tokio::net::UdpSocket; -use tokio::time::{timeout, Instant}; +use tokio::time::{Instant, timeout}; +use tokio_util::sync::CancellationToken; use uni_stream::addr::ToSocketAddrs; use uni_stream::stream::{ListenerProvider, TcpListenerProvider, UdpListenerProvider}; use uni_stream::stream::{StreamProvider, StreamSplit, TcpStreamProvider, UdpStreamProvider}; @@ -109,7 +112,20 @@ async fn run_udp_echo_server(addr: &str) -> Result<(), Box().ok()) + .unwrap_or_default(); + let auth_config = AuthConfig { + state_dir: std::env::temp_dir() + .join(format!("pb-mapper-delay-{}-{port}", std::process::id())), + max_temporary_keys: 64, + max_temporary_key_ttl: Duration::from_secs(3600), + legacy_protocol: LegacyProtocolPolicy::Allow, + }; + if let Err(e) = + run_server_with_auth_config(addr, CancellationToken::new(), None, false, auth_config).await + { eprintln!("pb-mapper server failed to start: {e}"); } } @@ -131,6 +147,8 @@ async fn run_pb_mapper_server_cli( need_codec, is_datagram: true, keep_alive: false, + namespace: None, + force_namespace: false, }, ) .await @@ -144,6 +162,8 @@ async fn run_pb_mapper_server_cli( need_codec, is_datagram: false, keep_alive: false, + namespace: None, + force_namespace: false, }, ) .await @@ -200,11 +220,11 @@ async fn run_udp_datagram_echo(addr: &str, rounds: usize, burst: usize) { let mut ready = false; for _ in 0..10 { socket.send(probe).await.unwrap(); - if let Ok(Ok(len)) = timeout(Duration::from_millis(300), socket.recv(&mut buf)).await { - if &buf[..len] == probe { - ready = true; - break; - } + if let Ok(Ok(len)) = timeout(Duration::from_millis(300), socket.recv(&mut buf)).await + && &buf[..len] == probe + { + ready = true; + break; } tokio::time::sleep(Duration::from_millis(50)).await; } diff --git a/crates/pb-mapper-client/Cargo.toml b/crates/pb-mapper-client/Cargo.toml new file mode 100644 index 0000000..771b2e8 --- /dev/null +++ b/crates/pb-mapper-client/Cargo.toml @@ -0,0 +1,21 @@ +[package] +name = "pb-mapper-client" +version.workspace = true +edition.workspace = true +authors.workspace = true + +[dependencies] +pb-mapper-core.workspace = true +pb-mapper-protocol.workspace = true + +serde_json.workspace = true +snafu.workspace = true +tokio.workspace = true +tracing.workspace = true +uni-stream.workspace = true + +[features] +udp-timeout = ["uni-stream/udp-timeout", "pb-mapper-protocol/udp-timeout"] + +[lints] +workspace = true diff --git a/src/local/client/error.rs b/crates/pb-mapper-client/src/client/error.rs similarity index 93% rename from src/local/client/error.rs rename to crates/pb-mapper-client/src/client/error.rs index 7a83d16..bdae3c3 100644 --- a/src/local/client/error.rs +++ b/crates/pb-mapper-client/src/client/error.rs @@ -2,7 +2,9 @@ use std::time::Duration; use snafu::Snafu; -use crate::common; +// The `common::error::Error` spellings below are the source type on nearly every +// variant; aliasing keeps them as they were. +use pb_mapper_core as common; #[derive(Debug, Snafu)] #[snafu(visibility(pub(super)))] diff --git a/src/local/client/mod.rs b/crates/pb-mapper-client/src/client/mod.rs similarity index 72% rename from src/local/client/mod.rs rename to crates/pb-mapper-client/src/client/mod.rs index 697cd40..768f5d5 100644 --- a/src/local/client/mod.rs +++ b/crates/pb-mapper-client/src/client/mod.rs @@ -13,17 +13,18 @@ use tokio::time::MissedTickBehavior; use uni_stream::udp::set_custom_timeout; use self::error::{AcceptLocalStreamSnafu, BindLocalListenerSnafu}; -use self::status::get_status; +use self::status::{get_status, get_status_scoped, get_status_with_credential}; use self::stream::handle_local_stream; -use crate::common::config::{ - client_health_check_interval, client_health_check_timeout, client_health_failure_threshold, - StatusOp, +use pb_mapper_core::checksum::{Credential, get_process_credential}; +use pb_mapper_core::config::{ + StatusOp, client_health_check_interval, client_health_check_timeout, + client_health_failure_threshold, }; -use crate::common::message::command::{PbConnStatusReq, PbConnStatusResp}; -use crate::common::message::forward::StreamForward; -use crate::snafu_error_get_or_return; -use crate::utils::timeout::RetryBackoff; -use uni_stream::addr::{each_addr, ToSocketAddrs}; +use pb_mapper_core::snafu_error_get_or_return; +use pb_mapper_core::timeout::RetryBackoff; +use pb_mapper_protocol::command::{PbConnStatusReq, PbConnStatusResp}; +use pb_mapper_protocol::forward::StreamForward; +use uni_stream::addr::{ToSocketAddrs, each_addr}; use uni_stream::stream::got_one_socket_addr; use uni_stream::stream::{ListenerProvider, StreamAccept}; @@ -48,6 +49,27 @@ pub async fn run_client_side_cli( + local_addr: A, + remote_addr: A, + key: Arc, + keep_alive: bool, + namespace: Option, +) where + ::Item: StreamForward, +{ + run_client_side_cli_with_callback_scoped::( + local_addr, + remote_addr, + key, + keep_alive, + namespace, + None, + None, + ) + .await +} + pub async fn run_client_side_cli_with_callback( local_addr: A, remote_addr: A, @@ -56,6 +78,57 @@ pub async fn run_client_side_cli_with_callback, ) where ::Item: StreamForward, +{ + run_client_side_cli_with_callback_scoped::( + local_addr, + remote_addr, + key, + keep_alive, + None, + status_callback, + None, + ) + .await +} + +pub async fn run_client_side_cli_with_pinned_credential< + LocalListener: ListenerProvider, + A: ToSocketAddrs, +>( + local_addr: A, + remote_addr: A, + key: Arc, + keep_alive: bool, + status_callback: Option, + credential: pb_mapper_core::checksum::Credential, +) where + ::Item: StreamForward, +{ + run_client_side_cli_with_callback_scoped::( + local_addr, + remote_addr, + key, + keep_alive, + None, + status_callback, + Some(credential), + ) + .await +} + +pub async fn run_client_side_cli_with_callback_scoped< + LocalListener: ListenerProvider, + A: ToSocketAddrs, +>( + local_addr: A, + remote_addr: A, + key: Arc, + keep_alive: bool, + namespace: Option, + status_callback: Option, + pinned_credential: Option, +) where + ::Item: StreamForward, { set_custom_timeout(Duration::from_secs(120)); @@ -79,6 +152,19 @@ pub async fn run_client_side_cli_with_callback credential, + None => match get_process_credential() { + Ok(credential) => credential, + Err(e) => { + tracing::error!("load client credential failed: {e}"); + if let Some(ref callback) = status_callback { + callback("failed"); + } + return; + } + }, + }; let mut retry_backoff = RetryBackoff::default(); @@ -92,7 +178,9 @@ pub async fn run_client_side_cli_with_callback { - if let Err(reason) = probe_remote_key(remote_addr, key.as_ref()).await { + if let Err(reason) = probe_remote_key(remote_addr, key.as_ref(), namespace, credential).await { consecutive_health_failures = consecutive_health_failures.saturating_add(1); if consecutive_health_failures < health_failure_threshold { tracing::warn!( @@ -234,7 +322,7 @@ pub async fn run_client_side_cli_with_callback std::result::Result<(), String> { +async fn probe_remote_key( + remote_addr: SocketAddr, + key: &str, + namespace: Option, + credential: Credential, +) -> std::result::Result<(), String> { let timeout = client_health_check_timeout(); - match tokio::time::timeout(timeout, probe_remote_key_once(remote_addr, key)).await { + match tokio::time::timeout( + timeout, + probe_remote_key_once(remote_addr, key, namespace, credential), + ) + .await + { Ok(result) => result, Err(_) => Err(format!("remote key probe timed out after {timeout:?}")), } @@ -279,12 +377,16 @@ async fn probe_remote_key(remote_addr: SocketAddr, key: &str) -> std::result::Re async fn probe_remote_key_once( remote_addr: SocketAddr, key: &str, + namespace: Option, + credential: Credential, ) -> std::result::Result<(), String> { match fetch_remote_status( remote_addr, PbConnStatusReq::Service { key: key.to_string(), }, + namespace, + credential, ) .await { @@ -312,7 +414,8 @@ async fn probe_remote_key_once( } } - let status_resp = fetch_remote_status(remote_addr, PbConnStatusReq::Keys).await?; + let status_resp = + fetch_remote_status(remote_addr, PbConnStatusReq::Keys, namespace, credential).await?; let PbConnStatusResp::Keys(keys) = status_resp else { return Err(format!( "expected keys status response, got {status_resp:?}" @@ -330,11 +433,13 @@ async fn probe_remote_key_once( async fn fetch_remote_status( remote_addr: SocketAddr, req: PbConnStatusReq, + namespace: Option, + credential: Credential, ) -> std::result::Result { let mut stream = each_addr(remote_addr, TcpStream::connect) .await .map_err(|e| format!("connect remote stream failed: {e}"))?; - get_status(&mut stream, req) + get_status_with_credential(&mut stream, req, namespace, &credential) .await .map_err(|e| format!("get status failed: {}", snafu::Report::from_error(e))) } @@ -357,8 +462,32 @@ pub async fn handle_status_cli op: StatusOp, addr: A, ) { + if let Err(error) = handle_status_cli_scoped(op, addr, None).await { + tracing::error!("{error}"); + } +} + +pub async fn handle_status_cli_scoped( + op: StatusOp, + addr: A, + namespace: Option, +) -> Result<(), Box> { match op { - StatusOp::RemoteId => show_status(addr, PbConnStatusReq::RemoteId).await, - StatusOp::Keys => show_status(addr, PbConnStatusReq::Keys).await, + StatusOp::RemoteId => show_status_scoped(addr, PbConnStatusReq::RemoteId, namespace).await, + StatusOp::Keys => show_status_scoped(addr, PbConnStatusReq::Keys, namespace).await, } } + +pub async fn show_status_scoped( + remote_addr: A, + req: PbConnStatusReq, + namespace: Option, +) -> Result<(), Box> { + let mut stream = each_addr(remote_addr, TcpStream::connect) + .await + .map_err(|error| format!("get status stream: {error}"))?; + let status = get_status_scoped(&mut stream, req, namespace).await?; + let status = serde_json::to_string_pretty(&status)?; + println!("Status:{status}"); + Ok(()) +} diff --git a/crates/pb-mapper-client/src/client/status.rs b/crates/pb-mapper-client/src/client/status.rs new file mode 100644 index 0000000..6f686d7 --- /dev/null +++ b/crates/pb-mapper-client/src/client/status.rs @@ -0,0 +1,99 @@ +use snafu::ResultExt; +use tokio::io::{AsyncReadExt, AsyncWriteExt}; + +use super::error::{ + CreateHeaderToolSnafu, DecodeStatusRespSnafu, EncodeStatusReqSnafu, StatusRespNotMatchSnafu, + WriteStatusReqSnafu, +}; +use pb_mapper_core::checksum::Credential; +use pb_mapper_core::config::control_io_timeout; +use pb_mapper_protocol::command::{ + MessageSerializer, PbConnRequest, PbConnResponse, PbConnStatusReq, PbConnStatusResp, +}; +use pb_mapper_protocol::secure::ClientHeaderSession; + +pub async fn get_status( + remote_stream: &mut S, + req: PbConnStatusReq, +) -> super::error::Result { + get_status_scoped(remote_stream, req, None).await +} + +pub async fn get_status_scoped( + remote_stream: &mut S, + req: PbConnStatusReq, + namespace: Option, +) -> super::error::Result { + let session = + ClientHeaderSession::from_process().context(CreateHeaderToolSnafu { action: "session" })?; + get_status_with_session(remote_stream, req, namespace, session).await +} + +pub async fn get_status_with_credential( + remote_stream: &mut S, + req: PbConnStatusReq, + namespace: Option, + credential: &Credential, +) -> super::error::Result { + let session = ClientHeaderSession::new_v2(credential) + .context(CreateHeaderToolSnafu { action: "session" })?; + get_status_with_session(remote_stream, req, namespace, session).await +} + +async fn get_status_with_session( + remote_stream: &mut S, + req: PbConnStatusReq, + namespace: Option, + session: ClientHeaderSession, +) -> super::error::Result { + let timeout = control_io_timeout(); + let request = match namespace { + Some(namespace) => PbConnRequest::StatusScoped { + status: req, + namespace, + }, + None => PbConnRequest::Status(req), + }; + let msg = request.encode().context(EncodeStatusReqSnafu)?; + let response = session + .exchange(remote_stream, &msg, timeout) + .await + .context(WriteStatusReqSnafu)?; + let resp = PbConnResponse::decode(&response).context(DecodeStatusRespSnafu)?; + match resp { + PbConnResponse::Status(status) => Ok(status), + PbConnResponse::Error(error) => StatusRespNotMatchSnafu { + resp: format!("{}: {}", error.code, error.message), + } + .fail(), + other => StatusRespNotMatchSnafu { + resp: format!("{other:?}"), + } + .fail(), + } +} + +#[cfg(test)] +mod tests { + use std::time::Duration; + + use super::*; + + #[tokio::test] + async fn get_status_times_out_when_peer_stalls_after_request() { + // SAFETY: no other thread in this test reads the environment. + unsafe { std::env::set_var("PB_MAPPER_CONTROL_IO_TIMEOUT", "20ms") }; + let (mut client, _server) = tokio::io::duplex(1024); + + let result = tokio::time::timeout( + Duration::from_millis(200), + get_status(&mut client, PbConnStatusReq::Keys), + ) + .await + .expect("get_status ignored PB_MAPPER_CONTROL_IO_TIMEOUT"); + + // SAFETY: as above. + unsafe { std::env::remove_var("PB_MAPPER_CONTROL_IO_TIMEOUT") }; + assert!(result.is_err()); + } +} diff --git a/src/local/client/stream.rs b/crates/pb-mapper-client/src/client/stream.rs similarity index 50% rename from src/local/client/stream.rs rename to crates/pb-mapper-client/src/client/stream.rs index e978536..ef8bf6b 100644 --- a/src/local/client/stream.rs +++ b/crates/pb-mapper-client/src/client/stream.rs @@ -6,20 +6,18 @@ use tokio::net::TcpStream; use tracing::{info_span, instrument}; use super::error::{ - ConnectRemoteStreamSnafu, ControlIoTimeoutSnafu, DecodeSubcribeRespSnafu, - EncodeSubcribeReqSnafu, ReadSubcribeRespSnafu, Result, SubcribeRespNotMatchSnafu, - WriteSubcribeReqSnafu, + ConnectRemoteStreamSnafu, DecodeSubcribeRespSnafu, EncodeSubcribeReqSnafu, Result, + SubcribeRespNotMatchSnafu, WriteSubcribeReqSnafu, }; -use crate::common::config::control_io_timeout; -use crate::common::message::command::{MessageSerializer, PbConnRequest, PbConnResponse}; -use crate::common::message::forward::StreamForward; -use crate::common::message::{ - get_header_msg_reader, get_header_msg_writer, MessageReader, MessageWriter, -}; -use crate::local::client::error::CreateHeaderToolSnafu; -use crate::snafu_error_handle; -use uni_stream::addr::{each_addr, ToSocketAddrs}; -use uni_stream::stream::{set_tcp_keep_alive, set_tcp_nodelay, NetworkStream}; +use crate::client::error::CreateHeaderToolSnafu; +use pb_mapper_core::checksum::Credential; +use pb_mapper_core::config::control_io_timeout; +use pb_mapper_core::snafu_error_handle; +use pb_mapper_protocol::command::{MessageSerializer, PbConnRequest, PbConnResponse}; +use pb_mapper_protocol::forward::StreamForward; +use pb_mapper_protocol::secure::ClientHeaderSession; +use uni_stream::addr::{ToSocketAddrs, each_addr}; +use uni_stream::stream::{NetworkStream, set_tcp_keep_alive, set_tcp_nodelay}; #[instrument(skip(local_stream))] pub async fn handle_local_stream< @@ -30,6 +28,8 @@ pub async fn handle_local_stream< key: Arc, remote_addr: A, keep_alive: bool, + namespace: Option, + credential: Credential, ) -> Result<()> { let mut remote_stream = each_addr(remote_addr, TcpStream::connect) .await @@ -47,39 +47,33 @@ pub async fn handle_local_stream< let (codec_key, client_id, server_id) = { let timeout = control_io_timeout(); // handle request - let msg = PbConnRequest::Subcribe { - key: key.to_string(), - } - .encode() - .context(EncodeSubcribeReqSnafu)?; - let mut msg_writer = get_header_msg_writer(&mut remote_stream) - .context(CreateHeaderToolSnafu { action: "writer" })?; - match tokio::time::timeout(timeout, msg_writer.write_msg(&msg)).await { - Ok(result) => result.context(WriteSubcribeReqSnafu)?, - Err(_) => ControlIoTimeoutSnafu { - action: "write subcribe request", - timeout, - } - .fail()?, - } - // handle response - let mut msg_reader = get_header_msg_reader(&mut remote_stream) - .context(CreateHeaderToolSnafu { action: "reader" })?; - let msg = match tokio::time::timeout(timeout, msg_reader.read_msg()).await { - Ok(result) => result.context(ReadSubcribeRespSnafu)?, - Err(_) => ControlIoTimeoutSnafu { - action: "read subcribe response", - timeout, - } - .fail()?, + let request = match namespace { + Some(namespace) => PbConnRequest::SubcribeScoped { + key: key.to_string(), + namespace, + }, + None => PbConnRequest::Subcribe { + key: key.to_string(), + }, }; - let resp = PbConnResponse::decode(msg).context(DecodeSubcribeRespSnafu)?; + let msg = request.encode().context(EncodeSubcribeReqSnafu)?; + let session = ClientHeaderSession::new_v2(&credential) + .context(CreateHeaderToolSnafu { action: "session" })?; + let response = session + .exchange(&mut remote_stream, &msg, timeout) + .await + .context(WriteSubcribeReqSnafu)?; + let resp = PbConnResponse::decode(&response).context(DecodeSubcribeRespSnafu)?; match resp { PbConnResponse::Subcribe { codec_key, client_id, server_id, } => (codec_key, client_id, server_id), + PbConnResponse::Error(error) => SubcribeRespNotMatchSnafu { + resp: format!("{}: {}", error.code, error.message), + } + .fail()?, resp => SubcribeRespNotMatchSnafu { resp: format!("{resp:?}"), } @@ -95,6 +89,7 @@ pub async fn handle_local_stream< snafu_error_handle!( ::forward_local_to_remote( codec_key, + *credential.key(), client_reader, client_writer, server_reader, diff --git a/src/local/mod.rs b/crates/pb-mapper-client/src/lib.rs similarity index 100% rename from src/local/mod.rs rename to crates/pb-mapper-client/src/lib.rs diff --git a/src/local/server/error.rs b/crates/pb-mapper-client/src/server/error.rs similarity index 94% rename from src/local/server/error.rs rename to crates/pb-mapper-client/src/server/error.rs index 1029071..19bdd4c 100644 --- a/src/local/server/error.rs +++ b/crates/pb-mapper-client/src/server/error.rs @@ -2,7 +2,9 @@ use std::time::Duration; use snafu::Snafu; -use crate::common::{self}; +// The `common::error::Error` spellings below are the source type on nearly every +// variant; aliasing keeps them as they were. +use pb_mapper_core as common; #[derive(Debug, Snafu)] #[snafu(visibility(pub(super)))] diff --git a/src/local/server/mod.rs b/crates/pb-mapper-client/src/server/mod.rs similarity index 79% rename from src/local/server/mod.rs rename to crates/pb-mapper-client/src/server/mod.rs index 9fca9df..b47ed2c 100644 --- a/src/local/server/mod.rs +++ b/crates/pb-mapper-client/src/server/mod.rs @@ -8,6 +8,7 @@ use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH}; use snafu::ResultExt; use tokio::net::TcpStream; +use tokio::task::JoinSet; use tokio::time::MissedTickBehavior; use tracing::instrument; @@ -16,27 +17,27 @@ use self::error::{ EncodeRegisterReqSnafu, EncodeStreamAckMsgSnafu, ReadRegisterRespSnafu, ReadStreamReqSnafu, RegisterRespNotMatchSnafu, SendRegisterReqSnafu, WritePingMsgSnafu, WriteStreamAckMsgSnafu, }; -use self::stream::handle_stream; -use crate::common::config::{ +use self::stream::{StreamConnect, handle_stream}; +use pb_mapper_core::checksum::{Credential, get_process_credential}; +use pb_mapper_core::config::{ control_conn_pool_size, control_heartbeat_interval, control_heartbeat_tolerance, control_io_timeout, control_suspect_grace, registration_probe_timeout, }; -use crate::common::message::command::{ - LocalServer, MessageSerializer, PbConnRequest, PbConnResponse, PbConnStatusReq, - PbConnStatusResp, PbServerRequest, CONTROL_PROTOCOL_V2, -}; -use crate::common::message::forward::StreamForward; -use crate::common::message::{ - get_header_msg_reader, get_header_msg_writer, MessageReader, MessageWriter, -}; -use crate::utils::timeout::RetryBackoff; -use crate::{ +use pb_mapper_core::timeout::RetryBackoff; +use pb_mapper_core::{ snafu_error_get_or_continue, snafu_error_get_or_return, snafu_error_get_or_return_ok, snafu_error_handle, }; -use uni_stream::addr::{each_addr, ToSocketAddrs}; +use pb_mapper_protocol::command::{ + CONTROL_PROTOCOL_V2, LocalServer, MessageSerializer, PbConnRequest, PbConnResponse, + PbConnStatusReq, PbConnStatusResp, PbServerRequest, +}; +use pb_mapper_protocol::forward::StreamForward; +use pb_mapper_protocol::secure::ClientHeaderSession; +use pb_mapper_protocol::{MessageReader, MessageWriter}; +use uni_stream::addr::{ToSocketAddrs, each_addr}; use uni_stream::stream::{ - got_one_socket_addr, set_tcp_keep_alive, set_tcp_nodelay, StreamProvider, + StreamProvider, got_one_socket_addr, set_tcp_keep_alive, set_tcp_nodelay, }; fn get_ping_message(protocol_version: u16, seq: u64) -> error::Result> { @@ -131,23 +132,32 @@ pub struct ServerTunnelOptions { pub need_codec: bool, pub is_datagram: bool, pub keep_alive: bool, + pub namespace: Option, + pub force_namespace: bool, } -#[derive(Clone, Debug)] +#[derive(Clone)] struct ServerCliRunConfig { local_addr: A, remote_addr: A, key: Arc, options: ServerTunnelOptions, worker_index: usize, + credential: Credential, } -/// Where a stream request should connect, and how. -#[derive(Clone, Copy, Debug)] -struct StreamTarget { - local_addr: A, - remote_addr: A, - keep_alive: bool, +impl Debug for ServerCliRunConfig { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("ServerCliRunConfig") + .field("local_addr", &self.local_addr) + .field("remote_addr", &self.remote_addr) + .field("key", &self.key) + .field("options", &self.options) + .field("worker_index", &self.worker_index) + .field("credential_key_id", &self.credential.key_id()) + .finish() + } } fn duration_to_millis(duration: Duration) -> u64 { @@ -166,17 +176,21 @@ async fn probe_remote_registration( remote_addr: SocketAddr, key: Arc, registration: ControlRegistration, + namespace: Option, + credential: Credential, ) -> RegistrationProbeResult { let timeout = registration_probe_timeout(); let result = tokio::time::timeout(timeout, async { let mut stream = each_addr(remote_addr, TcpStream::connect) .await .map_err(|e| format!("connect remote status stream failed: {e}"))?; - crate::local::client::status::get_status( + crate::client::status::get_status_with_credential( &mut stream, PbConnStatusReq::Service { key: key.to_string(), }, + namespace, + &credential, ) .await .map_err(|e| { @@ -194,7 +208,7 @@ async fn probe_remote_registration( Err(_) => { return RegistrationProbeResult::Failed(format!( "status probe timed out after {timeout:?}" - )) + )); } }; @@ -238,6 +252,29 @@ pub async fn run_server_side_cli( .await } +pub async fn run_server_side_cli_with_pinned_credential( + local_addr: A, + remote_addr: A, + key: Arc, + options: ServerTunnelOptions, + status_callback: Option, + credential: Credential, +) where + LocalStream: StreamProvider + Send + 'static, + LocalStream::Item: StreamForward, + A: ToSocketAddrs + Debug + Copy, +{ + run_server_side_cli_pool::( + local_addr, + remote_addr, + key, + options, + status_callback, + Some(credential), + ) + .await; +} + pub async fn run_server_side_cli_with_callback( local_addr: A, remote_addr: A, @@ -248,6 +285,45 @@ pub async fn run_server_side_cli_with_callback( LocalStream: StreamProvider + Send + 'static, LocalStream::Item: StreamForward, A: ToSocketAddrs + Debug + Copy, +{ + run_server_side_cli_pool::( + local_addr, + remote_addr, + key, + options, + status_callback, + None, + ) + .await; +} + +async fn resolve_registration_credential(pinned: Option) -> Credential { + if let Some(credential) = pinned { + return credential; + } + let mut retry_backoff = RetryBackoff::default(); + loop { + match get_process_credential() { + Ok(credential) => return credential, + Err(error) => { + tracing::error!("load registration credential failed: {error}"); + tokio::time::sleep(retry_backoff.next_delay()).await; + } + } + } +} + +async fn run_server_side_cli_pool( + local_addr: A, + remote_addr: A, + key: Arc, + options: ServerTunnelOptions, + status_callback: Option, + pinned_credential: Option, +) where + LocalStream: StreamProvider + Send + 'static, + LocalStream::Item: StreamForward, + A: ToSocketAddrs + Debug + Copy, { let local_addr = match got_one_socket_addr(local_addr).await { Ok(addr) => addr, @@ -263,41 +339,38 @@ pub async fn run_server_side_cli_with_callback( return; } }; - let pool_size = control_conn_pool_size(); + let credential = resolve_registration_credential(pinned_credential).await; + let pool_size = control_conn_pool_size().max(1); tracing::info!( event = "local_server_control_pool_starting", key = %key, pool_size, "starting local server control connection pool" ); - let mut worker_handles = Vec::new(); - if pool_size > 1 { - for worker_index in 1..pool_size { - let worker_key = key.clone(); - worker_handles.push(tokio::spawn(async move { - run_server_side_cli_worker::( - local_addr, - remote_addr, - worker_key, - options, - None, - worker_index, - ) - .await; - })); - } + let mut workers = JoinSet::new(); + let mut status_callback = status_callback; + for worker_index in 0..pool_size { + let worker_key = key.clone(); + let callback = if worker_index == 0 { + status_callback.take() + } else { + None + }; + workers.spawn(async move { + run_server_side_cli_worker::( + local_addr, + remote_addr, + worker_key, + options, + callback, + worker_index, + credential, + ) + .await; + }); } - run_server_side_cli_worker::( - local_addr, - remote_addr, - key, - options, - status_callback, - 0, - ) - .await; - for handle in worker_handles { - if let Err(e) = handle.await { + while let Some(result) = workers.join_next().await { + if let Err(e) = result { tracing::warn!( event = "local_server_control_worker_join_failed", error = %e, @@ -314,19 +387,21 @@ async fn run_server_side_cli_worker( options: ServerTunnelOptions, status_callback: Option, worker_index: usize, + credential: Credential, ) where LocalStream: StreamProvider + Send + 'static, LocalStream::Item: StreamForward, A: ToSocketAddrs + Debug + Copy + Send + 'static, { + let mut retry_backoff = RetryBackoff::default(); let run_config = ServerCliRunConfig { local_addr, remote_addr, key: key.clone(), options, worker_index, + credential, }; - let mut retry_backoff = RetryBackoff::default(); loop { let status = if let Err(status) = run_server_side_cli_inner::( &mut retry_backoff, @@ -392,8 +467,11 @@ where need_codec, is_datagram, keep_alive, + namespace, + force_namespace, }, worker_index, + credential, } = config; let local_addr = match got_one_socket_addr(local_addr).await { Ok(addr) => addr, @@ -437,40 +515,54 @@ where "manager stream set tcp nodelay" ); - // start register server with key - { - let timeout = control_io_timeout(); - let heartbeat_interval = control_heartbeat_interval(); - let heartbeat_tolerance = control_heartbeat_tolerance(); - let msg = snafu_error_get_or_return_ok!(PbConnRequest::Register { + // Start registration with a protocol-v2 first frame. The session is reused for all + // subsequent control messages on this TCP connection. The credential is pinned + // when the worker starts so a later process-key change cannot retarget reconnects. + let session = match ClientHeaderSession::new_v2(&credential) { + Ok(session) => session, + Err(error) => { + tracing::error!("create manager protocol-v2 session failed: {error}"); + return Err(Status::ConnectRemote); + } + }; + let timeout = control_io_timeout(); + let heartbeat_interval = control_heartbeat_interval(); + let heartbeat_tolerance = control_heartbeat_tolerance(); + let request = match namespace { + Some(namespace) => PbConnRequest::RegisterScoped { key: key.to_string(), + namespace, + force_namespace, need_codec, is_datagram, protocol_version: Some(CONTROL_PROTOCOL_V2), client_instance_id: Some(new_client_instance_id(worker_index)), heartbeat_interval_ms: Some(duration_to_millis(heartbeat_interval)), heartbeat_tolerance_ms: Some(duration_to_millis(heartbeat_tolerance)), - } - .encode() - .context(EncodeRegisterReqSnafu)); - let mut msg_writer = match get_header_msg_writer(&mut manager_stream) { - Ok(writer) => writer, - Err(e) => { - tracing::error!("create manager header writer failed: {e}"); - return Err(Status::ConnectRemote); - } - }; - match tokio::time::timeout(timeout, msg_writer.write_msg(&msg)).await { - Ok(result) => snafu_error_get_or_return_ok!(result.context(SendRegisterReqSnafu)), - Err(_) => snafu_error_get_or_return_ok!(ControlIoTimeoutSnafu { + }, + None => PbConnRequest::Register { + key: key.to_string(), + need_codec, + is_datagram, + protocol_version: Some(CONTROL_PROTOCOL_V2), + client_instance_id: Some(new_client_instance_id(worker_index)), + heartbeat_interval_ms: Some(duration_to_millis(heartbeat_interval)), + heartbeat_tolerance_ms: Some(duration_to_millis(heartbeat_tolerance)), + }, + }; + let msg = snafu_error_get_or_return_ok!(request.encode().context(EncodeRegisterReqSnafu)); + match tokio::time::timeout(timeout, session.write_initial(&mut manager_stream, &msg)).await { + Ok(result) => snafu_error_get_or_return_ok!(result.context(SendRegisterReqSnafu)), + Err(_) => snafu_error_get_or_return_ok!( + ControlIoTimeoutSnafu { action: "send register request", timeout, } - .fail()), - } + .fail() + ), } let (mut reader, mut writer) = manager_stream.into_split(); - let mut msg_reader = match get_header_msg_reader(&mut reader) { + let mut msg_reader = match session.response_reader(&mut reader) { Ok(reader) => reader, Err(e) => { tracing::error!("create manager header reader failed: {e}"); @@ -482,11 +574,13 @@ where let timeout = control_io_timeout(); let msg = match tokio::time::timeout(timeout, msg_reader.read_msg()).await { Ok(result) => snafu_error_get_or_return_ok!(result.context(ReadRegisterRespSnafu)), - Err(_) => snafu_error_get_or_return_ok!(ControlIoTimeoutSnafu { - action: "read register response", - timeout, - } - .fail()), + Err(_) => snafu_error_get_or_return_ok!( + ControlIoTimeoutSnafu { + action: "read register response", + timeout, + } + .fail() + ), }; let resp = snafu_error_get_or_return_ok!( PbConnResponse::decode(msg).context(DecodeRegisterRespSnafu) @@ -508,6 +602,16 @@ where protocol_version: 1, lease_ttl_ms: 0, }, + PbConnResponse::Error(error) => { + tracing::error!( + event = "local_server_registration_rejected", + reason = %error.code, + retryable = error.retryable, + message = %error.message, + "pb server rejected service registration" + ); + snafu_error_get_or_return_ok!(RegisterRespNotMatchSnafu {}.fail()) + } _ => snafu_error_get_or_return_ok!(RegisterRespNotMatchSnafu {}.fail()), }; tracing::info!( @@ -536,7 +640,7 @@ where let writer_key = key.clone(); let writer_registration = registration; let mut writer_handle = tokio::spawn(async move { - let mut msg_writer = match get_header_msg_writer(&mut writer) { + let mut msg_writer = match session.continuation_writer(&mut writer) { Ok(writer) => writer, Err(e) => { tracing::error!("create manager header writer failed: {e}"); @@ -617,10 +721,12 @@ where snafu_error_get_or_continue!( handle_request::( msg, - StreamTarget { + StreamConnect { local_addr, remote_addr, keep_alive, + namespace, + credential, }, key.clone(), registration.conn_id, @@ -671,7 +777,14 @@ where let probe_tx = probe_tx.clone(); let probe_key = key.clone(); tokio::spawn(async move { - let result = probe_remote_registration(remote_addr, probe_key, registration).await; + let result = probe_remote_registration( + remote_addr, + probe_key, + registration, + namespace, + credential, + ) + .await; let _ = probe_tx.send(result); }); } @@ -776,7 +889,7 @@ async fn handle_request< A: ToSocketAddrs + Debug + Copy + Clone + Send + 'static, >( msg: &[u8], - target: StreamTarget, + target: StreamConnect, key: Arc, conn_id: u32, write_tx: &tokio::sync::mpsc::UnboundedSender, @@ -811,15 +924,8 @@ where let key = key.clone(); tokio::spawn(async move { snafu_error_handle!( - handle_stream::( - target.local_addr, - target.remote_addr, - key, - client_id, - server_generation, - target.keep_alive - ) - .await + handle_stream::(key, client_id, server_generation, target,) + .await ) }); } diff --git a/src/local/server/stream.rs b/crates/pb-mapper-client/src/server/stream.rs similarity index 53% rename from src/local/server/stream.rs rename to crates/pb-mapper-client/src/server/stream.rs index 0acdcd7..c8b5201 100644 --- a/src/local/server/stream.rs +++ b/crates/pb-mapper-client/src/server/stream.rs @@ -7,19 +7,27 @@ use tracing::info_span; use super::error::{ ConnectLocalStreamSnafu, ConnectRemoteStreamSnafu, ControlIoTimeoutSnafu, - DecodePbConnStreamRespSnafu, EncodePbConnStreamReqSnafu, PbConnStreamRespNotMatchSnafu, - ReadPbConnStreamRespSnafu, Result, WritePbConnStreamReqSnafu, + DecodePbConnStreamRespSnafu, EncodePbConnStreamReqSnafu, PbConnStreamRespNotMatchSnafu, Result, + WritePbConnStreamReqSnafu, }; -use crate::common::config::control_io_timeout; -use crate::common::message::command::{MessageSerializer, PbConnRequest, PbConnResponse}; -use crate::common::message::forward::StreamForward; -use crate::common::message::{ - get_header_msg_reader, get_header_msg_writer, MessageReader, MessageWriter, -}; -use crate::local::server::error::CreateHeaderToolSnafu; -use crate::snafu_error_handle; -use uni_stream::addr::{each_addr, ToSocketAddrs}; -use uni_stream::stream::{set_tcp_keep_alive, set_tcp_nodelay, StreamProvider, StreamSplit}; +use crate::server::error::CreateHeaderToolSnafu; +use pb_mapper_core::checksum::Credential; +use pb_mapper_core::config::control_io_timeout; +use pb_mapper_core::snafu_error_handle; +use pb_mapper_protocol::command::{MessageSerializer, PbConnRequest, PbConnResponse}; +use pb_mapper_protocol::forward::StreamForward; +use pb_mapper_protocol::secure::ClientHeaderSession; +use uni_stream::addr::{ToSocketAddrs, each_addr}; +use uni_stream::stream::{StreamProvider, StreamSplit, set_tcp_keep_alive, set_tcp_nodelay}; + +#[derive(Clone, Copy, Debug)] +pub struct StreamConnect { + pub local_addr: A, + pub remote_addr: A, + pub keep_alive: bool, + pub namespace: Option, + pub credential: Credential, +} /// Handle a stream connection and establish a forward network traffic forwarding. /// This function handles both local and remote streams, sets up message writers and readers, @@ -28,27 +36,39 @@ pub async fn handle_stream< LocalStream: StreamProvider, A: ToSocketAddrs + Debug + Copy + Clone + Send, >( - local_addr: A, - remote_addr: A, key: Arc, client_id: u32, server_generation: u64, - keep_alive: bool, + connect: StreamConnect, ) -> Result<()> where LocalStream::Item: StreamForward, { + let StreamConnect { + local_addr, + remote_addr, + keep_alive, + namespace, + credential, + } = connect; let key_ref = key.as_ref(); let client_id_span = info_span!("client_id", key_ref, client_id); let _enter = client_id_span.enter(); - let msg = PbConnRequest::Stream { - key: key.to_string(), - dst_id: client_id, - server_generation, - } - .encode() - .context(EncodePbConnStreamReqSnafu)?; + let request = match namespace { + Some(namespace) => PbConnRequest::StreamScoped { + key: key.to_string(), + namespace, + dst_id: client_id, + server_generation, + }, + None => PbConnRequest::Stream { + key: key.to_string(), + dst_id: client_id, + server_generation, + }, + }; + let msg = request.encode().context(EncodePbConnStreamReqSnafu)?; let timeout = control_io_timeout(); let mut remote_stream = @@ -68,31 +88,22 @@ where } snafu_error_handle!(set_tcp_nodelay(&remote_stream), "remote stream set nodelay"); - // write stream request and read response + // Use the credential captured at registration. A later UI/process key + // change must not move these streams into another namespace. let codec_key = { - let mut msg_writer = get_header_msg_writer(&mut remote_stream) - .context(CreateHeaderToolSnafu { action: "writer" })?; - match tokio::time::timeout(timeout, msg_writer.write_msg(&msg)).await { - Ok(result) => result.context(WritePbConnStreamReqSnafu)?, - Err(_) => ControlIoTimeoutSnafu { - action: "write pb conn stream request", - timeout, - } - .fail()?, - } - let mut msg_reader = get_header_msg_reader(&mut remote_stream) - .context(CreateHeaderToolSnafu { action: "reader" })?; - let msg = match tokio::time::timeout(timeout, msg_reader.read_msg()).await { - Ok(result) => result.context(ReadPbConnStreamRespSnafu)?, - Err(_) => ControlIoTimeoutSnafu { - action: "read pb conn stream response", - timeout, - } - .fail()?, - }; - let resp = PbConnResponse::decode(msg).context(DecodePbConnStreamRespSnafu)?; + let session = ClientHeaderSession::new_v2(&credential) + .context(CreateHeaderToolSnafu { action: "session" })?; + let response = session + .exchange(&mut remote_stream, &msg, timeout) + .await + .context(WritePbConnStreamReqSnafu)?; + let resp = PbConnResponse::decode(&response).context(DecodePbConnStreamRespSnafu)?; match resp { PbConnResponse::Stream { codec_key } => codec_key, + PbConnResponse::Error(error) => PbConnStreamRespNotMatchSnafu { + resp: format!("{}: {}", error.code, error.message), + } + .fail()?, _ => PbConnStreamRespNotMatchSnafu { resp: format!("{resp:?}"), } @@ -111,6 +122,7 @@ where snafu_error_handle!( ::forward_local_to_remote( codec_key, + *credential.key(), server_reader, server_writer, client_reader, diff --git a/crates/pb-mapper-core/Cargo.toml b/crates/pb-mapper-core/Cargo.toml new file mode 100644 index 0000000..1ae3b9c --- /dev/null +++ b/crates/pb-mapper-core/Cargo.toml @@ -0,0 +1,21 @@ +[package] +name = "pb-mapper-core" +version.workspace = true +edition.workspace = true +authors.workspace = true + +[dependencies] +base64.workspace = true +clap.workspace = true +hickory-resolver.workspace = true +parking_lot.workspace = true +rand.workspace = true +ring.workspace = true +serde_json.workspace = true +snafu.workspace = true +tokio.workspace = true +tracing.workspace = true +tracing-subscriber.workspace = true + +[lints] +workspace = true diff --git a/src/utils/addr.rs b/crates/pb-mapper-core/src/addr.rs similarity index 82% rename from src/utils/addr.rs rename to crates/pb-mapper-core/src/addr.rs index a8485ca..618f151 100644 --- a/src/utils/addr.rs +++ b/crates/pb-mapper-core/src/addr.rs @@ -2,11 +2,12 @@ use std::future::{self, Future}; use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr, SocketAddrV4, SocketAddrV6}; use std::pin::Pin; use std::sync::LazyLock; -use std::task::{ready, Context, Poll}; +use std::task::{Context, Poll, ready}; +use hickory_resolver::config::{NameServerConfig, ResolverConfig, ResolverOpts}; +use hickory_resolver::net::runtime::TokioRuntimeProvider; +use hickory_resolver::{Resolver, TokioResolver}; use tokio::task::JoinHandle; -use trust_dns_resolver::config::{NameServerConfigGroup, ResolverConfig, ResolverOpts}; -use trust_dns_resolver::{Resolver, TokioAsyncResolver}; type Result = std::result::Result; type ReadyFuture = future::Ready>; @@ -243,42 +244,40 @@ const DEFAULT_DNS_SERVER_GROUP: &[IpAddr] = &[ IpAddr::V6(Ipv6Addr::new(0x2001, 0x4860, 0x4860, 0, 0, 0, 0, 0x8888)), // google ]; -const DNS_QUERY_PORT: u16 = 53; - #[inline] fn custom_resolver_config() -> ResolverConfig { + // `udp_and_tcp` uses the standard DNS port and trusts negative responses, + // matching what this passed explicitly before. ResolverConfig::from_parts( None, vec![], - NameServerConfigGroup::from_ips_clear(DEFAULT_DNS_SERVER_GROUP, DNS_QUERY_PORT, true), + DEFAULT_DNS_SERVER_GROUP + .iter() + .copied() + .map(NameServerConfig::udp_and_tcp) + .collect::>(), ) } +/// The custom resolver, or `None` if it could not be built. +/// +/// Built once: `build()` can fail, and retrying per lookup would repeat the +/// same failure. Callers fall back to the system resolver. #[inline] -pub fn get_custom_resolver() -> Option { - // The sync resolver uses `block_on` internally and will panic if called from a Tokio runtime - // thread. Keep a guard here so callers can fall back to system DNS, and use the async helpers - // below when running inside async code. - if tokio::runtime::Handle::try_current().is_ok() { - tracing::debug!("Skipping sync custom DNS resolver inside Tokio runtime thread"); - return None; - } - - match Resolver::new(custom_resolver_config(), ResolverOpts::default()) { - Ok(r) => Some(r), - Err(e) => { - tracing::error!( - "Create custom dns resolver error:{e},we will use default dns resolver" - ); - None +fn get_custom_async_resolver() -> Option { + static RESOLVER: LazyLock> = LazyLock::new(|| { + let mut builder = Resolver::builder_with_config( + custom_resolver_config(), + TokioRuntimeProvider::default(), + ); + *builder.options_mut() = ResolverOpts::default(); + match builder.build() { + Ok(resolver) => Some(resolver), + Err(e) => { + tracing::error!("Create custom dns resolver error:{e},falling back to system dns"); + None + } } - } -} - -#[inline] -fn get_custom_async_resolver() -> TokioAsyncResolver { - static RESOLVER: LazyLock = LazyLock::new(|| { - TokioAsyncResolver::tokio(custom_resolver_config(), ResolverOpts::default()) }); RESOLVER.clone() } @@ -298,24 +297,15 @@ macro_rules! try_opt { }; } -fn get_ip_addrs(s: &str) -> Result> { - thread_local! { - static RESOLVER:Option = get_custom_resolver(); - } - let result = RESOLVER.with(|r| r.as_ref().map(|r| r.lookup_ip(s))); - try_opt!(result, "custom resolver not exist") - .map(|v| v.into_iter().collect()) - .map_err(|e| invalid_input!(e)) -} - -/// Blocking DNS lookup. Avoid calling this from inside a Tokio runtime thread. +/// Blocking DNS lookup, via the system resolver. +/// +/// The custom DNS servers are only reachable from the async helpers below. +/// hickory has no blocking resolver, and the sync path never reached the +/// custom one in practice anyway: it was skipped inside a Tokio runtime, and +/// outside one this fell back to `std` whenever it was unavailable. #[inline] pub fn get_socket_addrs_from_host_port(host: &str, port: u16) -> Result> { - match get_ip_addrs(host) { - Ok(r) => Ok(r.into_iter().map(|ip| SocketAddr::new(ip, port)).collect()), - // Resolve dns properly with the standard library - Err(_) => std::net::ToSocketAddrs::to_socket_addrs(&(host, port)).map(|v| v.collect()), - } + std::net::ToSocketAddrs::to_socket_addrs(&(host, port)).map(|v| v.collect()) } /// Blocking DNS lookup. Avoid calling this from inside a Tokio runtime thread. @@ -328,7 +318,7 @@ pub fn get_socket_addrs(s: &str) -> Result> { /// Async DNS lookup using the custom resolver, safe to call inside Tokio runtimes. pub async fn get_ip_addrs_async(s: &str) -> Result> { - let resolver = get_custom_async_resolver(); + let resolver = try_opt!(get_custom_async_resolver(), "custom resolver not exist"); resolver .lookup_ip(s) .await diff --git a/src/common/checksum.rs b/crates/pb-mapper-core/src/checksum.rs similarity index 51% rename from src/common/checksum.rs rename to crates/pb-mapper-core/src/checksum.rs index 419c1fe..190ced2 100644 --- a/src/common/checksum.rs +++ b/crates/pb-mapper-core/src/checksum.rs @@ -3,62 +3,100 @@ use std::io; use std::os::unix::fs::PermissionsExt; use std::path::Path; use std::process::Command; +use std::sync::LazyLock; use std::sync::atomic::{AtomicU32, Ordering}; -use std::sync::{LazyLock, RwLock}; +use parking_lot::RwLock; + +use base64::Engine; +use base64::engine::general_purpose::URL_SAFE_NO_PAD; use rand::RngExt; -use ring::digest::{digest, SHA256}; +use ring::digest::{SHA256, digest}; -use super::message::DataLenType; +use crate::DataLenType; pub type ChecksumType = u32; -const DEFAULT_KEY: &str = "abcdefghijklmnopqlsn123456789j01"; /// Environment variable used by server/client processes to carry the 32-byte header key. pub const ENV_MSG_HEADER_KEY: &str = "MSG_HEADER_KEY"; /// Fixed file path used to persist a machine-derived key for operators to reuse. pub const MACHINE_MSG_HEADER_KEY_PATH: &str = "/var/lib/pb-mapper-server/msg_header_key"; +pub const ADMIN_KEY_PATH: &str = "/var/lib/pb-mapper/auth/admin.key"; +pub const TEMP_CREDENTIAL_PREFIX: &str = "pbmt1_"; +pub const ADMIN_KEY_LEN: usize = 32; + +/// Administrator keys are also stored in `MSG_HEADER_KEY`. `std::env::set_var` +/// panics on interior NUL, so the key must be printable ASCII with no whitespace. +pub fn is_env_safe_admin_key(bytes: &[u8]) -> bool { + bytes.len() == ADMIN_KEY_LEN && bytes.iter().all(|byte| byte.is_ascii_graphic()) +} + +pub fn env_safe_admin_key_error() -> String { + format!( + "`{ENV_MSG_HEADER_KEY}` administrator key must be 32 printable ASCII bytes without whitespace or NUL" + ) +} const DERIVE_MSG_HEADER_KEY_TAG: &str = "pb-mapper-msg-header-key-v1"; const DERIVE_MSG_HEADER_KEY_CHARSET: &[u8] = b"0123456789abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ"; struct MsgHeaderKeyState { - key: RwLock>, + credential: RwLock>, + load_error: RwLock>, hash: AtomicU32, } -fn key_len_error(input: &str) -> String { - format!("`{ENV_MSG_HEADER_KEY}` must have 256 bit(32 byte)!. current input key:{input}") +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum Credential { + Admin(AesKeyType), + Temporary { key_id: u64, key: AesKeyType }, } -fn load_msg_header_key_from_env_or_default() -> Vec { - let key = match std::env::var(ENV_MSG_HEADER_KEY) { - Ok(k) => { - let key = k.as_bytes(); - if key.len() != 32 { - tracing::warn!("{}", key_len_error(&k)); - std::process::exit(1); - } - key.to_vec() +impl Credential { + pub fn key_id(&self) -> u64 { + match self { + Self::Admin(_) => 0, + Self::Temporary { key_id, .. } => *key_id, } - Err(_) => { - tracing::warn!( - "No ENV:`{ENV_MSG_HEADER_KEY}` provided,we use default key:{DEFAULT_KEY}" - ); - DEFAULT_KEY.as_bytes().to_vec() + } + + pub fn key(&self) -> &AesKeyType { + match self { + Self::Admin(key) | Self::Temporary { key, .. } => key, } + } + + pub fn is_admin(&self) -> bool { + matches!(self, Self::Admin(_)) + } +} + +fn key_len_error(input: &str) -> String { + format!( + "`{ENV_MSG_HEADER_KEY}` administrator key must be exactly 32 bytes; received {} bytes", + input.len() + ) +} + +fn load_credential_from_env() -> Result, String> { + let Some(raw) = std::env::var_os(ENV_MSG_HEADER_KEY) else { + return Ok(None); }; - key + let raw = raw + .into_string() + .map_err(|_| format!("`{ENV_MSG_HEADER_KEY}` must contain valid UTF-8 credential text"))?; + parse_credential(raw.trim()).map(Some) } -fn update_runtime_msg_header_key(key: Vec) { - let hash = gen_checksum_by_key(&key); - let mut guard = MSG_HEADER_KEY_STATE - .key - .write() - .unwrap_or_else(|poisoned| poisoned.into_inner()); - *guard = key; +fn update_runtime_credential(credential: Option) { + let hash = credential + .as_ref() + .map(|credential| gen_checksum_by_key(credential.key())) + .unwrap_or_default(); + let mut guard = MSG_HEADER_KEY_STATE.credential.write(); + *guard = credential; + *MSG_HEADER_KEY_STATE.load_error.write() = None; MSG_HEADER_KEY_STATE.hash.store(hash, Ordering::Release); } @@ -67,45 +105,129 @@ fn update_runtime_msg_header_key(key: Vec) { /// This state is mutable so FFI/UI can update `MSG_HEADER_KEY` at runtime /// without restarting the process. static MSG_HEADER_KEY_STATE: LazyLock = LazyLock::new(|| { - let key = load_msg_header_key_from_env_or_default(); - let hash = gen_checksum_by_key(&key); + let (credential, load_error) = match load_credential_from_env() { + Ok(credential) => (credential, None), + Err(error) => { + tracing::error!(reason = "credential_invalid", %error, "invalid MSG_HEADER_KEY"); + (None, Some(error)) + } + }; + let hash = credential + .as_ref() + .map(|credential| gen_checksum_by_key(credential.key())) + .unwrap_or_default(); MsgHeaderKeyState { - key: RwLock::new(key), + credential: RwLock::new(credential), + load_error: RwLock::new(load_error), hash: AtomicU32::new(hash), } }); +/// Return the configured process credential, failing closed when none exists. +pub fn get_process_credential() -> Result { + if let Some(error) = MSG_HEADER_KEY_STATE.load_error.read().clone() { + return Err(error); + } + MSG_HEADER_KEY_STATE.credential.read().ok_or_else(|| { + format!("`{ENV_MSG_HEADER_KEY}` is required; no insecure default credential is available") + }) +} + /// Get current message header key bytes. -pub fn get_msg_header_key() -> Vec { - MSG_HEADER_KEY_STATE - .key - .read() - .unwrap_or_else(|poisoned| poisoned.into_inner()) - .clone() +pub fn get_msg_header_key() -> Result, String> { + get_process_credential().map(|credential| credential.key().to_vec()) } /// Set process `MSG_HEADER_KEY` and update runtime checksum/key state. /// -/// - `Some(non-empty)` => validate length 32, set env, apply immediately. -/// - `None` or empty => remove env and reset to default key. +/// - `Some(non-empty)` => validate an admin or temporary credential and apply it immediately. +/// - `None` or empty => remove the credential. Subsequent network operations fail closed. pub fn set_process_msg_header_key(msg_header_key: Option<&str>) -> Result<(), String> { let normalized = msg_header_key.map(str::trim).unwrap_or(""); if normalized.is_empty() { - std::env::remove_var(ENV_MSG_HEADER_KEY); - update_runtime_msg_header_key(DEFAULT_KEY.as_bytes().to_vec()); + // SAFETY: edition 2024 makes these unsafe because the environment is + // process-global and another thread reading it concurrently is a data + // race. Here the environment is only a mirror for child processes and + // for the initial read at startup; the credential every operation + // actually consults is `MSG_HEADER_KEY_STATE`, behind an `RwLock`, and + // it is updated immediately below. + unsafe { std::env::remove_var(ENV_MSG_HEADER_KEY) }; + update_runtime_credential(None); return Ok(()); } - let key = normalized.as_bytes(); - if key.len() != 32 { - return Err(key_len_error(normalized)); - } + let credential = parse_credential(normalized)?; - std::env::set_var(ENV_MSG_HEADER_KEY, normalized); - update_runtime_msg_header_key(key.to_vec()); + // SAFETY: as above. + unsafe { std::env::set_var(ENV_MSG_HEADER_KEY, normalized) }; + update_runtime_credential(Some(credential)); Ok(()) } +pub fn parse_credential(raw: &str) -> Result { + if let Some(encoded) = raw.strip_prefix(TEMP_CREDENTIAL_PREFIX) { + let payload = URL_SAFE_NO_PAD + .decode(encoded) + .map_err(|_| "temporary credential is not valid base64url".to_string())?; + if payload.len() != 45 { + return Err(format!( + "temporary credential payload must be 45 bytes, got {}", + payload.len() + )); + } + if payload[0] != 1 { + return Err(format!( + "unsupported temporary credential version {}", + payload[0] + )); + } + let expected = digest(&SHA256, &payload[..41]); + if expected.as_ref()[..4] != payload[41..45] { + return Err("temporary credential checksum mismatch".to_string()); + } + // The 45-byte check above already guarantees both widths, so neither arm + // is reachable — but this returns `Result` anyway, so saying so costs a + // line and removes a panic from a path that parses network input. + let key_id = u64::from_be_bytes( + payload[1..9] + .try_into() + .map_err(|_| "temporary credential key id is malformed".to_string())?, + ); + if key_id == 0 { + return Err("temporary credential key id must not be zero".to_string()); + } + let key = payload[9..41] + .try_into() + .map_err(|_| "temporary credential key is malformed".to_string())?; + return Ok(Credential::Temporary { key_id, key }); + } + + let bytes = raw.as_bytes(); + if bytes.len() != ADMIN_KEY_LEN { + return Err(key_len_error(raw)); + } + if !is_env_safe_admin_key(bytes) { + return Err(env_safe_admin_key_error()); + } + // Unreachable after the `ADMIN_KEY_LEN` check above, for the same reason. + Ok(Credential::Admin( + bytes.try_into().map_err(|_| key_len_error(raw))?, + )) +} + +pub fn encode_temporary_credential(key_id: u64, key: &AesKeyType) -> String { + let mut payload = Vec::with_capacity(45); + payload.push(1); + payload.extend_from_slice(&key_id.to_be_bytes()); + payload.extend_from_slice(key); + let checksum = digest(&SHA256, &payload); + payload.extend_from_slice(&checksum.as_ref()[..4]); + format!( + "{TEMP_CREDENTIAL_PREFIX}{}", + URL_SAFE_NO_PAD.encode(payload) + ) +} + /// Derive a stable machine-specific `MSG_HEADER_KEY` and persist it. /// /// The derivation seed is built from normalized hostname + normalized MAC list, @@ -133,18 +255,18 @@ fn get_machine_hostname() -> io::Result { return Ok(hostname); } - if let Ok(content) = std::fs::read_to_string("/etc/hostname") { - if let Some(hostname) = normalize_non_empty(Some(content.as_str())) { - return Ok(hostname); - } + if let Ok(content) = std::fs::read_to_string("/etc/hostname") + && let Some(hostname) = normalize_non_empty(Some(content.as_str())) + { + return Ok(hostname); } - if let Ok(output) = Command::new("hostname").output() { - if output.status.success() { - let stdout = String::from_utf8_lossy(&output.stdout); - if let Some(hostname) = normalize_non_empty(Some(stdout.as_ref())) { - return Ok(hostname); - } + if let Ok(output) = Command::new("hostname").output() + && output.status.success() + { + let stdout = String::from_utf8_lossy(&output.stdout); + if let Some(hostname) = normalize_non_empty(Some(stdout.as_ref())) { + return Ok(hostname); } } @@ -165,22 +287,22 @@ fn normalize_non_empty(input: Option<&str>) -> Option { } fn get_machine_mac_addresses() -> io::Result> { - if let Ok(mac_addresses) = get_machine_mac_addresses_from_sysfs() { - if !mac_addresses.is_empty() { - return Ok(mac_addresses); - } + if let Ok(mac_addresses) = get_machine_mac_addresses_from_sysfs() + && !mac_addresses.is_empty() + { + return Ok(mac_addresses); } - if let Ok(mac_addresses) = get_machine_mac_addresses_from_ip_link() { - if !mac_addresses.is_empty() { - return Ok(mac_addresses); - } + if let Ok(mac_addresses) = get_machine_mac_addresses_from_ip_link() + && !mac_addresses.is_empty() + { + return Ok(mac_addresses); } - if let Ok(mac_addresses) = get_machine_mac_addresses_from_ifconfig() { - if !mac_addresses.is_empty() { - return Ok(mac_addresses); - } + if let Ok(mac_addresses) = get_machine_mac_addresses_from_ifconfig() + && !mac_addresses.is_empty() + { + return Ok(mac_addresses); } Err(io::Error::new( @@ -363,7 +485,7 @@ fn write_machine_msg_header_key(key: &str) -> io::Result<()> { ) })?; #[cfg(unix)] - std::fs::set_permissions(path, std::fs::Permissions::from_mode(0o644))?; + std::fs::set_permissions(path, std::fs::Permissions::from_mode(0o600))?; Ok(()) } @@ -373,16 +495,37 @@ fn gen_checksum_by_key(key: &[u8]) -> ChecksumType { }) } +#[inline] +/// Compute frame checksum from payload length and an explicit header key. +pub fn get_checksum_for_key(datalen: DataLenType, key: &[u8]) -> ChecksumType { + datalen ^ gen_checksum_by_key(key) +} + +/// `true` when the process credential can key a length checksum. +pub fn process_checksum_is_ready() -> bool { + get_process_credential().is_ok() +} + #[inline] /// Compute frame checksum from payload length and the current header key hash. pub fn get_checksum(datalen: DataLenType) -> ChecksumType { datalen ^ MSG_HEADER_KEY_STATE.hash.load(Ordering::Acquire) } +#[inline] +/// Validate a frame checksum against an explicit header key. +pub fn valid_checksum_for_key(datalen: DataLenType, checksum: ChecksumType, key: &[u8]) -> bool { + checksum == get_checksum_for_key(datalen, key) +} + #[inline] /// Validate frame checksum generated by [`get_checksum`]. +/// +/// Missing or invalid process credentials fail closed instead of accepting +/// an unkeyed `datalen` checksum (`hash == 0`). pub fn valid_checksum(datalen: DataLenType, checksum: ChecksumType) -> bool { - datalen == (checksum ^ MSG_HEADER_KEY_STATE.hash.load(Ordering::Acquire)) + process_checksum_is_ready() + && datalen == (checksum ^ MSG_HEADER_KEY_STATE.hash.load(Ordering::Acquire)) } pub type AesKeyType = [u8; 32]; @@ -404,11 +547,64 @@ pub fn gen_random_key() -> [u8; 32] { random_key } +// This was missing its `#[cfg(test)]`, so it compiled into release builds. +#[cfg(test)] mod tests { #[test] fn test_random_checksum() { use super::*; - println!("{}", gen_checksum_by_key(DEFAULT_KEY.as_bytes())); + println!( + "{}", + gen_checksum_by_key(b"0123456789abcdefghijklmnopqrstuv") + ); + } + + #[test] + fn checksum_for_an_explicit_key_is_independent_of_process_state() { + use super::*; + let key = b"0123456789abcdefghijklmnopqrstuv"; + let datalen = 32; + let checksum = get_checksum_for_key(datalen, key); + assert!(valid_checksum_for_key(datalen, checksum, key)); + assert!(!valid_checksum_for_key( + datalen, + checksum, + b"abcdefghijklmnopqrstuvwxyz012345" + )); + } + + #[test] + fn env_safe_admin_key_rejects_nul_and_accepts_printable_ascii() { + use super::*; + assert!(is_env_safe_admin_key(b"0123456789abcdefghijklmnopqrstuv")); + let mut with_nul = *b"0123456789abcdefghijklmnopqrstuv"; + with_nul[8] = 0; + assert!(!is_env_safe_admin_key(&with_nul)); + assert!(!is_env_safe_admin_key(b"short")); + } + + #[tokio::test] + async fn clearing_the_process_credential_fails_closed_for_checksums() { + use super::*; + use crate::test_support::PROCESS_CREDENTIAL_TEST_LOCK; + + let _guard = PROCESS_CREDENTIAL_TEST_LOCK.lock().await; + set_process_msg_header_key(Some("0123456789abcdefghijklmnopqrstuv")).unwrap(); + let checksum = get_checksum(32); + assert!(valid_checksum(32, checksum)); + set_process_msg_header_key(None).unwrap(); + assert!(!valid_checksum(32, checksum)); + assert!(!valid_checksum(32, 32)); + assert!(get_process_credential().is_err()); + } + + #[test] + fn administrator_credentials_reject_nul_and_whitespace() { + use super::*; + let mut with_nul = *b"0123456789abcdefghijklmnopqrstuv"; + with_nul[8] = 0; + assert!(parse_credential(std::str::from_utf8(&with_nul).unwrap()).is_err()); + assert!(parse_credential("0123456789abcdefghijklmnopq rstuv").is_err()); } #[test] @@ -424,4 +620,24 @@ mod tests { assert_eq!(key1.len(), 32); assert!(key1.chars().all(|ch| ch.is_ascii_alphanumeric())); } + + #[test] + fn temporary_credential_round_trip_and_checksum() { + use super::*; + + let key = [7_u8; 32]; + let encoded = encode_temporary_credential(0x0000_0007_0000_002a, &key); + assert_eq!( + parse_credential(&encoded).unwrap(), + Credential::Temporary { + key_id: 0x0000_0007_0000_002a, + key + } + ); + + let mut corrupted = encoded.into_bytes(); + let last = corrupted.last_mut().unwrap(); + *last = if *last == b'A' { b'B' } else { b'A' }; + assert!(parse_credential(std::str::from_utf8(&corrupted).unwrap()).is_err()); + } } diff --git a/src/utils/codec.rs b/crates/pb-mapper-core/src/codec.rs similarity index 61% rename from src/utils/codec.rs rename to crates/pb-mapper-core/src/codec.rs index b375470..ae6c186 100644 --- a/src/utils/codec.rs +++ b/crates/pb-mapper-core/src/codec.rs @@ -1,12 +1,10 @@ use std::mem::size_of; use ring::aead::{ - Aad, BoundKey, Nonce, NonceSequence, OpeningKey, SealingKey, Tag, UnboundKey, AES_256_GCM, - NONCE_LEN, + AES_256_GCM, Aad, BoundKey, NONCE_LEN, Nonce, NonceSequence, OpeningKey, SealingKey, Tag, + UnboundKey, }; -use crate::common::checksum::get_msg_header_key; - #[derive(Clone, Copy, Default)] struct Counter(u32); @@ -55,11 +53,6 @@ impl Aes256GcmCodec { }) } - pub fn try_new_with_default_key() -> RingResult { - let key = get_msg_header_key(); - Aes256GcmCodec::try_new(key.as_ref()) - } - pub fn encrypt(&mut self, data: &mut [u8]) -> RingResult { self.seal.encrypt(data) } @@ -159,7 +152,7 @@ mod tests { use std::slice::from_raw_parts_mut; use std::time::Instant; - use crate::utils::codec::Aes256GcmCodec; + use crate::codec::Aes256GcmCodec; struct Timer { ins: Instant, @@ -197,8 +190,11 @@ mod tests { #[test] fn test_codec() { - let data = String::from("fdafas反对fdasfasfasfsdafdasfsdfasd范德萨发顺🤣❤️😁😍👍👍丰十大大师傅士大夫大撒发射点发士大夫大师傅大师傅士大夫士大夫阿斯蒂芬大师傅阿斯顿法大师傅看叫阿三的发就可是大家发开始打客服开始大幅喀什的开发点卡收费就开始打客服就是的咖啡肯定撒法开始打客服就是的咖啡就开始大幅扣税的急啊看发叫阿三的发生的开发就是大家可是大家发看大数据开发大数据开发大家ask发就是的咖啡的萨芬就卡死的房价开始打家开发商的JFK上的飞机卡上的纠纷开始打飞机宽带技术开发就开始大家开发建设的卡JFK大数据风控静安寺的看法角度看萨芬卡上的纠纷看静安寺的看法角度思考积分可是大家发卡是大家看法就大肆砍伐尽快打算减肥肯定是积分开始大幅技术大咖积分开始打飞机扣税的急啊看发的技术开发就是JFK十大福克斯大家开发大撒发射点幅度萨芬撒旦发发收范德萨发顺丰士大夫十大阿斯蒂芬大师傅阿斯顿附件是的客服对接撒巨大石块积分的课时费阿斯蒂芬法大师傅大师傅十大法大师傅阿斯蒂芬阿斯顿法大师傅阿斯蒂芬大师傅阿斯顿法大师傅大师傅阿斯蒂芬阿斯蒂芬士大夫阿斯蒂芬大师傅的萨芬打算减肥上岛咖啡加快速度大数据开发就是打客服看大数据开发就开始减肥卡萨丁JFK是大家看法加快速度JFK技术大咖积分喀什的开发独守空房技术大咖积分空手道解放扣税的开发商的开发接口是大家看法角度看是否扣税的急啊看发生的开发的快速减肥开始大幅就是打客服卡上的纠纷啊撒旦解放扣税的急啊看发加快速度点卡JFK啥的但是法大师傅技术大咖积分卡萨丁就反馈是大家看法啊是大家看法卡上的纠纷可是大家发喀什的开发大卡司喀什的开发就是打客服法大师傅士大夫的式咖啡机上岛咖啡就是的咖啡艰苦大师傅看上雕刻技法喀什的开发上岛咖啡就喀什的开发就是打客服卡上的纠纷技术的咖啡机肯定撒开发啊十大科技开发速度加啊反馈就是的咖啡开始大幅大师傅似的十大放假啊上岛咖啡就可是大家发空间的是否撒旦士大夫的撒娇开发是大家看法大肆砍伐就喀什的开发氨基酸的考虑非军事对抗疗法金克拉撒旦发艰苦拉萨的飞机喀什打开发就可是大家发可是大家看附件卡上的纠纷卡刷点卡技术的咖啡机可是大家发卡是大家看法静安寺的看法就可是大家发卡萨丁就开发商的急啊看飞机迪斯科发技术的咖啡机可是大家发看电视剧开发商大开始打到发大水发大水"); - let mut cryption = Aes256GcmCodec::try_new_with_default_key().unwrap(); + const TEST_KEY: [u8; 32] = [0x42; 32]; + let data = String::from( + "fdafas反对fdasfasfasfsdafdasfsdfasd范德萨发顺🤣❤️😁😍👍👍丰十大大师傅士大夫大撒发射点发士大夫大师傅大师傅士大夫士大夫阿斯蒂芬大师傅阿斯顿法大师傅看叫阿三的发就可是大家发开始打客服开始大幅喀什的开发点卡收费就开始打客服就是的咖啡肯定撒法开始打客服就是的咖啡就开始大幅扣税的急啊看发叫阿三的发生的开发就是大家可是大家发看大数据开发大数据开发大家ask发就是的咖啡的萨芬就卡死的房价开始打家开发商的JFK上的飞机卡上的纠纷开始打飞机宽带技术开发就开始大家开发建设的卡JFK大数据风控静安寺的看法角度看萨芬卡上的纠纷看静安寺的看法角度思考积分可是大家发卡是大家看法就大肆砍伐尽快打算减肥肯定是积分开始大幅技术大咖积分开始打飞机扣税的急啊看发的技术开发就是JFK十大福克斯大家开发大撒发射点幅度萨芬撒旦发发收范德萨发顺丰士大夫十大阿斯蒂芬大师傅阿斯顿附件是的客服对接撒巨大石块积分的课时费阿斯蒂芬法大师傅大师傅十大法大师傅阿斯蒂芬阿斯顿法大师傅阿斯蒂芬大师傅阿斯顿法大师傅大师傅阿斯蒂芬阿斯蒂芬士大夫阿斯蒂芬大师傅的萨芬打算减肥上岛咖啡加快速度大数据开发就是打客服看大数据开发就开始减肥卡萨丁JFK是大家看法加快速度JFK技术大咖积分喀什的开发独守空房技术大咖积分空手道解放扣税的开发商的开发接口是大家看法角度看是否扣税的急啊看发生的开发的快速减肥开始大幅就是打客服卡上的纠纷啊撒旦解放扣税的急啊看发加快速度点卡JFK啥的但是法大师傅技术大咖积分卡萨丁就反馈是大家看法啊是大家看法卡上的纠纷可是大家发喀什的开发大卡司喀什的开发就是打客服法大师傅士大夫的式咖啡机上岛咖啡就是的咖啡艰苦大师傅看上雕刻技法喀什的开发上岛咖啡就喀什的开发就是打客服卡上的纠纷技术的咖啡机肯定撒开发啊十大科技开发速度加啊反馈就是的咖啡开始大幅大师傅似的十大放假啊上岛咖啡就可是大家发空间的是否撒旦士大夫的撒娇开发是大家看法大肆砍伐就喀什的开发氨基酸的考虑非军事对抗疗法金克拉撒旦发艰苦拉萨的飞机喀什打开发就可是大家发可是大家看附件卡上的纠纷卡刷点卡技术的咖啡机可是大家发卡是大家看法静安寺的看法就可是大家发卡萨丁就开发商的急啊看飞机迪斯科发技术的咖啡机可是大家发看电视剧开发商大开始打到发大水发大水", + ); + let mut cryption = Aes256GcmCodec::try_new(&TEST_KEY).unwrap(); let mut out_buf = data.as_bytes().to_vec(); let tag = { let _timer = Timer::new_with_hint("Encrypt".into()); diff --git a/src/common/config.rs b/crates/pb-mapper-core/src/config.rs similarity index 89% rename from src/common/config.rs rename to crates/pb-mapper-core/src/config.rs index 0c0fcce..bef77f3 100644 --- a/src/common/config.rs +++ b/crates/pb-mapper-core/src/config.rs @@ -5,9 +5,9 @@ use std::time::Duration; use clap::ValueEnum; use snafu::ResultExt; use tracing_subscriber::layer::SubscriberExt; -use tracing_subscriber::{fmt, EnvFilter, Layer}; +use tracing_subscriber::{EnvFilter, Layer, fmt}; -use super::error::{CfgPbServerEnvNotExistSnafu, Result}; +use crate::error::{CfgPbServerEnvNotExistSnafu, Result}; #[derive(ValueEnum, Debug, Clone, Copy)] pub enum StatusOp { @@ -30,24 +30,24 @@ pub fn get_sockaddr(addr: &str) -> Result { Ok(mut socket_addrs) => { socket_addrs .next() - .ok_or_else(|| super::error::Error::CfgParseSockAddr { + .ok_or_else(|| crate::error::Error::CfgParseSockAddr { string: addr.to_string(), source: original_parse_error, }) } - Err(_) => Err(super::error::Error::CfgParseSockAddr { + Err(_) => Err(crate::error::Error::CfgParseSockAddr { string: addr.to_string(), source: original_parse_error, }), } } else { // For other hostnames, use the custom DNS resolution - use crate::utils::addr::get_socket_addrs; + use crate::addr::get_socket_addrs; match get_socket_addrs(addr) { Ok(socket_addrs) => { // Return the first resolved address socket_addrs.into_iter().next().ok_or_else(|| { - super::error::Error::CfgParseSockAddr { + crate::error::Error::CfgParseSockAddr { string: addr.to_string(), source: original_parse_error, } @@ -57,12 +57,12 @@ pub fn get_sockaddr(addr: &str) -> Result { // If custom DNS resolution fails, fallback to system resolver match std::net::ToSocketAddrs::to_socket_addrs(addr) { Ok(mut socket_addrs) => socket_addrs.next().ok_or_else(|| { - super::error::Error::CfgParseSockAddr { + crate::error::Error::CfgParseSockAddr { string: addr.to_string(), source: original_parse_error, } }), - Err(_) => Err(super::error::Error::CfgParseSockAddr { + Err(_) => Err(crate::error::Error::CfgParseSockAddr { string: addr.to_string(), source: original_parse_error, }), @@ -87,22 +87,22 @@ pub async fn get_sockaddr_async(addr: &str) -> Result { Ok(mut socket_addrs) => { socket_addrs .next() - .ok_or_else(|| super::error::Error::CfgParseSockAddr { + .ok_or_else(|| crate::error::Error::CfgParseSockAddr { string: addr.to_string(), source: original_parse_error, }) } - Err(_) => Err(super::error::Error::CfgParseSockAddr { + Err(_) => Err(crate::error::Error::CfgParseSockAddr { string: addr.to_string(), source: original_parse_error, }), } } else { // For other hostnames, use the custom DNS resolution - use crate::utils::addr::get_socket_addrs_async; + use crate::addr::get_socket_addrs_async; match get_socket_addrs_async(addr).await { Ok(socket_addrs) => socket_addrs.into_iter().next().ok_or_else(|| { - super::error::Error::CfgParseSockAddr { + crate::error::Error::CfgParseSockAddr { string: addr.to_string(), source: original_parse_error, } @@ -111,12 +111,12 @@ pub async fn get_sockaddr_async(addr: &str) -> Result { // If custom DNS resolution fails, fallback to system resolver match tokio::net::lookup_host(addr).await { Ok(mut socket_addrs) => socket_addrs.next().ok_or_else(|| { - super::error::Error::CfgParseSockAddr { + crate::error::Error::CfgParseSockAddr { string: addr.to_string(), source: original_parse_error, } }), - Err(_) => Err(super::error::Error::CfgParseSockAddr { + Err(_) => Err(crate::error::Error::CfgParseSockAddr { string: addr.to_string(), source: original_parse_error, }), @@ -410,17 +410,27 @@ mod tests { fn keep_alive_reads_the_environment_every_time() { let restore = std::env::var(PB_MAPPER_KEEP_ALIVE).ok(); - std::env::remove_var(PB_MAPPER_KEEP_ALIVE); + // SAFETY: mutating the environment is unsafe in edition 2024 because + // it is process-global. This is the only test that touches + // `PB_MAPPER_KEEP_ALIVE` — which is why it is one test and not several + // — and it restores the original value before returning. + unsafe { + std::env::remove_var(PB_MAPPER_KEEP_ALIVE); + } assert!(!keep_alive_from_env(), "absent means off"); - std::env::set_var(PB_MAPPER_KEEP_ALIVE, "ON"); + unsafe { + std::env::set_var(PB_MAPPER_KEEP_ALIVE, "ON"); + } assert!(keep_alive_from_env(), "the documented spelling"); // The regression. This used to be a `LazyLock`, so the answer was // whatever the first caller in the process saw and could never change — // which is why the UI's per-service toggle did nothing after the first // tunnel started. - std::env::set_var(PB_MAPPER_KEEP_ALIVE, "OFF"); + unsafe { + std::env::set_var(PB_MAPPER_KEEP_ALIVE, "OFF"); + } assert!( !keep_alive_from_env(), "OFF must mean off; the old check was `is_ok()`, so any value at \ @@ -428,17 +438,23 @@ mod tests { ); for truthy in ["on", "1", "true", "yes", " ON "] { - std::env::set_var(PB_MAPPER_KEEP_ALIVE, truthy); + unsafe { + std::env::set_var(PB_MAPPER_KEEP_ALIVE, truthy); + } assert!(keep_alive_from_env(), "{truthy:?} should enable"); } for falsy in ["", "off", "0", "false", "no"] { - std::env::set_var(PB_MAPPER_KEEP_ALIVE, falsy); + unsafe { + std::env::set_var(PB_MAPPER_KEEP_ALIVE, falsy); + } assert!(!keep_alive_from_env(), "{falsy:?} should not enable"); } - match restore { - Some(value) => std::env::set_var(PB_MAPPER_KEEP_ALIVE, value), - None => std::env::remove_var(PB_MAPPER_KEEP_ALIVE), + unsafe { + match restore { + Some(value) => std::env::set_var(PB_MAPPER_KEEP_ALIVE, value), + None => std::env::remove_var(PB_MAPPER_KEEP_ALIVE), + } } } } diff --git a/src/common/conn_id.rs b/crates/pb-mapper-core/src/conn_id.rs similarity index 100% rename from src/common/conn_id.rs rename to crates/pb-mapper-core/src/conn_id.rs diff --git a/crates/pb-mapper-core/src/durable_file.rs b/crates/pb-mapper-core/src/durable_file.rs new file mode 100644 index 0000000..95b730d --- /dev/null +++ b/crates/pb-mapper-core/src/durable_file.rs @@ -0,0 +1,76 @@ +//! Atomic replace and parent-directory durability. +//! +//! These are general file primitives, not credential ones — they lived in the +//! auth persistence layer only because that is where the first caller was. Both +//! report `io::Result` and leave it to the caller to map into its own error +//! type. + +use std::fs::File; +use std::path::Path; + +/// Replaces `to` with `from`, atomically where the platform allows it. +pub fn replace_file(from: &Path, to: &Path) -> std::io::Result<()> { + #[cfg(windows)] + { + use std::os::windows::ffi::OsStrExt; + + const MOVEFILE_REPLACE_EXISTING: u32 = 0x1; + const MOVEFILE_WRITE_THROUGH: u32 = 0x8; + unsafe extern "system" { + fn MoveFileExW( + lp_existing_file_name: *const u16, + lp_new_file_name: *const u16, + dw_flags: u32, + ) -> i32; + } + fn wide(path: &Path) -> Vec { + path.as_os_str() + .encode_wide() + .chain(std::iter::once(0)) + .collect() + } + let from_w = wide(from); + let to_w = wide(to); + let ok = unsafe { + MoveFileExW( + from_w.as_ptr(), + to_w.as_ptr(), + MOVEFILE_REPLACE_EXISTING | MOVEFILE_WRITE_THROUGH, + ) + }; + if ok == 0 { + Err(std::io::Error::last_os_error()) + } else { + Ok(()) + } + } + #[cfg(not(windows))] + std::fs::rename(from, to) +} + +/// Fsyncs the directory holding `path`, so a rename into it survives a crash. +/// +/// A path with no parent is a no-op rather than an error. +pub fn sync_parent_directory(path: &Path) -> std::io::Result<()> { + let Some(parent) = path.parent() else { + return Ok(()); + }; + open_directory_for_sync(parent).and_then(|directory| directory.sync_all()) +} + +fn open_directory_for_sync(path: &Path) -> std::io::Result { + #[cfg(windows)] + { + use std::fs::OpenOptions; + use std::os::windows::fs::OpenOptionsExt; + const GENERIC_READ: u32 = 0x8000_0000; + const GENERIC_WRITE: u32 = 0x4000_0000; + const FILE_FLAG_BACKUP_SEMANTICS: u32 = 0x0200_0000; + OpenOptions::new() + .access_mode(GENERIC_READ | GENERIC_WRITE) + .custom_flags(FILE_FLAG_BACKUP_SEMANTICS) + .open(path) + } + #[cfg(not(windows))] + File::open(path) +} diff --git a/src/common/error.rs b/crates/pb-mapper-core/src/error.rs similarity index 85% rename from src/common/error.rs rename to crates/pb-mapper-core/src/error.rs index 66bb572..b4fc24f 100644 --- a/src/common/error.rs +++ b/crates/pb-mapper-core/src/error.rs @@ -3,11 +3,13 @@ use std::net::AddrParseError; use snafu::Snafu; -use super::checksum::ChecksumType; -use super::message::DataLenType; +use crate::DataLenType; +use crate::checksum::ChecksumType; #[derive(Debug, Snafu)] -#[snafu(visibility(pub(super)))] +// The generated context selectors are constructed by the protocol, client, and +// server crates, so they have to be reachable from outside this one. +#[snafu(visibility(pub))] pub enum Error { /// Error handling for message #[snafu(display("read `checksum` from network error"))] @@ -59,42 +61,17 @@ pub enum Error { // specific error explanation detail: String, }, + #[snafu(display("protocol-v2 error: {detail}"))] + MsgProtocol { detail: String }, #[snafu(display("`{action}` forward message failed: {source}"))] MsgForward { // must be "read" or "write" action: &'static str, source: std::io::Error, }, - /// Error for manager - #[snafu(display("`TaskManager` fails while waiting for a task"))] - MngWaitForTask { source: kanal::ReceiveError }, /// Error for forward #[snafu(display("failed to forward message to write in normal text"))] FwdNetworkWriteWithNormal { source: std::io::Error }, - /// Error for stream - #[snafu(display("failed to connect stream, type:`{stream_type}`"))] - StmConnectStream { - // must be "UDP" or "TCP" - stream_type: &'static str, - source: std::io::Error, - }, - #[snafu(display("failed to got one addr from iter"))] - StmGotOneAddrFromIter, - #[snafu(display("failed to got one addr when parsing address"))] - StmGotOneAddr { source: std::io::Error }, - /// Error for listener - #[snafu(display("listener failed to bind addr, type:`{listener_type}`"))] - LsnListenerBind { - // must be "UDP" or "TCP" - listener_type: &'static str, - source: std::io::Error, - }, - #[snafu(display("listener failed to accept stream, type:`{listener_type}`"))] - LsnListenerAccept { - // must be "UDP" or "TCP" - listener_type: &'static str, - source: std::io::Error, - }, /// Error for config #[snafu(display("parse socket address from string:`{string}` error"))] CfgParseSockAddr { diff --git a/crates/pb-mapper-core/src/lib.rs b/crates/pb-mapper-core/src/lib.rs new file mode 100644 index 0000000..0d13345 --- /dev/null +++ b/crates/pb-mapper-core/src/lib.rs @@ -0,0 +1,20 @@ +//! The bottom layer: credential primitives, framing checksums, configuration, +//! address resolution, and the file primitives the durable stores build on. +//! +//! Nothing here depends on another `pb-mapper` crate, which is what makes it +//! the bottom. `DataLenType` lives here rather than with the message framing +//! that names it, so that `checksum` and `error` can use it without depending +//! on the protocol layer. + +pub mod addr; +pub mod checksum; +pub mod codec; +pub mod config; +pub mod conn_id; +pub mod durable_file; +pub mod error; +pub mod test_support; +pub mod timeout; + +/// The width of the length prefix on a framed message. +pub type DataLenType = u32; diff --git a/crates/pb-mapper-core/src/test_support.rs b/crates/pb-mapper-core/src/test_support.rs new file mode 100644 index 0000000..12f5cee --- /dev/null +++ b/crates/pb-mapper-core/src/test_support.rs @@ -0,0 +1,14 @@ +//! Shared serialisation for tests that mutate process-global credential state. + +/// Serialises tests that set or clear the process credential. +/// +/// The credential is process-global, so two such tests running on different +/// runner threads would see each other's writes. Every test that calls +/// `set_process_msg_header_key` — in this crate and in the auth, protocol, and +/// server crates — takes this first. +/// +/// It lives here, next to the state it guards, and is unconditionally `pub` +/// rather than `#[cfg(test)]`: a test-only item is not visible to another +/// crate's tests, because each crate compiles its own test configuration. +pub static PROCESS_CREDENTIAL_TEST_LOCK: std::sync::LazyLock> = + std::sync::LazyLock::new(|| tokio::sync::Mutex::new(())); diff --git a/src/utils/timeout.rs b/crates/pb-mapper-core/src/timeout.rs similarity index 100% rename from src/utils/timeout.rs rename to crates/pb-mapper-core/src/timeout.rs diff --git a/crates/pb-mapper-protocol/Cargo.toml b/crates/pb-mapper-protocol/Cargo.toml new file mode 100644 index 0000000..b4a948a --- /dev/null +++ b/crates/pb-mapper-protocol/Cargo.toml @@ -0,0 +1,26 @@ +[package] +name = "pb-mapper-protocol" +version.workspace = true +edition.workspace = true +authors.workspace = true + +[dependencies] +pb-mapper-auth.workspace = true +pb-mapper-core.workspace = true + +bytes.workspace = true +parking_lot.workspace = true +rand.workspace = true +ring.workspace = true +serde.workspace = true +serde_json.workspace = true +snafu.workspace = true +tokio.workspace = true +tracing.workspace = true +uni-stream.workspace = true + +[features] +udp-timeout = ["uni-stream/udp-timeout"] + +[lints] +workspace = true diff --git a/src/common/buffer.rs b/crates/pb-mapper-protocol/src/buffer.rs similarity index 92% rename from src/common/buffer.rs rename to crates/pb-mapper-protocol/src/buffer.rs index 57e3326..b53f572 100644 --- a/src/common/buffer.rs +++ b/crates/pb-mapper-protocol/src/buffer.rs @@ -3,7 +3,7 @@ use snafu::ResultExt; use tokio::io::AsyncReadExt; -use super::error::MsgNetworkReadBufferdRawDataSnafu; +use pb_mapper_core::error::MsgNetworkReadBufferdRawDataSnafu; const INIT_BUF_SIZE: usize = 8 * 1024; const MAX_BUF_SIZE: usize = 8 * 1024 * 1024; @@ -104,7 +104,7 @@ impl BufferGetter for CommonBuffer { /// This trait is used for buffered reads where the packet length is not known pub trait BufferedReader { - async fn read(&mut self) -> super::error::Result<&'_ [u8]>; + async fn read(&mut self) -> pb_mapper_core::error::Result<&'_ [u8]>; } pub struct BufferReader<'a, T> { @@ -119,7 +119,7 @@ impl<'reader, T: AsyncReadExt + Unpin> BufferReader<'reader, T> { } } - async fn read_inner(&mut self) -> super::error::Result<&[u8]> { + async fn read_inner(&mut self) -> pb_mapper_core::error::Result<&[u8]> { if self.buffer.need_resize() { self.buffer.dyn_resize() } @@ -134,7 +134,7 @@ impl<'reader, T: AsyncReadExt + Unpin> BufferReader<'reader, T> { } impl<'reader, T: AsyncReadExt + Unpin> BufferedReader for BufferReader<'reader, T> { - async fn read(&mut self) -> super::error::Result<&'_ [u8]> { + async fn read(&mut self) -> pb_mapper_core::error::Result<&'_ [u8]> { self.read_inner().await } } diff --git a/crates/pb-mapper-protocol/src/command.rs b/crates/pb-mapper-protocol/src/command.rs new file mode 100644 index 0000000..0da74fb --- /dev/null +++ b/crates/pb-mapper-protocol/src/command.rs @@ -0,0 +1,363 @@ +use serde::{Deserialize, Serialize}; +use snafu::ResultExt; + +use pb_mapper_auth::{ + AuthStatus, IssuedTemporaryKey, KeyPage, LegacyProtocolPolicy, TemporaryKeyMetadata, +}; +use pb_mapper_core::checksum::AesKeyType; +use pb_mapper_core::error::{MsgSerializeSnafu, Result}; + +pub const CONTROL_PROTOCOL_V2: u16 = 2; + +pub trait MessageSerializer { + fn encode(&self) -> Result>; + fn decode(msg: &[u8]) -> Result + where + Self: Sized; +} + +#[derive(Debug, Clone, PartialEq, Eq, Deserialize, Serialize)] +pub enum PbConnStatusReq { + RemoteId, + Keys, + Service { key: String }, +} + +#[derive(Debug, Clone, Deserialize, Serialize)] +pub enum PbConnStatusResp { + RemoteId { + server_map: String, + active: String, + idle: String, + }, + Keys(Vec), + Service { + key: String, + connections: Vec, + }, +} + +#[derive(Debug, Clone, PartialEq, Eq, Deserialize, Serialize)] +pub struct PbServiceConnStatus { + pub conn_id: u32, + pub generation: u64, + pub protocol_version: u16, + pub healthy: bool, + pub last_rx_age_ms: u64, +} + +#[derive(Debug, Clone, PartialEq, Eq, Deserialize, Serialize)] +pub enum PbConnRequest { + Register { + need_codec: bool, + is_datagram: bool, + key: String, + #[serde(default, skip_serializing_if = "Option::is_none")] + protocol_version: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + client_instance_id: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + heartbeat_interval_ms: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + heartbeat_tolerance_ms: Option, + }, + RegisterScoped { + need_codec: bool, + is_datagram: bool, + key: String, + namespace: u64, + force_namespace: bool, + #[serde(default, skip_serializing_if = "Option::is_none")] + protocol_version: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + client_instance_id: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + heartbeat_interval_ms: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + heartbeat_tolerance_ms: Option, + }, + Subcribe { + key: String, + }, + SubcribeScoped { + key: String, + namespace: u64, + }, + Status(PbConnStatusReq), + StatusScoped { + status: PbConnStatusReq, + namespace: u64, + }, + Stream { + key: String, + dst_id: u32, + #[serde(default)] + server_generation: u64, + }, + StreamScoped { + key: String, + namespace: u64, + dst_id: u32, + #[serde(default)] + server_generation: u64, + }, + Admin(AdminRequest), +} + +#[derive(Debug, Clone, PartialEq, Eq, Deserialize, Serialize)] +pub enum AdminRequest { + KeyIssue { + ttl_seconds: u64, + #[serde(default, skip_serializing_if = "Option::is_none")] + label: Option, + }, + KeyList { + #[serde(default)] + page: u32, + #[serde(default = "default_page_size")] + page_size: u16, + }, + KeyShow { + key_id: u64, + }, + KeyReveal { + key_id: u64, + }, + KeyRenew { + key_id: u64, + ttl_seconds: u64, + }, + KeyRevoke { + key_id: u64, + }, + KeyGc, + AuthStatus, + AuthStateReset { + confirm: bool, + }, + RootKeyRotate { + new_admin_key: String, + }, + LegacyProtocolSet { + policy: LegacyProtocolPolicy, + }, + ConnectionList { + #[serde(default, skip_serializing_if = "Option::is_none")] + key_id: Option, + #[serde(default)] + page: u32, + #[serde(default = "default_page_size")] + page_size: u16, + }, + ServiceList { + #[serde(default, skip_serializing_if = "Option::is_none")] + key_id: Option, + #[serde(default)] + page: u32, + #[serde(default = "default_page_size")] + page_size: u16, + }, +} + +impl AdminRequest { + pub fn is_mutating(&self) -> bool { + matches!( + self, + Self::KeyIssue { .. } + | Self::KeyRenew { .. } + | Self::KeyRevoke { .. } + | Self::KeyGc + | Self::AuthStateReset { .. } + | Self::RootKeyRotate { .. } + | Self::LegacyProtocolSet { .. } + ) + } +} + +const fn default_page_size() -> u16 { + 100 +} + +#[derive(Debug, Clone, Deserialize, Serialize)] +pub struct PbErrorResponse { + pub code: String, + pub message: String, + pub retryable: bool, + pub server_time: u64, +} + +#[derive(Debug, Clone, Deserialize, Serialize)] +pub struct AdminServiceInfo { + pub key_id: u64, + pub namespace: u64, + pub service_name: String, + pub transport: String, + pub codec_enabled: bool, + pub connection_count: u32, +} + +#[derive(Debug, Clone, Deserialize, Serialize)] +pub struct AdminConnectionInfo { + pub key_id: u64, + pub namespace: u64, + pub service_name: String, + pub conn_id: u32, + pub generation: u64, + pub protocol_version: u16, + pub healthy: bool, + pub transport: String, + pub codec_enabled: bool, + pub last_rx_age_ms: u64, +} + +#[derive(Debug, Clone, Deserialize, Serialize)] +pub struct AdminServicePage { + pub schema_version: u16, + pub items: Vec, + pub next_page: Option, +} + +#[derive(Debug, Clone, Deserialize, Serialize)] +pub struct AdminConnectionPage { + pub schema_version: u16, + pub items: Vec, + pub next_page: Option, +} + +#[derive(Debug, Clone, Deserialize, Serialize)] +pub enum AdminResponse { + KeyIssued(IssuedTemporaryKey), + KeyList(KeyPage), + KeyShown(IssuedTemporaryKey), + KeyRenewed(IssuedTemporaryKey), + KeyRevoked(TemporaryKeyMetadata), + KeyGc { removed: u64 }, + AuthStatus(AuthStatus), + Services(AdminServicePage), + Connections(AdminConnectionPage), + Ok { action: String }, +} + +#[derive(Debug, Clone, Deserialize, Serialize)] +pub enum PbConnResponse { + Register(u32), + RegisterV2 { + conn_id: u32, + generation: u64, + lease_ttl_ms: u64, + }, + Subcribe { + codec_key: Option, + client_id: u32, + server_id: u32, + }, + Stream { + codec_key: Option, + }, + Status(PbConnStatusResp), + Admin(AdminResponse), + Error(PbErrorResponse), +} + +impl PbConnResponse { + pub fn error(code: impl Into, message: impl Into, retryable: bool) -> Self { + Self::Error(PbErrorResponse { + code: code.into(), + message: message.into(), + retryable, + server_time: std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap_or_default() + .as_secs(), + }) + } +} + +#[derive(Debug, Clone, Deserialize, Serialize)] +pub enum PbServerRequest { + Ping, + PingV2 { + seq: u64, + }, + StreamAck { + client_id: u32, + #[serde(default)] + server_generation: u64, + }, +} + +#[derive(Debug, Clone, Deserialize, Serialize)] +pub enum LocalServer { + /// pb server makes a stream request to local server + Stream { + client_id: u32, + #[serde(default)] + server_generation: u64, + }, + /// pb server response a pong msg when it receive a ping request + Pong, + PongV2 { + seq: u64, + }, + Retire { + reason: String, + conn_id: u32, + #[serde(default)] + server_generation: u64, + }, +} + +macro_rules! gen_impl_msg_serializer { + ($struct_name:ident) => { + impl MessageSerializer for $struct_name { + fn encode(&self) -> Result> { + serde_json::to_vec(self).with_context(|_| MsgSerializeSnafu { + action: "encode", + struct_name: stringify!($struct_name), + content: "payload redacted".to_string(), + }) + } + + fn decode(msg: &[u8]) -> Result { + serde_json::from_slice(msg).with_context(|_| MsgSerializeSnafu { + action: "decode", + struct_name: stringify!($struct_name), + content: format!("{}-byte payload redacted", msg.len()), + }) + } + } + }; +} + +gen_impl_msg_serializer!(PbConnRequest); +gen_impl_msg_serializer!(PbConnResponse); +gen_impl_msg_serializer!(PbServerRequest); +gen_impl_msg_serializer!(LocalServer); + +#[cfg(test)] +mod tests { + use super::PbConnRequest; + + /// The wire form of `Register` is load-bearing: a running peer on the other + /// side of an upgrade has to keep parsing it. The `None` fields must stay + /// absent from the JSON rather than serialise as null. + #[test] + fn test_serde_mapper_header() { + let mapper = PbConnRequest::Register { + key: "test".into(), + need_codec: false, + is_datagram: false, + protocol_version: None, + client_instance_id: None, + heartbeat_interval_ms: None, + heartbeat_tolerance_ms: None, + }; + let json_value = serde_json::to_string(&mapper).unwrap(); + let raw_json_str = + r##"{"Register":{"need_codec":false,"is_datagram":false,"key":"test"}}"##; + assert_eq!(raw_json_str, json_value); + + let value: PbConnRequest = serde_json::from_str(raw_json_str).unwrap(); + assert_eq!(mapper, value) + } +} diff --git a/src/common/message/forward.rs b/crates/pb-mapper-protocol/src/forward.rs similarity index 80% rename from src/common/message/forward.rs rename to crates/pb-mapper-protocol/src/forward.rs index d25cd82..d8af9a7 100644 --- a/src/common/message/forward.rs +++ b/crates/pb-mapper-protocol/src/forward.rs @@ -6,20 +6,16 @@ use std::time::Duration; use tokio::io::{AsyncReadExt, AsyncWriteExt}; use tokio::time::Instant; -use super::super::buffer::{BufferReader, BufferedReader}; -use super::error::{FwdNetworkWriteWithNormalSnafu, Result}; use super::{ CodecMessageReader, CodecMessageWriter, MessageReader, MessageWriter, NormalMessageReader, NormalMessageWriter, }; -use crate::common::checksum::AesKeyType; -use crate::common::config::duration_from_env; -use crate::common::message::{get_decodec, get_encodec}; -use crate::utils::codec::{Decryptor, Encryptor}; -use crate::{ - create_component, snafu_error_get_or_return_ok, start_datagram_forward_with_codec_key, - start_forward_with_codec_key, -}; +use crate::buffer::{BufferReader, BufferedReader}; +use pb_mapper_core::checksum::AesKeyType; +use pb_mapper_core::codec::{Decryptor, Encryptor}; +use pb_mapper_core::config::duration_from_env; +use pb_mapper_core::error::{FwdNetworkWriteWithNormalSnafu, Result}; +use pb_mapper_core::snafu_error_get_or_return_ok; use uni_stream::stream::{StreamSplit, TcpStreamImpl, UdpStreamImpl}; use uni_stream::udp::{UdpStreamReadHalf, UdpStreamWriteHalf}; @@ -96,6 +92,12 @@ impl<'a, T: AsyncReadExt + Unpin + Send> NormalDatagramReader<'a, T> { reader: NormalMessageReader::new(reader), } } + + pub fn with_checksum_key(self, key: AesKeyType) -> Self { + Self { + reader: self.reader.with_checksum_key(key), + } + } } impl<'a, T: AsyncReadExt + Unpin + Send> DatagramReader for NormalDatagramReader<'a, T> { @@ -142,6 +144,12 @@ impl<'a, T: AsyncWriteExt + Unpin + Send> NormalDatagramWriter<'a, T> { writer: NormalMessageWriter::new(writer), } } + + pub fn with_checksum_key(self, key: AesKeyType) -> Self { + Self { + writer: self.writer.with_checksum_key(key), + } + } } impl<'a, T: AsyncWriteExt + Unpin + Send> DatagramWriter for NormalDatagramWriter<'a, T> { @@ -158,6 +166,10 @@ impl<'a, T: AsyncReadExt + Send + Unpin, D: Decryptor> CodecForwardReader<'a, T, pub fn new(reader: &'a mut T, decryptor: D) -> Self { Self(CodecMessageReader::new(reader, decryptor)) } + + pub fn with_checksum_key(self, key: AesKeyType) -> Self { + Self(self.0.with_checksum_key(key)) + } } impl<'a, T: AsyncReadExt + Send + Unpin, D: Decryptor> ForwardReader @@ -176,6 +188,10 @@ impl<'a, T: AsyncReadExt + Send + Unpin, D: Decryptor> CodecDatagramReader<'a, T pub fn new(reader: &'a mut T, decryptor: D) -> Self { Self(CodecMessageReader::new(reader, decryptor)) } + + pub fn with_checksum_key(self, key: AesKeyType) -> Self { + Self(self.0.with_checksum_key(key)) + } } impl<'a, T: AsyncReadExt + Send + Unpin, D: Decryptor> DatagramReader @@ -195,6 +211,10 @@ impl<'a, T: AsyncWriteExt + Send + Unpin, E: Encryptor> CodecForwardWriter<'a, T pub fn new(writer: &'a mut T, encryptor: E) -> Self { Self(CodecMessageWriter::new(writer, encryptor)) } + + pub fn with_checksum_key(self, key: AesKeyType) -> Self { + Self(self.0.with_checksum_key(key)) + } } impl<'a, T: AsyncWriteExt + Send + Unpin, E: Encryptor> ForwardWriter @@ -218,6 +238,10 @@ impl<'a, T: AsyncWriteExt + Send + Unpin, E: Encryptor> CodecDatagramWriter<'a, pub fn new(writer: &'a mut T, encryptor: E) -> Self { Self(CodecMessageWriter::new(writer, encryptor)) } + + pub fn with_checksum_key(self, key: AesKeyType) -> Self { + Self(self.0.with_checksum_key(key)) + } } impl<'a, T: AsyncWriteExt + Send + Unpin, E: Encryptor> DatagramWriter @@ -509,7 +533,7 @@ impl DatagramReader for UdpStreamReadHalf { async fn recv(&mut self) -> Result { self.recv_datagram() .await - .map_err(|e| super::error::Error::MsgForward { + .map_err(|e| pb_mapper_core::error::Error::MsgForward { action: "read", source: e, }) @@ -520,7 +544,7 @@ impl DatagramWriter for UdpStreamWriteHalf<'_> { async fn send(&mut self, src: &[u8]) -> Result<()> { self.send_datagram(src) .await - .map_err(|e| super::error::Error::MsgForward { + .map_err(|e| pb_mapper_core::error::Error::MsgForward { action: "write", source: e, }) @@ -530,6 +554,7 @@ impl DatagramWriter for UdpStreamWriteHalf<'_> { pub trait StreamForward: StreamSplit + Sized { fn forward_local_to_remote<'a, R, W>( codec_key: Option, + framing_key: AesKeyType, local_reader: Self::ReaderRef<'a>, local_writer: Self::WriterRef<'a>, remote_reader: R, @@ -543,6 +568,7 @@ pub trait StreamForward: StreamSplit + Sized { impl StreamForward for TcpStreamImpl { fn forward_local_to_remote<'a, R, W>( codec_key: Option, + framing_key: AesKeyType, local_reader: Self::ReaderRef<'a>, local_writer: Self::WriterRef<'a>, remote_reader: R, @@ -557,17 +583,40 @@ impl StreamForward for TcpStreamImpl { let mut local_writer = local_writer; let mut remote_reader = remote_reader; let mut remote_writer = remote_writer; - start_forward_with_codec_key!( - codec_key, - &mut local_reader, - &mut local_writer, - &mut remote_reader, - &mut remote_writer, - false, - false, - true, - true - ); + match codec_key { + Some(key) => { + start_forward( + NormalForwardReader::new(&mut local_reader), + NormalForwardWriter::new(&mut local_writer), + CodecForwardReader::new( + &mut remote_reader, + snafu_error_get_or_return_ok!( + super::get_decodec(&key), + "failed to create decoder when remote forward" + ), + ) + .with_checksum_key(framing_key), + CodecForwardWriter::new( + &mut remote_writer, + snafu_error_get_or_return_ok!( + super::get_encodec(&key), + "failed to create encoder when remote forward" + ), + ) + .with_checksum_key(framing_key), + ) + .await; + } + None => { + start_forward( + NormalForwardReader::new(&mut local_reader), + NormalForwardWriter::new(&mut local_writer), + NormalForwardReader::new(&mut remote_reader), + NormalForwardWriter::new(&mut remote_writer), + ) + .await; + } + } Ok(()) }) } @@ -576,6 +625,7 @@ impl StreamForward for TcpStreamImpl { impl StreamForward for UdpStreamImpl { fn forward_local_to_remote<'a, R, W>( codec_key: Option, + framing_key: AesKeyType, local_reader: Self::ReaderRef<'a>, local_writer: Self::WriterRef<'a>, remote_reader: R, @@ -588,174 +638,59 @@ impl StreamForward for UdpStreamImpl { Box::pin(async move { let mut remote_reader = remote_reader; let mut remote_writer = remote_writer; - start_datagram_forward_with_codec_key!( - codec_key, - local_reader, - local_writer, - &mut remote_reader, - &mut remote_writer - ); + match codec_key { + Some(key) => { + start_datagram_forward( + local_reader, + local_writer, + CodecDatagramReader::new( + &mut remote_reader, + snafu_error_get_or_return_ok!( + super::get_decodec(&key), + "failed to create decoder when datagram forward" + ), + ) + .with_checksum_key(framing_key), + CodecDatagramWriter::new( + &mut remote_writer, + snafu_error_get_or_return_ok!( + super::get_encodec(&key), + "failed to create encoder when datagram forward" + ), + ) + .with_checksum_key(framing_key), + ) + .await; + } + None => { + start_datagram_forward( + local_reader, + local_writer, + NormalDatagramReader::new(&mut remote_reader) + .with_checksum_key(framing_key), + NormalDatagramWriter::new(&mut remote_writer) + .with_checksum_key(framing_key), + ) + .await; + } + } Ok(()) }) } } -#[macro_export] -macro_rules! create_component { - (Reader, $stream:expr,true, $key:expr, $get_codec:ident, $name:expr) => { - CodecForwardReader::new( - $stream, - snafu_error_get_or_return_ok!( - $get_codec(&$key), - concat!("failed to create decoder when `", $name, "` forward msg") - ), - ) - }; - (Reader, $stream:expr,false, $key:expr, $get_codec:ident, $name:expr) => { - NormalForwardReader::new($stream) - }; - (Writer, $stream:expr,true, $key:expr, $get_codec:ident, $name:expr) => { - CodecForwardWriter::new( - $stream, - snafu_error_get_or_return_ok!( - $get_codec(&$key), - concat!("failed to create encoder when `", $name, "` forward msg") - ), - ) - }; - (Writer, $stream:expr,false, $key:expr, $get_codec:ident, $name:expr) => { - NormalForwardWriter::new($stream) - }; -} - -/// When using it, please remember to manually import the following symbols: -/// - [`start_forward`] -/// - [`crate::create_component`] -/// - [`ForwardReader`] -/// - [`ForwardWriter`] -/// - [`CodecForwardReader`] -/// - [`CodecForwardWriter`] -/// - [`crate::snafu_error_get_or_return_ok`] -/// - [`super::get_decodec`] -/// - [`super::get_encodec`] -#[macro_export] -macro_rules! start_forward_with_codec_key { - ( - $codec_key:expr, - $client_reader:expr, - $client_writer:expr, - $server_reader:expr, - $server_writer:expr, - $client_reader_codec:tt, - $client_writer_codec:tt, - $server_reader_codec:tt, - $server_writer_codec:tt - ) => { - match $codec_key { - Some(key) => { - (start_forward( - create_component!( - Reader, - $client_reader, - $client_reader_codec, - key, - get_decodec, - "client_reader" - ), - create_component!( - Writer, - $client_writer, - $client_writer_codec, - key, - get_encodec, - "client_writer" - ), - create_component!( - Reader, - $server_reader, - $server_reader_codec, - key, - get_decodec, - "server_reader" - ), - create_component!( - Writer, - $server_writer, - $server_writer_codec, - key, - get_encodec, - "server_writer" - ), - ) - .await) - } - None => { - (start_forward( - NormalForwardReader::new($client_reader), - NormalForwardWriter::new($client_writer), - NormalForwardReader::new($server_reader), - NormalForwardWriter::new($server_writer), - ) - .await) - } - } - }; -} - -#[macro_export] -macro_rules! start_datagram_forward_with_codec_key { - ( - $codec_key:expr, - $udp_reader:expr, - $udp_writer:expr, - $tcp_reader:expr, - $tcp_writer:expr - ) => { - match $codec_key { - Some(key) => { - (start_datagram_forward( - $udp_reader, - $udp_writer, - CodecDatagramReader::new( - $tcp_reader, - snafu_error_get_or_return_ok!( - $crate::common::message::get_decodec(&key), - "failed to create decoder when datagram forward" - ), - ), - CodecDatagramWriter::new( - $tcp_writer, - snafu_error_get_or_return_ok!( - $crate::common::message::get_encodec(&key), - "failed to create encoder when datagram forward" - ), - ), - ) - .await) - } - None => { - (start_datagram_forward( - $udp_reader, - $udp_writer, - NormalDatagramReader::new($tcp_reader), - NormalDatagramWriter::new($tcp_writer), - ) - .await) - } - } - }; -} - #[cfg(test)] mod tests { use std::collections::VecDeque; use std::io; - use std::sync::{Arc, Mutex}; + use std::sync::Arc; + + use parking_lot::Mutex; use std::time::Duration; use super::*; - use crate::common::config::parse_duration; - use crate::common::error::Error; + use pb_mapper_core::config::parse_duration; + use pb_mapper_core::error::Error; use tokio::sync::Notify; enum ReadAction { @@ -812,22 +747,22 @@ mod tests { impl ScriptedWriter { fn chunks(&self) -> Vec> { - self.state.lock().unwrap().chunks.clone() + self.state.lock().chunks.clone() } fn shutdowns(&self) -> usize { - self.state.lock().unwrap().shutdowns + self.state.lock().shutdowns } } impl ForwardWriter for ScriptedWriter { async fn write(&mut self, src: &[u8]) -> Result<()> { - self.state.lock().unwrap().chunks.push(src.to_vec()); + self.state.lock().chunks.push(src.to_vec()); Ok(()) } async fn shutdown(&mut self) { - self.state.lock().unwrap().shutdowns += 1; + self.state.lock().shutdowns += 1; } } @@ -875,7 +810,7 @@ mod tests { } fn chunks(&self) -> Vec> { - self.state.lock().unwrap().chunks.clone() + self.state.lock().chunks.clone() } } @@ -883,12 +818,12 @@ mod tests { async fn write(&mut self, src: &[u8]) -> Result<()> { self.write_started.notify_one(); tokio::time::sleep(self.delay).await; - self.state.lock().unwrap().chunks.push(src.to_vec()); + self.state.lock().chunks.push(src.to_vec()); Ok(()) } async fn shutdown(&mut self) { - self.state.lock().unwrap().shutdowns += 1; + self.state.lock().shutdowns += 1; } } diff --git a/src/common/message/mod.rs b/crates/pb-mapper-protocol/src/lib.rs similarity index 67% rename from src/common/message/mod.rs rename to crates/pb-mapper-protocol/src/lib.rs index 0477ed8..29ae36c 100644 --- a/src/common/message/mod.rs +++ b/crates/pb-mapper-protocol/src/lib.rs @@ -1,20 +1,31 @@ //! Define message protocols and tools for reading and writing //! messages +//! +//! The reader and writer traits are `async fn` in a public trait, which cannot +//! state its auto-trait bounds. That is deliberate: these are only ever awaited +//! on the connection task that owns the stream, never sent across one. +#![allow(async_fn_in_trait)] + +pub mod buffer; pub mod command; pub mod forward; -use snafu::{ensure, ResultExt}; +pub mod secure; +use snafu::{ResultExt, ensure}; use tokio::io::{AsyncReadExt, AsyncWriteExt}; -use super::buffer::{BufferGetter, CommonBuffer, FixedSizeBuffer}; -use super::checksum::{get_checksum, get_msg_header_key, valid_checksum}; -use super::error::{ +use crate::buffer::{BufferGetter, CommonBuffer, FixedSizeBuffer}; +use pb_mapper_core::checksum::{ + AesKeyType, get_checksum, get_checksum_for_key, get_msg_header_key, process_checksum_is_ready, + valid_checksum, valid_checksum_for_key, +}; +use pb_mapper_core::codec::{Aes256GcmDeCodec, Aes256GcmEnCodec, Decryptor, Encryptor}; +use pb_mapper_core::error::MsgDatalenExceededSnafu; +use pb_mapper_core::error::{ self, MsgDatalenValidateSnafu, MsgNetworkReadBodySnafu, MsgNetworkReadCheckSumSnafu, MsgNetworkReadDatalenSnafu, MsgNetworkWriteBodySnafu, MsgNetworkWriteCheckSumSnafu, MsgNetworkWriteCodecMsgSnafu, MsgNetworkWriteCodecTagSnafu, MsgNetworkWriteDatalenSnafu, Result, }; -use crate::common::error::MsgDatalenExceededSnafu; -use crate::utils::codec::{Aes256GcmDeCodec, Aes256GcmEnCodec, Decryptor, Encryptor}; /// This message protocol contains header and body, and the header /// includes checksum, datalen,respectively, u32, u32, where datalen @@ -65,7 +76,11 @@ const CODEC_TAG_LEN: DataLenType = 16; /// For encrypted frames, the tag is appended to the payload. const MAX_MSG_LEN: DataLenType = MAX_PLAINTEXT_LEN + CODEC_TAG_LEN; -pub type DataLenType = u32; +// Defined in `pb-mapper-core` so that the checksum and error types can name it +// without depending on this module. Re-exported rather than redeclared: a second +// `pub type` would be a distinct name for the same width, and the two would read +// as unrelated at the crate boundary. +pub use pb_mapper_core::DataLenType; macro_rules! gen_read_network_with_error { ($func_name:ident, $read_method:ident, $error:expr, $return_ty:ty) => { @@ -129,11 +144,44 @@ gen_write_network_with_error!( &[u8] ); +fn checksum_key_bytes(key: &Option) -> Option<&[u8]> { + key.as_ref().map(|key| key.as_slice()) +} + +#[inline] +fn checksum_matches(datalen: DataLenType, checksum: u32, key: Option<&[u8]>) -> bool { + match key { + Some(key) => valid_checksum_for_key(datalen, checksum, key), + None => valid_checksum(datalen, checksum), + } +} + #[inline] -async fn get_msg_len(reader: &mut T) -> Result { +fn checksum_for(len: DataLenType, key: Option<&[u8]>) -> Result { + match key { + Some(key) => Ok(get_checksum_for_key(len, key)), + None => { + if !process_checksum_is_ready() { + return Err(error::Error::MsgCodec { + action: "load configured credential", + detail: + "`MSG_HEADER_KEY` is required; no insecure default checksum is available" + .to_string(), + }); + } + Ok(get_checksum(len)) + } + } +} + +#[inline] +async fn get_msg_len( + reader: &mut T, + checksum_key: Option<&[u8]>, +) -> Result { let checksum = read_checksum(reader).await?; let datalen = read_datalen(reader).await?; - if valid_checksum(datalen, checksum) { + if checksum_matches(datalen, checksum, checksum_key) { ensure!( datalen <= MAX_MSG_LEN, MsgDatalenExceededSnafu { @@ -148,14 +196,19 @@ async fn get_msg_len(reader: &mut T) -> Result(writer: &mut T, len: DataLenType) -> Result<()> { - write_checksum(writer, get_checksum(len)).await?; +async fn set_msg_len( + writer: &mut T, + len: DataLenType, + checksum_key: Option<&[u8]>, +) -> Result<()> { + write_checksum(writer, checksum_for(len, checksum_key)?).await?; write_datalen(writer, len).await } pub struct NormalMessageReader<'a, T: AsyncReadExt + Unpin> { reader: &'a mut T, buffer: CommonBuffer, + checksum_key: Option, } impl<'a, T: AsyncReadExt + Unpin> NormalMessageReader<'a, T> { @@ -163,11 +216,17 @@ impl<'a, T: AsyncReadExt + Unpin> NormalMessageReader<'a, T> { Self { reader, buffer: CommonBuffer::new(), + checksum_key: None, } } + pub fn with_checksum_key(mut self, key: AesKeyType) -> Self { + self.checksum_key = Some(key); + self + } + async fn read_msg_inner(&mut self) -> Result<&'_ [u8]> { - let datalen = get_msg_len(&mut self.reader).await?; + let datalen = get_msg_len(&mut self.reader, checksum_key_bytes(&self.checksum_key)).await?; self.buffer.fixed_resize(datalen as usize); let n = read_msg_body(&mut self.reader, self.buffer.buffer_mut()).await?; Ok(&self.buffer.buffer()[0..n]) @@ -182,15 +241,29 @@ impl<'a, T: AsyncReadExt + Unpin> MessageReader for NormalMessageReader<'a, T> { pub struct NormalMessageWriter<'a, T: AsyncWriteExt> { writer: &'a mut T, + checksum_key: Option, } impl<'a, T: AsyncWriteExt + Unpin> NormalMessageWriter<'a, T> { pub fn new(writer: &'a mut T) -> Self { - Self { writer } + Self { + writer, + checksum_key: None, + } + } + + pub fn with_checksum_key(mut self, key: AesKeyType) -> Self { + self.checksum_key = Some(key); + self } async fn write_msg_inner(&mut self, msg: &[u8]) -> Result<()> { - set_msg_len(&mut self.writer, msg.len() as u32).await?; + set_msg_len( + &mut self.writer, + msg.len() as u32, + checksum_key_bytes(&self.checksum_key), + ) + .await?; write_msg_body(&mut self.writer, msg).await } @@ -214,6 +287,18 @@ impl<'a, T: AsyncReadExt + Unpin, D: Decryptor> CodecMessageReader<'a, T, D> { decryptor, } } + + /// Bind the length checksum to `key` instead of the process credential. + /// Isolated relays keep a remote `MSG_HEADER_KEY` while speaking with a + /// different local administrator key. + pub fn for_session_key(reader: &'a mut T, decryptor: D, key: AesKeyType) -> Self { + Self::new(reader, decryptor).with_checksum_key(key) + } + + pub fn with_checksum_key(mut self, key: AesKeyType) -> Self { + self.reader.checksum_key = Some(key); + self + } } impl<'a, T: AsyncReadExt + Unpin, D: Decryptor> MessageReader for CodecMessageReader<'a, T, D> { @@ -235,11 +320,27 @@ impl<'a, T: AsyncReadExt + Unpin, D: Decryptor> MessageReader for CodecMessageRe pub struct CodecMessageWriter<'a, T: AsyncWriteExt + Unpin, E: Encryptor> { writer: &'a mut T, encryptor: E, + /// `None` uses the process `MSG_HEADER_KEY` hash. Isolated relays must set + /// this to the session key so continuation frames stay decryptable. + checksum_key: Option, } impl<'a, T: AsyncWriteExt + Unpin, E: Encryptor> CodecMessageWriter<'a, T, E> { pub fn new(writer: &'a mut T, encryptor: E) -> Self { - Self { writer, encryptor } + Self { + writer, + encryptor, + checksum_key: None, + } + } + + pub fn for_session_key(writer: &'a mut T, encryptor: E, key: AesKeyType) -> Self { + Self::new(writer, encryptor).with_checksum_key(key) + } + + pub fn with_checksum_key(mut self, key: AesKeyType) -> Self { + self.checksum_key = Some(key); + self } pub async fn shutdown(&mut self) -> std::io::Result<()> { @@ -259,7 +360,7 @@ impl<'a, T: AsyncWriteExt + Unpin, E: Encryptor> MessageWriter for CodecMessageW })?; let msg_len = (buf.len() + tag.as_ref().len()) as DataLenType; - set_msg_len(self.writer, msg_len).await?; + set_msg_len(self.writer, msg_len, checksum_key_bytes(&self.checksum_key)).await?; write_codec_msg(self.writer, &buf).await?; write_codec_tag(self.writer, tag.as_ref()).await } @@ -281,7 +382,10 @@ pub fn get_header_msg_writer( #[inline] pub fn get_default_encodec() -> Result { - let key = get_msg_header_key(); + let key = get_msg_header_key().map_err(|detail| error::Error::MsgCodec { + action: "load configured credential", + detail, + })?; Aes256GcmEnCodec::try_new(&key).map_err(|e| error::Error::MsgCodec { action: "create default encodec", detail: format!("{e}"), @@ -290,7 +394,10 @@ pub fn get_default_encodec() -> Result { #[inline] pub fn get_default_decodec() -> Result { - let key = get_msg_header_key(); + let key = get_msg_header_key().map_err(|detail| error::Error::MsgCodec { + action: "load configured credential", + detail, + })?; Aes256GcmDeCodec::try_new(&key).map_err(|e| error::Error::MsgCodec { action: "create default decodec", detail: format!("{e}"), diff --git a/crates/pb-mapper-protocol/src/secure.rs b/crates/pb-mapper-protocol/src/secure.rs new file mode 100644 index 0000000..98a6c16 --- /dev/null +++ b/crates/pb-mapper-protocol/src/secure.rs @@ -0,0 +1,711 @@ +//! Protocol-v2 single-flight authentication framing. +//! +//! The first client frame carries a clear-text routing prefix and an authenticated encrypted +//! request. It does not add a handshake or round trip. All following control messages on the +//! same TCP connection use independently derived directional keys and monotonically increasing +//! 64-bit counters. +//! +//! ```text +//! first flight: PBM2 | version | key id | timestamp+salt | counter | len | ciphertext +//! | | | +//! | +-> replay/time checks +-> bounded AEAD open +//! +-> derive directional session keys +//! +//! continuation: counter(n+1) | len | ciphertext -> same authenticated session +//! ``` +//! +//! This root module coordinates client/server sessions. Frame mechanics, replay admission, +//! log suppression, and protocol tests are isolated in focused child modules. + +use std::sync::Arc; + +use parking_lot::Mutex; +use std::time::{Duration, SystemTime, UNIX_EPOCH}; + +use rand::RngExt; +use ring::aead::{AES_256_GCM, Aad, LessSafeKey, Nonce, UnboundKey}; +use ring::digest::{SHA256, digest}; +use ring::hkdf::{HKDF_SHA256, Salt}; +use tokio::io::{AsyncReadExt, AsyncWriteExt}; + +use super::{ + CodecMessageReader, CodecMessageWriter, DataLenType, MAX_MSG_LEN, MessageReader, MessageWriter, +}; +use pb_mapper_auth::{ + ADMIN_KEY_ID, AuthContext, AuthFailure, AuthRuntime, KeyId, LegacyConnectionGuard, +}; +use pb_mapper_core::checksum::{ + AesKeyType, Credential, get_process_credential, valid_checksum_for_key, +}; +use pb_mapper_core::codec::{Aes256GcmDeCodec, Aes256GcmEnCodec, Decryptor}; +use pb_mapper_core::error::{Error, Result}; + +pub const PROTOCOL_V2_MAGIC: [u8; 4] = *b"PBM2"; +pub const PROTOCOL_V2_VERSION: u8 = 2; +const CONNECTION_SALT_LEN: usize = 16; +const FIRST_PREFIX_REMAINDER_LEN: usize = 28; +const FRAME_HEADER_LEN: usize = 12; +const DIRECTION_CLIENT_TO_SERVER: u8 = 0; +const DIRECTION_SERVER_TO_CLIENT: u8 = 1; +const MAX_CONNECTION_CLOCK_SKEW_SECONDS: u64 = 5 * 60; +/// Each Bloom generation must outlive the accepted clock-skew interval. A +/// salt inserted at the end of a window with `ts = now + skew` stays valid +/// until `insert + 2*skew`, so one generation is `2 * skew`. +const DEFAULT_REPLAY_WINDOW_SECONDS: u64 = MAX_CONNECTION_CLOCK_SKEW_SECONDS.saturating_mul(2); +const DEFAULT_REPLAY_FILTER_BYTES: usize = 1024 * 1024; +const MAX_INITIAL_PLAINTEXT_LEN: u32 = 64 * 1024; +const MAX_INITIAL_CIPHERTEXT_LEN: u32 = MAX_INITIAL_PLAINTEXT_LEN + 16; + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub enum HeaderProtocol { + Legacy, + V2, +} + +pub struct ClientHeaderSession { + protocol: HeaderProtocol, + legacy_key: AesKeyType, + v2: Option, +} + +impl ClientHeaderSession { + /// The v2 material, which is `Some` exactly when `protocol` is `V2`. + /// + /// The type does not tie the two together, so this reports instead of + /// panicking on a state that construction never produces. + fn v2_material(&self) -> Result<&V2Material> { + self.v2 + .as_ref() + .ok_or_else(|| protocol_error("v2 session is missing its key material")) + } + + /// New clients always use protocol v2, for both administrator and temporary credentials. + pub fn from_process() -> Result { + let credential = get_process_credential().map_err(protocol_error)?; + Self::new_v2(&credential) + } + + pub fn new_v2(credential: &Credential) -> Result { + let mut salt = [0_u8; CONNECTION_SALT_LEN]; + salt[..8].copy_from_slice(&unix_seconds().to_be_bytes()); + let mut rng = rand::rng(); + for byte in &mut salt[8..] { + *byte = rng.random(); + } + let material = + derive_material(KeyId::from_u64(credential.key_id()), credential.key(), salt)?; + Ok(Self { + protocol: HeaderProtocol::V2, + legacy_key: *credential.key(), + v2: Some(material), + }) + } + + #[cfg(test)] + pub fn new_legacy(key: AesKeyType) -> Self { + Self { + protocol: HeaderProtocol::Legacy, + legacy_key: key, + v2: None, + } + } + + pub fn protocol(&self) -> HeaderProtocol { + self.protocol + } + + pub async fn write_initial( + &self, + writer: &mut T, + message: &[u8], + ) -> Result<()> { + match self.protocol { + HeaderProtocol::Legacy => { + legacy_message_writer(writer, &self.legacy_key, "legacy writer")? + .write_msg(message) + .await + } + HeaderProtocol::V2 => { + let material = self.v2_material()?; + writer + .write_all(&first_prefix(material)) + .await + .map_err(|error| { + protocol_error(format!("failed to write v2 prefix: {error}")) + })?; + V2MessageWriter::new(writer, material.clone(), DIRECTION_CLIENT_TO_SERVER, 0)? + .write_msg(message) + .await + } + } + } + + pub fn response_reader<'a, T: AsyncReadExt + Unpin>( + &self, + reader: &'a mut T, + ) -> Result> { + match self.protocol { + HeaderProtocol::Legacy => Ok(HeaderMessageReader::Legacy(legacy_message_reader( + reader, + &self.legacy_key, + "legacy reader", + )?)), + HeaderProtocol::V2 => Ok(HeaderMessageReader::V2(V2MessageReader::new( + reader, + self.v2_material()?.clone(), + DIRECTION_SERVER_TO_CLIENT, + 0, + )?)), + } + } + + pub async fn exchange( + &self, + stream: &mut T, + payload: &[u8], + timeout: Duration, + ) -> Result> { + match tokio::time::timeout(timeout, self.write_initial(stream, payload)).await { + Ok(result) => result?, + Err(_) => { + return Err(protocol_error(format!( + "timed out writing first-flight request after {timeout:?}" + ))); + } + } + let mut reader = self.response_reader(stream)?; + let message = match tokio::time::timeout(timeout, reader.read_msg()).await { + Ok(result) => result?, + Err(_) => { + return Err(protocol_error(format!( + "timed out reading first-flight response after {timeout:?}" + ))); + } + }; + Ok(message.to_vec()) + } + + pub fn continuation_writer<'a, T: AsyncWriteExt + Unpin>( + &self, + writer: &'a mut T, + ) -> Result> { + match self.protocol { + HeaderProtocol::Legacy => Ok(HeaderMessageWriter::Legacy(legacy_message_writer( + writer, + &self.legacy_key, + "legacy writer", + )?)), + HeaderProtocol::V2 => Ok(HeaderMessageWriter::V2(V2MessageWriter::new( + writer, + self.v2_material()?.clone(), + DIRECTION_CLIENT_TO_SERVER, + 1, + )?)), + } + } +} + +pub struct ServerHeaderSession { + protocol: HeaderProtocol, + legacy_key: AesKeyType, + v2: Option, + context: Option, + _legacy_guard: Option, +} + +impl fmt::Debug for ServerHeaderSession { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("ServerHeaderSession") + .field("protocol", &self.protocol) + .field("key_id", &self.key_id()) + .field("authenticated", &self.context.is_some()) + .finish() + } +} + +impl ServerHeaderSession { + /// The v2 material, `Some` exactly when `protocol` is `V2` — see + /// [`ClientHeaderSession::v2_material`]. + fn v2_material(&self) -> Result<&V2Material> { + self.v2 + .as_ref() + .ok_or_else(|| protocol_error("v2 session is missing its key material")) + } + + pub fn protocol(&self) -> HeaderProtocol { + self.protocol + } + + pub fn framing_key(&self) -> AesKeyType { + self.legacy_key + } + + pub fn key_id(&self) -> KeyId { + self.context + .as_ref() + .map(|context| context.key_id) + .unwrap_or_else(|| { + self.v2 + .as_ref() + .map(|material| material.key_id) + .unwrap_or(ADMIN_KEY_ID) + }) + } + + pub fn context(&self) -> Result<&AuthContext> { + self.context + .as_ref() + .ok_or_else(|| protocol_error("server session was not authenticated")) + } + + pub fn take_context(&mut self) -> Result { + self.context + .take() + .ok_or_else(|| protocol_error("server session was not authenticated")) + } + + pub fn response_writer<'a, T: AsyncWriteExt + Unpin>( + &self, + writer: &'a mut T, + ) -> Result> { + match self.protocol { + HeaderProtocol::Legacy => Ok(HeaderMessageWriter::Legacy(legacy_message_writer( + writer, + &self.legacy_key, + "legacy response writer", + )?)), + HeaderProtocol::V2 => Ok(HeaderMessageWriter::V2(V2MessageWriter::new( + writer, + self.v2_material()?.clone(), + DIRECTION_SERVER_TO_CLIENT, + 0, + )?)), + } + } + + pub fn continuation_reader<'a, T: AsyncReadExt + Unpin>( + &self, + reader: &'a mut T, + ) -> Result> { + match self.protocol { + HeaderProtocol::Legacy => Ok(HeaderMessageReader::Legacy(legacy_message_reader( + reader, + &self.legacy_key, + "legacy reader", + )?)), + HeaderProtocol::V2 => Ok(HeaderMessageReader::V2(V2MessageReader::new( + reader, + self.v2_material()?.clone(), + DIRECTION_CLIENT_TO_SERVER, + 1, + )?)), + } + } +} + +pub struct ServerInitialMessage { + pub payload: Vec, + pub session: ServerHeaderSession, + pub replay_fingerprint: Option<[u8; 32]>, + pub client_timestamp: Option, +} + +pub struct ServerInitialError { + pub failure: AuthFailure, + pub response_session: Option, + pub presented_key_id: Option, +} + +impl ServerInitialError { + fn new(failure: AuthFailure) -> Self { + Self { + failure, + response_session: None, + presented_key_id: None, + } + } + + fn fail(code: &'static str, message: impl Into, retryable: bool) -> Self { + Self::new(AuthFailure::new(code, message, retryable)) + } + + fn fail_key( + code: &'static str, + message: impl Into, + retryable: bool, + key_id: KeyId, + ) -> Self { + Self { + failure: AuthFailure::new(code, message, retryable), + response_session: None, + presented_key_id: Some(key_id), + } + } + + fn from_failure_key(failure: AuthFailure, key_id: KeyId) -> Self { + Self { + failure, + response_session: None, + presented_key_id: Some(key_id), + } + } +} + +impl fmt::Debug for ServerInitialError { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("ServerInitialError") + .field("failure", &self.failure) + .field("has_response_session", &self.response_session.is_some()) + .field("presented_key_id", &self.presented_key_id) + .finish() + } +} + +impl fmt::Display for ServerInitialError { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + self.failure.fmt(formatter) + } +} + +impl std::error::Error for ServerInitialError {} + +use std::fmt; + +#[derive(Clone)] +pub struct ServerSecurity { + auth: AuthRuntime, + replay: Arc>, + failure_logs: Arc>, +} + +// `ServerInitialError` is 256 bytes, but the success type it is paired with, +// `ServerInitialMessage`, is 264 — so the `Result` is already sized by its `Ok` +// variant and boxing the error would buy an allocation for no size reduction. +#[allow(clippy::result_large_err)] +impl ServerSecurity { + pub fn new(auth: AuthRuntime) -> Self { + let replay_path = auth.config().state_dir.join("connection.replay"); + Self { + auth, + replay: Arc::new(Mutex::new(ReplayGuard::open( + Some(replay_path), + DEFAULT_REPLAY_FILTER_BYTES, + DEFAULT_REPLAY_WINDOW_SECONDS, + ))), + failure_logs: Arc::new(Mutex::new(FailureLogLimiter::default())), + } + } + + pub fn auth(&self) -> &AuthRuntime { + &self.auth + } + + pub fn record_failure_log( + &self, + peer_ip: std::net::IpAddr, + key_id: KeyId, + reason: &str, + ) -> FailureLogDecision { + self.failure_logs + .lock() + .record(peer_ip, key_id, reason, unix_seconds()) + } + + pub async fn read_initial( + &self, + reader: &mut T, + ) -> std::result::Result { + let mut first = [0_u8; 4]; + reader.read_exact(&mut first).await.map_err(|error| { + ServerInitialError::new(AuthFailure::new( + "protocol_header_read_failed", + format!("failed to read initial protocol header: {error}"), + true, + )) + })?; + if first == PROTOCOL_V2_MAGIC { + self.read_v2_initial(reader).await + } else { + self.read_legacy_initial(reader, first).await + } + } + + async fn read_legacy_initial( + &self, + reader: &mut T, + checksum_bytes: [u8; 4], + ) -> std::result::Result { + if !self.auth.legacy_protocol_allowed().unwrap_or(false) { + return Err(ServerInitialError::fail( + "legacy_protocol_disabled", + "legacy protocol is disabled by the administrator", + false, + )); + } + let key = self.auth.admin_key().map_err(ServerInitialError::new)?; + let checksum = u32::from_be_bytes(checksum_bytes); + let datalen = reader.read_u32().await.map_err(|error| { + ServerInitialError::fail( + "legacy_frame_invalid", + format!("failed to read legacy frame length: {error}"), + true, + ) + })?; + if !valid_checksum_for_key(datalen, checksum, &key) || datalen > MAX_INITIAL_CIPHERTEXT_LEN + { + return Err(ServerInitialError::fail( + "legacy_frame_invalid", + "legacy frame checksum or length is invalid", + false, + )); + } + let mut encrypted = vec![0_u8; datalen as usize]; + reader.read_exact(&mut encrypted).await.map_err(|error| { + ServerInitialError::fail( + "legacy_frame_invalid", + format!("failed to read legacy frame body: {error}"), + true, + ) + })?; + let mut codec = Aes256GcmDeCodec::try_new(&key).map_err(|_| { + ServerInitialError::fail( + "legacy_decrypt_failed", + "failed to initialize legacy decryption", + false, + ) + })?; + let plain = codec.decrypt(&mut encrypted).map_err(|_| { + ServerInitialError::fail( + "legacy_decrypt_failed", + "legacy credential or encrypted frame is invalid", + false, + ) + })?; + let context = self + .auth + .authenticate_presented(ADMIN_KEY_ID, &key) + .map_err(ServerInitialError::new)?; + let legacy_guard = self + .auth + .record_legacy_connection() + .map_err(ServerInitialError::new)?; + Ok(ServerInitialMessage { + payload: plain.to_vec(), + session: ServerHeaderSession { + protocol: HeaderProtocol::Legacy, + legacy_key: key, + v2: None, + context: Some(context), + _legacy_guard: Some(legacy_guard), + }, + replay_fingerprint: None, + client_timestamp: None, + }) + } + + async fn read_v2_initial( + &self, + reader: &mut T, + ) -> std::result::Result { + let mut remainder = [0_u8; FIRST_PREFIX_REMAINDER_LEN]; + reader.read_exact(&mut remainder).await.map_err(|error| { + ServerInitialError::fail( + "protocol_v2_header_invalid", + format!("failed to read protocol-v2 header: {error}"), + true, + ) + })?; + let version = remainder[0]; + let flags = remainder[1]; + let reserved = u16::from_be_bytes([remainder[2], remainder[3]]); + if version != PROTOCOL_V2_VERSION || flags != 0 || reserved != 0 { + return Err(ServerInitialError::fail( + if version != PROTOCOL_V2_VERSION { + "protocol_version_unsupported" + } else { + "protocol_v2_header_invalid" + }, + format!( + "unsupported protocol header version={version} flags={flags} reserved={reserved}" + ), + false, + )); + } + // The prefix-length check above fixes all three widths, so none of these + // can fail. Reported rather than asserted: this parses the first bytes an + // unauthenticated peer sends, and a panic there is a remote abort. + let malformed = + || ServerInitialError::fail("protocol_error", "v2 prefix is malformed", false); + let key_id = KeyId::from_u64(u64::from_be_bytes( + remainder[4..12].try_into().map_err(|_| malformed())?, + )); + let salt: [u8; CONNECTION_SALT_LEN] = + remainder[12..28].try_into().map_err(|_| malformed())?; + let client_timestamp = u64::from_be_bytes(salt[..8].try_into().map_err(|_| malformed())?); + let now = unix_seconds(); + if now.abs_diff(client_timestamp) > MAX_CONNECTION_CLOCK_SKEW_SECONDS { + return Err(ServerInitialError::fail_key( + "connection_timestamp_invalid", + "protocol-v2 connection timestamp is outside the accepted clock-skew window", + false, + key_id, + )); + } + let key = self + .auth + .derive_key(key_id) + .map_err(|failure| ServerInitialError::from_failure_key(failure, key_id))?; + let material = derive_material(key_id, &key, salt).map_err(|error| { + ServerInitialError::fail_key( + "protocol_v2_key_derivation_failed", + error.to_string(), + false, + key_id, + ) + })?; + let mut session = v2_session(key, material.clone()); + let (counter, ciphertext) = read_v2_frame(reader, 0, MAX_INITIAL_PLAINTEXT_LEN) + .await + .map_err(|error| { + ServerInitialError::fail_key( + "protocol_v2_decrypt_failed", + error.to_string(), + false, + key_id, + ) + })?; + let mut current_ciphertext = ciphertext.clone(); + let fingerprint = replay_fingerprint(key_id, &salt); + let work = match open_v2_payload( + &material, + DIRECTION_CLIENT_TO_SERVER, + counter, + &mut current_ciphertext, + ) { + Ok(payload) => FirstFlightWork::Live { + key, + payload, + error_session: session_without_context(&session), + }, + Err(error) => { + match stale_root_first_flight(&self.auth, key_id, salt, counter, &ciphertext) { + Some(stale) => FirstFlightWork::Stale(stale), + None => { + return Err(first_flight_error( + "protocol_v2_decrypt_failed", + error.to_string(), + false, + key_id, + )); + } + } + } + }; + let replay = self.replay.clone(); + let auth = self.auth.clone(); + let (payload, context) = tokio::task::spawn_blocking(move || { + evaluate_first_flight(&auth, &replay, key_id, fingerprint, work) + }) + .await + .unwrap_or_else(|_| { + Err(first_flight_error( + "connection_replay_store_unavailable", + "failed to evaluate first-flight admission", + true, + key_id, + )) + })?; + session.context = Some(context); + Ok(ServerInitialMessage { + payload, + session, + replay_fingerprint: Some(fingerprint), + client_timestamp: Some(client_timestamp), + }) + } +} + +mod limiter; +pub use limiter::FailureLogDecision; +use limiter::FailureLogLimiter; + +pub enum HeaderMessageReader<'a, T: AsyncReadExt + Unpin> { + Legacy(CodecMessageReader<'a, T, Aes256GcmDeCodec>), + V2(V2MessageReader<'a, T>), +} + +impl MessageReader for HeaderMessageReader<'_, T> { + async fn read_msg(&mut self) -> Result<&'_ [u8]> { + match self { + Self::Legacy(reader) => reader.read_msg().await, + Self::V2(reader) => reader.read_msg().await, + } + } +} + +pub enum HeaderMessageWriter<'a, T: AsyncWriteExt + Unpin> { + Legacy(CodecMessageWriter<'a, T, Aes256GcmEnCodec>), + V2(V2MessageWriter<'a, T>), +} + +impl MessageWriter for HeaderMessageWriter<'_, T> { + async fn write_msg(&mut self, message: &[u8]) -> Result<()> { + match self { + Self::Legacy(writer) => writer.write_msg(message).await, + Self::V2(writer) => writer.write_msg(message).await, + } + } +} + +mod frame; +use frame::{V2Material, derive_material, first_prefix, open_v2_payload, read_v2_frame}; +pub use frame::{V2MessageReader, V2MessageWriter}; +mod replay; +#[cfg(test)] +use replay::RotatingBloom; +use replay::{FirstFlightAdmit, ReplayGuard, replay_fingerprint}; +mod first_flight; +use first_flight::*; +fn legacy_message_reader<'a, T: AsyncReadExt + Unpin>( + reader: &'a mut T, + key: &AesKeyType, + action: &str, +) -> Result> { + Ok(CodecMessageReader::for_session_key( + reader, + Aes256GcmDeCodec::try_new(key) + .map_err(|_| protocol_error(format!("failed to initialize {action}")))?, + *key, + )) +} + +fn legacy_message_writer<'a, T: AsyncWriteExt + Unpin>( + writer: &'a mut T, + key: &AesKeyType, + action: &str, +) -> Result> { + Ok(CodecMessageWriter::for_session_key( + writer, + Aes256GcmEnCodec::try_new(key) + .map_err(|_| protocol_error(format!("failed to initialize {action}")))?, + *key, + )) +} + +fn protocol_error(detail: impl Into) -> Error { + Error::MsgProtocol { + detail: detail.into(), + } +} + +fn unix_seconds() -> u64 { + SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_secs() +} + +#[cfg(test)] +mod tests; diff --git a/crates/pb-mapper-protocol/src/secure/first_flight.rs b/crates/pb-mapper-protocol/src/secure/first_flight.rs new file mode 100644 index 0000000..fb9183c --- /dev/null +++ b/crates/pb-mapper-protocol/src/secure/first_flight.rs @@ -0,0 +1,144 @@ +//! First-flight authentication, replay admission, and stale-root classification. +use super::*; + +pub(super) enum FirstFlightWork { + Live { + key: AesKeyType, + payload: Vec, + error_session: ServerHeaderSession, + }, + Stale(ServerInitialError), +} + +pub(super) fn first_flight_error( + code: &'static str, + message: impl Into, + retryable: bool, + key_id: KeyId, +) -> ServerInitialError { + ServerInitialError::fail_key(code, message, retryable, key_id) +} + +fn reserved_error_session( + replay: &mut ReplayGuard, + fingerprint: &[u8; 32], + session: ServerHeaderSession, +) -> Option { + matches!( + replay.claim(fingerprint, unix_seconds()), + FirstFlightAdmit::Fresh + ) + .then_some(session) +} + +#[allow(clippy::result_large_err)] +pub(super) fn evaluate_first_flight( + auth: &AuthRuntime, + replay: &parking_lot::Mutex, + key_id: KeyId, + fingerprint: [u8; 32], + work: FirstFlightWork, +) -> std::result::Result<(Vec, AuthContext), ServerInitialError> { + let mut replay = replay.lock(); + match work { + FirstFlightWork::Live { + key, + payload, + error_session, + } => { + let context = auth + .authenticate_presented(key_id, &key) + .map_err(|failure| ServerInitialError { + failure, + response_session: reserved_error_session( + &mut replay, + &fingerprint, + error_session, + ), + presented_key_id: Some(key_id), + })?; + match replay.admit(key_id, &fingerprint, unix_seconds()) { + FirstFlightAdmit::Fresh => Ok((payload, context)), + FirstFlightAdmit::Replayed => Err(first_flight_error( + "connection_salt_replayed", + "protocol-v2 connection salt was already accepted", + true, + key_id, + )), + FirstFlightAdmit::Limited => Err(first_flight_error( + "connection_admission_limited", + "this credential has opened too many new connections in the current window", + true, + key_id, + )), + FirstFlightAdmit::Unavailable => Err(first_flight_error( + "connection_replay_store_unavailable", + "failed to persist first-flight replay admission", + true, + key_id, + )), + } + } + FirstFlightWork::Stale(mut error) => { + if let Some(session) = error.response_session.take() { + error.response_session = reserved_error_session(&mut replay, &fingerprint, session); + } + Err(error) + } + } +} + +pub(super) fn stale_root_first_flight( + auth: &AuthRuntime, + key_id: KeyId, + salt: [u8; CONNECTION_SALT_LEN], + counter: u64, + ciphertext: &[u8], +) -> Option { + let previous_key = auth.derive_previous_key(key_id)?; + let previous_material = derive_material(key_id, &previous_key, salt).ok()?; + let mut previous_ciphertext = ciphertext.to_vec(); + open_v2_payload( + &previous_material, + DIRECTION_CLIENT_TO_SERVER, + counter, + &mut previous_ciphertext, + ) + .ok()?; + let (code, message) = if key_id.is_admin() { + ( + "administrator_key_invalid", + "administrator credential does not match the active root key", + ) + } else { + ( + "temporary_key_rotated", + "temporary credential was invalidated by administrator root rotation or auth-state reset", + ) + }; + Some(ServerInitialError { + failure: AuthFailure::new(code, message, false), + response_session: Some(v2_session(previous_key, previous_material)), + presented_key_id: Some(key_id), + }) +} + +pub(super) fn v2_session(key: AesKeyType, material: V2Material) -> ServerHeaderSession { + ServerHeaderSession { + protocol: HeaderProtocol::V2, + legacy_key: key, + v2: Some(material), + context: None, + _legacy_guard: None, + } +} + +pub(super) fn session_without_context(session: &ServerHeaderSession) -> ServerHeaderSession { + ServerHeaderSession { + protocol: session.protocol, + legacy_key: session.legacy_key, + v2: session.v2.clone(), + context: None, + _legacy_guard: None, + } +} diff --git a/crates/pb-mapper-protocol/src/secure/frame.rs b/crates/pb-mapper-protocol/src/secure/frame.rs new file mode 100644 index 0000000..09a3c8f --- /dev/null +++ b/crates/pb-mapper-protocol/src/secure/frame.rs @@ -0,0 +1,280 @@ +//! Protocol-v2 key derivation and directional authenticated frame codecs. +//! +//! ```text +//! credential + connection salt -> HKDF -> c2s key / s2c key +//! plaintext -> counter + length + AEAD(AAD) -> encrypted frame +//! encrypted frame -> bound length -> verify counter/tag -> plaintext +//! ``` +//! +//! Counters are monotonic per direction and are included in both the nonce and AAD. +//! The initial reader can impose a smaller pre-authentication limit before allocating +//! a body; continuation frames retain the normal protocol maximum. + +use pb_mapper_auth::KeyId; + +use super::*; + +#[derive(Clone)] +pub(super) struct V2Material { + pub(super) key_id: KeyId, + pub(super) flags: u8, + pub(super) salt: [u8; CONNECTION_SALT_LEN], + pub(super) client_to_server: AesKeyType, + pub(super) server_to_client: AesKeyType, +} + +pub struct V2MessageReader<'a, T: AsyncReadExt + Unpin> { + reader: &'a mut T, + material: V2Material, + key: LessSafeKey, + direction: u8, + expected_counter: u64, + buffer: Vec, +} + +impl<'a, T: AsyncReadExt + Unpin> V2MessageReader<'a, T> { + pub(super) fn new( + reader: &'a mut T, + material: V2Material, + direction: u8, + expected_counter: u64, + ) -> Result { + let key_bytes = direction_key(&material, direction); + let key = LessSafeKey::new( + UnboundKey::new(&AES_256_GCM, key_bytes) + .map_err(|_| protocol_error("invalid protocol-v2 read key"))?, + ); + Ok(Self { + reader, + material, + key, + direction, + expected_counter, + buffer: Vec::new(), + }) + } + + pub(super) async fn read_msg_with_limit(&mut self, max_plaintext_len: u32) -> Result<&'_ [u8]> { + let (counter, ciphertext) = + read_v2_frame(self.reader, self.expected_counter, max_plaintext_len).await?; + self.buffer = ciphertext; + let datalen = u32::try_from(self.buffer.len()) + .map_err(|_| protocol_error("protocol-v2 payload exceeds u32 length"))?; + let aad = frame_aad(&self.material, self.direction, counter, datalen); + let plain = self + .key + .open_in_place(nonce(counter), Aad::from(aad.as_slice()), &mut self.buffer) + .map_err(|_| protocol_error("protocol-v2 payload authentication failed"))?; + let plain_len = plain.len(); + self.buffer.truncate(plain_len); + self.expected_counter = self + .expected_counter + .checked_add(1) + .ok_or_else(|| protocol_error("protocol-v2 receive counter exhausted"))?; + Ok(&self.buffer) + } +} + +pub(super) async fn read_v2_frame( + reader: &mut T, + expected_counter: u64, + max_plaintext_len: u32, +) -> Result<(u64, Vec)> { + let counter = reader + .read_u64() + .await + .map_err(|error| protocol_error(format!("failed to read v2 counter: {error}")))?; + if counter != expected_counter { + return Err(protocol_error(format!( + "protocol-v2 counter mismatch: expected {expected_counter}, got {counter}" + ))); + } + let datalen = reader + .read_u32() + .await + .map_err(|error| protocol_error(format!("failed to read v2 length: {error}")))?; + let max_encrypted_len = max_plaintext_len.saturating_add(AES_256_GCM.tag_len() as u32); + if datalen < AES_256_GCM.tag_len() as u32 || datalen > max_encrypted_len { + return Err(protocol_error(format!( + "protocol-v2 payload length {datalen} exceeds the {max_plaintext_len}-byte limit" + ))); + } + let mut ciphertext = vec![0_u8; datalen as usize]; + reader + .read_exact(&mut ciphertext) + .await + .map_err(|error| protocol_error(format!("failed to read v2 payload: {error}")))?; + Ok((counter, ciphertext)) +} + +impl MessageReader for V2MessageReader<'_, T> { + async fn read_msg(&mut self) -> Result<&'_ [u8]> { + self.read_msg_with_limit(MAX_MSG_LEN - AES_256_GCM.tag_len() as u32) + .await + } +} + +pub struct V2MessageWriter<'a, T: AsyncWriteExt + Unpin> { + writer: &'a mut T, + material: V2Material, + key: LessSafeKey, + direction: u8, + counter: u64, +} + +impl<'a, T: AsyncWriteExt + Unpin> V2MessageWriter<'a, T> { + pub(super) fn new( + writer: &'a mut T, + material: V2Material, + direction: u8, + counter: u64, + ) -> Result { + let key_bytes = direction_key(&material, direction); + let key = LessSafeKey::new( + UnboundKey::new(&AES_256_GCM, key_bytes) + .map_err(|_| protocol_error("invalid protocol-v2 write key"))?, + ); + Ok(Self { + writer, + material, + key, + direction, + counter, + }) + } +} + +impl MessageWriter for V2MessageWriter<'_, T> { + async fn write_msg(&mut self, message: &[u8]) -> Result<()> { + let encrypted_len = message + .len() + .checked_add(AES_256_GCM.tag_len()) + .and_then(|len| DataLenType::try_from(len).ok()) + .ok_or_else(|| protocol_error("protocol-v2 message is too large"))?; + if encrypted_len > MAX_MSG_LEN { + return Err(protocol_error( + "protocol-v2 message exceeds the maximum length", + )); + } + let counter = self.counter; + let aad = frame_aad(&self.material, self.direction, counter, encrypted_len); + let mut encrypted = message.to_vec(); + self.key + .seal_in_place_append_tag(nonce(counter), Aad::from(aad.as_slice()), &mut encrypted) + .map_err(|_| protocol_error("failed to encrypt protocol-v2 message"))?; + self.writer + .write_u64(counter) + .await + .map_err(|error| protocol_error(format!("failed to write v2 frame header: {error}")))?; + self.writer + .write_u32(encrypted_len) + .await + .map_err(|error| protocol_error(format!("failed to write v2 frame header: {error}")))?; + self.writer + .write_all(&encrypted) + .await + .map_err(|error| protocol_error(format!("failed to write v2 frame body: {error}")))?; + self.counter = self + .counter + .checked_add(1) + .ok_or_else(|| protocol_error("protocol-v2 send counter exhausted"))?; + Ok(()) + } +} + +pub(super) fn open_v2_payload( + material: &V2Material, + direction: u8, + counter: u64, + ciphertext: &mut [u8], +) -> Result> { + let key_bytes = direction_key(material, direction); + let key = LessSafeKey::new( + UnboundKey::new(&AES_256_GCM, key_bytes) + .map_err(|_| protocol_error("invalid protocol-v2 read key"))?, + ); + let datalen = u32::try_from(ciphertext.len()) + .map_err(|_| protocol_error("protocol-v2 payload length is invalid"))?; + let aad = frame_aad(material, direction, counter, datalen); + let plain = key + .open_in_place(nonce(counter), Aad::from(aad.as_slice()), ciphertext) + .map_err(|_| protocol_error("protocol-v2 payload authentication failed"))?; + Ok(plain.to_vec()) +} + +pub(super) fn derive_material( + key_id: KeyId, + credential_key: &AesKeyType, + salt_bytes: [u8; CONNECTION_SALT_LEN], +) -> Result { + let salt = Salt::new(HKDF_SHA256, &salt_bytes); + let pseudo_random_key = salt.extract(credential_key); + let client_to_server = expand_direction(&pseudo_random_key, b"pb-mapper-v2-c2s")?; + let server_to_client = expand_direction(&pseudo_random_key, b"pb-mapper-v2-s2c")?; + Ok(V2Material { + key_id, + flags: 0, + salt: salt_bytes, + client_to_server, + server_to_client, + }) +} + +fn expand_direction( + pseudo_random_key: &ring::hkdf::Prk, + label: &'static [u8], +) -> Result { + let info = [label]; + let output = pseudo_random_key + .expand(&info, HkdfLen(32)) + .map_err(|_| protocol_error("failed to derive protocol-v2 direction key"))?; + let mut key = [0_u8; 32]; + output + .fill(&mut key) + .map_err(|_| protocol_error("failed to fill protocol-v2 direction key"))?; + Ok(key) +} + +struct HkdfLen(usize); + +impl ring::hkdf::KeyType for HkdfLen { + fn len(&self) -> usize { + self.0 + } +} + +fn direction_key(material: &V2Material, direction: u8) -> &AesKeyType { + if direction == DIRECTION_CLIENT_TO_SERVER { + &material.client_to_server + } else { + &material.server_to_client + } +} + +pub(super) fn first_prefix(material: &V2Material) -> Vec { + let mut prefix = Vec::with_capacity(PROTOCOL_V2_MAGIC.len() + FIRST_PREFIX_REMAINDER_LEN); + prefix.extend_from_slice(&PROTOCOL_V2_MAGIC); + prefix.push(PROTOCOL_V2_VERSION); + prefix.push(material.flags); + prefix.extend_from_slice(&0_u16.to_be_bytes()); + prefix.extend_from_slice(&material.key_id.to_be_bytes()); + prefix.extend_from_slice(&material.salt); + prefix +} + +fn frame_aad(material: &V2Material, direction: u8, counter: u64, datalen: u32) -> Vec { + let mut aad = Vec::with_capacity( + PROTOCOL_V2_MAGIC.len() + FIRST_PREFIX_REMAINDER_LEN + 1 + FRAME_HEADER_LEN, + ); + aad.extend_from_slice(&first_prefix(material)); + aad.push(direction); + aad.extend_from_slice(&counter.to_be_bytes()); + aad.extend_from_slice(&datalen.to_be_bytes()); + aad +} + +fn nonce(counter: u64) -> Nonce { + let mut bytes = [0_u8; 12]; + bytes[4..].copy_from_slice(&counter.to_be_bytes()); + Nonce::assume_unique_for_key(bytes) +} diff --git a/crates/pb-mapper-protocol/src/secure/limiter.rs b/crates/pb-mapper-protocol/src/secure/limiter.rs new file mode 100644 index 0000000..e67a7a6 --- /dev/null +++ b/crates/pb-mapper-protocol/src/secure/limiter.rs @@ -0,0 +1,89 @@ +//! Cardinality-bounded suppression for repeated authentication failure logs. +//! +//! ```text +//! (peer IP, key id, reason) -> per-window counter -> emit first / suppress repeats +//! too many distinct keys ---------> shared overflow bucket +//! ``` +//! +//! This limits log amplification from the public relay port without changing protocol +//! decisions: every authentication failure is still rejected, only duplicate logging +//! is coalesced. + +use pb_mapper_auth::KeyId; + +#[derive(Clone, Copy, Debug)] +pub struct FailureLogDecision { + pub emit: bool, + pub suppressed: u64, +} + +pub(super) struct FailureLogEntry { + window_started_at: u64, + emitted: u8, + suppressed: u64, +} + +#[derive(Default)] +pub(super) struct FailureLogLimiter { + pub(super) entries: + std::collections::HashMap<(std::net::IpAddr, KeyId, String), FailureLogEntry>, + pub(super) overflow: Option, +} + +impl FailureLogLimiter { + pub(super) fn record( + &mut self, + peer_ip: std::net::IpAddr, + key_id: KeyId, + reason: &str, + now: u64, + ) -> FailureLogDecision { + let key = (peer_ip, key_id, reason.to_string()); + if !self.entries.contains_key(&key) && self.entries.len() >= 4096 { + self.entries + .retain(|_, entry| now.saturating_sub(entry.window_started_at) < 120); + if self.entries.len() >= 4096 { + let entry = self.overflow.get_or_insert(FailureLogEntry { + window_started_at: now, + emitted: 0, + suppressed: 0, + }); + return record_failure_entry(entry, now); + } + } + let entry = self.entries.entry(key).or_insert(FailureLogEntry { + window_started_at: now, + emitted: 0, + suppressed: 0, + }); + record_failure_entry(entry, now) + } +} + +fn record_failure_entry(entry: &mut FailureLogEntry, now: u64) -> FailureLogDecision { + if now.saturating_sub(entry.window_started_at) >= 60 { + let suppressed = entry.suppressed; + *entry = FailureLogEntry { + window_started_at: now, + emitted: 1, + suppressed: 0, + }; + return FailureLogDecision { + emit: true, + suppressed, + }; + } + if entry.emitted < 5 { + entry.emitted += 1; + FailureLogDecision { + emit: true, + suppressed: 0, + } + } else { + entry.suppressed = entry.suppressed.saturating_add(1); + FailureLogDecision { + emit: false, + suppressed: 0, + } + } +} diff --git a/crates/pb-mapper-protocol/src/secure/replay.rs b/crates/pb-mapper-protocol/src/secure/replay.rs new file mode 100644 index 0000000..8d240c0 --- /dev/null +++ b/crates/pb-mapper-protocol/src/secure/replay.rs @@ -0,0 +1,423 @@ +//! Fast process-local duplicate admission guard for protocol-v2 first flights. +//! +//! ```text +//! key id + salt -> SHA-256 fingerprint -> current Bloom window +//! -> previous Bloom window +//! key id ---------> per-credential first-flight count (before Bloom insert) +//! ``` +//! +//! `admit` is called while one mutex is held, making concurrent admission +//! atomic. This Bloom filter protects all connection types from immediate duplicates; +//! administrator mutations additionally use the exact durable replay set in `auth`. +//! +//! Each generation lasts `2 *` the accepted clock-skew so a salt inserted at the +//! end of a window with a max-future timestamp cannot be replayed after rotation. +//! Per-credential counts stop one tenant from filling the shared filter with +//! unique salts before the request payload is decoded. +//! +//! The `expect`s here are all slices of a `[u8; REPLAY_RECORD_LEN]` at constant +//! offsets, or SHA-256's 32-byte output — widths the array types already fix, so +//! the conversions cannot fail. Unlike the parsing paths, these sit in functions +//! that return `()`, so there is nothing to report a width error to. +#![allow(clippy::expect_used)] + +use pb_mapper_auth::KeyId; + +use std::collections::HashMap; +use std::fs::{File, OpenOptions}; +use std::io::{ErrorKind, Read, Write}; +use std::path::PathBuf; + +use rand::RngExt; + +use super::*; + +const DEFAULT_NEW_STREAMS_PER_SECOND: u32 = 100; +const REPLAY_RECORD_LEN: usize = 40; +const REPLAY_COMPACT_INTERVAL_SECONDS: u64 = 60; + +fn first_flight_budget(window_seconds: u64) -> u32 { + let streams_per_sec = std::env::var("PB_MAPPER_NEW_STREAMS_PER_SECOND") + .ok() + .and_then(|value| value.parse().ok()) + .filter(|value| *value > 0) + .unwrap_or(DEFAULT_NEW_STREAMS_PER_SECOND) + .min(1_000_000); + let window = u32::try_from(window_seconds).unwrap_or(u32::MAX); + streams_per_sec + .saturating_mul(2) + .saturating_mul(window) + .saturating_mul(2) + .max(8_192) +} + +fn bloom_insert_capacity(bytes: usize) -> u32 { + let bits = (bytes as u64).saturating_mul(8); + u32::try_from(bits / 16).unwrap_or(u32::MAX).max(1) +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(super) enum FirstFlightAdmit { + Fresh, + Replayed, + Limited, + Unavailable, +} + +pub(super) fn replay_fingerprint(key_id: KeyId, salt: &[u8; CONNECTION_SALT_LEN]) -> [u8; 32] { + let mut input = [0_u8; 8 + CONNECTION_SALT_LEN]; + input[..8].copy_from_slice(&key_id.to_be_bytes()); + input[8..].copy_from_slice(salt); + digest(&SHA256, &input) + .as_ref() + .try_into() + .expect("SHA-256 width") +} + +pub(super) struct RotatingBloom { + current: Vec, + previous: Vec, + pub(super) current_started_at: u64, + window_seconds: u64, +} + +impl RotatingBloom { + pub(super) fn new(bytes: usize, window_seconds: u64) -> Self { + Self { + current: vec![0; bytes], + previous: vec![0; bytes], + current_started_at: unix_seconds(), + window_seconds, + } + } + + pub(super) fn contains(&mut self, fingerprint: &[u8; 32], now: u64) -> bool { + self.rotate(now); + bloom_contains(&self.current, fingerprint) || bloom_contains(&self.previous, fingerprint) + } + + pub(super) fn insert(&mut self, fingerprint: &[u8; 32], now: u64) { + self.rotate(now); + bloom_insert(&mut self.current, fingerprint); + } + + fn rotate(&mut self, now: u64) { + let elapsed = now.saturating_sub(self.current_started_at); + if elapsed < self.window_seconds { + return; + } + if elapsed >= self.window_seconds.saturating_mul(2) { + self.current.fill(0); + self.previous.fill(0); + } else { + std::mem::swap(&mut self.current, &mut self.previous); + self.current.fill(0); + } + self.current_started_at = now; + } +} + +pub(super) struct ReplayGuard { + bloom: RotatingBloom, + counts: HashMap, + counts_started_at: u64, + window_seconds: u64, + max_per_key: u32, + total: u32, + total_started_at: u64, + max_total: u32, + log_path: Option, + last_compact_at: u64, + log_failed: bool, +} + +impl ReplayGuard { + pub(super) fn open(log_path: Option, bytes: usize, window_seconds: u64) -> Self { + let now = unix_seconds(); + let configured = first_flight_budget(window_seconds); + let mut guard = Self { + bloom: RotatingBloom::new(bytes, window_seconds), + counts: HashMap::new(), + counts_started_at: now, + window_seconds, + max_per_key: configured, + total: 0, + total_started_at: now, + max_total: configured.min(bloom_insert_capacity(bytes)), + log_path, + last_compact_at: now, + log_failed: false, + }; + guard.load_persisted(); + guard + } + + #[cfg(test)] + pub(super) fn with_max_per_key(mut self, max_per_key: u32) -> Self { + self.max_per_key = max_per_key; + self + } + + #[cfg(test)] + pub(super) fn with_max_total(mut self, max_total: u32) -> Self { + self.max_total = max_total; + self + } + + pub(super) fn admit( + &mut self, + key_id: KeyId, + fingerprint: &[u8; 32], + now: u64, + ) -> FirstFlightAdmit { + self.rotate_counts(now); + if self.bloom.contains(fingerprint, now) { + return FirstFlightAdmit::Replayed; + } + self.reset_total_if_rotated(); + if self.counts.get(&key_id).copied().unwrap_or(0) >= self.max_per_key + || self.total >= self.max_total + { + return FirstFlightAdmit::Limited; + } + let result = self.reserve(fingerprint, now); + if result == FirstFlightAdmit::Fresh { + *self.counts.entry(key_id).or_insert(0) += 1; + if now.saturating_sub(self.last_compact_at) >= REPLAY_COMPACT_INTERVAL_SECONDS { + self.compact(now); + } + } + result + } + + pub(super) fn claim(&mut self, fingerprint: &[u8; 32], now: u64) -> FirstFlightAdmit { + if self.bloom.contains(fingerprint, now) { + return FirstFlightAdmit::Replayed; + } + self.reset_total_if_rotated(); + if self.total >= self.max_total { + return FirstFlightAdmit::Unavailable; + } + self.reserve(fingerprint, now) + } + + fn reset_total_if_rotated(&mut self) { + if self.bloom.current_started_at != self.total_started_at { + self.total = 0; + self.total_started_at = self.bloom.current_started_at; + } + } + + fn reserve(&mut self, fingerprint: &[u8; 32], now: u64) -> FirstFlightAdmit { + if self.persist(fingerprint, now).is_err() { + return FirstFlightAdmit::Unavailable; + } + self.bloom.insert(fingerprint, now); + self.total = self.total.saturating_add(1); + FirstFlightAdmit::Fresh + } + + fn rotate_counts(&mut self, now: u64) { + if now.saturating_sub(self.counts_started_at) < self.window_seconds { + return; + } + self.counts.clear(); + self.counts_started_at = now; + } + + fn persist(&mut self, fingerprint: &[u8; 32], now: u64) -> std::io::Result<()> { + if self.log_failed { + return Err(std::io::Error::other( + "durable first-flight replay log is unavailable after a failed rollback", + )); + } + let Some(path) = &self.log_path else { + return Ok(()); + }; + if let Some(parent) = path.parent() { + std::fs::create_dir_all(parent)?; + } + let created = !path.exists(); + let mut file = OpenOptions::new().create(true).append(true).open(path)?; + let start_len = file.metadata()?.len(); + let mut record = [0_u8; REPLAY_RECORD_LEN]; + record[..32].copy_from_slice(fingerprint); + record[32..].copy_from_slice(&now.to_be_bytes()); + if let Err(error) = file.write_all(&record).and_then(|()| file.sync_data()) { + if file + .set_len(start_len) + .and_then(|()| file.sync_data()) + .is_err() + { + self.log_failed = true; + } + return Err(error); + } + if created && let Err(error) = pb_mapper_core::durable_file::sync_parent_directory(path) { + self.log_failed = true; + return Err(error); + } + Ok(()) + } + + fn load_persisted(&mut self) { + let Some(path) = self.log_path.clone() else { + return; + }; + if !path.exists() { + return; + } + let mut file = match File::open(&path) { + Ok(file) => file, + Err(_) => { + self.log_failed = true; + return; + } + }; + let now = unix_seconds(); + let records = match read_complete_replay_records(&mut file) { + Ok(records) => records, + Err(_) => { + self.log_failed = true; + return; + } + }; + let mut live = Vec::new(); + for record in records { + let timestamp = u64::from_be_bytes(record[32..].try_into().expect("timestamp width")); + if now.saturating_sub(timestamp) > self.window_seconds { + continue; + } + let fingerprint: [u8; 32] = record[..32].try_into().expect("fingerprint width"); + self.bloom.insert(&fingerprint, timestamp); + live.push(record); + } + if let Some(oldest) = live + .iter() + .map(|record| u64::from_be_bytes(record[32..].try_into().expect("timestamp width"))) + .min() + { + self.bloom.current_started_at = oldest; + } + self.total = u32::try_from(live.len()) + .unwrap_or(u32::MAX) + .min(self.max_total); + self.total_started_at = self.bloom.current_started_at; + if self.rewrite_live(&live).is_err() { + self.log_failed = true; + return; + } + self.last_compact_at = now; + } + + fn compact(&mut self, now: u64) { + let Some(path) = &self.log_path else { + self.last_compact_at = now; + return; + }; + let mut file = match File::open(path) { + Ok(file) => file, + Err(_) => { + self.log_failed = true; + return; + } + }; + let records = match read_complete_replay_records(&mut file) { + Ok(records) => records, + Err(_) => { + self.log_failed = true; + return; + } + }; + let live = records + .into_iter() + .filter(|record| { + let timestamp = + u64::from_be_bytes(record[32..].try_into().expect("timestamp width")); + now.saturating_sub(timestamp) <= self.window_seconds + }) + .collect::>(); + if self.rewrite_live(&live).is_err() { + self.log_failed = true; + return; + } + self.last_compact_at = now; + } + + fn rewrite_live(&self, live: &[[u8; REPLAY_RECORD_LEN]]) -> std::io::Result<()> { + let Some(path) = &self.log_path else { + return Ok(()); + }; + let mut random_suffix = [0_u8; 8]; + let mut rng = rand::rng(); + for byte in &mut random_suffix { + *byte = rng.random(); + } + let temporary = path.with_file_name(format!( + ".{}.tmp-{}-{:016x}", + path.file_name() + .and_then(|name| name.to_str()) + .unwrap_or("connection.replay"), + std::process::id(), + u64::from_be_bytes(random_suffix) + )); + let result = (|| { + let mut file = OpenOptions::new() + .create_new(true) + .write(true) + .open(&temporary)?; + file.write_all(&live.concat())?; + file.sync_all()?; + drop(file); + pb_mapper_core::durable_file::replace_file(&temporary, path)?; + pb_mapper_core::durable_file::sync_parent_directory(path) + })(); + if result.is_err() { + let _ = std::fs::remove_file(&temporary); + } + result + } +} + +fn read_complete_replay_records(file: &mut File) -> std::io::Result> { + let mut records = Vec::new(); + loop { + let mut record = [0_u8; REPLAY_RECORD_LEN]; + let read = file.read(&mut record)?; + if read == 0 { + return Ok(records); + } + if read < REPLAY_RECORD_LEN { + return Err(std::io::Error::new( + ErrorKind::UnexpectedEof, + "truncated first-flight replay record", + )); + } + records.push(record); + } +} + +fn bloom_positions(filter_len: usize, fingerprint: &[u8; 32]) -> [usize; 4] { + let bits = filter_len * 8; + std::array::from_fn(|index| { + let offset = index * 8; + let hash = u64::from_be_bytes( + fingerprint[offset..offset + 8] + .try_into() + .expect("fingerprint chunk"), + ); + hash as usize % bits + }) +} + +fn bloom_contains(filter: &[u8], fingerprint: &[u8; 32]) -> bool { + bloom_positions(filter.len(), fingerprint) + .into_iter() + .all(|position| filter[position / 8] & (1 << (position % 8)) != 0) +} + +fn bloom_insert(filter: &mut [u8], fingerprint: &[u8; 32]) { + for position in bloom_positions(filter.len(), fingerprint) { + filter[position / 8] |= 1 << (position % 8); + } +} diff --git a/crates/pb-mapper-protocol/src/secure/tests.rs b/crates/pb-mapper-protocol/src/secure/tests.rs new file mode 100644 index 0000000..a1f982c --- /dev/null +++ b/crates/pb-mapper-protocol/src/secure/tests.rs @@ -0,0 +1,775 @@ +//! End-to-end protocol-v2 framing and admission invariants. +//! +//! ```text +//! client session -> duplex transport -> ServerSecurity -> authenticated context +//! captured frame -- concurrent replay -----------------> exactly one admission +//! oversized first header ------------------------------> reject before body read +//! ``` +//! +//! These tests intentionally exercise both administrator and derived temporary +//! credentials, while lifecycle persistence remains covered by `common::auth::tests`. + +use super::*; +use pb_mapper_auth::{AuthConfig, LegacyProtocolPolicy}; +use pb_mapper_core::checksum::{ + encode_temporary_credential, parse_credential, set_process_msg_header_key, +}; +use pb_mapper_core::test_support::PROCESS_CREDENTIAL_TEST_LOCK; + +fn temp_config() -> AuthConfig { + let mut random = [0_u8; 8]; + let mut rng = rand::rng(); + for byte in &mut random { + *byte = rng.random(); + } + AuthConfig { + state_dir: std::env::temp_dir() + .join(format!("pb-mapper-v2-{}", u64::from_be_bytes(random))), + max_temporary_keys: 8, + max_temporary_key_ttl: std::time::Duration::from_secs(3600), + legacy_protocol: LegacyProtocolPolicy::Allow, + } +} + +#[tokio::test] +async fn v2_round_trip_uses_directional_counters() { + let credential = Credential::Admin(*b"0123456789abcdefghijklmnopqrstuv"); + let client = ClientHeaderSession::new_v2(&credential).unwrap(); + let config = temp_config(); + let auth = AuthRuntime::start(*credential.key(), config.clone()) + .await + .unwrap(); + let security = ServerSecurity::new(auth); + let (mut client_io, mut server_io) = tokio::io::duplex(4096); + + let client_task = async { + client + .write_initial(&mut client_io, b"request") + .await + .unwrap(); + let mut reader = client.response_reader(&mut client_io).unwrap(); + assert_eq!(reader.read_msg().await.unwrap(), b"response"); + }; + let server_task = async { + let initial = security.read_initial(&mut server_io).await.unwrap(); + assert_eq!(initial.payload, b"request"); + let mut writer = initial.session.response_writer(&mut server_io).unwrap(); + writer.write_msg(b"response").await.unwrap(); + }; + tokio::join!(client_task, server_task); + let _ = std::fs::remove_dir_all(config.state_dir); +} + +#[tokio::test] +async fn temporary_credential_authenticates_without_storing_secret() { + let admin = *b"0123456789abcdefghijklmnopqrstuv"; + let config = temp_config(); + let auth = AuthRuntime::start(admin, config.clone()).await.unwrap(); + let admin_context = auth.authenticate_presented(ADMIN_KEY_ID, &admin).unwrap(); + let issued = auth + .issue(&admin_context, std::time::Duration::from_secs(60), None) + .await + .unwrap(); + let Credential::Temporary { key_id, key } = + pb_mapper_core::checksum::parse_credential(&issued.credential).unwrap() + else { + panic!("expected temporary credential") + }; + assert_eq!(issued.credential, encode_temporary_credential(key_id, &key)); + let client = ClientHeaderSession::new_v2(&Credential::Temporary { key_id, key }).unwrap(); + let security = ServerSecurity::new(auth); + let (mut client_io, mut server_io) = tokio::io::duplex(4096); + let client_task = client.write_initial(&mut client_io, b"temporary"); + let server_task = security.read_initial(&mut server_io); + let (client_result, server_result) = tokio::join!(client_task, server_task); + client_result.unwrap(); + let initial = server_result.unwrap(); + assert_eq!(initial.payload, b"temporary"); + assert_eq!(initial.session.context().unwrap().namespace, key_id); + let _ = std::fs::remove_dir_all(config.state_dir); +} + +async fn encode_initial(session: &ClientHeaderSession, payload: &[u8]) -> Vec { + let (mut writer, mut reader) = tokio::io::duplex(128 * 1024); + let write = async { + session.write_initial(&mut writer, payload).await.unwrap(); + drop(writer); + }; + let read = async { + let mut bytes = Vec::new(); + reader.read_to_end(&mut bytes).await.unwrap(); + bytes + }; + let (_, bytes) = tokio::join!(write, read); + bytes +} + +#[tokio::test] +async fn identical_initial_frames_are_admitted_only_once() { + let credential = Credential::Admin(*b"0123456789abcdefghijklmnopqrstuv"); + let session = ClientHeaderSession::new_v2(&credential).unwrap(); + let bytes = encode_initial(&session, b"same-request").await; + let config = temp_config(); + let auth = AuthRuntime::start(*credential.key(), config.clone()) + .await + .unwrap(); + let security = ServerSecurity::new(auth); + let mut first = std::io::Cursor::new(bytes.clone()); + let mut second = std::io::Cursor::new(bytes); + + let (first, second) = tokio::join!( + security.read_initial(&mut first), + security.read_initial(&mut second) + ); + let results = [first, second]; + assert_eq!(results.iter().filter(|result| result.is_ok()).count(), 1); + assert_eq!( + results + .iter() + .filter_map(|result| result.as_ref().err()) + .next() + .unwrap() + .failure + .code, + "connection_salt_replayed" + ); + assert!( + results + .iter() + .filter_map(|result| result.as_ref().err()) + .next() + .unwrap() + .response_session + .is_none() + ); + + let _ = std::fs::remove_dir_all(config.state_dir); +} + +#[tokio::test] +async fn revoked_first_flights_do_not_consume_the_replay_filter() { + let admin = *b"0123456789abcdefghijklmnopqrstuv"; + let config = temp_config(); + let auth = AuthRuntime::start(admin, config.clone()).await.unwrap(); + let admin_context = auth.authenticate_presented(ADMIN_KEY_ID, &admin).unwrap(); + let issued = auth + .issue(&admin_context, std::time::Duration::from_secs(60), None) + .await + .unwrap(); + let Credential::Temporary { key_id, key } = parse_credential(&issued.credential).unwrap() + else { + panic!("expected temporary credential"); + }; + let client = ClientHeaderSession::new_v2(&Credential::Temporary { key_id, key }).unwrap(); + let bytes = encode_initial(&client, b"revoked").await; + auth.revoke(&admin_context, KeyId::from_u64(key_id)) + .await + .unwrap(); + let security = ServerSecurity::new(auth); + + let first = match security + .read_initial(&mut std::io::Cursor::new(bytes.clone())) + .await + { + Ok(_) => panic!("revoked credential should fail"), + Err(error) => error, + }; + let second = match security + .read_initial(&mut std::io::Cursor::new(bytes)) + .await + { + Ok(_) => panic!("revoked credential should fail again"), + Err(error) => error, + }; + assert_eq!(first.failure.code, "temporary_key_revoked"); + assert!(first.response_session.is_some()); + assert_eq!(second.failure.code, "temporary_key_revoked"); + assert!( + second.response_session.is_none(), + "a claimed revoked salt must not reuse nonce 0" + ); + + let _ = std::fs::remove_dir_all(config.state_dir); +} + +#[tokio::test] +async fn accepted_then_revoked_replay_does_not_reuse_the_session_nonce() { + let admin = *b"0123456789abcdefghijklmnopqrstuv"; + let config = temp_config(); + let auth = AuthRuntime::start(admin, config.clone()).await.unwrap(); + let admin_context = auth.authenticate_presented(ADMIN_KEY_ID, &admin).unwrap(); + let issued = auth + .issue(&admin_context, std::time::Duration::from_secs(60), None) + .await + .unwrap(); + let Credential::Temporary { key_id, key } = parse_credential(&issued.credential).unwrap() + else { + panic!("expected temporary credential"); + }; + let client = ClientHeaderSession::new_v2(&Credential::Temporary { key_id, key }).unwrap(); + let bytes = encode_initial(&client, b"accepted-then-revoked").await; + let security = ServerSecurity::new(auth); + security + .read_initial(&mut std::io::Cursor::new(bytes.clone())) + .await + .expect("first flight should be accepted"); + security + .auth() + .revoke(&admin_context, KeyId::from_u64(key_id)) + .await + .unwrap(); + let replayed = match security + .read_initial(&mut std::io::Cursor::new(bytes)) + .await + { + Ok(_) => panic!("replay after revoke should fail"), + Err(error) => error, + }; + assert_eq!(replayed.failure.code, "temporary_key_revoked"); + assert!( + replayed.response_session.is_none(), + "a previously accepted salt must not reuse nonce 0 for the revoke error" + ); + + let _ = std::fs::remove_dir_all(config.state_dir); +} + +#[tokio::test] +async fn rotated_temporary_first_flight_returns_a_readable_rotated_error() { + let _process_credential_guard = PROCESS_CREDENTIAL_TEST_LOCK.lock().await; + let admin = *b"0123456789abcdefghijklmnopqrstuv"; + let new_admin = *b"abcdefghijklmnopqrstuvwxyz012345"; + let config = temp_config(); + let auth = AuthRuntime::start(admin, config.clone()).await.unwrap(); + let admin_context = auth.authenticate_presented(ADMIN_KEY_ID, &admin).unwrap(); + let issued = auth + .issue(&admin_context, std::time::Duration::from_secs(60), None) + .await + .unwrap(); + let Credential::Temporary { key_id, key } = parse_credential(&issued.credential).unwrap() + else { + panic!("expected temporary credential"); + }; + let client = ClientHeaderSession::new_v2(&Credential::Temporary { key_id, key }).unwrap(); + auth.rotate_root(&admin_context, new_admin).await.unwrap(); + + let security = ServerSecurity::new(auth); + let (mut client_io, mut server_io) = tokio::io::duplex(4096); + let client_task = async { + client + .write_initial(&mut client_io, b"stale-after-rotate") + .await + .unwrap(); + let mut reader = client.response_reader(&mut client_io).unwrap(); + reader.read_msg().await.unwrap().to_vec() + }; + let server_task = async { + let error = match security.read_initial(&mut server_io).await { + Ok(_) => panic!("rotated credential should fail"), + Err(error) => error, + }; + assert_eq!(error.failure.code, "temporary_key_rotated"); + let session = error.response_session.expect("readable error session"); + let mut writer = session.response_writer(&mut server_io).unwrap(); + writer.write_msg(b"temporary_key_rotated").await.unwrap(); + error.failure.code + }; + let (plaintext, code) = tokio::join!(client_task, server_task); + assert_eq!(code, "temporary_key_rotated"); + assert_eq!(plaintext, b"temporary_key_rotated"); + + let _ = std::fs::remove_dir_all(config.state_dir); +} + +#[tokio::test] +async fn stale_root_replay_omits_the_rotated_error_session() { + let _process_credential_guard = PROCESS_CREDENTIAL_TEST_LOCK.lock().await; + let admin = *b"0123456789abcdefghijklmnopqrstuv"; + let new_admin = *b"abcdefghijklmnopqrstuvwxyz012345"; + let config = temp_config(); + let auth = AuthRuntime::start(admin, config.clone()).await.unwrap(); + let admin_context = auth.authenticate_presented(ADMIN_KEY_ID, &admin).unwrap(); + let issued = auth + .issue(&admin_context, std::time::Duration::from_secs(60), None) + .await + .unwrap(); + let Credential::Temporary { key_id, key } = parse_credential(&issued.credential).unwrap() + else { + panic!("expected temporary credential"); + }; + let client = ClientHeaderSession::new_v2(&Credential::Temporary { key_id, key }).unwrap(); + let bytes = encode_initial(&client, b"stale-replay").await; + auth.rotate_root(&admin_context, new_admin).await.unwrap(); + let security = ServerSecurity::new(auth); + let first = match security + .read_initial(&mut std::io::Cursor::new(bytes.clone())) + .await + { + Ok(_) => panic!("rotated credential should fail"), + Err(error) => error, + }; + let second = match security + .read_initial(&mut std::io::Cursor::new(bytes)) + .await + { + Ok(_) => panic!("rotated credential replay should fail"), + Err(error) => error, + }; + assert_eq!(first.failure.code, "temporary_key_rotated"); + assert!(first.response_session.is_some()); + assert_eq!(second.failure.code, "temporary_key_rotated"); + assert!( + second.response_session.is_none(), + "a claimed stale-root salt must not reuse nonce 0" + ); + + let _ = std::fs::remove_dir_all(config.state_dir); +} + +#[tokio::test] +async fn reset_temporary_first_flight_returns_a_readable_rotated_error() { + let _process_credential_guard = PROCESS_CREDENTIAL_TEST_LOCK.lock().await; + let admin = *b"0123456789abcdefghijklmnopqrstuv"; + let config = temp_config(); + let auth = AuthRuntime::start(admin, config.clone()).await.unwrap(); + let admin_context = auth.authenticate_presented(ADMIN_KEY_ID, &admin).unwrap(); + let issued = auth + .issue(&admin_context, std::time::Duration::from_secs(60), None) + .await + .unwrap(); + let Credential::Temporary { key_id, key } = parse_credential(&issued.credential).unwrap() + else { + panic!("expected temporary credential"); + }; + let client = ClientHeaderSession::new_v2(&Credential::Temporary { key_id, key }).unwrap(); + auth.reset(&admin_context).await.unwrap(); + + let security = ServerSecurity::new(auth); + let (mut client_io, mut server_io) = tokio::io::duplex(4096); + let client_task = async { + client + .write_initial(&mut client_io, b"stale-after-reset") + .await + .unwrap(); + let mut reader = client.response_reader(&mut client_io).unwrap(); + reader.read_msg().await.unwrap().to_vec() + }; + let server_task = async { + let error = match security.read_initial(&mut server_io).await { + Ok(_) => panic!("reset credential should fail"), + Err(error) => error, + }; + assert_eq!(error.failure.code, "temporary_key_rotated"); + let session = error.response_session.expect("readable error session"); + let mut writer = session.response_writer(&mut server_io).unwrap(); + writer.write_msg(b"temporary_key_rotated").await.unwrap(); + error.failure.code + }; + let (plaintext, code) = tokio::join!(client_task, server_task); + assert_eq!(code, "temporary_key_rotated"); + assert_eq!(plaintext, b"temporary_key_rotated"); + + let _ = std::fs::remove_dir_all(config.state_dir); +} + +#[tokio::test] +async fn mistyped_temporary_first_flight_does_not_send_an_unreadable_error() { + let admin = *b"0123456789abcdefghijklmnopqrstuv"; + let config = temp_config(); + let auth = AuthRuntime::start(admin, config.clone()).await.unwrap(); + let admin_context = auth.authenticate_presented(ADMIN_KEY_ID, &admin).unwrap(); + let issued = auth + .issue(&admin_context, std::time::Duration::from_secs(60), None) + .await + .unwrap(); + let Credential::Temporary { key_id, mut key } = parse_credential(&issued.credential).unwrap() + else { + panic!("expected temporary credential"); + }; + key[0] ^= 0x01; + let client = ClientHeaderSession::new_v2(&Credential::Temporary { key_id, key }).unwrap(); + let security = ServerSecurity::new(auth); + let (mut client_io, mut server_io) = tokio::io::duplex(4096); + let client_task = client.write_initial(&mut client_io, b"mistyped"); + let server_task = security.read_initial(&mut server_io); + let (client_result, server_result) = tokio::join!(client_task, server_task); + client_result.unwrap(); + let error = match server_result { + Ok(_) => panic!("mistyped temporary credential should fail decryption"), + Err(error) => error, + }; + assert_eq!(error.failure.code, "protocol_v2_decrypt_failed"); + assert!( + error.response_session.is_none(), + "the presenter cannot open a session derived from the live key" + ); + assert_eq!(error.presented_key_id, Some(KeyId::from_u64(key_id))); + + let _ = std::fs::remove_dir_all(config.state_dir); +} + +#[tokio::test] +async fn oversized_initial_frame_is_rejected_before_reading_its_body() { + let credential = Credential::Admin(*b"0123456789abcdefghijklmnopqrstuv"); + let session = ClientHeaderSession::new_v2(&credential).unwrap(); + let material = session.v2.as_ref().unwrap(); + let mut bytes = first_prefix(material); + bytes.extend_from_slice(&0_u64.to_be_bytes()); + bytes.extend_from_slice(&(MAX_INITIAL_PLAINTEXT_LEN + 17).to_be_bytes()); + let config = temp_config(); + let auth = AuthRuntime::start(*credential.key(), config.clone()) + .await + .unwrap(); + let security = ServerSecurity::new(auth); + let mut input = std::io::Cursor::new(bytes); + + let error = match security.read_initial(&mut input).await { + Ok(_) => panic!("oversized initial frame was accepted"), + Err(error) => error, + }; + assert_eq!(error.failure.code, "protocol_v2_decrypt_failed"); + assert!(error.failure.message.contains("65536-byte limit")); + + let _ = std::fs::remove_dir_all(config.state_dir); +} + +#[tokio::test] +async fn oversized_legacy_initial_frame_is_rejected_before_reading_its_body() { + use pb_mapper_core::checksum::get_checksum_for_key; + + let admin = *b"0123456789abcdefghijklmnopqrstuv"; + let config = temp_config(); + let auth = AuthRuntime::start(admin, config.clone()).await.unwrap(); + let datalen = MAX_INITIAL_CIPHERTEXT_LEN + 1; + let checksum = get_checksum_for_key(datalen, &admin); + let mut bytes = checksum.to_be_bytes().to_vec(); + bytes.extend_from_slice(&datalen.to_be_bytes()); + let security = ServerSecurity::new(auth); + let error = match security + .read_initial(&mut std::io::Cursor::new(bytes)) + .await + { + Ok(_) => panic!("oversized legacy frame was accepted"), + Err(error) => error, + }; + assert_eq!(error.failure.code, "legacy_frame_invalid"); + + let _ = std::fs::remove_dir_all(config.state_dir); +} + +#[test] +fn rotating_bloom_covers_current_and_previous_window() { + let mut bloom = RotatingBloom::new(1024, DEFAULT_REPLAY_WINDOW_SECONDS); + let value = [7_u8; 32]; + let start = bloom.current_started_at; + assert!(!bloom.contains(&value, start)); + bloom.insert(&value, start); + assert!(bloom.contains(&value, start + DEFAULT_REPLAY_WINDOW_SECONDS)); + assert!(!bloom.contains( + &value, + start + DEFAULT_REPLAY_WINDOW_SECONDS.saturating_mul(2) + 1 + )); +} + +#[test] +fn rotating_bloom_retains_fingerprints_for_the_clock_skew_window() { + let mut bloom = RotatingBloom::new(1024, DEFAULT_REPLAY_WINDOW_SECONDS); + let value = [9_u8; 32]; + let start = bloom.current_started_at; + bloom.insert(&value, start); + assert!(bloom.contains(&value, start + 120)); + assert!(bloom.contains(&value, start + MAX_CONNECTION_CLOCK_SKEW_SECONDS)); +} + +#[test] +fn rotating_bloom_retains_a_max_future_timestamp_past_the_next_rotation() { + let mut bloom = RotatingBloom::new(1024, DEFAULT_REPLAY_WINDOW_SECONDS); + let value = [11_u8; 32]; + let start = bloom.current_started_at; + let insert_at = start + DEFAULT_REPLAY_WINDOW_SECONDS - 1; + bloom.insert(&value, insert_at); + assert!(bloom.contains( + &value, + insert_at + MAX_CONNECTION_CLOCK_SKEW_SECONDS.saturating_mul(2) - 1 + )); +} + +#[test] +fn per_credential_admission_limit_does_not_consume_other_keys() { + let now = unix_seconds(); + let mut guard = + ReplayGuard::open(None, 1024, DEFAULT_REPLAY_WINDOW_SECONDS).with_max_per_key(2); + assert_eq!( + guard.admit(KeyId::from_u64(1), &[1_u8; 32], now), + FirstFlightAdmit::Fresh + ); + assert_eq!( + guard.admit(KeyId::from_u64(1), &[2_u8; 32], now), + FirstFlightAdmit::Fresh + ); + assert_eq!( + guard.admit(KeyId::from_u64(1), &[3_u8; 32], now), + FirstFlightAdmit::Limited + ); + assert_eq!( + guard.admit(KeyId::from_u64(1), &[3_u8; 32], now), + FirstFlightAdmit::Limited + ); + assert_eq!( + guard.admit(KeyId::from_u64(2), &[4_u8; 32], now), + FirstFlightAdmit::Fresh + ); + assert_eq!( + guard.admit(KeyId::from_u64(1), &[1_u8; 32], now), + FirstFlightAdmit::Replayed + ); +} + +#[test] +fn aggregate_admission_limit_covers_all_keys() { + let now = unix_seconds(); + let mut guard = ReplayGuard::open(None, 1024, DEFAULT_REPLAY_WINDOW_SECONDS) + .with_max_per_key(100) + .with_max_total(2); + assert_eq!( + guard.admit(KeyId::from_u64(1), &[1_u8; 32], now), + FirstFlightAdmit::Fresh + ); + assert_eq!( + guard.admit(KeyId::from_u64(2), &[2_u8; 32], now), + FirstFlightAdmit::Fresh + ); + assert_eq!( + guard.admit(KeyId::from_u64(3), &[3_u8; 32], now), + FirstFlightAdmit::Limited + ); + assert_eq!( + guard.admit(KeyId::from_u64(3), &[3_u8; 32], now), + FirstFlightAdmit::Limited + ); + assert_eq!( + guard.admit(KeyId::from_u64(4), &[4_u8; 32], now), + FirstFlightAdmit::Limited + ); +} + +#[test] +fn persisted_first_flights_survive_a_torn_trailing_record() { + let mut random = [0_u8; 8]; + let mut rng = rand::rng(); + for byte in &mut random { + *byte = rng.random(); + } + let path = std::env::temp_dir().join(format!( + "pb-mapper-replay-torn-{}", + u64::from_be_bytes(random) + )); + let now = unix_seconds(); + let fingerprint = [17_u8; 32]; + { + let mut guard = ReplayGuard::open(Some(path.clone()), 1024, DEFAULT_REPLAY_WINDOW_SECONDS); + assert_eq!( + guard.admit(KeyId::from_u64(7), &fingerprint, now), + FirstFlightAdmit::Fresh + ); + } + let mut torn = std::fs::read(&path).unwrap(); + torn.extend_from_slice(&[0_u8; 10]); + std::fs::write(&path, torn).unwrap(); + let mut restored = ReplayGuard::open(Some(path.clone()), 1024, DEFAULT_REPLAY_WINDOW_SECONDS); + assert_eq!( + restored.admit(KeyId::from_u64(7), &fingerprint, now), + FirstFlightAdmit::Unavailable + ); + let _ = std::fs::remove_file(path); +} + +#[test] +fn replay_rewrite_succeeds_when_a_pid_temporary_file_already_exists() { + let mut random = [0_u8; 8]; + let mut rng = rand::rng(); + for byte in &mut random { + *byte = rng.random(); + } + let path = std::env::temp_dir().join(format!( + "pb-mapper-replay-tmp-{}", + u64::from_be_bytes(random) + )); + let now = unix_seconds(); + let fingerprint = [19_u8; 32]; + { + let mut guard = ReplayGuard::open(Some(path.clone()), 1024, DEFAULT_REPLAY_WINDOW_SECONDS); + assert_eq!( + guard.admit(KeyId::from_u64(9), &fingerprint, now), + FirstFlightAdmit::Fresh + ); + } + let leftover = path.with_file_name(format!( + ".{}.tmp-{}", + path.file_name().unwrap().to_str().unwrap(), + std::process::id() + )); + std::fs::write(&leftover, b"stale").unwrap(); + let mut restored = ReplayGuard::open(Some(path.clone()), 1024, DEFAULT_REPLAY_WINDOW_SECONDS); + assert_eq!( + restored.admit(KeyId::from_u64(9), &fingerprint, now), + FirstFlightAdmit::Replayed + ); + let _ = std::fs::remove_file(leftover); + let _ = std::fs::remove_file(path); +} + +#[test] +fn persisted_first_flights_survive_replay_guard_restart() { + let mut random = [0_u8; 8]; + let mut rng = rand::rng(); + for byte in &mut random { + *byte = rng.random(); + } + let path = + std::env::temp_dir().join(format!("pb-mapper-replay-{}", u64::from_be_bytes(random))); + let now = unix_seconds(); + let fingerprint = [13_u8; 32]; + { + let mut guard = ReplayGuard::open(Some(path.clone()), 1024, DEFAULT_REPLAY_WINDOW_SECONDS); + assert_eq!( + guard.admit(KeyId::from_u64(7), &fingerprint, now), + FirstFlightAdmit::Fresh + ); + } + let mut restored = ReplayGuard::open(Some(path.clone()), 1024, DEFAULT_REPLAY_WINDOW_SECONDS); + assert_eq!( + restored.admit(KeyId::from_u64(7), &fingerprint, now), + FirstFlightAdmit::Replayed + ); + let _ = std::fs::remove_file(path); +} + +#[test] +fn restored_replay_log_consumes_the_aggregate_budget() { + let mut random = [0_u8; 8]; + let mut rng = rand::rng(); + for byte in &mut random { + *byte = rng.random(); + } + let path = std::env::temp_dir().join(format!( + "pb-mapper-replay-total-{}", + u64::from_be_bytes(random) + )); + let now = unix_seconds(); + { + let mut guard = ReplayGuard::open(Some(path.clone()), 1024, DEFAULT_REPLAY_WINDOW_SECONDS) + .with_max_per_key(100) + .with_max_total(2); + assert_eq!( + guard.admit(KeyId::from_u64(1), &[1_u8; 32], now), + FirstFlightAdmit::Fresh + ); + assert_eq!( + guard.admit(KeyId::from_u64(2), &[2_u8; 32], now), + FirstFlightAdmit::Fresh + ); + } + let mut restored = ReplayGuard::open(Some(path.clone()), 1024, DEFAULT_REPLAY_WINDOW_SECONDS) + .with_max_per_key(100) + .with_max_total(2); + assert_eq!( + restored.admit(KeyId::from_u64(3), &[3_u8; 32], now), + FirstFlightAdmit::Limited + ); + let _ = std::fs::remove_file(path); +} + +#[test] +fn restored_replay_generation_ages_with_loaded_records() { + let mut random = [0_u8; 8]; + let mut rng = rand::rng(); + for byte in &mut random { + *byte = rng.random(); + } + let path = std::env::temp_dir().join(format!( + "pb-mapper-replay-age-{}", + u64::from_be_bytes(random) + )); + let start = unix_seconds().saturating_sub(DEFAULT_REPLAY_WINDOW_SECONDS / 2); + { + let mut guard = ReplayGuard::open(Some(path.clone()), 1024, DEFAULT_REPLAY_WINDOW_SECONDS) + .with_max_per_key(100) + .with_max_total(2); + assert_eq!( + guard.admit(KeyId::from_u64(1), &[1_u8; 32], start), + FirstFlightAdmit::Fresh + ); + assert_eq!( + guard.admit(KeyId::from_u64(2), &[2_u8; 32], start), + FirstFlightAdmit::Fresh + ); + } + let mut restored = ReplayGuard::open(Some(path.clone()), 1024, DEFAULT_REPLAY_WINDOW_SECONDS) + .with_max_per_key(100) + .with_max_total(2); + assert_eq!( + restored.admit( + KeyId::from_u64(3), + &[3_u8; 32], + start.saturating_add(DEFAULT_REPLAY_WINDOW_SECONDS) + ), + FirstFlightAdmit::Fresh + ); + let _ = std::fs::remove_file(path); +} + +#[tokio::test] +async fn legacy_initial_frame_validates_against_isolated_relay_key() { + let _process_credential_guard = PROCESS_CREDENTIAL_TEST_LOCK.lock().await; + let config = temp_config(); + let isolated_admin = *b"isolated-admin-key-0123456789abc"; + std::fs::create_dir_all(&config.state_dir).unwrap(); + std::fs::write(config.state_dir.join("admin.key"), isolated_admin).unwrap(); + let auth = AuthRuntime::from_isolated_state(config.clone()) + .await + .unwrap(); + + let temporary_key_id = 1; + let temporary_key = *b"temporary-remote-key-0123456789a"; + set_process_msg_header_key(Some(&encode_temporary_credential( + temporary_key_id, + &temporary_key, + ))) + .unwrap(); + + let security = ServerSecurity::new(auth); + let (mut client_io, mut server_io) = tokio::io::duplex(4096); + let client = ClientHeaderSession::new_legacy(isolated_admin); + let client_task = async { + client + .write_initial(&mut client_io, b"legacy-isolated") + .await + .unwrap(); + let mut reader = client.response_reader(&mut client_io).unwrap(); + reader.read_msg().await.unwrap().to_vec() + }; + let server_task = async { + let initial = security.read_initial(&mut server_io).await.unwrap(); + assert_eq!(initial.payload, b"legacy-isolated"); + assert_eq!(initial.session.protocol(), HeaderProtocol::Legacy); + let mut writer = initial.session.response_writer(&mut server_io).unwrap(); + writer.write_msg(b"legacy-response").await.unwrap(); + }; + let (response, _) = tokio::join!(client_task, server_task); + assert_eq!(response, b"legacy-response"); + + set_process_msg_header_key(None).unwrap(); + let _ = std::fs::remove_dir_all(config.state_dir); +} + +#[test] +fn failure_log_limiter_has_a_hard_cardinality_bound() { + let mut limiter = FailureLogLimiter::default(); + let peer = "127.0.0.1".parse().unwrap(); + for key_id in 0..10_000 { + limiter.record(peer, KeyId::from_u64(key_id), "invalid", 1_000); + } + assert_eq!(limiter.entries.len(), 4096); + assert!(limiter.overflow.is_some()); +} diff --git a/crates/pb-mapper-server/Cargo.toml b/crates/pb-mapper-server/Cargo.toml new file mode 100644 index 0000000..6efb20a --- /dev/null +++ b/crates/pb-mapper-server/Cargo.toml @@ -0,0 +1,27 @@ +[package] +name = "pb-mapper-server" +version.workspace = true +edition.workspace = true +authors.workspace = true + +[dependencies] +pb-mapper-auth.workspace = true +pb-mapper-core.workspace = true +pb-mapper-protocol.workspace = true + +hashbrown.workspace = true +kanal.workspace = true +snafu.workspace = true +tokio.workspace = true +tokio-util.workspace = true +tracing.workspace = true +uni-stream.workspace = true + +[dev-dependencies] +rand.workspace = true + +[features] +udp-timeout = ["uni-stream/udp-timeout", "pb-mapper-protocol/udp-timeout"] + +[lints] +workspace = true diff --git a/crates/pb-mapper-server/src/admin.rs b/crates/pb-mapper-server/src/admin.rs new file mode 100644 index 0000000..b119d6d --- /dev/null +++ b/crates/pb-mapper-server/src/admin.rs @@ -0,0 +1,402 @@ +//! Server-side administrator request execution. +//! +//! ```text +//! authenticated AdminRequest -> revalidated AuthContext -> auth actor / manager +//! -> AdminResponse +//! ``` +//! +//! Credential lifecycle operations go to `AuthRuntime`; service and connection +//! inventory requests go to the routing manager. Read operations are audited without +//! weakening the primary response when only audit emission fails. + +use std::time::Duration; + +use tokio::net::TcpStream; + +use super::error::Error; +use super::{ManagerTask, ManagerTaskSender, Result}; +use pb_mapper_auth::{AuthContext, AuthFailure, AuthRuntime, KeyId}; +use pb_mapper_core::checksum::{Credential, parse_credential}; +use pb_mapper_core::conn_id::RemoteConnId; +use pb_mapper_protocol::MessageWriter; +use pb_mapper_protocol::command::{AdminRequest, AdminResponse, MessageSerializer, PbConnResponse}; +use pb_mapper_protocol::secure::ServerHeaderSession; + +pub async fn handle_admin_request( + request: AdminRequest, + authorization: AuthContext, + auth: AuthRuntime, + manager: ManagerTaskSender, + conn_id: RemoteConnId, + mut conn: TcpStream, + session: ServerHeaderSession, +) -> Result<()> { + let result = execute(request, &authorization, auth, manager).await; + let response = match result { + Ok(response) => PbConnResponse::Admin(response), + Err(failure) => { + tracing::warn!( + event = "admin_operation_failed", + auth_stage = "permission_or_state", + conn_id = %conn_id, + reason = %failure.code, + retryable = failure.retryable, + error = %failure.message, + "administrator operation failed" + ); + PbConnResponse::error(failure.code, failure.message, failure.retryable) + } + }; + let message = response.encode().map_err(|error| Error::AdminOperation { + detail: format!("failed to encode response: {error}"), + })?; + let mut writer = session + .response_writer(&mut conn) + .map_err(|error| Error::AdminOperation { + detail: format!("failed to create response writer: {error}"), + })?; + writer + .write_msg(&message) + .await + .map_err(|error| Error::AdminOperation { + detail: format!("failed to write response: {error}"), + }) +} + +async fn execute( + request: AdminRequest, + authorization: &AuthContext, + auth: AuthRuntime, + manager: ManagerTaskSender, +) -> std::result::Result { + match request { + AdminRequest::KeyIssue { ttl_seconds, label } => auth + .issue(authorization, Duration::from_secs(ttl_seconds), label) + .await + .map(AdminResponse::KeyIssued), + AdminRequest::KeyList { page, page_size } => { + audit_read( + &auth, + authorization, + "temporary_key_list", + None, + Some(format!("page={page},page_size={page_size}")), + ) + .await; + auth.list(authorization, page, page_size) + .await + .map(AdminResponse::KeyList) + } + AdminRequest::KeyShow { key_id } => auth + .show(authorization, KeyId::from_u64(key_id), false) + .await + .map(AdminResponse::KeyShown), + AdminRequest::KeyReveal { key_id } => auth + .show(authorization, KeyId::from_u64(key_id), true) + .await + .map(AdminResponse::KeyShown), + AdminRequest::KeyRenew { + key_id, + ttl_seconds, + } => auth + .renew( + authorization, + KeyId::from_u64(key_id), + Duration::from_secs(ttl_seconds), + ) + .await + .map(AdminResponse::KeyRenewed), + AdminRequest::KeyRevoke { key_id } => auth + .revoke(authorization, KeyId::from_u64(key_id)) + .await + .map(AdminResponse::KeyRevoked), + AdminRequest::KeyGc => auth + .gc(authorization) + .await + .map(|removed| AdminResponse::KeyGc { removed }), + AdminRequest::AuthStatus => { + audit_read(&auth, authorization, "auth_status", None, None).await; + auth.status(authorization) + .await + .map(AdminResponse::AuthStatus) + } + AdminRequest::AuthStateReset { confirm } => { + if !confirm { + return Err(AuthFailure::new( + "confirmation_required", + "auth-state reset requires explicit confirmation", + false, + )); + } + auth.reset(authorization).await?; + Ok(AdminResponse::Ok { + action: "auth_state_reset".to_string(), + }) + } + AdminRequest::RootKeyRotate { new_admin_key } => { + let Credential::Admin(new_key) = parse_credential(new_admin_key.trim()) + .map_err(|error| AuthFailure::new("administrator_key_invalid", error, false))? + else { + return Err(AuthFailure::new( + "administrator_key_invalid", + "root rotation requires a 32-byte administrator key", + false, + )); + }; + auth.rotate_root(authorization, new_key).await?; + Ok(AdminResponse::Ok { + action: "administrator_key_rotated".to_string(), + }) + } + AdminRequest::LegacyProtocolSet { policy } => { + auth.set_legacy_protocol(authorization, policy).await?; + Ok(AdminResponse::Ok { + action: "legacy_protocol_updated".to_string(), + }) + } + AdminRequest::ServiceList { + key_id, + page, + page_size, + } => { + audit_read( + &auth, + authorization, + "service_list", + key_id.map(KeyId::from_u64), + Some(format!("page={page},page_size={page_size}")), + ) + .await; + let (response_sender, receiver) = tokio::sync::oneshot::channel(); + let response = query_inventory( + authorization, + &manager, + ManagerTask::AdminServiceList { + key_id, + page, + page_size, + response_sender, + }, + receiver, + "service", + ) + .await?; + Ok(AdminResponse::Services(response)) + } + AdminRequest::ConnectionList { + key_id, + page, + page_size, + } => { + audit_read( + &auth, + authorization, + "connection_list", + key_id.map(KeyId::from_u64), + Some(format!("page={page},page_size={page_size}")), + ) + .await; + let (response_sender, receiver) = tokio::sync::oneshot::channel(); + let response = query_inventory( + authorization, + &manager, + ManagerTask::AdminConnectionList { + key_id, + page, + page_size, + response_sender, + }, + receiver, + "connection", + ) + .await?; + Ok(AdminResponse::Connections(response)) + } + } +} + +/// Dispatch an inventory read only while the administrator lease remains current. +/// +/// Root rotation cancels the old lease. Racing both channel operations against that +/// cancellation prevents a request authenticated under the old root from waiting for or +/// returning relay inventory after the rotation has taken effect. The final revalidation +/// establishes the successful read's authorization point after the manager produced its page. +async fn query_inventory( + authorization: &AuthContext, + manager: &ManagerTaskSender, + task: ManagerTask, + receiver: tokio::sync::oneshot::Receiver, + inventory: &'static str, +) -> std::result::Result { + let cancellation = authorization.admin_cancellation_token()?; + tokio::select! { + biased; + _ = cancellation.cancelled() => return Err(cancelled_authorization(authorization)), + result = manager.send(task) => result.map_err(|_| { + AuthFailure::new( + "server_state_unavailable", + "relay connection manager is unavailable", + true, + ) + })?, + } + let response = tokio::select! { + biased; + _ = cancellation.cancelled() => return Err(cancelled_authorization(authorization)), + result = receiver => result.map_err(|_| { + AuthFailure::new( + "server_state_unavailable", + format!("relay connection manager dropped the {inventory} query"), + true, + ) + })?, + }; + authorization.ensure_active()?; + Ok(response) +} + +fn cancelled_authorization(authorization: &AuthContext) -> AuthFailure { + match authorization.ensure_active() { + Err(error) => error, + Ok(_) => AuthFailure::new( + "administrator_key_rotated", + "administrator credential lease was cancelled during the inventory query", + false, + ), + } +} + +async fn audit_read( + auth: &AuthRuntime, + authorization: &AuthContext, + action: &str, + key_id: Option, + detail: Option, +) { + if let Err(error) = auth + .audit_admin(authorization, action, key_id, detail) + .await + { + tracing::warn!( + event = "admin_audit_failed", + auth_stage = "audit", + action, + reason = %error.code, + error = %error.message, + "administrator read operation could not be audited" + ); + } +} + +#[cfg(test)] +mod tests { + use super::*; + use pb_mapper_auth::ADMIN_KEY_ID; + use pb_mapper_auth::{AuthConfig, LegacyProtocolPolicy}; + + fn temp_state_dir(name: &str) -> std::path::PathBuf { + std::env::temp_dir().join(format!( + "pb-mapper-admin-{name}-{}", + std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .map(|duration| duration.as_nanos()) + .unwrap_or_default() + )) + } + + async fn inventory_query_rejects_rotation(connection_query: bool) { + let _process_credential_guard = pb_mapper_core::test_support::PROCESS_CREDENTIAL_TEST_LOCK + .lock() + .await; + let state_dir = temp_state_dir(if connection_query { + "connections" + } else { + "services" + }); + let old_key = *b"0123456789abcdefghijklmnopqrstuv"; + let new_key = *b"abcdefghijklmnopqrstuvwxyz012345"; + let runtime = AuthRuntime::start( + old_key, + AuthConfig { + state_dir: state_dir.clone(), + max_temporary_keys: 4, + max_temporary_key_ttl: Duration::from_secs(3600), + legacy_protocol: LegacyProtocolPolicy::Allow, + }, + ) + .await + .expect("authentication runtime should start"); + let admin = runtime + .authenticate_presented(ADMIN_KEY_ID, &old_key) + .expect("old administrator key should authenticate"); + let request = if connection_query { + AdminRequest::ConnectionList { + key_id: None, + page: 0, + page_size: 100, + } + } else { + AdminRequest::ServiceList { + key_id: None, + page: 0, + page_size: 100, + } + }; + let (manager, receiver) = kanal::unbounded_async(); + let request_admin = admin.clone(); + let request_runtime = runtime.clone(); + let pending = tokio::spawn(async move { + execute(request, &request_admin, request_runtime, manager).await + }); + + let manager_task = receiver + .recv() + .await + .expect("inventory request should reach the manager"); + runtime + .rotate_root(&admin, new_key) + .await + .expect("root rotation should succeed"); + match manager_task { + ManagerTask::AdminServiceList { + response_sender, .. + } => { + let _ = response_sender.send(pb_mapper_protocol::command::AdminServicePage { + schema_version: 1, + items: Vec::new(), + next_page: None, + }); + } + ManagerTask::AdminConnectionList { + response_sender, .. + } => { + let _ = response_sender.send(pb_mapper_protocol::command::AdminConnectionPage { + schema_version: 1, + items: Vec::new(), + next_page: None, + }); + } + _ => panic!("expected an administrator inventory manager task"), + } + + let failure = pending + .await + .expect("inventory task should not panic") + .expect_err("rotated administrator must not receive inventory"); + assert_eq!(failure.code, "administrator_key_rotated"); + + drop(runtime); + tokio::time::sleep(Duration::from_millis(20)).await; + let _ = std::fs::remove_dir_all(state_dir); + } + + #[tokio::test] + async fn service_inventory_rejects_root_rotation_after_dispatch() { + inventory_query_rejects_rotation(false).await; + } + + #[tokio::test] + async fn connection_inventory_rejects_root_rotation_after_dispatch() { + inventory_query_rejects_rotation(true).await; + } +} diff --git a/src/pb_server/client.rs b/crates/pb-mapper-server/src/client.rs similarity index 73% rename from src/pb_server/client.rs rename to crates/pb-mapper-server/src/client.rs index 317412a..960cad8 100644 --- a/src/pb_server/client.rs +++ b/crates/pb-mapper-server/src/client.rs @@ -1,8 +1,9 @@ +use std::sync::Arc; use std::time::Duration; use snafu::ResultExt; use tokio::net::TcpStream; -use tokio::time::{timeout, Instant}; +use tokio::time::{Instant, timeout}; use tracing::instrument; use super::error::{ @@ -12,21 +13,22 @@ use super::error::{ ClientConnSubcribeRespNotMatchSnafu, ClientConnWriteSubcribeRespSnafu, }; use super::{ConnTask, ImutableKey, ManagerTask, ManagerTaskSender, Result}; -use crate::common::checksum::{gen_random_key, AesKeyType}; -use crate::common::config::{stream_ack_timeout, stream_ready_timeout, stream_recovery_timeout}; -use crate::common::conn_id::RemoteConnId; -use crate::common::message::command::{MessageSerializer, PbConnResponse}; -use crate::common::message::forward::{ - start_datagram_forward, start_forward, CodecDatagramReader, CodecDatagramWriter, - CodecForwardReader, CodecForwardWriter, NormalDatagramReader, NormalDatagramWriter, - NormalForwardReader, NormalForwardWriter, -}; -use crate::common::message::{get_decodec, get_encodec, get_header_msg_writer, MessageWriter}; -use crate::pb_server::error::{ +use crate::error::{ ClientConnCreateHeaderToolSnafu, ClientConnEncodeStreamRespSnafu, ClientConnWriteStreamRespSnafu, }; -use crate::{create_component, snafu_error_get_or_return_ok, start_forward_with_codec_key}; +use pb_mapper_core::checksum::{AesKeyType, gen_random_key}; +use pb_mapper_core::config::{stream_ack_timeout, stream_ready_timeout, stream_recovery_timeout}; +use pb_mapper_core::conn_id::RemoteConnId; +use pb_mapper_core::snafu_error_get_or_return_ok; +use pb_mapper_protocol::command::{MessageSerializer, PbConnResponse}; +use pb_mapper_protocol::forward::{ + CodecDatagramReader, CodecDatagramWriter, CodecForwardReader, CodecForwardWriter, + NormalDatagramReader, NormalDatagramWriter, NormalForwardReader, NormalForwardWriter, + start_datagram_forward, start_forward, +}; +use pb_mapper_protocol::secure::ServerHeaderSession; +use pb_mapper_protocol::{MessageWriter, get_decodec, get_encodec}; /// Ensure that client-side connections are properly deregistered before a normal connection is /// disconnected or an exception occurs @@ -137,17 +139,26 @@ const CLIENT_CONN_CONTROL_TIMEOUT: Duration = Duration::from_secs(30); /// 1. Request server stream /// 2. Forward the traffic between client stream and server stream -#[instrument(skip(task_sender, conn))] +#[instrument(skip(task_sender, conn, session))] pub async fn handle_client_conn( key: ImutableKey, conn_id: RemoteConnId, task_sender: ManagerTaskSender, mut conn: TcpStream, + session: ServerHeaderSession, ) -> Result<()> { let prev_time = Instant::now(); let mut guard = ClientConnGuard::new(conn_id, None, task_sender.clone(), key.clone()); - let (mut server_stream, server_id, codec_key, is_datagram) = - match get_server_stream(&mut conn, key.clone(), conn_id, task_sender.clone()).await { + let (mut server_stream, server_session, server_id, codec_key, is_datagram) = + match get_server_stream( + &mut conn, + &session, + key.clone(), + conn_id, + task_sender.clone(), + ) + .await + { Ok(res) => res, Err(e) => { tracing::warn!( @@ -162,8 +173,17 @@ pub async fn handle_client_conn( } }; guard.set_server_id(server_id); + let server_cancellation = server_session + .context() + .map_err(|error| super::error::Error::ClientConnAuthInactive { + detail: error.to_string(), + })? + .cancellation_token() + .map_err(|error| super::error::Error::ClientConnAuthInactive { + detail: error.to_string(), + })?; - let result = async { + let forwarding = async { let duration = Instant::now() - prev_time; tracing::info!( @@ -182,7 +202,8 @@ pub async fn handle_client_conn( // response message to server to indicate that stream handling has finished { - let mut msg_writer = get_header_msg_writer(&mut server_writer) + let mut msg_writer = server_session + .response_writer(&mut server_writer) .context(ClientConnCreateHeaderToolSnafu { tool: "writer" })?; let msg = PbConnResponse::Stream { codec_key }.encode().context( ClientConnEncodeStreamRespSnafu { @@ -199,6 +220,8 @@ pub async fn handle_client_conn( })?; } + let client_framing = session.framing_key(); + let server_framing = server_session.framing_key(); if is_datagram { match codec_key { Some(key) => { @@ -206,49 +229,94 @@ pub async fn handle_client_conn( CodecDatagramReader::new( &mut client_reader, snafu_error_get_or_return_ok!(get_decodec(&key)), - ), + ) + .with_checksum_key(client_framing), CodecDatagramWriter::new( &mut client_writer, snafu_error_get_or_return_ok!(get_encodec(&key)), - ), + ) + .with_checksum_key(client_framing), CodecDatagramReader::new( &mut server_reader, snafu_error_get_or_return_ok!(get_decodec(&key)), - ), + ) + .with_checksum_key(server_framing), CodecDatagramWriter::new( &mut server_writer, snafu_error_get_or_return_ok!(get_encodec(&key)), - ), + ) + .with_checksum_key(server_framing), ) .await; } None => { start_datagram_forward( - NormalDatagramReader::new(&mut client_reader), - NormalDatagramWriter::new(&mut client_writer), - NormalDatagramReader::new(&mut server_reader), - NormalDatagramWriter::new(&mut server_writer), + NormalDatagramReader::new(&mut client_reader) + .with_checksum_key(client_framing), + NormalDatagramWriter::new(&mut client_writer) + .with_checksum_key(client_framing), + NormalDatagramReader::new(&mut server_reader) + .with_checksum_key(server_framing), + NormalDatagramWriter::new(&mut server_writer) + .with_checksum_key(server_framing), ) .await; } } } else { - start_forward_with_codec_key!( - codec_key, - &mut client_reader, - &mut client_writer, - &mut server_reader, - &mut server_writer, - true, - true, - true, - true - ); + match codec_key { + Some(key) => { + start_forward( + CodecForwardReader::new( + &mut client_reader, + snafu_error_get_or_return_ok!(get_decodec(&key)), + ) + .with_checksum_key(client_framing), + CodecForwardWriter::new( + &mut client_writer, + snafu_error_get_or_return_ok!(get_encodec(&key)), + ) + .with_checksum_key(client_framing), + CodecForwardReader::new( + &mut server_reader, + snafu_error_get_or_return_ok!(get_decodec(&key)), + ) + .with_checksum_key(server_framing), + CodecForwardWriter::new( + &mut server_writer, + snafu_error_get_or_return_ok!(get_encodec(&key)), + ) + .with_checksum_key(server_framing), + ) + .await; + } + None => { + start_forward( + NormalForwardReader::new(&mut client_reader), + NormalForwardWriter::new(&mut client_writer), + NormalForwardReader::new(&mut server_reader), + NormalForwardWriter::new(&mut server_writer), + ) + .await; + } + } } Ok(()) - } - .await; + }; + let result = tokio::select! { + result = forwarding => result, + _ = server_cancellation.cancelled() => { + tracing::info!( + event = "connection_auth_expired", + key = %key, + client_conn_id = %conn_id, + server_conn_id = %server_id, + "closing active data stream because its registering credential expired or was revoked" + ); + Ok(()) + } + }; match &result { Ok(()) => tracing::info!( event = "client_forward_finished", @@ -297,14 +365,16 @@ async fn retire_server_conn( async fn write_subscribe_response( conn: &mut TcpStream, + session: &ServerHeaderSession, key: &ImutableKey, conn_id: RemoteConnId, server_id: RemoteConnId, codec_key: Option, is_datagram: bool, ) -> Result<()> { - let mut msg_writer = - get_header_msg_writer(conn).context(ClientConnCreateHeaderToolSnafu { tool: "writer" })?; + let mut msg_writer = session + .response_writer(conn) + .context(ClientConnCreateHeaderToolSnafu { tool: "writer" })?; let msg = PbConnResponse::Subcribe { codec_key, client_id: conn_id.into(), @@ -336,10 +406,17 @@ async fn write_subscribe_response( async fn get_server_stream( conn: &mut TcpStream, + session: &ServerHeaderSession, key: ImutableKey, conn_id: RemoteConnId, task_sender: ManagerTaskSender, -) -> Result<(TcpStream, RemoteConnId, Option, bool)> { +) -> Result<( + TcpStream, + ServerHeaderSession, + RemoteConnId, + Option, + bool, +)> { let (tx, rx) = kanal::bounded_async(DEFAULT_CLIENT_CHAN_CAP); let ack_timeout = stream_ack_timeout(); let ready_timeout = stream_ready_timeout(); @@ -407,7 +484,11 @@ async fn get_server_stream( ); (codec_key, is_datagram, server_conn_id, server_generation) } - ConnTask::SubcribeFailed { reason } => { + ConnTask::SubcribeFailed { + code, + reason, + retryable, + } => { tracing::warn!( event = "subscribe_failed", key = %key, @@ -415,6 +496,7 @@ async fn get_server_stream( reason = %reason, "subscribe failed before stream forwarding" ); + write_subscribe_error(conn, session, &code, &reason, retryable).await?; ClientConnSubcribeFailedSnafu { key: key.clone(), conn_id, @@ -492,10 +574,19 @@ async fn get_server_stream( server_id, server_generation: response_generation, stream, + session: stream_session, } if response_generation == server_generation => { - write_subscribe_response(conn, &key, conn_id, server_id, codec_key, is_datagram) - .await?; - return Ok((stream, server_id, codec_key, is_datagram)); + write_subscribe_response( + conn, + session, + &key, + conn_id, + server_id, + codec_key, + is_datagram, + ) + .await?; + return Ok((stream, stream_session, server_id, codec_key, is_datagram)); } ConnTask::StreamAck { server_id, @@ -557,12 +648,21 @@ async fn get_server_stream( server_id, server_generation: response_generation, stream, + session: stream_session, } = resp { if response_generation == server_generation { - write_subscribe_response(conn, &key, conn_id, server_id, codec_key, is_datagram) - .await?; - return Ok((stream, server_id, codec_key, is_datagram)); + write_subscribe_response( + conn, + session, + &key, + conn_id, + server_id, + codec_key, + is_datagram, + ) + .await?; + return Ok((stream, stream_session, server_id, codec_key, is_datagram)); } tracing::warn!( event = "server_stream_generation_mismatch", @@ -585,3 +685,28 @@ async fn get_server_stream( .fail()? } } + +async fn write_subscribe_error( + conn: &mut TcpStream, + session: &ServerHeaderSession, + code: &str, + reason: &str, + retryable: bool, +) -> Result<()> { + let message = PbConnResponse::error(code, reason, retryable) + .encode() + .context(ClientConnEncodeSubcribeRespSnafu { + key: Arc::from(""), + conn_id: RemoteConnId::default(), + })?; + let mut writer = session + .response_writer(conn) + .context(ClientConnCreateHeaderToolSnafu { tool: "writer" })?; + writer + .write_msg(&message) + .await + .context(ClientConnWriteSubcribeRespSnafu { + key: Arc::from(""), + conn_id: RemoteConnId::default(), + }) +} diff --git a/crates/pb-mapper-server/src/connection.rs b/crates/pb-mapper-server/src/connection.rs new file mode 100644 index 0000000..1134200 --- /dev/null +++ b/crates/pb-mapper-server/src/connection.rs @@ -0,0 +1,588 @@ +//! Per-connection admission, authentication, namespace resolution, and role dispatch. +//! +//! ```text +//! accepted TCP socket +//! | +//! v +//! bounded V2/legacy first frame -> AuthContext -> namespace policy +//! | | +//! +-> structured auth error +-> register / subscribe / stream +//! +-> status / administrator request +//! ``` +//! +//! Long-lived register, subscribe, and status futures are raced against the +//! credential's cancellation token here. This outer guard closes a subscriber even +//! when the paired service stream belongs to a different credential. + +use super::*; + +pub(super) async fn handle_listener( + task_sender: ManagerTaskSender, + listener: TcpListener, + keep_alive: bool, +) -> Result<()> { + loop { + let (stream, addr) = listener.accept().await.context(ServerListenSnafu)?; + tracing::debug!( + event = "tcp_conn_accepted", + peer_addr = %addr, + "accepted tcp connection" + ); + // set keepalive (optional) and nodelay + if keep_alive { + snafu_error_handle!(set_tcp_keep_alive(&stream).context(TaskCenterSetKeepAliveSnafu)); + } + snafu_error_handle!(set_tcp_nodelay(&stream), "remote stream set nodelay"); + task_sender + .send(ManagerTask::Accept { + stream, + peer_addr: addr, + }) + .await + .map_err(|_| kanal::SendError(())) + .context(TaskCenterSendListenerSnafu)? + } +} + +#[instrument(skip(manager_task_sender, conn, security), fields(conn_id = %conn_id, peer_addr = %peer_addr))] +pub(super) async fn handle_conn( + conn_id: RemoteConnId, + peer_addr: SocketAddr, + manager_task_sender: ManagerTaskSender, + mut conn: TcpStream, + security: ServerSecurity, +) -> Result<()> { + let timeout = control_io_timeout(); + let initial = match tokio::time::timeout(timeout, security.read_initial(&mut conn)).await { + Err(_) => TaskCenterInitRequestTimeoutSnafu { conn_id, timeout }.fail()?, + Ok(Err(error)) => { + let key_id = error + .response_session + .as_ref() + .map(|session| session.key_id()) + .or(error.presented_key_id) + .unwrap_or(ADMIN_KEY_ID); + let decision = security.record_failure_log(peer_addr.ip(), key_id, &error.failure.code); + if decision.suppressed > 0 { + tracing::warn!( + event = "auth_failures_suppressed", + peer_ip = %peer_addr.ip(), + key_id = key_id.as_u64(), + reason = %error.failure.code, + suppressed = decision.suppressed, + "suppressed repeated authentication failures in the previous window" + ); + } + if decision.emit { + tracing::warn!( + event = "auth_failed", + auth_stage = "initial_frame", + conn_id = %conn_id, + peer_addr = %peer_addr, + key_id = key_id.as_u64(), + reason = %error.failure.code, + retryable = error.failure.retryable, + error = %error.failure.message, + "connection authentication failed" + ); + } + if let Some(session) = error.response_session { + write_protocol_error(&mut conn, &session, &error.failure).await; + } + return Ok(()); + } + Ok(Ok(initial)) => initial, + }; + let replay_fingerprint = initial.replay_fingerprint; + let client_timestamp = initial.client_timestamp; + let init_request = match PbConnRequest::decode(&initial.payload) { + Ok(request) => request, + Err(error) => { + tracing::warn!( + event = "auth_failed", + auth_stage = "request_decode", + conn_id = %conn_id, + peer_addr = %peer_addr, + error = %error, + "authenticated request could not be decoded" + ); + write_protocol_error( + &mut conn, + &initial.session, + &pb_mapper_auth::AuthFailure::new( + "request_decode_failed", + "authenticated request payload is malformed", + false, + ), + ) + .await; + return Ok(()); + } + }; + let mut requested_namespace = None; + let mut force_register_namespace = false; + let init_request = match init_request { + PbConnRequest::RegisterScoped { + need_codec, + is_datagram, + key, + namespace, + force_namespace, + protocol_version, + client_instance_id, + heartbeat_interval_ms, + heartbeat_tolerance_ms, + } => { + requested_namespace = Some(namespace); + force_register_namespace = force_namespace; + PbConnRequest::Register { + need_codec, + is_datagram, + key, + protocol_version, + client_instance_id, + heartbeat_interval_ms, + heartbeat_tolerance_ms, + } + } + PbConnRequest::SubcribeScoped { key, namespace } => { + requested_namespace = Some(namespace); + PbConnRequest::Subcribe { key } + } + PbConnRequest::StatusScoped { status, namespace } => { + requested_namespace = Some(namespace); + PbConnRequest::Status(status) + } + PbConnRequest::StreamScoped { + key, + namespace, + dst_id, + server_generation, + } => { + requested_namespace = Some(namespace); + PbConnRequest::Stream { + key, + dst_id, + server_generation, + } + } + request => request, + }; + let session = initial.session; + let auth_context = match session.context() { + Ok(context) => context.clone(), + Err(error) => { + tracing::warn!(conn_id = %conn_id, peer_addr = %peer_addr, %error, "missing auth context"); + return Ok(()); + } + }; + tracing::info!( + event = "auth_succeeded", + auth_stage = "session", + conn_id = %conn_id, + peer_addr = %peer_addr, + key_id = auth_context.key_id.as_u64(), + namespace = auth_context.namespace, + protocol = ?session.protocol(), + is_admin = auth_context.is_admin, + "connection authentication succeeded" + ); + let effective_namespace = match resolve_namespace( + &auth_context, + requested_namespace, + force_register_namespace, + matches!(&init_request, PbConnRequest::Register { .. }), + ) { + Ok(namespace) => namespace, + Err(failure) => { + write_protocol_error(&mut conn, &session, &failure).await; + return Ok(()); + } + }; + match init_request { + PbConnRequest::Register { + key, + need_codec, + is_datagram, + protocol_version, + client_instance_id, + heartbeat_interval_ms, + heartbeat_tolerance_ms, + } => { + let protocol_version = protocol_version.unwrap_or(1); + tracing::info!( + event = "init_request", + request = "register", + conn_id = %conn_id, + peer_addr = %peer_addr, + key = %key, + protocol_version, + client_instance_id = ?client_instance_id, + heartbeat_interval_ms = ?heartbeat_interval_ms, + heartbeat_tolerance_ms = ?heartbeat_tolerance_ms, + need_codec, + is_datagram, + "received pb init request" + ); + let Some(key) = scope_service_or_reject( + &mut conn, + &session, + &auth_context, + effective_namespace, + &key, + ) + .await? + else { + return Ok(()); + }; + run_while_credential_active( + conn, + session, + &auth_context, + conn_id, + "registered service connection", + |conn, session| { + handle_server_conn( + ServerRegistration { + key, + need_codec, + is_datagram, + protocol_version, + conn_id, + }, + manager_task_sender, + conn, + session, + ) + }, + ) + .await?; + } + PbConnRequest::Subcribe { key } => { + tracing::info!( + event = "init_request", + request = "subscribe", + conn_id = %conn_id, + peer_addr = %peer_addr, + key = %key, + "received pb init request" + ); + let Some(key) = scope_service_or_reject( + &mut conn, + &session, + &auth_context, + effective_namespace, + &key, + ) + .await? + else { + return Ok(()); + }; + run_while_credential_active( + conn, + session, + &auth_context, + conn_id, + "subscribed data connection", + |conn, session| { + handle_client_conn(key, conn_id, manager_task_sender, conn, session) + }, + ) + .await?; + } + PbConnRequest::Stream { + key, + dst_id, + server_generation, + } => { + tracing::debug!( + event = "init_request", + request = "stream", + conn_id = %conn_id, + peer_addr = %peer_addr, + key = %key, + client_conn_id = dst_id, + server_generation, + "received pb init request" + ); + let Some(key) = scope_service_or_reject( + &mut conn, + &session, + &auth_context, + effective_namespace, + &key, + ) + .await? + else { + return Ok(()); + }; + manager_task_sender + .send(ManagerTask::Stream { + key: key.clone(), + stream: conn, + session, + server_id: conn_id, + client_id: dst_id.into(), + server_generation, + }) + .await + .map_err(|_| kanal::SendError(())) + .context(TaskCenterSendStreamRespToManagerSnafu { key, conn_id })?; + } + PbConnRequest::Status(status) => { + tracing::debug!( + event = "init_request", + request = "status", + conn_id = %conn_id, + peer_addr = %peer_addr, + status = ?status, + "received pb init request" + ); + run_while_credential_active( + conn, + session, + &auth_context, + conn_id, + "status request", + |conn, session| { + handle_show_status( + status, + effective_namespace, + manager_task_sender, + conn_id, + conn, + session, + ) + }, + ) + .await?; + } + PbConnRequest::Admin(request) => { + if !auth_context.is_admin { + write_protocol_error( + &mut conn, + &session, + &pb_mapper_auth::AuthFailure::new( + "admin_permission_required", + "administrator credential is required for this operation", + false, + ), + ) + .await; + return Ok(()); + } + if session.protocol() != HeaderProtocol::V2 { + write_protocol_error( + &mut conn, + &session, + &pb_mapper_auth::AuthFailure::new( + "admin_protocol_v2_required", + "administrator operations require protocol v2", + false, + ), + ) + .await; + return Ok(()); + } + if request.is_mutating() { + let Some(fingerprint) = replay_fingerprint else { + unreachable!("protocol-v2 sessions always carry a replay fingerprint"); + }; + let Some(client_timestamp) = client_timestamp else { + unreachable!("protocol-v2 sessions always carry a client timestamp"); + }; + if let Err(failure) = security + .auth() + .claim_admin_mutation(&auth_context, fingerprint, client_timestamp) + .await + { + write_protocol_error(&mut conn, &session, &failure).await; + return Ok(()); + } + } + handle_admin_request( + request, + auth_context, + security.auth().clone(), + manager_task_sender, + conn_id, + conn, + session, + ) + .await?; + } + PbConnRequest::RegisterScoped { .. } + | PbConnRequest::SubcribeScoped { .. } + | PbConnRequest::StatusScoped { .. } + | PbConnRequest::StreamScoped { .. } => unreachable!("scoped request was normalized"), + } + Ok(()) +} + +async fn scope_service_or_reject( + conn: &mut TcpStream, + session: &ServerHeaderSession, + auth_context: &AuthContext, + namespace: u64, + service_name: &str, +) -> Result> { + match scoped_service_key(auth_context, namespace, service_name) { + Ok(key) => Ok(Some(key)), + Err(failure) => { + write_protocol_error(conn, session, &failure).await; + Ok(None) + } + } +} + +async fn run_while_credential_active( + mut conn: TcpStream, + session: ServerHeaderSession, + auth_context: &AuthContext, + conn_id: RemoteConnId, + closed_what: &'static str, + work: F, +) -> Result<()> +where + F: FnOnce(TcpStream, ServerHeaderSession) -> Fut, + Fut: std::future::Future>, +{ + let cancellation = match auth_context.cancellation_token() { + Ok(token) => token, + Err(failure) => { + write_protocol_error(&mut conn, &session, &failure).await; + return Ok(()); + } + }; + tokio::select! { + result = work(conn, session) => result?, + _ = cancellation.cancelled() => { + tracing::info!( + event = "connection_auth_expired", + key_id = auth_context.key_id.as_u64(), + conn_id = %conn_id, + "closing {closed_what}" + ); + } + } + Ok(()) +} + +async fn write_protocol_error( + conn: &mut TcpStream, + session: &ServerHeaderSession, + failure: &pb_mapper_auth::AuthFailure, +) { + let response = PbConnResponse::error( + failure.code.clone(), + failure.message.clone(), + failure.retryable, + ); + let Ok(message) = response.encode() else { + return; + }; + let Ok(mut writer) = session.response_writer(conn) else { + return; + }; + if let Err(error) = writer.write_msg(&message).await { + tracing::debug!(%error, reason = %failure.code, "failed to write structured protocol error"); + } +} + +fn resolve_namespace( + context: &AuthContext, + requested: Option, + force_register_namespace: bool, + is_register: bool, +) -> std::result::Result { + let namespace = requested.unwrap_or(context.namespace); + if !context.is_admin && namespace != context.namespace { + return Err(pb_mapper_auth::AuthFailure::new( + "namespace_access_denied", + "temporary credentials can only access their own namespace", + false, + )); + } + if context.is_admin && is_register && namespace != 0 && !force_register_namespace { + return Err(pb_mapper_auth::AuthFailure::new( + "namespace_force_required", + "administrator registration in a temporary namespace requires --force", + false, + )); + } + Ok(namespace) +} + +fn scoped_service_key( + context: &AuthContext, + namespace: u64, + service_name: &str, +) -> std::result::Result { + if service_name.is_empty() || service_name.len() > 1024 || service_name.contains('\0') { + return Err(pb_mapper_auth::AuthFailure::new( + "service_name_invalid", + "service names must be 1-1024 bytes and must not contain NUL", + false, + )); + } + if !context.is_admin + && (service_name.len() > 128 + || !service_name + .bytes() + .all(|byte| byte.is_ascii_alphanumeric() || b"._:-".contains(&byte))) + { + return Err(pb_mapper_auth::AuthFailure::new( + "service_name_invalid", + "temporary-key service names must be 1-128 ASCII bytes from [A-Za-z0-9._:-]", + false, + )); + } + if namespace == 0 { + Ok(Arc::from(service_name)) + } else { + Ok(Arc::from(format!("@{namespace:016x}\u{0}{service_name}"))) + } +} + +pub(super) fn split_scoped_service_key(key: &str) -> (u64, &str) { + let Some((prefix, name)) = key.split_once('\0') else { + return (0, key); + }; + let Some(hex) = prefix.strip_prefix('@') else { + return (0, key); + }; + match u64::from_str_radix(hex, 16) { + Ok(namespace) => (namespace, name), + Err(_) => (0, key), + } +} + +pub(super) fn decrement_namespace_stream_count( + namespace_stream_counts: &mut hashbrown::HashMap, + namespace: u64, +) { + let Some(count) = namespace_stream_counts.get_mut(&namespace) else { + return; + }; + *count = count.saturating_sub(1); + if *count == 0 { + namespace_stream_counts.remove(&namespace); + } +} + +pub(super) fn release_namespace_rate_limit_if_idle( + namespace: u64, + server_conn_map: &ServerConnMap, + pending_streams: &hashbrown::HashMap, + namespace_rate_limits: &mut hashbrown::HashMap, +) { + let has_registered_service = server_conn_map + .keys() + .any(|key| split_scoped_service_key(key).0 == namespace); + let has_pending_stream = pending_streams + .values() + .any(|(_, _, key)| split_scoped_service_key(key).0 == namespace); + if !has_registered_service && !has_pending_stream { + namespace_rate_limits.remove(&namespace); + } +} diff --git a/src/pb_server/error.rs b/crates/pb-mapper-server/src/error.rs similarity index 85% rename from src/pb_server/error.rs rename to crates/pb-mapper-server/src/error.rs index dc896da..6149158 100644 --- a/src/pb_server/error.rs +++ b/crates/pb-mapper-server/src/error.rs @@ -2,18 +2,16 @@ use std::{sync::Arc, time::Duration}; use snafu::Snafu; -use crate::common::conn_id::RemoteConnId; -use crate::common::{self}; +use pb_mapper_core::conn_id::RemoteConnId; +// The `common::error::Error` spellings below are the source type on nearly every +// variant; aliasing keeps them as they were. +use pb_mapper_core as common; #[derive(Debug, Snafu)] #[snafu(visibility(pub(super)))] pub enum Error { - /// server task center error - #[snafu(display("read pb conn init request with `conn_id:{conn_id}`"))] - TaskCenterReadInitRequest { - conn_id: RemoteConnId, - source: common::error::Error, - }, + #[snafu(display("administrator operation failed: {detail}"))] + AdminOperation { detail: String }, #[snafu(display( "timed out reading pb conn init request with `conn_id:{conn_id}` after {timeout:?}" ))] @@ -21,11 +19,6 @@ pub enum Error { conn_id: RemoteConnId, timeout: Duration, }, - #[snafu(display("decode pb conn init request with `conn_id:{conn_id}`"))] - TaskCenterDecodeInitRequest { - conn_id: RemoteConnId, - source: common::error::Error, - }, #[snafu(display("send listener task error, type:{source:?} detail:{source}"))] TaskCenterSendListener { source: kanal::SendError<()> }, @@ -128,9 +121,7 @@ pub enum Error { conn_id: RemoteConnId, source: common::error::Error, }, - #[snafu(display( - "server conn write register resp error with `key:{key}` `conn_id:{conn_id}`" - ))] + #[snafu(display("server conn write register resp error with `key:{key}` `conn_id:{conn_id}`"))] ServerConnWriteRegisteredOk { key: Arc, conn_id: RemoteConnId, @@ -176,15 +167,6 @@ pub enum Error { conn_id: RemoteConnId, source: common::error::Error, }, - #[snafu(display( - "send deregister server task error with `key:{key}` `conn_id:{conn_id}`, \ - type:{source:?} detail:{source}" - ))] - ServerConnSendDeregisterServer { - key: Arc, - conn_id: RemoteConnId, - source: kanal::SendTimeoutError<()>, - }, #[snafu(display( "send register task error with `key:{key}` `conn_id:{conn_id}`, \ type:{source:?} detail:{source}" @@ -201,16 +183,8 @@ pub enum Error { tool: &'static str, source: common::error::Error, }, - #[snafu(display( - "send deregister client task error with `key:{key}` `server:{server_id:?}` <-> \ - `client:{client_id}`, type:{source:?} detail:{source}" - ))] - ClientConnSendDeregisterClient { - key: Arc, - server_id: Option, - client_id: RemoteConnId, - source: kanal::SendTimeoutError<()>, - }, + #[snafu(display("client data stream credential is inactive: {detail}"))] + ClientConnAuthInactive { detail: String }, #[snafu(display( "send subcribe task error with `key:{key}` `conn_id:{conn_id}`, type:{source:?} \ detail:{source}" @@ -288,9 +262,7 @@ pub enum Error { conn_id: RemoteConnId, source: common::error::Error, }, - #[snafu(display( - "client conn write subcribe resp error with `key:{key}` `conn_id:{conn_id}`" - ))] + #[snafu(display("client conn write subcribe resp error with `key:{key}` `conn_id:{conn_id}`"))] ClientConnWriteSubcribeResp { key: Arc, conn_id: RemoteConnId, @@ -321,14 +293,6 @@ pub enum Error { StatusEncodeResp { source: common::error::Error }, #[snafu(display("write status response error"))] StatusWriteResp { source: common::error::Error }, - #[snafu(display( - "send deregister request error with `conn_id:{conn_id}`, type:{source:?} \ - detail:{source}" - ))] - StatusSendDeregister { - conn_id: RemoteConnId, - source: kanal::SendTimeoutError<()>, - }, #[snafu(display("Server listen error"))] ServerListen { source: std::io::Error }, } diff --git a/crates/pb-mapper-server/src/lib.rs b/crates/pb-mapper-server/src/lib.rs new file mode 100644 index 0000000..9d4c9a5 --- /dev/null +++ b/crates/pb-mapper-server/src/lib.rs @@ -0,0 +1,386 @@ +//! Relay server domain model and module wiring. +//! +//! ```text +//! TCP listener -> connection authentication/dispatch -> ManagerTask queue +//! -> routing runtime +//! registered control connection <------ ConnTask ----+------> subscriber +//! ``` +//! +//! `connection` owns per-socket protocol/authentication concerns, while `runtime` +//! serializes global routing maps and quotas. Service-side and client-side tunnel loops +//! remain isolated in `server` and `client`. + +mod admin; +mod client; +mod error; +// Moved here from `common`: the routing runtime is its only caller. +pub mod manager; +mod server; +mod status; + +use std::net::SocketAddr; +use std::sync::Arc; +use std::time::Instant; + +use error::Result; +use snafu::{OptionExt, ResultExt}; +use tokio::net::{TcpListener, TcpStream, ToSocketAddrs}; +use tokio_util::sync::CancellationToken; +use tracing::instrument; + +use self::admin::handle_admin_request; +use self::client::handle_client_conn; +use self::error::{ + TaskCenterInitRequestTimeoutSnafu, TaskCenterSendListenerSnafu, TaskCenterSendStatusRespSnafu, + TaskCenterSendStreamRespToManagerSnafu, TaskCenterSetKeepAliveSnafu, +}; +use self::server::{ServerRegistration, handle_server_conn}; +use self::status::handle_show_status; +use crate::error::{ + ServerListenSnafu, TaskCenterClientSendStreamSnafu, TaskCenterSendRegisterRespSnafu, + TaskCenterSendStreamRespToClientSnafu, TaskCenterSendSubcribeRespSnafu, + TaskCenterStreamConnIdNotExistSnafu, +}; +use crate::manager::{ForwardMessage, SenderChan, TaskManager}; +use pb_mapper_auth::{ADMIN_KEY_ID, AuthConfig, AuthContext, AuthRuntime}; +use pb_mapper_core::config::{control_io_timeout, keep_alive_from_env, server_lease_timeout}; +use pb_mapper_core::conn_id::{ConnIdProvider, RemoteConnId}; +use pb_mapper_core::{snafu_error_get_or_continue, snafu_error_handle}; +use pb_mapper_protocol::MessageWriter; +use pb_mapper_protocol::command::{ + AdminConnectionInfo, AdminConnectionPage, AdminServiceInfo, AdminServicePage, + MessageSerializer, PbConnRequest, PbConnResponse, PbConnStatusReq, PbConnStatusResp, + PbServiceConnStatus, +}; +use pb_mapper_protocol::secure::{HeaderProtocol, ServerHeaderSession, ServerSecurity}; +use uni_stream::stream::{set_tcp_keep_alive, set_tcp_nodelay}; + +pub enum ManagerTask { + Accept { + stream: TcpStream, + peer_addr: SocketAddr, + }, + Register { + key: ImutableKey, + conn_id: RemoteConnId, + need_codec: bool, + is_datagram: bool, + protocol_version: u16, + conn_sender: ConnTaskSender, + }, + ServerConnActivity { + key: ImutableKey, + conn_id: RemoteConnId, + }, + Subcribe { + key: ImutableKey, + conn_id: RemoteConnId, + conn_sender: ConnTaskSender, + excluded_server_conns: Vec<(RemoteConnId, u64)>, + }, + Stream { + key: ImutableKey, + stream: TcpStream, + session: ServerHeaderSession, + server_id: RemoteConnId, + client_id: RemoteConnId, + server_generation: u64, + }, + StreamAck { + server_id: RemoteConnId, + client_id: RemoteConnId, + server_generation: u64, + }, + Status { + conn_sender: ConnTaskSender, + status: PbConnStatusReq, + namespace: u64, + conn_id: RemoteConnId, + }, + StatusQuery { + response_sender: tokio::sync::oneshot::Sender, + }, + AdminServiceList { + key_id: Option, + page: u32, + page_size: u16, + response_sender: tokio::sync::oneshot::Sender, + }, + AdminConnectionList { + key_id: Option, + page: u32, + page_size: u16, + response_sender: tokio::sync::oneshot::Sender, + }, + DeRegisterServerConn { + key: ImutableKey, + conn_id: RemoteConnId, + }, + RetireServerConn { + key: ImutableKey, + conn_id: RemoteConnId, + reason: String, + }, + DeRegisterClientConn { + server_id: Option, + client_id: RemoteConnId, + }, + Shutdown, +} + +#[derive(Debug)] +pub enum ConnTask { + Forward(ForwardMessage), + RegisterResp { + generation: u64, + protocol_version: u16, + lease_ttl_ms: u64, + }, + RegisterFailed { + code: String, + reason: String, + retryable: bool, + }, + SubcribeResp { + server_conn_id: RemoteConnId, + server_generation: u64, + need_codec: bool, + is_datagram: bool, + }, + SubcribeFailed { + code: String, + reason: String, + retryable: bool, + }, + SubcribeRetry { + reason: String, + }, + StreamReq { + client_id: RemoteConnId, + server_generation: u64, + }, + Retire { + reason: String, + }, + StreamAck { + server_id: RemoteConnId, + server_generation: u64, + }, + StreamResp { + server_id: RemoteConnId, + server_generation: u64, + stream: TcpStream, + session: ServerHeaderSession, + }, + StatusResp(PbConnResponse), +} + +pub(crate) type ManagerTaskSender = SenderChan; +pub(crate) type ConnTaskSender = SenderChan; + +pub type ImutableKey = Arc; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum ServerConnHealth { + Healthy, + Suspect, +} + +#[derive(Debug, Clone, Copy)] +pub struct ServerConnInfo { + pub conn_id: RemoteConnId, + pub generation: u64, + pub health: ServerConnHealth, + pub need_codec: bool, + pub is_datagram: bool, + pub protocol_version: u16, + pub last_rx_at: Instant, +} + +pub type ServerConnMap = hashbrown::HashMap>; + +struct NamespaceRateLimit { + tokens: f64, + last_refill: Instant, + rate_per_second: f64, + burst: f64, +} + +impl NamespaceRateLimit { + fn new(rate_per_second: usize, burst: usize) -> Self { + Self { + tokens: burst as f64, + last_refill: Instant::now(), + rate_per_second: rate_per_second as f64, + burst: burst as f64, + } + } + + fn allow(&mut self) -> bool { + let now = Instant::now(); + self.tokens = (self.tokens + + now.duration_since(self.last_refill).as_secs_f64() * self.rate_per_second) + .min(self.burst); + self.last_refill = now; + if self.tokens < 1.0 { + false + } else { + self.tokens -= 1.0; + true + } + } +} + +fn env_limit(name: &str, default: usize) -> usize { + std::env::var(name) + .ok() + .and_then(|value| value.parse::().ok()) + .filter(|value| *value > 0) + .unwrap_or(default) +} + +#[derive(Debug, Clone)] +pub struct ServerStatusInfo { + pub active_connections: u32, + pub registered_services: u32, + pub uptime_seconds: u64, +} + +fn remove_server_conn( + server_conn_map: &mut ServerConnMap, + key: &ImutableKey, + conn_id: RemoteConnId, +) -> bool { + if let Some(ids) = server_conn_map.get_mut(key) + && let Some(idx) = ids.iter().position(|info| info.conn_id == conn_id) + { + ids.remove(idx); + if ids.is_empty() { + server_conn_map.remove(key); + } + return true; + } + false +} + +fn registered_server_conn_count(server_conn_map: &ServerConnMap) -> usize { + server_conn_map.values().map(Vec::len).sum() +} + +fn service_conn_count(server_conn_map: &ServerConnMap, key: &ImutableKey) -> usize { + server_conn_map.get(key).map(Vec::len).unwrap_or_default() +} + +fn service_status_connections( + server_conn_map: &ServerConnMap, + key: &ImutableKey, +) -> Vec { + let now = Instant::now(); + server_conn_map + .get(key) + .map(|infos| { + infos + .iter() + .map(|info| PbServiceConnStatus { + conn_id: info.conn_id.into(), + generation: info.generation, + protocol_version: info.protocol_version, + healthy: info.health == ServerConnHealth::Healthy, + last_rx_age_ms: now.duration_since(info.last_rx_at).as_millis() as u64, + }) + .collect() + }) + .unwrap_or_default() +} + +fn record_server_conn_activity( + server_conn_map: &mut ServerConnMap, + key: &ImutableKey, + conn_id: RemoteConnId, +) -> bool { + let Some(infos) = server_conn_map.get_mut(key) else { + return false; + }; + let Some(info) = infos.iter_mut().find(|info| info.conn_id == conn_id) else { + return false; + }; + info.last_rx_at = Instant::now(); + info.health = ServerConnHealth::Healthy; + true +} + +fn record_server_conn_activity_by_conn_id( + server_conn_map: &mut ServerConnMap, + conn_id: RemoteConnId, +) -> bool { + let Some(info) = server_conn_map + .values_mut() + .flat_map(|infos| infos.iter_mut()) + .find(|info| info.conn_id == conn_id) + else { + return false; + }; + info.last_rx_at = Instant::now(); + info.health = ServerConnHealth::Healthy; + true +} + +async fn send_subcribe_failed( + conn_sender: &ConnTaskSender, + key: &ImutableKey, + conn_id: RemoteConnId, + reason: impl Into, +) { + let reason = reason.into(); + if conn_sender + .send(ConnTask::SubcribeFailed { + code: "service_not_available".to_string(), + reason: reason.clone(), + retryable: true, + }) + .await + .is_err() + { + tracing::debug!( + event = "subscribe_failure_receiver_dropped", + key = %key, + client_conn_id = %conn_id, + reason = %reason, + "subscribe failure receiver dropped" + ); + } +} + +async fn send_subcribe_retry( + conn_sender: &ConnTaskSender, + key: &ImutableKey, + conn_id: RemoteConnId, + reason: impl Into, +) { + let reason = reason.into(); + if conn_sender + .send(ConnTask::SubcribeRetry { + reason: reason.clone(), + }) + .await + .is_err() + { + tracing::debug!( + event = "subscribe_retry_receiver_dropped", + key = %key, + client_conn_id = %conn_id, + reason = %reason, + "subscribe retry receiver dropped" + ); + } +} + +mod runtime; +pub use runtime::{ + run_server, run_server_on_listener, run_server_with_auth_config, run_server_with_shutdown, +}; +mod connection; +use connection::{ + decrement_namespace_stream_count, handle_conn, handle_listener, + release_namespace_rate_limit_if_idle, split_scoped_service_key, +}; diff --git a/src/common/manager.rs b/crates/pb-mapper-server/src/manager.rs similarity index 91% rename from src/common/manager.rs rename to crates/pb-mapper-server/src/manager.rs index c6348d6..d94a34b 100644 --- a/src/common/manager.rs +++ b/crates/pb-mapper-server/src/manager.rs @@ -1,8 +1,17 @@ -use snafu::ResultExt; +use snafu::{ResultExt, Snafu}; use tracing::instrument; -use super::conn_id::{ConnId, ConnIdProvider, ConnIdTrait}; -use super::error::{MngWaitForTaskSnafu, Result}; +use pb_mapper_core::conn_id::{ConnId, ConnIdProvider, ConnIdTrait}; + +/// The manager owns this rather than the core error enum: it is the only thing +/// that waits on a task channel, and it keeps `kanal` out of the bottom layer. +#[derive(Debug, Snafu)] +#[snafu(display("`TaskManager` fails while waiting for a task"))] +pub struct MngWaitForTaskError { + source: kanal::ReceiveError, +} + +type Result = std::result::Result; /// The [`ConnId::local_id`] of the server is the same as the client. and it is only generated /// by the client. The [`ConnId::remote_id`] and [`ConnId::local_id`] of the client can be used to @@ -31,11 +40,11 @@ pub struct TaskManager, - > TaskManager + MangerChanType, + ConnChanType, + ConnIdType: ConnIdTrait, + ConnIdProviderType: ConnIdProvider, +> TaskManager { pub fn new( conn_id_provider: ConnIdProviderType, diff --git a/src/pb_server/mod.rs b/crates/pb-mapper-server/src/runtime.rs similarity index 55% rename from src/pb_server/mod.rs rename to crates/pb-mapper-server/src/runtime.rs index dd622cd..4efaee5 100644 --- a/src/pb_server/mod.rs +++ b/crates/pb-mapper-server/src/runtime.rs @@ -1,296 +1,17 @@ -mod client; -mod error; -mod server; -mod status; - -use std::net::SocketAddr; -use std::sync::Arc; -use std::time::Instant; - -use error::Result; -use snafu::{OptionExt, ResultExt}; -use tokio::net::{TcpListener, TcpStream, ToSocketAddrs}; -use tokio_util::sync::CancellationToken; -use tracing::instrument; - -use self::client::handle_client_conn; -use self::error::{ - TaskCenterDecodeInitRequestSnafu, TaskCenterInitRequestTimeoutSnafu, - TaskCenterReadInitRequestSnafu, TaskCenterSendListenerSnafu, TaskCenterSendStatusRespSnafu, - TaskCenterSendStreamRespToManagerSnafu, TaskCenterSetKeepAliveSnafu, -}; -use self::server::handle_server_conn; -use self::status::handle_show_status; -use crate::common::config::{control_io_timeout, keep_alive_from_env, server_lease_timeout}; -use crate::common::conn_id::{ConnIdProvider, RemoteConnId}; -use crate::common::manager::{ForwardMessage, SenderChan, TaskManager}; -use crate::common::message::command::{ - MessageSerializer, PbConnRequest, PbConnResponse, PbConnStatusReq, PbConnStatusResp, - PbServiceConnStatus, -}; -use crate::common::message::{get_header_msg_reader, MessageReader}; -use crate::pb_server::error::{ - ServerListenSnafu, TaskCenterClientSendStreamSnafu, TaskCenterSendRegisterRespSnafu, - TaskCenterSendStreamRespToClientSnafu, TaskCenterSendSubcribeRespSnafu, - TaskCenterStreamConnIdNotExistSnafu, -}; -use crate::{snafu_error_get_or_continue, snafu_error_handle}; -use uni_stream::stream::{set_tcp_keep_alive, set_tcp_nodelay}; - -pub enum ManagerTask { - Accept { - stream: TcpStream, - peer_addr: SocketAddr, - }, - Register { - key: ImutableKey, - conn_id: RemoteConnId, - need_codec: bool, - is_datagram: bool, - protocol_version: u16, - conn_sender: ConnTaskSender, - }, - ServerConnActivity { - key: ImutableKey, - conn_id: RemoteConnId, - }, - Subcribe { - key: ImutableKey, - conn_id: RemoteConnId, - conn_sender: ConnTaskSender, - excluded_server_conns: Vec<(RemoteConnId, u64)>, - }, - Stream { - stream: TcpStream, - server_id: RemoteConnId, - client_id: RemoteConnId, - server_generation: u64, - }, - StreamAck { - server_id: RemoteConnId, - client_id: RemoteConnId, - server_generation: u64, - }, - Status { - conn_sender: ConnTaskSender, - status: PbConnStatusReq, - conn_id: RemoteConnId, - }, - StatusQuery { - response_sender: tokio::sync::oneshot::Sender, - }, - DeRegisterServerConn { - key: ImutableKey, - conn_id: RemoteConnId, - }, - RetireServerConn { - key: ImutableKey, - conn_id: RemoteConnId, - reason: String, - }, - DeRegisterClientConn { - server_id: Option, - client_id: RemoteConnId, - }, - Shutdown, -} - -#[derive(Debug)] -pub enum ConnTask { - Forward(ForwardMessage), - RegisterResp { - generation: u64, - protocol_version: u16, - lease_ttl_ms: u64, - }, - SubcribeResp { - server_conn_id: RemoteConnId, - server_generation: u64, - need_codec: bool, - is_datagram: bool, - }, - SubcribeFailed { - reason: String, - }, - SubcribeRetry { - reason: String, - }, - StreamReq { - client_id: RemoteConnId, - server_generation: u64, - }, - Retire { - reason: String, - }, - StreamAck { - server_id: RemoteConnId, - server_generation: u64, - }, - StreamResp { - server_id: RemoteConnId, - server_generation: u64, - stream: TcpStream, - }, - StatusResp(PbConnResponse), -} - -pub(crate) type ManagerTaskSender = SenderChan; -pub(crate) type ConnTaskSender = SenderChan; - -pub type ImutableKey = Arc; - -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -pub enum ServerConnHealth { - Healthy, - Suspect, -} - -#[derive(Debug, Clone, Copy)] -pub struct ServerConnInfo { - pub conn_id: RemoteConnId, - pub generation: u64, - pub health: ServerConnHealth, - pub need_codec: bool, - pub is_datagram: bool, - pub protocol_version: u16, - pub last_rx_at: Instant, -} - -pub type ServerConnMap = hashbrown::HashMap>; - -#[derive(Debug, Clone)] -pub struct ServerStatusInfo { - pub active_connections: u32, - pub registered_services: u32, - pub uptime_seconds: u64, -} - -fn remove_server_conn( - server_conn_map: &mut ServerConnMap, - key: &ImutableKey, - conn_id: RemoteConnId, -) -> bool { - if let Some(ids) = server_conn_map.get_mut(key) { - if let Some(idx) = ids.iter().position(|info| info.conn_id == conn_id) { - ids.remove(idx); - if ids.is_empty() { - server_conn_map.remove(key); - } - return true; - } - } - false -} - -fn registered_server_conn_count(server_conn_map: &ServerConnMap) -> usize { - server_conn_map.values().map(Vec::len).sum() -} - -fn service_conn_count(server_conn_map: &ServerConnMap, key: &ImutableKey) -> usize { - server_conn_map.get(key).map(Vec::len).unwrap_or_default() -} - -fn service_status_connections( - server_conn_map: &ServerConnMap, - key: &ImutableKey, -) -> Vec { - let now = Instant::now(); - server_conn_map - .get(key) - .map(|infos| { - infos - .iter() - .map(|info| PbServiceConnStatus { - conn_id: info.conn_id.into(), - generation: info.generation, - protocol_version: info.protocol_version, - healthy: info.health == ServerConnHealth::Healthy, - last_rx_age_ms: now.duration_since(info.last_rx_at).as_millis() as u64, - }) - .collect() - }) - .unwrap_or_default() -} - -fn record_server_conn_activity( - server_conn_map: &mut ServerConnMap, - key: &ImutableKey, - conn_id: RemoteConnId, -) -> bool { - let Some(infos) = server_conn_map.get_mut(key) else { - return false; - }; - let Some(info) = infos.iter_mut().find(|info| info.conn_id == conn_id) else { - return false; - }; - info.last_rx_at = Instant::now(); - info.health = ServerConnHealth::Healthy; - true -} - -fn record_server_conn_activity_by_conn_id( - server_conn_map: &mut ServerConnMap, - conn_id: RemoteConnId, -) -> bool { - let Some(info) = server_conn_map - .values_mut() - .flat_map(|infos| infos.iter_mut()) - .find(|info| info.conn_id == conn_id) - else { - return false; - }; - info.last_rx_at = Instant::now(); - info.health = ServerConnHealth::Healthy; - true -} - -async fn send_subcribe_failed( - conn_sender: &ConnTaskSender, - key: &ImutableKey, - conn_id: RemoteConnId, - reason: impl Into, -) { - let reason = reason.into(); - if conn_sender - .send(ConnTask::SubcribeFailed { - reason: reason.clone(), - }) - .await - .is_err() - { - tracing::debug!( - event = "subscribe_failure_receiver_dropped", - key = %key, - client_conn_id = %conn_id, - reason = %reason, - "subscribe failure receiver dropped" - ); - } -} - -async fn send_subcribe_retry( - conn_sender: &ConnTaskSender, - key: &ImutableKey, - conn_id: RemoteConnId, - reason: impl Into, -) { - let reason = reason.into(); - if conn_sender - .send(ConnTask::SubcribeRetry { - reason: reason.clone(), - }) - .await - .is_err() - { - tracing::debug!( - event = "subscribe_retry_receiver_dropped", - key = %key, - client_conn_id = %conn_id, - reason = %reason, - "subscribe retry receiver dropped" - ); - } -} +//! Relay orchestration and the serialized routing-manager event loop. +//! +//! ```text +//! listener task ---- Accept -------+ +//! control tasks ---- Register -----+-> ManagerTask loop -> routing maps / quotas +//! subscriber ------- Subcribe -----+ -> ConnTask responses +//! provider stream -- Stream/Ack ---+ +//! ``` +//! +//! The manager loop is the single writer for connection IDs, registrations, pending +//! streams, per-namespace counts, and rate limits. Socket I/O runs in spawned connection +//! tasks and communicates with this state only through typed tasks. + +use super::*; struct RemoteIdProvider { next_id: RemoteConnId, @@ -333,13 +54,58 @@ pub async fn run_server_with_shutdown( >, keep_alive: bool, ) -> std::io::Result<()> { + run_server_with_auth_config( + addr, + shutdown_token, + status_channel, + keep_alive, + AuthConfig::default(), + ) + .await +} + +pub async fn run_server_with_auth_config( + addr: A, + shutdown_token: CancellationToken, + status_channel: Option< + tokio::sync::mpsc::UnboundedReceiver>, + >, + keep_alive: bool, + auth_config: AuthConfig, +) -> std::io::Result<()> { + let auth = AuthRuntime::from_process(auth_config) + .await + .map_err(|error| std::io::Error::other(error.to_string()))?; + let listener = TcpListener::bind(addr).await?; + run_server_on_listener(listener, shutdown_token, status_channel, keep_alive, auth).await +} + +pub async fn run_server_on_listener( + listener: TcpListener, + shutdown_token: CancellationToken, + status_channel: Option< + tokio::sync::mpsc::UnboundedReceiver>, + >, + keep_alive: bool, + auth: AuthRuntime, +) -> std::io::Result<()> { + let security = ServerSecurity::new(auth); let mut manager = ServerMananger::new(RemoteIdProvider::new()); // represent the mapping of the `key` to the id of the server-side conn let mut server_conn_map = ServerConnMap::new(); - let mut pending_streams = hashbrown::HashMap::::new(); + let mut pending_streams = + hashbrown::HashMap::::new(); + let mut namespace_stream_counts = hashbrown::HashMap::::new(); + let mut namespace_rate_limits = hashbrown::HashMap::::new(); + let max_services_per_namespace = env_limit("PB_MAPPER_MAX_SERVICES_PER_NAMESPACE", 256); + let max_register_connections_per_service = + env_limit("PB_MAPPER_MAX_REGISTER_CONNECTIONS_PER_SERVICE", 16); + let max_streams_per_namespace = env_limit("PB_MAPPER_MAX_STREAMS_PER_NAMESPACE", 1024); + let new_streams_per_second = env_limit("PB_MAPPER_NEW_STREAMS_PER_SECOND", 100); + let new_streams_burst = env_limit("PB_MAPPER_NEW_STREAMS_BURST", 200); let mut next_server_generation = 1_u64; + let mut connection_tasks = tokio::task::JoinSet::new(); - let listener = TcpListener::bind(addr).await?; let listen_addr = listener.local_addr()?; tracing::info!( event = "pb_server_listening", @@ -399,6 +165,97 @@ pub async fn run_server_with_shutdown( }; match task { + ManagerTask::AdminServiceList { + key_id, + page, + page_size, + response_sender, + } => { + let page_size = page_size.clamp(1, 1000) as usize; + let start = (page as usize).saturating_mul(page_size); + let mut all = server_conn_map + .iter() + .filter_map(|(key, connections)| { + let (namespace, service_name) = split_scoped_service_key(key); + if key_id.is_some_and(|key_id| key_id != namespace) { + return None; + } + let first = connections.first()?; + Some(AdminServiceInfo { + key_id: namespace, + namespace, + service_name: service_name.to_string(), + transport: if first.is_datagram { "udp" } else { "tcp" }.to_string(), + codec_enabled: first.need_codec, + connection_count: connections.len() as u32, + }) + }) + .collect::>(); + all.sort_by(|left, right| { + left.namespace + .cmp(&right.namespace) + .then_with(|| left.service_name.cmp(&right.service_name)) + }); + let items = all.iter().skip(start).take(page_size).cloned().collect(); + let next_page = + (start.saturating_add(page_size) < all.len()).then_some(page.saturating_add(1)); + let _ = response_sender.send(AdminServicePage { + schema_version: 1, + items, + next_page, + }); + } + ManagerTask::AdminConnectionList { + key_id, + page, + page_size, + response_sender, + } => { + let now = Instant::now(); + let page_size = page_size.clamp(1, 1000) as usize; + let start = (page as usize).saturating_mul(page_size); + let mut all = server_conn_map + .iter() + .flat_map(|(key, connections)| { + let (namespace, service_name) = split_scoped_service_key(key); + connections.iter().filter_map(move |connection| { + if key_id.is_some_and(|key_id| key_id != namespace) { + return None; + } + Some(AdminConnectionInfo { + key_id: namespace, + namespace, + service_name: service_name.to_string(), + conn_id: connection.conn_id.into(), + generation: connection.generation, + protocol_version: connection.protocol_version, + healthy: connection.health == ServerConnHealth::Healthy, + transport: if connection.is_datagram { "udp" } else { "tcp" } + .to_string(), + codec_enabled: connection.need_codec, + last_rx_age_ms: now + .duration_since(connection.last_rx_at) + .as_millis() + as u64, + }) + }) + }) + .collect::>(); + all.sort_by(|left, right| { + left.namespace + .cmp(&right.namespace) + .then_with(|| left.service_name.cmp(&right.service_name)) + .then_with(|| left.conn_id.cmp(&right.conn_id)) + }); + let items = all.iter().skip(start).take(page_size).cloned().collect(); + let next_page = + (start.saturating_add(page_size) < all.len()).then_some(page.saturating_add(1)); + let _ = response_sender.send(AdminConnectionPage { + schema_version: 1, + items, + next_page, + }); + } ManagerTask::StatusQuery { response_sender } => { let total_connections = server_conn_map .values() @@ -425,32 +282,66 @@ pub async fn run_server_with_shutdown( ManagerTask::Status { conn_sender, status, + namespace, conn_id, } => { let resp = match status { PbConnStatusReq::RemoteId => { + let scoped = server_conn_map + .iter() + .filter(|(key, _)| split_scoped_service_key(key).0 == namespace) + .map(|(key, value)| (split_scoped_service_key(key).1, value)) + .collect::>(); + let registered_ids = scoped + .iter() + .flat_map(|(_, connections)| { + connections.iter().map(|connection| connection.conn_id) + }) + .collect::>(); + let client_ids = pending_streams + .iter() + .filter_map(|(client_id, (_, _, key))| { + (split_scoped_service_key(key).0 == namespace).then_some(*client_id) + }) + .collect::>(); PbConnResponse::Status(PbConnStatusResp::RemoteId { - server_map: format!("{server_conn_map:?}"), - active: manager.active_conn_id_msg(), - idle: manager.idle_conn_id_msg(), + server_map: format!("{scoped:?}"), + active: format!( + "registered={registered_ids:?}, clients={client_ids:?}" + ), + idle: "namespace scoped; use `pb-mapper admin connection list` for global inspection" + .to_string(), }) } PbConnStatusReq::Keys => PbConnResponse::Status(PbConnStatusResp::Keys( - server_conn_map.keys().map(|k| k.to_string()).collect(), + server_conn_map + .keys() + .filter_map(|key| { + let (key_namespace, service_name) = split_scoped_service_key(key); + (key_namespace == namespace).then(|| service_name.to_string()) + }) + .collect(), )), PbConnStatusReq::Service { key } => { - let key: ImutableKey = key.into(); + let display_key = key.clone(); + let key: ImutableKey = if namespace == 0 { + key.into() + } else { + Arc::from(format!("@{namespace:016x}\u{0}{key}")) + }; PbConnResponse::Status(PbConnStatusResp::Service { - key: key.to_string(), + key: display_key, connections: service_status_connections(&server_conn_map, &key), }) } }; - snafu_error_get_or_continue!(conn_sender - .send(ConnTask::StatusResp(resp)) - .await - .map_err(|_| kanal::SendError(())) - .context(TaskCenterSendStatusRespSnafu { conn_id })); + snafu_error_get_or_continue!( + conn_sender + .send(ConnTask::StatusResp(resp)) + .await + .map_err(|_| kanal::SendError(())) + .context(TaskCenterSendStatusRespSnafu { conn_id }) + ); } ManagerTask::Accept { stream, peer_addr } => { let conn_id = manager.get_conn_id( @@ -469,9 +360,12 @@ pub async fn run_server_with_shutdown( "accepted pb connection" ); let manager_task_sender = manager.get_task_sender(); - tokio::spawn(async move { + let security = security.clone(); + while connection_tasks.try_join_next().is_some() {} + connection_tasks.spawn(async move { snafu_error_handle!( - handle_conn(conn_id, peer_addr, manager_task_sender, stream).await + handle_conn(conn_id, peer_addr, manager_task_sender, stream, security) + .await ); }); } @@ -479,7 +373,12 @@ pub async fn run_server_with_shutdown( let removed_from_service_map = remove_server_conn(&mut server_conn_map, &key, conn_id); let removed_from_active_map = manager.deregister_conn(conn_id); - pending_streams.retain(|_, (server_id, _)| *server_id != conn_id); + release_namespace_rate_limit_if_idle( + split_scoped_service_key(&key).0, + &server_conn_map, + &pending_streams, + &mut namespace_rate_limits, + ); tracing::info!( event = "server_conn_deregistered", key = %key, @@ -512,14 +411,12 @@ pub async fn run_server_with_shutdown( let removed_from_service_map = remove_server_conn(&mut server_conn_map, &key, conn_id); let removed_from_active_map = manager.deregister_conn(conn_id); - let mut removed_pending_streams = 0usize; - pending_streams.retain(|_, (server_id, _)| { - let keep = *server_id != conn_id; - if !keep { - removed_pending_streams += 1; - } - keep - }); + release_namespace_rate_limit_if_idle( + split_scoped_service_key(&key).0, + &server_conn_map, + &pending_streams, + &mut namespace_rate_limits, + ); let retire_notified = conn_sender .as_ref() .and_then(|sender| { @@ -537,7 +434,6 @@ pub async fn run_server_with_shutdown( reason = %reason, removed_from_service_map, removed_from_active_map, - removed_pending_streams, retire_notified, registered_services = server_conn_map.len(), server_connections = registered_server_conn_count(&server_conn_map), @@ -550,13 +446,25 @@ pub async fn run_server_with_shutdown( server_id, client_id, } => { - pending_streams.remove(&client_id); + let removed_namespace = pending_streams.remove(&client_id).map(|(_, _, key)| { + let namespace = split_scoped_service_key(&key).0; + decrement_namespace_stream_count(&mut namespace_stream_counts, namespace); + namespace + }); let removed_server_conn = if let Some(server_id) = server_id { manager.deregister_conn(server_id) } else { false }; let removed_client_conn = manager.deregister_conn(client_id); + if let Some(namespace) = removed_namespace { + release_namespace_rate_limit_if_idle( + namespace, + &server_conn_map, + &pending_streams, + &mut namespace_rate_limits, + ); + } if removed_server_conn || removed_client_conn { tracing::info!( event = "client_conn_deregistered", @@ -591,6 +499,51 @@ pub async fn run_server_with_shutdown( is_datagram, protocol_version, } => { + let namespace = split_scoped_service_key(&key).0; + let existing = server_conn_map.get(&key); + let failure = if existing.is_some_and(|connections| { + connections + .first() + .is_some_and(|connection| connection.is_datagram != is_datagram) + }) { + Some(( + "service_transport_mismatch", + "the service name is already registered with a different transport", + false, + )) + } else if existing.is_some_and(|connections| { + connections.len() >= max_register_connections_per_service + }) { + Some(( + "service_connection_limit_exceeded", + "the service has reached its register connection limit", + true, + )) + } else if existing.is_none() + && server_conn_map + .keys() + .filter(|registered| split_scoped_service_key(registered).0 == namespace) + .count() + >= max_services_per_namespace + { + Some(( + "namespace_service_limit_exceeded", + "the namespace has reached its service name limit", + true, + )) + } else { + None + }; + if let Some((code, reason, retryable)) = failure { + let _ = conn_sender + .send(ConnTask::RegisterFailed { + code: code.to_string(), + reason: reason.to_string(), + retryable, + }) + .await; + continue; + } let generation = next_server_generation; next_server_generation = next_server_generation.saturating_add(1).max(1); let now = Instant::now(); @@ -637,24 +590,28 @@ pub async fn run_server_with_shutdown( idle_connections = manager.idle_conn_count(), "server connection registered" ); - snafu_error_get_or_continue!(conn_sender - .send(ConnTask::RegisterResp { - generation, - protocol_version, - lease_ttl_ms: server_lease_timeout().as_millis() as u64, - }) - .await - .map_err(|_| kanal::SendError(())) - .context(TaskCenterSendRegisterRespSnafu { key, conn_id })); + snafu_error_get_or_continue!( + conn_sender + .send(ConnTask::RegisterResp { + generation, + protocol_version, + lease_ttl_ms: server_lease_timeout().as_millis() as u64, + }) + .await + .map_err(|_| kanal::SendError(())) + .context(TaskCenterSendRegisterRespSnafu { key, conn_id }) + ); } ManagerTask::Stream { + key, stream, + session, server_id, client_id, server_generation, } => { - let Some((expected_control_conn_id, expected_generation)) = - pending_streams.get(&client_id).copied() + let Some((expected_control_conn_id, expected_generation, expected_key)) = + pending_streams.get(&client_id).cloned() else { tracing::warn!( event = "stale_stream_without_pending_client", @@ -665,6 +622,17 @@ pub async fn run_server_with_shutdown( ); continue; }; + if key != expected_key { + tracing::warn!( + event = "stream_namespace_mismatch", + stream_conn_id = %server_id, + client_conn_id = %client_id, + expected_key = %expected_key, + actual_key = %key, + "dropping stream that does not belong to the pending namespace and service" + ); + continue; + } if server_generation != 0 && expected_generation != server_generation { tracing::warn!( event = "stale_stream_generation_mismatch", @@ -696,18 +664,23 @@ pub async fn run_server_with_shutdown( active_connections = manager.active_conn_count(), "server stream ready for client" ); - let client_sender = snafu_error_get_or_continue!(manager - .get_conn_sender_chan(&client_id) - .context(TaskCenterStreamConnIdNotExistSnafu { conn_id: client_id })); - snafu_error_handle!(client_sender - .send(ConnTask::StreamResp { - server_id, - server_generation: expected_generation, - stream - }) - .await - .map_err(|_| kanal::SendError(())) - .context(TaskCenterSendStreamRespToClientSnafu { conn_id: client_id })); + let client_sender = snafu_error_get_or_continue!( + manager + .get_conn_sender_chan(&client_id) + .context(TaskCenterStreamConnIdNotExistSnafu { conn_id: client_id }) + ); + snafu_error_handle!( + client_sender + .send(ConnTask::StreamResp { + server_id, + server_generation: expected_generation, + stream, + session, + }) + .await + .map_err(|_| kanal::SendError(())) + .context(TaskCenterSendStreamRespToClientSnafu { conn_id: client_id }) + ); } ManagerTask::StreamAck { server_id, @@ -716,8 +689,8 @@ pub async fn run_server_with_shutdown( } => { let recorded_activity = record_server_conn_activity_by_conn_id(&mut server_conn_map, server_id); - let Some((expected_server_id, expected_generation)) = - pending_streams.get(&client_id).copied() + let Some((expected_server_id, expected_generation, _)) = + pending_streams.get(&client_id).cloned() else { tracing::warn!( event = "stale_stream_ack_without_pending_client", @@ -749,17 +722,21 @@ pub async fn run_server_with_shutdown( { info.health = ServerConnHealth::Healthy; } - let client_sender = snafu_error_get_or_continue!(manager - .get_conn_sender_chan(&client_id) - .context(TaskCenterStreamConnIdNotExistSnafu { conn_id: client_id })); - snafu_error_handle!(client_sender - .send(ConnTask::StreamAck { - server_id, - server_generation, - }) - .await - .map_err(|_| kanal::SendError(())) - .context(TaskCenterSendStreamRespToClientSnafu { conn_id: client_id })); + let client_sender = snafu_error_get_or_continue!( + manager + .get_conn_sender_chan(&client_id) + .context(TaskCenterStreamConnIdNotExistSnafu { conn_id: client_id }) + ); + snafu_error_handle!( + client_sender + .send(ConnTask::StreamAck { + server_id, + server_generation, + }) + .await + .map_err(|_| kanal::SendError(())) + .context(TaskCenterSendStreamRespToClientSnafu { conn_id: client_id }) + ); } ManagerTask::Subcribe { key, @@ -767,6 +744,22 @@ pub async fn run_server_with_shutdown( conn_sender, excluded_server_conns, } => { + let namespace = split_scoped_service_key(&key).0; + if namespace_stream_counts + .get(&namespace) + .copied() + .unwrap_or_default() + >= max_streams_per_namespace + { + let _ = conn_sender + .send(ConnTask::SubcribeFailed { + code: "namespace_stream_limit_exceeded".to_string(), + reason: "the namespace has reached its active stream limit".to_string(), + retryable: true, + }) + .await; + continue; + } let Some(server_conn_id_list) = server_conn_map.get(&key).cloned() else { let reason = format!("server key `{key}` is not registered"); tracing::warn!( @@ -785,6 +778,22 @@ pub async fn run_server_with_shutdown( } continue; }; + if !namespace_rate_limits + .entry(namespace) + .or_insert_with(|| { + NamespaceRateLimit::new(new_streams_per_second, new_streams_burst) + }) + .allow() + { + let _ = conn_sender + .send(ConnTask::SubcribeFailed { + code: "namespace_stream_rate_exceeded".to_string(), + reason: "the namespace new-stream rate limit was exceeded".to_string(), + retryable: true, + }) + .await; + continue; + } let mut selected = false; let mut candidates = Vec::new(); candidates.extend(server_conn_id_list.iter().rev().copied().filter(|info| { @@ -845,7 +854,12 @@ pub async fn run_server_with_shutdown( if manager.get_conn_sender_chan(&conn_id).is_none() { manager.sign_up_conn_sender(conn_id, conn_sender.clone()); } - pending_streams.insert(conn_id, (server_conn_id, server_generation)); + let is_new_stream = pending_streams + .insert(conn_id, (server_conn_id, server_generation, key.clone())) + .is_none(); + if is_new_stream { + *namespace_stream_counts.entry(namespace).or_default() += 1; + } // 2. Response subcribe ok if let Err(e) = conn_sender .send(ConnTask::SubcribeResp { @@ -915,152 +929,86 @@ pub async fn run_server_with_shutdown( } } - // Gracefully shutdown the listener - listener_handle.abort(); - shutdown_handle.abort(); - if let Some(handle) = status_forward_handle { - handle.abort(); - } + // Abort first, then wait. Dropping a JoinHandle after abort() does not + // wait for the task to drop its AuthRuntime clone, so a UI restart can + // still see auth.lock held. + connection_tasks.abort_all(); + while connection_tasks.join_next().await.is_some() {} + abort_and_wait( + std::iter::once(listener_handle) + .chain(std::iter::once(shutdown_handle)) + .chain(status_forward_handle), + ) + .await; + security.auth().shutdown_actor().await; tracing::info!("Server shutdown completed"); Ok(()) } -async fn handle_listener( - task_sender: ManagerTaskSender, - listener: TcpListener, - keep_alive: bool, -) -> Result<()> { - loop { - let (stream, addr) = listener.accept().await.context(ServerListenSnafu)?; - tracing::debug!( - event = "tcp_conn_accepted", - peer_addr = %addr, - "accepted tcp connection" - ); - // set keepalive (optional) and nodelay - if keep_alive { - snafu_error_handle!(set_tcp_keep_alive(&stream).context(TaskCenterSetKeepAliveSnafu)); - } - snafu_error_handle!(set_tcp_nodelay(&stream), "remote stream set nodelay"); - task_sender - .send(ManagerTask::Accept { - stream, - peer_addr: addr, - }) - .await - .map_err(|_| kanal::SendError(())) - .context(TaskCenterSendListenerSnafu)? +async fn abort_and_wait(handles: impl IntoIterator>) { + let handles: Vec<_> = handles.into_iter().collect(); + for handle in &handles { + handle.abort(); + } + for handle in handles { + let _ = handle.await; } } -#[instrument(skip(manager_task_sender, conn), fields(conn_id = %conn_id, peer_addr = %peer_addr))] -async fn handle_conn( - conn_id: RemoteConnId, - peer_addr: SocketAddr, - manager_task_sender: ManagerTaskSender, - mut conn: TcpStream, -) -> Result<()> { - // handle by action - let init_request = get_init_request(&mut conn, conn_id).await?; - match init_request { - PbConnRequest::Register { - key, - need_codec, - is_datagram, - protocol_version, - client_instance_id, - heartbeat_interval_ms, - heartbeat_tolerance_ms, - } => { - let protocol_version = protocol_version.unwrap_or(1); - tracing::info!( - event = "init_request", - request = "register", - conn_id = %conn_id, - peer_addr = %peer_addr, - key = %key, - protocol_version, - client_instance_id = ?client_instance_id, - heartbeat_interval_ms = ?heartbeat_interval_ms, - heartbeat_tolerance_ms = ?heartbeat_tolerance_ms, - need_codec, - is_datagram, - "received pb init request" - ); - handle_server_conn( - key.into(), - need_codec, - is_datagram, - protocol_version, - conn_id, - manager_task_sender, - conn, - ) - .await?; - } - PbConnRequest::Subcribe { key } => { - tracing::info!( - event = "init_request", - request = "subscribe", - conn_id = %conn_id, - peer_addr = %peer_addr, - key = %key, - "received pb init request" - ); - handle_client_conn(key.into(), conn_id, manager_task_sender, conn).await?; - } - PbConnRequest::Stream { - key, - dst_id, - server_generation, - } => { - tracing::debug!( - event = "init_request", - request = "stream", - conn_id = %conn_id, - peer_addr = %peer_addr, - key = %key, - client_conn_id = dst_id, - server_generation, - "received pb init request" - ); - let key = ImutableKey::from(key); - manager_task_sender - .send(ManagerTask::Stream { - stream: conn, - server_id: conn_id, - client_id: dst_id.into(), - server_generation, - }) - .await - .map_err(|_| kanal::SendError(())) - .context(TaskCenterSendStreamRespToManagerSnafu { key, conn_id })?; - } - PbConnRequest::Status(status) => { - tracing::debug!( - event = "init_request", - request = "status", - conn_id = %conn_id, - peer_addr = %peer_addr, - status = ?status, - "received pb init request" - ); - handle_show_status(status, manager_task_sender, conn_id, conn).await?; +#[cfg(test)] +mod tests { + use super::*; + use pb_mapper_auth::{AuthConfig, AuthRuntime, LegacyProtocolPolicy}; + use pb_mapper_core::test_support::PROCESS_CREDENTIAL_TEST_LOCK; + use rand::RngExt; + use std::path::PathBuf; + use std::time::Duration; + + fn temp_state_dir(name: &str) -> PathBuf { + let mut suffix = [0_u8; 8]; + let mut rng = rand::rng(); + for byte in &mut suffix { + *byte = rng.random(); } + std::env::temp_dir().join(format!( + "pb-mapper-{name}-{}", + suffix + .iter() + .map(|byte| format!("{byte:02x}")) + .collect::() + )) } - Ok(()) -} -pub async fn get_init_request( - conn: &mut TcpStream, - conn_id: RemoteConnId, -) -> Result { - let mut reader = - get_header_msg_reader(conn).context(TaskCenterReadInitRequestSnafu { conn_id })?; - let timeout = control_io_timeout(); - let msg = match tokio::time::timeout(timeout, reader.read_msg()).await { - Ok(result) => result.context(TaskCenterReadInitRequestSnafu { conn_id })?, - Err(_) => TaskCenterInitRequestTimeoutSnafu { conn_id, timeout }.fail()?, - }; - PbConnRequest::decode(msg).context(TaskCenterDecodeInitRequestSnafu { conn_id }) + #[tokio::test] + async fn shutdown_releases_auth_lock_while_a_connection_is_open() { + let _process_credential_guard = PROCESS_CREDENTIAL_TEST_LOCK.lock().await; + let state_dir = temp_state_dir("shutdown-lock"); + let admin_key = *b"0123456789abcdefghijklmnopqrstuv"; + let config = AuthConfig { + state_dir: state_dir.clone(), + max_temporary_keys: 4, + max_temporary_key_ttl: Duration::from_secs(3600), + legacy_protocol: LegacyProtocolPolicy::Allow, + }; + let auth = AuthRuntime::start(admin_key, config.clone()).await.unwrap(); + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let shutdown_token = CancellationToken::new(); + let server = tokio::spawn({ + let shutdown_token = shutdown_token.clone(); + async move { run_server_on_listener(listener, shutdown_token, None, false, auth).await } + }); + let _client = TcpStream::connect(addr).await.unwrap(); + tokio::time::sleep(Duration::from_millis(80)).await; + shutdown_token.cancel(); + tokio::time::timeout(Duration::from_secs(2), server) + .await + .expect("server shutdown should finish after aborting connections") + .unwrap() + .unwrap(); + let restarted = AuthRuntime::start(admin_key, config).await.unwrap(); + drop(restarted); + tokio::time::sleep(Duration::from_millis(20)).await; + let _ = std::fs::remove_dir_all(state_dir); + } } diff --git a/src/pb_server/server.rs b/crates/pb-mapper-server/src/server.rs similarity index 90% rename from src/pb_server/server.rs rename to crates/pb-mapper-server/src/server.rs index f330e4f..3fb8d2a 100644 --- a/src/pb_server/server.rs +++ b/crates/pb-mapper-server/src/server.rs @@ -12,14 +12,13 @@ use super::error::{ ServerConnWriteRegisteredOkSnafu, ServerConnWriteStreamRequestSnafu, }; use super::{ConnTask, ImutableKey, ManagerTask, ManagerTaskSender, Result}; -use crate::common::config::server_lease_timeout; -use crate::common::conn_id::RemoteConnId; -use crate::common::message::command::{ - LocalServer, MessageSerializer, PbConnResponse, PbServerRequest, CONTROL_PROTOCOL_V2, -}; -use crate::common::message::{ - get_header_msg_reader, get_header_msg_writer, MessageReader, MessageWriter, +use pb_mapper_core::config::server_lease_timeout; +use pb_mapper_core::conn_id::RemoteConnId; +use pb_mapper_protocol::command::{ + CONTROL_PROTOCOL_V2, LocalServer, MessageSerializer, PbConnResponse, PbServerRequest, }; +use pb_mapper_protocol::secure::ServerHeaderSession; +use pb_mapper_protocol::{MessageReader, MessageWriter}; /// Ensure that server-side connections are properly deregistered before a normal connection is /// disconnected or an exception occurs @@ -132,18 +131,30 @@ enum ServerControlWrite { Pong(Vec), } +pub struct ServerRegistration { + pub key: ImutableKey, + pub need_codec: bool, + pub is_datagram: bool, + pub protocol_version: u16, + pub conn_id: RemoteConnId, +} + /// Maintaining a connection to the server. /// This connection is used to send channel request -#[instrument(skip(task_sender))] +#[instrument(skip(registration, task_sender, session))] pub async fn handle_server_conn( - key: ImutableKey, - need_codec: bool, - is_datagram: bool, - protocol_version: u16, - conn_id: RemoteConnId, + registration: ServerRegistration, task_sender: ManagerTaskSender, - conn: TcpStream, + mut conn: TcpStream, + session: ServerHeaderSession, ) -> Result<()> { + let ServerRegistration { + key, + need_codec, + is_datagram, + protocol_version, + conn_id, + } = registration; let (tx, rx) = kanal::bounded_async(DEFAULT_SERVER_CHAN_CAP); // register metadate @@ -187,6 +198,30 @@ pub async fn handle_server_conn( lease_ttl_ms, } = response else { + if let ConnTask::RegisterFailed { + code, + reason, + retryable, + } = response + { + let response = PbConnResponse::error(code, reason, retryable) + .encode() + .context(ServerConnEncodeRegisterRespSnafu { + key: key.clone(), + conn_id, + })?; + let mut writer = session + .response_writer(&mut conn) + .context(ServerConnCreateHeaderToolSnafu { tool: "writer" })?; + writer + .write_msg(&response) + .await + .context(ServerConnWriteRegisteredOkSnafu { + key: key.clone(), + conn_id, + })?; + return Ok(()); + } ServerConnRegisteredRespNotMatchSnafu { key: key.clone(), conn_id, @@ -203,7 +238,8 @@ pub async fn handle_server_conn( ); let (mut reader, mut writer) = conn.into_split(); - let mut msg_reader = get_header_msg_reader(&mut reader) + let mut msg_reader = session + .continuation_reader(&mut reader) .context(ServerConnCreateHeaderToolSnafu { tool: "reader" })?; // Keep one header writer for the register response and all later control frames. The // encrypted header codec is stateful; recreating it between frames breaks peer decoding. @@ -224,7 +260,8 @@ pub async fn handle_server_conn( let (write_tx, mut write_rx) = tokio::sync::mpsc::unbounded_channel::(); let writer_key = key.clone(); let mut writer_handle = tokio::spawn(async move { - let mut msg_writer = get_header_msg_writer(&mut writer) + let mut msg_writer = session + .response_writer(&mut writer) .context(ServerConnCreateHeaderToolSnafu { tool: "writer" })?; msg_writer.write_msg(®ister_response).await.context( ServerConnWriteRegisteredOkSnafu { @@ -542,10 +579,10 @@ mod tests { use tokio::time::Instant; - use crate::common::conn_id::RemoteConnId; - use crate::pb_server::ManagerTask; + use crate::ManagerTask; + use pb_mapper_core::conn_id::RemoteConnId; - use super::{ServerConnGuard, SERVER_TIMEOUT}; + use super::{SERVER_TIMEOUT, ServerConnGuard}; #[test] fn server_timeout_has_slack_over_local_server_ping_interval() { diff --git a/src/pb_server/status.rs b/crates/pb-mapper-server/src/status.rs similarity index 90% rename from src/pb_server/status.rs rename to crates/pb-mapper-server/src/status.rs index 789a7c3..e7372bd 100644 --- a/src/pb_server/status.rs +++ b/crates/pb-mapper-server/src/status.rs @@ -7,9 +7,10 @@ use super::error::{ StatusRecvConnTaskSnafu, StatusSendManagerTaskSnafu, StatusWriteRespSnafu, }; use super::{ConnTask, ManagerTask, ManagerTaskSender}; -use crate::common::conn_id::RemoteConnId; -use crate::common::message::command::{MessageSerializer, PbConnStatusReq}; -use crate::common::message::{get_header_msg_writer, MessageWriter}; +use pb_mapper_core::conn_id::RemoteConnId; +use pb_mapper_protocol::MessageWriter; +use pb_mapper_protocol::command::{MessageSerializer, PbConnStatusReq}; +use pb_mapper_protocol::secure::ServerHeaderSession; struct StatusConnGuard { conn_id: RemoteConnId, @@ -85,9 +86,11 @@ impl Drop for StatusConnGuard { pub async fn handle_show_status( status: PbConnStatusReq, + namespace: u64, manager_sender: ManagerTaskSender, conn_id: RemoteConnId, mut conn: TcpStream, + session: ServerHeaderSession, ) -> Result<()> { let info_span = info_span!("show status", "{status:?},{conn_id:?}"); let mut guard = StatusConnGuard::new(conn_id, manager_sender.clone()); @@ -97,6 +100,7 @@ pub async fn handle_show_status( let req = ManagerTask::Status { conn_sender: tx, status, + namespace, conn_id, }; manager_sender @@ -108,7 +112,8 @@ pub async fn handle_show_status( let resp = rx.recv().await.context(StatusRecvConnTaskSnafu)?; if let ConnTask::StatusResp(resp) = resp { let msg = resp.encode().context(StatusEncodeRespSnafu)?; - let mut msg_writer = get_header_msg_writer(&mut conn) + let mut msg_writer = session + .response_writer(&mut conn) .context(StatusCreateHeaderToolSnafu { tool: "writer" })?; msg_writer .write_msg(&msg) diff --git a/docker/docker-compose.yml b/docker/docker-compose.yml index c172cb2..699519f 100644 --- a/docker/docker-compose.yml +++ b/docker/docker-compose.yml @@ -5,8 +5,17 @@ services: image: ackingliu/pb-mapper:x86_64_musl environment: PB_MAPPER_PORT: 7666 + # Keep the machine-derived key on first boot of an empty auth volume so + # recreating this service does not mint a new administrator key. USE_MACHINE_MSG_HEADER_KEY: true RUST_LOG: error + volumes: + - pb-mapper-auth:/var/lib/pb-mapper/auth + - pb-mapper-legacy:/var/lib/pb-mapper-server ports: - "7666:7666" restart: unless-stopped + +volumes: + pb-mapper-auth: + pb-mapper-legacy: diff --git a/docker/pb-mapper.dockerfile b/docker/pb-mapper.dockerfile index 4e9f1cd..0cf2f6d 100644 --- a/docker/pb-mapper.dockerfile +++ b/docker/pb-mapper.dockerfile @@ -9,7 +9,8 @@ RUN chmod +x ./pb-mapper ./pb-mapper.sh ENV PB_MAPPER_PORT=7666 ENV USE_IPV6=false -ENV USE_MACHINE_MSG_HEADER_KEY=true +ENV USE_MACHINE_MSG_HEADER_KEY=false +VOLUME ["/var/lib/pb-mapper/auth"] EXPOSE $PB_MAPPER_PORT ENTRYPOINT [ "./pb-mapper.sh" ] diff --git a/docs/authentication-v2.md b/docs/authentication-v2.md new file mode 100644 index 0000000..900a0b2 --- /dev/null +++ b/docs/authentication-v2.md @@ -0,0 +1,348 @@ +# Authentication and Protocol V2 + +## Background and goals + +pb-mapper uses one public relay port for registration, subscription, status, and +administration. Version 0.4 keeps that transport model and adds two credential +levels without adding a TLS-style handshake: + +- one 32-byte administrator key owns the relay; +- renewable `pbmt1_` temporary credentials can inspect, register, and connect + only inside their own namespace; +- the first protocol-v2 frame authenticates and carries the request in one TCP + flight; +- revocation, expiry, and root-key rotation close affected live control and data + connections. + +TLS is still appropriate when endpoint identity, certificate trust, or traffic +analysis resistance is required. Protocol v2 protects pb-mapper frames with a +pre-shared credential; it is not a replacement for a public-key PKI. + +## Model and terminology + +| Term | Meaning | +| --- | --- | +| Administrator key | The sole 32-byte root credential. It can manage keys and inspect every namespace. | +| Temporary credential | A printable `pbmt1_...` value containing a key ID and derived 32-byte secret. | +| Key ID | A 64-bit `generation:u32 | slot:u32` identifier used for direct slot lookup. | +| Namespace | `0` for the administrator, otherwise the temporary key ID. | +| Service name | A user-facing name within one namespace. Equal names in different temporary namespaces do not collide. | +| Credential lease | The cancellation object shared by connections authenticated with one credential. | + +The administrator key is never copied into a temporary credential. A temporary +secret is derived with HKDF-SHA256 from the administrator key, the persistent +server instance ID, and the key ID. The fixed slot table stores lifecycle +metadata and a weak lease reference, not the temporary secret. + +## End-to-end architecture + +```mermaid +sequenceDiagram + participant C as register/connect/admin CLI + participant R as pb-mapper relay + participant A as auth actor + participant M as connection manager + + C->>R: V2 prefix + encrypted first request + R->>R: derive directional keys and authenticate frame + R->>A: validate key ID, generation, state, expiry + A-->>R: namespace + weak credential lease + alt administrator operation + R->>A: issue/renew/revoke/status + A-->>R: durable result after WAL fsync + else register/connect/status + R->>M: namespace-scoped service operation + M-->>R: scoped result or stable error + end + R-->>C: encrypted response on the same connection +``` + +Only the long-lived registration control connection and each independently +opened subscribe/data connection carry a V2 first frame. The relay does not add +an extra authentication exchange when the register process opens a data TCP +connection for a request. + +## Protocol-v2 framing + +### Initial prefix + +Every new client writes this 32-byte clear-text routing prefix: + +| Bytes | Field | +| ---: | --- | +| 4 | Magic `PBM2` | +| 1 | Version `2` | +| 1 | Flags, currently `0` | +| 2 | Reserved, currently `0` | +| 8 | Big-endian key ID; `0` means administrator | +| 16 | Connection salt: 8-byte Unix timestamp plus 8 random bytes | + +The prefix is not secret. It is authenticated as associated data on every +encrypted frame. Unsupported flags, versions, non-zero reserved bytes, and +timestamps outside the five-minute clock-skew window are rejected before +request dispatch. The encrypted first request is capped at 64 KiB before +authentication, while authenticated continuation frames retain the normal +protocol limit. + +### Directional frame keys + +HKDF-SHA256 uses the connection salt as salt and the credential's 32-byte secret +as input key material. Two independent outputs are expanded with +`pb-mapper-v2-c2s` and `pb-mapper-v2-s2c`. This prevents nonce reuse across +directions even though both directions begin with counter zero. + +### Encrypted frames + +Each frame is encoded as: + +| Bytes | Field | +| ---: | --- | +| 8 | Big-endian monotonically increasing counter | +| 4 | Big-endian ciphertext length, including the 16-byte GCM tag | +| variable | AES-256-GCM ciphertext and tag | + +The 96-bit AES-GCM nonce is four zero bytes followed by the 64-bit counter. AAD +contains the complete initial prefix, one direction byte, the counter, and the +ciphertext length. Counter mismatch, authentication failure, oversized frames, +and counter exhaustion close the connection. + +The first client request uses client-to-server counter `0`. The first response +uses server-to-client counter `0`. Later control frames continue from counter +`1` through one stateful reader/writer per direction. + +### Replay resistance + +The relay fingerprints `(key_id, connection_salt)` and atomically checks and +inserts it in two rotating 1 MiB Bloom filters covering the current and previous +600-second windows, so a max-future first-flight timestamp cannot outlive replay +retention. A probable duplicate is `connection_salt_replayed`. The relay does not encrypt +that error with the already-used first-response nonce, so the presenter may +see a decrypt/EOF instead of a readable frame. One-shot administrator CLI +operations retry once with a fresh salt only for a readable +`connection_salt_replayed` result or a failure before the request is written. +They do not resend after a dropped response, which would duplicate +non-idempotent commands such as `key issue`. Mutating administrator requests additionally claim their +exact fingerprint in the encrypted WAL before dispatch. Those claims survive +restart and compaction for ten minutes after the server accepted them, so an +old captured mutation cannot be replayed after the Bloom window or a process +restart. The client-supplied first-flight timestamp is still checked for +freshness, but it does not control how long the claim is retained. + +After a live root rotation or auth-state reset, the relay keeps the immediately +previous root key and instance id in memory. A first flight still encrypted with +the old temporary credential is decrypted with that previous material so the +client can read the stable `temporary_key_rotated` error instead of a decrypt +failure. + +## Credential lifecycle + +### Issuance and renewal + +`key issue` allocates a free fixed-table slot, increments its generation, +derives the secret, appends an encrypted WAL mutation, calls `fsync`, and only +then exposes the credential. `key renew` keeps the same credential and key ID, +updates its absolute expiry, and inserts a new versioned timing-wheel entry. +Stale wheel entries are ignored. + +```bash +export MSG_HEADER_KEY="$(sudo cat /var/lib/pb-mapper/auth/admin.key)" + +pb-mapper admin --server relay.example.com:7666 \ + key issue --ttl 24h --label home-web + +pb-mapper admin --server relay.example.com:7666 \ + key renew 4294967296 --ttl 7d +``` + +Temporary TTLs are at least 10 seconds and at most 30 days by default. The +server maximum is configurable. + +### Expiry, revocation, and garbage collection + +A four-level hierarchical timing wheel owns the strong `Arc` for every active +temporary lease. Foreground authentication state contains only `Weak` +references. Expiry or explicit revocation cancels the lease, immediately +causing authenticated control and data tasks to drop their TCP streams. +Tombstones remain briefly for stable diagnostics, then become reusable slots. + +```bash +pb-mapper admin --server relay.example.com:7666 key revoke 4294967296 +pb-mapper admin --server relay.example.com:7666 key gc +``` + +### Root rotation and state reset + +Root rotation writes an empty snapshot encrypted with the new key, preserves +the bounded audit history, persists `admin.key`, and then switches the key and +administrator lease as one state transition. It is a global invalidate: every +temporary credential stops authenticating, including unexpired keys in other +namespaces, and live connections using those keys are cancelled. A later +first flight that still decrypts under the previous root returns +`temporary_key_rotated` using that previous session, so the client can read +the structured error. A first flight that cannot decrypt — a mistyped, +foreign, or corrupted credential — fails as `protocol_v2_decrypt_failed` +without an encrypted error frame, because the relay cannot derive a session +the presenter can open. `temporary_key_invalid` is the in-process result when +presented material does not match after derivation. The CLI +stages the candidate key before the request and verifies the new key with an +authenticated status call. When `--key-file` is omitted, the recovery copy is +written below `$XDG_CONFIG_HOME/pb-mapper` (or `$HOME/.config/pb-mapper`) +rather than requiring local `/var/lib` access. + +An explicit auth-state reset also invalidates all temporary credentials. It +rotates the server instance ID so credentials from a corrupted or lost slot +table cannot become valid again if a key ID is later reused. + +## Namespace authorization + +Temporary credentials may perform `register`, `connect`, and `status` only in +their own namespace. They cannot issue keys, reveal credentials, inspect other +namespaces, alter protocol policy, reset auth state, or rotate the root key. + +The administrator defaults to namespace `0`. It may inspect or connect to a +temporary namespace with `--namespace `. Registering into another +namespace additionally requires `--force` to avoid accidental ownership +confusion. + +Temporary-key service names are 1-128 ASCII bytes from +`[A-Za-z0-9._:-]`. The relay enforces per-namespace caps for services, +registration connections, active streams, and new-stream rate. + +| Approach | Memory and lookup | Revocation | Namespace isolation | Wire cost | +| --- | --- | --- | --- | --- | +| Stateless signed token | Minimal server state | Requires a deny list | Token claim based | One request | +| Stateful hash map | Proportional allocations and hashing | Direct | Direct | One request | +| Fixed slots plus derived secrets | Fixed hot memory and O(1) lookup | Direct slot cancellation | Key ID is namespace | One request | + +The fixed-slot design deliberately accepts bounded server state to make early +revocation and hard connection closure deterministic. + +## Persistence and safe mode + +The Linux system-service state directory is `/var/lib/pb-mapper/auth`. +Unprivileged Linux, macOS, and Windows desktop binaries default to a +user-writable application directory instead of `/var/lib`: + +| File | Purpose | +| --- | --- | +| `admin.key` | Root credential, mode `0600` | +| `admin.key.next` | Staged root key for an in-flight rotation | +| `server-instance-id` | 16-byte persistent derivation identity | +| `server-instance-id.next` | Staged instance id for an in-flight reset | +| `auth.snapshot` | AES-256-GCM encrypted compact slot state | +| `auth.wal` | Length-prefixed, individually encrypted mutations and audit records | + +The directory is mode `0700`. Mutating operations acknowledge only after the +WAL record is synced. The actor compacts state every five minutes with an atomic +snapshot replacement and WAL truncation. The snapshot carries the bounded +audit history and active administrator replay claims, so compaction does not +discard either security record. + +The Flutter server uses its application config directory's `auth/` child and +does not report itself running until both the TCP listener and authentication +state have initialized successfully. This keeps desktop/mobile starts writable +without pretending that a failed `/var/lib` initialization succeeded. + +Invalid authentication-state headers, failed integrity checks, truncated WAL +records, schema mismatch, and failed compaction place temporary authentication +in safe mode. Administrator authentication stays available for inspection and +explicit reset; temporary authentication fails closed. + +## Administration and output contracts + +```bash +pb-mapper admin --server relay.example.com:7666 status +pb-mapper admin --server relay.example.com:7666 key list --page-size 100 +pb-mapper admin --server relay.example.com:7666 key show 4294967296 +pb-mapper admin --server relay.example.com:7666 key reveal 4294967296 +pb-mapper admin --server relay.example.com:7666 service list --key-id 4294967296 +pb-mapper admin --server relay.example.com:7666 connection list --all +pb-mapper admin --server relay.example.com:7666 legacy-protocol set deny +pb-mapper admin --server relay.example.com:7666 auth-state reset --confirm +pb-mapper admin --server relay.example.com:7666 root-key rotate +``` + +`--output human|json|ndjson` controls rendering. Pages default to 100 and are +capped at 1000. `--all` follows every page while preserving the selected output +format. NDJSON is the streaming choice for large inventories; JSON emits one +combined document and human output emits one combined table. + +Stable structured errors contain `code`, `message`, `retryable`, and +`server_time`. Authentication failure logs include stage, key ID, peer, and +reason but never credential material. Repeated failures are emitted five times +per minute per `(peer IP, key ID, reason)`, followed by a suppression summary. + +## Migration and compatibility + +New clients always emit protocol v2. A v0.4 server defaults to accepting legacy +framing so older clients can be upgraded without an outage. Operators can view +legacy connection counters, upgrade all clients, and then set the policy to +`deny`. Upgrade the relay before any client because v0.3 relays do not understand +the v2 first-frame magic. An explicitly configured `PB_MAPPER_LEGACY_PROTOCOL` +is trimmed and must be `allow` or `deny`; malformed values fail closed to `deny`. + +Fresh servers generate a random administrator key. Both the relay and install +scripts preserve an existing `/var/lib/pb-mapper-server/msg_header_key` by +copying it to the new `admin.key` path when no new key or environment credential +is configured. `--use-machine-msg-header-key` remains available only as an +explicit legacy compatibility option. + +Docker deployments must persist `/var/lib/pb-mapper/auth`; otherwise a recreated +container generates a different root key and cannot decrypt previous auth +state. + +## Operations playbook + +### Temporary credential rejected after renewal + +1. Run `pb-mapper admin status` and confirm `safe_mode=false`. +2. Run `key show ` and verify the key is `active` and its absolute expiry. +3. Check structured logs for `temporary_key_generation_mismatch`, + `temporary_key_expired`, or `protocol_v2_decrypt_failed`. +4. If the credential text was lost or copied incorrectly, run `key reveal ` + and replace the endpoint configuration. Renewal does not change the value. + +### Relay starts in safe mode + +1. Preserve the entire auth directory for diagnosis. +2. Confirm `admin.key`, `server-instance-id`, snapshot, and WAL belong to the + same server instance and were not partially restored. +3. Use `pb-mapper admin status`; administrator access remains available. +4. If recovery is impossible, run `auth-state reset --confirm`, then issue new + temporary credentials. Reset rotates the server instance ID and closes old + workloads. + +### Legacy clients stop connecting + +1. Check `admin status` for the current legacy policy and active legacy count. +2. If policy is `deny`, upgrade the client or temporarily set it to `allow`. +3. New clients should log protocol `V2`; a continuing legacy count identifies + an old binary or integration that still needs replacement. + +## Code index + +- Credential format and process configuration: + `crates/pb-mapper-core/src/checksum.rs` +- Authentication facade and shared model: `crates/pb-mapper-auth/src/lib.rs` +- Lifecycle actor: `crates/pb-mapper-auth/src/actor/` +- Persistence: `crates/pb-mapper-auth/src/persistence/` +- Runtime and timing wheel: `crates/pb-mapper-auth/src/runtime.rs` and + `crates/pb-mapper-auth/src/timing_wheel.rs` +- V2 session facade plus frame, limiter, and replay modules: + `crates/pb-mapper-protocol/src/secure.rs` and + `crates/pb-mapper-protocol/src/secure/` +- Relay state, runtime loop, and connection dispatch: + `crates/pb-mapper-server/src/lib.rs`, `runtime.rs`, and `connection.rs` +- Administrator request execution: `crates/pb-mapper-server/src/admin.rs` +- Unified CLI and administrator command module: + `crates/pb-mapper-cli/src/bin/pb-mapper.rs` and + `crates/pb-mapper-cli/src/bin/pb-mapper/admin.rs` + +## Summary + +Version 0.4 retains pb-mapper's one-port, long-lived-control-connection model +while separating root administration from scoped workload access. Temporary +keys are renewable but revocable, namespace collisions are eliminated, and +authentication remains part of the first request. The explicit operational +boundary is unchanged: protocol v2 is symmetric pre-shared-key security, while +TLS remains the layer for certificate-based endpoint identity. diff --git a/docs/authentication-v2.zh-CN.md b/docs/authentication-v2.zh-CN.md new file mode 100644 index 0000000..af9b2e2 --- /dev/null +++ b/docs/authentication-v2.zh-CN.md @@ -0,0 +1,272 @@ +# 认证体系与 V2 协议 + +## 背景与目标 + +pb-mapper 的注册、订阅、状态与管理流量共用一个公网端口。0.4 版本不改变这一 +连接模型,也不引入类似 TLS 的额外握手,而是在原有对称密钥体系内增加两级权限: + +- 一把 32 字节管理员密钥拥有中继的全部权限; +- 可续期、可提前吊销、自动过期的 `pbmt1_` 临时凭据只能查看、注册和连接自己的 + 命名空间; +- V2 第一个加密帧同时完成鉴权与请求传输,不增加一次网络往返; +- 过期、吊销与根密钥轮换会主动关闭受影响的控制连接和数据连接。 + +V2 是预共享密钥协议。当系统还需要证书身份、公开信任链或对流量分析的额外防护时, +TLS 仍然有独立价值,V2 不替代公钥 PKI。 + +## 核心模型 + +| 概念 | 含义 | +| --- | --- | +| 管理员密钥 | 唯一的 32 字节根凭据,可管理密钥并查看全部命名空间。 | +| 临时凭据 | `pbmt1_...` 字符串,携带 key ID 与派生后的 32 字节 secret。 | +| Key ID | 64 位 `generation:u32 | slot:u32`,用于直接定位固定槽位。 | +| 命名空间 | 管理员默认为 `0`;临时凭据的命名空间就是自己的 key ID。 | +| 凭据租约 | 同一凭据认证出的连接共同观察的取消对象。 | + +临时 secret 由管理员密钥、持久化 server instance ID 与 key ID 通过 +HKDF-SHA256 派生。服务端固定槽位只存生命周期元数据与弱引用,不存临时 secret。 + +## 端到端数据流 + +```mermaid +sequenceDiagram + participant C as register/connect/admin CLI + participant R as pb-mapper relay + participant A as auth actor + participant M as connection manager + + C->>R: V2 前缀 + 加密后的首个请求 + R->>R: 派生双向密钥并验证加密帧 + R->>A: 校验 key ID、generation、状态与过期时间 + A-->>R: 命名空间 + 凭据租约弱引用 + alt 管理操作 + R->>A: 签发/续期/吊销/状态查询 + A-->>R: WAL fsync 后的结果 + else 业务操作 + R->>M: 命名空间内的注册/订阅/查询 + M-->>R: 结果或稳定错误码 + end + R-->>C: 同一连接上的加密响应 +``` + +register 进程长期保持的控制连接只在建立时认证一次。后续每次业务请求拉起的新数据 +TCP 连接仍有自己的 V2 首帧,但不会在其上再做多轮鉴权交换。 + +## V2 帧结构 + +### 首帧前缀 + +新客户端先写入 32 字节明文路由前缀: + +| 字节数 | 字段 | +| ---: | --- | +| 4 | Magic `PBM2` | +| 1 | 版本 `2` | +| 1 | Flags,当前必须为 `0` | +| 2 | Reserved,当前必须为 `0` | +| 8 | 大端 key ID;`0` 表示管理员 | +| 16 | connection salt:8 字节 Unix 时间戳 + 8 字节随机数 | + +前缀不承担保密作用,但会作为每个加密帧的 AAD 被完整认证。未知版本、flags、 +reserved 值或超出五分钟时钟偏差窗口的时间戳会在请求分发前被拒绝。未认证的首个 +加密请求上限为 64 KiB;鉴权完成后的后续帧仍沿用正常协议上限。 + +### 双向密钥与计数器 + +HKDF-SHA256 使用 connection salt 作为 salt,凭据的 32 字节 secret 作为 IKM, +分别以 `pb-mapper-v2-c2s` 与 `pb-mapper-v2-s2c` 派生两个 AES-256-GCM 密钥。 +因此两个方向都从计数器 0 开始,也不会重复使用同一密钥与 nonce 组合。 + +每个加密帧由 8 字节大端计数器、4 字节密文长度、密文与 16 字节 GCM tag 组成。 +96 位 nonce 是四个零字节加 64 位计数器。AAD 包含完整首帧前缀、方向字节、计数器 +与密文长度。计数器不连续、认证失败、帧过大或计数器耗尽都会关闭连接。 + +首个请求使用 C2S counter 0,首个响应使用 S2C counter 0;后续控制消息由同一组 +有状态 reader/writer 从 counter 1 继续。 + +### 重放检测 + +服务端对 `(key_id, connection_salt)` 做指纹,并在同一个临界区内完成两个轮换的 +1 MiB Bloom filter 的检查与写入,覆盖当前与上一个 600 秒窗口,使首帧允许的 +最大未来时间戳无法在过滤器遗忘后继续重放。疑似重复是 +`connection_salt_replayed`。中继不会用已经用过的首响 nonce 加密该错误,所以 +对端可能看到解密失败/EOF 而不是可读错误帧。一次性 admin CLI 只在读到该错误、 +或请求尚未写出时换 salt 重试一次,不会在响应丢失后再发,以免 `key issue` +这类非幂等命令被执行两次。 +会修改状态的管理员请求还会在分发前把精确指纹写入加密 WAL;该记录自服务端接受起 +在十分钟内跨重启、跨 compact 保留,不能通过等待 Bloom 窗口结束或重启进程来重放 +旧操作。客户端首帧时间戳仍用于新鲜度检查,但不决定这条记录保留多久。 + +根密钥轮换或认证状态重置之后,进程会在内存里保留上一份根密钥和 instance id。 +仍用旧临时凭据加密的首帧会用这份材料解密,从而把稳定的 +`temporary_key_rotated` 错误返回给客户端,而不是解密失败。 + +## 临时凭据生命周期 + +### 签发与续期 + +签发会寻找空槽位、递增 generation、派生 secret、写入加密 WAL 并 `fsync`,随后才 +把凭据返回给管理员。续期不换 key ID 与凭据文本,只更新绝对过期时间并把新版本任务 +放入时间轮,旧任务到期时因版本不匹配而被忽略。 + +```bash +export MSG_HEADER_KEY="$(sudo cat /var/lib/pb-mapper/auth/admin.key)" +pb-mapper admin --server relay.example.com:7666 \ + key issue --ttl 24h --label home-web +pb-mapper admin --server relay.example.com:7666 \ + key renew 4294967296 --ttl 7d +``` + +默认最短 TTL 为 10 秒,最长为 30 天,服务端可调整最大值。 + +### 到期、吊销与 GC + +四层层级时间轮持有每个活动凭据租约的强 `Arc`;前台认证状态只持有 `Weak`。 +到期或管理员吊销会取消租约,相关控制任务和数据转发任务立即释放 TCP 连接。短暂保留 +tombstone 以给出稳定错误后,槽位可以复用。显式 `key gc` 可立即清理非活动槽位。 + +### 根密钥轮换与状态重置 + +根密钥轮换先用新密钥写空 snapshot,同时保留有上限的审计历史,持久化 +`admin.key`,再把密钥与管理员 lease 作为一次状态变更切换。这是全局作废:所有临时 +凭据都会立刻失效,包括其他命名空间里尚未过期的 key,并用这些 key 建立的活动连接 +会被取消。之后如果 first flight 仍能用上一轮根密钥解密,会得到客户端可读的 +`temporary_key_rotated`。无法解密的 first flight(输错、外站或损坏的凭据)则是 +`protocol_v2_decrypt_failed`,且不会附带加密错误帧,因为中继推导不出对端能打开 +的 session。`temporary_key_invalid` 只用于派生之后材料仍不匹配的进程内校验。 +CLI 在发请求前保存候选 key,完成后再用 +新 key 执行一次 `admin status` 验证。未指定 `--key-file` 时,恢复副本默认写到 +`$XDG_CONFIG_HOME/pb-mapper`(或 `$HOME/.config/pb-mapper`),不要求本机能写 +`/var/lib`。 + +`auth-state reset --confirm` 同样会清空临时凭据,并轮换 server instance ID。这样即使 +原槽位表损坏或丢失,旧凭据也不会因为未来复用了相同 key ID 而重新有效。 + +## 命名空间与权限边界 + +临时凭据只能在自己的命名空间执行 `register`、`connect` 与 `status`。它不能签发或 +查看其他 key、进入其他命名空间、修改 legacy 策略、重置状态或轮换管理员密钥。 + +管理员默认使用命名空间 0;通过 `--namespace ` 可以查看或连接临时命名空间。 +管理员要在临时命名空间内注册服务时还必须显式使用 `--force`,避免误把业务服务挂到 +错误租户。 + +临时凭据的 service name 限制为 1 到 128 个 ASCII 字节,字符集为 +`[A-Za-z0-9._:-]`。服务端分别按命名空间限制 service 数、单 service 注册连接数、 +活动 stream 数与新建 stream 速率。 + +| 方案 | 内存与查询 | 提前吊销 | 命名空间隔离 | 网络成本 | +| --- | --- | --- | --- | --- | +| 无状态签名 token | 服务端状态少 | 仍需 deny list | 依赖 token claim | 一个请求 | +| 通用 HashMap | 动态分配与哈希 | 直接删除 | 直接 | 一个请求 | +| 固定槽位 + 派生 secret | 固定热内存、O(1) 查找 | 直接取消槽位租约 | key ID 即 namespace | 一个请求 | + +当前方案明确接受有上限的服务端状态,以换取确定的提前吊销与活动连接硬关闭。 + +## 持久化与安全模式 + +Linux 系统服务默认目录是 `/var/lib/pb-mapper/auth`,权限为 `0700`。无特权 +Linux、macOS 与 Windows 桌面二进制默认写到用户可写的应用目录,而不是 +`/var/lib`: + +| 文件 | 用途 | +| --- | --- | +| `admin.key` | 根凭据,权限 `0600` | +| `admin.key.next` | 轮换进行中的暂存根密钥 | +| `server-instance-id` | 16 字节持久派生身份 | +| `server-instance-id.next` | reset 进行中的暂存实例 ID | +| `auth.snapshot` | AES-256-GCM 加密的紧凑槽位快照 | +| `auth.wal` | 带长度前缀、逐条加密的 mutation 与 audit | + +变更只有在 WAL 同步成功后才对外确认。后台 actor 每五分钟原子替换 snapshot 并截断 +WAL;snapshot 同时保存有上限的审计历史和仍有效的管理员重放声明,compact 不会丢弃 +这些安全记录。无效文件头、完整性验证失败、WAL 截断、schema 不匹配或 compact 失败 +都会进入 safe mode:临时凭据全部 fail closed,管理员仍可查看状态并执行显式 reset。 + +Flutter 启动服务端时使用应用配置目录下的 `auth/` 子目录,并且只有 TCP listener 与 +认证状态都初始化成功后才会报告 running;桌面和移动端无需写 `/var/lib`,初始化失败 +时也不会出现虚假的运行状态。 + +## 管理命令与输出 + +```bash +pb-mapper admin --server relay.example.com:7666 status +pb-mapper admin --server relay.example.com:7666 key list --page-size 100 +pb-mapper admin --server relay.example.com:7666 key show 4294967296 +pb-mapper admin --server relay.example.com:7666 key reveal 4294967296 +pb-mapper admin --server relay.example.com:7666 service list --key-id 4294967296 +pb-mapper admin --server relay.example.com:7666 connection list --all +pb-mapper admin --server relay.example.com:7666 legacy-protocol set deny +pb-mapper admin --server relay.example.com:7666 auth-state reset --confirm +pb-mapper admin --server relay.example.com:7666 root-key rotate +``` + +`--output human|json|ndjson` 控制展示格式。默认每页 100,最大 1000;`--all` 自动翻完 +所有页面并保留选定的输出格式。大列表应选择 NDJSON 流式输出;JSON 输出单个合并文档, +human 输出单个合并表格。稳定错误结构包含 `code`、`message`、 +`retryable` 与 `server_time`。 + +日志记录 auth stage、key ID、peer 与 reason,但不记录凭据。相同 +`(peer IP, key ID, reason)` 每分钟最多直接输出 5 次,下一窗口汇总被抑制的数量。 + +## 迁移与兼容性 + +新客户端固定发送 V2。0.4 服务端默认暂时接受旧帧,方便滚动升级;确认 +`active_legacy_connections` 归零后,可执行 `legacy-protocol set deny`。必须先升级中继、 +再升级客户端,因为 0.3 中继无法识别 V2 首帧 magic。显式配置的 +`PB_MAPPER_LEGACY_PROTOCOL` 会先去除首尾空白,并且只能是 `allow` 或 `deny`; +无效值会 fail closed 为 `deny`。 + +新安装会随机生成管理员密钥。中继自身与安装脚本在未配置新 key 或环境变量时,如果 +发现旧的 `/var/lib/pb-mapper-server/msg_header_key`,会将其复制到新路径,保留现有 +业务连接。`--use-machine-msg-header-key` 只作为明确的兼容选项继续存在。 + +Docker 必须持久化 `/var/lib/pb-mapper/auth`;否则重建容器会产生新管理员密钥,并且 +无法读取先前认证状态。 + +## 运维排障 + +### 续期后临时凭据仍被拒绝 + +1. 执行 `admin status`,确认 `safe_mode=false`。 +2. 执行 `key show `,确认状态为 active 并核对绝对过期时间。 +3. 从结构化日志区分 generation 不匹配、已过期与 V2 解密失败。 +4. 如果只是凭据文本复制错误,执行 `key reveal ` 重新配置;续期本身不会换凭据。 + +### 服务端进入 safe mode + +1. 先完整保留 auth 目录用于诊断。 +2. 确认 key、instance ID、snapshot 与 WAL 是否来自同一份服务器状态。 +3. 使用管理员 key 查询 `admin status`;管理员通道仍然可用。 +4. 无法恢复时执行 `auth-state reset --confirm`,再重新签发业务凭据。该操作会轮换 + instance ID,并断开旧业务。 + +### 禁用 legacy 后仍有旧客户端 + +1. 在 `admin status` 查看 legacy policy、当前连接数与最后连接时间。 +2. 如果业务尚未升级,可短暂改回 `allow`,但应尽快升级客户端。 +3. 新客户端的服务端日志应显示协议 `V2`;仍增长的 legacy 计数可以定位旧 binary。 + +## 代码索引 + +- 凭据格式与进程配置:`crates/pb-mapper-core/src/checksum.rs` +- 认证 facade 与共享模型:`crates/pb-mapper-auth/src/lib.rs` +- 生命周期 actor:`crates/pb-mapper-auth/src/actor/` +- 持久化:`crates/pb-mapper-auth/src/persistence/` +- runtime 与时间轮:`crates/pb-mapper-auth/src/runtime.rs`、 + `crates/pb-mapper-auth/src/timing_wheel.rs` +- V2 session facade、frame、限流与 replay 模块: + `crates/pb-mapper-protocol/src/secure.rs` 与 + `crates/pb-mapper-protocol/src/secure/` +- 中继状态、runtime loop 与连接分发:`crates/pb-mapper-server/src/lib.rs`、 + `runtime.rs`、`connection.rs` +- 管理请求执行:`crates/pb-mapper-server/src/admin.rs` +- 统一 CLI 与管理员命令模块:`crates/pb-mapper-cli/src/bin/pb-mapper.rs`、 + `crates/pb-mapper-cli/src/bin/pb-mapper/admin.rs` + +## 总结 + +0.4 在保持单端口与长控制连接模型的同时,把根管理权限与业务访问权限拆开。临时凭据 +可续期、可吊销、按命名空间隔离,鉴权仍然包含在第一个业务请求内。需要明确保留的 +边界是:V2 解决预共享密钥下的帧认证与权限控制,证书身份仍属于 TLS 层。 diff --git a/docs/pb-mapper-intro.zh-CN.md b/docs/pb-mapper-intro.zh-CN.md index e8c6606..0f3dc6d 100644 --- a/docs/pb-mapper-intro.zh-CN.md +++ b/docs/pb-mapper-intro.zh-CN.md @@ -45,7 +45,7 @@ flowchart LR ## 技术栈 -- **语言和运行时**:Rust 2021 + Tokio 异步运行时 +- **语言和运行时**:Rust 2024 edition + Tokio 异步运行时 - **内存分配器**:自己 fork 的 [`better_mimalloc_rs`](https://github.com/acking-you/better_mimalloc_rs),后面会细说为什么 - **网络抽象**:自研 [`uni-stream`](https://github.com/acking-you/uni-stream),把 TCP 和 UDP 统一成一套流接口;底层用 `socket2` 控制 socket 选项,`trust-dns-resolver` 做 DNS - **协议**:serde_json 序列化,自定义帧格式(checksum + 长度头),可选 `ring` 做 AES-256-GCM 端到端加密 @@ -87,11 +87,11 @@ VPS 能直连 GitHub 的话: curl -fsSL https://raw.githubusercontent.com/acking-you/pb-mapper/master/scripts/install-server-github.sh | bash ``` -装完默认监听 `7666`,开启 `--use-machine-msg-header-key`,密钥写到 `/var/lib/pb-mapper-server/msg_header_key`。client 侧 `export MSG_HEADER_KEY="$(cat /var/lib/pb-mapper-server/msg_header_key)"` 就能对上。 +装完默认监听 `7666`,首次启动会在 `/var/lib/pb-mapper/auth/admin.key` 生成随机管理员密钥。管理员用它签发带过期时间的 `pbmt1_` 临时凭据,再把临时凭据交给 register/connect 两端;不同临时凭据的同名 service 不会撞名。 ### 方式三:手动跑 CLI 或者用 Flutter UI -三个二进制,名字就是功能: +同一个 `pb-mapper` 二进制通过子命令切换功能: - `pb-mapper server`:公网中继 - `pb-mapper register`:跑在本地服务那一侧,把 `127.0.0.1:xxx` 注册成一个 service key @@ -102,6 +102,11 @@ curl -fsSL https://raw.githubusercontent.com/acking-you/pb-mapper/master/scripts ```bash # VPS 上 pb-mapper server --port 7666 +export MSG_HEADER_KEY="$(sudo cat /var/lib/pb-mapper/auth/admin.key)" +pb-mapper admin --server 127.0.0.1:7666 key issue --ttl 24h --label web + +# 家里和咖啡店都先导入上一步输出的 pbmt1_ 临时凭据 +export MSG_HEADER_KEY='' # 家里 pb-mapper register tcp --server :7666 --key web --addr 127.0.0.1:8080 diff --git a/docs/rust-async-send-sync-pin-deep-dive.md b/docs/rust-async-send-sync-pin-deep-dive.md index 3af9ca6..6a4aa5e 100644 --- a/docs/rust-async-send-sync-pin-deep-dive.md +++ b/docs/rust-async-send-sync-pin-deep-dive.md @@ -1,5 +1,9 @@ # Rust 并发安全(Send/Sync/Pin)与 async 状态机深度解析 +> 历史设计文档:反映 2025 年某时期的实现,其中的行号与代码引用可能已漂移 +> (拆分为多 crate 后所有路径都已移到 `crates/` 下),仅供设计意图参考。 +> 当前实现请以代码为准。 + 面向场景:你在实现网络转发(如 pb-mapper)时需要理解 **为什么某些 future 必须 `Send + 'static`**、为什么 `async fn` 能跨 `.await` 持有借用、以及 `Pin` 如何保证自引用安全。 > **Code Version**: pb-mapper 本地工作区(2026-01-17) diff --git a/docs/rust-shared-mutability-and-locks.zh-CN.md b/docs/rust-shared-mutability-and-locks.zh-CN.md new file mode 100644 index 0000000..2b83971 --- /dev/null +++ b/docs/rust-shared-mutability-and-locks.zh-CN.md @@ -0,0 +1,370 @@ +# 共享可变性、Arc 与锁:从时间轮重构中得到的判断顺序 + +面向场景:你在设计一个数据结构时,发现「好像得加个锁」,但不确定这个锁到底是需求还是自己造出来的。 + +> **Code Version**: pb-mapper 工作区,`feat/temporary-credential-auth`(2026-08-21) +> +> 相关文档:[Send/Sync/Pin 与 async 状态机深度解析](./rust-async-send-sync-pin-deep-dive.md) +> 覆盖 `async fn` 如何编译成状态机、`Pin` 为什么必要。本文讲的是它的另一面: +> **数据结构层面**该不该共享、该不该加锁。 + +## 1. 问题从哪里来 + +重构 `crates/pb-mapper-auth/src/timing_wheel.rs` 时,中间某一版长成这样: + +```rust +struct Queues { + levels: Vec>>>, // 每一级一把锁 +} +``` + +时间轮的每一级都上了一把锁。而这个时间轮是被 auth actor **独占**的——`Leases` +从头到尾以 `&mut Leases` 传递,从未进过 `Arc`。既然没有任何并发,这些锁是从哪来的? + +答案是我自己造出来的:我让 `Link::drop` 自己把下一跳投递进桶里。`Drop::drop` +只有 `&mut self`(指向 Link 自己),拿不到轮子的 `&mut`,所以只能让 Link 持有一个 +`Weak` 共享轮子——**一旦共享,就必须加锁**。 + +去掉共享,锁就自己消失了:让 `Drop` 只保留数据,由 `tick` 拿着 `&mut self` 投递。 + +这件事暴露出一个常见的思维捷径: + +> ❌ 「共享可变 → 加锁」 + +它跳过了两个更该先问的问题。本文把正确的判断顺序拆出来。 + +## 2. 三个正交的机制 + +先把三件经常被混为一谈的事分开。它们各管一件事,互不替代。 + +| 机制 | 管什么 | 典型工具 | +|---|---|---| +| **所有权 / 借用** | 谁能改、谁能读 | `&mut T` / `&T` | +| **`Arc`** | 生命周期:owner 何时死 | `Arc` / `Rc` | +| **`Send` / `Sync`** | 跨线程访问的安全性 | marker trait,编译器自动推导 | +| **内部可变性** | 通过 `&T` 修改 | `Cell` / `RefCell` / `Atomic` / `Mutex` | + +``` + 「我能改它吗」 「它还活着吗」 「换线程安全吗」 + │ │ │ + 借用规则 Arc/Rc Send/Sync + │ │ │ + └──── 三者独立,缺一不可,且不能互相代替 ────────┘ +``` + +一个具体的反直觉例子:`Vec` 本身就是 `Send + Sync`,但这**不代表**你可以 +不用 `Arc` 就把它共享给 `tokio::spawn` 的任务。`Send + Sync` 解决的是安全性, +`'static` 生命周期要求得靠 `Arc` 解决。反过来,`Arc>` 生命周期没问题, +但因为 `Cell` 不是 `Sync`,一样过不了 `spawn`。 + +## 3. Send 与 Sync 到底是什么 + +两个 marker trait,不含任何方法,编译器**自动推导**(结构体所有字段都满足则它满足): + +- **`Send`**:这个值可以**移动**到另一个线程。 +- **`Sync`**:这个值可以被**多个线程同时引用**。 + +第二条有个更精确、也更好用的等价定义: + +``` +T: Sync ⟺ &T: Send +``` + +即「把它的引用发给别的线程是否安全」。这比「线程安全」这种模糊说法准确得多。 + +### 3.1 为什么必须区分这两件事 + +看 `Cell`: + +```rust +let c = Cell::new(0); +thread::spawn(move || c.set(1)); // 独占移过去 —— 安全,故 Cell: Send +// 但: +thread::scope(|s| { + s.spawn(|| c.set(1)); // 两个线程同时 set + s.spawn(|| c.set(2)); // 数据竞争,故 Cell: !Sync +}); +``` + +`Cell::set` 只是一条普通写入,没有任何同步。**移动**给一个线程完全安全(原线程 +再也碰不到它);**同时借给两个线程**就是 UB。这恰好是 `Send` 与 `Sync` 的分界。 + +### 3.2 Sync 只赋予「共享读」,不赋予「共享写」 + +这是最容易混淆的一点。大多数类型(`u64`、`String`、`Vec`、`HashMap`)都是 +`Send + Sync`——但你拿到 `&Vec` 依然改不了它: + +```rust +let shared = Arc::new(vec![1_u8, 2, 3]); +shared.push(4); +// error[E0596]: cannot borrow data in an `Arc` as mutable +``` + +`Vec` 是 `Sync` 的**原因**正是「通过 `&T` 改不了它」。它 `Sync` 是因为它老实, +不是因为它做了同步。于是 `Sync` 的类型分成两类: + +| 类别 | 例子 | 为什么 Sync | +|---|---|---| +| **老实类型** | `u64`, `Vec`, `HashMap` | 通过 `&T` 根本改不了,无从竞争 | +| **内部可变 + 自带同步** | `Mutex`, `AtomicU64`, `RwLock` | 通过 `&T` 能改,但访问被串行化 | + +`Cell` 是第三类:通过 `&T` 能改,**且不带同步**——所以它被排除在 `Sync` 之外。 +三类划清,`Sync` 的定义就自洽了。 + +> 💡 **Key Point**:需要锁的条件不是「T 不是 Sync」,而是**「我要通过 `&T` 修改它」**。 + +### 3.3 Mutex 的类型学意义 + +``` +T: Send ──[ 包一层 Mutex ]──> Mutex: Sync +``` + +**锁的作用就是把「可移动」升级成「可共享」。** 所以 `Mutex>` 是 `Sync` +的——一个类型不是 `Sync` 从来不是死路。 + +## 4. 为什么 Rc 既不 Send 也不 Sync + +根因只有一个:**`Rc` 的引用计数是普通 `usize` 加减,非原子**。那正是它比 `Arc` 快的 +全部原因。但这一个根因导致两个独立后果: + +**`!Sync`**(直觉的那一半):`Rc::clone` 只要 `&self`,而它会 `count += 1`。两个线程 +各持 `&Rc` 同时 clone,两次非原子递增丢一次计数 → 提前释放 → use-after-free。 + +**`!Send`**(更微妙,也更关键):你可能想「整个移过去不就独占了吗」。但 +**`Rc` 从来不是唯一的那一份**: + +```rust +let here = Rc::new(0); +let there = here.clone(); // 两个 handle,一个共享的非原子计数 +thread::spawn(move || drop(there)); // 那边 count -= 1 +drop(here); // 这边 count -= 1,同时进行 +``` + +`there` 移走了,`here` 还在原线程。两边同时递减同一个非原子计数 → 泄漏或双重释放。 + +> 💡 **Key Point**:`Rc: Send` 不安全,不是因为被移动的那一份,而是因为 +> **留在原地的那些**。类型系统无法表达「仅当这是最后一份 handle 时才允许移动」, +> 所以只能整个禁掉。 + +对比 `Cell` 就完整了: + +| | 计数 | Send | Sync | +|---|---|---|---| +| `Rc` | 非原子 | ✗ 别的 handle 会同时改计数 | ✗ `&Rc` 就能 clone | +| `Arc` | 原子 | ✓(需 `T: Send + Sync`)| ✓(同)| +| `Cell` | — | ✓ 移走后原线程什么都不剩 | ✗ `&Cell` 就能 set | + +`Cell: Send` 而 `Rc: !Send`,差别正在于「移走后原线程手里还有没有东西」。 + +## 5. Arc 管的是生命周期,不是安全性 + +既然 `Vec` 本身就 `Send + Sync`,为什么还需要 `Arc`?直接 `&Vec` 不行吗? + +**在能证明作用域的场景里,确实不需要**: + +```rust +let data = vec![1_u8, 2, 3]; +thread::scope(|s| { + s.spawn(|| println!("{:?}", &data)); // 零 Arc + s.spawn(|| println!("{:?}", &data)); +}); +println!("still owned: {:?}", data); // 依然是 owner +``` + +`thread::scope` 保证所有子线程在它返回前 join,所以编译器**能证明** `data` 活得更久。 + +`tokio::spawn` 和 `thread::spawn` 则不然——任务是 detached 的,何时结束由运行时决定。 +编译器无法证明任何栈上的东西活得比它久,于是有了 `'static` 约束。满足它只有两条路: + +1. 把所有权 `move` 进去 → 只有一个任务能拿到,没法共享; +2. 用 `Arc` → 所有权归引用计数集体所有,**没有任何栈帧是它的 owner**, + 于是每个持有者天然满足 `'static`。 + +``` + 能证明作用域 不能(detached / 'static) + 只读 &T(零成本) Arc + 要改 &mut T(零成本) Arc> +``` + +> 💡 **Key Point**:`Arc` 把生命周期问题转成运行时引用计数,代价是一次原子加减。 +> 这跟 `T` 是否 `Sync` 无关——`Sync` 只决定 `Arc` 能不能 `Send`。 + +## 6. tokio::spawn 的签名从哪来 + +```rust +pub fn spawn(future: F) -> JoinHandle +where F: Future + Send + 'static, F::Output: Send + 'static +``` + +三个约束各有其因: + +- **`Send`**:tokio 多线程调度器有 work-stealing——空闲 worker 会从别的 worker + 队列里偷任务。你的 future 可能在线程 A 上 poll 一次、挂起、然后在线程 B 上 poll + 下一次。它是被**移动**过去的,故需 `Send`。 +- **`'static`**:future 存活时间由运行时决定,不受调用处作用域约束,不能借用栈上的东西。 +- **`F::Output: Send`**:结果要从 worker 线程送回 `JoinHandle` 的等待方。 + +对 `async fn`,编译器把它编译成状态机,**所有跨 `.await` 存活的局部变量都成为该状态机 +的字段**。于是「future 是否 `Send`」= 「这些字段是否全部 `Send`」。这就是为什么一个 +持有 `Rc` 的 async 函数无法 `spawn`——哪怕只在两个 `.await` 之间用了一下。 + +(状态机的展开细节见 +[Send/Sync/Pin 与 async 状态机深度解析 §3](./rust-async-send-sync-pin-deep-dive.md)。) + +逃逸口:`spawn_local` 没有 `Send` 约束,代价是任务被钉在单线程 `LocalSet` 上, +拿不到 work-stealing。 + +## 7. 判断顺序 + +把上面几节合起来,得到一个可执行的检查表。**顺序很重要**——跳过前两问就会得到 +第 1 节那种每级一把锁的东西。 + +``` +① 这里为什么是共享的?能不能改成独占? + │ + ├─ 能 ──> 用 &mut T,到此结束(零成本,无锁) + │ + ↓ 不能 +② 共享是否真的要跨线程? + │ + ├─ 不必 ──> Cell / RefCell(零成本 / 一个计数器) + │ + ↓ 必须(Send/Sync 约束逼上来了) +③ 要改的是什么粒度? + │ + ├─ 单个整数/指针 ──> Atomic(无锁) + │ + └─ 一段临界区 ────> Mutex / RwLock +``` + +第 ① 步最容易被跳过,而它恰恰是收益最大的一步:**共享是可以被设计掉的**。 + +## 8. 落到 pb-mapper 的真实代码 + +### 8.1 被设计掉的锁:时间轮的 Queues + +第 1 节那版每级一把 `Mutex`,走的是「`Drop` 里投递 → 需要共享 `Queues` → 加锁」。 +现在 `Link::Relay` 只是纯数据,`tick` 拿 `&mut self` 自己投递 +(`crates/pb-mapper-auth/src/timing_wheel.rs`): + +```rust +Link::Relay { level, slot, next } => self.file(level as usize, slot as usize, *next), +``` + +停在第 ① 步。`Queues` 类型、`Weak`、以及那一排锁全部消失。 + +### 8.2 无锁的共享可变:AuthLease.expires_at + +```rust +pub struct AuthLease { + expires_at: AtomicU64, // crates/pb-mapper-auth/src/lib.rs + ... +} +``` + +lease 通过 `Arc` 共享(请求侧持 `Weak`,时间轮持强引用),续期要改 `expires_at`, +所以是货真价实的「跨线程共享可变」。但它只读写一个 `u64`,停在第 ③ 步的 `Atomic` 分支, +不需要锁。 + +### 8.3 无法避免的锁:Timer.callback + +```rust +pub(super) struct Timer { + callback: Mutex>>, +} +``` + +这把锁走完了全部三步,每一步都无路可退: + +``` +① 能独占吗? 不能 —— 续期时 retire/reap 两条路径共享同一个 Timer, + 且「最后一个引用被丢弃时触发」这个语义本身就是引用计数 +② 能只单线程吗? 不能 —— 推导链如下 +③ 能用 Atomic 吗? 不能 —— FnOnce 只能按值调用,必须把整个 Box 移出来 +``` + +第 ② 步的推导链值得完整写出来,它是本文所有概念的汇合点: + +``` +tokio::spawn(run_auth_actor(...)) 要求 future: Send + → actor future 跨 .await 持有 Leases ⇒ Leases: Send + → Leases 持有 Arc ⇒ Arc: Send + → Arc: Send 需要 T: Send + Sync ⇒ Timer: Sync + → Timer: Sync 需要 callback 字段: Sync + → Cell 不是 Sync,Mutex 是(当 T: Send) +``` + +(`Arc: Send` 为何要 `T: Sync`:clone 出去后另一个线程通过它拿到 `&T`, +那正是 `Sync` 管的事;同时也要 `T: Send`,因为那个线程可能持有最后一份引用 +并在自己那里析构 `T`。) + +不过实际开销接近零:**绝大多数 timer 从不加锁触发**。`Drop for Timer` 有 +`&mut self`,走 `Mutex::get_mut()`: + +```rust +impl Drop for Timer { + fn drop(&mut self) { + let callback = self.callback.get_mut().take(); // 无锁 + run(callback); + } +} +``` + +只有显式提前 `fire()`(revoke、GC)才真正 lock,而那时也无人竞争。 + +### 8.4 停在第一步:Leases.stages + +```rust +pub(super) struct Leases { + stages: HashMap, // 没有 Arc,没有锁 +} +``` + +`HashMap` 是 `Send + Sync`,但这在这里毫不相关——actor 独占 `Leases`,所有方法都是 +`&mut self`。**`HashMap` 是不是 `Sync` 根本不影响这个决定。** + +### 8.5 必要的锁:AuthStateInner + +```rust +struct AuthStateInner { + slots: RwLock>, + cold: RwLock>, + ... +} +``` + +它被 `Arc` 共享给请求处理路径和 actor 两边,双方都要改,且要保护的是「查槽位 → +校验 generation → 改状态」这样的临界区而非单个整数。三步走完,`RwLock` 是对的。 + +## 9. 为什么用 parking_lot + +标准库的 `Mutex`/`RwLock` 带 **poisoning**:持锁线程 panic 后,锁被标记为「有毒」, +后续 `lock()` 返回 `Err`。本项目从不利用这个信号——迁移前每个调用点都是同一句样板: + +```rust +lock.read().unwrap_or_else(|poisoned| poisoned.into_inner()) +``` + +即「无论如何都取出内部值」,等于把 poisoning 显式关掉。`crates/pb-mapper-auth/src/lib.rs` 里还为此 +养了一个 `recover_lock` 辅助函数专门抹掉 `LockResult`。 + +`parking_lot` 不做 poisoning,于是: + +- `lock()` 直接返回 guard,所有样板和 `recover_lock` 一起消失; +- FFI 侧 `claim_key` 里那条「state is poisoned」错误分支变成不可达,直接删掉—— + 少一个永远不会发生的错误码; +- 未竞争时是纯自旋 + 无系统调用,比 std 的 futex 路径更快; +- 锁本身不必为 poisoning 保留状态,`Mutex` 只占一个字节加 `T`。 + +代价:panic 后锁会正常释放,其他线程可能看到中间状态的值。本项目原先就用 +`into_inner()` 接受了这个行为,所以迁移是纯简化,语义不变。 + +## 10. 小结 + +- **`Arc` 管生命周期,`Send`/`Sync` 管跨线程安全性,内部可变性管「通过 `&T` 修改」。** + 三件事正交,不能互相代替。 +- `Sync` 只赋予共享**读**。需要锁的条件是「我要通过 `&T` 改它」,而非「T 不是 Sync」。 +- `Mutex` 的类型学意义:把 `Send` 升级成 `Sync`。 +- 判断顺序:**能不能不共享 → 是否真要跨线程 → 粒度是整数还是临界区**。 + 第一步收益最大,也最常被跳过。 +- 一把锁如果无法说清它走完了这三步,它大概是被设计出来的,而不是需求。 diff --git a/docs/ui-cli-mode-spec.md b/docs/ui-cli-mode-spec.md index c94d47c..6c5bca3 100644 --- a/docs/ui-cli-mode-spec.md +++ b/docs/ui-cli-mode-spec.md @@ -1,5 +1,10 @@ # pb-mapper UI as a CLI — design spec +> Historical design document: this reflects the implementation at some point in +> 2025. Its line numbers and code references may have drifted — the crate split +> moved every path under `crates/` — and it is kept for design intent only. For +> current behaviour, read the code. + > Historical design note: the standalone CLI was consolidated in v0.3.0. In > current commands, `pb-mapper-server` maps to `pb-mapper server`, > `pb-mapper-server-cli` maps to `pb-mapper register`, and diff --git a/docs/user-guide.md b/docs/user-guide.md index 67ae576..dc4a4b1 100644 --- a/docs/user-guide.md +++ b/docs/user-guide.md @@ -4,7 +4,7 @@ ## Overview -pb-mapper exposes local TCP/UDP services through a public relay using a service key. One `pb-mapper` binary provides the `server`, `register`, `connect`, and `status` commands, alongside an optional Flutter GUI. +pb-mapper exposes local TCP/UDP services through a public relay using a service key. One `pb-mapper` binary provides the `server`, `register`, `connect`, `status`, and `admin` commands, alongside an optional Flutter GUI. ## How it works @@ -112,31 +112,52 @@ Optional flags: - `--ipv6`: enable IPv6 listening - `--keep-alive`: enable TCP keep-alive -- `--use-machine-msg-header-key`: derive `MSG_HEADER_KEY` from current machine hostname + MAC, - and write it to `/var/lib/pb-mapper-server/msg_header_key` +- `--auth-state-dir`: authentication state directory (Linux services default `/var/lib/pb-mapper/auth`; unprivileged Linux, macOS, and Windows use a user-writable application directory) +- `--max-temporary-keys`: fixed temporary-key slot capacity (default `65536`) +- `--max-temporary-key-ttl`: maximum issued TTL (default `30d`) +- `--legacy-protocol allow|deny`: initial legacy-client policy +- `--use-machine-msg-header-key`: explicit legacy compatibility mode + +### Administrator and temporary credentials + +On first start, the relay creates a random administrator key in its +authentication state directory (`admin.key`). On Linux system services that +is `/var/lib/pb-mapper/auth/admin.key`. Unprivileged Linux, macOS, and +Windows builds use a user-writable application directory instead. There is +no built-in default credential. +Keep the administrator key on the relay host and use it to issue a temporary +credential for a workload: -### Machine-derived `MSG_HEADER_KEY` (optional) +```bash +export MSG_HEADER_KEY="$(sudo cat /var/lib/pb-mapper/auth/admin.key)" +pb-mapper admin --server "your-server:7666" \ + key issue --ttl 24h --label my-service +``` -When you want each deployed server to use a host-specific key (instead of the built-in default), -start server with: +Export the printed `pbmt1_...` credential on both the register and connect +machines. The temporary key can see and use only its own namespace: ```bash -pb-mapper server --port 7666 --use-machine-msg-header-key +export MSG_HEADER_KEY='pbmt1_...' +pb-mapper register tcp --server "your-server:7666" --key "my-service" --addr "127.0.0.1:8080" ``` -This will: - -- derive a stable 32-byte key from hostname + MAC addresses -- set server process `MSG_HEADER_KEY` automatically -- persist the key to `/var/lib/pb-mapper-server/msg_header_key` +Renewing a key preserves the credential text. Revocation or expiry immediately +closes its active control and data connections. See +[`authentication-v2.md`](authentication-v2.md) for the full lifecycle, +namespace model, protocol framing, and migration procedure. -Then use the same key for the `register` and `connect` commands: +The machine-derived option remains available for an existing deployment, but +it is not recommended for new installations: ```bash -export MSG_HEADER_KEY="$(cat /var/lib/pb-mapper-server/msg_header_key)" -pb-mapper register tcp --server "your-server:7666" --key "my-service" --addr "127.0.0.1:8080" +pb-mapper server --port 7666 --use-machine-msg-header-key ``` +On upgrade, if no new administrator key or `MSG_HEADER_KEY` is present, the +relay automatically imports `/var/lib/pb-mapper-server/msg_header_key` so +legacy clients keep working. + ### 2) Register a local service Register a TCP service: @@ -177,6 +198,30 @@ pb-mapper status remote-id --server "your-server:7666" pb-mapper status keys --server "your-server:7666" ``` +An administrator can explicitly inspect or connect to a temporary namespace: + +```bash +pb-mapper status keys --server "your-server:7666" --namespace 4294967296 +pb-mapper connect tcp --server "your-server:7666" --namespace 4294967296 \ + --key "my-service" --addr "127.0.0.1:9090" +``` + +Registering as administrator inside a temporary namespace additionally requires +`--force`. + +### Administrator commands + +```bash +pb-mapper admin --server "your-server:7666" status +pb-mapper admin --server "your-server:7666" key list +pb-mapper admin --server "your-server:7666" key reveal 4294967296 +pb-mapper admin --server "your-server:7666" service list --all +pb-mapper admin --server "your-server:7666" connection list --all +``` + +Use `--output json` for one JSON document or `--output ndjson` for streaming +automation. Page size defaults to 100 and is capped at 1000. + ## Run (GUI) The Flutter UI can start the server, register services, and connect clients through a graphical workflow. Start it from `ui/`: @@ -189,6 +234,16 @@ flutter run ## Environment variables - `PB_MAPPER_SERVER`: default server address for the CLI +- `MSG_HEADER_KEY`: 32-character administrator key or a `pbmt1_` temporary credential +- `PB_MAPPER_AUTH_STATE_DIR`: relay auth-state directory (Linux services default `/var/lib/pb-mapper/auth`; unprivileged Linux, macOS, and Windows use a user-writable application directory) +- `PB_MAPPER_AUTH_MAX_TEMP_KEYS`: fixed temporary-key capacity, default `65536` +- `PB_MAPPER_AUTH_MAX_TEMP_TTL_SECS`: maximum temporary-key TTL, default 30 days +- `PB_MAPPER_LEGACY_PROTOCOL`: `allow` or `deny`, default `allow` +- `PB_MAPPER_MAX_SERVICES_PER_NAMESPACE`: service names per namespace, default `256` +- `PB_MAPPER_MAX_REGISTER_CONNECTIONS_PER_SERVICE`: control connections per service, default `16` +- `PB_MAPPER_MAX_STREAMS_PER_NAMESPACE`: active streams per namespace, default `1024` +- `PB_MAPPER_NEW_STREAMS_PER_SECOND`: sustained new-stream rate per namespace, default `100` +- `PB_MAPPER_NEW_STREAMS_BURST`: new-stream burst per namespace, default `200` - `PB_MAPPER_KEEP_ALIVE`: enable TCP keep-alive (set to `ON`) - `PB_MAPPER_LOG_FORMAT`: tracing output format, one of `pretty` (default), `compact`, or `json` - `PB_MAPPER_CONTROL_IO_TIMEOUT`: close stalled control-plane handshakes after this duration, default `30s` diff --git a/docs/user-guide.zh-CN.md b/docs/user-guide.zh-CN.md index 4278365..01a80e8 100644 --- a/docs/user-guide.zh-CN.md +++ b/docs/user-guide.zh-CN.md @@ -4,7 +4,7 @@ ## 概览 -pb-mapper 通过“服务 key”将本地 TCP/UDP 服务暴露到公网中继。统一的 `pb-mapper` 二进制提供 `server`、`register`、`connect`、`status` 四类命令,并保留可选的 Flutter GUI。 +pb-mapper 通过“服务 key”将本地 TCP/UDP 服务暴露到公网中继。统一的 `pb-mapper` 二进制提供 `server`、`register`、`connect`、`status`、`admin` 五类命令,并保留可选的 Flutter GUI。 ## 运转机制 @@ -112,30 +112,46 @@ pb-mapper server --port 7666 - `--ipv6`:开启 IPv6 监听 - `--keep-alive`:开启 TCP keep-alive -- `--use-machine-msg-header-key`:基于当前机器 hostname + MAC 派生 `MSG_HEADER_KEY`, - 并写入 `/var/lib/pb-mapper-server/msg_header_key` +- `--auth-state-dir`:认证状态目录(Linux 系统服务默认 `/var/lib/pb-mapper/auth`;无特权 Linux、macOS 与 Windows 使用当前用户可写的应用目录) +- `--max-temporary-keys`:临时 key 固定槽位容量,默认 `65536` +- `--max-temporary-key-ttl`:临时 key 最大 TTL,默认 `30d` +- `--legacy-protocol allow|deny`:旧协议初始接入策略 +- `--use-machine-msg-header-key`:明确启用旧版机器派生 key 兼容模式 -### 基于机器信息派生 `MSG_HEADER_KEY`(可选) +### 管理员密钥与临时凭据 -如果你希望每台部署机器都使用各自唯一的 key(而不是内置默认 key),可以这样启动服务端: +中继首次启动时会在认证状态目录生成随机管理员密钥(`admin.key`)。Linux 系统服务 +默认写到 `/var/lib/pb-mapper/auth/admin.key`;无特权 Linux、macOS 与 Windows +桌面构建则使用当前用户可写的应用目录。系统不再提供内置默认 key。管理员密钥留在 +中继机器上,用它为业务签发临时凭据: ```bash -pb-mapper server --port 7666 --use-machine-msg-header-key +export MSG_HEADER_KEY="$(sudo cat /var/lib/pb-mapper/auth/admin.key)" +pb-mapper admin --server "your-server:7666" \ + key issue --ttl 24h --label my-service ``` -该参数会完成: +把输出的 `pbmt1_...` 凭据导入 register 与 connect 两端。临时 key 只能看到并操作 +自己的命名空间: + +```bash +export MSG_HEADER_KEY='pbmt1_...' +pb-mapper register tcp --server "your-server:7666" --key "my-service" --addr "127.0.0.1:8080" +``` -- 基于 hostname + MAC 地址派生稳定的 32 字节 key -- 自动设置当前服务端进程的 `MSG_HEADER_KEY` -- 将 key 持久化到 `/var/lib/pb-mapper-server/msg_header_key` +续期不会改变凭据文本;吊销或到期会立即关闭对应的控制连接与数据连接。完整生命周期、 +命名空间、V2 帧格式和迁移步骤见 +[`authentication-v2.zh-CN.md`](authentication-v2.zh-CN.md)。 -随后在 `register` 与 `connect` 命令中使用同一 key: +已有部署仍可显式启用机器派生 key,但不建议新安装继续使用: ```bash -export MSG_HEADER_KEY="$(cat /var/lib/pb-mapper-server/msg_header_key)" -pb-mapper register tcp --server "your-server:7666" --key "my-service" --addr "127.0.0.1:8080" +pb-mapper server --port 7666 --use-machine-msg-header-key ``` +升级时,如果新管理员密钥和 `MSG_HEADER_KEY` 都不存在,中继会自动导入 +`/var/lib/pb-mapper-server/msg_header_key`,旧客户端无需立刻换 key。 + ### 2)注册本地服务 注册 TCP 服务: @@ -176,6 +192,29 @@ pb-mapper status remote-id --server "your-server:7666" pb-mapper status keys --server "your-server:7666" ``` +管理员可明确查看或连接某个临时命名空间: + +```bash +pb-mapper status keys --server "your-server:7666" --namespace 4294967296 +pb-mapper connect tcp --server "your-server:7666" --namespace 4294967296 \ + --key "my-service" --addr "127.0.0.1:9090" +``` + +管理员要在临时命名空间注册服务,还必须增加 `--force`。 + +### 管理命令 + +```bash +pb-mapper admin --server "your-server:7666" status +pb-mapper admin --server "your-server:7666" key list +pb-mapper admin --server "your-server:7666" key reveal 4294967296 +pb-mapper admin --server "your-server:7666" service list --all +pb-mapper admin --server "your-server:7666" connection list --all +``` + +自动化场景可使用 `--output json` 输出单个 JSON 文档,或用 `--output ndjson` 流式输出。 +默认每页 100 条,最大 1000 条。 + ## 运行(GUI) Flutter UI 可用于启动服务器、注册服务与建立连接。启动方式: @@ -188,6 +227,16 @@ flutter run ## 环境变量 - `PB_MAPPER_SERVER`:CLI 默认服务器地址 +- `MSG_HEADER_KEY`:32 字符管理员密钥或 `pbmt1_` 临时凭据 +- `PB_MAPPER_AUTH_STATE_DIR`:中继认证状态目录(Linux 系统服务默认 `/var/lib/pb-mapper/auth`;无特权 Linux、macOS 与 Windows 使用当前用户可写的应用目录) +- `PB_MAPPER_AUTH_MAX_TEMP_KEYS`:临时 key 固定容量,默认 `65536` +- `PB_MAPPER_AUTH_MAX_TEMP_TTL_SECS`:临时 key 最大 TTL,默认 30 天 +- `PB_MAPPER_LEGACY_PROTOCOL`:`allow` 或 `deny`,默认 `allow` +- `PB_MAPPER_MAX_SERVICES_PER_NAMESPACE`:每命名空间 service 数,默认 `256` +- `PB_MAPPER_MAX_REGISTER_CONNECTIONS_PER_SERVICE`:每 service 控制连接数,默认 `16` +- `PB_MAPPER_MAX_STREAMS_PER_NAMESPACE`:每命名空间活动 stream 数,默认 `1024` +- `PB_MAPPER_NEW_STREAMS_PER_SECOND`:每命名空间持续新建 stream 速率,默认 `100` +- `PB_MAPPER_NEW_STREAMS_BURST`:每命名空间新建 stream 突发量,默认 `200` - `PB_MAPPER_KEEP_ALIVE`:启用 TCP keep-alive(设置为 `ON`) - `PB_MAPPER_LOG_FORMAT`:tracing 输出格式,可选 `pretty`(默认)、`compact` 或 `json` - `PB_MAPPER_CONTROL_IO_TIMEOUT`:控制面握手卡住后的关闭时间,默认 `30s` diff --git a/rust-toolchain.toml b/rust-toolchain.toml index e88baf1..b73c15e 100644 --- a/rust-toolchain.toml +++ b/rust-toolchain.toml @@ -1,2 +1,2 @@ [toolchain] -channel = "1.88.0" +channel = "1.98.0" diff --git a/scripts/install-server-gitee.sh b/scripts/install-server-gitee.sh index 51cd1c8..d8acf0a 100755 --- a/scripts/install-server-gitee.sh +++ b/scripts/install-server-gitee.sh @@ -2,7 +2,7 @@ set -euo pipefail # Configuration -VERSION="${PB_MAPPER_VERSION:-0.3.0}" +VERSION="${PB_MAPPER_VERSION:-0.4.0}" ARCH="${PB_MAPPER_ARCH:-x86_64-unknown-linux-musl}" TARBALL="pb-mapper-${ARCH}.tar.gz" DOWNLOAD_URL="https://gitee.com/acking-you/pb-mapper/releases/download/v${VERSION}/${TARBALL}" @@ -10,6 +10,37 @@ INSTALL_DIR="/usr/local/bin" SERVICE_NAME="pb-mapper-server" SERVICE_PATH="/etc/systemd/system/${SERVICE_NAME}.service" PORT="${PB_MAPPER_PORT:-7666}" +AUTH_DIR="/var/lib/pb-mapper/auth" +ADMIN_KEY_PATH="${AUTH_DIR}/admin.key" +LEGACY_KEY_PATH="/var/lib/pb-mapper-server/msg_header_key" +SERVER_ENV_FILE="/etc/pb-mapper/server.env" + +admin_key_is_env_safe() { + local key="$1" + local bytes + bytes=$(printf '%s' "$key" | wc -c) + [ "$bytes" -eq 32 ] || return 1 + printf '%s' "$key" | LC_ALL=C grep -qx '[[:graph:]]\{32\}' +} + +configured_msg_header_key() { + if [ -n "${MSG_HEADER_KEY:-}" ]; then + printf '%s' "$MSG_HEADER_KEY" + return 0 + fi + if [ ! -f "$SERVER_ENV_FILE" ]; then + return 0 + fi + awk -F= ' + $1 ~ /^[[:space:]]*#/ { next } + $1 ~ /^[[:space:]]*MSG_HEADER_KEY[[:space:]]*$/ { + val = substr($0, index($0, "=") + 1) + sub(/\r$/, "", val) + key = val + } + END { printf "%s", key } + ' "$SERVER_ENV_FILE" +} # Re-run with sudo if needed if [ "${EUID:-$(id -u)}" -ne 0 ]; then @@ -67,6 +98,39 @@ fi mkdir -p "$INSTALL_DIR" install -m 0755 "$BIN_PATH" "${INSTALL_DIR}/pb-mapper" +# Preserve the former machine-derived credential on upgrade only when neither +# admin.key nor an explicit MSG_HEADER_KEY is already configured. An explicit +# key in the environment or /etc/pb-mapper/server.env must win; otherwise the +# runtime would prefer the newly copied admin.key and lock operators out. +install -d -m 0700 "$AUTH_DIR" +INSTALLER_KEY="$(configured_msg_header_key)" +if [ -n "$INSTALLER_KEY" ] && [ ! -s "$ADMIN_KEY_PATH" ]; then + case "$INSTALLER_KEY" in + pbmt1_*) + echo "MSG_HEADER_KEY is a temporary credential; write a 32-character administrator key to $ADMIN_KEY_PATH" >&2 + exit 1 + ;; + esac + if ! admin_key_is_env_safe "$INSTALLER_KEY"; then + echo "MSG_HEADER_KEY must be exactly 32 printable ASCII bytes without whitespace or NUL" >&2 + exit 1 + fi + if [ -s "${AUTH_DIR}/auth.snapshot" ] || [ -s "${AUTH_DIR}/auth.wal" ]; then + echo "Leaving $ADMIN_KEY_PATH unset so the service can verify MSG_HEADER_KEY against existing authentication state" + else + printf '%s\n' "$INSTALLER_KEY" > "$ADMIN_KEY_PATH" + chmod 0600 "$ADMIN_KEY_PATH" + echo "Persisted installer MSG_HEADER_KEY to $ADMIN_KEY_PATH" + fi +elif [ -z "$INSTALLER_KEY" ] && [ ! -s "$ADMIN_KEY_PATH" ] && [ -s "$LEGACY_KEY_PATH" ]; then + if [ -s "${AUTH_DIR}/auth.snapshot" ] || [ -s "${AUTH_DIR}/auth.wal" ]; then + echo "Leaving $ADMIN_KEY_PATH unset so the service can verify the legacy key against existing authentication state" + else + install -m 0600 "$LEGACY_KEY_PATH" "$ADMIN_KEY_PATH" + echo "Migrated the legacy machine-derived key into $ADMIN_KEY_PATH" + fi +fi + # Stop and remove existing service if present if systemctl is-active --quiet "${SERVICE_NAME}.service"; then systemctl stop "${SERVICE_NAME}.service" @@ -86,8 +150,11 @@ After=network.target [Service] Type=simple -ExecStart=${INSTALL_DIR}/pb-mapper server --port ${PORT} --use-machine-msg-header-key +ExecStart=${INSTALL_DIR}/pb-mapper server --port ${PORT} Environment=RUST_LOG=info +EnvironmentFile=-/etc/pb-mapper/server.env +StateDirectory=pb-mapper +StateDirectoryMode=0700 Restart=on-failure RestartSec=3 LimitNOFILE=65535 @@ -102,4 +169,5 @@ systemctl enable --now "${SERVICE_NAME}.service" echo "pb-mapper server is installed and running." echo "Service name: ${SERVICE_NAME}.service" -echo "Machine-derived key file: /var/lib/pb-mapper-server/msg_header_key" +echo "Administrator key file: /var/lib/pb-mapper/auth/admin.key" +echo "Read it locally as root and issue temporary credentials for register/connect clients." diff --git a/scripts/install-server-github.sh b/scripts/install-server-github.sh index beb7d9b..13acdc4 100755 --- a/scripts/install-server-github.sh +++ b/scripts/install-server-github.sh @@ -2,7 +2,7 @@ set -euo pipefail # Configuration -VERSION="${PB_MAPPER_VERSION:-0.3.0}" +VERSION="${PB_MAPPER_VERSION:-0.4.0}" ARCH="${PB_MAPPER_ARCH:-x86_64-unknown-linux-musl}" TARBALL="pb-mapper-${ARCH}.tar.gz" DOWNLOAD_URL="https://github.com/acking-you/pb-mapper/releases/download/v${VERSION}/${TARBALL}" @@ -10,6 +10,37 @@ INSTALL_DIR="/usr/local/bin" SERVICE_NAME="pb-mapper-server" SERVICE_PATH="/etc/systemd/system/${SERVICE_NAME}.service" PORT="${PB_MAPPER_PORT:-7666}" +AUTH_DIR="/var/lib/pb-mapper/auth" +ADMIN_KEY_PATH="${AUTH_DIR}/admin.key" +LEGACY_KEY_PATH="/var/lib/pb-mapper-server/msg_header_key" +SERVER_ENV_FILE="/etc/pb-mapper/server.env" + +admin_key_is_env_safe() { + local key="$1" + local bytes + bytes=$(printf '%s' "$key" | wc -c) + [ "$bytes" -eq 32 ] || return 1 + printf '%s' "$key" | LC_ALL=C grep -qx '[[:graph:]]\{32\}' +} + +configured_msg_header_key() { + if [ -n "${MSG_HEADER_KEY:-}" ]; then + printf '%s' "$MSG_HEADER_KEY" + return 0 + fi + if [ ! -f "$SERVER_ENV_FILE" ]; then + return 0 + fi + awk -F= ' + $1 ~ /^[[:space:]]*#/ { next } + $1 ~ /^[[:space:]]*MSG_HEADER_KEY[[:space:]]*$/ { + val = substr($0, index($0, "=") + 1) + sub(/\r$/, "", val) + key = val + } + END { printf "%s", key } + ' "$SERVER_ENV_FILE" +} # Re-run with sudo if needed if [ "${EUID:-$(id -u)}" -ne 0 ]; then @@ -67,6 +98,39 @@ fi mkdir -p "$INSTALL_DIR" install -m 0755 "$BIN_PATH" "${INSTALL_DIR}/pb-mapper" +# Preserve the former machine-derived credential on upgrade only when neither +# admin.key nor an explicit MSG_HEADER_KEY is already configured. An explicit +# key in the environment or /etc/pb-mapper/server.env must win; otherwise the +# runtime would prefer the newly copied admin.key and lock operators out. +install -d -m 0700 "$AUTH_DIR" +INSTALLER_KEY="$(configured_msg_header_key)" +if [ -n "$INSTALLER_KEY" ] && [ ! -s "$ADMIN_KEY_PATH" ]; then + case "$INSTALLER_KEY" in + pbmt1_*) + echo "MSG_HEADER_KEY is a temporary credential; write a 32-character administrator key to $ADMIN_KEY_PATH" >&2 + exit 1 + ;; + esac + if ! admin_key_is_env_safe "$INSTALLER_KEY"; then + echo "MSG_HEADER_KEY must be exactly 32 printable ASCII bytes without whitespace or NUL" >&2 + exit 1 + fi + if [ -s "${AUTH_DIR}/auth.snapshot" ] || [ -s "${AUTH_DIR}/auth.wal" ]; then + echo "Leaving $ADMIN_KEY_PATH unset so the service can verify MSG_HEADER_KEY against existing authentication state" + else + printf '%s\n' "$INSTALLER_KEY" > "$ADMIN_KEY_PATH" + chmod 0600 "$ADMIN_KEY_PATH" + echo "Persisted installer MSG_HEADER_KEY to $ADMIN_KEY_PATH" + fi +elif [ -z "$INSTALLER_KEY" ] && [ ! -s "$ADMIN_KEY_PATH" ] && [ -s "$LEGACY_KEY_PATH" ]; then + if [ -s "${AUTH_DIR}/auth.snapshot" ] || [ -s "${AUTH_DIR}/auth.wal" ]; then + echo "Leaving $ADMIN_KEY_PATH unset so the service can verify the legacy key against existing authentication state" + else + install -m 0600 "$LEGACY_KEY_PATH" "$ADMIN_KEY_PATH" + echo "Migrated the legacy machine-derived key into $ADMIN_KEY_PATH" + fi +fi + # Stop and remove existing service if present if systemctl is-active --quiet "${SERVICE_NAME}.service"; then systemctl stop "${SERVICE_NAME}.service" @@ -86,8 +150,11 @@ After=network.target [Service] Type=simple -ExecStart=${INSTALL_DIR}/pb-mapper server --port ${PORT} --use-machine-msg-header-key +ExecStart=${INSTALL_DIR}/pb-mapper server --port ${PORT} Environment=RUST_LOG=info +EnvironmentFile=-/etc/pb-mapper/server.env +StateDirectory=pb-mapper +StateDirectoryMode=0700 Restart=on-failure RestartSec=3 LimitNOFILE=65535 @@ -102,4 +169,5 @@ systemctl enable --now "${SERVICE_NAME}.service" echo "pb-mapper server is installed and running." echo "Service name: ${SERVICE_NAME}.service" -echo "Machine-derived key file: /var/lib/pb-mapper-server/msg_header_key" +echo "Administrator key file: /var/lib/pb-mapper/auth/admin.key" +echo "Read it locally as root and issue temporary credentials for register/connect clients." diff --git a/scripts/release/entrypoint/pb-mapper.sh b/scripts/release/entrypoint/pb-mapper.sh index b750c00..2c4258e 100644 --- a/scripts/release/entrypoint/pb-mapper.sh +++ b/scripts/release/entrypoint/pb-mapper.sh @@ -7,9 +7,22 @@ if [ -z "${PB_MAPPER_PORT:-}" ]; then exit 1 fi -USE_MACHINE_MSG_HEADER_KEY=${USE_MACHINE_MSG_HEADER_KEY:-true} +USE_MACHINE_MSG_HEADER_KEY=${USE_MACHINE_MSG_HEADER_KEY:-false} ARGS=(-p "$PB_MAPPER_PORT") +AUTH_DIR="${PB_MAPPER_AUTH_STATE_DIR:-/var/lib/pb-mapper/auth}" +ADMIN_KEY_PATH="$AUTH_DIR/admin.key" +LEGACY_KEY_PATH="/var/lib/pb-mapper-server/msg_header_key" + +install -d -m 0700 "$AUTH_DIR" +if [ -z "${MSG_HEADER_KEY:-}" ] && [ ! -s "$ADMIN_KEY_PATH" ] && [ -s "$LEGACY_KEY_PATH" ]; then + if [ -s "$AUTH_DIR/auth.snapshot" ] || [ -s "$AUTH_DIR/auth.wal" ]; then + echo "Leaving $ADMIN_KEY_PATH unset so the service can verify the legacy key against existing authentication state" + else + install -m 0600 "$LEGACY_KEY_PATH" "$ADMIN_KEY_PATH" + echo "Migrated the legacy machine-derived key into $ADMIN_KEY_PATH" + fi +fi if [ "${USE_IPV6:-false}" = "true" ]; then echo "USE_IPV6 is set to true" @@ -19,10 +32,20 @@ else fi if [ "$USE_MACHINE_MSG_HEADER_KEY" = "true" ]; then - echo "USE_MACHINE_MSG_HEADER_KEY is set to true" - ARGS+=(--use-machine-msg-header-key) + if [ -s "$ADMIN_KEY_PATH" ]; then + echo "admin.key already exists; skipping --use-machine-msg-header-key" + else + echo "WARNING: USE_MACHINE_MSG_HEADER_KEY is a legacy compatibility mode" + ARGS+=(--use-machine-msg-header-key) + fi else echo "USE_MACHINE_MSG_HEADER_KEY is set to false" fi +if [ -n "${MSG_HEADER_KEY:-}" ]; then + echo "Using the configured administrator credential" +else + echo "Using or initializing $ADMIN_KEY_PATH" +fi + exec ./pb-mapper server "${ARGS[@]}" diff --git a/services/pb-mapper-server.service b/services/pb-mapper-server.service index 822a251..9f13c92 100644 --- a/services/pb-mapper-server.service +++ b/services/pb-mapper-server.service @@ -7,7 +7,9 @@ Wants=network-online.target Type=simple Environment=RUST_LOG=info EnvironmentFile=-/etc/pb-mapper/server.env -ExecStart=/usr/local/bin/pb-mapper server --port 7666 --use-machine-msg-header-key +ExecStart=/usr/local/bin/pb-mapper server --port 7666 +StateDirectory=pb-mapper +StateDirectoryMode=0700 Restart=on-failure RestartSec=3 LimitNOFILE=65535 diff --git a/services/readme.md b/services/readme.md index e4dda8e..49d6942 100644 --- a/services/readme.md +++ b/services/readme.md @@ -9,9 +9,20 @@ sudo install -m 0644 services/pb-mapper-register@.service /etc/systemd/system/ sudo install -m 0644 services/pb-mapper-connect@.service /etc/systemd/system/ ``` -The relay unit runs `pb-mapper server` directly. Override it with a systemd -drop-in if the default `7666` port or machine-derived key behavior is not -appropriate. +The relay unit runs `pb-mapper server` directly. On first start it creates a +random administrator key at `/var/lib/pb-mapper/auth/admin.key` with mode +`0600`. Keep that directory persistent. If upgrading from the old +machine-derived mode, the first v0.4 start automatically copies +`/var/lib/pb-mapper-server/msg_header_key` to the new path when neither an +administrator key file nor `MSG_HEADER_KEY` is already configured. + +Use the administrator key locally for management, then issue a scoped temporary +credential for each tenant or workload: + +```bash +export MSG_HEADER_KEY="$(sudo cat /var/lib/pb-mapper/auth/admin.key)" +pb-mapper admin --server relay.example.com:7666 key issue --ttl 24h --label home-web +``` Registration instances read `/etc/pb-mapper/register/.env`: @@ -21,7 +32,7 @@ SERVICE_KEY=home-web LOCAL_ADDR=127.0.0.1:8080 TRANSPORT=tcp REGISTER_EXTRA_ARGS=--codec --keep-alive -MSG_HEADER_KEY=replace-with-the-shared-32-byte-key +MSG_HEADER_KEY=pbmt1_replace-with-an-issued-temporary-credential ``` Connect instances read `/etc/pb-mapper/connect/.env`: @@ -32,7 +43,7 @@ SERVICE_KEY=home-web LOCAL_ADDR=127.0.0.1:9090 TRANSPORT=tcp CONNECT_EXTRA_ARGS=--keep-alive -MSG_HEADER_KEY=replace-with-the-shared-32-byte-key +MSG_HEADER_KEY=pbmt1_replace-with-the-same-temporary-credential ``` Create the matching directory and env file, then enable the instance: diff --git a/skills/pb-mapper-connect-deploy/SKILL.md b/skills/pb-mapper-connect-deploy/SKILL.md index 5c6f7b7..15141a4 100644 --- a/skills/pb-mapper-connect-deploy/SKILL.md +++ b/skills/pb-mapper-connect-deploy/SKILL.md @@ -38,7 +38,7 @@ Prompt the user for each value below. Do NOT assume or hardcode any value. | `SERVICE_KEY` | pb-mapper service name to subscribe | — | Required | | `LISTEN_IP` | Listen on localhost only (`127.0.0.1`) or all interfaces (`0.0.0.0`)? | `127.0.0.1` | Required | | `LISTEN_PORT` | Local listening port on remote host | — | Required | -| `MSG_HEADER_KEY` | Encryption key (exactly 32 chars) | *(empty)* | Optional, **confidential** — never log or echo | +| `MSG_HEADER_KEY` | Administrator key or `pbmt1_...` temporary credential | *(empty)* | Required for authenticated v2 connections; **confidential** — never log or echo | | `PUBLIC_CHECK_URL` | URL for external validation | *(empty)* | Optional | | `VERSION` | Release version (without `v` prefix) | Latest release | Auto-detect or user-specified | | `TARGET_TRIPLE` | Build target | `x86_64-unknown-linux-musl` | User can override | @@ -143,7 +143,10 @@ REMOTE_SYSTEMD ### 4. Write instance env file and start service -`MSG_HEADER_KEY` must be omitted when empty; never write an empty value to env file. +Prefer a temporary credential issued by the relay administrator. It can register, +connect, and inspect only its own namespace, and it expires automatically. Use the +administrator key only for relay administration or an intentional namespace-0 +deployment. Never write an empty credential to the environment file. ```bash ssh ${SSH_PORT_OPT} "${SSH_TARGET}" \ @@ -159,14 +162,21 @@ RUST_LOG=info PB_MAPPER_KEEP_ALIVE=ON EOF -if [ -n "${MSG_HEADER_KEY}" ]; then - CLEAN_KEY="$(printf '%s' "${MSG_HEADER_KEY}" | tr -d '\r\n')" - if [ "${#CLEAN_KEY}" -ne 32 ]; then - echo "MSG_HEADER_KEY must be exactly 32 characters" >&2 - exit 1 - fi - echo "MSG_HEADER_KEY=${CLEAN_KEY}" | sudo tee -a "${ENV_FILE}" >/dev/null +if [ -z "${MSG_HEADER_KEY}" ]; then + echo "MSG_HEADER_KEY is required" >&2 + exit 1 fi +CLEAN_KEY="$(printf '%s' "${MSG_HEADER_KEY}" | tr -d '\r\n')" +case "${CLEAN_KEY}" in + pbmt1_*) ;; + *) + if [ "${#CLEAN_KEY}" -ne 32 ]; then + echo "MSG_HEADER_KEY must be a 32-byte administrator key or pbmt1_ temporary credential" >&2 + exit 1 + fi + ;; +esac +echo "MSG_HEADER_KEY=${CLEAN_KEY}" | sudo tee -a "${ENV_FILE}" >/dev/null sudo systemctl daemon-reload sudo systemctl enable --now "pb-mapper-connect@${INSTANCE_NAME}.service" @@ -191,7 +201,10 @@ If `jq` is available, pipe through `jq .` for formatted JSON output. Use this quick triage when startup or forwarding fails: -- `datalen not valid`: likely `MSG_HEADER_KEY` mismatch or hidden newline; verify both sides use the same 32-byte key. +- `protocol_v2_decrypt_failed`: credential does not belong to this relay or was corrupted in transit. +- `temporary_key_expired` / `temporary_key_revoked`: ask an administrator to renew the same key ID or issue a replacement. +- `namespace_access_denied`: a temporary credential can only use its own namespace; remove an incorrect `--namespace` value. +- Legacy `datalen not valid`: administrator key mismatch or a hidden newline in an old protocol-v1 client. - Service restarts immediately: inspect logs with `journalctl -u pb-mapper-connect@${INSTANCE_NAME} -n 200 --no-pager`. - Remote port not listening: confirm `LOCAL_ADDR` host/port and no port conflict. - Public URL fails but localhost works: investigate reverse proxy (for example, Caddy route/TLS config). diff --git a/skills/pb-mapper-server-deploy/SKILL.md b/skills/pb-mapper-server-deploy/SKILL.md index d3ddcd7..ec24086 100644 --- a/skills/pb-mapper-server-deploy/SKILL.md +++ b/skills/pb-mapper-server-deploy/SKILL.md @@ -39,7 +39,6 @@ Prompt the user for each value below. Do NOT assume or hardcode any value. | `SERVER_PORT` | `pb-mapper server` listening port | `7666` | User can override | | `USE_IPV6` | Listen on IPv6 (`::`) instead of IPv4? | `false` | Optional | | `ENABLE_KEEP_ALIVE` | Enable TCP keep-alive? | `false` | Optional | -| `USE_MACHINE_KEY` | Use machine-derived msg header key? | `true` | Recommended; generates and persists a 32-char key | | `VERSION` | Release version (without `v` prefix) | Latest release | Auto-detect or user-specified | | `TARGET_TRIPLE` | Build target | `x86_64-unknown-linux-musl` | User can override | @@ -67,9 +66,6 @@ fi if [ "${ENABLE_KEEP_ALIVE}" = "true" ]; then EXTRA_FLAGS="${EXTRA_FLAGS} --keep-alive" fi -if [ "${USE_MACHINE_KEY}" = "true" ]; then - EXTRA_FLAGS="${EXTRA_FLAGS} --use-machine-msg-header-key" -fi export EXTRA_FLAGS ``` @@ -138,6 +134,8 @@ After=network.target Type=simple ExecStart=/usr/local/bin/pb-mapper server --port ${SERVER_PORT} ${EXTRA_FLAGS} Environment=RUST_LOG=info +StateDirectory=pb-mapper +StateDirectoryMode=0700 Restart=on-failure RestartSec=3 LimitNOFILE=65535 @@ -165,13 +163,25 @@ ssh ${SSH_PORT_OPT} "${SSH_TARGET}" "sudo systemctl --no-pager --full status pb- ssh ${SSH_PORT_OPT} "${SSH_TARGET}" "ss -tlnp | grep ':${SERVER_PORT}'" ``` -If `USE_MACHINE_KEY=true`, retrieve the generated key for use with `register` and `connect`: +On a fresh installation, the relay creates a random administrator key with mode +`0600`. Retrieve it once to initialize the administrator CLI: + +```bash +ssh ${SSH_PORT_OPT} "${SSH_TARGET}" "sudo cat /var/lib/pb-mapper/auth/admin.key" +``` + +Store it securely. Do not distribute it to ordinary register/connect instances. +Instead, load it only in the administrator shell and issue a scoped temporary +credential: ```bash -ssh ${SSH_PORT_OPT} "${SSH_TARGET}" "sudo cat /var/lib/pb-mapper-server/msg_header_key" +export MSG_HEADER_KEY='' +pb-mapper admin --server "${SSH_HOST}:${SERVER_PORT}" key issue --ttl 30d --label '' ``` -Store this key securely — it is needed by `pb-mapper register` and `pb-mapper connect` when `--codec` or `MSG_HEADER_KEY` is used. +Copy the returned `pbmt1_...` credential to the business instance. Use +`pb-mapper admin key renew --ttl 30d` to extend it in place without +changing deployed credentials. ## Troubleshooting Checklist @@ -180,7 +190,8 @@ Use this quick triage when the server fails to start or accept connections: - Port already in use: check with `ss -tlnp | grep :${SERVER_PORT}` and stop the conflicting process. - Firewall blocking: ensure the server port is open (`ufw allow ${SERVER_PORT}/tcp` or equivalent). - Service crashes on start: inspect logs with `journalctl -u pb-mapper-server -n 200 --no-pager`. -- Machine key not generated: verify `/var/lib/pb-mapper-server/` directory exists and is writable; check logs for key derivation errors. +- Administrator key missing: verify `/var/lib/pb-mapper/auth/` is writable and inspect startup logs for `administrator_key_initialized` or an `auth_stage` error. +- Temporary clients rejected: run `pb-mapper admin status`, then `pb-mapper admin key show ` and correlate the stable error code in the server log. - Permission denied: binary must be owned by root with `0755` permissions. ## Safe Update Procedure diff --git a/src/bin/pb-mapper.rs b/src/bin/pb-mapper.rs deleted file mode 100644 index 686e445..0000000 --- a/src/bin/pb-mapper.rs +++ /dev/null @@ -1,286 +0,0 @@ -use std::error::Error; -use std::net::{IpAddr, Ipv4Addr, Ipv6Addr}; - -use better_mimalloc_rs::MiMalloc; -use clap::{Args, Parser, Subcommand, ValueEnum}; -use pb_mapper::common::checksum::{setup_machine_msg_header_key, MACHINE_MSG_HEADER_KEY_PATH}; -use pb_mapper::common::config::{ - get_pb_mapper_server_async, get_sockaddr_async, init_tracing, keep_alive_from_env, StatusOp, -}; -use pb_mapper::common::message::forward::StreamForward; -use pb_mapper::local::client::{handle_status_cli, run_client_side_cli}; -use pb_mapper::local::server::{run_server_side_cli, ServerTunnelOptions}; -use pb_mapper::pb_server::run_server_with_shutdown; -use tokio_util::sync::CancellationToken; -use uni_stream::stream::{ - StreamProvider, TcpListenerProvider, TcpStreamProvider, UdpListenerProvider, UdpStreamProvider, -}; - -#[global_allocator] -static GLOBAL_MIMALLOC: MiMalloc = MiMalloc; - -#[derive(Debug, Parser)] -#[command( - author = "L_B__", - version, - about = "Expose and consume keyed TCP/UDP services through a pb-mapper relay", - subcommand_required = true, - arg_required_else_help = true -)] -struct Cli { - #[command(subcommand)] - command: Command, -} - -#[derive(Debug, Subcommand)] -enum Command { - /// Run the public relay server. - Server(ServerArgs), - /// Register a local service with a relay. - Register(RegisterArgs), - /// Expose a registered service on a local listening address. - Connect(ConnectArgs), - /// Query relay status. - Status(StatusArgs), -} - -#[derive(Debug, Args)] -struct ServerArgs { - /// Port exposed to registering services and connecting clients. - #[arg(short, long, visible_alias = "pb-mapper-port", default_value_t = 7666)] - port: u16, - /// Listen on IPv6 (::) instead of IPv4 (0.0.0.0). - #[arg(long, visible_alias = "use-ipv6", default_value_t = false)] - ipv6: bool, - /// Enable TCP keep-alive. PB_MAPPER_KEEP_ALIVE=ON is also supported. - #[arg(long, default_value_t = false)] - keep_alive: bool, - /// Derive MSG_HEADER_KEY from this machine and persist it for other roles. - #[arg(long, default_value_t = false)] - use_machine_msg_header_key: bool, -} - -#[derive(Debug, Args)] -struct RegisterArgs { - /// Transport used by the local service. - #[arg(value_enum)] - transport: Transport, - /// Service key registered with the relay. - #[arg(short, long)] - key: String, - /// Local service address to forward to. - #[arg(short, long, visible_alias = "local")] - addr: String, - #[command(flatten)] - relay: RelayArgs, - /// Encrypt forwarded traffic with the configured MSG_HEADER_KEY. - #[arg(short, long, default_value_t = false)] - codec: bool, -} - -#[derive(Debug, Args)] -struct ConnectArgs { - /// Transport exposed by the local listener. - #[arg(value_enum)] - transport: Transport, - /// Registered service key to subscribe to. - #[arg(short, long)] - key: String, - /// Local address on which downstream clients connect. - #[arg(short, long, visible_alias = "local")] - addr: String, - #[command(flatten)] - relay: RelayArgs, -} - -#[derive(Debug, Args)] -struct StatusArgs { - /// Status query to execute. - #[arg(value_enum)] - op: StatusOp, - /// Relay address. Falls back to PB_MAPPER_SERVER. - #[arg(short, long, visible_alias = "pb-mapper-server", value_name = "ADDR")] - server: Option, -} - -#[derive(Debug, Args)] -struct RelayArgs { - /// Relay address. Falls back to PB_MAPPER_SERVER. - #[arg(short, long, visible_alias = "pb-mapper-server", value_name = "ADDR")] - server: Option, - /// Enable TCP keep-alive. PB_MAPPER_KEEP_ALIVE=ON is also supported. - #[arg(long, default_value_t = false)] - keep_alive: bool, -} - -#[derive(Clone, Copy, Debug, Eq, PartialEq, ValueEnum)] -enum Transport { - Tcp, - Udp, -} - -#[tokio::main] -async fn main() { - MiMalloc::init(); - let cli = Cli::parse(); - init_tracing(); - - if let Err(error) = run(cli).await { - tracing::error!(%error, "pb-mapper command failed"); - std::process::exit(1); - } -} - -async fn run(cli: Cli) -> Result<(), Box> { - match cli.command { - Command::Server(args) => run_server(args).await?, - Command::Register(args) => run_register(args).await?, - Command::Connect(args) => run_connect(args).await?, - Command::Status(args) => run_status(args).await?, - } - Ok(()) -} - -async fn run_server(args: ServerArgs) -> Result<(), Box> { - if args.use_machine_msg_header_key { - setup_machine_msg_header_key()?; - tracing::info!( - path = MACHINE_MSG_HEADER_KEY_PATH, - "derived and persisted machine MSG_HEADER_KEY" - ); - } - - let ip_addr = if args.ipv6 { - IpAddr::V6(Ipv6Addr::UNSPECIFIED) - } else { - IpAddr::V4(Ipv4Addr::UNSPECIFIED) - }; - run_server_with_shutdown( - (ip_addr, args.port), - CancellationToken::new(), - None, - args.keep_alive || keep_alive_from_env(), - ) - .await?; - Ok(()) -} - -async fn run_register(args: RegisterArgs) -> Result<(), Box> { - let local_addr = get_sockaddr_async(&args.addr).await?; - let remote_addr = get_pb_mapper_server_async(args.relay.server.as_deref()).await?; - let options = ServerTunnelOptions { - need_codec: args.codec, - is_datagram: args.transport == Transport::Udp, - keep_alive: args.relay.keep_alive || keep_alive_from_env(), - }; - - match args.transport { - Transport::Tcp => { - register::(local_addr, remote_addr, args.key, options).await - } - Transport::Udp => { - register::(local_addr, remote_addr, args.key, options).await - } - } - Ok(()) -} - -async fn register( - local_addr: std::net::SocketAddr, - remote_addr: std::net::SocketAddr, - key: String, - options: ServerTunnelOptions, -) where - LocalStream::Item: StreamForward, -{ - run_server_side_cli::(local_addr, remote_addr, key.into(), options).await; -} - -async fn run_connect(args: ConnectArgs) -> Result<(), Box> { - let local_addr = get_sockaddr_async(&args.addr).await?; - let remote_addr = get_pb_mapper_server_async(args.relay.server.as_deref()).await?; - let key = args.key.into(); - let keep_alive = args.relay.keep_alive || keep_alive_from_env(); - - match args.transport { - Transport::Tcp => { - run_client_side_cli::(local_addr, remote_addr, key, keep_alive) - .await; - } - Transport::Udp => { - run_client_side_cli::(local_addr, remote_addr, key, keep_alive) - .await; - } - } - Ok(()) -} - -async fn run_status(args: StatusArgs) -> Result<(), Box> { - let remote_addr = get_pb_mapper_server_async(args.server.as_deref()).await?; - handle_status_cli(args.op, remote_addr).await; - Ok(()) -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn parses_each_runtime_role() { - let cases = [ - vec!["pb-mapper", "server", "--port", "7666", "--ipv6"], - vec![ - "pb-mapper", - "register", - "tcp", - "--key", - "web", - "--addr", - "127.0.0.1:8080", - "--server", - "relay:7666", - "--codec", - ], - vec![ - "pb-mapper", - "connect", - "udp", - "--key", - "game", - "--addr", - "127.0.0.1:8211", - "--server", - "relay:7666", - ], - vec!["pb-mapper", "status", "keys", "--server", "relay:7666"], - ]; - - for args in cases { - Cli::try_parse_from(args).expect("unified command should parse"); - } - } - - #[test] - fn accepts_documented_option_aliases() { - Cli::try_parse_from([ - "pb-mapper", - "server", - "--pb-mapper-port", - "7666", - "--use-ipv6", - ]) - .expect("server aliases should parse"); - Cli::try_parse_from([ - "pb-mapper", - "register", - "tcp", - "--key", - "web", - "--local", - "127.0.0.1:8080", - "--pb-mapper-server", - "relay:7666", - ]) - .expect("relay and local aliases should parse"); - } -} diff --git a/src/common/message/command.rs b/src/common/message/command.rs deleted file mode 100644 index e4feef0..0000000 --- a/src/common/message/command.rs +++ /dev/null @@ -1,162 +0,0 @@ -use serde::{Deserialize, Serialize}; -use snafu::ResultExt; - -use super::super::error::{MsgSerializeSnafu, Result}; -use crate::common::checksum::AesKeyType; - -pub const CONTROL_PROTOCOL_V2: u16 = 2; - -pub trait MessageSerializer { - fn encode(&self) -> Result>; - fn decode(msg: &[u8]) -> Result - where - Self: Sized; -} - -#[derive(Debug, Clone, PartialEq, Eq, Deserialize, Serialize)] -pub enum PbConnStatusReq { - RemoteId, - Keys, - Service { key: String }, -} - -#[derive(Debug, Clone, Deserialize, Serialize)] -pub enum PbConnStatusResp { - RemoteId { - server_map: String, - active: String, - idle: String, - }, - Keys(Vec), - Service { - key: String, - connections: Vec, - }, -} - -#[derive(Debug, Clone, PartialEq, Eq, Deserialize, Serialize)] -pub struct PbServiceConnStatus { - pub conn_id: u32, - pub generation: u64, - pub protocol_version: u16, - pub healthy: bool, - pub last_rx_age_ms: u64, -} - -#[derive(Debug, Clone, PartialEq, Eq, Deserialize, Serialize)] -pub enum PbConnRequest { - Register { - need_codec: bool, - is_datagram: bool, - key: String, - #[serde(default, skip_serializing_if = "Option::is_none")] - protocol_version: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - client_instance_id: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - heartbeat_interval_ms: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - heartbeat_tolerance_ms: Option, - }, - Subcribe { - key: String, - }, - Status(PbConnStatusReq), - Stream { - key: String, - dst_id: u32, - #[serde(default)] - server_generation: u64, - }, -} - -#[derive(Debug, Clone, Deserialize, Serialize)] -pub enum PbConnResponse { - Register(u32), - RegisterV2 { - conn_id: u32, - generation: u64, - lease_ttl_ms: u64, - }, - Subcribe { - codec_key: Option, - client_id: u32, - server_id: u32, - }, - Stream { - codec_key: Option, - }, - Status(PbConnStatusResp), -} - -#[derive(Debug, Clone, Deserialize, Serialize)] -pub enum PbServerRequest { - Ping, - PingV2 { - seq: u64, - }, - StreamAck { - client_id: u32, - #[serde(default)] - server_generation: u64, - }, -} - -#[derive(Debug, Clone, Deserialize, Serialize)] -pub enum LocalServer { - /// pb server makes a stream request to local server - Stream { - client_id: u32, - #[serde(default)] - server_generation: u64, - }, - /// pb server response a pong msg when it receive a ping request - Pong, - PongV2 { - seq: u64, - }, - Retire { - reason: String, - conn_id: u32, - #[serde(default)] - server_generation: u64, - }, -} - -const CONTENT_SHOW_LIMIT_SIZE: usize = 1024; - -#[inline] -fn get_content(raw_content: String) -> String { - if raw_content.len() > CONTENT_SHOW_LIMIT_SIZE { - raw_content[0..CONTENT_SHOW_LIMIT_SIZE].to_string() - } else { - raw_content - } -} - -macro_rules! gen_impl_msg_serializer { - ($struct_name:ident) => { - impl MessageSerializer for $struct_name { - fn encode(&self) -> Result> { - serde_json::to_vec(self).with_context(|_| MsgSerializeSnafu { - action: "encode", - struct_name: stringify!($struct_name), - content: get_content(format!("{self:?}")), - }) - } - - fn decode(msg: &[u8]) -> Result { - serde_json::from_slice(msg).with_context(|_| MsgSerializeSnafu { - action: "decode", - struct_name: stringify!($struct_name), - content: get_content(format!("{}", String::from_utf8_lossy(msg))), - }) - } - } - }; -} - -gen_impl_msg_serializer!(PbConnRequest); -gen_impl_msg_serializer!(PbConnResponse); -gen_impl_msg_serializer!(PbServerRequest); -gen_impl_msg_serializer!(LocalServer); diff --git a/src/common/mod.rs b/src/common/mod.rs deleted file mode 100644 index 7dad4ac..0000000 --- a/src/common/mod.rs +++ /dev/null @@ -1,7 +0,0 @@ -pub mod buffer; -pub mod checksum; -pub mod config; -pub mod conn_id; -pub mod error; -pub mod manager; -pub mod message; diff --git a/src/lib.rs b/src/lib.rs deleted file mode 100644 index 7bd3c9f..0000000 --- a/src/lib.rs +++ /dev/null @@ -1,29 +0,0 @@ -#[allow(async_fn_in_trait)] -pub mod common; -pub mod local; -pub mod pb_server; -pub mod utils; - -mod tests { - - #[test] - fn test_serde_mapper_header() { - use crate::common::message::command::PbConnRequest; - let mapper = PbConnRequest::Register { - key: "test".into(), - need_codec: false, - is_datagram: false, - protocol_version: None, - client_instance_id: None, - heartbeat_interval_ms: None, - heartbeat_tolerance_ms: None, - }; - let json_value = serde_json::to_string(&mapper).unwrap(); - let raw_json_str = - r##"{"Register":{"need_codec":false,"is_datagram":false,"key":"test"}}"##; - assert_eq!(raw_json_str, json_value); - - let value: PbConnRequest = serde_json::from_str(raw_json_str).unwrap(); - assert_eq!(mapper, value) - } -} diff --git a/src/local/client/status.rs b/src/local/client/status.rs deleted file mode 100644 index ea47acc..0000000 --- a/src/local/client/status.rs +++ /dev/null @@ -1,81 +0,0 @@ -use snafu::ResultExt; -use tokio::io::{AsyncReadExt, AsyncWriteExt}; - -use super::error::{ - ControlIoTimeoutSnafu, CreateHeaderToolSnafu, DecodeStatusRespSnafu, EncodeStatusReqSnafu, - ReadStatusRespSnafu, StatusRespNotMatchSnafu, WriteStatusReqSnafu, -}; -use crate::common::config::control_io_timeout; -use crate::common::message::command::{ - MessageSerializer, PbConnRequest, PbConnResponse, PbConnStatusReq, PbConnStatusResp, -}; -use crate::common::message::{ - get_header_msg_reader, get_header_msg_writer, MessageReader, MessageWriter, -}; - -pub async fn get_status( - remote_stream: &mut S, - req: PbConnStatusReq, -) -> super::error::Result { - let timeout = control_io_timeout(); - let msg = PbConnRequest::Status(req) - .encode() - .context(EncodeStatusReqSnafu)?; - - // send status request - { - let mut msg_writer = get_header_msg_writer(remote_stream) - .context(CreateHeaderToolSnafu { action: "writer" })?; - match tokio::time::timeout(timeout, msg_writer.write_msg(&msg)).await { - Ok(result) => result.context(WriteStatusReqSnafu)?, - Err(_) => ControlIoTimeoutSnafu { - action: "write status request", - timeout, - } - .fail()?, - } - } - - // get status - let mut msg_reader = - get_header_msg_reader(remote_stream).context(CreateHeaderToolSnafu { action: "reader" })?; - let msg = match tokio::time::timeout(timeout, msg_reader.read_msg()).await { - Ok(result) => result.context(ReadStatusRespSnafu)?, - Err(_) => ControlIoTimeoutSnafu { - action: "read status response", - timeout, - } - .fail()?, - }; - let resp = PbConnResponse::decode(msg).context(DecodeStatusRespSnafu)?; - match resp { - PbConnResponse::Status(status) => Ok(status), - _ => StatusRespNotMatchSnafu { - resp: String::from_utf8_lossy(msg), - } - .fail(), - } -} - -#[cfg(test)] -mod tests { - use std::time::Duration; - - use super::*; - - #[tokio::test] - async fn get_status_times_out_when_peer_stalls_after_request() { - std::env::set_var("PB_MAPPER_CONTROL_IO_TIMEOUT", "20ms"); - let (mut client, _server) = tokio::io::duplex(1024); - - let result = tokio::time::timeout( - Duration::from_millis(200), - get_status(&mut client, PbConnStatusReq::Keys), - ) - .await - .expect("get_status ignored PB_MAPPER_CONTROL_IO_TIMEOUT"); - - std::env::remove_var("PB_MAPPER_CONTROL_IO_TIMEOUT"); - assert!(result.is_err()); - } -} diff --git a/src/utils/mod.rs b/src/utils/mod.rs deleted file mode 100644 index 4313e7c..0000000 --- a/src/utils/mod.rs +++ /dev/null @@ -1,3 +0,0 @@ -pub mod addr; -pub mod codec; -pub mod timeout; diff --git a/ui/lib/l10n/app_en.arb b/ui/lib/l10n/app_en.arb index b58cf4f..6d56cf7 100644 --- a/ui/lib/l10n/app_en.arb +++ b/ui/lib/l10n/app_en.arb @@ -41,9 +41,9 @@ "setupServerBody": "Host and port. All traffic goes through it.", "setupServerHint": "example.com:7666", "setupServerInvalid": "Use host:port", - "setupKeyLabel": "Header key (optional)", - "setupKeyBody": "Only if your server was started with one.", - "setupKeyInvalid": "Must be exactly 32 characters", + "setupKeyLabel": "Credential", + "setupKeyBody": "Use a temporary credential issued by the relay administrator. The administrator key also works.", + "setupKeyInvalid": "Use a 32-character administrator key or a pbmt1_ temporary credential", "setupCheckingServer": "Checking…", "setupServerOk": "Server reached", "setupServerFailed": "Not reachable. You can continue anyway.", @@ -161,7 +161,10 @@ "saving": "Saving…", "checkServer": "Check Server Connectivity", "serverAddressHelp": "Address of the pb-mapper server to connect to", - "msgHeaderKeyHelp": "Used for the message checksum and encryption handshake. 32 characters when set.", + "msgHeaderKeyHelp": "Required. Prefer a pbmt1_ temporary credential issued by the relay; the 32-character administrator key also works.", + "isolatedRelayAdminKey": "Embedded relay administrator key", + "isolatedRelayAdminKeyHelp": "Generated for this app's local relay only. Use it to register or connect to that relay; it is not the outbound MSG_HEADER_KEY.", + "isolatedRelayReveal": "Reveal embedded relay key", "keepAliveHelp": "Enable TCP keep-alive for connections", "configServerAddress": "Server Address: {value}", "@configServerAddress": { @@ -206,7 +209,7 @@ }, "saveFailed": "Failed to save the configuration", "serverCheckFailed": "Server check failed", - "keyLengthInvalid": "The header key must be exactly 32 characters", + "keyLengthInvalid": "Use a 32-character administrator key or a pbmt1_ temporary credential", "registerTitle": "Register Service", "registerAction": "Register & Start", "registeredList": "Registered Services ({count})", diff --git a/ui/lib/l10n/app_zh.arb b/ui/lib/l10n/app_zh.arb index cf946cd..b938b9a 100644 --- a/ui/lib/l10n/app_zh.arb +++ b/ui/lib/l10n/app_zh.arb @@ -41,9 +41,9 @@ "setupServerBody": "地址和端口,所有流量都经过它。", "setupServerHint": "example.com:7666", "setupServerInvalid": "请用 host:port 格式", - "setupKeyLabel": "消息头密钥(可选)", - "setupKeyBody": "仅当服务器启动时带了密钥。", - "setupKeyInvalid": "必须是 32 个字符", + "setupKeyLabel": "连接凭据", + "setupKeyBody": "请使用中继管理员签发的临时凭据;管理员密钥也可以连接。", + "setupKeyInvalid": "请输入 32 字符管理员密钥或 pbmt1_ 临时凭据", "setupCheckingServer": "检测中…", "setupServerOk": "已连上服务器", "setupServerFailed": "连不上,也可以先继续。", @@ -161,7 +161,10 @@ "saving": "保存中…", "checkServer": "检测服务器连通性", "serverAddressHelp": "要连接的 pb-mapper 服务器地址", - "msgHeaderKeyHelp": "用于消息校验与加密握手,填写时必须是 32 个字符。", + "msgHeaderKeyHelp": "必填。优先使用中继签发的 pbmt1_ 临时凭据,也可使用 32 字符管理员密钥。", + "isolatedRelayAdminKey": "内嵌中继管理员密钥", + "isolatedRelayAdminKeyHelp": "仅用于本应用内嵌中继。用它向本地中继注册或连接,不要把它当成对外的 MSG_HEADER_KEY。", + "isolatedRelayReveal": "显示内嵌中继密钥", "keepAliveHelp": "为连接启用 TCP keep-alive", "configServerAddress": "服务器地址:{value}", "@configServerAddress": { @@ -206,7 +209,7 @@ }, "saveFailed": "保存配置失败", "serverCheckFailed": "服务器检测失败", - "keyLengthInvalid": "消息头密钥必须是 32 个字符", + "keyLengthInvalid": "请输入 32 字符管理员密钥或 pbmt1_ 临时凭据", "registerTitle": "注册服务", "registerAction": "注册并启动", "registeredList": "已注册服务({count})", diff --git a/ui/lib/src/ffi/pb_mapper_api.dart b/ui/lib/src/ffi/pb_mapper_api.dart index 454e4a6..fb099b7 100644 --- a/ui/lib/src/ffi/pb_mapper_api.dart +++ b/ui/lib/src/ffi/pb_mapper_api.dart @@ -16,11 +16,15 @@ class ConfigStatus { final String serverAddress; final bool keepAliveEnabled; final String msgHeaderKey; + final bool isolatedRelayAdminKeySet; + final String isolatedRelayAdminKey; const ConfigStatus({ required this.serverAddress, required this.keepAliveEnabled, required this.msgHeaderKey, + this.isolatedRelayAdminKeySet = false, + this.isolatedRelayAdminKey = '', }); factory ConfigStatus.fromMap(Map map) { @@ -31,6 +35,11 @@ class ConfigStatus { ), keepAliveEnabled: _asBool(map['keepAliveEnabled'], fallback: true), msgHeaderKey: _asString(map['msgHeaderKey']), + isolatedRelayAdminKeySet: _asBool( + map['isolatedRelayAdminKeySet'], + fallback: false, + ), + isolatedRelayAdminKey: _asString(map['isolatedRelayAdminKey']), ); } } @@ -241,6 +250,7 @@ abstract interface class PbMapperApiClient { Future setAppDirectoryPath(String path); Future fetchConfig(); + Future revealIsolatedRelayAdminKey(); Future updateConfig({ required String serverAddress, required bool keepAlive, @@ -310,6 +320,15 @@ class PbMapperApi implements PbMapperApiClient { ); } + @override + Future revealIsolatedRelayAdminKey() async { + final result = await _service.revealIsolatedRelayAdminKey(); + if (result['success'] == true) { + return _asString(_asMap(result['data'])['isolatedRelayAdminKey']); + } + return ''; + } + @override Future updateConfig({ required String serverAddress, diff --git a/ui/lib/src/ffi/pb_mapper_ffi.dart b/ui/lib/src/ffi/pb_mapper_ffi.dart index 30b9529..b0b38f5 100644 --- a/ui/lib/src/ffi/pb_mapper_ffi.dart +++ b/ui/lib/src/ffi/pb_mapper_ffi.dart @@ -296,6 +296,11 @@ class PbMapperFFI { 'pb_mapper_get_config_json', ); + late final pbMapperRevealIsolatedAdminKey = lib + .lookupFunction<_PbMapperGetConfigNative, _PbMapperGetConfigDart>( + 'pb_mapper_reveal_isolated_admin_key', + ); + late final pbMapperUpdateConfig = lib .lookupFunction<_PbMapperUpdateConfigNative, _PbMapperUpdateConfigDart>( 'pb_mapper_update_config', diff --git a/ui/lib/src/ffi/pb_mapper_service.dart b/ui/lib/src/ffi/pb_mapper_service.dart index 2054d34..fe25a8f 100644 --- a/ui/lib/src/ffi/pb_mapper_service.dart +++ b/ui/lib/src/ffi/pb_mapper_service.dart @@ -240,6 +240,10 @@ class PbMapperService { return _runJsonOnWorker('getConfig', {}); } + Future> revealIsolatedRelayAdminKey() { + return _runJsonOnWorker('revealIsolatedRelayAdminKey', {}); + } + Future> updateConfig({ required String serverAddress, required bool keepAlive, @@ -383,6 +387,9 @@ Map _callJsonIsolate(Map params) { case 'getConfig': result = ffi.pbMapperGetConfig(handle); break; + case 'revealIsolatedRelayAdminKey': + result = ffi.pbMapperRevealIsolatedAdminKey(handle); + break; case 'updateConfig': arg1 = (params['serverAddress'] as String).toNativeUtf8(); arg2 = (params['msgHeaderKey'] as String).toNativeUtf8(); diff --git a/ui/lib/src/views/configuration_view.dart b/ui/lib/src/views/configuration_view.dart index de7059a..a639f83 100644 --- a/ui/lib/src/views/configuration_view.dart +++ b/ui/lib/src/views/configuration_view.dart @@ -32,14 +32,12 @@ class _ConfigurationViewState extends State { bool? _serverReachable; String _serverCheckMessage = ''; ConfigStatus? _currentConfig; + String _revealedIsolatedRelayAdminKey = ''; ChangeSubscription? _changes; - @override - void initState() { - super.initState(); // Reload when anything changes this list, including a change made @@ -47,20 +45,19 @@ class _ConfigurationViewState extends State { // from a terminal while this window was open. _changes = ChangeSubscription.listen( - PbMapperService.changeStream, {StateChangeKind.config}, - (_) { if (mounted) _loadConfig(); }, - + (_) { + if (mounted) _loadConfig(); + }, ); _loadConfig(); } @override void dispose() { - _changes?.cancel(); _serverAddressController.dispose(); _msgHeaderKeyController.dispose(); @@ -90,7 +87,9 @@ class _ConfigurationViewState extends State { Future _saveConfiguration() async { if (_isSaving) return; // Prevent multiple simultaneous saves final msgHeaderKey = _msgHeaderKeyController.text.trim(); - if (msgHeaderKey.isNotEmpty && msgHeaderKey.length != 32) { + if (msgHeaderKey.isNotEmpty && + msgHeaderKey.length != 32 && + !msgHeaderKey.startsWith('pbmt1_')) { showToast(context, context.l10n.keyLengthInvalid, kind: ToastKind.error); return; } @@ -127,6 +126,14 @@ class _ConfigurationViewState extends State { } } + Future _revealIsolatedRelayAdminKey() async { + final key = await _api.revealIsolatedRelayAdminKey(); + if (!mounted) return; + setState(() { + _revealedIsolatedRelayAdminKey = key; + }); + } + Future _checkServerConnection() async { if (_isCheckingServer) return; setState(() { @@ -308,9 +315,11 @@ class _ConfigurationViewState extends State { if (serverAddress.isEmpty) { throw const FormatException('serverAddress is required'); } - if (msgHeaderKey.isNotEmpty && msgHeaderKey.length != 32) { + if (msgHeaderKey.isNotEmpty && + msgHeaderKey.length != 32 && + !msgHeaderKey.startsWith('pbmt1_')) { throw const FormatException( - 'MSG_HEADER_KEY must be exactly 32 characters', + 'MSG_HEADER_KEY must be a 32-character administrator key or a pbmt1_ temporary credential', ); } @@ -364,11 +373,28 @@ class _ConfigurationViewState extends State { controller: _msgHeaderKeyController, decoration: InputDecoration( labelText: 'MSG_HEADER_KEY', - hintText: '32 characters, or empty', + hintText: '32-character admin key or pbmt1_ credential', border: OutlineInputBorder(), helperText: context.l10n.msgHeaderKeyHelp, ), ), + if (_currentConfig?.isolatedRelayAdminKeySet == true) ...[ + const SizedBox(height: 16), + if (_revealedIsolatedRelayAdminKey.isEmpty) + OutlinedButton( + onPressed: _revealIsolatedRelayAdminKey, + child: Text(context.l10n.isolatedRelayReveal), + ) + else + InputDecorator( + decoration: InputDecoration( + labelText: context.l10n.isolatedRelayAdminKey, + border: const OutlineInputBorder(), + helperText: context.l10n.isolatedRelayAdminKeyHelp, + ), + child: SelectableText(_revealedIsolatedRelayAdminKey), + ), + ], const SizedBox(height: 16), SwitchListTile( title: const Text('PB_MAPPER_KEEP_ALIVE'), diff --git a/ui/lib/src/views/setup_wizard_view.dart b/ui/lib/src/views/setup_wizard_view.dart index 364c3e8..51aa3ed 100644 --- a/ui/lib/src/views/setup_wizard_view.dart +++ b/ui/lib/src/views/setup_wizard_view.dart @@ -158,7 +158,7 @@ class _SetupWizardViewState extends State { setState(() => _error = l10n.setupServerInvalid); return; } - if (key.isNotEmpty && key.length != 32) { + if (key.isNotEmpty && key.length != 32 && !key.startsWith('pbmt1_')) { setState(() => _error = l10n.setupKeyInvalid); return; } diff --git a/ui/native/pb_mapper_ffi/Cargo.toml b/ui/native/pb_mapper_ffi/Cargo.toml index 2b30fdc..be8a7aa 100644 --- a/ui/native/pb_mapper_ffi/Cargo.toml +++ b/ui/native/pb_mapper_ffi/Cargo.toml @@ -1,14 +1,13 @@ [package] name = "pb-mapper-ffi" version = "0.1.0" -edition = "2021" +edition = "2024" [lib] crate-type = ["cdylib", "staticlib"] -[lints.clippy] -unwrap_used = "deny" -expect_used = "deny" +[lints] +workspace = true [dependencies] serde = { version = "1.0.219", features = ["derive"] } @@ -18,8 +17,13 @@ tracing-subscriber.workspace = true tokio.workspace = true clap.workspace = true better_mimalloc_rs.workspace = true -tokio-util = "0.7" -dirs = "5.0" +tokio-util.workspace = true +dirs.workspace = true -pb-mapper = { path = "../../../" } +pb-mapper-auth.workspace = true +pb-mapper-client.workspace = true +pb-mapper-core.workspace = true +pb-mapper-protocol.workspace = true +pb-mapper-server.workspace = true uni-stream.workspace = true +parking_lot.workspace = true diff --git a/ui/native/pb_mapper_ffi/src/cli.rs b/ui/native/pb_mapper_ffi/src/cli.rs index e270994..d749c4e 100644 --- a/ui/native/pb_mapper_ffi/src/cli.rs +++ b/ui/native/pb_mapper_ffi/src/cli.rs @@ -11,12 +11,12 @@ //! Dart's whole part is: hand argv over, print nothing, exit with what comes //! back. -use std::ffi::{c_char, c_int, CStr}; +use std::ffi::{CStr, c_char, c_int}; use clap::{CommandFactory, Parser}; use crate::ctl::proto::Response; -use crate::ctl::{endpoint, server, Command}; +use crate::ctl::{Command, endpoint, server}; /// Returned when argv is not a command at all, so Dart knows to run the GUI. /// Chosen so it cannot collide with a real exit code. @@ -143,7 +143,7 @@ fn run(args: Vec) -> c_int { /// /// # Safety /// `argv` must point to `argc` valid, NUL-terminated C strings. -#[no_mangle] +#[unsafe(no_mangle)] pub unsafe extern "C" fn pb_mapper_cli_main(argc: c_int, argv: *const *const c_char) -> c_int { if argv.is_null() || argc <= 0 { return NOT_A_COMMAND; diff --git a/ui/native/pb_mapper_ffi/src/client.rs b/ui/native/pb_mapper_ffi/src/client.rs index 7db2abd..3caec10 100644 --- a/ui/native/pb_mapper_ffi/src/client.rs +++ b/ui/native/pb_mapper_ffi/src/client.rs @@ -12,7 +12,7 @@ use crate::response::{err_ctl, err_null_handle, ok_data, ok_message, parse_c_str use crate::state::{ClientConfigInfo, ClientStatusResponse}; /// Connect client to service. -#[no_mangle] +#[unsafe(no_mangle)] pub unsafe extern "C" fn pb_mapper_connect_service( handle: *mut PbMapperHandle, service_key: *const c_char, @@ -65,7 +65,7 @@ pub unsafe extern "C" fn pb_mapper_connect_service( } /// Disconnect client from service. -#[no_mangle] +#[unsafe(no_mangle)] pub unsafe extern "C" fn pb_mapper_disconnect_service( handle: *mut PbMapperHandle, service_key: *const c_char, @@ -96,7 +96,7 @@ pub unsafe extern "C" fn pb_mapper_disconnect_service( } /// Delete client config (also stops client if running). -#[no_mangle] +#[unsafe(no_mangle)] pub unsafe extern "C" fn pb_mapper_delete_client_config( handle: *mut PbMapperHandle, service_key: *const c_char, @@ -127,7 +127,7 @@ pub unsafe extern "C" fn pb_mapper_delete_client_config( } /// Get client configs list. -#[no_mangle] +#[unsafe(no_mangle)] pub unsafe extern "C" fn pb_mapper_get_client_configs_json( handle: *mut PbMapperHandle, ) -> *mut c_char { @@ -146,7 +146,7 @@ pub unsafe extern "C" fn pb_mapper_get_client_configs_json( } /// Get client status for a specific key. -#[no_mangle] +#[unsafe(no_mangle)] pub unsafe extern "C" fn pb_mapper_get_client_status_json( handle: *mut PbMapperHandle, service_key: *const c_char, diff --git a/ui/native/pb_mapper_ffi/src/config.rs b/ui/native/pb_mapper_ffi/src/config.rs index 2ba14ed..b6f542a 100644 --- a/ui/native/pb_mapper_ffi/src/config.rs +++ b/ui/native/pb_mapper_ffi/src/config.rs @@ -9,10 +9,9 @@ use crate::ctl::Origin; use crate::events; use crate::handle::PbMapperHandle; use crate::response::{err_ctl, err_null_handle, ok_data, ok_message, parse_c_string}; -use crate::state::AppConfig; /// Get current app config. -#[no_mangle] +#[unsafe(no_mangle)] pub unsafe extern "C" fn pb_mapper_get_config_json(handle: *mut PbMapperHandle) -> *mut c_char { if handle.is_null() { return err_null_handle(); @@ -20,20 +19,46 @@ pub unsafe extern "C" fn pb_mapper_get_config_json(handle: *mut PbMapperHandle) let handle = unsafe { &mut *handle }; let state = handle.state.clone(); - let config: AppConfig = handle.runtime.block_on(async move { + let (config, isolated_admin_key_set) = handle.runtime.block_on(async move { let state = state.lock().await; - state.get_config_status().await + ( + state.get_config_status().await, + state.isolated_admin_key().is_some(), + ) }); ok_data(json!({ "serverAddress": config.server_address, "keepAliveEnabled": config.keep_alive_enabled, - "msgHeaderKey": config.msg_header_key + "msgHeaderKey": config.msg_header_key, + "isolatedRelayAdminKeySet": isolated_admin_key_set, + })) +} + +/// Reveal the embedded relay administrator key. This is a separate call so +/// routine config fetches cannot leak the root secret. +#[unsafe(no_mangle)] +pub unsafe extern "C" fn pb_mapper_reveal_isolated_admin_key( + handle: *mut PbMapperHandle, +) -> *mut c_char { + if handle.is_null() { + return err_null_handle(); + } + + let handle = unsafe { &mut *handle }; + let state = handle.state.clone(); + let key = handle.runtime.block_on(async move { + let state = state.lock().await; + state.isolated_admin_key() + }); + + ok_data(json!({ + "isolatedRelayAdminKey": key.unwrap_or_default(), })) } /// Update app config. -#[no_mangle] +#[unsafe(no_mangle)] pub unsafe extern "C" fn pb_mapper_update_config( handle: *mut PbMapperHandle, server_address: *const c_char, diff --git a/ui/native/pb_mapper_ffi/src/ctl/endpoint.rs b/ui/native/pb_mapper_ffi/src/ctl/endpoint.rs index 27e1e23..b808c2c 100644 --- a/ui/native/pb_mapper_ffi/src/ctl/endpoint.rs +++ b/ui/native/pb_mapper_ffi/src/ctl/endpoint.rs @@ -105,10 +105,10 @@ mod imp { } pub fn endpoint() -> String { - if let Ok(custom) = std::env::var(super::ENDPOINT_ENV) { - if !custom.is_empty() { - return custom; - } + if let Ok(custom) = std::env::var(super::ENDPOINT_ENV) + && !custom.is_empty() + { + return custom; } // Per user, so two accounts on one machine do not collide. let user = std::env::var("USERNAME").unwrap_or_else(|_| "default".into()); @@ -160,10 +160,10 @@ mod imp { const SUN_PATH_MAX: usize = 100; pub fn endpoint() -> String { - if let Ok(custom) = std::env::var(super::ENDPOINT_ENV) { - if !custom.is_empty() { - return custom; - } + if let Ok(custom) = std::env::var(super::ENDPOINT_ENV) + && !custom.is_empty() + { + return custom; } // XDG_RUNTIME_DIR is per-user and cleaned on logout, which is what a // socket wants. TMPDIR is the macOS equivalent and matters more there @@ -195,7 +195,7 @@ mod imp { format!("/tmp/pb-mapper-ui-{uid}.sock") } - extern "C" { + unsafe extern "C" { #[link_name = "getuid"] fn libc_getuid() -> u32; } diff --git a/ui/native/pb_mapper_ffi/src/ctl/mod.rs b/ui/native/pb_mapper_ffi/src/ctl/mod.rs index 1e4866e..034b9a7 100644 --- a/ui/native/pb_mapper_ffi/src/ctl/mod.rs +++ b/ui/native/pb_mapper_ffi/src/ctl/mod.rs @@ -206,12 +206,15 @@ async fn run( } Command::ConfigGet => { - let config = state.lock().await.get_config_status().await; + let guard = state.lock().await; + let config = guard.get_config_status().await; + let isolated_admin_key = guard.isolated_admin_key(); Ok(proto::Response::ok( Some(json!({ "serverAddress": config.server_address, "keepAliveEnabled": config.keep_alive_enabled, "msgHeaderKeySet": !config.msg_header_key.is_empty(), + "isolatedRelayAdminKeySet": isolated_admin_key.is_some(), })), None, )) diff --git a/ui/native/pb_mapper_ffi/src/ctl/server.rs b/ui/native/pb_mapper_ffi/src/ctl/server.rs index a4570e7..a4ac90a 100644 --- a/ui/native/pb_mapper_ffi/src/ctl/server.rs +++ b/ui/native/pb_mapper_ffi/src/ctl/server.rs @@ -11,8 +11,8 @@ use tokio::sync::Mutex; use tokio_util::sync::CancellationToken; use crate::ctl::endpoint; -use crate::ctl::proto::{self, Request, Response, PROTOCOL_VERSION}; -use crate::ctl::{dispatch, Origin}; +use crate::ctl::proto::{self, PROTOCOL_VERSION, Request, Response}; +use crate::ctl::{Origin, dispatch}; use crate::error::CtlError; use crate::state::PbMapperState; diff --git a/ui/native/pb_mapper_ffi/src/events.rs b/ui/native/pb_mapper_ffi/src/events.rs index d05abfa..a5521e4 100644 --- a/ui/native/pb_mapper_ffi/src/events.rs +++ b/ui/native/pb_mapper_ffi/src/events.rs @@ -7,7 +7,7 @@ //! mean taking the state lock on the emit path and guessing which projection //! the receiver wants. -use std::ffi::{c_char, CString}; +use std::ffi::{CString, c_char}; use std::sync::atomic::{AtomicU64, Ordering}; use serde::Serialize; @@ -71,7 +71,7 @@ pub fn emit(kind: ChangeKind, key: Option<&str>, origin: Origin) { /// /// # Safety /// `callback` must stay valid until it is replaced or cleared with null. -#[no_mangle] +#[unsafe(no_mangle)] pub extern "C" fn pb_mapper_set_change_callback(callback: Option) { CHANGE_CALLBACK.store(callback); } diff --git a/ui/native/pb_mapper_ffi/src/handle.rs b/ui/native/pb_mapper_ffi/src/handle.rs index b97a07e..a6ada13 100644 --- a/ui/native/pb_mapper_ffi/src/handle.rs +++ b/ui/native/pb_mapper_ffi/src/handle.rs @@ -1,7 +1,7 @@ //! FFI handle lifecycle and app directory configuration. #![allow(clippy::missing_safety_doc)] -use std::ffi::{c_char, CStr}; +use std::ffi::{CStr, c_char}; use std::ptr; use std::sync::Arc; @@ -31,7 +31,7 @@ impl Drop for PbMapperHandle { /// /// # Safety /// Returns a pointer that must be freed with `pb_mapper_destroy`. -#[no_mangle] +#[unsafe(no_mangle)] pub extern "C" fn pb_mapper_create() -> *mut PbMapperHandle { let runtime = match Runtime::new() { Ok(rt) => rt, @@ -57,7 +57,7 @@ pub extern "C" fn pb_mapper_create() -> *mut PbMapperHandle { /// /// # Safety /// `handle` must be a valid pointer returned by `pb_mapper_create`. -#[no_mangle] +#[unsafe(no_mangle)] pub unsafe extern "C" fn pb_mapper_start_control_server( handle: *mut PbMapperHandle, ) -> *mut c_char { @@ -82,7 +82,7 @@ pub unsafe extern "C" fn pb_mapper_start_control_server( /// /// # Safety /// `handle` must be a valid pointer returned by `pb_mapper_create`. -#[no_mangle] +#[unsafe(no_mangle)] pub unsafe extern "C" fn pb_mapper_destroy(handle: *mut PbMapperHandle) { if !handle.is_null() { unsafe { drop(Box::from_raw(handle)) }; @@ -93,7 +93,7 @@ pub unsafe extern "C" fn pb_mapper_destroy(handle: *mut PbMapperHandle) { /// /// # Safety /// `handle` must be valid. `path` must be valid C string or null. -#[no_mangle] +#[unsafe(no_mangle)] pub unsafe extern "C" fn pb_mapper_set_app_dir( handle: *mut PbMapperHandle, path: *const c_char, diff --git a/ui/native/pb_mapper_ffi/src/lib.rs b/ui/native/pb_mapper_ffi/src/lib.rs index c3db9e8..5853746 100644 --- a/ui/native/pb_mapper_ffi/src/lib.rs +++ b/ui/native/pb_mapper_ffi/src/lib.rs @@ -17,16 +17,18 @@ mod state; // Re-export public FFI functions and handle type. use better_mimalloc_rs::MiMalloc; -pub use cli::{pb_mapper_cli_main, NOT_A_COMMAND}; +pub use cli::{NOT_A_COMMAND, pb_mapper_cli_main}; pub use client::{ pb_mapper_connect_service, pb_mapper_delete_client_config, pb_mapper_disconnect_service, pb_mapper_get_client_configs_json, pb_mapper_get_client_status_json, }; -pub use config::{pb_mapper_get_config_json, pb_mapper_update_config}; +pub use config::{ + pb_mapper_get_config_json, pb_mapper_reveal_isolated_admin_key, pb_mapper_update_config, +}; pub use events::pb_mapper_set_change_callback; pub use handle::{ - pb_mapper_create, pb_mapper_destroy, pb_mapper_set_app_dir, pb_mapper_start_control_server, - PbMapperHandle, + PbMapperHandle, pb_mapper_create, pb_mapper_destroy, pb_mapper_set_app_dir, + pb_mapper_start_control_server, }; pub use logging::{pb_mapper_free_string, pb_mapper_init_logging, pb_mapper_set_log_callback}; pub use server::{ diff --git a/ui/native/pb_mapper_ffi/src/logging.rs b/ui/native/pb_mapper_ffi/src/logging.rs index 2419029..2768e21 100644 --- a/ui/native/pb_mapper_ffi/src/logging.rs +++ b/ui/native/pb_mapper_ffi/src/logging.rs @@ -1,6 +1,6 @@ //! Logging system for FFI interface. -use std::ffi::{c_char, c_int, CString}; +use std::ffi::{CString, c_char, c_int}; use std::time::{SystemTime, UNIX_EPOCH}; use crate::callback::CallbackSlot; @@ -34,7 +34,7 @@ pub(crate) fn send_log(level: c_int, message: &str) { /// /// # Safety /// `callback` must be a valid function pointer or null to disable logging. -#[no_mangle] +#[unsafe(no_mangle)] pub extern "C" fn pb_mapper_set_log_callback(callback: Option) { LOG_CALLBACK.store(callback); } @@ -43,7 +43,7 @@ pub extern "C" fn pb_mapper_set_log_callback(callback: Option) { /// /// # Safety /// `s` must be a valid pointer returned from this library, or null. -#[no_mangle] +#[unsafe(no_mangle)] pub unsafe extern "C" fn pb_mapper_free_string(s: *mut c_char) { if !s.is_null() { unsafe { drop(CString::from_raw(s)) }; @@ -112,11 +112,11 @@ impl tracing::field::Visit for MessageVisitor { /// /// # Safety /// Can be called multiple times safely. -#[no_mangle] +#[unsafe(no_mangle)] pub extern "C" fn pb_mapper_init_logging() { + use tracing_subscriber::Layer; use tracing_subscriber::layer::SubscriberExt; use tracing_subscriber::util::SubscriberInitExt; - use tracing_subscriber::Layer; let _ = tracing_subscriber::registry() .with(FfiLogLayer) diff --git a/ui/native/pb_mapper_ffi/src/response.rs b/ui/native/pb_mapper_ffi/src/response.rs index aa92880..180ae4a 100644 --- a/ui/native/pb_mapper_ffi/src/response.rs +++ b/ui/native/pb_mapper_ffi/src/response.rs @@ -1,6 +1,6 @@ //! Shared helpers for FFI response formatting and argument parsing. -use std::ffi::{c_char, CStr, CString}; +use std::ffi::{CStr, CString, c_char}; use std::ptr; use serde_json::json; diff --git a/ui/native/pb_mapper_ffi/src/server.rs b/ui/native/pb_mapper_ffi/src/server.rs index ef0df8f..aac58c2 100644 --- a/ui/native/pb_mapper_ffi/src/server.rs +++ b/ui/native/pb_mapper_ffi/src/server.rs @@ -11,7 +11,7 @@ use crate::handle::PbMapperHandle; use crate::response::{err_ctl, err_null_handle, ok_data, ok_message, parse_c_string}; /// Start pb-mapper server. -#[no_mangle] +#[unsafe(no_mangle)] pub unsafe extern "C" fn pb_mapper_start_server( handle: *mut PbMapperHandle, port: u16, @@ -38,7 +38,7 @@ pub unsafe extern "C" fn pb_mapper_start_server( } /// Stop pb-mapper server. -#[no_mangle] +#[unsafe(no_mangle)] pub unsafe extern "C" fn pb_mapper_stop_server(handle: *mut PbMapperHandle) -> *mut c_char { if handle.is_null() { return err_null_handle(); @@ -61,7 +61,7 @@ pub unsafe extern "C" fn pb_mapper_stop_server(handle: *mut PbMapperHandle) -> * } /// Get local server status (running/uptime). -#[no_mangle] +#[unsafe(no_mangle)] pub unsafe extern "C" fn pb_mapper_get_local_server_status_json( handle: *mut PbMapperHandle, ) -> *mut c_char { @@ -84,7 +84,7 @@ pub unsafe extern "C" fn pb_mapper_get_local_server_status_json( /// The status detail's `serverMap` is a Debug dump of the whole map and is not /// something a UI should be parsing. This answers the same question with the /// protocol's own structured query, one key at a time. -#[no_mangle] +#[unsafe(no_mangle)] pub unsafe extern "C" fn pb_mapper_get_service_conns_json( handle: *mut PbMapperHandle, service_key: *const c_char, @@ -112,7 +112,7 @@ pub unsafe extern "C" fn pb_mapper_get_service_conns_json( } /// Get server status detail (remote server). -#[no_mangle] +#[unsafe(no_mangle)] pub unsafe extern "C" fn pb_mapper_get_server_status_detail_json( handle: *mut PbMapperHandle, ) -> *mut c_char { @@ -134,7 +134,7 @@ pub unsafe extern "C" fn pb_mapper_get_server_status_detail_json( } /// Force-refresh server status (blocks until network result is available). -#[no_mangle] +#[unsafe(no_mangle)] pub unsafe extern "C" fn pb_mapper_force_refresh_server_status_json( handle: *mut PbMapperHandle, ) -> *mut c_char { diff --git a/ui/native/pb_mapper_ffi/src/service.rs b/ui/native/pb_mapper_ffi/src/service.rs index 5b479d6..9dc4d53 100644 --- a/ui/native/pb_mapper_ffi/src/service.rs +++ b/ui/native/pb_mapper_ffi/src/service.rs @@ -12,7 +12,7 @@ use crate::response::{err_ctl, err_null_handle, ok_data, ok_message, parse_c_str use crate::state::{ServiceConfigInfo, ServiceStatusResponse}; /// Register a service. -#[no_mangle] +#[unsafe(no_mangle)] pub unsafe extern "C" fn pb_mapper_register_service( handle: *mut PbMapperHandle, service_key: *const c_char, @@ -65,7 +65,7 @@ pub unsafe extern "C" fn pb_mapper_register_service( } /// Unregister a service (stop running but keep config). -#[no_mangle] +#[unsafe(no_mangle)] pub unsafe extern "C" fn pb_mapper_unregister_service( handle: *mut PbMapperHandle, service_key: *const c_char, @@ -96,7 +96,7 @@ pub unsafe extern "C" fn pb_mapper_unregister_service( } /// Delete service config (also stops service if running). -#[no_mangle] +#[unsafe(no_mangle)] pub unsafe extern "C" fn pb_mapper_delete_service_config( handle: *mut PbMapperHandle, service_key: *const c_char, @@ -127,7 +127,7 @@ pub unsafe extern "C" fn pb_mapper_delete_service_config( } /// Get service configs list. -#[no_mangle] +#[unsafe(no_mangle)] pub unsafe extern "C" fn pb_mapper_get_service_configs_json( handle: *mut PbMapperHandle, ) -> *mut c_char { @@ -146,7 +146,7 @@ pub unsafe extern "C" fn pb_mapper_get_service_configs_json( } /// Get service status for a specific key. -#[no_mangle] +#[unsafe(no_mangle)] pub unsafe extern "C" fn pb_mapper_get_service_status_json( handle: *mut PbMapperHandle, service_key: *const c_char, diff --git a/ui/native/pb_mapper_ffi/src/state.rs b/ui/native/pb_mapper_ffi/src/state.rs index 71d4bf5..6c437b2 100644 --- a/ui/native/pb_mapper_ffi/src/state.rs +++ b/ui/native/pb_mapper_ffi/src/state.rs @@ -1,27 +1,42 @@ +//! Shared Flutter-FFI application state and module boundaries. +//! +//! ```text +//! Flutter command -> Arc> -> configuration / runtime / status +//! -> change events back to Flutter +//! ``` +//! +//! Slow DNS, bind, and connectivity work is deliberately performed outside the global +//! state lock. Per-key claims prevent duplicate setup while keeping unrelated UI reads +//! and operations responsive. + use std::collections::{HashMap, HashSet}; use std::fs; use std::net::{IpAddr, Ipv4Addr, SocketAddr}; use std::path::PathBuf; +use std::sync::Arc; use std::sync::atomic::{AtomicBool, Ordering}; -use std::sync::{Arc, Mutex as StdMutex}; use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH}; +use parking_lot::Mutex as SyncMutex; use serde::{Deserialize, Serialize}; use tokio::net::{TcpListener, TcpStream}; use tokio::sync::{Mutex, RwLock}; use tokio::task::JoinHandle; use tokio_util::sync::CancellationToken; -use pb_mapper::common::checksum::set_process_msg_header_key; -use pb_mapper::common::config::{get_pb_mapper_server_async, get_sockaddr_async}; -use pb_mapper::common::message::command::{PbConnStatusReq, PbConnStatusResp}; -use pb_mapper::local::client::status::get_status; -use pb_mapper::local::client::{run_client_side_cli_with_callback, ClientStatusCallback}; -use pb_mapper::local::server::{ - run_server_side_cli_with_callback, ServerTunnelOptions, StatusCallback, +use pb_mapper_auth::{AuthConfig, AuthRuntime}; +use pb_mapper_client::client::status::{get_status, get_status_with_credential}; +use pb_mapper_client::client::{ClientStatusCallback, run_client_side_cli_with_pinned_credential}; +use pb_mapper_client::server::{ + ServerTunnelOptions, StatusCallback, run_server_side_cli_with_pinned_credential, +}; +use pb_mapper_core::addr::each_addr; +use pb_mapper_core::checksum::{ + Credential, get_process_credential, parse_credential, set_process_msg_header_key, }; -use pb_mapper::pb_server::{run_server_with_shutdown, ServerStatusInfo}; -use pb_mapper::utils::addr::each_addr; +use pb_mapper_core::config::{get_pb_mapper_server_async, get_sockaddr_async}; +use pb_mapper_protocol::command::{PbConnStatusReq, PbConnStatusResp}; +use pb_mapper_server::{ServerStatusInfo, run_server_on_listener}; use uni_stream::stream::got_one_socket_addr; use uni_stream::stream::{ ListenerProvider, StreamProvider, TcpListenerProvider, TcpStreamProvider, UdpListenerProvider, @@ -46,6 +61,7 @@ struct StatusCacheEntry { async fn check_service_with_get_status( server_addr: &str, service_key: &str, + credential: Option, ) -> Result { let addr = get_sockaddr_async(server_addr) .await @@ -54,7 +70,13 @@ async fn check_service_with_get_status( match TcpStreamProvider::from_addr(addr).await { Ok(mut stream) => { let status_req = PbConnStatusReq::Keys; - match get_status(&mut stream, status_req).await { + let status = match credential { + Some(credential) => { + get_status_with_credential(&mut stream, status_req, None, &credential).await + } + None => get_status(&mut stream, status_req).await, + }; + match status { Ok(status_resp) => match status_resp { PbConnStatusResp::Keys(keys) => { if keys.contains(&service_key.to_string()) { @@ -205,11 +227,7 @@ fn normalize_msg_header_key(msg_header_key: String) -> Result if normalized.is_empty() { return Ok(normalized); } - if normalized.len() != 32 { - return Err(CtlError::invalid_argument( - "MSG_HEADER_KEY must be exactly 32 bytes (256-bit) when provided", - )); - } + parse_credential(&normalized).map_err(CtlError::invalid_argument)?; Ok(normalized) } @@ -371,30 +389,26 @@ struct ConnectionInfo { /// the first's [`JoinHandle`], leaving a tunnel running that nothing could /// abort. Holding the key across the gap is what closes it again. /// -/// The set is behind a `std::sync::Mutex` rather than tokio's so that `Drop` can +/// The set is behind a blocking mutex rather than tokio's so that `Drop` can /// release it; it is only ever held for a set insert or remove. struct KeyClaim { key: String, - claims: Arc>>, + claims: Arc>>, } impl Drop for KeyClaim { fn drop(&mut self) { - if let Ok(mut claims) = self.claims.lock() { - claims.remove(&self.key); - } + self.claims.lock().remove(&self.key); } } /// Claims `key`, or reports that someone else is already setting it up. fn claim_key( - claims: &Arc>>, + claims: &Arc>>, key: &str, what: &str, ) -> Result { - let mut guard = claims - .lock() - .map_err(|_| CtlError::internal(format!("{what} state for '{key}' is poisoned")))?; + let mut guard = claims.lock(); if !guard.insert(key.to_string()) { return Err(CtlError::already_in_progress(format!( "'{key}' is already {what}" @@ -406,6 +420,33 @@ fn claim_key( }) } +#[derive(Clone, Copy)] +struct PinnedTunnel { + credential: Credential, + endpoint: SocketAddr, +} + +struct TunnelRuntime { + handle: JoinHandle<()>, + pin: PinnedTunnel, +} + +fn abort_runtime(map: &mut HashMap, key: &str) -> bool { + match map.remove(key) { + Some(runtime) => { + runtime.handle.abort(); + true + } + None => false, + } +} + +fn replace_runtime(map: &mut HashMap, key: &str, runtime: TunnelRuntime) { + if let Some(previous) = map.insert(key.to_string(), runtime) { + previous.handle.abort(); + } +} + /// Everything [`PbMapperState::finish_register`] needs once the slow work is done. struct RegisterCommit { service_key: String, @@ -415,6 +456,7 @@ struct RegisterCommit { enable_keep_alive: bool, local_sock_addr: SocketAddr, remote_sock_addr: SocketAddr, + credential: Credential, } /// Everything [`PbMapperState::finish_connect`] needs once the slow work is done. @@ -425,18 +467,20 @@ struct ConnectCommit { enable_keep_alive: bool, local_sock_addr: SocketAddr, remote_sock_addr: SocketAddr, + credential: Credential, } pub struct PbMapperState { server_handle: Option>, + server_auth: Option, server_shutdown_token: Option, server_status_sender: Option>>, server_start_time: Option, registered_services: Arc>>, active_connections: Arc>>, - service_handles: HashMap>, - client_handles: HashMap>, + service_runtime: HashMap, + client_runtime: HashMap, config: AppConfig, config_dir: PathBuf, app_directory_path: Option, @@ -449,1210 +493,13 @@ pub struct PbMapperState { client_status_refreshing: Arc>>, /// Keys currently being set up. See [`KeyClaim`]. Separate sets because a /// key can legitimately be registered and connected to at the same time. - registering: Arc>>, - connecting: Arc>>, + registering: Arc>>, + connecting: Arc>>, } -impl PbMapperState { - async fn reset_status_caches(&self) { - { - let mut cache = self.local_server_status_cache.write().await; - *cache = LocalServerStatus { - is_running: false, - active_connections: 0, - registered_services: 0, - uptime_seconds: 0, - }; - } - { - let mut last_update = self.local_server_status_last_update.write().await; - *last_update = None; - } - self.local_server_status_refreshing - .store(false, Ordering::Release); - - self.service_status_cache.write().await.clear(); - self.client_status_cache.write().await.clear(); - self.service_status_refreshing.write().await.clear(); - self.client_status_refreshing.write().await.clear(); - } - pub fn new(app_directory_path: Option) -> Self { - let config_dir = Self::get_config_dir(&app_directory_path); - tracing::info!("Using config directory: {:?}", config_dir); - - let local_server_status_cache = Arc::new(RwLock::new(LocalServerStatus { - is_running: false, - active_connections: 0, - registered_services: 0, - uptime_seconds: 0, - })); - - let temp_state = Self { - server_handle: None, - server_shutdown_token: None, - server_status_sender: None, - server_start_time: None, - registered_services: Arc::new(RwLock::new(HashMap::new())), - active_connections: Arc::new(RwLock::new(HashMap::new())), - service_handles: HashMap::new(), - client_handles: HashMap::new(), - config: AppConfig::default(), - config_dir: config_dir.clone(), - app_directory_path: app_directory_path.clone(), - local_server_status_cache: local_server_status_cache.clone(), - local_server_status_last_update: Arc::new(RwLock::new(None)), - local_server_status_refreshing: Arc::new(AtomicBool::new(false)), - service_status_cache: Arc::new(RwLock::new(HashMap::new())), - client_status_cache: Arc::new(RwLock::new(HashMap::new())), - service_status_refreshing: Arc::new(RwLock::new(HashSet::new())), - client_status_refreshing: Arc::new(RwLock::new(HashSet::new())), - registering: Arc::new(StdMutex::new(HashSet::new())), - connecting: Arc::new(StdMutex::new(HashSet::new())), - }; - - let config = temp_state.load_config().unwrap_or_else(|e| { - tracing::warn!("Could not load config: {}, using defaults", e); - AppConfig::default() - }); - - tracing::info!( - "Loaded configuration: server_address={}, keep_alive={}, msg_header_key_set={}", - config.server_address, - config.keep_alive_enabled, - !config.msg_header_key.is_empty() - ); - - let state = Self { - server_handle: None, - server_shutdown_token: None, - server_status_sender: None, - server_start_time: None, - registered_services: Arc::new(RwLock::new(HashMap::new())), - active_connections: Arc::new(RwLock::new(HashMap::new())), - service_handles: HashMap::new(), - client_handles: HashMap::new(), - config, - config_dir, - app_directory_path, - local_server_status_cache, - local_server_status_last_update: Arc::new(RwLock::new(None)), - local_server_status_refreshing: Arc::new(AtomicBool::new(false)), - service_status_cache: Arc::new(RwLock::new(HashMap::new())), - client_status_cache: Arc::new(RwLock::new(HashMap::new())), - service_status_refreshing: Arc::new(RwLock::new(HashSet::new())), - client_status_refreshing: Arc::new(RwLock::new(HashSet::new())), - registering: Arc::new(StdMutex::new(HashSet::new())), - connecting: Arc::new(StdMutex::new(HashSet::new())), - }; - if let Err(e) = state.apply_msg_header_key_env() { - tracing::error!("Failed to apply MSG_HEADER_KEY during init: {}", e); - } - state - } - - pub fn set_app_directory_path(&mut self, path: Option) -> Result<(), CtlError> { - self.app_directory_path = path; - self.config_dir = Self::get_config_dir(&self.app_directory_path); - - // Reload config from new location if exists - match self.load_config() { - Ok(config) => self.config = config, - Err(e) => { - tracing::warn!("Failed to reload config after setting app dir: {}", e); - } - } - self.apply_msg_header_key_env()?; - - Ok(()) - } - - fn apply_msg_header_key_env(&self) -> Result<(), CtlError> { - let key = (!self.config.msg_header_key.is_empty()).then_some(&*self.config.msg_header_key); - // The library validates the key's length and shape; a rejection here is - // the stored setting being wrong, not something going wrong. - set_process_msg_header_key(key).map_err(CtlError::invalid_argument) - } - - #[allow(unused_variables)] - fn get_config_dir(app_directory_path: &Option) -> PathBuf { - // An explicit path wins everywhere. Mobile is where it normally comes - // from — Flutter hands it over, because there is no OS config dir to - // discover — but honouring it on desktop too is what lets a test point - // a state at a temporary directory instead of the user's real config. - if let Some(app_dir) = app_directory_path { - let path = PathBuf::from(app_dir).join("pb-mapper-ui"); - tracing::info!("Using caller-provided app directory: {:?}", path); - return path; - } - #[cfg(any(target_os = "android", target_os = "ios"))] - { - tracing::warn!("No app directory provided for mobile platform, using relative path"); - PathBuf::from("pb-mapper-ui") - } - #[cfg(not(any(target_os = "android", target_os = "ios")))] - { - if let Some(config_dir) = dirs::config_dir() { - config_dir.join("pb-mapper-ui") - } else if let Some(home_dir) = dirs::home_dir() { - home_dir.join(".config").join("pb-mapper-ui") - } else { - tracing::warn!("Could not determine home directory, using current directory"); - std::env::current_dir() - .unwrap_or_else(|_| PathBuf::from(".")) - .join("pb-mapper-ui-config") - } - } - } - - fn get_config_file_path(&self) -> PathBuf { - let config_dir = Self::get_config_dir(&self.app_directory_path); - - if let Err(e) = std::fs::create_dir_all(&config_dir) { - tracing::warn!( - "Failed to create config directory {:?}: {}, using current directory", - config_dir, - e - ); - return PathBuf::from("pb_mapper_config.json"); - } - - let config_file = config_dir.join("config.json"); - tracing::info!("Using config file path: {:?}", config_file); - config_file - } - - pub fn load_config(&self) -> Result { - let config_path = self.get_config_file_path(); - if config_path.exists() { - let contents = - fs::read_to_string(config_path).map_err(|e| CtlError::io(e.to_string()))?; - let mut config: AppConfig = - serde_json::from_str(&contents).map_err(|e| CtlError::io(e.to_string()))?; - config.msg_header_key = normalize_msg_header_key(config.msg_header_key)?; - Ok(config) - } else { - Ok(AppConfig::default()) - } - } - - pub fn save_config(&self) -> Result<(), CtlError> { - let config_path = self.get_config_file_path(); - let contents = - serde_json::to_string_pretty(&self.config).map_err(|e| CtlError::io(e.to_string()))?; - fs::write(config_path, contents).map_err(|e| CtlError::io(e.to_string()))?; - Ok(()) - } - - fn get_service_config_path(&self) -> PathBuf { - self.config_dir.join("services.json") - } - - fn get_client_config_path(&self) -> PathBuf { - self.config_dir.join("clients.json") - } - - pub fn load_service_configs(&self) -> ServiceConfigStore { - let path = self.get_service_config_path(); - match fs::read_to_string(&path) { - Ok(content) => serde_json::from_str(&content).unwrap_or_else(|_| ServiceConfigStore { - services: HashMap::new(), - }), - Err(_) => ServiceConfigStore { - services: HashMap::new(), - }, - } - } - - pub fn save_service_configs(&self, store: &ServiceConfigStore) -> Result<(), CtlError> { - let path = self.get_service_config_path(); - if let Some(parent) = path.parent() { - fs::create_dir_all(parent) - .map_err(|e| CtlError::io(format!("Failed to create config dir: {e}")))?; - } - - let content = serde_json::to_string_pretty(store) - .map_err(|e| CtlError::io(format!("Failed to serialize config: {e}")))?; - - fs::write(&path, content) - .map_err(|e| CtlError::io(format!("Failed to write config file: {e}")))?; - Ok(()) - } - - pub fn save_service_config( - &self, - service_key: &str, - local_address: &str, - protocol: &str, - enable_encryption: bool, - enable_keep_alive: bool, - ) -> Result<(), CtlError> { - let mut store = self.load_service_configs(); - let now = SystemTime::now(); - - let config = ServiceConfigData { - service_key: service_key.to_string(), - local_address: local_address.to_string(), - protocol: protocol.to_string(), - enable_encryption, - enable_keep_alive, - created_at: if store.services.contains_key(service_key) { - store.services[service_key].created_at - } else { - now - }, - }; - - store.services.insert(service_key.to_string(), config); - self.save_service_configs(&store) - } - - pub fn delete_service_config(&self, service_key: &str) -> Result<(), CtlError> { - let mut store = self.load_service_configs(); - store.services.remove(service_key); - self.save_service_configs(&store) - } - - pub fn load_client_configs(&self) -> ClientConfigStore { - let path = self.get_client_config_path(); - match fs::read_to_string(&path) { - Ok(content) => serde_json::from_str(&content).unwrap_or_else(|_| ClientConfigStore { - clients: HashMap::new(), - }), - Err(_) => ClientConfigStore { - clients: HashMap::new(), - }, - } - } - - pub fn save_client_configs(&self, store: &ClientConfigStore) -> Result<(), CtlError> { - let path = self.get_client_config_path(); - if let Some(parent) = path.parent() { - fs::create_dir_all(parent) - .map_err(|e| CtlError::io(format!("Failed to create config dir: {e}")))?; - } - - let content = serde_json::to_string_pretty(store) - .map_err(|e| CtlError::io(format!("Failed to serialize client config: {e}")))?; - - fs::write(&path, content) - .map_err(|e| CtlError::io(format!("Failed to write client config file: {e}")))?; - Ok(()) - } - - pub fn save_client_config( - &self, - service_key: &str, - local_address: &str, - protocol: &str, - enable_keep_alive: bool, - ) -> Result<(), CtlError> { - let mut store = self.load_client_configs(); - let now = SystemTime::now(); - - let config = ClientConfigData { - service_key: service_key.to_string(), - local_address: local_address.to_string(), - protocol: protocol.to_string(), - enable_keep_alive, - created_at: if store.clients.contains_key(service_key) { - store.clients[service_key].created_at - } else { - now - }, - }; - - store.clients.insert(service_key.to_string(), config); - self.save_client_configs(&store) - } - - pub fn delete_client_config(&self, service_key: &str) -> Result<(), CtlError> { - let mut store = self.load_client_configs(); - store.clients.remove(service_key); - self.save_client_configs(&store) - } - - pub async fn start_server( - &mut self, - port: u16, - enable_keep_alive: bool, - ) -> Result<(), CtlError> { - if self.server_handle.is_some() { - return Err(CtlError::already_exists("Server is already running")); - } - - let ip_addr = IpAddr::V4(Ipv4Addr::new(0, 0, 0, 0)); - let bind_addr = std::net::SocketAddr::new(ip_addr, port); - - // Preflight bind to surface "port already in use" errors before spawning. - let listener = TcpListener::bind(bind_addr).await.map_err(|e| { - CtlError::address_in_use(format!("Failed to bind server on {bind_addr}: {e}")) - })?; - drop(listener); - - tracing::info!("Starting pb-mapper server on {}:{}", ip_addr, port); - - let shutdown_token = CancellationToken::new(); - let shutdown_token_clone = shutdown_token.clone(); - - let (status_sender, status_receiver) = tokio::sync::mpsc::unbounded_channel(); - - let handle = tokio::spawn(async move { - if let Err(e) = run_server_with_shutdown( - (ip_addr, port), - shutdown_token_clone, - Some(status_receiver), - enable_keep_alive, - ) - .await - { - tracing::error!("pb-mapper server stopped with error: {e}"); - } - }); - - self.server_handle = Some(handle); - self.server_shutdown_token = Some(shutdown_token); - self.server_status_sender = Some(status_sender); - self.server_start_time = Some(SystemTime::now()); - - { - let mut cache = self.local_server_status_cache.write().await; - *cache = LocalServerStatus { - is_running: true, - active_connections: 0, - registered_services: 0, - uptime_seconds: 0, - }; - } - { - let mut last_update = self.local_server_status_last_update.write().await; - *last_update = Some(Instant::now()); - } - - tracing::info!("pb-mapper server started successfully"); - Ok(()) - } - - pub async fn stop_server(&mut self) -> Result<(), CtlError> { - if let (Some(handle), Some(shutdown_token)) = - (self.server_handle.take(), self.server_shutdown_token.take()) - { - self.server_status_sender = None; - - shutdown_token.cancel(); - - let shutdown_timeout = tokio::time::Duration::from_secs(5); - - match tokio::time::timeout(shutdown_timeout, handle).await { - Ok(_) => { - tracing::info!("Server shutdown gracefully"); - } - Err(_) => { - tracing::warn!("Server shutdown timed out, may not have closed gracefully"); - } - } - - self.server_start_time = None; - - { - let mut cache = self.local_server_status_cache.write().await; - *cache = LocalServerStatus { - is_running: false, - active_connections: 0, - registered_services: 0, - uptime_seconds: 0, - }; - } - { - let mut last_update = self.local_server_status_last_update.write().await; - *last_update = Some(Instant::now()); - } - - for (_, handle) in self.service_handles.drain() { - handle.abort(); - } - - for (_, handle) in self.client_handles.drain() { - handle.abort(); - } - - self.registered_services.write().await.clear(); - self.active_connections.write().await.clear(); - - tracing::info!("pb-mapper server stopped, all services and connections terminated"); - Ok(()) - } else { - Err(CtlError::not_found("Server is not running")) - } - } - - async fn finish_register(&mut self, commit: RegisterCommit) -> Result<(), CtlError> { - let RegisterCommit { - service_key, - local_address, - protocol, - enable_encryption, - enable_keep_alive, - local_sock_addr, - remote_sock_addr, - } = commit; - - if let Some(previous) = self.service_handles.remove(&service_key) { - tracing::warn!( - "Service '{service_key}' is already registered, replacing existing handle" - ); - // Dropping a `JoinHandle` does not stop the task. Without this the - // replaced tunnel kept running and retrying, with nothing left - // holding a handle able to abort it. - previous.abort(); - } - - tracing::info!( - "Registering service '{}' with protocol {}, local address {}, server address {}", - service_key, - protocol, - local_address, - self.config.server_address - ); - - self.save_service_config( - &service_key, - &local_address, - &protocol, - enable_encryption, - enable_keep_alive, - ) - .map_err(|e| CtlError::io(format!("Failed to save service configuration: {e}")))?; - - let key_clone = service_key.clone(); - let service_key_for_status = service_key.clone(); - - let callback: StatusCallback = Box::new(move |status: &str| { - tracing::info!( - "Service {} status update: {}", - service_key_for_status, - status - ); - }); - - let handle = if protocol.to_uppercase() == "TCP" { - tokio::spawn(async move { - let _ = run_server_side_cli_with_callback::( - local_sock_addr, - remote_sock_addr, - key_clone.into(), - ServerTunnelOptions { - need_codec: enable_encryption, - is_datagram: false, - keep_alive: enable_keep_alive, - }, - Some(callback), - ) - .await; - }) - } else { - tokio::spawn(async move { - let _ = run_server_side_cli_with_callback::( - local_sock_addr, - remote_sock_addr, - key_clone.into(), - ServerTunnelOptions { - need_codec: enable_encryption, - is_datagram: true, - keep_alive: enable_keep_alive, - }, - Some(callback), - ) - .await; - }) - }; - - self.service_handles.insert(service_key.clone(), handle); - - { - let mut cache = self.service_status_cache.write().await; - cache.insert( - service_key.clone(), - StatusCacheEntry { - status: "retrying".to_string(), - message: "Connecting to pb-mapper server...".to_string(), - updated_at: Instant::now(), - }, - ); - } - self.schedule_service_status_refresh(&service_key).await; - - let service_info = ServiceInfo { - service_key: service_key.clone(), - protocol, - local_address, - status: "Registering".to_string(), - }; - - self.registered_services - .write() - .await - .insert(service_key.clone(), service_info); - - tracing::info!("Service '{}' registration initiated", service_key); - Ok(()) - } - - pub async fn unregister_service(&mut self, service_key: String) -> Result<(), CtlError> { - if let Some(handle) = self.service_handles.remove(&service_key) { - handle.abort(); - } - - if self - .registered_services - .write() - .await - .remove(&service_key) - .is_some() - { - tracing::info!("Service '{}' unregistered successfully", service_key); - Ok(()) - } else { - Err(CtlError::not_found(format!( - "Service '{service_key}' is not registered" - ))) - } - } - - pub async fn delete_service_config_and_stop( - &mut self, - service_key: String, - ) -> Result<(), CtlError> { - if let Some(handle) = self.service_handles.remove(&service_key) { - handle.abort(); - } - - self.registered_services.write().await.remove(&service_key); - - self.delete_service_config(&service_key) - } - - async fn finish_connect(&mut self, commit: ConnectCommit) -> Result<(), CtlError> { - let ConnectCommit { - service_key, - local_address, - protocol, - enable_keep_alive, - local_sock_addr, - remote_sock_addr, - } = commit; - - if let Some(previous) = self.client_handles.remove(&service_key) { - tracing::warn!( - "Client for service '{service_key}' is already connected, replacing handle" - ); - // As in `finish_register`: dropping the handle leaves the old - // client's retry loop running with nothing able to stop it. - previous.abort(); - } - - let protocol_upper = protocol.to_uppercase(); - - tracing::info!( - "Connecting to service '{}' with protocol {}, local address {}, server address {}", - service_key, - protocol, - local_address, - self.config.server_address - ); - - let key_clone = service_key.clone(); - - let status_callback: ClientStatusCallback = { - let service_key_for_callback = service_key.clone(); - Box::new(move |status: &str| { - tracing::info!("Client {} status: {}", service_key_for_callback, status); - }) - }; - - let handle = if protocol_upper == "TCP" { - tokio::spawn(async move { - run_client_side_cli_with_callback::( - local_sock_addr, - remote_sock_addr, - key_clone.into(), - enable_keep_alive, - Some(status_callback), - ) - .await; - }) - } else { - tokio::spawn(async move { - run_client_side_cli_with_callback::( - local_sock_addr, - remote_sock_addr, - key_clone.into(), - enable_keep_alive, - Some(status_callback), - ) - .await; - }) - }; - - self.client_handles.insert(service_key.clone(), handle); - - { - let mut cache = self.client_status_cache.write().await; - cache.insert( - service_key.clone(), - StatusCacheEntry { - status: "retrying".to_string(), - message: "Connecting to pb-mapper server...".to_string(), - updated_at: Instant::now(), - }, - ); - } - self.schedule_client_status_refresh(&service_key).await; - - let connection_info = ConnectionInfo { - service_key: service_key.clone(), - client_id: format!("client-{service_key}"), - status: "Connected".to_string(), - }; - - self.active_connections - .write() - .await - .insert(service_key.clone(), connection_info); - - // Persist here rather than at the FFI boundary, so a connection made - // from a terminal is remembered exactly like one made from the window. - // `finish_register` has always done this; leaving it out here meant a - // CLI `connect` started a client that never appeared in the list. - if let Err(e) = - self.save_client_config(&service_key, &local_address, &protocol, enable_keep_alive) - { - // The client is up either way, so this is a warning and not a - // failure: losing the config costs the entry after a restart. - tracing::warn!("Failed to save client config for '{service_key}': {e}"); - } - - tracing::info!("Connected to service '{}' successfully", service_key); - Ok(()) - } - - /// Claims a service key for a registration. See [`KeyClaim`]. - fn claim_registering(&self, service_key: &str) -> Result { - claim_key(&self.registering, service_key, "being registered") - } - - /// Claims a service key for a client connection. See [`KeyClaim`]. - fn claim_connecting(&self, service_key: &str) -> Result { - claim_key(&self.connecting, service_key, "being connected") - } - - pub async fn disconnect_service(&mut self, service_key: String) -> Result<(), CtlError> { - // Aborting the task is the part that matters: it is what stops the - // retry loop still dialling in the background. - let aborted = match self.client_handles.remove(&service_key) { - Some(handle) => { - handle.abort(); - true - } - None => false, - }; - - let was_listed = self - .active_connections - .write() - .await - .remove(&service_key) - .is_some(); - - // Reported failure only when there was nothing to stop. It used to key - // off the bookkeeping map alone, so a client whose task had been - // aborted could still be reported as "not connected" — an error for an - // operation that had in fact just done its job. - if aborted || was_listed { - tracing::info!("Disconnected from service '{}'", service_key); - Ok(()) - } else { - Err(CtlError::not_found(format!( - "Service '{service_key}' is not connected" - ))) - } - } - - pub async fn delete_client_config_and_stop( - &mut self, - service_key: String, - ) -> Result<(), CtlError> { - if let Some(handle) = self.client_handles.remove(&service_key) { - handle.abort(); - } - - self.active_connections.write().await.remove(&service_key); - - self.delete_client_config(&service_key) - } - - pub async fn get_config_status(&self) -> AppConfig { - self.config.clone() - } - - pub async fn update_config( - &mut self, - server_address: String, - keep_alive: bool, - msg_header_key: String, - ) -> Result<(), CtlError> { - let msg_header_key = normalize_msg_header_key(msg_header_key)?; - self.config.server_address = server_address; - self.config.keep_alive_enabled = keep_alive; - self.config.msg_header_key = msg_header_key; - self.apply_msg_header_key_env()?; - self.save_config()?; - self.reset_status_caches().await; - Ok(()) - } - - pub async fn get_service_configs(&self) -> Vec { - let store = self.load_service_configs(); - let mut services = Vec::new(); - - let mut sorted_configs: Vec<_> = store.services.values().collect(); - sorted_configs.sort_by_key(|config| config.created_at); - - for config in sorted_configs { - let (status, message) = self.calculate_service_status(&config.service_key).await; - - services.push(ServiceConfigInfo { - service_key: config.service_key.clone(), - local_address: config.local_address.clone(), - protocol: config.protocol.clone(), - enable_encryption: config.enable_encryption, - enable_keep_alive: config.enable_keep_alive, - status, - status_message: message, - created_at_ms: config - .created_at - .duration_since(SystemTime::UNIX_EPOCH) - .unwrap_or_default() - .as_millis() as u64, - updated_at_ms: SystemTime::now() - .duration_since(SystemTime::UNIX_EPOCH) - .unwrap_or_default() - .as_millis() as u64, - }); - } - - services - } - - pub async fn get_service_status(&self, service_key: String) -> ServiceStatusResponse { - let (status, message) = self.calculate_service_status(&service_key).await; - ServiceStatusResponse { - service_key, - status, - message, - } - } - - pub async fn get_client_configs(&self) -> Vec { - let store = self.load_client_configs(); - let mut client_infos = Vec::new(); - - for (service_key, config) in store.clients.iter() { - let (status, status_message) = self.calculate_client_status(service_key).await; - - client_infos.push(ClientConfigInfo { - service_key: config.service_key.clone(), - local_address: config.local_address.clone(), - protocol: config.protocol.clone(), - enable_keep_alive: config.enable_keep_alive, - status, - status_message, - created_at_ms: config - .created_at - .duration_since(UNIX_EPOCH) - .unwrap_or_default() - .as_millis() as u64, - updated_at_ms: config - .created_at - .duration_since(UNIX_EPOCH) - .unwrap_or_default() - .as_millis() as u64, - }); - } - - client_infos.sort_by_key(|info| info.created_at_ms); - client_infos - } - - pub async fn get_client_status(&self, service_key: String) -> ClientStatusResponse { - let (status, message) = self.calculate_client_status(&service_key).await; - ClientStatusResponse { - service_key, - status, - message, - } - } - - pub async fn get_local_server_status(&self) -> LocalServerStatus { - let is_running = self.server_handle.is_some(); - if !is_running { - let status = LocalServerStatus { - is_running: false, - active_connections: 0, - registered_services: 0, - uptime_seconds: 0, - }; - { - let mut cache = self.local_server_status_cache.write().await; - *cache = status.clone(); - } - { - let mut last_update = self.local_server_status_last_update.write().await; - *last_update = Some(Instant::now()); - } - return status; - } - - let should_refresh = { - let last_update = self.local_server_status_last_update.read().await; - cache_is_stale(*last_update, STATUS_CACHE_TTL) - }; - - if should_refresh { - self.schedule_local_server_status_refresh(); - } - - let cache = self.local_server_status_cache.read().await; - cache.clone() - } - - fn schedule_local_server_status_refresh(&self) { - if self - .local_server_status_refreshing - .swap(true, Ordering::AcqRel) - { - return; - } - - let sender = self.server_status_sender.clone(); - let cache = self.local_server_status_cache.clone(); - let last_update = self.local_server_status_last_update.clone(); - let refreshing = self.local_server_status_refreshing.clone(); - let start_time = self.server_start_time; - - tokio::spawn(async move { - let mut status = LocalServerStatus { - is_running: true, - active_connections: 0, - registered_services: 0, - uptime_seconds: start_time - .and_then(|ts| SystemTime::now().duration_since(ts).ok()) - .map(|d| d.as_secs()) - .unwrap_or(0), - }; - - if let Some(sender) = sender { - let (response_sender, response_receiver) = tokio::sync::oneshot::channel(); - if sender.send(response_sender).is_ok() { - if let Ok(Ok(info)) = - tokio::time::timeout(Duration::from_millis(200), response_receiver).await - { - status.active_connections = info.active_connections; - status.registered_services = info.registered_services; - status.uptime_seconds = info.uptime_seconds; - } - } - } - - { - let mut cache = cache.write().await; - *cache = status; - } - { - let mut last_update = last_update.write().await; - *last_update = Some(Instant::now()); - } - refreshing.store(false, Ordering::Release); - }); - } - - pub async fn get_server_status_detail(&self) -> Result { - self.force_refresh_server_status().await - } - - /// The connections the server holds for one key, from the protocol's own - /// structured query rather than the Debug dump in `server_map`. - pub async fn get_service_conns( - &self, - service_key: String, - ) -> Result, CtlError> { - let server_addr = self.config.server_address.clone(); - match tokio::time::timeout( - FORCE_REFRESH_TIMEOUT, - get_service_conns_with_addr(&server_addr, &service_key), - ) - .await - { - Ok(result) => result, - Err(_) => Err(CtlError::timeout(format!( - "Timed out asking {server_addr} about {service_key}" - ))), - } - } - - /// Perform a blocking status refresh — waits for the actual network result - /// instead of returning stale cache. - pub async fn force_refresh_server_status(&self) -> Result { - let server_addr = self.config.server_address.clone(); - - let detail = match tokio::time::timeout( - FORCE_REFRESH_TIMEOUT, - fetch_real_status_with_addr(&server_addr), - ) - .await - { - Ok(Ok((services, remote_id_data))) => ServerStatusDetail { - server_available: true, - registered_services: services, - server_map: remote_id_data.server_map, - active_connections: remote_id_data.active, - idle_connections: remote_id_data.idle, - }, - Ok(Err(e)) => { - tracing::warn!("Force refresh failed: {}", e); - ServerStatusDetail { - server_available: false, - registered_services: Vec::new(), - server_map: String::new(), - active_connections: String::new(), - idle_connections: String::new(), - } - } - Err(_) => { - tracing::warn!("Force refresh timed out after {:?}", FORCE_REFRESH_TIMEOUT); - ServerStatusDetail { - server_available: false, - registered_services: Vec::new(), - server_map: String::new(), - active_connections: String::new(), - idle_connections: String::new(), - } - } - }; - - Ok(detail) - } - - // Cache service status to avoid blocking UI with network checks on every paint. - async fn get_cached_service_status(&self, service_key: &str) -> (String, String) { - if let Some(handle) = self.service_handles.get(service_key) { - if handle.is_finished() { - return ( - "failed".to_string(), - "Service connection terminated".to_string(), - ); - } - - let cached = { - let cache = self.service_status_cache.read().await; - cache.get(service_key).cloned() - }; - - let should_refresh = cached - .as_ref() - .map(|entry| entry.updated_at.elapsed() > STATUS_CACHE_TTL) - .unwrap_or(true); - - if should_refresh { - self.schedule_service_status_refresh(service_key).await; - } - - if let Some(entry) = cached { - return (entry.status, entry.message); - } - - return ( - "retrying".to_string(), - "Checking service status...".to_string(), - ); - } - - ( - "stopped".to_string(), - "Service is not registered".to_string(), - ) - } - - // Cache client status to avoid blocking UI with network checks on every paint. - async fn get_cached_client_status(&self, service_key: &str) -> (String, String) { - if let Some(handle) = self.client_handles.get(service_key) { - if handle.is_finished() { - return ( - "failed".to_string(), - "Client connection terminated".to_string(), - ); - } - - let cached = { - let cache = self.client_status_cache.read().await; - cache.get(service_key).cloned() - }; - - let should_refresh = cached - .as_ref() - .map(|entry| entry.updated_at.elapsed() > STATUS_CACHE_TTL) - .unwrap_or(true); - - if should_refresh { - self.schedule_client_status_refresh(service_key).await; - } - - if let Some(entry) = cached { - return (entry.status, entry.message); - } - - return ( - "retrying".to_string(), - "Checking client status...".to_string(), - ); - } - - ("stopped".to_string(), "Client is not connected".to_string()) - } - - async fn schedule_service_status_refresh(&self, service_key: &str) { - { - let mut refreshing = self.service_status_refreshing.write().await; - if refreshing.contains(service_key) { - return; - } - refreshing.insert(service_key.to_string()); - } - - let server_addr = self.config.server_address.clone(); - let cache = self.service_status_cache.clone(); - let refreshing = self.service_status_refreshing.clone(); - let key = service_key.to_string(); - - tokio::spawn(async move { - let result = tokio::time::timeout( - STATUS_REFRESH_TIMEOUT, - check_service_with_get_status(&server_addr, &key), - ) - .await; - - let (status, message) = match result { - Ok(Ok(true)) => ( - "running".to_string(), - "Service is running normally".to_string(), - ), - Ok(Ok(false)) => ( - "retrying".to_string(), - "Service is in retry connection loop".to_string(), - ), - Ok(Err(_)) | Err(_) => ( - "failed".to_string(), - "Cannot connect to pb-server".to_string(), - ), - }; - - let changed = { - let mut cache = cache.write().await; - let changed = cache - .get(&key) - .is_none_or(|entry| entry.status != status || entry.message != message); - cache.insert( - key.clone(), - StatusCacheEntry { - status, - message, - updated_at: Instant::now(), - }, - ); - changed - }; - // Only transitions the user can perceive. These run on a timer for - // every configured entry, so emitting on every refresh would reload - // the list several times a second for no visible reason. - if changed { - events::emit(events::ChangeKind::Services, Some(&key), Origin::Internal); - } - - let mut refreshing = refreshing.write().await; - refreshing.remove(&key); - }); - } - - async fn schedule_client_status_refresh(&self, service_key: &str) { - { - let mut refreshing = self.client_status_refreshing.write().await; - if refreshing.contains(service_key) { - return; - } - refreshing.insert(service_key.to_string()); - } - - let server_addr = self.config.server_address.clone(); - let cache = self.client_status_cache.clone(); - let refreshing = self.client_status_refreshing.clone(); - let key = service_key.to_string(); - - tokio::spawn(async move { - let result = tokio::time::timeout( - STATUS_REFRESH_TIMEOUT, - check_service_with_get_status(&server_addr, &key), - ) - .await; - - let (status, message) = match result { - Ok(Ok(true)) => ( - "running".to_string(), - "Client is connected normally".to_string(), - ), - Ok(Ok(false)) => ( - "retrying".to_string(), - "Client is in retry connection loop".to_string(), - ), - Ok(Err(_)) | Err(_) => ( - "failed".to_string(), - "Cannot connect to pb-server".to_string(), - ), - }; - - let changed = { - let mut cache = cache.write().await; - let changed = cache - .get(&key) - .is_none_or(|entry| entry.status != status || entry.message != message); - cache.insert( - key.clone(), - StatusCacheEntry { - status, - message, - updated_at: Instant::now(), - }, - ); - changed - }; - // Only transitions the user can perceive. These run on a timer for - // every configured entry, so emitting on every refresh would reload - // the list several times a second for no visible reason. - if changed { - events::emit(events::ChangeKind::Clients, Some(&key), Origin::Internal); - } - - let mut refreshing = refreshing.write().await; - refreshing.remove(&key); - }); - } - - async fn calculate_service_status(&self, service_key: &str) -> (String, String) { - self.get_cached_service_status(service_key).await - } - - async fn calculate_client_status(&self, service_key: &str) -> (String, String) { - self.get_cached_client_status(service_key).await - } -} +mod configuration; +mod runtime; +mod status; /// Registers a service, holding the state lock only for the bookkeeping. /// @@ -1671,10 +518,13 @@ pub async fn register_service( // 1. Claim the key and take what the slow work needs. Microseconds. // `_claim` is held to the end of the function on purpose: dropping it // early would release the key while the setup is still running. - let (_claim, server_address) = { + // The credential is captured with the relay address so a later config + // change cannot pair a new key with the already-resolved socket. + let (_claim, server_address, credential) = { let state = state.lock().await; let claim = state.claim_registering(&service_key)?; - (claim, state.config.server_address.clone()) + let credential = get_process_credential().map_err(CtlError::invalid_argument)?; + (claim, state.config.server_address.clone(), credential) }; // 2. The slow parts, with the lock released. @@ -1701,6 +551,7 @@ pub async fn register_service( enable_keep_alive, local_sock_addr, remote_sock_addr, + credential, }) .await } @@ -1714,10 +565,11 @@ pub async fn connect_service( protocol: String, enable_keep_alive: bool, ) -> Result<(), CtlError> { - let (_claim, server_address) = { + let (_claim, server_address, credential) = { let state = state.lock().await; let claim = state.claim_connecting(&service_key)?; - (claim, state.config.server_address.clone()) + let credential = get_process_credential().map_err(CtlError::invalid_argument)?; + (claim, state.config.server_address.clone(), credential) }; let local_sock_addr = get_sockaddr_async(&local_address) @@ -1758,6 +610,7 @@ pub async fn connect_service( enable_keep_alive, local_sock_addr, remote_sock_addr, + credential, }) .await } @@ -1784,6 +637,37 @@ mod tests { (Arc::new(Mutex::new(state)), root) } + #[tokio::test] + async fn ui_server_uses_its_writable_config_directory_and_reports_readiness() { + let (state, root) = temp_state("server-auth-path"); + let auth_dir = { + let mut state = state.lock().await; + let auth_dir = state.config_dir.join("auth"); + state + .start_server(0, false) + .await + .expect("UI server should bind and initialize authentication"); + assert!(state.server_handle.is_some()); + assert!(state.get_local_server_status().await.is_running); + auth_dir + }; + + assert!(auth_dir.join("admin.key").is_file()); + let isolated = state + .lock() + .await + .isolated_admin_key() + .expect("embedded relay should expose its administrator key"); + assert_eq!(isolated.len(), 32); + state + .lock() + .await + .stop_server() + .await + .expect("UI server should stop cleanly"); + let _ = std::fs::remove_dir_all(&root); + } + /// The claim is what stands in for the lock that registration no longer /// holds across its slow phase. Without it, two callers — the window and a /// terminal, say — could both finish the preflight and both insert, and the @@ -1834,6 +718,15 @@ mod tests { #[tokio::test] async fn a_failed_registration_releases_its_claim() { let (state, root) = temp_state("release"); + struct RestoreProcessKey; + impl Drop for RestoreProcessKey { + fn drop(&mut self) { + let _ = set_process_msg_header_key(None); + } + } + set_process_msg_header_key(Some("0123456789abcdefghijklmnopqrstuv")) + .expect("test credential"); + let _restore_process_key = RestoreProcessKey; // Fails in phase 2, while the claim is held. let first = register_service( diff --git a/ui/native/pb_mapper_ffi/src/state/configuration.rs b/ui/native/pb_mapper_ffi/src/state/configuration.rs new file mode 100644 index 0000000..feb4bd5 --- /dev/null +++ b/ui/native/pb_mapper_ffi/src/state/configuration.rs @@ -0,0 +1,307 @@ +//! User-writable configuration storage and in-memory state initialization. +//! +//! ```text +//! app directory -> pb-mapper-ui/config.json -> AppConfig -> process credential +//! -> services.json / clients.json -> remembered tunnel definitions +//! ``` +//! +//! An explicit Flutter app directory wins on every platform. This same directory is +//! the root for relay authentication state, so desktop and mobile UI processes never +//! depend on root-owned `/var/lib` paths. + +use super::*; + +impl PbMapperState { + pub(super) async fn reset_status_caches(&self) { + { + let mut cache = self.local_server_status_cache.write().await; + *cache = LocalServerStatus { + is_running: false, + active_connections: 0, + registered_services: 0, + uptime_seconds: 0, + }; + } + { + let mut last_update = self.local_server_status_last_update.write().await; + *last_update = None; + } + self.local_server_status_refreshing + .store(false, Ordering::Release); + + self.service_status_cache.write().await.clear(); + self.client_status_cache.write().await.clear(); + self.service_status_refreshing.write().await.clear(); + self.client_status_refreshing.write().await.clear(); + } + pub fn new(app_directory_path: Option) -> Self { + let config_dir = Self::get_config_dir(&app_directory_path); + tracing::info!("Using config directory: {:?}", config_dir); + + let local_server_status_cache = Arc::new(RwLock::new(LocalServerStatus { + is_running: false, + active_connections: 0, + registered_services: 0, + uptime_seconds: 0, + })); + + let mut state = Self { + server_handle: None, + server_auth: None, + server_shutdown_token: None, + server_status_sender: None, + server_start_time: None, + registered_services: Arc::new(RwLock::new(HashMap::new())), + active_connections: Arc::new(RwLock::new(HashMap::new())), + service_runtime: HashMap::new(), + client_runtime: HashMap::new(), + config: AppConfig::default(), + config_dir, + app_directory_path, + local_server_status_cache, + local_server_status_last_update: Arc::new(RwLock::new(None)), + local_server_status_refreshing: Arc::new(AtomicBool::new(false)), + service_status_cache: Arc::new(RwLock::new(HashMap::new())), + client_status_cache: Arc::new(RwLock::new(HashMap::new())), + service_status_refreshing: Arc::new(RwLock::new(HashSet::new())), + client_status_refreshing: Arc::new(RwLock::new(HashSet::new())), + registering: Arc::new(SyncMutex::new(HashSet::new())), + connecting: Arc::new(SyncMutex::new(HashSet::new())), + }; + state.config = state.load_config().unwrap_or_else(|e| { + tracing::warn!("Could not load config: {}, using defaults", e); + AppConfig::default() + }); + tracing::info!( + "Loaded configuration: server_address={}, keep_alive={}, msg_header_key_set={}", + state.config.server_address, + state.config.keep_alive_enabled, + !state.config.msg_header_key.is_empty() + ); + if let Err(e) = state.apply_msg_header_key_env() { + tracing::error!("Failed to apply MSG_HEADER_KEY during init: {}", e); + } + state + } + + pub fn set_app_directory_path(&mut self, path: Option) -> Result<(), CtlError> { + self.app_directory_path = path; + self.config_dir = Self::get_config_dir(&self.app_directory_path); + + // Reload config from new location if exists + match self.load_config() { + Ok(config) => self.config = config, + Err(e) => { + tracing::warn!("Failed to reload config after setting app dir: {}", e); + } + } + self.apply_msg_header_key_env()?; + + Ok(()) + } + + pub(super) fn apply_msg_header_key_env(&self) -> Result<(), CtlError> { + let key = (!self.config.msg_header_key.is_empty()).then_some(&*self.config.msg_header_key); + // The library validates the key's length and shape; a rejection here is + // the stored setting being wrong, not something going wrong. + set_process_msg_header_key(key).map_err(CtlError::invalid_argument) + } + + #[allow(unused_variables)] + fn get_config_dir(app_directory_path: &Option) -> PathBuf { + // An explicit path wins everywhere. Mobile is where it normally comes + // from — Flutter hands it over, because there is no OS config dir to + // discover — but honouring it on desktop too is what lets a test point + // a state at a temporary directory instead of the user's real config. + if let Some(app_dir) = app_directory_path { + let path = PathBuf::from(app_dir).join("pb-mapper-ui"); + tracing::info!("Using caller-provided app directory: {:?}", path); + return path; + } + #[cfg(any(target_os = "android", target_os = "ios"))] + { + tracing::warn!("No app directory provided for mobile platform, using relative path"); + PathBuf::from("pb-mapper-ui") + } + #[cfg(not(any(target_os = "android", target_os = "ios")))] + { + if let Some(config_dir) = dirs::config_dir() { + config_dir.join("pb-mapper-ui") + } else if let Some(home_dir) = dirs::home_dir() { + home_dir.join(".config").join("pb-mapper-ui") + } else { + tracing::warn!("Could not determine home directory, using current directory"); + std::env::current_dir() + .unwrap_or_else(|_| PathBuf::from(".")) + .join("pb-mapper-ui-config") + } + } + } + + fn get_config_file_path(&self) -> PathBuf { + let config_dir = Self::get_config_dir(&self.app_directory_path); + + if let Err(e) = std::fs::create_dir_all(&config_dir) { + tracing::warn!( + "Failed to create config directory {:?}: {}, using current directory", + config_dir, + e + ); + return PathBuf::from("pb_mapper_config.json"); + } + + let config_file = config_dir.join("config.json"); + tracing::info!("Using config file path: {:?}", config_file); + config_file + } + + pub fn load_config(&self) -> Result { + let config_path = self.get_config_file_path(); + if config_path.exists() { + let contents = + fs::read_to_string(config_path).map_err(|e| CtlError::io(e.to_string()))?; + let mut config: AppConfig = + serde_json::from_str(&contents).map_err(|e| CtlError::io(e.to_string()))?; + config.msg_header_key = normalize_msg_header_key(config.msg_header_key)?; + Ok(config) + } else { + Ok(AppConfig::default()) + } + } + + pub(super) fn write_config_file(&self, config: &AppConfig) -> Result<(), CtlError> { + let config_path = self.get_config_file_path(); + let contents = + serde_json::to_string_pretty(config).map_err(|e| CtlError::io(e.to_string()))?; + fs::write(config_path, contents).map_err(|e| CtlError::io(e.to_string()))?; + Ok(()) + } + + fn get_service_config_path(&self) -> PathBuf { + self.config_dir.join("services.json") + } + + fn get_client_config_path(&self) -> PathBuf { + self.config_dir.join("clients.json") + } + + pub fn load_service_configs(&self) -> ServiceConfigStore { + let path = self.get_service_config_path(); + match fs::read_to_string(&path) { + Ok(content) => serde_json::from_str(&content).unwrap_or_else(|_| ServiceConfigStore { + services: HashMap::new(), + }), + Err(_) => ServiceConfigStore { + services: HashMap::new(), + }, + } + } + + pub fn save_service_configs(&self, store: &ServiceConfigStore) -> Result<(), CtlError> { + let path = self.get_service_config_path(); + if let Some(parent) = path.parent() { + fs::create_dir_all(parent) + .map_err(|e| CtlError::io(format!("Failed to create config dir: {e}")))?; + } + + let content = serde_json::to_string_pretty(store) + .map_err(|e| CtlError::io(format!("Failed to serialize config: {e}")))?; + + fs::write(&path, content) + .map_err(|e| CtlError::io(format!("Failed to write config file: {e}")))?; + Ok(()) + } + + pub fn save_service_config( + &self, + service_key: &str, + local_address: &str, + protocol: &str, + enable_encryption: bool, + enable_keep_alive: bool, + ) -> Result<(), CtlError> { + let mut store = self.load_service_configs(); + let now = SystemTime::now(); + + let config = ServiceConfigData { + service_key: service_key.to_string(), + local_address: local_address.to_string(), + protocol: protocol.to_string(), + enable_encryption, + enable_keep_alive, + created_at: if store.services.contains_key(service_key) { + store.services[service_key].created_at + } else { + now + }, + }; + + store.services.insert(service_key.to_string(), config); + self.save_service_configs(&store) + } + + pub fn delete_service_config(&self, service_key: &str) -> Result<(), CtlError> { + let mut store = self.load_service_configs(); + store.services.remove(service_key); + self.save_service_configs(&store) + } + + pub fn load_client_configs(&self) -> ClientConfigStore { + let path = self.get_client_config_path(); + match fs::read_to_string(&path) { + Ok(content) => serde_json::from_str(&content).unwrap_or_else(|_| ClientConfigStore { + clients: HashMap::new(), + }), + Err(_) => ClientConfigStore { + clients: HashMap::new(), + }, + } + } + + pub fn save_client_configs(&self, store: &ClientConfigStore) -> Result<(), CtlError> { + let path = self.get_client_config_path(); + if let Some(parent) = path.parent() { + fs::create_dir_all(parent) + .map_err(|e| CtlError::io(format!("Failed to create config dir: {e}")))?; + } + + let content = serde_json::to_string_pretty(store) + .map_err(|e| CtlError::io(format!("Failed to serialize client config: {e}")))?; + + fs::write(&path, content) + .map_err(|e| CtlError::io(format!("Failed to write client config file: {e}")))?; + Ok(()) + } + + pub fn save_client_config( + &self, + service_key: &str, + local_address: &str, + protocol: &str, + enable_keep_alive: bool, + ) -> Result<(), CtlError> { + let mut store = self.load_client_configs(); + let now = SystemTime::now(); + + let config = ClientConfigData { + service_key: service_key.to_string(), + local_address: local_address.to_string(), + protocol: protocol.to_string(), + enable_keep_alive, + created_at: if store.clients.contains_key(service_key) { + store.clients[service_key].created_at + } else { + now + }, + }; + + store.clients.insert(service_key.to_string(), config); + self.save_client_configs(&store) + } + + pub fn delete_client_config(&self, service_key: &str) -> Result<(), CtlError> { + let mut store = self.load_client_configs(); + store.clients.remove(service_key); + self.save_client_configs(&store) + } +} diff --git a/ui/native/pb_mapper_ffi/src/state/runtime.rs b/ui/native/pb_mapper_ffi/src/state/runtime.rs new file mode 100644 index 0000000..9786fd3 --- /dev/null +++ b/ui/native/pb_mapper_ffi/src/state/runtime.rs @@ -0,0 +1,500 @@ +//! Runtime lifecycle for the embedded relay, registered services, and local clients. +//! +//! ```text +//! start relay: bind listener -> initialize app-local auth -> spawn -> mark running +//! register: resolved addresses -> spawn control pool -> retain JoinHandle +//! connect: preflight local bind -> spawn listener ----> retain JoinHandle +//! stop: cancel relay + abort owned tunnel tasks + clear runtime maps +//! ``` +//! +//! Readiness is published only after both listener binding and authentication +//! initialization succeed, preventing the UI from displaying a phantom running relay. + +use super::*; + +impl PbMapperState { + pub async fn start_server( + &mut self, + port: u16, + enable_keep_alive: bool, + ) -> Result<(), CtlError> { + if self.server_handle.is_some() { + return Err(CtlError::already_exists("Server is already running")); + } + + let ip_addr = IpAddr::V4(Ipv4Addr::new(0, 0, 0, 0)); + let bind_addr = std::net::SocketAddr::new(ip_addr, port); + + let listener = TcpListener::bind(bind_addr).await.map_err(|e| { + CtlError::address_in_use(format!("Failed to bind server on {bind_addr}: {e}")) + })?; + let auth_config = AuthConfig { + state_dir: self.config_dir.join("auth"), + ..AuthConfig::default() + }; + let auth = AuthRuntime::from_isolated_state(auth_config) + .await + .map_err(|error| { + CtlError::io(format!( + "Failed to initialize relay authentication: {error}" + )) + })?; + + tracing::info!("Starting pb-mapper server on {}:{}", ip_addr, port); + + let shutdown_token = CancellationToken::new(); + let shutdown_token_clone = shutdown_token.clone(); + + let (status_sender, status_receiver) = tokio::sync::mpsc::unbounded_channel(); + + self.server_auth = Some(auth.clone()); + let handle = tokio::spawn(async move { + if let Err(e) = run_server_on_listener( + listener, + shutdown_token_clone, + Some(status_receiver), + enable_keep_alive, + auth, + ) + .await + { + tracing::error!("pb-mapper server stopped with error: {e}"); + } + }); + + self.server_handle = Some(handle); + self.server_shutdown_token = Some(shutdown_token); + self.server_status_sender = Some(status_sender); + self.server_start_time = Some(SystemTime::now()); + + { + let mut cache = self.local_server_status_cache.write().await; + *cache = LocalServerStatus { + is_running: true, + active_connections: 0, + registered_services: 0, + uptime_seconds: 0, + }; + } + { + let mut last_update = self.local_server_status_last_update.write().await; + *last_update = Some(Instant::now()); + } + + tracing::info!("pb-mapper server started successfully"); + Ok(()) + } + + pub async fn stop_server(&mut self) -> Result<(), CtlError> { + if let (Some(handle), Some(shutdown_token)) = + (self.server_handle.take(), self.server_shutdown_token.take()) + { + self.server_status_sender = None; + + shutdown_token.cancel(); + + let shutdown_timeout = tokio::time::Duration::from_secs(5); + + let mut handle = handle; + match tokio::time::timeout(shutdown_timeout, &mut handle).await { + Ok(_) => { + tracing::info!("Server shutdown gracefully"); + } + Err(_) => { + if let Some(auth) = self.server_auth.as_ref() + && let Err(error) = auth.abort_actor().await + { + tracing::warn!( + "timed out waiting for the authentication actor to drop: {error}" + ); + } + handle.abort(); + let _ = handle.await; + tracing::warn!("Server shutdown timed out; aborted the relay task"); + } + } + if let Some(auth) = self.server_auth.take() { + auth.shutdown_actor().await; + } + + self.server_start_time = None; + + { + let mut cache = self.local_server_status_cache.write().await; + *cache = LocalServerStatus { + is_running: false, + active_connections: 0, + registered_services: 0, + uptime_seconds: 0, + }; + } + { + let mut last_update = self.local_server_status_last_update.write().await; + *last_update = Some(Instant::now()); + } + + for (_, runtime) in self.service_runtime.drain() { + runtime.handle.abort(); + } + for (_, runtime) in self.client_runtime.drain() { + runtime.handle.abort(); + } + + self.registered_services.write().await.clear(); + self.active_connections.write().await.clear(); + + tracing::info!("pb-mapper server stopped, all services and connections terminated"); + Ok(()) + } else { + Err(CtlError::not_found("Server is not running")) + } + } + + pub(super) async fn finish_register(&mut self, commit: RegisterCommit) -> Result<(), CtlError> { + let RegisterCommit { + service_key, + local_address, + protocol, + enable_encryption, + enable_keep_alive, + local_sock_addr, + remote_sock_addr, + credential, + } = commit; + + if abort_runtime(&mut self.service_runtime, &service_key) { + tracing::warn!( + "Service '{service_key}' is already registered, replacing existing handle" + ); + } + + tracing::info!( + "Registering service '{}' with protocol {}, local address {}, server address {}", + service_key, + protocol, + local_address, + self.config.server_address + ); + + self.save_service_config( + &service_key, + &local_address, + &protocol, + enable_encryption, + enable_keep_alive, + ) + .map_err(|e| CtlError::io(format!("Failed to save service configuration: {e}")))?; + + let key_clone = service_key.clone(); + let service_key_for_status = service_key.clone(); + + let callback: StatusCallback = Box::new(move |status: &str| { + tracing::info!( + "Service {} status update: {}", + service_key_for_status, + status + ); + }); + + let handle = spawn_register_tunnel( + &protocol, + local_sock_addr, + remote_sock_addr, + key_clone, + ServerTunnelOptions { + need_codec: enable_encryption, + is_datagram: !protocol.eq_ignore_ascii_case("TCP"), + keep_alive: enable_keep_alive, + namespace: None, + force_namespace: false, + }, + callback, + credential, + ); + replace_runtime( + &mut self.service_runtime, + &service_key, + TunnelRuntime { + handle, + pin: PinnedTunnel { + credential, + endpoint: remote_sock_addr, + }, + }, + ); + + { + let mut cache = self.service_status_cache.write().await; + cache.insert( + service_key.clone(), + StatusCacheEntry { + status: "retrying".to_string(), + message: "Connecting to pb-mapper server...".to_string(), + updated_at: Instant::now(), + }, + ); + } + self.schedule_service_status_refresh(&service_key).await; + + let service_info = ServiceInfo { + service_key: service_key.clone(), + protocol, + local_address, + status: "Registering".to_string(), + }; + + self.registered_services + .write() + .await + .insert(service_key.clone(), service_info); + + tracing::info!("Service '{}' registration initiated", service_key); + Ok(()) + } + + pub async fn unregister_service(&mut self, service_key: String) -> Result<(), CtlError> { + abort_runtime(&mut self.service_runtime, &service_key); + + if self + .registered_services + .write() + .await + .remove(&service_key) + .is_some() + { + tracing::info!("Service '{}' unregistered successfully", service_key); + Ok(()) + } else { + Err(CtlError::not_found(format!( + "Service '{service_key}' is not registered" + ))) + } + } + + pub async fn delete_service_config_and_stop( + &mut self, + service_key: String, + ) -> Result<(), CtlError> { + abort_runtime(&mut self.service_runtime, &service_key); + + self.registered_services.write().await.remove(&service_key); + + self.delete_service_config(&service_key) + } + + pub(super) async fn finish_connect(&mut self, commit: ConnectCommit) -> Result<(), CtlError> { + let ConnectCommit { + service_key, + local_address, + protocol, + enable_keep_alive, + local_sock_addr, + remote_sock_addr, + credential, + } = commit; + + if abort_runtime(&mut self.client_runtime, &service_key) { + tracing::warn!( + "Client for service '{service_key}' is already connected, replacing handle" + ); + } + + let protocol_upper = protocol.to_uppercase(); + + tracing::info!( + "Connecting to service '{}' with protocol {}, local address {}, server address {}", + service_key, + protocol, + local_address, + self.config.server_address + ); + + let key_clone = service_key.clone(); + + let status_callback: ClientStatusCallback = { + let service_key_for_callback = service_key.clone(); + Box::new(move |status: &str| { + tracing::info!("Client {} status: {}", service_key_for_callback, status); + }) + }; + + let handle = spawn_connect_tunnel( + &protocol_upper, + local_sock_addr, + remote_sock_addr, + key_clone, + enable_keep_alive, + status_callback, + credential, + ); + replace_runtime( + &mut self.client_runtime, + &service_key, + TunnelRuntime { + handle, + pin: PinnedTunnel { + credential, + endpoint: remote_sock_addr, + }, + }, + ); + + { + let mut cache = self.client_status_cache.write().await; + cache.insert( + service_key.clone(), + StatusCacheEntry { + status: "retrying".to_string(), + message: "Connecting to pb-mapper server...".to_string(), + updated_at: Instant::now(), + }, + ); + } + self.schedule_client_status_refresh(&service_key).await; + + let connection_info = ConnectionInfo { + service_key: service_key.clone(), + client_id: format!("client-{service_key}"), + status: "Connected".to_string(), + }; + + self.active_connections + .write() + .await + .insert(service_key.clone(), connection_info); + + // Persist here rather than at the FFI boundary, so a connection made + // from a terminal is remembered exactly like one made from the window. + // `finish_register` has always done this; leaving it out here meant a + // CLI `connect` started a client that never appeared in the list. + if let Err(e) = + self.save_client_config(&service_key, &local_address, &protocol, enable_keep_alive) + { + // The client is up either way, so this is a warning and not a + // failure: losing the config costs the entry after a restart. + tracing::warn!("Failed to save client config for '{service_key}': {e}"); + } + + tracing::info!("Connected to service '{}' successfully", service_key); + Ok(()) + } + + /// Claims a service key for a registration. See [`KeyClaim`]. + pub(super) fn claim_registering(&self, service_key: &str) -> Result { + claim_key(&self.registering, service_key, "being registered") + } + + /// Claims a service key for a client connection. See [`KeyClaim`]. + pub(super) fn claim_connecting(&self, service_key: &str) -> Result { + claim_key(&self.connecting, service_key, "being connected") + } + + pub async fn disconnect_service(&mut self, service_key: String) -> Result<(), CtlError> { + // Aborting the task is the part that matters: it is what stops the + // retry loop still dialling in the background. + let aborted = abort_runtime(&mut self.client_runtime, &service_key); + + let was_listed = self + .active_connections + .write() + .await + .remove(&service_key) + .is_some(); + + // Reported failure only when there was nothing to stop. It used to key + // off the bookkeeping map alone, so a client whose task had been + // aborted could still be reported as "not connected" — an error for an + // operation that had in fact just done its job. + if aborted || was_listed { + tracing::info!("Disconnected from service '{}'", service_key); + Ok(()) + } else { + Err(CtlError::not_found(format!( + "Service '{service_key}' is not connected" + ))) + } + } + + pub async fn delete_client_config_and_stop( + &mut self, + service_key: String, + ) -> Result<(), CtlError> { + abort_runtime(&mut self.client_runtime, &service_key); + + self.active_connections.write().await.remove(&service_key); + + self.delete_client_config(&service_key) + } +} + +fn spawn_register_tunnel( + protocol: &str, + local_sock_addr: SocketAddr, + remote_sock_addr: SocketAddr, + key: String, + options: ServerTunnelOptions, + callback: StatusCallback, + credential: Credential, +) -> JoinHandle<()> { + if protocol.eq_ignore_ascii_case("TCP") { + tokio::spawn(async move { + let _ = run_server_side_cli_with_pinned_credential::( + local_sock_addr, + remote_sock_addr, + key.into(), + options, + Some(callback), + credential, + ) + .await; + }) + } else { + tokio::spawn(async move { + let _ = run_server_side_cli_with_pinned_credential::( + local_sock_addr, + remote_sock_addr, + key.into(), + options, + Some(callback), + credential, + ) + .await; + }) + } +} + +fn spawn_connect_tunnel( + protocol: &str, + local_sock_addr: SocketAddr, + remote_sock_addr: SocketAddr, + key: String, + enable_keep_alive: bool, + callback: ClientStatusCallback, + credential: Credential, +) -> JoinHandle<()> { + if protocol.eq_ignore_ascii_case("TCP") { + tokio::spawn(async move { + run_client_side_cli_with_pinned_credential::( + local_sock_addr, + remote_sock_addr, + key.into(), + enable_keep_alive, + Some(callback), + credential, + ) + .await; + }) + } else { + tokio::spawn(async move { + run_client_side_cli_with_pinned_credential::( + local_sock_addr, + remote_sock_addr, + key.into(), + enable_keep_alive, + Some(callback), + credential, + ) + .await; + }) + } +} diff --git a/ui/native/pb_mapper_ffi/src/state/status.rs b/ui/native/pb_mapper_ffi/src/state/status.rs new file mode 100644 index 0000000..b4219cf --- /dev/null +++ b/ui/native/pb_mapper_ffi/src/state/status.rs @@ -0,0 +1,490 @@ +//! Non-blocking status views and bounded asynchronous refresh scheduling for Flutter. +//! +//! ```text +//! UI read -> cached snapshot -> immediate response +//! | +//! +-- stale? -> one deduplicated network refresh -> cache + change event +//! force refresh -----------------------------------------> awaited network result +//! ``` +//! +//! Service, client, and embedded-relay caches are independent. Refresh markers prevent +//! duplicate probes, while visible events are emitted only when the displayed state +//! actually changes. + +use super::*; + +impl PbMapperState { + pub async fn get_config_status(&self) -> AppConfig { + self.config.clone() + } + + pub fn isolated_admin_key(&self) -> Option { + let path = self.config_dir.join("auth").join("admin.key"); + std::fs::read_to_string(path).ok().and_then(|raw| { + let key = raw.trim().to_string(); + (!key.is_empty()).then_some(key) + }) + } + + pub async fn update_config( + &mut self, + server_address: String, + keep_alive: bool, + msg_header_key: String, + ) -> Result<(), CtlError> { + let msg_header_key = normalize_msg_header_key(msg_header_key)?; + let previous = self.config.clone(); + let candidate = AppConfig { + server_address, + keep_alive_enabled: keep_alive, + msg_header_key, + }; + self.write_config_file(&candidate)?; + self.config = candidate; + if let Err(error) = self.apply_msg_header_key_env() { + self.config = previous.clone(); + let _ = self.write_config_file(&previous); + let _ = self.apply_msg_header_key_env(); + return Err(error); + } + self.reset_status_caches().await; + Ok(()) + } + + pub async fn get_service_configs(&self) -> Vec { + let store = self.load_service_configs(); + let mut services = Vec::new(); + + let mut sorted_configs: Vec<_> = store.services.values().collect(); + sorted_configs.sort_by_key(|config| config.created_at); + + for config in sorted_configs { + let (status, message) = self.get_cached_service_status(&config.service_key).await; + + services.push(ServiceConfigInfo { + service_key: config.service_key.clone(), + local_address: config.local_address.clone(), + protocol: config.protocol.clone(), + enable_encryption: config.enable_encryption, + enable_keep_alive: config.enable_keep_alive, + status, + status_message: message, + created_at_ms: config + .created_at + .duration_since(SystemTime::UNIX_EPOCH) + .unwrap_or_default() + .as_millis() as u64, + updated_at_ms: SystemTime::now() + .duration_since(SystemTime::UNIX_EPOCH) + .unwrap_or_default() + .as_millis() as u64, + }); + } + + services + } + + pub async fn get_service_status(&self, service_key: String) -> ServiceStatusResponse { + let (status, message) = self.get_cached_service_status(&service_key).await; + ServiceStatusResponse { + service_key, + status, + message, + } + } + + pub async fn get_client_configs(&self) -> Vec { + let store = self.load_client_configs(); + let mut client_infos = Vec::new(); + + for (service_key, config) in store.clients.iter() { + let (status, status_message) = self.get_cached_client_status(service_key).await; + + client_infos.push(ClientConfigInfo { + service_key: config.service_key.clone(), + local_address: config.local_address.clone(), + protocol: config.protocol.clone(), + enable_keep_alive: config.enable_keep_alive, + status, + status_message, + created_at_ms: config + .created_at + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_millis() as u64, + updated_at_ms: config + .created_at + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_millis() as u64, + }); + } + + client_infos.sort_by_key(|info| info.created_at_ms); + client_infos + } + + pub async fn get_client_status(&self, service_key: String) -> ClientStatusResponse { + let (status, message) = self.get_cached_client_status(&service_key).await; + ClientStatusResponse { + service_key, + status, + message, + } + } + + pub async fn get_local_server_status(&self) -> LocalServerStatus { + let is_running = self.server_handle.is_some(); + if !is_running { + let status = LocalServerStatus { + is_running: false, + active_connections: 0, + registered_services: 0, + uptime_seconds: 0, + }; + { + let mut cache = self.local_server_status_cache.write().await; + *cache = status.clone(); + } + { + let mut last_update = self.local_server_status_last_update.write().await; + *last_update = Some(Instant::now()); + } + return status; + } + + let should_refresh = { + let last_update = self.local_server_status_last_update.read().await; + cache_is_stale(*last_update, STATUS_CACHE_TTL) + }; + + if should_refresh { + self.schedule_local_server_status_refresh(); + } + + let cache = self.local_server_status_cache.read().await; + cache.clone() + } + + fn schedule_local_server_status_refresh(&self) { + if self + .local_server_status_refreshing + .swap(true, Ordering::AcqRel) + { + return; + } + + let sender = self.server_status_sender.clone(); + let cache = self.local_server_status_cache.clone(); + let last_update = self.local_server_status_last_update.clone(); + let refreshing = self.local_server_status_refreshing.clone(); + let start_time = self.server_start_time; + + tokio::spawn(async move { + let mut status = LocalServerStatus { + is_running: true, + active_connections: 0, + registered_services: 0, + uptime_seconds: start_time + .and_then(|ts| SystemTime::now().duration_since(ts).ok()) + .map(|d| d.as_secs()) + .unwrap_or(0), + }; + + if let Some(sender) = sender { + let (response_sender, response_receiver) = tokio::sync::oneshot::channel(); + if sender.send(response_sender).is_ok() + && let Ok(Ok(info)) = + tokio::time::timeout(Duration::from_millis(200), response_receiver).await + { + status.active_connections = info.active_connections; + status.registered_services = info.registered_services; + status.uptime_seconds = info.uptime_seconds; + } + } + + { + let mut cache = cache.write().await; + *cache = status; + } + { + let mut last_update = last_update.write().await; + *last_update = Some(Instant::now()); + } + refreshing.store(false, Ordering::Release); + }); + } + + pub async fn get_server_status_detail(&self) -> Result { + self.force_refresh_server_status().await + } + + /// The connections the server holds for one key, from the protocol's own + /// structured query rather than the Debug dump in `server_map`. + pub async fn get_service_conns( + &self, + service_key: String, + ) -> Result, CtlError> { + let server_addr = self.config.server_address.clone(); + match tokio::time::timeout( + FORCE_REFRESH_TIMEOUT, + get_service_conns_with_addr(&server_addr, &service_key), + ) + .await + { + Ok(result) => result, + Err(_) => Err(CtlError::timeout(format!( + "Timed out asking {server_addr} about {service_key}" + ))), + } + } + + /// Perform a blocking status refresh — waits for the actual network result + /// instead of returning stale cache. + pub async fn force_refresh_server_status(&self) -> Result { + let server_addr = self.config.server_address.clone(); + + let detail = match tokio::time::timeout( + FORCE_REFRESH_TIMEOUT, + fetch_real_status_with_addr(&server_addr), + ) + .await + { + Ok(Ok((services, remote_id_data))) => ServerStatusDetail { + server_available: true, + registered_services: services, + server_map: remote_id_data.server_map, + active_connections: remote_id_data.active, + idle_connections: remote_id_data.idle, + }, + Ok(Err(e)) => { + tracing::warn!("Force refresh failed: {}", e); + ServerStatusDetail { + server_available: false, + registered_services: Vec::new(), + server_map: String::new(), + active_connections: String::new(), + idle_connections: String::new(), + } + } + Err(_) => { + tracing::warn!("Force refresh timed out after {:?}", FORCE_REFRESH_TIMEOUT); + ServerStatusDetail { + server_available: false, + registered_services: Vec::new(), + server_map: String::new(), + active_connections: String::new(), + idle_connections: String::new(), + } + } + }; + + Ok(detail) + } + + // Cache service status to avoid blocking UI with network checks on every paint. + async fn get_cached_service_status(&self, service_key: &str) -> (String, String) { + let Some(runtime) = self.service_runtime.get(service_key) else { + return ( + "stopped".to_string(), + "Service is not registered".to_string(), + ); + }; + if runtime.handle.is_finished() { + return ( + "failed".to_string(), + "Service connection terminated".to_string(), + ); + } + let cached = { + let cache = self.service_status_cache.read().await; + cache.get(service_key).cloned() + }; + if cached + .as_ref() + .map(|entry| entry.updated_at.elapsed() > STATUS_CACHE_TTL) + .unwrap_or(true) + { + self.schedule_service_status_refresh(service_key).await; + } + cached + .map(|entry| (entry.status, entry.message)) + .unwrap_or_else(|| { + ( + "retrying".to_string(), + "Checking service status...".to_string(), + ) + }) + } + + async fn get_cached_client_status(&self, service_key: &str) -> (String, String) { + let Some(runtime) = self.client_runtime.get(service_key) else { + return ("stopped".to_string(), "Client is not connected".to_string()); + }; + if runtime.handle.is_finished() { + return ( + "failed".to_string(), + "Client connection terminated".to_string(), + ); + } + let cached = { + let cache = self.client_status_cache.read().await; + cache.get(service_key).cloned() + }; + if cached + .as_ref() + .map(|entry| entry.updated_at.elapsed() > STATUS_CACHE_TTL) + .unwrap_or(true) + { + self.schedule_client_status_refresh(service_key).await; + } + cached + .map(|entry| (entry.status, entry.message)) + .unwrap_or_else(|| { + ( + "retrying".to_string(), + "Checking client status...".to_string(), + ) + }) + } + + pub(super) async fn schedule_service_status_refresh(&self, service_key: &str) { + { + let mut refreshing = self.service_status_refreshing.write().await; + if refreshing.contains(service_key) { + return; + } + refreshing.insert(service_key.to_string()); + } + + let tunnel = self + .service_runtime + .get(service_key) + .map(|runtime| runtime.pin); + let server_addr = tunnel + .map(|tunnel| tunnel.endpoint.to_string()) + .unwrap_or_else(|| self.config.server_address.clone()); + let cache = self.service_status_cache.clone(); + let refreshing = self.service_status_refreshing.clone(); + let key = service_key.to_string(); + let credential = tunnel.map(|tunnel| tunnel.credential); + + tokio::spawn(async move { + let result = tokio::time::timeout( + STATUS_REFRESH_TIMEOUT, + check_service_with_get_status(&server_addr, &key, credential), + ) + .await; + + let (status, message) = match result { + Ok(Ok(true)) => ( + "running".to_string(), + "Service is running normally".to_string(), + ), + Ok(Ok(false)) => ( + "retrying".to_string(), + "Service is in retry connection loop".to_string(), + ), + Ok(Err(_)) | Err(_) => ( + "failed".to_string(), + "Cannot connect to pb-server".to_string(), + ), + }; + + let changed = { + let mut cache = cache.write().await; + let changed = cache + .get(&key) + .is_none_or(|entry| entry.status != status || entry.message != message); + cache.insert( + key.clone(), + StatusCacheEntry { + status, + message, + updated_at: Instant::now(), + }, + ); + changed + }; + // Only transitions the user can perceive. These run on a timer for + // every configured entry, so emitting on every refresh would reload + // the list several times a second for no visible reason. + if changed { + events::emit(events::ChangeKind::Services, Some(&key), Origin::Internal); + } + + let mut refreshing = refreshing.write().await; + refreshing.remove(&key); + }); + } + + pub(super) async fn schedule_client_status_refresh(&self, service_key: &str) { + { + let mut refreshing = self.client_status_refreshing.write().await; + if refreshing.contains(service_key) { + return; + } + refreshing.insert(service_key.to_string()); + } + + let tunnel = self + .client_runtime + .get(service_key) + .map(|runtime| runtime.pin); + let server_addr = tunnel + .map(|tunnel| tunnel.endpoint.to_string()) + .unwrap_or_else(|| self.config.server_address.clone()); + let cache = self.client_status_cache.clone(); + let refreshing = self.client_status_refreshing.clone(); + let key = service_key.to_string(); + let credential = tunnel.map(|tunnel| tunnel.credential); + + tokio::spawn(async move { + let result = tokio::time::timeout( + STATUS_REFRESH_TIMEOUT, + check_service_with_get_status(&server_addr, &key, credential), + ) + .await; + + let (status, message) = match result { + Ok(Ok(true)) => ( + "running".to_string(), + "Client is connected normally".to_string(), + ), + Ok(Ok(false)) => ( + "retrying".to_string(), + "Client is in retry connection loop".to_string(), + ), + Ok(Err(_)) | Err(_) => ( + "failed".to_string(), + "Cannot connect to pb-server".to_string(), + ), + }; + + let changed = { + let mut cache = cache.write().await; + let changed = cache + .get(&key) + .is_none_or(|entry| entry.status != status || entry.message != message); + cache.insert( + key.clone(), + StatusCacheEntry { + status, + message, + updated_at: Instant::now(), + }, + ); + changed + }; + // Only transitions the user can perceive. These run on a timer for + // every configured entry, so emitting on every refresh would reload + // the list several times a second for no visible reason. + if changed { + events::emit(events::ChangeKind::Clients, Some(&key), Origin::Internal); + } + + let mut refreshing = refreshing.write().await; + refreshing.remove(&key); + }); + } +} diff --git a/ui/test/fake_pb_mapper_api.dart b/ui/test/fake_pb_mapper_api.dart index fa4220d..58b384c 100644 --- a/ui/test/fake_pb_mapper_api.dart +++ b/ui/test/fake_pb_mapper_api.dart @@ -36,6 +36,12 @@ class FakePbMapperApi implements PbMapperApiClient { @override Future fetchConfig() async => config; + @override + Future revealIsolatedRelayAdminKey() async { + calls.add('revealIsolatedRelayAdminKey'); + return config.isolatedRelayAdminKey; + } + @override Future updateConfig({ required String serverAddress, diff --git a/ui/test/widget_test.dart b/ui/test/widget_test.dart index 6c2e1ea..c7e2dda 100644 --- a/ui/test/widget_test.dart +++ b/ui/test/widget_test.dart @@ -818,7 +818,12 @@ void main() { await tester.tap(find.text('Next')); await tester.pump(); - expect(find.text('Must be exactly 32 characters'), findsOneWidget); + expect( + find.text( + 'Use a 32-character administrator key or a pbmt1_ temporary credential', + ), + findsOneWidget, + ); }); testWidgets('server-only mode starts at the server question', (tester) async {