From 84258829eded2f714ef51fd4b2c1243dff160a48 Mon Sep 17 00:00:00 2001 From: Simo Lin <25425177+slin1237@users.noreply.github.com> Date: Wed, 7 Oct 2026 09:19:44 -0700 Subject: [PATCH] feat(kv_index): the chain index, a lossless KV-event relay with state snapshots, and liveness-aware cache routing The gateway's cache-aware routing keeps an index of which worker holds which prefix of which prompt, fed by the engines' KV-cache event streams. This series replaces the design of that index, the relay that feeds it and the routing decisions around it. The chain index stores the engines' block-hash chains run-length compressed: a path-compressed trie of runs in an arena, one content hash per position, one coverage bit per worker per run plus a table of partial holders, children at any offset so a divergence inside a run does not split it, lock-free readers that write nothing, writers that lock one run at a time and descend without locks, recycled runs, arrays and tables, a per-lane open-addressing block map, a lane pool that schedules whole workers, and a sharded form (one index per NUMA node, lookups unioned). It is exact by construction against a single-threaded reference indexer added as a standing test, on recorded engine streams and on seeded corpora with evicted middle blocks, inside the crate and inside the gateway's own apply path. The positional indexer stays the default behind `--kv-index {positional,chain}` until the chain index has passed its gateway soak; its removal is a later change. The servicers' KV-event relay decodes both wire layouts of both engines into one model and normalizes per stream (tiers, cache groups, locality, ownership, namespaces, bigram pages, one counter per drop reason), forwards stores and removals one for one while the gateway counts physical copies per tier and rank, reproduces both engines' chain hashes so a misconfigured worker shows as a mismatch rate, subscribes every data-parallel rank, keeps a bounded history so a resume is served from it and a dropped batch is refilled from the engine's replay socket, keeps the engine's live-block record and serves it as a state snapshot to a subscriber whose cursor predates the history, starts at the servicer's boot and primes itself from the engine's replay, reads a publisher restart by rules that hold on every wire, and attaches the engine's load to every batch with heartbeats while it is quiet, the load poll kept only as a fallback. The vLLM servicers report queued uncached token-work, generation throughput and hit rate from their own bookkeeping. The gateway admits events through one cursor per (worker, rank) with bounded gap handling and snapshot resyncs; a liveness tracker beside the health check vetoes a worker whose connection failed and stayed silent or that holds requests without a token, steers around it without ever emptying the pool (the keepalive stays at 30 s pings, which the engines' grpc-core servers accept), and re-admits it on first contact with a closed circuit breaker; a thin or returned worker receives a slice of cache-miss traffic; every request's end reaches the policy that placed it through the worker's load guard; worker selection runs through a cost-function selection layer whose default reproduces the existing decision; worker overload protection is on by default as steering, never as shedding. The mock worker becomes a vLLM-style engine with the engines' ZMQ wires, fault hooks and engine truth, and its replay binary scores every routing decision against the fleet's arrival-time oracle and the engines' own cached-token counts. The stdout log sink no longer blocks runtime threads, and jemalloc purges on schedule. Measured in the strongest open-source KV router's own benchmark binary, both indexers behind the same lanes on the same host and the same cores, 20 fresh-process trials per point with interleaved same-binary controls and bootstrap intervals: the chain index sustains 1.44x that router's block operations per second under the published memory policy and 1.50x with node-local memory, at about a third of its lookup p99 and about one fifteenth of its bytes per indexed block; in the gateway it matches the positional path's hit rate with the lookup's p99 fifty times lower, and every fault drill of the recovery protocol passes in both topologies. Design: crates/kv_index/README.md. Benchmarks: crates/kv_index/benches/README.md. Signed-off-by: Simo Lin <25425177+slin1237@users.noreply.github.com> --- .gitignore | 2 + .pre-commit-config.yaml | 2 + Cargo.lock | 36 + bindings/python/src/lib.rs | 72 +- bindings/python/src/servicer.rs | 12 + bindings/python/src/smg/router_args.py | 200 +- bindings/python/tests/test_arg_parser.py | 101 +- bindings/python/tests/test_router_config.py | 33 + crates/engine_servicer/Cargo.toml | 9 + .../engine_servicer/benches/kv_relay_apply.rs | 321 ++ .../scripts/generate_kv_events_golden.py | 2 +- crates/engine_servicer/src/engine_hash.rs | 495 +++ crates/engine_servicer/src/kv_events.rs | 3418 +++++++++++++++-- crates/engine_servicer/src/kv_history.rs | 352 ++ crates/engine_servicer/src/kv_state.rs | 1022 +++++ crates/engine_servicer/src/kv_wire.rs | 2276 +++++++++++ .../src/kv_wire/shapes_tests.rs | 1296 +++++++ crates/engine_servicer/src/lib.rs | 15 +- crates/engine_servicer/src/load_tracker.rs | 283 ++ crates/engine_servicer/src/sglang/engine.rs | 4 +- crates/engine_servicer/src/sglang/mod.rs | 54 +- crates/engine_servicer/src/sglang/service.rs | 15 +- crates/engine_servicer/src/sglang/tests.rs | 350 +- crates/engine_servicer/src/testing.rs | 24 + .../engine_servicer/src/tokenspeed/engine.rs | 4 +- crates/engine_servicer/src/tokenspeed/mod.rs | 48 +- .../engine_servicer/src/tokenspeed/service.rs | 11 +- .../engine_servicer/src/tokenspeed/tests.rs | 306 +- crates/engine_servicer/src/vllm/engine.rs | 6 +- crates/engine_servicer/src/vllm/generate.rs | 48 +- crates/engine_servicer/src/vllm/info.rs | 35 +- crates/engine_servicer/src/vllm/mod.rs | 57 +- crates/engine_servicer/src/vllm/service.rs | 11 +- crates/engine_servicer/src/vllm/tests.rs | 417 +- .../tests/kv_event_snapshot.rs | 676 ++++ crates/engine_zmq_adapter/src/client.rs | 14 +- crates/engine_zmq_adapter/src/lib.rs | 2 +- crates/engine_zmq_adapter/src/sockets.rs | 100 +- crates/engine_zmq_client/src/mock_engine.rs | 46 +- crates/grpc_client/proto/common.proto | 183 +- .../grpc_client/proto/sglang_scheduler.proto | 4 +- .../proto/tokenspeed_scheduler.proto | 2 + crates/grpc_client/proto/vllm_engine.proto | 9 +- crates/grpc_client/src/channel.rs | 22 +- crates/grpc_client/src/engine_load.rs | 380 ++ crates/grpc_client/src/lib.rs | 1 + crates/grpc_client/src/sglang_scheduler.rs | 1 + .../grpc_client/src/tokenspeed_scheduler.rs | 1 + crates/grpc_client/src/vllm_engine.rs | 4 +- crates/kv_index/Cargo.toml | 17 + crates/kv_index/README.md | 232 ++ crates/kv_index/benches/README.md | 140 + crates/kv_index/benches/churn.rs | 320 ++ crates/kv_index/benches/match_insert.rs | 158 +- crates/kv_index/benches/mooncake_replay.rs | 2783 ++++++++++++++ crates/kv_index/src/chain_index.rs | 2132 ++++++++++ crates/kv_index/src/chain_index/arena.rs | 670 ++++ crates/kv_index/src/chain_index/slab.rs | 298 ++ crates/kv_index/src/chain_index/tests.rs | 1043 +++++ crates/kv_index/src/chain_index/walk.rs | 225 ++ crates/kv_index/src/churn.rs | 487 +++ crates/kv_index/src/event_tree.rs | 790 ++-- crates/kv_index/src/lane_map.rs | 504 +++ crates/kv_index/src/lane_pool.rs | 870 +++++ crates/kv_index/src/lib.rs | 21 +- crates/kv_index/src/prefetch.rs | 70 + crates/kv_index/src/reference.rs | 276 ++ crates/kv_index/src/salt.rs | 129 + crates/kv_index/src/sharded.rs | 604 +++ crates/kv_index/tests/churn_gate.rs | 94 + crates/kv_index/tests/common/mod.rs | 57 + crates/kv_index/tests/concurrency_chain.rs | 474 +++ crates/kv_index/tests/exactness_chain.rs | 937 +++++ crates/kv_index/tests/exactness_positional.rs | 731 ++++ crates/kv_index/tests/split_counters.rs | 75 + crates/mock_worker/Cargo.toml | 15 +- crates/mock_worker/README.md | 250 +- crates/mock_worker/src/admin.rs | 336 ++ crates/mock_worker/src/bin/replay.rs | 1535 ++++++++ crates/mock_worker/src/config.rs | 280 +- crates/mock_worker/src/engine.rs | 2808 ++++++++++++-- crates/mock_worker/src/grpc.rs | 230 +- crates/mock_worker/src/http.rs | 23 +- crates/mock_worker/src/kv_zmq.rs | 1006 +++++ crates/mock_worker/src/lib.rs | 2 + crates/mock_worker/src/main.rs | 5 +- crates/mock_worker/src/replay.rs | 6 +- crates/mock_worker/src/zmq.rs | 43 +- crates/mock_worker/tests/capture.rs | 11 +- crates/protocols/src/worker.rs | 15 + grpc_servicer/DEVELOPMENT.md | 15 + grpc_servicer/README.md | 17 +- grpc_servicer/scripts/gen_proto_stubs.py | 80 + grpc_servicer/smg_grpc_servicer/kv_relay.py | 1169 ++++++ .../smg_grpc_servicer/sglang/kv_events.py | 154 +- .../smg_grpc_servicer/sglang/rust.py | 29 + .../smg_grpc_servicer/sglang/servicer.py | 72 +- .../smg_grpc_servicer/tokenspeed/kv_events.py | 20 +- .../smg_grpc_servicer/tokenspeed/rust.py | 1 + .../smg_grpc_servicer/vllm/kv_events.py | 53 +- grpc_servicer/smg_grpc_servicer/vllm/loads.py | 167 + .../smg_grpc_servicer/vllm/servicer.py | 112 +- grpc_servicer/tests/conftest.py | 9 + grpc_servicer/tests/test_kv_relay.py | 1265 ++++++ grpc_servicer/tests/test_sglang_kv_events.py | 408 +- .../tests/test_sglang_rust_servicer.py | 29 + .../tests/test_tokenspeed_kv_events.py | 25 + .../tests/test_tokenspeed_rust_servicer.py | 4 +- grpc_servicer/tests/test_vllm_loads.py | 121 + model_gateway/Cargo.toml | 14 + model_gateway/benches/kv_index_decision.rs | 513 +++ model_gateway/benches/policy_selection.rs | 120 + model_gateway/benches/workers_endpoint.rs | 1 + model_gateway/src/app_context.rs | 44 +- model_gateway/src/config/builder.rs | 45 +- model_gateway/src/config/types.rs | 267 +- model_gateway/src/config/validation.rs | 32 + model_gateway/src/main.rs | 301 +- model_gateway/src/mesh/wiring.rs | 2 + model_gateway/src/observability/logging.rs | 311 +- model_gateway/src/observability/metrics.rs | 386 +- model_gateway/src/policies/cache_aware.rs | 2270 ++++++++++- model_gateway/src/policies/cache_namespace.rs | 52 +- model_gateway/src/policies/cost/accounting.rs | 423 ++ model_gateway/src/policies/cost/catalog.rs | 58 + model_gateway/src/policies/cost/default.rs | 92 + model_gateway/src/policies/cost/inputs.rs | 46 + model_gateway/src/policies/cost/mod.rs | 29 + model_gateway/src/policies/cost/policy.rs | 161 + model_gateway/src/policies/cost/sim_tests.rs | 324 ++ model_gateway/src/policies/cost/softmax.rs | 32 + model_gateway/src/policies/factory.rs | 6 + model_gateway/src/policies/least_load.rs | 405 +- model_gateway/src/policies/mod.rs | 24 + model_gateway/src/policies/registry.rs | 181 +- model_gateway/src/routers/common/overload.rs | 321 +- model_gateway/src/routers/common/placement.rs | 192 +- .../src/routers/common/worker_selection.rs | 65 +- .../grpc/common/stages/client_acquisition.rs | 23 +- .../grpc/common/stages/request_execution.rs | 17 +- .../grpc/common/stages/worker_selection.rs | 82 +- model_gateway/src/routers/grpc/context.rs | 70 + model_gateway/src/routers/grpc/pipeline.rs | 262 +- .../src/routers/grpc/proto_wrapper.rs | 223 +- .../grpc/regular/streaming/eof_tests.rs | 4 +- model_gateway/src/routers/http/pd_router.rs | 14 +- model_gateway/src/routers/http/router.rs | 20 +- model_gateway/src/worker/builder.rs | 1 + model_gateway/src/worker/kv_event_monitor.rs | 1375 +++---- .../src/worker/kv_event_monitor/admission.rs | 1214 ++++++ .../src/worker/kv_event_monitor/apply.rs | 1037 +++++ .../worker/kv_event_monitor/subscription.rs | 734 ++++ model_gateway/src/worker/kv_event_recovery.rs | 581 +++ model_gateway/src/worker/kv_index_backend.rs | 701 ++++ .../src/worker/kv_index_backend/exactness.rs | 1275 ++++++ .../worker/kv_index_backend/mock_streams.rs | 806 ++++ model_gateway/src/worker/liveness.rs | 994 +++++ model_gateway/src/worker/manager.rs | 26 +- model_gateway/src/worker/mod.rs | 10 +- model_gateway/src/worker/monitor.rs | 1000 ++++- model_gateway/src/worker/overload.rs | 84 +- model_gateway/src/worker/prefill_admission.rs | 33 + model_gateway/src/worker/registry.rs | 58 +- model_gateway/src/worker/worker.rs | 734 +++- .../src/workflow/steps/local/create_worker.rs | 14 +- .../steps/local/update_policies_for_worker.rs | 3 + .../steps/local/update_worker_properties.rs | 12 +- .../workflow/steps/shared/update_policies.rs | 3 + .../tests/allocator_artifact_test.rs | 76 +- model_gateway/tests/api/api_endpoints_test.rs | 8 +- model_gateway/tests/common/mock_worker.rs | 67 +- model_gateway/tests/common/mod.rs | 104 +- model_gateway/tests/common/test_app.rs | 4 +- .../tests/grpc_context_length_test.rs | 4 +- model_gateway/tests/grpc_pd_fanout_test.rs | 4 +- model_gateway/tests/pushed_loads_test.rs | 193 + model_gateway/tests/routing/mod.rs | 1 + .../tests/routing/model_alias_test.rs | 12 +- .../tests/routing/pd_routing_test.rs | 3 +- .../tests/routing/policy_completion_test.rs | 129 + .../tests/routing/stream_request_body_test.rs | 2 + .../tests/routing/test_openai_routing.rs | 11 +- .../tests/routing/test_pd_routing.rs | 2 + .../tests/tenant_rate_limiting_grpc_test.rs | 4 +- model_gateway/tests/zmq_backend_test.rs | 4 +- 185 files changed, 52451 insertions(+), 3538 deletions(-) create mode 100644 crates/engine_servicer/benches/kv_relay_apply.rs create mode 100644 crates/engine_servicer/src/engine_hash.rs create mode 100644 crates/engine_servicer/src/kv_history.rs create mode 100644 crates/engine_servicer/src/kv_state.rs create mode 100644 crates/engine_servicer/src/kv_wire.rs create mode 100644 crates/engine_servicer/src/kv_wire/shapes_tests.rs create mode 100644 crates/engine_servicer/src/load_tracker.rs create mode 100644 crates/engine_servicer/src/testing.rs create mode 100644 crates/engine_servicer/tests/kv_event_snapshot.rs create mode 100644 crates/grpc_client/src/engine_load.rs create mode 100644 crates/kv_index/README.md create mode 100644 crates/kv_index/benches/README.md create mode 100644 crates/kv_index/benches/churn.rs create mode 100644 crates/kv_index/benches/mooncake_replay.rs create mode 100644 crates/kv_index/src/chain_index.rs create mode 100644 crates/kv_index/src/chain_index/arena.rs create mode 100644 crates/kv_index/src/chain_index/slab.rs create mode 100644 crates/kv_index/src/chain_index/tests.rs create mode 100644 crates/kv_index/src/chain_index/walk.rs create mode 100644 crates/kv_index/src/churn.rs create mode 100644 crates/kv_index/src/lane_map.rs create mode 100644 crates/kv_index/src/lane_pool.rs create mode 100644 crates/kv_index/src/prefetch.rs create mode 100644 crates/kv_index/src/reference.rs create mode 100644 crates/kv_index/src/salt.rs create mode 100644 crates/kv_index/src/sharded.rs create mode 100644 crates/kv_index/tests/churn_gate.rs create mode 100644 crates/kv_index/tests/common/mod.rs create mode 100644 crates/kv_index/tests/concurrency_chain.rs create mode 100644 crates/kv_index/tests/exactness_chain.rs create mode 100644 crates/kv_index/tests/exactness_positional.rs create mode 100644 crates/kv_index/tests/split_counters.rs create mode 100644 crates/mock_worker/src/admin.rs create mode 100644 crates/mock_worker/src/bin/replay.rs create mode 100644 crates/mock_worker/src/kv_zmq.rs create mode 100755 grpc_servicer/scripts/gen_proto_stubs.py create mode 100644 grpc_servicer/smg_grpc_servicer/kv_relay.py create mode 100644 grpc_servicer/smg_grpc_servicer/vllm/loads.py create mode 100644 grpc_servicer/tests/test_kv_relay.py create mode 100644 grpc_servicer/tests/test_vllm_loads.py create mode 100644 model_gateway/benches/kv_index_decision.rs create mode 100644 model_gateway/benches/policy_selection.rs create mode 100644 model_gateway/src/policies/cost/accounting.rs create mode 100644 model_gateway/src/policies/cost/catalog.rs create mode 100644 model_gateway/src/policies/cost/default.rs create mode 100644 model_gateway/src/policies/cost/inputs.rs create mode 100644 model_gateway/src/policies/cost/mod.rs create mode 100644 model_gateway/src/policies/cost/policy.rs create mode 100644 model_gateway/src/policies/cost/sim_tests.rs create mode 100644 model_gateway/src/policies/cost/softmax.rs create mode 100644 model_gateway/src/worker/kv_event_monitor/admission.rs create mode 100644 model_gateway/src/worker/kv_event_monitor/apply.rs create mode 100644 model_gateway/src/worker/kv_event_monitor/subscription.rs create mode 100644 model_gateway/src/worker/kv_event_recovery.rs create mode 100644 model_gateway/src/worker/kv_index_backend.rs create mode 100644 model_gateway/src/worker/kv_index_backend/exactness.rs create mode 100644 model_gateway/src/worker/kv_index_backend/mock_streams.rs create mode 100644 model_gateway/src/worker/liveness.rs create mode 100644 model_gateway/tests/pushed_loads_test.rs create mode 100644 model_gateway/tests/routing/policy_completion_test.rs diff --git a/.gitignore b/.gitignore index eba0e97588..7aa05af323 100644 --- a/.gitignore +++ b/.gitignore @@ -116,6 +116,8 @@ work_dirs/ # Rust target directory target/ +# Older builds of the kv-index churn bench wrote this into the crate directory under cargo test +churn_result.json # VSCode .vscode diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 27395f5d7a..e5bebb2927 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -7,6 +7,8 @@ repos: - id: check-symlinks - id: destroyed-symlinks - id: trailing-whitespace + # Patch files keep their whitespace: a blank context line is a single space. + exclude: \.patch$ - id: end-of-file-fixer - id: check-yaml args: [--allow-multiple-documents, --unsafe] diff --git a/Cargo.lock b/Cargo.lock index 610eaf5e32..b7371a0a58 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2216,15 +2216,18 @@ version = "0.1.0" dependencies = [ "bytes", "chrono", + "criterion", "engine-zmq-adapter", "engine-zmq-client", "futures", + "kv-index", "llm-tokenizer", "openai-protocol", "portpicker", "prost", "prost-types", "rmp-serde", + "rmpv", "serde", "serde_json", "sha2 0.11.0", @@ -3823,11 +3826,16 @@ dependencies = [ name = "kv-index" version = "1.5.0" dependencies = [ + "anyhow", "bincode", "blake3", "clap", "criterion", + "crossbeam-queue", + "crossbeam-utils", "dashmap", + "libc", + "mimalloc", "once_cell", "parking_lot", "rand 0.10.2", @@ -3921,6 +3929,15 @@ version = "0.2.16" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b6d2cec3eae94f9f509c767b45932f1ada8350c4bdb85af2fcab4a3c14807981" +[[package]] +name = "libmimalloc-sys" +version = "0.1.49" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6a45a52f43e1c16f667ccfe4dd8c85b7f7c204fd5e3bf46c5b0db9a5c3c0b8e9" +dependencies = [ + "cc", +] + [[package]] name = "libredox" version = "0.1.20" @@ -4296,6 +4313,15 @@ dependencies = [ "sketches-ddsketch", ] +[[package]] +name = "mimalloc" +version = "0.1.52" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2d4139bb28d14ad1facf21d5eb8825051b326e172d216b39f6d31df53cc97862" +dependencies = [ + "libmimalloc-sys", +] + [[package]] name = "mime" version = "0.3.17" @@ -4365,9 +4391,16 @@ dependencies = [ name = "mock-worker" version = "0.1.0" dependencies = [ + "anyhow", "axum", + "clap", + "engine-servicer", "engine-zmq-client", "futures", + "reqwest 0.13.4", + "rmp-serde", + "rmpv", + "serde", "serde_json", "smg-grpc-client", "tempfile", @@ -4376,6 +4409,7 @@ dependencies = [ "tonic", "tracing", "tracing-subscriber", + "zeromq", ] [[package]] @@ -7528,6 +7562,7 @@ dependencies = [ "criterion", "dashmap", "data-connector", + "engine-servicer", "engine-zmq-adapter", "engine-zmq-client", "futures", @@ -7571,6 +7606,7 @@ dependencies = [ "regex", "reqwest 0.13.4", "rmcp", + "rmp-serde", "rsa", "rustix", "rustls", diff --git a/bindings/python/src/lib.rs b/bindings/python/src/lib.rs index 2d2cafae6b..11772fab90 100755 --- a/bindings/python/src/lib.rs +++ b/bindings/python/src/lib.rs @@ -549,6 +549,17 @@ struct Router { prefill_max_inflight_requests_per_worker: i32, prefill_queue_size: Option, prefill_queue_timeout_secs: Option, + worker_overload_shed: bool, + kv_index: String, + worker_stall_secs: u64, + worker_wedge_secs: u64, + worker_warmup_secs: u64, + worker_warmup_share: f32, + worker_warmup_blocks: usize, + worker_warmup_thin_ratio: f32, + worker_warmup_divert_every: u64, + selection_policy: String, + selection_accounting_ttl_ms: u64, /// The keyword-only `discovery` mapping, read by the same rules as /// `RouterConfig.discovery`. discovery: Option, @@ -663,6 +674,14 @@ impl Router { }) .transpose()?; + let kv_index = config::KvIndexKind::parse(&self.kv_index).ok_or_else(|| { + config::ConfigError::InvalidValue { + field: "kv_index".to_string(), + value: self.kv_index.clone(), + reason: "expected 'positional' or 'chain'".to_string(), + } + })?; + let convert_policy = |policy: &PolicyType| -> config::ConfigResult { Ok(match policy { PolicyType::Random => ConfigPolicyConfig::Random, @@ -682,6 +701,9 @@ impl Router { cache_index: self.parse_cache_index()?, cache_ttl_secs: self.cache_ttl_secs, cache_boundaries: self.cache_boundaries.clone(), + selection_policy: (self.selection_policy != policies::cost::DEFAULT_POLICY) + .then(|| self.selection_policy.clone()), + selection_accounting_ttl_ms: self.selection_accounting_ttl_ms, }, PolicyType::PowerOfTwo => ConfigPolicyConfig::PowerOfTwo { load_check_interval_secs: self.load_monitor_interval, @@ -899,6 +921,17 @@ impl Router { .worker_overload_waiting_requests(self.worker_overload_waiting_requests) .worker_overload_token_usage(self.worker_overload_token_usage) .worker_overload_protection(self.worker_overload_protection) + .worker_overload_shed(self.worker_overload_shed) + .worker_stall_secs(self.worker_stall_secs) + .worker_wedge_secs(self.worker_wedge_secs) + .worker_warmup( + self.worker_warmup_secs, + self.worker_warmup_share, + self.worker_warmup_blocks, + self.worker_warmup_thin_ratio, + self.worker_warmup_divert_every, + ) + .kv_index(kv_index) .disable_load_monitoring(self.disable_load_monitoring) .load_monitor_interval_secs(self.load_monitor_interval) .pd_admission_wait_secs(self.pd_admission_wait_secs) @@ -1159,9 +1192,9 @@ impl Router { cache_ttl_secs = 180, job_queue_capacity = 1000, job_queue_concurrency = 200, - worker_overload_waiting_requests = None, - worker_overload_token_usage = None, - worker_overload_protection = false, + worker_overload_waiting_requests = Some(8), + worker_overload_token_usage = Some(0.8), + worker_overload_protection = true, disable_load_monitoring = false, max_buffered_request_bytes = 1_048_576, kv_connector_annotation = String::from("smg.ai/kv-connector"), @@ -1184,6 +1217,17 @@ impl Router { prefill_max_inflight_requests_per_worker = -1, prefill_queue_size = None, prefill_queue_timeout_secs = None, + worker_overload_shed = false, + kv_index = String::from("positional"), + worker_stall_secs = 2, + worker_wedge_secs = 3, + worker_warmup_secs = 60, + worker_warmup_share = 0.25, + worker_warmup_blocks = 1024, + worker_warmup_thin_ratio = 0.5, + worker_warmup_divert_every = 8, + selection_policy = String::from("cache-aware-default"), + selection_accounting_ttl_ms = 0, // Keyword-only, so it never takes a positional slot. *, discovery = None, @@ -1352,6 +1396,17 @@ impl Router { prefill_max_inflight_requests_per_worker: i32, prefill_queue_size: Option, prefill_queue_timeout_secs: Option, + worker_overload_shed: bool, + kv_index: String, + worker_stall_secs: u64, + worker_wedge_secs: u64, + worker_warmup_secs: u64, + worker_warmup_share: f32, + worker_warmup_blocks: usize, + worker_warmup_thin_ratio: f32, + worker_warmup_divert_every: u64, + selection_policy: String, + selection_accounting_ttl_ms: u64, discovery: Option>, ) -> PyResult { // Two spellings of one choice: refuse both rather than pick one. @@ -1546,6 +1601,17 @@ impl Router { prefill_max_inflight_requests_per_worker, prefill_queue_size, prefill_queue_timeout_secs, + worker_overload_shed, + kv_index, + worker_stall_secs, + worker_wedge_secs, + worker_warmup_secs, + worker_warmup_share, + worker_warmup_blocks, + worker_warmup_thin_ratio, + worker_warmup_divert_every, + selection_policy, + selection_accounting_ttl_ms, discovery, }) } diff --git a/bindings/python/src/servicer.rs b/bindings/python/src/servicer.rs index 0c07bac1c2..e736eacc63 100644 --- a/bindings/python/src/servicer.rs +++ b/bindings/python/src/servicer.rs @@ -787,6 +787,7 @@ impl PyTokenSpeedGrpcServer { max_running_requests = 0, data_parallel_size = 1, kv_events_endpoint = String::new(), + kv_events_replay_endpoint = String::new(), kv_events_topic = String::new(), engine_startup_timeout_secs = None, ))] @@ -822,6 +823,7 @@ impl PyTokenSpeedGrpcServer { max_running_requests: i32, data_parallel_size: i32, kv_events_endpoint: String, + kv_events_replay_endpoint: String, kv_events_topic: String, engine_startup_timeout_secs: Option, ) -> PyResult { @@ -850,6 +852,7 @@ impl PyTokenSpeedGrpcServer { max_running_requests, data_parallel_size, kv_events_endpoint, + kv_events_replay_endpoint, kv_events_topic, }; let config = TokenSpeedServicerConfig { @@ -952,6 +955,9 @@ impl PySglangGrpcServer { sglang_version = String::new(), max_running_requests = 0, data_parallel_size = 1, + kv_events_endpoint = String::new(), + kv_events_replay_endpoint = String::new(), + kv_events_topic = String::new(), engine_startup_timeout_secs = None, ))] #[expect(clippy::too_many_arguments)] @@ -985,6 +991,9 @@ impl PySglangGrpcServer { sglang_version: String, max_running_requests: i32, data_parallel_size: i32, + kv_events_endpoint: String, + kv_events_replay_endpoint: String, + kv_events_topic: String, engine_startup_timeout_secs: Option, ) -> PyResult { let model = SglangModelInfo { @@ -1011,6 +1020,9 @@ impl PySglangGrpcServer { sglang_version, max_running_requests, data_parallel_size, + kv_events_endpoint, + kv_events_replay_endpoint, + kv_events_topic, }; let config = SglangServicerConfig { bind_address, diff --git a/bindings/python/src/smg/router_args.py b/bindings/python/src/smg/router_args.py index 76c0cde0dd..8d1d7b0128 100644 --- a/bindings/python/src/smg/router_args.py +++ b/bindings/python/src/smg/router_args.py @@ -248,11 +248,13 @@ class RouterArgs: # Control-plane job queue sizing (worker registration/removal jobs) job_queue_capacity: int = 1000 job_queue_concurrency: int = 200 - # Absolute per-worker overload thresholds; both None disables the feature - worker_overload_waiting_requests: int | None = None - worker_overload_token_usage: float | None = None - # Enable overload protection with the gateway default token ceiling (0.9) - worker_overload_protection: bool = False + # Absolute per-worker overload thresholds, as the Rust CLI defaults them; + # None switches one signal off + worker_overload_waiting_requests: int | None = 8 + worker_overload_token_usage: float | None = 0.8 + # Worker overload protection is on by default; False switches both + # thresholds off (--disable-worker-overload-protection) + worker_overload_protection: bool = True # Restore the conditional load-monitor poll gate (default: poll always) disable_load_monitoring: bool = False # Most bytes the router may buffer for a request it holds only to keep @@ -291,6 +293,25 @@ class RouterArgs: # The worker discovery provider. service_discovery=True is the legacy # spelling of discovery_provider="kubernetes"; set one or the other. discovery_provider: str | None = None + # Refuse with a 503 when every worker a request could use is overloaded, + # instead of routing it to the least-loaded one + worker_overload_shed: bool = False + # The event-driven KV index behind cache-aware routing: "positional" + # (one entry per block position) or "chain" (chains as runs) + kv_index: str = "positional" + # Liveness: seconds without contact after a transport failure before a + # worker is vetoed; seconds without progress before a loaded one is wedged + worker_stall_secs: int = 2 + worker_wedge_secs: int = 3 + # Warm-up slice for cache-aware routing (see the --worker-warmup-* flags) + worker_warmup_secs: int = 60 + worker_warmup_share: float = 0.25 + worker_warmup_blocks: int = 1024 + worker_warmup_thin_ratio: float = 0.5 + worker_warmup_divert_every: int = 8 + # The cache-aware selection policy and the optimistic accounting TTL + selection_policy: str = "cache-aware-default" + selection_accounting_ttl_ms: int = 0 @staticmethod def add_cli_args( @@ -607,12 +628,13 @@ def add_cli_args( default=RouterArgs.worker_overload_waiting_requests, help=( "Queued-request count AT OR ABOVE which a worker is considered" - " overloaded and excluded from routing until the signal recovers;" - " when every worker is overloaded, requests are shed immediately" - " rather than queued. Unset disables overload protection. This" - " signal is the queued (waiting) request count, summed across DP" - " ranks. Must be >= 1: the comparison is inclusive, so 0 would veto" - " every worker unconditionally." + " overloaded and left out of routing while another worker is" + " under the thresholds; when every worker is over them the" + " request goes to the least-loaded one (see" + " --worker-overload-shed). The signal is the queued (waiting)" + " request count, summed across DP ranks. Must be >= 1: the" + " comparison is inclusive, so 0 would veto every worker" + " unconditionally. Defaults to 8." ), ) routing_group.add_argument( @@ -621,10 +643,10 @@ def add_cli_args( default=RouterArgs.worker_overload_token_usage, help=( "KV-cache token usage AT OR ABOVE which a worker is considered" - " overloaded and excluded from routing until the signal recovers;" - " when every worker is overloaded, requests are shed immediately" - " rather than queued. Unset disables overload protection. This" - " signal is mean KV-cache token usage across DP ranks, the same one" + " overloaded and left out of routing while another worker is" + " under the thresholds (see --worker-overload-waiting-requests)." + " Defaults to 0.8. This signal is mean KV-cache token usage" + " across DP ranks, the same one" " --balance-token-usage-threshold reads, applied as an absolute" " per-worker CEILING rather than a fleet-relative spread. Backend" " must report token_usage. Must be in (0.0, 1.0]: the comparison is" @@ -637,16 +659,137 @@ def add_cli_args( routing_group.add_argument( f"--{prefix}worker-overload-protection", action="store_true", + default=None, + help=( + "Worker overload protection is on by default (8 waiting requests," + " 0.8 KV usage); this flag keeps it on and is accepted so older" + " command lines still parse. --disable-worker-overload-protection" + " switches both thresholds off." + ), + ) + routing_group.add_argument( + f"--{prefix}disable-worker-overload-protection", + action="store_true", + default=None, + help=( + "Switch worker overload protection off: no worker is left out of" + " routing for its waiting queue or KV usage (per-worker overload" + " blocks on a WorkerSpec still apply)." + ), + ) + routing_group.add_argument( + f"--{prefix}worker-overload-shed", + action="store_true", + help=( + "Refuse a request with a 503 (worker_overload_protection_shed," + " Retry-After the load poll interval) when every worker it could" + " use is overloaded, instead of routing it to the least-loaded one;" + " also sheds a worker that crossed a threshold between selection" + " and dispatch. Off by default." + ), + ) + routing_group.add_argument( + f"--{prefix}kv-index", + type=str, + default=RouterArgs.kv_index, + help=( + "The event-driven KV index behind cache-aware routing: 'positional'" + " (one entry per block position, the default) or 'chain' (chains as" + " runs with per-run worker coverage; lock-free, store-free lookups)." + ), + ) + routing_group.add_argument( + f"--{prefix}worker-stall-secs", + type=int, + default=RouterArgs.worker_stall_secs, + help=( + "Seconds without any contact from a worker (a load poll, a health" + " probe, a KV event, a response) after which a transport failure" + " excludes it from routing; the first successful contact re-admits" + " it. Defaults to 2." + ), + ) + routing_group.add_argument( + f"--{prefix}worker-wedge-secs", + type=int, + default=RouterArgs.worker_wedge_secs, + help=( + "Seconds without a token or a completion from a worker with requests" + " in flight whose waiting queue grows, or whose in-flight pile grows" + " or is four deep, after which new requests stop being routed to it" + " until it makes progress; the bound stretches to the time its" + " in-flight prompts may still need in prefill, up to 120 seconds." + " Defaults to 3." + ), + ) + routing_group.add_argument( + f"--{prefix}worker-warmup-secs", + type=int, + default=RouterArgs.worker_warmup_secs, + help=( + "Warm-up slice for cache-aware routing: for this many seconds after" + " a worker becomes routable, until its index has grown by" + " --worker-warmup-blocks blocks, one cache miss in" + " 1/--worker-warmup-share is routed to it so it builds a cache" + " instead of idling. 0 disables. Defaults to 60." + ), + ) + routing_group.add_argument( + f"--{prefix}worker-warmup-share", + type=float, + default=RouterArgs.worker_warmup_share, + help="Share of cache misses offered to warming workers (0.0 to 1.0). Defaults to 0.25.", + ) + routing_group.add_argument( + f"--{prefix}worker-warmup-blocks", + type=int, + default=RouterArgs.worker_warmup_blocks, + help=( + "A worker whose index has grown by this many blocks since it became" + " thin is warm. Defaults to 1024." + ), + ) + routing_group.add_argument( + f"--{prefix}worker-warmup-thin-ratio", + type=float, + default=RouterArgs.worker_warmup_thin_ratio, help=( - "Enable worker overload protection with the gateway default" - " thresholds. This flag alone applies" - " --worker-overload-token-usage 0.9 and leaves" - " --worker-overload-waiting-requests unset: KV token usage means" - " the same thing on every engine, while a sensible" - " waiting-requests ceiling is workload-dependent, so it has no" - " universal default. Explicit thresholds override the default," - " and either threshold set on its own enables protection without" - " this flag." + "A worker whose index holds less than this share of the fleet's" + " median (or nothing) is thin and receives the warm-up slice until" + " it has grown by --worker-warmup-blocks, whatever emptied it. 0" + " keeps the age rule alone. Defaults to 0.5." + ), + ) + routing_group.add_argument( + f"--{prefix}worker-warmup-divert-every", + type=int, + default=RouterArgs.worker_warmup_divert_every, + help=( + "One cache hit in this many is diverted to a thin worker although" + " another worker holds its prefix (shallow overlaps first, one in" + " flight per thin worker), so an index emptied by a resync refills" + " on a workload where every request has a holder. 0 disables." + " Defaults to 8." + ), + ) + routing_group.add_argument( + f"--{prefix}selection-policy", + type=str, + default=RouterArgs.selection_policy, + help=( + "The cache-aware worker selection policy. 'cache-aware-default', the" + " pre-policy decision, is the one policy." + ), + ) + routing_group.add_argument( + f"--{prefix}selection-accounting-ttl-ms", + type=int, + default=RouterArgs.selection_accounting_ttl_ms, + help=( + "Lifetime in milliseconds of optimistic dispatch bookings for" + " cache_aware: predicted prefill and prefix placement are charged to" + " the chosen worker until the engine reports them or the booking" + " expires. 0 disables; set a little above the engine's KV-event lag." ), ) routing_group.add_argument( @@ -1843,6 +1986,15 @@ def from_cli_args(cls, args: argparse.Namespace, use_router_prefix: bool = False # CLI args are tls_cert_path/tls_key_path # We need to manually map them if names don't match + # --disable-worker-overload-protection is the CLI spelling of + # worker_overload_protection=False; the prefixed form wins, the + # unprefixed one applies unless the fallback is disabled. + disable_protection = cli_args_dict.get(f"{prefix}disable_worker_overload_protection") + if disable_protection is None and not disable_arg_fallback: + disable_protection = cli_args_dict.get("disable_worker_overload_protection") + if disable_protection: + args_dict["worker_overload_protection"] = False + # Map tls args to server cert/key path if f"{prefix}tls_cert_path" in cli_args_dict: args_dict["server_cert_path"] = cli_args_dict[f"{prefix}tls_cert_path"] diff --git a/bindings/python/tests/test_arg_parser.py b/bindings/python/tests/test_arg_parser.py index c7ca855721..c6f054fbb0 100644 --- a/bindings/python/tests/test_arg_parser.py +++ b/bindings/python/tests/test_arg_parser.py @@ -645,7 +645,7 @@ def test_parse_cache_index_args(self): assert defaults.cache_ttl_secs == 180 def test_parse_worker_overload_args(self): - """Both overload flags round-trip, and both default to unset. + """Both overload flags round-trip, and both default as the Rust CLI does. The argparse names are built from an f-string prefix, so a typo or a dest/field mismatch would leave the field at its default and silently @@ -664,8 +664,8 @@ def test_parse_worker_overload_args(self): assert router_args.worker_overload_token_usage == pytest.approx(0.9) defaults = parse_router_args([]) - assert defaults.worker_overload_waiting_requests is None - assert defaults.worker_overload_token_usage is None + assert defaults.worker_overload_waiting_requests == 8 + assert defaults.worker_overload_token_usage == pytest.approx(0.8) def test_prefixed_worker_overload_args(self): """The --router-prefixed aliases reach the same fields.""" @@ -686,10 +686,11 @@ def test_prefixed_worker_overload_args(self): assert router_args.worker_overload_token_usage == pytest.approx(0.75) def test_parse_overload_protection_and_monitoring_flags(self): - """The enable/opt-out flags round-trip, and both default to False. + """Protection is on by default, as in the Rust CLI; the disable flag + switches it off and the legacy enable flag keeps it on. Same failure mode as the threshold flags: a dest/field mismatch would - silently disable the feature from Python. + silently change the feature from Python. """ router_args = parse_router_args( ["--worker-overload-protection", "--disable-load-monitoring"] @@ -698,9 +699,75 @@ def test_parse_overload_protection_and_monitoring_flags(self): assert router_args.disable_load_monitoring is True defaults = parse_router_args([]) - assert defaults.worker_overload_protection is False + assert defaults.worker_overload_protection is True assert defaults.disable_load_monitoring is False + disabled = parse_router_args(["--disable-worker-overload-protection"]) + assert disabled.worker_overload_protection is False + + def test_parse_overload_shed_liveness_warmup_index_and_selection_flags(self): + """Every RouterConfig field the Rust CLI exposes reaches RouterArgs, + with the CLI's defaults when the flags are absent.""" + router_args = parse_router_args( + [ + "--worker-overload-shed", + "--kv-index", + "chain", + "--worker-stall-secs", + "5", + "--worker-wedge-secs", + "7", + "--worker-warmup-secs", + "30", + "--worker-warmup-share", + "0.5", + "--worker-warmup-blocks", + "256", + "--worker-warmup-thin-ratio", + "0.25", + "--worker-warmup-divert-every", + "4", + "--selection-policy", + "cache-aware-default", + "--selection-accounting-ttl-ms", + "250", + ] + ) + assert router_args.worker_overload_shed is True + assert router_args.kv_index == "chain" + assert router_args.worker_stall_secs == 5 + assert router_args.worker_wedge_secs == 7 + assert router_args.worker_warmup_secs == 30 + assert router_args.worker_warmup_share == pytest.approx(0.5) + assert router_args.worker_warmup_blocks == 256 + assert router_args.worker_warmup_thin_ratio == pytest.approx(0.25) + assert router_args.worker_warmup_divert_every == 4 + assert router_args.selection_policy == "cache-aware-default" + assert router_args.selection_accounting_ttl_ms == 250 + + defaults = parse_router_args([]) + assert defaults.worker_overload_shed is False + assert defaults.kv_index == "positional" + assert defaults.worker_stall_secs == 2 + assert defaults.worker_wedge_secs == 3 + assert defaults.worker_warmup_secs == 60 + assert defaults.worker_warmup_share == pytest.approx(0.25) + assert defaults.worker_warmup_blocks == 1024 + assert defaults.worker_warmup_thin_ratio == pytest.approx(0.5) + assert defaults.worker_warmup_divert_every == 8 + assert defaults.selection_policy == "cache-aware-default" + assert defaults.selection_accounting_ttl_ms == 0 + + def test_prefixed_disable_overload_protection_flag(self): + """The --router-prefixed disable flag reaches the same field.""" + parser = argparse.ArgumentParser() + RouterArgs.add_cli_args(parser, use_router_prefix=True) + namespace = parser.parse_args(["--router-disable-worker-overload-protection"]) + + router_args = RouterArgs.from_cli_args(namespace, use_router_prefix=True) + + assert router_args.worker_overload_protection is False + def test_prefixed_overload_protection_and_monitoring_flags(self): """The --router-prefixed aliases reach the same fields.""" parser = argparse.ArgumentParser() @@ -1578,6 +1645,17 @@ class TestRouterArgsFieldOrder: "prefill_queue_size", "prefill_queue_timeout_secs", "discovery_provider", + "worker_overload_shed", + "kv_index", + "worker_stall_secs", + "worker_wedge_secs", + "worker_warmup_secs", + "worker_warmup_share", + "worker_warmup_blocks", + "worker_warmup_thin_ratio", + "worker_warmup_divert_every", + "selection_policy", + "selection_accounting_ttl_ms", ] def test_complete_field_sequence_is_frozen(self): @@ -1616,6 +1694,17 @@ def test_new_fields_appended_after_positional_reserve(self): "enable_rl", "rl_control_timeout_secs", "rl_fanout_concurrency", + "worker_overload_shed", + "kv_index", + "worker_stall_secs", + "worker_wedge_secs", + "worker_warmup_secs", + "worker_warmup_share", + "worker_warmup_blocks", + "worker_warmup_thin_ratio", + "worker_warmup_divert_every", + "selection_policy", + "selection_accounting_ttl_ms", ): assert names.index(appended) > marker, ( f"{appended} must be appended after worker_startup_delay to " diff --git a/bindings/python/tests/test_router_config.py b/bindings/python/tests/test_router_config.py index 5564292bc3..d0e5c9c179 100644 --- a/bindings/python/tests/test_router_config.py +++ b/bindings/python/tests/test_router_config.py @@ -469,3 +469,36 @@ def test_mapping_and_service_discovery_conflict(self): def test_unknown_provider_is_rejected(self): with pytest.raises(ValueError, match="Invalid discovery mapping"): _Router(worker_urls=[], discovery={"provider": "zookeeper"}) + + +class TestRouterPositionalSignature: + """External callers construct ``_Router(...)`` positionally, so the order + of its parameters is a public contract: a new parameter is appended after + the last positional one and before the keyword-only ``discovery``.""" + + @staticmethod + def positional_names(): + import inspect + + params = list(inspect.signature(_Router).parameters.values()) + assert params[-1].name == "discovery" + assert params[-1].kind is inspect.Parameter.KEYWORD_ONLY + return [p.name for p in params if p.kind is not inspect.Parameter.KEYWORD_ONLY] + + def test_appended_parameters_follow_prefill_queue_timeout(self): + names = self.positional_names() + tail = names[names.index("prefill_queue_timeout_secs") :] + assert tail == [ + "prefill_queue_timeout_secs", + "worker_overload_shed", + "kv_index", + "worker_stall_secs", + "worker_wedge_secs", + "worker_warmup_secs", + "worker_warmup_share", + "worker_warmup_blocks", + "worker_warmup_thin_ratio", + "worker_warmup_divert_every", + "selection_policy", + "selection_accounting_ttl_ms", + ] diff --git a/crates/engine_servicer/Cargo.toml b/crates/engine_servicer/Cargo.toml index c7f6cd7d52..29c9ef353c 100644 --- a/crates/engine_servicer/Cargo.toml +++ b/crates/engine_servicer/Cargo.toml @@ -36,8 +36,17 @@ zeromq.workspace = true [dev-dependencies] engine-zmq-client = { workspace = true, features = ["mock-engine"] } +criterion = { version = "0.8", features = ["html_reports"] } +# The normalized KV-event streams' round trip into the gateway's reference index. +kv-index.workspace = true +# The engines' KV-event wire shapes, built in code for the normalizer tests. +rmpv.workspace = true portpicker = "0.1" tempfile = "3" +[[bench]] +name = "kv_relay_apply" +harness = false + [lints] workspace = true diff --git a/crates/engine_servicer/benches/kv_relay_apply.rs b/crates/engine_servicer/benches/kv_relay_apply.rs new file mode 100644 index 0000000000..edc2eb14be --- /dev/null +++ b/crates/engine_servicer/benches/kv_relay_apply.rs @@ -0,0 +1,321 @@ +//! The relay's publisher-side cost and footprint with the live-block record +//! ([`engine_servicer::kv_state`]) next to the history and the normalizer's +//! own per-hash record, at the fleet's scale and above: eight data-parallel +//! ranks holding 676k blocks each (the 8B workers' pool token count taken as +//! a block count, sixteen times their 42k blocks of 16 tokens). +//! +//! Printed once at the start: the resident set after filling the state +//! without and with the live-block record (bytes per live block), and the +//! end-to-end rate of a real [`KvEventRelay`] absorbing a ZMQ publisher. +//! Then criterion times one incoming batch on the publisher task's path +//! (msgpack decode, normalize, history push; plus the record) in steady +//! state, every batch storing a 64-block chain and evicting one, so the live +//! set stays at its size: the throughput is reported in blocks per second. +//! +//! Run: `cargo bench -p engine-servicer --bench kv_relay_apply`. +#![allow( + clippy::expect_used, + clippy::unwrap_used, + clippy::print_stderr, + clippy::cast_possible_truncation, + clippy::cast_possible_wrap, + clippy::cast_precision_loss +)] + +use std::{ + fs, + sync::Arc, + time::{Duration, Instant}, +}; + +use criterion::{criterion_group, criterion_main, BatchSize, Criterion, Throughput}; +use engine_servicer::{ + kv_events::{KvEventRelay, RelayConfig, DEFAULT_HISTORY_BATCHES, DEFAULT_HISTORY_BYTES}, + kv_history::History, + kv_state::LiveState, + kv_wire::{Normalizer, WireBatch}, +}; +use engine_zmq_client::codec::TrailingTolerant; +use serde_json::json; +use zeromq::{prelude::*, PubSocket, ZmqMessage}; + +const RANKS: i32 = 8; +/// Live blocks per rank: 676,144 / 64 chains of 64 blocks. +const CHAINS_PER_RANK: u64 = 676_144 / BLOCKS_PER_BATCH as u64; +const BLOCKS_PER_BATCH: usize = 64; +const BLOCK_SIZE: usize = 16; +const TOKENS_PER_BATCH: usize = BLOCKS_PER_BATCH * BLOCK_SIZE; +/// Batches the ZMQ run publishes. +const ZMQ_BATCHES: u64 = 20_000; + +fn hashes(rank: i32, chain: u64) -> Vec { + let first = chain * BLOCKS_PER_BATCH as u64 + 1; + (first..first + BLOCKS_PER_BATCH as u64) + .map(|hash| (hash as i64) | (i64::from(rank) << 48)) + .collect() +} + +fn token(seed: u64) -> u32 { + let mut x = seed.wrapping_mul(0x9E37_79B9_7F4A_7C15) ^ 0xD1B5_4A32_D192_ED03; + x ^= x >> 29; + x = x.wrapping_mul(0xBF58_476D_1CE4_E5B9); + x ^= x >> 32; + (x & 0x7FFF) as u32 +} + +/// A vLLM-style publisher batch on `rank`: one `BlockStored` of chain +/// `chain` (64 blocks, 1,024 tokens, the engine's event keys) and, when +/// `evict` names a chain, one `BlockRemoved` of its 64 hashes. +fn payload(rank: i32, chain: u64, evict: Option) -> Vec { + let stored = hashes(rank, chain); + let token_ids: Vec = (0..TOKENS_PER_BATCH as u64) + .map(|i| token(chain * 4_096 + i)) + .collect(); + let mut events = vec![json!({ + "type": "BlockStored", + "block_hashes": stored, + "parent_block_hash": null, + "token_ids": token_ids, + "block_size": BLOCK_SIZE, + "lora_id": null, + "medium": "GPU", + "lora_name": null, + "group_idx": 0, + "kv_cache_spec_kind": "full_attention", + })]; + if let Some(old) = evict { + events.push(json!({ + "type": "BlockRemoved", + "block_hashes": hashes(rank, old), + "medium": "GPU", + "group_idx": 0, + })); + } + rmp_serde::to_vec_named(&json!([1_700_000_000.0, events, rank])).expect("encodes") +} + +/// The publisher task's per-batch work, with or without the live-block record. +struct Path { + normalizer: Normalizer, + history: History, + state: Option, + event_id: u64, + seq: u64, +} + +impl Path { + fn new(with_state: bool) -> Self { + Self { + normalizer: Normalizer::new(), + history: History::new(DEFAULT_HISTORY_BATCHES, DEFAULT_HISTORY_BYTES), + state: with_state.then(LiveState::new), + event_id: 0, + seq: 0, + } + } + + fn apply(&mut self, payload: &[u8]) { + let wire = rmp_serde::from_slice::>(payload) + .expect("decodes") + .0; + self.seq += 1; + let batch = Arc::new( + self.normalizer + .normalize_batch(wire, self.seq, &mut self.event_id), + ); + self.history.push(self.seq, Arc::clone(&batch)); + if let Some(state) = &mut self.state { + state.apply(&batch); + } + } + + /// Every rank at its full live set. + fn fill(&mut self) { + for chain in 0..CHAINS_PER_RANK { + for rank in 0..RANKS { + self.apply(&payload(rank, chain, None)); + } + } + } +} + +/// Resident set in bytes, from `/proc/self/status`. +fn rss_bytes() -> u64 { + let status = fs::read_to_string("/proc/self/status").expect("/proc/self/status"); + status + .lines() + .find_map(|line| line.strip_prefix("VmRSS:")) + .and_then(|rest| rest.split_whitespace().next()) + .and_then(|kib| kib.parse::().ok()) + .map_or(0, |kib| kib * 1024) +} + +fn mib(bytes: u64) -> f64 { + bytes as f64 / (1024.0 * 1024.0) +} + +/// Fill the publisher path without and with the record; print what each +/// adds to the resident set. +fn footprint() -> (Path, Path) { + let live_blocks = u64::from(RANKS as u32) * CHAINS_PER_RANK * BLOCKS_PER_BATCH as u64; + let baseline = rss_bytes(); + let started = Instant::now(); + let mut without = Path::new(false); + without.fill(); + let fill_without = started.elapsed(); + let after_without = rss_bytes(); + let started = Instant::now(); + let mut with = Path::new(true); + with.fill(); + let fill_with = started.elapsed(); + let after_with = rss_bytes(); + let history_and_normalizer = after_without.saturating_sub(baseline); + let record = after_with + .saturating_sub(after_without) + .saturating_sub(history_and_normalizer); + eprintln!( + "footprint: {} ranks x {} live blocks ({} batches of {} blocks per rank)\n \ + history + normalizer record: {:.0} MiB ({:.0} B per live block), filled in {:.1?} \ + ({:.0} blocks/s)\n \ + live-block record on top: {:.0} MiB ({:.0} B per live block), filled in {:.1?} \ + ({:.0} blocks/s)\n \ + record entries {} (copies {}), history {} batches / {:.0} MiB encoded", + RANKS, + CHAINS_PER_RANK * BLOCKS_PER_BATCH as u64, + CHAINS_PER_RANK, + BLOCKS_PER_BATCH, + mib(history_and_normalizer), + history_and_normalizer as f64 / live_blocks as f64, + fill_without, + live_blocks as f64 / fill_without.as_secs_f64(), + mib(record), + record as f64 / live_blocks as f64, + fill_with, + live_blocks as f64 / fill_with.as_secs_f64(), + with.state.as_ref().map_or(0, LiveState::entries), + with.state.as_ref().map_or(0, LiveState::blocks), + with.history.len(), + mib(with.history.bytes() as u64), + ); + (without, with) +} + +/// A real relay behind a ZMQ publisher, absorbing `ZMQ_BATCHES` batches of +/// 64 blocks as fast as the publisher sends them (a ring of 512 chains, so +/// the record sees repeats as capped copies). +fn relay_over_zmq() { + let runtime = tokio::runtime::Builder::new_multi_thread() + .worker_threads(2) + .enable_all() + .build() + .expect("runtime"); + runtime.block_on(async { + let mut publisher = PubSocket::new(); + let endpoint = publisher + .bind("tcp://127.0.0.1:0") + .await + .expect("publisher binds") + .to_string(); + let relay = KvEventRelay::new(RelayConfig { + endpoint, + replay_endpoint: None, + topic: "kv".to_string(), + history_batches: DEFAULT_HISTORY_BATCHES, + history_bytes: DEFAULT_HISTORY_BYTES, + replay_timeout: Duration::from_secs(2), + load_tick: Duration::from_millis(100), + heartbeat_interval: Duration::from_secs(1), + heartbeat_backoff: Duration::from_secs(5), + }); + relay.start(); + let frame = |seq: u64, body: &[u8]| { + let mut message = ZmqMessage::from(b"kv".to_vec()); + message.push_back(seq.to_be_bytes().to_vec().into()); + message.push_back(body.to_vec().into()); + message + }; + let ring: Vec> = (0..512u64).map(|chain| payload(0, chain, None)).collect(); + // The subscription lands a moment after the connect. + for _ in 0..500 { + publisher.send(frame(0, &ring[0])).await.expect("publish"); + if relay.counts().relayed >= 1 { + break; + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + assert_eq!(relay.counts().relayed, 1, "the relay joined the publisher"); + let started = Instant::now(); + for seq in 1..=ZMQ_BATCHES { + publisher + .send(frame(seq, &ring[(seq % 512) as usize])) + .await + .expect("publish"); + } + let published = started.elapsed(); + let counts = loop { + let counts = relay.counts(); + if counts.relayed + counts.gap_batches_lost > ZMQ_BATCHES + || started.elapsed() > Duration::from_secs(120) + { + break counts; + } + tokio::time::sleep(Duration::from_millis(5)).await; + }; + let elapsed = started.elapsed(); + let blocks = (counts.relayed - 1) * BLOCKS_PER_BATCH as u64; + eprintln!( + "relay over zmq: {} batches of {} blocks published in {:.2?}, relayed {} in {:.2?} \ + ({:.0} batches/s, {:.0} blocks/s); publisher gaps {}, batches lost {}", + ZMQ_BATCHES, + BLOCKS_PER_BATCH, + published, + counts.relayed - 1, + elapsed, + (counts.relayed - 1) as f64 / elapsed.as_secs_f64(), + blocks as f64 / elapsed.as_secs_f64(), + counts.publisher_gaps, + counts.gap_batches_lost, + ); + }); +} + +/// Steady state: each batch stores a new chain on the next rank and evicts +/// that rank's oldest, keeping every rank at its live set. +fn bench_apply(c: &mut Criterion, name: &str, mut path: Path) { + let mut next = [CHAINS_PER_RANK; RANKS as usize]; + let mut rank = 0usize; + let mut group = c.benchmark_group("relay_apply"); + group.throughput(Throughput::Elements(BLOCKS_PER_BATCH as u64)); + group.sample_size(30); + group.bench_function(name, |b| { + b.iter_batched( + || { + let r = rank; + rank = (rank + 1) % RANKS as usize; + let chain = next[r]; + next[r] += 1; + payload(r as i32, chain, Some(chain - CHAINS_PER_RANK)) + }, + |payload| path.apply(&payload), + BatchSize::SmallInput, + ); + }); + group.finish(); + if let Some(state) = &path.state { + assert_eq!( + state.blocks(), + u64::from(RANKS as u32) * CHAINS_PER_RANK * BLOCKS_PER_BATCH as u64, + "the live set stayed at its size" + ); + } +} + +fn benches(c: &mut Criterion) { + relay_over_zmq(); + let (without, with) = footprint(); + bench_apply(c, "decode_normalize_push", without); + bench_apply(c, "decode_normalize_push_record", with); +} + +criterion_group!(relay, benches); +criterion_main!(relay); diff --git a/crates/engine_servicer/scripts/generate_kv_events_golden.py b/crates/engine_servicer/scripts/generate_kv_events_golden.py index 20c301357d..6557ce0503 100644 --- a/crates/engine_servicer/scripts/generate_kv_events_golden.py +++ b/crates/engine_servicer/scripts/generate_kv_events_golden.py @@ -1,6 +1,6 @@ """Golden msgpack bytes for the Rust KV-event relay tests. -Regenerates the hex payloads in `crates/engine_servicer/src/vllm/kv_events.rs` +Regenerates the hex payloads in `crates/engine_servicer/src/kv_events.rs` (`golden` module) from vLLM's own `KVEventBatch` encoder, and prints the Python relay's conversion of them, which the tests expect byte for byte. Run it whenever vLLM changes the event layout, with a Python that has vllm, diff --git a/crates/engine_servicer/src/engine_hash.rs b/crates/engine_servicer/src/engine_hash.rs new file mode 100644 index 0000000000..8738997f96 --- /dev/null +++ b/crates/engine_servicer/src/engine_hash.rs @@ -0,0 +1,495 @@ +//! The engines' own KV block hashes, reproduced for verification. +//! +//! The gateway indexes SMG content hashes, never these; but recomputing what +//! an engine published tells whether a worker hashes the way its peers do +//! (same algorithm, seed, page size, token layout). The relay's opt-in check +//! ([`crate::kv_wire::Normalizer::with_hash_check`]) counts mismatches and +//! never drops an event, so a worker with a different hash setup shows up as +//! a counter long before it shows up as a routing miss. +//! +//! SGLang (`python/sglang/srt/mem_cache/cpp_utils/hash_binding.cpp` and +//! `mem_cache/utils.py` at b1bbd74f28): per page, `SHA256(prior digest || +//! page tokens as u32 little-endian)`; the first page of an unseeded chain +//! hashes the page bytes alone; a request's `cache_salt` seeds the chain with +//! `SHA256(b"sglang-cache-salt-v1\0" || salt)`; the prior of a child node is +//! the last page digest of its parent; the published integer is the first +//! eight digest bytes big-endian as a signed i64. Under Eagle-family bigram +//! hashing a page's words are `t, t+1` per position. +//! +//! vLLM (`vllm/v1/core/kv_cache_utils.py` and `vllm/utils/hashing.py` at +//! 0c16eee3f1) with `--prefix-caching-hash-algo sha256_cbor` and +//! `PYTHONHASHSEED` unset: `NONE = sha256(cbor("vllm-none-hash"))`, and +//! `H_i = sha256(cbor([bstr(H_{i-1} or NONE), [uint tokens...], null | extra +//! keys]))` where the extra keys are arrays in the order LoRA +//! `["lora", name, path]`, multimodal `["mm", identifier, offset]`, +//! `["cache_salt", salt]` (block 0 only) and `["prompt_embeds", bstr]`; +//! the published integer is the last eight digest bytes big-endian, unsigned +//! (carried here as the same 64 bits in an i64). The other algorithms +//! (`sha256` over pickle, `xxhash*`) are not reproducible outside the +//! engine's process and are not attempted. The LoRA adapter path is inside +//! the hash but not in the events, so a LoRA block cannot be verified from +//! its event alone. + +use sha2::{Digest, Sha256}; + +/// A full SHA-256 digest, what the engines chain on. +pub type Digest32 = [u8; 32]; + +/// Which engine hash a worker publishes. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum EngineHash { + /// SGLang's per-page SHA-256 chain. + Sglang, + /// vLLM's `sha256_cbor` with the default seed. + VllmSha256Cbor, +} + +impl EngineHash { + /// The setting's spelling: `sglang` or `vllm-sha256-cbor` (underscores + /// accepted), case-insensitive; `None` for anything else. + pub fn parse(name: &str) -> Option { + match name.trim().to_ascii_lowercase().replace('_', "-").as_str() { + "sglang" => Some(Self::Sglang), + "vllm-sha256-cbor" => Some(Self::VllmSha256Cbor), + _ => None, + } + } + + pub fn as_str(self) -> &'static str { + match self { + Self::Sglang => "sglang", + Self::VllmSha256Cbor => "vllm-sha256-cbor", + } + } +} + +// --------------------------------------------------------------------------- +// SGLang +// --------------------------------------------------------------------------- + +/// The chain seed of a salted request. +pub fn sglang_salt_seed(cache_salt: &str) -> Digest32 { + let mut hasher = Sha256::new(); + hasher.update(b"sglang-cache-salt-v1\0"); + hasher.update(cache_salt.as_bytes()); + hasher.finalize().into() +} + +/// One page: the prior digest (if any) then the page's words, each as four +/// little-endian bytes. Under bigram hashing pass `t, t+1` per position. +pub fn sglang_page(prior: Option<&Digest32>, words: &[u32]) -> Digest32 { + let mut hasher = Sha256::new(); + if let Some(prior) = prior { + hasher.update(prior); + } + for word in words { + hasher.update(word.to_le_bytes()); + } + hasher.finalize().into() +} + +/// The integer SGLang publishes for a digest. +pub fn sglang_event_int(digest: &Digest32) -> i64 { + let mut top = [0u8; 8]; + top.copy_from_slice(&digest[..8]); + i64::from_be_bytes(top) +} + +/// Every full page of `tokens` chained from `prior`: `(digest, published +/// integer)` per page, as a request starting at `prior` would be stored. +pub fn sglang_chain( + tokens: &[u32], + page_size: usize, + prior: Option<&Digest32>, +) -> Vec<(Digest32, i64)> { + if page_size == 0 { + return Vec::new(); + } + let mut prior = prior.copied(); + tokens + .chunks_exact(page_size) + .map(|page| { + let digest = sglang_page(prior.as_ref(), page); + prior = Some(digest); + (digest, sglang_event_int(&digest)) + }) + .collect() +} + +// --------------------------------------------------------------------------- +// vLLM sha256_cbor +// --------------------------------------------------------------------------- + +/// One of the keys vLLM folds into a block hash, as the hash sees it. +#[derive(Clone, Debug, PartialEq, Eq)] +pub enum VllmExtraKey { + Lora { name: String, path: String }, + Mm { identifier: String, offset: i64 }, + CacheSalt(String), + PromptEmbeds(Vec), +} + +const VLLM_NONE_HASH_SEED: &str = "vllm-none-hash"; + +/// `NONE_HASH`: the parent of a chain's first block. +pub fn vllm_none_hash() -> Digest32 { + let mut encoded = Vec::with_capacity(VLLM_NONE_HASH_SEED.len() + 1); + cbor::text(&mut encoded, VLLM_NONE_HASH_SEED); + Sha256::digest(&encoded).into() +} + +/// One block: `sha256(cbor([parent, tokens, extra_keys]))`. +pub fn vllm_block( + parent: Option<&Digest32>, + tokens: &[u32], + extra_keys: Option<&[VllmExtraKey]>, +) -> Digest32 { + let none = vllm_none_hash(); + let parent = parent.unwrap_or(&none); + let mut encoded = Vec::with_capacity(64 + tokens.len() * 3); + cbor::array(&mut encoded, 3); + cbor::bytes(&mut encoded, parent); + cbor::array(&mut encoded, tokens.len()); + for &token in tokens { + cbor::uint(&mut encoded, u64::from(token)); + } + match extra_keys { + None => cbor::null(&mut encoded), + Some(keys) => { + cbor::array(&mut encoded, keys.len()); + for key in keys { + match key { + VllmExtraKey::Lora { name, path } => { + cbor::array(&mut encoded, 3); + cbor::text(&mut encoded, "lora"); + cbor::text(&mut encoded, name); + cbor::text(&mut encoded, path); + } + VllmExtraKey::Mm { identifier, offset } => { + cbor::array(&mut encoded, 3); + cbor::text(&mut encoded, "mm"); + cbor::text(&mut encoded, identifier); + cbor::int(&mut encoded, *offset); + } + VllmExtraKey::CacheSalt(salt) => { + cbor::array(&mut encoded, 2); + cbor::text(&mut encoded, "cache_salt"); + cbor::text(&mut encoded, salt); + } + VllmExtraKey::PromptEmbeds(digest) => { + cbor::array(&mut encoded, 2); + cbor::text(&mut encoded, "prompt_embeds"); + cbor::bytes(&mut encoded, digest); + } + } + } + } + } + Sha256::digest(&encoded).into() +} + +/// The integer vLLM publishes for a digest (`VLLM_KV_EVENTS_USE_INT_BLOCK_HASHES`): +/// the low 64 bits, carried as the same bit pattern in an i64. +pub fn vllm_event_int(digest: &Digest32) -> i64 { + let mut low = [0u8; 8]; + low.copy_from_slice(&digest[24..]); + i64::from_be_bytes(low) +} + +/// Every full block of `tokens` chained from `parent` with no extra keys. +pub fn vllm_chain( + tokens: &[u32], + block_size: usize, + parent: Option<&Digest32>, +) -> Vec<(Digest32, i64)> { + if block_size == 0 { + return Vec::new(); + } + let mut parent = parent.copied(); + tokens + .chunks_exact(block_size) + .map(|block| { + let digest = vllm_block(parent.as_ref(), block, None); + parent = Some(digest); + (digest, vllm_event_int(&digest)) + }) + .collect() +} + +/// The subset of canonical CBOR (RFC 8949 §4.2) that vLLM's hash input uses: +/// definite-length arrays, byte and text strings, integers and null. `cbor2` +/// with `canonical=True` produces exactly these bytes for those values. +pub(crate) mod cbor { + fn head(out: &mut Vec, major: u8, value: u64) { + let major = major << 5; + if value < 24 { + out.push(major | value as u8); + } else if value <= u64::from(u8::MAX) { + out.push(major | 24); + out.push(value as u8); + } else if value <= u64::from(u16::MAX) { + out.push(major | 25); + out.extend_from_slice(&(value as u16).to_be_bytes()); + } else if value <= u64::from(u32::MAX) { + out.push(major | 26); + out.extend_from_slice(&(value as u32).to_be_bytes()); + } else { + out.push(major | 27); + out.extend_from_slice(&value.to_be_bytes()); + } + } + + pub(crate) fn uint(out: &mut Vec, value: u64) { + head(out, 0, value); + } + + pub(crate) fn int(out: &mut Vec, value: i64) { + if value >= 0 { + head(out, 0, value as u64); + } else { + head(out, 1, !(value as u64)); + } + } + + pub(crate) fn bytes(out: &mut Vec, value: &[u8]) { + head(out, 2, value.len() as u64); + out.extend_from_slice(value); + } + + pub(crate) fn text(out: &mut Vec, value: &str) { + head(out, 3, value.len() as u64); + out.extend_from_slice(value.as_bytes()); + } + + pub(crate) fn array(out: &mut Vec, len: usize) { + head(out, 4, len as u64); + } + + pub(crate) fn null(out: &mut Vec) { + out.push(0xf6); + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::kv_events::golden::bytes as hex; + + fn digest(hex_text: &str) -> Digest32 { + hex(hex_text).try_into().expect("32 bytes") + } + + // SGLang vectors published with the router's hash tests + // (experimental/sgl-router/src/state/kv_events/hash.rs). + #[test] + fn sglang_reproduces_the_published_vectors() { + let ints = |tokens: &[u32], page| -> Vec { + sglang_chain(tokens, page, None) + .into_iter() + .map(|(_, int)| int) + .collect() + }; + assert_eq!(ints(&[1, 2, 3, 4], 4), vec![-3488128144981237669]); + assert_eq!( + ints(&[10, 20, 30, 40, 50, 60, 70, 80], 2), + vec![ + 978178666101069530, + -895308556211281782, + -8033692805846017938, + 835415944263129316 + ] + ); + assert_eq!( + sglang_event_int(&Sha256::digest(b"").into()), + -2039914840885289964 + ); + assert!(ints(&[1, 2, 3], 4).is_empty(), "only full pages are hashed"); + assert!(sglang_chain(&[1, 2, 3, 4], 0, None).is_empty()); + } + + // Derived with hashlib from the algorithm above (vectors file in the + // branch notes): salted and bigram chains. + #[test] + fn sglang_salted_and_bigram_chains() { + let seed = sglang_salt_seed("tenant-a"); + assert_eq!( + seed, + digest("f5d0f785efe6042a4e4b0297a4d712917e1763c850835e93df58b373fd47fd2c") + ); + let salted: Vec = sglang_chain(&[1, 2, 3, 4, 5, 6, 7, 8], 4, Some(&seed)) + .into_iter() + .map(|(_, int)| int) + .collect(); + assert_eq!(salted, vec![3718046898735569995, 3308615664605479373]); + let plain = sglang_chain(&[1, 2, 3, 4, 5, 6, 7, 8], 4, None); + assert_eq!( + plain[0].0, + digest("cf97adeedb59e05bfd73a2b4c2a8885708c4f4f70c84c64b27120e72ab733b72") + ); + assert_eq!( + plain[1].0, + digest("4ebfa8a1f3c341517621838c6e1b9aa350307e3f00b3cbd1a07ef740f54396d6") + ); + // Bigram page (1,2) (2,3) (3,4) (4,5): both words of every pair. + let bigram = sglang_page(None, &[1, 2, 2, 3, 3, 4, 4, 5]); + assert_eq!(sglang_event_int(&bigram), -638950109823820341); + } + + // vLLM vectors: `hash_block_tokens` with `sha256_cbor` from the checkout + // at 0c16eee3f1, run under cbor2 with PYTHONHASHSEED unset. + #[test] + fn vllm_reproduces_the_reference_run() { + assert_eq!( + vllm_none_hash(), + digest("9bd96a485ad84efdafb72ee48a1d7a69bcead0f8f0433173941b276b9581eef0") + ); + let a = vllm_block(None, &[1, 2, 3, 4], None); + assert_eq!( + a, + digest("58d0879dff3800f65f8c5fd449d73048c0f7699dd238a03b84b146b151d55111") + ); + assert_eq!(vllm_event_int(&a), -8885242862429187823); + assert_eq!(vllm_event_int(&a) as u64, 9561501211280363793); + let b = vllm_block(Some(&a), &[5, 6, 7, 8], None); + assert_eq!( + b, + digest("e1bc29393a2da9fd798f35ff04d215f846aee6df4b803ce3d43b57df56b58fb1") + ); + assert_eq!(vllm_event_int(&b), -3153830497298837583); + let chain = vllm_chain(&[1, 2, 3, 4, 5, 6, 7, 8], 4, None); + assert_eq!( + chain, + vec![(a, vllm_event_int(&a)), (b, vllm_event_int(&b))] + ); + + let lora = VllmExtraKey::Lora { + name: "adapter".into(), + path: "/adapters/adapter".into(), + }; + let c = vllm_block( + None, + &[1, 2, 3, 4], + Some(&[ + lora.clone(), + VllmExtraKey::Mm { + identifier: "mm-abc".into(), + offset: 0, + }, + VllmExtraKey::CacheSalt("salt-1".into()), + VllmExtraKey::PromptEmbeds((0u8..32).collect()), + ]), + ); + assert_eq!( + c, + digest("280bce663478b44320e68be593605214c94f7648ac34b4b94274fdcd0f2a8a04") + ); + assert_eq!(vllm_event_int(&c), 4788731360966248964); + let d = vllm_block( + Some(&c), + &[5, 6, 7, 8], + Some(&[ + lora, + VllmExtraKey::Mm { + identifier: "mm-abc".into(), + offset: -4, + }, + ]), + ); + assert_eq!( + d, + digest("33e7394e25b7be46c7bf0c090da02e903b4e9d9536b1c18eed4d6e57e7a3658a") + ); + assert_eq!(vllm_event_int(&d), -1347299389686454902); + + let e = vllm_block( + None, + &[1, 2, 3, 4], + Some(&[ + VllmExtraKey::Mm { + identifier: "mm-abc".into(), + offset: 0, + }, + VllmExtraKey::CacheSalt("salt-1".into()), + VllmExtraKey::PromptEmbeds((0u8..32).collect()), + ]), + ); + assert_eq!(e, digest(VECTOR_E)); + let f = vllm_block( + Some(&e), + &[5, 6, 7, 8], + Some(&[VllmExtraKey::Mm { + identifier: "mm-abc".into(), + offset: -4, + }]), + ); + assert_eq!(f, digest(VECTOR_F)); + let g = vllm_block(Some(&f), &[9, 10, 11, 12], None); + assert_eq!(g, digest(VECTOR_G)); + } + + // The CBOR bytes `cbor2.dumps(value, canonical=True)` produced for the + // same inputs. + #[test] + fn cbor_matches_cbor2_canonical_bytes() { + let mut seed = Vec::new(); + cbor::text(&mut seed, VLLM_NONE_HASH_SEED); + assert_eq!(seed, hex("6e766c6c6d2d6e6f6e652d68617368")); + + let mut ints = Vec::new(); + cbor::array(&mut ints, 14); + for value in [ + 0i64, + 23, + 24, + 255, + 256, + 65535, + 65536, + 4_294_967_295, + 4_294_967_296, + ] { + cbor::int(&mut ints, value); + } + for value in [-1i64, -24, -25, -256, -257] { + cbor::int(&mut ints, value); + } + assert_eq!(ints, hex(CBOR_INTS)); + + // `[NONE, [1, 2, 3, 4], null]` and block D's input. + let mut a_input = Vec::new(); + cbor::array(&mut a_input, 3); + cbor::bytes(&mut a_input, &vllm_none_hash()); + cbor::array(&mut a_input, 4); + for token in 1..=4u64 { + cbor::uint(&mut a_input, token); + } + cbor::null(&mut a_input); + assert_eq!( + a_input, + hex("8358209bd96a485ad84efdafb72ee48a1d7a69bcead0f8f0433173941b276b9581eef08401020304f6") + ); + } + + #[test] + fn setting_names_parse() { + assert_eq!(EngineHash::parse("sglang"), Some(EngineHash::Sglang)); + assert_eq!(EngineHash::parse(" SGLang "), Some(EngineHash::Sglang)); + assert_eq!( + EngineHash::parse("vllm_sha256_cbor"), + Some(EngineHash::VllmSha256Cbor) + ); + assert_eq!( + EngineHash::parse("vllm-sha256-cbor"), + Some(EngineHash::VllmSha256Cbor) + ); + assert_eq!(EngineHash::parse("vllm"), None); + assert_eq!(EngineHash::parse(""), None); + assert_eq!(EngineHash::VllmSha256Cbor.as_str(), "vllm-sha256-cbor"); + } + + const VECTOR_E: &str = "0d7b0d8344b79ad5e183b117cacc04aeb415bfdfa968fd28f611bc6971f4abba"; + const VECTOR_F: &str = "b81a4631bec909b72bcd0081cc2d7c87bcce5652c48318f9615e4aa3370c1d69"; + const VECTOR_G: &str = "024cf74ecdb93c6049fe59a66321192aab04701a75ac3e86619395104df937c4"; + const CBOR_INTS: &str = + "8e0017181818ff19010019ffff1a000100001affffffff1b00000001000000002037381838ff390100"; +} diff --git a/crates/engine_servicer/src/kv_events.rs b/crates/engine_servicer/src/kv_events.rs index 0d5ef72ece..3c67dbff6f 100644 --- a/crates/engine_servicer/src/kv_events.rs +++ b/crates/engine_servicer/src/kv_events.rs @@ -1,65 +1,255 @@ -//! `SubscribeKvEvents`: vLLM's ZMQ KV-cache event publisher relayed into the -//! gRPC stream with the Python servicer's semantics (`kv_events.py`): the -//! publisher's own sequence numbers, one event-id counter per stream, bad -//! frames skipped, no replay, and the stream ending with the client. +//! `SubscribeKvEvents`: an engine's ZMQ KV-cache event publisher relayed into +//! gRPC streams through one [`KvEventRelay`] per engine. //! -//! Wire format (`vllm/distributed/kv_events.py`, `ZmqEventPublisher`): one -//! PUB multipart message per scheduler step, `[topic, sequence as u64 -//! big-endian, msgpack KVEventBatch]`. The batch is a msgspec `array_like` -//! struct, `[ts, events, data_parallel_rank]`; each event is a tagged map, -//! `{"type": "BlockStored" | "BlockRemoved" | "AllBlocksCleared", ...}`, whose -//! block hashes are sha256 bytes or 64-bit ints. +//! The relay subscribes to the publisher once, for the servicer's lifetime, +//! when the servicer starts serving (so it sees the engine from its first +//! batch, and a servicer that outlives a gateway outage or restarts beside a +//! warm engine is not blind until eviction; `SMG_KV_EVENT_RELAY_START=lazy` +//! defers the subscription to the first `SubscribeKvEvents` for a +//! memory-constrained host), keeps what it relays in a bounded [`History`] +//! (the last +//! `SMG_KV_EVENT_HISTORY_BATCHES` batches, 10,000 by default like the +//! engines' own `buffer_steps`, within `SMG_KV_EVENT_HISTORY_BYTES`, 256 MiB +//! by default) and folds every batch into a record of the engine's live +//! blocks per rank ([`LiveState`]). A `SubscribeKvEvents` call is served +//! from them: +//! +//! - `start_sequence_number` is the last sequence the gateway applied. Inside +//! the window, the batches after it come first, then live events; below +//! the window (or beyond the newest sequence, which means another publisher +//! incarnation) the call fails with `OUT_OF_RANGE`, and the gateway clears +//! and resubscribes from zero. A relay that started after the publisher +//! holds no history from before its start, so a cursor from before it is +//! `OUT_OF_RANGE` too. +//! - zero means no cursor. When the window holds every batch the publisher +//! ever numbered (it starts at 0, nothing lost or evicted) the subscriber +//! gets all of it, which is the publisher's whole state. Once the window no +//! longer starts there (it rolled past its caps, has a hole, or the relay +//! started after the publisher) the subscriber gets a state snapshot +//! instead: the live blocks as the relay recorded them, cut at the last +//! relayed sequence under the same lock that admits batches, sent as +//! chunks marked `KvSnapshotChunk` with consecutive sequence numbers +//! ending at that cut (chunk 0 begins with `AllBlocksCleared`, the stores +//! follow parents first), then live events from the next sequence: no +//! gap, no duplicate. The gateway applies the chunks as a `snapshot` +//! resync. A stale cursor below the window stays `OUT_OF_RANGE`: the +//! gateway clears, resubscribes from zero and receives the snapshot. Before +//! anything was relayed, zero is live only. +//! +//! The publisher's own sequence numbers are kept. A sequence the relay did +//! not receive is asked of the engine's replay socket (vLLM's and SGLang's +//! ROUTER, the mock engine's too); what the replay cannot give leaves a hole +//! in the window, counted, which a subscriber skips the way the live stream +//! did. A payload that does not decode relays as an empty batch under its +//! sequence, so nobody sees a gap for it. The batches before the first one +//! the relay sees are treated the same way: ZMQ delivers nothing from before +//! a subscription, so the relay asks the replay for everything from 0 when +//! it starts (retrying every second until the engine's replay answers, since +//! the engine may still be coming up), and again before relaying a first +//! live batch (of an incarnation) past sequence 1 (the publishers count from +//! 0, the mock engine from 1) unless that batch carries the engine's startup +//! clear, after which nothing earlier matters. The start replay is what +//! covers an engine that published a few batches at registration and nothing +//! since: without it the relay would never hear of them; and because those +//! batches may follow the first (empty) answer, a subscription from zero +//! that finds the relay still holding nothing asks the replay once more. +//! What the replay no longer reaches back to is counted as unknown before +//! the record's start (`unknown_before_start`): the window is then not +//! complete from the publisher's first batch, a subscriber from zero gets +//! the snapshot, and the snapshot's chunks carry `unknown_before` so the +//! gateway marks the worker degraded instead of trusting a silently partial +//! state. +//! +//! A publisher restart is read the same way on every wire, from three signs +//! (vLLM and SGLang count from 0 per process; SGLang's first batch after a +//! start carries `AllBlocksCleared`, vLLM's does not): the sequence goes +//! backwards on the same socket; the engine's startup clear arrives under a +//! sequence the relay already passed; the counter is back at 0 or 1 after a +//! cursor above them. Each starts a new incarnation: the history is cleared, +//! live subscribers end with `DATA_LOSS`, and the gateway clears and +//! resubscribes from zero, where the new incarnation's complete history is +//! waiting for it. A repeated sequence without a clear is a duplicate. +//! +//! Every batch a subscriber receives carries the servicer's load record +//! ([`LoadSource`], the figures `GetLoads` answers with, as `KvEventBatch. +//! load`), and a subscriber hears from the relay even when the publisher is +//! quiet: a batch with no events and the last sequence repeated, marked +//! `load_only`, goes out when the record changed (checked every +//! `load_tick`, 100 ms) and as a heartbeat after `heartbeat_interval` (1 s) +//! of silence, backing off to `heartbeat_backoff` (5 s) once the engine has +//! been idle for two heartbeats. The record's routing core rides every +//! batch; the engine's telemetry (cache hit rate, token counts, SGLang's +//! sections) rides the heartbeats and the stream's first record. The +//! gateway feeds the record where its `GetLoads` poll goes, admits nothing +//! from a `load_only` batch, and does not poll a worker while its records +//! arrive: `GetLoads` is the fallback. +//! +//! `SMG_KV_EVENT_HASH_CHECK=sglang|vllm-sha256-cbor` turns on the relay's +//! engine-hash verification ([`crate::engine_hash`]); mismatches are counted, +//! never dropped. The relay's counters ([`RelayCounts`]) are logged on every +//! gap, restart and refusal, and together with the normalizer's +//! ([`crate::kv_wire::Counts`]: forwarded, dropped by reason, the hash +//! check's tally) in a summary line every 500 relayed batches and when the +//! relay closes. +//! +//! Framing (`ZmqEventPublisher` in both engines): one PUB multipart message +//! per scheduler step, `[topic, sequence as u64 big-endian, msgpack batch]`. +//! The wire format and the normalization each event goes through live in +//! [`crate::kv_wire`]. -use std::fmt; +use std::{ + collections::VecDeque, + ops::RangeInclusive, + sync::{Arc, Mutex, MutexGuard, OnceLock, PoisonError}, + time::{Duration, Instant, SystemTime, UNIX_EPOCH}, +}; use engine_zmq_client::codec::TrailingTolerant; use futures::stream; -use serde::{ - de::{self, Visitor}, - Deserialize, Deserializer, -}; -use smg_grpc_client::common_proto::{self as common, kv_cache_event}; +use smg_grpc_client::common_proto::{self as common}; +use tokio::sync::{broadcast, Notify}; use tonic::Status; use tracing::{debug, info, warn}; use zeromq::{ - prelude::{Socket, SocketRecv}, - SocketOptions, SubSocket, ZmqError, ZmqMessage, + prelude::{Socket, SocketRecv, SocketSend}, + DealerSocket, SocketOptions, SubSocket, ZmqError, ZmqMessage, }; -use crate::BoxStream; +use crate::{ + kv_history::{History, Window}, + kv_state::{LiveState, Snapshot, SnapshotChunks}, + kv_wire::{low64_big_endian, Counts as WireCounts, Normalizer, WireBatch, WireEvent}, + BoxStream, +}; /// The Python vLLM servicer's refusal when vLLM runs without a ZMQ publisher. pub(crate) const VLLM_DISABLED_MESSAGE: &str = "KV cache events not enabled. Start vLLM with \ --kv-events-config '{\"enable_kv_cache_events\": true, \"publisher\": \"zmq\"}'"; +/// The Python SGLang servicer's refusal when SGLang runs without a ZMQ publisher. +pub(crate) const SGLANG_DISABLED_MESSAGE: &str = "KV cache events not enabled. Start SGLang \ + with --kv-events-config '{\"publisher\": \"zmq\"}'"; + /// The Python TokenSpeed servicer's refusal without a publisher. pub(crate) const TOKENSPEED_DISABLED_MESSAGE: &str = "KV cache events not enabled. Start \ TokenSpeed with --kv-events-config '{\"enable_kv_cache_events\": true, \"publisher\": \ \"zmq\"}'"; -/// Handle one `SubscribeKvEvents` call against a configured publisher: a -/// relay of rank 0's publisher at `kv_events_endpoint`. Both engines publish -/// the same batch shape (`[ts, events, rank]`, the rank named -/// `data_parallel_rank` by vLLM and `attn_dp_rank` by TokenSpeed; positional -/// on the wire). -pub(crate) fn subscribe( - kv_events_endpoint: &str, - topic: String, - request: common::SubscribeKvEventsRequest, -) -> BoxStream { - // For DP attention each rank publishes on port + rank with independent - // sequence counters; subscribing to several on one socket interleaves - // them and breaks gap detection. Subscribe to rank 0 only for now. - let endpoint = endpoint_for_rank(kv_events_endpoint, 0); - if request.start_sequence_number != 0 { - // As on the Python relay: no replay, the stream starts at the - // publisher's current position and the Router dedups by sequence. - debug!( - start_sequence_number = request.start_sequence_number, - "SubscribeKvEvents: replay is not supported; streaming live events" - ); - } - relay(endpoint, topic) +/// When the relay subscribes to the publisher: at the servicer's start +/// (default, any other value) or `lazy`, at the first `SubscribeKvEvents`. +pub(crate) const RELAY_START_ENV: &str = "SMG_KV_EVENT_RELAY_START"; +/// How many relayed batches the history keeps (the engines' `buffer_steps`). +pub(crate) const HISTORY_BATCHES_ENV: &str = "SMG_KV_EVENT_HISTORY_BATCHES"; +/// The history's byte budget over the encoded batches. +pub(crate) const HISTORY_BYTES_ENV: &str = "SMG_KV_EVENT_HISTORY_BYTES"; +pub const DEFAULT_HISTORY_BATCHES: usize = 10_000; +pub const DEFAULT_HISTORY_BYTES: usize = 256 << 20; +const DEFAULT_REPLAY_TIMEOUT: Duration = Duration::from_secs(5); +/// How often the start replay is retried while the engine's replay socket +/// does not answer yet. +const PRIME_RETRY: Duration = Duration::from_secs(1); +/// Live batches a slow subscriber may fall behind before it is refilled from +/// the history. +const LIVE_CHANNEL: usize = 4_096; +/// A relay summary line (relay and normalizer counters) every this many +/// relayed batches, so a live run shows them before the relay closes. +const SUMMARY_EVERY_BATCHES: u64 = 500; + +/// How often a subscriber's stream checks the load record for a change. +pub(crate) const DEFAULT_LOAD_TICK: Duration = Duration::from_millis(100); +/// Silence after which a `load_only` heartbeat goes out. +pub(crate) const DEFAULT_HEARTBEAT_INTERVAL: Duration = Duration::from_secs(1); +/// The heartbeat interval once the engine has been idle for two heartbeats. +pub(crate) const DEFAULT_HEARTBEAT_BACKOFF: Duration = Duration::from_secs(5); +/// The replay socket's end marker, eight 0xff bytes on both wires. +const END_SEQUENCE: [u8; 8] = [0xff; 8]; + +/// One publisher's relay settings. +#[derive(Clone, Debug)] +pub struct RelayConfig { + /// The connectable SUB endpoint of rank 0's publisher. + pub endpoint: String, + /// Its replay ROUTER, when the engine runs one. + pub replay_endpoint: Option, + pub topic: String, + pub history_batches: usize, + pub history_bytes: usize, + /// How long a replay reply may take before the rest of a gap is lost. + pub replay_timeout: Duration, + /// How often a subscriber's stream checks the load record for a change + /// while no batch flows. + pub load_tick: Duration, + /// Silence after which a `load_only` heartbeat goes out. + pub heartbeat_interval: Duration, + /// The heartbeat interval after two heartbeats with an unchanged record. + pub heartbeat_backoff: Duration, +} + +impl RelayConfig { + /// Rank 0 of the publisher at `kv_events_endpoint` (bind wildcards + /// resolved), with the history caps from the environment. + pub(crate) fn for_publisher( + kv_events_endpoint: &str, + replay_endpoint: Option<&str>, + topic: &str, + ) -> Self { + Self { + endpoint: endpoint_for_rank(kv_events_endpoint, 0), + replay_endpoint: replay_endpoint + .filter(|endpoint| !endpoint.is_empty()) + .map(|endpoint| endpoint_for_rank(endpoint, 0)), + topic: topic.to_string(), + history_batches: env_usize(HISTORY_BATCHES_ENV, DEFAULT_HISTORY_BATCHES), + history_bytes: env_usize(HISTORY_BYTES_ENV, DEFAULT_HISTORY_BYTES), + replay_timeout: DEFAULT_REPLAY_TIMEOUT, + load_tick: DEFAULT_LOAD_TICK, + heartbeat_interval: DEFAULT_HEARTBEAT_INTERVAL, + heartbeat_backoff: DEFAULT_HEARTBEAT_BACKOFF, + } + } +} + +/// Where the relay reads the engine's load for the record it attaches to +/// every batch: the servicer's `GetLoads` bookkeeping. +pub(crate) trait LoadSource: Send + Sync { + /// The load for `dp_rank` (the batch's rank; `None` for a publisher that + /// names none) as `GetLoads` would report it now, or `None` while nothing + /// is known (the engine is not up yet). `sample` and `load_only` are the + /// relay's to set. + fn load(&self, dp_rank: Option) -> Option; +} + +/// Whether a record moved enough from `last` to be worth a `load_only` +/// batch: any queue, running or window change, KV usage by half a percent, +/// the rate by 5 % or 50 tokens per second. +fn load_changed(last: &common::EngineLoad, current: &common::EngineLoad) -> bool { + last.running_requests != current.running_requests + || last.waiting_requests != current.waiting_requests + || last.waiting_uncached_tokens != current.waiting_uncached_tokens + || last.max_running_requests != current.max_running_requests + || (last.token_usage - current.token_usage).abs() > 0.005 + || { + let delta = (last.gen_throughput - current.gen_throughput).abs(); + delta > 50.0 || delta > 0.05 * last.gen_throughput.max(current.gen_throughput) + } +} + +/// Whether [`RELAY_START_ENV`] set to `value` keeps the start at boot. +fn starts_at_boot(value: Option<&str>) -> bool { + !value.is_some_and(|value| value.trim().eq_ignore_ascii_case("lazy")) +} + +fn env_usize(name: &str, default: usize) -> usize { + match std::env::var(name) { + Ok(value) if !value.trim().is_empty() => match value.trim().parse::() { + Ok(parsed) => parsed, + Err(_) => { + warn!(%value, "{name} is not a number; using {default}"); + default + } + }, + _ => default, + } } /// Resolve a KV-events PUB endpoint to a connectable SUB address: bind @@ -85,369 +275,1317 @@ pub(crate) fn endpoint_for_rank(endpoint: &str, dp_rank: u32) -> String { } } -type Item = Result; +/// What one relay has done since it started. +#[derive(Clone, Debug, Default, PartialEq, Eq)] +pub struct RelayCounts { + /// Batches relayed under the publisher's sequence (live or replayed). + pub relayed: u64, + /// Payloads that did not decode and were relayed as empty batches. + pub undecodable_batches: u64, + /// Subscriptions whose first batches came from the history. + pub served_from_history: u64, + /// Subscriptions that began with a state snapshot. + pub served_snapshots: u64, + /// Subscriptions refused because their cursor was outside the window. + pub out_of_range: u64, + /// Sequence gaps seen on the publisher's socket. + pub publisher_gaps: u64, + /// Batches of those gaps the engine's replay gave back. + pub gap_batches_recovered: u64, + /// Batches of those gaps nobody had: holes in the window. + pub gap_batches_lost: u64, + /// Publisher sequences before the relay's record that neither the stream + /// nor the engine's replay gave (summed over incarnations): how late the + /// history and the snapshots start. + pub unknown_before_start: u64, + /// Batches taken from the publisher's replay at the relay's start, before + /// any live batch reached it. + pub primed_batches: u64, + /// Publisher incarnations after the first. + pub publisher_restarts: u64, + /// Subscribers that fell behind the live channel and were refilled. + pub subscribers_lagged: u64, +} -/// The relay as a response stream: it connects on its first poll, so the -/// RPC's headers go out as soon as the handler returns (the Python relay -/// sends its initial metadata before its first receive), then yields one -/// proto batch per publisher message. Dropping it closes the socket. -fn relay(endpoint: String, topic: String) -> BoxStream { - Box::pin(stream::unfold( - Relay::Connecting { endpoint, topic }, - |relay| async move { relay.step().await }, - )) -} - -enum Relay { - Connecting { endpoint: String, topic: String }, - Live(Live), - Ended, -} - -impl Relay { - /// The next stream item and the state after it; `None` ends the stream. - async fn step(self) -> Option<(Item, Self)> { - let mut live = match self { - Self::Ended => return None, - Self::Connecting { endpoint, topic } => match connect(&endpoint, &topic).await { - Ok(socket) => Live { - endpoint, - socket, - event_id: 0, - }, - Err(status) => return Some((Err(status), Self::Ended)), - }, - Self::Live(live) => live, - }; - match live.next_batch().await { - Ok(batch) => Some((Ok(batch), Self::Live(live))), - Err(status) => Some((Err(status), Self::Ended)), - } - } +/// The state the publisher task and the subscribers share. +struct Shared { + history: History, + /// The engine's live blocks as the relayed stream describes them, + /// current through `cursor`. + state: LiveState, + /// The last sequence relayed in this incarnation, or none yet. + cursor: Option, + /// The first sequence relayed in this incarnation: the relay holds + /// nothing from before it, which is how a relay that started (or + /// restarted) after the publisher tells a resume from before its time. + started_at: Option, + /// Sequences of this incarnation before the record that even the + /// engine's replay no longer had: what a snapshot cannot cover. + unknown_before: u64, + /// Counts the publisher's restarts; sequences compare within one. + generation: u64, + counts: RelayCounts, + /// The normalizer's counters as of the last relayed batch: what was + /// forwarded, dropped by reason, and the engine-hash check's tally. + wire: WireCounts, + /// Why the publisher task gave up, when it did. + failed: Option, } -/// A connected subscription and the stream's event-id counter. -struct Live { - endpoint: String, - socket: SubSocket, - /// Advances once per publisher event, convertible or not, so ids stay - /// monotonic as the Python relay's do. - event_id: u64, +/// How the publisher task classified a sequence against the cursor. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +enum Admission { + Accept, + Duplicate, + Gap { from: u64, to: u64 }, + Restart { reason: RestartReason, last: u64 }, } -impl Drop for Live { - fn drop(&mut self) { - info!(endpoint = %self.endpoint, "SubscribeKvEvents: stream closed"); +/// What showed that the publisher started over. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub(crate) enum RestartReason { + /// The sequence went backwards on the same socket. + SequenceRegression, + /// The engine's startup `AllBlocksCleared` arrived under a sequence the + /// relay had already passed (SGLang's first batch after a start). + StartupClear, + /// The counter is back at 0 or 1 after a cursor above them (the mock + /// engine's publisher restart, a vLLM process restart). + CounterRestarted, +} + +impl RestartReason { + pub fn as_str(self) -> &'static str { + match self { + Self::SequenceRegression => "sequence_regression", + Self::StartupClear => "startup_clear", + Self::CounterRestarted => "counter_restarted", + } } } -impl Live { - /// The next decodable batch; what the Python relay skips (fewer than - /// three frames, undecodable payloads) is skipped here too. - async fn next_batch(&mut self) -> Item { - loop { - let message = self.socket.recv().await.map_err(|error| { - Status::internal(format!( - "SubscribeKvEvents: receive from {} failed: {error}", - self.endpoint - )) - })?; - let Some((sequence_number, payload)) = split_frames(&message) else { - continue; - }; - let batch = match rmp_serde::from_slice::>(payload) { - Ok(batch) => batch.0, - Err(error) => { - warn!(%error, sequence_number, "Failed to decode KV event batch"); - continue; - } +impl Shared { + /// Classify `seq` against the cursor; `startup_clear` says the batch + /// begins with the engine's `AllBlocksCleared`. + fn admit(&self, seq: u64, startup_clear: bool) -> Admission { + let Some(last) = self.cursor else { + return Admission::Accept; + }; + if seq == last + 1 { + return Admission::Accept; + } + if seq > last + 1 { + return Admission::Gap { + from: last + 1, + to: seq - 1, }; - return Ok(convert_batch(batch, sequence_number, &mut self.event_id)); } + let restart = |reason| Admission::Restart { reason, last }; + if seq <= 1 && last >= 2 { + return restart(RestartReason::CounterRestarted); + } + if seq < last { + return restart(RestartReason::SequenceRegression); + } + if startup_clear { + return restart(RestartReason::StartupClear); + } + Admission::Duplicate + } + + /// Start a new incarnation: the window is the old publisher's. + fn restart(&mut self) -> u64 { + self.generation += 1; + self.history.clear(); + self.state.clear(); + self.cursor = None; + self.started_at = None; + self.unknown_before = 0; + self.counts.publisher_restarts += 1; + self.generation } } -/// A SUB socket subscribed to `topic` and connected to `endpoint`. The -/// subscription is recorded first and sent on connect (and on the crate's -/// reconnects), as libzmq does; a refused publisher is retried until the -/// stream is dropped, as libzmq's background connect would. -async fn connect(endpoint: &str, topic: &str) -> Result { - let mut options = SocketOptions::default(); - options.no_connect_timeout(); - let mut socket = SubSocket::with_options(options); - let failed = |step: &str, error: ZmqError| { - Status::internal(format!("SubscribeKvEvents: {step} {endpoint}: {error}")) - }; - socket - .subscribe(topic) - .await - .map_err(|error| failed("could not subscribe to", error))?; - socket - .connect(endpoint) - .await - .map_err(|error| failed("could not connect to", error))?; - info!(%endpoint, "SubscribeKvEvents: connected to ZMQ endpoint"); - Ok(socket) +fn lock(shared: &Mutex) -> MutexGuard<'_, Shared> { + shared.lock().unwrap_or_else(PoisonError::into_inner) } -/// A publisher message's `[topic, sequence, payload, ...]` as the sequence -/// number and payload, or `None` for fewer than three frames. -fn split_frames(message: &ZmqMessage) -> Option<(u64, &[u8])> { - if message.len() < 3 { - return None; - } - let sequence = message.get(1)?; - let payload = message.get(2)?; - Some((low64_big_endian(sequence), payload.as_ref())) +/// What the live channel carries to subscribers. +#[derive(Clone)] +enum Live { + Batch(Arc), + Restart { generation: u64 }, } -/// `int.from_bytes(bytes, "big")` kept to 64 bits: the whole value for the -/// publisher's eight-byte sequence frame, the low 64 bits of a longer hash. -fn low64_big_endian(bytes: &[u8]) -> u64 { - bytes - .iter() - .fold(0, |value, &byte| (value << 8) | u64::from(byte)) -} - -/// vLLM's `KVEventBatch`, a msgspec `array_like` struct: `[ts, events, -/// data_parallel_rank]`, the rank omittable and later fields tolerated. -#[derive(Deserialize)] -struct WireBatch { - ts: f64, - events: Vec, - #[serde(default)] - data_parallel_rank: Option, -} - -/// vLLM's `KVCacheEvent` subclasses (msgspec `tag=True`, map layout): the -/// class name under `"type"`; fields without a default are always present -/// (so `medium` and `lora_name` ride along), defaulted ones may be omitted. -/// Only what the relay converts is declared; the rest is ignored. -#[derive(Deserialize)] -#[serde(tag = "type")] -enum WireEvent { - BlockStored { - block_hashes: Vec, - parent_block_hash: Option, - token_ids: Vec, - block_size: i64, - lora_id: Option, - }, - BlockRemoved { - block_hashes: Vec, - }, - AllBlocksCleared, - /// An event type this relay does not convert (one a newer vLLM added): - /// skipped on its own, like the Python relay's unknown types, so the - /// batch's other events still go through. - #[serde(other)] - Unknown, +/// One engine's KV-event relay: the publisher subscription, its history and +/// the live channel its subscribers read. +pub struct KvEventRelay { + config: RelayConfig, + shared: Arc>, + live: broadcast::Sender, + task: Mutex>>, + /// A subscriber came while the relay holds nothing: the publisher task + /// asks the engine's replay for the publisher's start again. + prime_request: Arc, + /// The servicer's load, attached to every batch sent; none until the + /// servicer installs it. + load_source: OnceLock>, } -/// A block hash as the proto's signed 64-bit identity (the Python relay's -/// `to_int64`): sha256 bytes keep their low 64 bits read big-endian; an int -/// is already masked to 64 bits by vLLM. -#[derive(Clone, Copy)] -struct BlockHash(i64); +impl Drop for KvEventRelay { + fn drop(&mut self) { + if let Some(task) = self + .task + .lock() + .unwrap_or_else(PoisonError::into_inner) + .take() + { + task.abort(); + } + let counts = self.counts(); + let shared = lock(&self.shared); + info!( + endpoint = %self.config.endpoint, + ?counts, + wire = ?shared.wire, + window = shared.history.len(), + holes = shared.history.holes(), + bytes = shared.history.bytes(), + live_blocks = shared.state.blocks(), + live_entries = shared.state.entries(), + state = ?shared.state.counts(), + "KV event relay closed" + ); + } +} -impl<'de> Deserialize<'de> for BlockHash { - fn deserialize>(deserializer: D) -> Result { - struct HashVisitor; +impl KvEventRelay { + pub fn new(config: RelayConfig) -> Arc { + let (live, _) = broadcast::channel(LIVE_CHANNEL); + Arc::new(Self { + shared: Arc::new(Mutex::new(Shared { + history: History::new(config.history_batches, config.history_bytes), + state: LiveState::new(), + cursor: None, + started_at: None, + wire: WireCounts::default(), + unknown_before: 0, + generation: 0, + counts: RelayCounts::default(), + failed: None, + })), + config, + live, + task: Mutex::new(None), + prime_request: Arc::new(Notify::new()), + load_source: OnceLock::new(), + }) + } - impl Visitor<'_> for HashVisitor { - type Value = BlockHash; + /// Install where the load record on every sent batch comes from (once; + /// a second call is ignored). + pub(crate) fn set_load_source(&self, source: Arc) { + let _ = self.load_source.set(source); + } - fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { - formatter.write_str("a block hash as bytes or an integer") - } + /// The relay for an engine's publisher, or `None` when events are off + /// (an empty endpoint). + pub(crate) fn for_publisher( + kv_events_endpoint: &str, + replay_endpoint: Option<&str>, + topic: &str, + ) -> Option> { + (!kv_events_endpoint.is_empty()).then(|| { + Self::new(RelayConfig::for_publisher( + kv_events_endpoint, + replay_endpoint, + topic, + )) + }) + } - fn visit_bytes(self, bytes: &[u8]) -> Result { - Ok(BlockHash(low64_big_endian(bytes) as i64)) - } + pub fn counts(&self) -> RelayCounts { + lock(&self.shared).counts.clone() + } - fn visit_u64(self, value: u64) -> Result { - Ok(BlockHash(value as i64)) - } + /// The normalizer's counters as of the last relayed batch (forwarded, + /// dropped by reason, the engine-hash check's tally); what the summary + /// and closing log lines print. + #[cfg(test)] + pub(crate) fn wire_counts(&self) -> WireCounts { + lock(&self.shared).wire.clone() + } - fn visit_i64(self, value: i64) -> Result { - Ok(BlockHash(value)) + /// Record `batches` as relayed without a publisher: the history, the + /// live state and the cursor as [`Relaying::relay`] would leave them. + #[cfg(test)] + fn preload(&self, batches: impl IntoIterator) { + let mut shared = lock(&self.shared); + for batch in batches { + let sequence = batch.sequence_number; + let batch = Arc::new(batch); + shared.history.push(sequence, Arc::clone(&batch)); + shared.state.apply(&batch); + if shared.started_at.is_none() { + shared.started_at = Some(sequence); } + shared.cursor = Some(sequence); + shared.counts.relayed += 1; } - - deserializer.deserialize_any(HashVisitor) } -} -/// The proto batch under the publisher's sequence number, `event_id` -/// advancing once per event whether or not it converts. -fn convert_batch( - batch: WireBatch, - sequence_number: u64, - event_id: &mut u64, -) -> common::KvEventBatch { - let mut events = Vec::with_capacity(batch.events.len()); - for event in batch.events { - *event_id += 1; - if let Some(converted) = convert_event(event, *event_id) { - events.push(converted); + /// [`Self::start`] when the servicer begins serving, unless + /// [`RELAY_START_ENV`] is `lazy`; needs a Tokio runtime. + pub(crate) fn start_at_boot(&self) { + if starts_at_boot(std::env::var(RELAY_START_ENV).ok().as_deref()) { + self.start(); + } else { + info!( + endpoint = %self.config.endpoint, + "{RELAY_START_ENV}=lazy: the KV event relay subscribes at the first SubscribeKvEvents" + ); } } - common::KvEventBatch { - sequence_number, - timestamp: batch.ts, - events, - dp_rank: batch.data_parallel_rank, + + /// Subscribe to the publisher if not yet subscribed. Called at the + /// servicer's start ([`Self::start_at_boot`]) and by + /// [`Self::subscribe`]; needs a Tokio runtime. + pub fn start(&self) { + let mut task = self.task.lock().unwrap_or_else(PoisonError::into_inner); + if task.as_ref().is_some_and(|task| !task.is_finished()) { + return; + } + let config = self.config.clone(); + let shared = Arc::clone(&self.shared); + let live = self.live.clone(); + let prime_request = Arc::clone(&self.prime_request); + #[expect( + clippy::disallowed_methods, + reason = "the publisher subscription outlives any one RPC; aborted when the relay drops" + )] + let handle = tokio::spawn(run(config, shared, live, prime_request)); + *task = Some(handle); } -} -/// One event as its proto, or `None` for a store whose hashes and tokens do -/// not form whole blocks (ordinal slicing needs dense, complete blocks). -fn convert_event(event: WireEvent, event_id: u64) -> Option { - let data = match event { - WireEvent::BlockStored { - block_hashes, - parent_block_hash, - token_ids, - block_size, - lora_id, - } => { - let width = usize::try_from(block_size).ok().filter(|&width| { - width > 0 - && i32::try_from(width).is_ok() - && block_hashes.len().checked_mul(width) == Some(token_ids.len()) - }); - let Some(width) = width else { - warn!( - hashes = block_hashes.len(), - block_size, - tokens = token_ids.len(), - "Skipping BlockStored: the hashes of this block size cannot map to the tokens" + /// Handle one `SubscribeKvEvents` call: the history after the cursor, or + /// a state snapshot, then live events; see the module docs for the + /// cursor rules. + pub fn subscribe( + &self, + request: common::SubscribeKvEventsRequest, + ) -> Result, Status> { + self.start(); + // Subscribed before the history is read, so nothing published between + // the two is missed; the subscriber skips what it already replayed. + let rx = self.live.subscribe(); + let cursor = request.start_sequence_number; + let serve = { + let mut shared = lock(&self.shared); + if let Some(error) = &shared.failed { + return Err(Status::internal(format!( + "SubscribeKvEvents: the relay for {} is down: {error}", + self.config.endpoint + ))); + } + if cursor == 0 { + if shared.history.complete_from_start() { + shared.counts.served_from_history += 1; + Serve::History { + batches: shared.history.all(), + last_sent: None, + } + } else if let Some(through) = shared.cursor { + // The window no longer reaches back to the publisher's + // first batch: the live set as of `through`, taken under + // the guard that admits batches, so every later batch + // reaches this subscriber through the live channel. + let started = Instant::now(); + let snapshot = shared.state.snapshot(); + let collected = started.elapsed(); + shared.counts.served_snapshots += 1; + Serve::Snapshot { + snapshot, + through, + collected, + oldest: shared.history.oldest(), + holes: shared.history.holes(), + unknown_before: shared.unknown_before, + } + } else { + // Nothing relayed yet: the publisher may have spoken + // before the subscription reached it and after the + // start replay answered; ask its replay once more. + if self.config.replay_endpoint.is_some() { + self.prime_request.notify_one(); + } + Serve::Live + } + } else { + match shared.history.after(cursor) { + Ok(batches) => { + shared.counts.served_from_history += 1; + Serve::History { + batches, + last_sent: Some(cursor), + } + } + Err(window) => { + shared.counts.out_of_range += 1; + let endpoint = &self.config.endpoint; + let message = match window { + Window::Empty => format!( + "SubscribeKvEvents: the relay for {endpoint} holds no history \ + yet (it started after the publisher); resubscribe from zero" + ), + Window::Behind { oldest } => match shared.started_at { + Some(started) if cursor + 1 < started => format!( + "SubscribeKvEvents: the relay for {endpoint} started at \ + sequence {started} and holds nothing before it; resubscribe \ + from zero" + ), + _ => format!( + "SubscribeKvEvents: the relay for {endpoint} keeps history \ + from sequence {oldest}, after cursor {cursor}; resubscribe \ + from zero for a state snapshot" + ), + }, + Window::Ahead { newest } => format!( + "SubscribeKvEvents: cursor {cursor} is beyond the publisher's \ + last sequence {newest} at {endpoint}: the publisher restarted; \ + resubscribe from zero" + ), + }; + info!( + counts = ?shared.counts, + window = shared.history.len(), + holes = shared.history.holes(), + "{message}" + ); + return Err(Status::out_of_range(message)); + } + } + } + }; + let (replay, last_sent, snapshot) = match serve { + Serve::History { batches, last_sent } => { + if !batches.is_empty() { + debug!( + endpoint = %self.config.endpoint, + cursor, + batches = batches.len(), + "SubscribeKvEvents: serving from history" + ); + } + (batches, last_sent, SnapshotPhase::None) + } + Serve::Snapshot { + snapshot, + through, + collected, + oldest, + holes, + unknown_before, + } => { + info!( + endpoint = %self.config.endpoint, + through, + blocks = snapshot.blocks, + entries = snapshot.entries(), + ranks = snapshot.ranks.len(), + collected_us = u64::try_from(collected.as_micros()).unwrap_or(u64::MAX), + oldest, + holes, + unknown_before, + "SubscribeKvEvents: the history no longer starts at the publisher's first \ + batch; serving a state snapshot" ); - return None; - }; - let blocks = block_hashes - .iter() - .zip(token_ids.chunks_exact(width)) - .map(|(hash, tokens)| common::KvBlock { - block_hash: hash.0, - token_ids: tokens.to_vec(), - block_size: i32::try_from(width).unwrap_or(i32::MAX), - lora_id, - cache_level: None, - }) - .collect(); - kv_cache_event::Data::Stored(common::KvBlocksStored { - blocks, - parent_block_hash: parent_block_hash.map(|hash| hash.0), - }) - } - WireEvent::BlockRemoved { block_hashes } => { - kv_cache_event::Data::Removed(common::KvBlocksRemoved { - block_hashes: block_hashes.into_iter().map(|hash| hash.0).collect(), - cache_level: None, - }) - } - WireEvent::AllBlocksCleared => kv_cache_event::Data::Cleared(common::KvCacheCleared {}), - WireEvent::Unknown => { - debug!( - event_id, - "Skipping a KV event of a type this relay does not convert" - ); - return None; - } - }; - Some(common::KvCacheEvent { - event_id, - data: Some(data), - }) + let timestamp = unix_seconds(); + // Ordering and sizing the chunks is a pass over every live + // block: off the runtime's workers, and off the lock. + let ordering = tokio::task::spawn_blocking(move || { + SnapshotChunks::new(snapshot, through, timestamp, unknown_before) + }); + (Vec::new(), Some(through), SnapshotPhase::Ordering(ordering)) + } + Serve::Live => (Vec::new(), None, SnapshotPhase::None), + }; + let subscriber = Subscriber { + snapshot, + replay: replay.into(), + rx, + last_sent, + shared: Arc::clone(&self.shared), + endpoint: self.config.endpoint.clone(), + done: false, + source: self.load_source.get().cloned(), + load_tick: self.config.load_tick, + heartbeat_interval: self.config.heartbeat_interval, + heartbeat_backoff: self.config.heartbeat_backoff, + last_rank: None, + last_record: None, + last_sent_at: Instant::now(), + idle_heartbeats: 0, + sample: 0, + }; + Ok(Box::pin(stream::unfold( + subscriber, + |mut subscriber| async move { subscriber.next().await.map(|item| (item, subscriber)) }, + ))) + } } -/// Golden publisher payloads encoded by vLLM 0.30.1rc1 (msgspec 0.22) with -/// `crates/engine_servicer/scripts/generate_kv_events_golden.py`, which also -/// prints the Python relay's conversion of them (the expected protos below). -#[cfg(test)] -pub(crate) mod golden { - use zeromq::ZmqMessage; +/// What a subscription begins with, decided under the relay's lock. +enum Serve { + History { + batches: Vec>, + last_sent: Option, + }, + Snapshot { + snapshot: Snapshot, + through: u64, + collected: Duration, + oldest: Option, + holes: usize, + unknown_before: u64, + }, + Live, +} - /// `KVEventBatch(ts=1700000000.5, data_parallel_rank=None)` with a - /// `BlockStored` of two sha256-byte hashes (`00..00 80 00..00`, - /// `00..00 ff..fe`), parent 7, tokens 1..=8, block size 4, medium GPU; a - /// `BlockStored` of int hash 0x1234, no parent, tokens [9, 10], block - /// size 2, lora_id 3, group_idx 0, kv_cache_spec_kind full_attention; an - /// unaligned `BlockStored` (one hash, block size 4, three tokens); a - /// `BlockRemoved` of [0x1234, ff..ff]; an `AllBlocksCleared`. - pub(crate) const BATCH1: &str = "93cb41d954fc402000009588a474797065ab426c6f636b53746f726564ac626c6f636b5f68617368657392c4200000000000000000000000000000000000000000000000008000000000000000c420000000000000000000000000000000000000000000000000fffffffffffffffeb1706172656e745f626c6f636b5f6861736807a9746f6b656e5f696473980102030405060708aa626c6f636b5f73697a6504a76c6f72615f6964c0a66d656469756da3475055a96c6f72615f6e616d65c08aa474797065ab426c6f636b53746f726564ac626c6f636b5f68617368657391cd1234b1706172656e745f626c6f636b5f68617368c0a9746f6b656e5f69647392090aaa626c6f636b5f73697a6502a76c6f72615f696403a66d656469756dc0a96c6f72615f6e616d65c0a967726f75705f69647800b26b765f63616368655f737065635f6b696e64ae66756c6c5f617474656e74696f6e88a474797065ab426c6f636b53746f726564ac626c6f636b5f6861736865739105b1706172656e745f626c6f636b5f68617368c0a9746f6b656e5f69647393010203aa626c6f636b5f73697a6504a76c6f72615f6964c0a66d656469756da3475055a96c6f72615f6e616d65c083a474797065ac426c6f636b52656d6f766564ac626c6f636b5f68617368657392cd1234c420ffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffa66d656469756da347505581a474797065b0416c6c426c6f636b73436c6561726564c0"; +/// Where a subscription's snapshot stands. +enum SnapshotPhase { + None, + /// Being ordered and sized on the blocking pool; awaited at the first poll. + Ordering(tokio::task::JoinHandle), + Serving(SnapshotChunks), +} - /// `KVEventBatch(ts=1700000001.0, data_parallel_rank=1)` with one - /// `BlockStored`: int hash 42, parent 41, tokens [100, 101], block size 2. - pub(crate) const BATCH2: &str = "93cb41d954fc404000009188a474797065ab426c6f636b53746f726564ac626c6f636b5f686173686573912ab1706172656e745f626c6f636b5f6861736829a9746f6b656e5f696473926465aa626c6f636b5f73697a6502a76c6f72615f6964c0a66d656469756da3475055a96c6f72615f6e616d65c001"; +/// Seconds since the Unix epoch, the way the engines stamp their batches. +fn unix_seconds() -> f64 { + SystemTime::now() + .duration_since(UNIX_EPOCH) + .map(|elapsed| elapsed.as_secs_f64()) + .unwrap_or(0.0) +} - pub(crate) fn bytes(hex: &str) -> Vec { - (0..hex.len()) - .step_by(2) - .map(|i| u8::from_str_radix(&hex[i..i + 2], 16).expect("hex")) - .collect() - } +type Item = Result; - /// A publisher message: `[topic, sequence (u64 big-endian), payload]`. - pub(crate) fn frame(topic: &[u8], sequence: u64, payload: &[u8]) -> ZmqMessage { - let mut message = ZmqMessage::from(topic.to_vec()); - message.push_back(sequence.to_be_bytes().to_vec().into()); - message.push_back(payload.to_vec().into()); - message - } +/// One `SubscribeKvEvents` stream: the snapshot it owes, what it still owes +/// from the history, then the live channel, deduplicated by sequence. +struct Subscriber { + snapshot: SnapshotPhase, + replay: VecDeque>, + rx: broadcast::Receiver, + last_sent: Option, + shared: Arc>, + endpoint: String, + done: bool, + /// The servicer's load for the record on every batch; `None` sends + /// batches without one and no `load_only` batches. + source: Option>, + load_tick: Duration, + heartbeat_interval: Duration, + heartbeat_backoff: Duration, + /// The rank of the last batch sent: what a `load_only` batch names. + last_rank: Option, + /// The last record sent, for the change check. + last_record: Option, + /// When the last batch of any kind went out. + last_sent_at: Instant, + /// Heartbeats in a row with an unchanged record. + idle_heartbeats: u32, + /// Records sent on this stream. + sample: u64, } -#[cfg(test)] -mod tests { - use std::time::Duration; +/// What the live wait ended with. +enum Waited { + Live(Result), + Tick, +} - use futures::StreamExt; - use tokio::time::timeout; - use zeromq::{prelude::*, PubSocket, SocketEvent}; +impl Subscriber { + /// Attach the load record to a batch about to go out and note it went. + fn stamped(&mut self, mut batch: common::KvEventBatch) -> common::KvEventBatch { + self.last_rank = batch.dp_rank; + self.last_sent_at = Instant::now(); + self.idle_heartbeats = 0; + batch.load = self.record(batch.dp_rank, false); + batch + } - use super::{golden, *}; + /// The next record for `rank` from the source, numbered. + fn record(&mut self, rank: Option, load_only: bool) -> Option { + let mut record = self.source.as_ref()?.load(rank)?; + self.sample += 1; + record.sample = self.sample; + record.load_only = load_only; + // The telemetry rides the heartbeats and the stream's first record; + // an event batch, one per scheduler step, carries the core only. + if !load_only && self.sample > 1 { + smg_grpc_client::engine_load::core_only(&mut record); + } + self.last_record = Some(record.clone()); + Some(record) + } - fn decode(hex: &str) -> WireBatch { - rmp_serde::from_slice::>(&golden::bytes(hex)) - .expect("golden batch decodes") - .0 + /// A `load_only` batch when the record moved since the last one sent, + /// or when the stream has been silent for the heartbeat interval (the + /// backoff interval after two unchanged heartbeats); `None` otherwise. + fn load_only_batch(&mut self) -> Option { + // Nothing sent yet, nothing to repeat: a heartbeat numbered 0 before + // the publisher's first batch would make a gateway that predates + // the field take that batch for a duplicate. + let sequence_number = self.last_sent?; + let current = self.source.as_ref()?.load(self.last_rank)?; + let changed = self + .last_record + .as_ref() + .is_none_or(|last| load_changed(last, ¤t)); + let due_after = if self.idle_heartbeats >= 2 { + self.heartbeat_backoff + } else { + self.heartbeat_interval + }; + if !changed && self.last_sent_at.elapsed() < due_after { + return None; + } + let rank = self.last_rank; + let record = self.record(rank, true)?; + self.last_sent_at = Instant::now(); + self.idle_heartbeats = if changed { + 0 + } else { + self.idle_heartbeats.saturating_add(1) + }; + Some(common::KvEventBatch { + sequence_number, + timestamp: unix_seconds(), + events: Vec::new(), + dp_rank: rank, + snapshot: None, + load: Some(record), + }) } - fn stored(event: &common::KvCacheEvent) -> &common::KvBlocksStored { - match &event.data { - Some(kv_cache_event::Data::Stored(stored)) => stored, - other => panic!("expected a stored event, got {other:?}"), + async fn next(&mut self) -> Option { + if self.done { + return None; + } + loop { + match &mut self.snapshot { + SnapshotPhase::Ordering(ordering) => { + let chunks = match ordering.await { + Ok(chunks) => chunks, + Err(error) => { + self.done = true; + return Some(Err(Status::internal(format!( + "SubscribeKvEvents: the snapshot of {} could not be ordered: \ + {error}", + self.endpoint + )))); + } + }; + debug!( + endpoint = %self.endpoint, + chunks = chunks.chunk_count(), + blocks = chunks.blocks(), + through = chunks.through(), + "SubscribeKvEvents: snapshot ordered" + ); + self.snapshot = SnapshotPhase::Serving(chunks); + continue; + } + SnapshotPhase::Serving(chunks) => { + if let Some(chunk) = chunks.next_chunk() { + return Some(Ok(self.stamped(chunk))); + } + self.snapshot = SnapshotPhase::None; + } + SnapshotPhase::None => {} + } + if let Some(batch) = self.replay.pop_front() { + self.last_sent = Some(batch.sequence_number); + return Some(Ok(self.stamped((*batch).clone()))); + } + let waited = if self.source.is_some() { + let tick = self.load_tick; + tokio::select! { + received = self.rx.recv() => Waited::Live(received), + () = tokio::time::sleep(tick) => Waited::Tick, + } + } else { + Waited::Live(self.rx.recv().await) + }; + let received = match waited { + Waited::Tick => { + if let Some(batch) = self.load_only_batch() { + return Some(Ok(batch)); + } + continue; + } + Waited::Live(received) => received, + }; + match received { + Ok(Live::Batch(batch)) => { + if self + .last_sent + .is_some_and(|last| batch.sequence_number <= last) + { + continue; + } + self.last_sent = Some(batch.sequence_number); + return Some(Ok(self.stamped((*batch).clone()))); + } + Ok(Live::Restart { generation }) => { + self.done = true; + return Some(Err(Status::data_loss(format!( + "SubscribeKvEvents: publisher at {} restarted (incarnation \ + {generation}); resubscribe from zero", + self.endpoint + )))); + } + Err(broadcast::error::RecvError::Lagged(skipped)) => { + let refill = { + let mut shared = lock(&self.shared); + shared.counts.subscribers_lagged += 1; + match self.last_sent { + Some(last) => shared.history.after(last), + None => Err(Window::Empty), + } + }; + match refill { + Ok(batches) => { + debug!( + endpoint = %self.endpoint, + skipped, + refilled = batches.len(), + "SubscribeKvEvents: subscriber fell behind; refilled from history" + ); + self.replay = batches.into(); + } + Err(_) => { + self.done = true; + return Some(Err(Status::data_loss(format!( + "SubscribeKvEvents: subscriber of {} fell {skipped} batches \ + behind and the history no longer reaches back; resubscribe \ + from zero", + self.endpoint + )))); + } + } + } + Err(broadcast::error::RecvError::Closed) => return None, + } } } +} - fn block(block_hash: i64, token_ids: Vec, lora_id: Option) -> common::KvBlock { - common::KvBlock { - block_hash, - block_size: i32::try_from(token_ids.len()).unwrap(), - token_ids, - lora_id, - cache_level: None, +/// A publisher payload as the relay read it. +enum Decoded { + Batch(WireBatch), + Undecodable, +} + +fn decode(payload: &[u8], sequence: u64) -> Decoded { + match rmp_serde::from_slice::>(payload) { + Ok(batch) => Decoded::Batch(batch.0), + Err(error) => { + warn!(%error, sequence, "Failed to decode KV event batch; relaying it empty"); + Decoded::Undecodable } } +} - #[test] - fn endpoint_for_rank_mirrors_the_python_helper() { - assert_eq!(endpoint_for_rank("tcp://*:5557", 0), "tcp://127.0.0.1:5557"); - assert_eq!( - endpoint_for_rank("tcp://0.0.0.0:5557", 0), - "tcp://127.0.0.1:5557" +/// The publisher task: one SUB socket for the relay's lifetime. +async fn run( + config: RelayConfig, + shared: Arc>, + live: broadcast::Sender, + prime_request: Arc, +) { + let mut socket = match connect(&config.endpoint, &config.topic).await { + Ok(socket) => socket, + Err(status) => { + lock(&shared).failed = Some(status.message().to_string()); + return; + } + }; + let mut relay = Relaying { + config: &config, + shared: &shared, + live: &live, + normalizer: Normalizer::from_env(), + event_id: 0, + }; + if let Some(check) = relay.normalizer.hash_check() { + info!( + endpoint = %config.endpoint, + check = check.as_str(), + "KV event engine-hash check on; its tally is in the relay's summary and closing lines" ); - assert_eq!(endpoint_for_rank("tcp://*:5557", 2), "tcp://127.0.0.1:5559"); - assert_eq!( - endpoint_for_rank("tcp://10.0.0.1:5557", 1), - "tcp://10.0.0.1:5558" + } + // The publisher may have been counting before the subscription reached + // it: take what its replay still holds before the first live batch, and + // keep asking until the engine answers (it may still be starting). A + // live batch arriving first settles the start on its own path. + let mut primed = config.replay_endpoint.is_none(); + let mut next_prime = tokio::time::Instant::now(); + loop { + let received = tokio::select! { + received = socket.recv() => received, + () = tokio::time::sleep_until(next_prime), if !primed => { + primed = relay.prime_from_replay().await; + next_prime = tokio::time::Instant::now() + PRIME_RETRY; + continue; + } + () = prime_request.notified(), if config.replay_endpoint.is_some() => { + if lock(&shared).cursor.is_none() { + primed = false; + next_prime = tokio::time::Instant::now(); + } + continue; + } + }; + primed = true; + let message = match received { + Ok(message) => message, + Err(error) => { + warn!(endpoint = %config.endpoint, %error, "KV event receive failed; retrying"); + tokio::time::sleep(Duration::from_millis(100)).await; + continue; + } + }; + let Some((sequence, payload)) = split_frames(&message) else { + continue; + }; + let decoded = decode(payload, sequence); + let startup_clear = matches!( + &decoded, + Decoded::Batch(batch) + if matches!(batch.events.first(), Some(WireEvent::AllBlocksCleared { .. })) ); - assert_eq!(endpoint_for_rank("tcp://host:port", 1), "tcp://host:port"); - assert_eq!(endpoint_for_rank("ipc:///tmp/kv", 1), "ipc:///tmp/kv"); + let (first, mut admission) = { + let shared = lock(&shared); + ( + shared.cursor.is_none(), + shared.admit(sequence, startup_clear), + ) + }; + if first && admission == Admission::Accept && joined_late(sequence, startup_clear) { + relay.recover_start(sequence).await; + // The replay may have given this sequence too. + admission = lock(&shared).admit(sequence, startup_clear); + } + if let Admission::Gap { from, to } = admission { + relay.recover_gap(from..=to).await; + // The replay may have run past the live sequence. + admission = lock(&shared).admit(sequence, startup_clear); + } + match admission { + Admission::Accept => relay.relay(sequence, decoded), + Admission::Gap { .. } | Admission::Duplicate => {} + Admission::Restart { reason, last } => { + let (generation, counts) = { + let mut shared = lock(&shared); + let generation = shared.restart(); + (generation, shared.counts.clone()) + }; + warn!( + endpoint = %config.endpoint, + sequence, + last, + generation, + reason = reason.as_str(), + ?counts, + "KV event publisher restarted; live subscribers end with DATA_LOSS" + ); + let _ = live.send(Live::Restart { generation }); + relay.normalizer = Normalizer::from_env(); + // The SUB re-joins a restarted publisher a moment late too. + if joined_late(sequence, startup_clear) { + relay.recover_start(sequence).await; + if lock(&shared).admit(sequence, startup_clear) != Admission::Accept { + continue; + } + } + relay.relay(sequence, decoded); + } + } } +} - /// The golden batches convert to exactly what the Python relay produced +/// Whether the first batch of an incarnation shows the relay joined after the +/// publisher's start: the publishers count from 0 (vLLM, SGLang) or 1 (the +/// mock engine), and a batch carrying the engine's startup clear needs +/// nothing before it. +fn joined_late(sequence: u64, startup_clear: bool) -> bool { + sequence > 1 && !startup_clear +} + +/// What a replay request gave back for a range of sequences. +struct Filled { + recovered: u64, + lost: u64, + /// Holes before the first recovered batch (the whole range when nothing + /// came back). + unknown_before: u64, +} + +/// The publisher task's per-batch work. +struct Relaying<'a> { + config: &'a RelayConfig, + shared: &'a Mutex, + live: &'a broadcast::Sender, + normalizer: Normalizer, + event_id: u64, +} + +impl Relaying<'_> { + /// Normalize, remember and broadcast one batch under `sequence`. + fn relay(&mut self, sequence: u64, decoded: Decoded) { + let (batch, undecodable) = match decoded { + Decoded::Batch(batch) => ( + self.normalizer + .normalize_batch(batch, sequence, &mut self.event_id), + false, + ), + Decoded::Undecodable => ( + common::KvEventBatch { + sequence_number: sequence, + ..common::KvEventBatch::default() + }, + true, + ), + }; + let batch = Arc::new(batch); + { + let mut shared = lock(self.shared); + shared.history.push(sequence, Arc::clone(&batch)); + shared.state.apply(&batch); + if shared.started_at.is_none() { + shared.started_at = Some(sequence); + } + shared.cursor = Some(sequence); + shared.counts.relayed += 1; + if undecodable { + shared.counts.undecodable_batches += 1; + } + shared.wire = self.normalizer.counts().clone(); + if shared.counts.relayed.is_multiple_of(SUMMARY_EVERY_BATCHES) { + info!( + endpoint = %self.config.endpoint, + counts = ?shared.counts, + wire = ?shared.wire, + live_blocks = shared.state.blocks(), + live_entries = shared.state.entries(), + "KV event relay summary" + ); + } + } + let _ = self.live.send(Live::Batch(batch)); + } + + /// The publisher skipped `gap` on the socket: fill it from the engine's + /// replay; what it does not give becomes holes. + async fn recover_gap(&mut self, gap: RangeInclusive) { + let (from, to) = (*gap.start(), *gap.end()); + lock(self.shared).counts.publisher_gaps += 1; + let filled = self.fill(gap).await; + let counts = lock(self.shared).counts.clone(); + warn!( + endpoint = %self.config.endpoint, + from, + to, + recovered = filled.recovered, + lost = filled.lost, + ?counts, + "KV event publisher skipped sequences" + ); + } + + /// The publisher was already past `first_seen` when the relay's + /// subscription reached it (ZMQ delivers nothing from before a join): + /// ask its replay for everything before, and record what even the replay + /// no longer had as unknown before the record's start. + async fn recover_start(&mut self, first_seen: u64) { + let filled = self.fill(0..=first_seen - 1).await; + let counts = { + let mut shared = lock(self.shared); + shared.unknown_before = filled.unknown_before; + shared.counts.unknown_before_start += filled.unknown_before; + shared.counts.clone() + }; + if filled.unknown_before == 0 { + info!( + endpoint = %self.config.endpoint, + first_seen, + recovered = filled.recovered, + ?counts, + "KV event relay joined a publisher already counting; its replay covered the \ + batches before" + ); + } else { + warn!( + endpoint = %self.config.endpoint, + first_seen, + recovered = filled.recovered, + unknown_before = filled.unknown_before, + ?counts, + "KV event relay joined a publisher already counting and its replay does not \ + reach back to the start; the history and the snapshots start late" + ); + } + } + + /// Before any live batch of an incarnation: everything the engine's + /// replay still holds, so a publisher that counted before the + /// subscription reached it is known even if it never publishes again. + /// `true` once the replay answered (batches, or an empty buffer: the + /// engine is up and has nothing yet), `false` when it could not be + /// reached or did not finish, to be asked again. + async fn prime_from_replay(&mut self) -> bool { + let Some(endpoint) = &self.config.replay_endpoint else { + return true; + }; + let timeout = self.config.replay_timeout; + let (replies, complete) = match tokio::time::timeout( + timeout * 2, + replay(endpoint, 0, timeout), + ) + .await + { + Ok(Ok(answer)) => answer, + Ok(Err(error)) => { + debug!(endpoint = %self.config.endpoint, %error, "KV event replay not answering yet"); + return false; + } + Err(_) => return false, + }; + if lock(self.shared).cursor.is_some() { + // A live batch got in first and settled the start. + return true; + } + if replies.is_empty() { + if complete { + info!( + endpoint = %self.config.endpoint, + "KV event relay asked the publisher's replay at start; it holds nothing yet" + ); + } + return complete; + } + let first = replies[0].0; + // The publishers count from 0 (vLLM, SGLang) or 1 (the mock engine). + let unknown_before = if first <= 1 { 0 } else { first }; + if unknown_before > 0 { + self.lose(0, first - 1); + } + let mut expected = first; + let mut taken = 0u64; + let mut last = first; + for (sequence, payload) in replies { + if sequence < expected { + continue; + } + if expected < sequence { + self.lose(expected, sequence - 1); + } + self.relay(sequence, decode(&payload, sequence)); + taken += 1; + last = sequence; + expected = sequence + 1; + } + let counts = { + let mut shared = lock(self.shared); + shared.unknown_before = unknown_before; + shared.counts.unknown_before_start += unknown_before; + shared.counts.primed_batches += taken; + shared.counts.clone() + }; + if unknown_before == 0 { + info!( + endpoint = %self.config.endpoint, + first, + last, + batches = taken, + ?counts, + "KV event relay primed from the publisher's replay at start" + ); + } else { + warn!( + endpoint = %self.config.endpoint, + first, + last, + batches = taken, + unknown_before, + ?counts, + "KV event relay primed from the publisher's replay at start; the replay no \ + longer reaches back to the publisher's first batch" + ); + } + true + } + + /// Fill `gap` from the engine's replay socket: what it gives back is + /// relayed under its sequence (past the gap's end too, when the replay + /// ran ahead of the live socket), what it does not becomes holes. + async fn fill(&mut self, gap: RangeInclusive) -> Filled { + let (from, to) = (*gap.start(), *gap.end()); + let replies = match &self.config.replay_endpoint { + Some(endpoint) => match replay(endpoint, from, self.config.replay_timeout).await { + Ok((replies, _)) => replies, + Err(error) => { + warn!( + endpoint = %self.config.endpoint, + %error, + "KV event replay failed; the gap stays" + ); + Vec::new() + } + }, + None => Vec::new(), + }; + let mut expected = from; + let mut filled = Filled { + recovered: 0, + lost: 0, + unknown_before: 0, + }; + let mut first_known = None; + for (sequence, payload) in replies { + if sequence < expected { + continue; + } + if expected < sequence { + filled.lost += sequence - expected; + self.lose(expected, sequence - 1); + } + if first_known.is_none() { + first_known = Some(sequence); + } + self.relay(sequence, decode(&payload, sequence)); + filled.recovered += 1; + expected = sequence + 1; + } + if expected <= to { + filled.lost += to + 1 - expected; + self.lose(expected, to); + } + filled.unknown_before = first_known.map_or(to + 1 - from, |first| first - from); + lock(self.shared).counts.gap_batches_recovered += filled.recovered; + filled + } + + /// `from..=to` passed without batches: holes in the window. + fn lose(&mut self, from: u64, to: u64) { + let mut shared = lock(self.shared); + shared.history.push_lost_range(from, to); + shared.cursor = Some(to); + shared.counts.gap_batches_lost += to + 1 - from; + } +} + +/// Ask a publisher's replay ROUTER for its buffered batches from `from`: +/// `[b"", from as 8 bytes big-endian]` on a DEALER; replies are +/// `[b"", topic, seq, payload]` (vLLM) or `[b"", seq, payload]` (SGLang), +/// ending with the all-ones sequence. Returned in arrival order with whether +/// the end marker arrived; a stop before it returns what arrived. +async fn replay( + endpoint: &str, + from: u64, + timeout: Duration, +) -> Result<(Vec<(u64, Vec)>, bool), String> { + let mut dealer = DealerSocket::new(); + dealer + .connect(endpoint) + .await + .map_err(|error| format!("could not connect to replay {endpoint}: {error}"))?; + let mut request = ZmqMessage::from(Vec::new()); + request.push_back(from.to_be_bytes().to_vec().into()); + dealer + .send(request) + .await + .map_err(|error| format!("could not send the replay request to {endpoint}: {error}"))?; + let mut replies = Vec::new(); + loop { + let reply = match tokio::time::timeout(timeout, dealer.recv()).await { + Ok(Ok(reply)) => reply, + Ok(Err(error)) => return Err(format!("replay from {endpoint} failed: {error}")), + Err(_) => { + warn!( + endpoint, + from, + received = replies.len(), + "KV event replay timed out" + ); + return Ok((replies, false)); + } + }; + let (sequence, payload) = match reply.len() { + 3 => (reply.get(1), reply.get(2)), + 4 => (reply.get(2), reply.get(3)), + frames => { + return Err(format!( + "malformed replay reply from {endpoint}: {frames} frames" + )) + } + }; + let (Some(sequence), Some(payload)) = (sequence, payload) else { + return Err(format!("malformed replay reply from {endpoint}")); + }; + if sequence.as_ref() == &END_SEQUENCE[..] { + return Ok((replies, true)); + } + if sequence.len() != 8 { + return Err(format!("malformed replay sequence from {endpoint}")); + } + replies.push((low64_big_endian(sequence.as_ref()), payload.to_vec())); + } +} + +/// A SUB socket subscribed to `topic` and connected to `endpoint`. The +/// subscription is recorded first and sent on connect (and on the crate's +/// reconnects), as libzmq does; a refused publisher is retried until the +/// socket is dropped, as libzmq's background connect would. +async fn connect(endpoint: &str, topic: &str) -> Result { + let mut options = SocketOptions::default(); + options.no_connect_timeout(); + let mut socket = SubSocket::with_options(options); + let failed = |step: &str, error: ZmqError| { + Status::internal(format!("SubscribeKvEvents: {step} {endpoint}: {error}")) + }; + socket + .subscribe(topic) + .await + .map_err(|error| failed("could not subscribe to", error))?; + socket + .connect(endpoint) + .await + .map_err(|error| failed("could not connect to", error))?; + info!(%endpoint, "SubscribeKvEvents: connected to ZMQ endpoint"); + Ok(socket) +} + +/// A publisher message's `[topic, sequence, payload, ...]` as the sequence +/// number and payload, or `None` for fewer than three frames. +fn split_frames(message: &ZmqMessage) -> Option<(u64, &[u8])> { + if message.len() < 3 { + return None; + } + let sequence = message.get(1)?; + let payload = message.get(2)?; + Some((low64_big_endian(sequence), payload.as_ref())) +} + +/// Golden publisher payloads encoded by vLLM 0.30.1rc1 (msgspec 0.22) with +/// `crates/engine_servicer/scripts/generate_kv_events_golden.py`, which also +/// prints the Python relay's conversion of them (the expected protos below). +#[cfg(test)] +pub(crate) mod golden { + use zeromq::ZmqMessage; + + /// `KVEventBatch(ts=1700000000.5, data_parallel_rank=None)` with a + /// `BlockStored` of two sha256-byte hashes (`00..00 80 00..00`, + /// `00..00 ff..fe`), parent 7, tokens 1..=8, block size 4, medium GPU; a + /// `BlockStored` of int hash 0x1234, no parent, tokens [9, 10], block + /// size 2, lora_id 3, group_idx 0, kv_cache_spec_kind full_attention; an + /// unaligned `BlockStored` (one hash, block size 4, three tokens); a + /// `BlockRemoved` of [0x1234, ff..ff]; an `AllBlocksCleared`. + pub(crate) const BATCH1: &str = "93cb41d954fc402000009588a474797065ab426c6f636b53746f726564ac626c6f636b5f68617368657392c4200000000000000000000000000000000000000000000000008000000000000000c420000000000000000000000000000000000000000000000000fffffffffffffffeb1706172656e745f626c6f636b5f6861736807a9746f6b656e5f696473980102030405060708aa626c6f636b5f73697a6504a76c6f72615f6964c0a66d656469756da3475055a96c6f72615f6e616d65c08aa474797065ab426c6f636b53746f726564ac626c6f636b5f68617368657391cd1234b1706172656e745f626c6f636b5f68617368c0a9746f6b656e5f69647392090aaa626c6f636b5f73697a6502a76c6f72615f696403a66d656469756dc0a96c6f72615f6e616d65c0a967726f75705f69647800b26b765f63616368655f737065635f6b696e64ae66756c6c5f617474656e74696f6e88a474797065ab426c6f636b53746f726564ac626c6f636b5f6861736865739105b1706172656e745f626c6f636b5f68617368c0a9746f6b656e5f69647393010203aa626c6f636b5f73697a6504a76c6f72615f6964c0a66d656469756da3475055a96c6f72615f6e616d65c083a474797065ac426c6f636b52656d6f766564ac626c6f636b5f68617368657392cd1234c420ffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffa66d656469756da347505581a474797065b0416c6c426c6f636b73436c6561726564c0"; + + /// `KVEventBatch(ts=1700000001.0, data_parallel_rank=1)` with one + /// `BlockStored`: int hash 42, parent 41, tokens [100, 101], block size 2. + pub(crate) const BATCH2: &str = "93cb41d954fc404000009188a474797065ab426c6f636b53746f726564ac626c6f636b5f686173686573912ab1706172656e745f626c6f636b5f6861736829a9746f6b656e5f696473926465aa626c6f636b5f73697a6502a76c6f72615f6964c0a66d656469756da3475055a96c6f72615f6e616d65c001"; + + pub(crate) fn bytes(hex: &str) -> Vec { + (0..hex.len()) + .step_by(2) + .map(|i| u8::from_str_radix(&hex[i..i + 2], 16).expect("hex")) + .collect() + } + + /// A publisher message: `[topic, sequence (u64 big-endian), payload]`. + pub(crate) fn frame(topic: &[u8], sequence: u64, payload: &[u8]) -> ZmqMessage { + let mut message = ZmqMessage::from(topic.to_vec()); + message.push_back(sequence.to_be_bytes().to_vec().into()); + message.push_back(payload.to_vec().into()); + message + } +} + +#[cfg(test)] +mod tests { + use std::{ + sync::atomic::{AtomicU64, Ordering}, + time::{Duration, Instant, SystemTime, UNIX_EPOCH}, + }; + + use futures::StreamExt; + use smg_grpc_client::common_proto::{kv_cache_event, KvCacheLocality, KvCacheTier}; + use tokio::time::timeout; + use zeromq::{PubSocket, RouterSocket, SocketEvent}; + + use super::{golden, *}; + use crate::kv_wire::WireBatch; + + fn decode(hex: &str) -> WireBatch { + rmp_serde::from_slice::>(&golden::bytes(hex)) + .expect("golden batch decodes") + .0 + } + + fn convert_batch( + batch: WireBatch, + sequence_number: u64, + event_id: &mut u64, + ) -> common::KvEventBatch { + Normalizer::new().normalize_batch(batch, sequence_number, event_id) + } + + fn stored(event: &common::KvCacheEvent) -> &common::KvBlocksStored { + match &event.data { + Some(kv_cache_event::Data::Stored(stored)) => stored, + other => panic!("expected a stored event, got {other:?}"), + } + } + + fn block(block_hash: i64, token_ids: Vec, lora_id: Option) -> common::KvBlock { + common::KvBlock { + block_hash, + block_size: i32::try_from(token_ids.len()).unwrap(), + token_ids, + lora_id, + cache_level: None, + ..Default::default() + } + } + + #[test] + fn the_relay_starts_at_boot_unless_told_to_wait() { + assert!(starts_at_boot(None)); + assert!(starts_at_boot(Some(""))); + assert!(starts_at_boot(Some("eager"))); + assert!(!starts_at_boot(Some("lazy"))); + assert!(!starts_at_boot(Some(" Lazy "))); + } + + #[test] + fn endpoint_for_rank_mirrors_the_python_helper() { + assert_eq!(endpoint_for_rank("tcp://*:5557", 0), "tcp://127.0.0.1:5557"); + assert_eq!( + endpoint_for_rank("tcp://0.0.0.0:5557", 0), + "tcp://127.0.0.1:5557" + ); + assert_eq!(endpoint_for_rank("tcp://*:5557", 2), "tcp://127.0.0.1:5559"); + assert_eq!( + endpoint_for_rank("tcp://10.0.0.1:5557", 1), + "tcp://10.0.0.1:5558" + ); + assert_eq!(endpoint_for_rank("tcp://host:port", 1), "tcp://host:port"); + assert_eq!(endpoint_for_rank("ipc:///tmp/kv", 1), "ipc:///tmp/kv"); + } + + /// The golden batches convert to exactly what the Python relay produced /// for them: sha256 hashes reduced to their low 64 bits, an unaligned /// store skipped but its event id consumed, `lora_id` and the parent /// carried, `dp_rank` set only when the publisher set it. @@ -484,11 +1622,17 @@ mod tests { Some(kv_cache_event::Data::Removed(common::KvBlocksRemoved { block_hashes: vec![0x1234, -1], cache_level: None, + tier: Some(KvCacheTier::Device as i32), + medium: Some("GPU".to_string()), + locality: Some(KvCacheLocality::Local as i32), + ..Default::default() })) ); assert_eq!( batch.events[3].data, - Some(kv_cache_event::Data::Cleared(common::KvCacheCleared {})) + Some(kv_cache_event::Data::Cleared( + common::KvCacheCleared::default() + )) ); let batch = convert_batch(decode(golden::BATCH2), 10, &mut event_id); @@ -511,13 +1655,13 @@ mod tests { .unwrap() .0; assert!(batch.events.is_empty()); - assert_eq!(batch.data_parallel_rank, None); + assert_eq!(batch.dp_rank, None); let long = rmp_serde::to_vec(&(1.5f64, Vec::::new(), 2i32, "future")).unwrap(); let batch = rmp_serde::from_slice::>(&long) .unwrap() .0; - assert_eq!(batch.data_parallel_rank, Some(2)); + assert_eq!(batch.dp_rank, Some(2)); let unknown = rmp_serde::to_vec(&serde_json::json!([ 1.5, @@ -560,80 +1704,1526 @@ mod tests { ); } - /// A local publisher's frames come out as proto batches under the - /// publisher's sequence numbers; short frames, undecodable payloads and - /// other topics are skipped; dropping the stream closes the connection. - #[tokio::test] - async fn relays_a_local_publisher_and_closes_with_the_stream() { + struct Lab { + publisher: PubSocket, + router: Option, + relay: Arc, + } + + /// A local publisher (and replay ROUTER when asked) with a started relay. + async fn start_lab(history_batches: usize, with_replay: bool) -> Lab { let mut publisher = PubSocket::new(); - let mut monitor = publisher.monitor(); let endpoint = publisher .bind("tcp://127.0.0.1:0") .await .expect("publisher binds") .to_string(); - let mut stream = relay(endpoint, "kv".to_string()); - let batch1 = golden::bytes(golden::BATCH1); - let batch2 = golden::bytes(golden::BATCH2); - - // The subscription reaches the publisher a moment after the connect; - // probe with sequence 0 until a batch comes through. - let mut first = None; - for _ in 0..200 { - publisher - .send(golden::frame(b"kv", 0, &batch1)) + let (router, replay_endpoint) = if with_replay { + let mut router = RouterSocket::new(); + let endpoint = router + .bind("tcp://127.0.0.1:0") .await - .expect("publish"); - if let Ok(item) = timeout(Duration::from_millis(50), stream.next()).await { - first = Some(item.expect("stream open").expect("a batch")); - break; - } + .expect("router binds") + .to_string(); + (Some(router), Some(endpoint)) + } else { + (None, None) + }; + let relay = KvEventRelay::new(RelayConfig { + endpoint, + replay_endpoint, + topic: "kv".to_string(), + history_batches, + history_bytes: 64 << 20, + replay_timeout: Duration::from_secs(2), + load_tick: DEFAULT_LOAD_TICK, + heartbeat_interval: DEFAULT_HEARTBEAT_INTERVAL, + heartbeat_backoff: DEFAULT_HEARTBEAT_BACKOFF, + }); + relay.start(); + let mut lab = Lab { + publisher, + router, + relay, + }; + if with_replay { + // The relay asks the replay socket for the publisher's start + // before anything else; this publisher has nothing yet. + lab.answer_replay(0, &[]).await; } - let first = first.expect("the subscription went live"); - assert_eq!(first.sequence_number, 0); - assert_eq!(first.events.len(), 4); + lab + } - let mut short = ZmqMessage::from(b"kv".to_vec()); - short.push_back(1u64.to_be_bytes().to_vec().into()); - publisher.send(short).await.expect("publish"); - publisher - .send(golden::frame(b"kv", 2, b"not msgpack")) - .await - .expect("publish"); - publisher - .send(golden::frame(b"other", 3, &batch2)) - .await - .expect("publish"); - publisher - .send(golden::frame(b"kv", 4, &batch2)) + /// A load source the tests steer: the record the relay attaches, and how + /// often it was asked. + struct StubLoads { + record: Mutex, + asked: AtomicU64, + } + + impl StubLoads { + fn new(running: u32, waiting: u32) -> Arc { + Arc::new(Self { + record: Mutex::new(common::EngineLoad { + running_requests: running, + waiting_requests: waiting, + waiting_uncached_tokens: Some(4_096), + token_usage: 0.25, + gen_throughput: 1_200.0, + max_running_requests: 64, + cache_hit_rate: Some(0.5), + num_used_tokens: Some(1_024), + max_total_num_tokens: Some(4_096), + ..Default::default() + }), + asked: AtomicU64::new(0), + }) + } + + fn set(&self, running: u32, waiting: u32) { + let mut record = self.record.lock().unwrap_or_else(PoisonError::into_inner); + record.running_requests = running; + record.waiting_requests = waiting; + } + } + + impl LoadSource for StubLoads { + fn load(&self, _dp_rank: Option) -> Option { + self.asked.fetch_add(1, Ordering::Relaxed); + Some( + self.record + .lock() + .unwrap_or_else(PoisonError::into_inner) + .clone(), + ) + } + } + + /// A relay with a load source and short record intervals (`tick`, + /// heartbeat, backoff), primed with sequence 0. + async fn start_lab_with_loads( + loads: Arc, + tick: Duration, + heartbeat: Duration, + backoff: Duration, + ) -> Lab { + start_lab_with_loads_at(loads, tick, heartbeat, backoff, true).await + } + + /// [`start_lab_with_loads`], with or without the publisher's first batch. + async fn start_lab_with_loads_at( + loads: Arc, + tick: Duration, + heartbeat: Duration, + backoff: Duration, + prime: bool, + ) -> Lab { + let mut publisher = PubSocket::new(); + let endpoint = publisher + .bind("tcp://127.0.0.1:0") .await - .expect("publish"); - // Probe duplicates may still be queued; the next new batch is 4. - let next = loop { - let batch = timeout(Duration::from_secs(5), stream.next()) - .await - .expect("a batch in time") - .expect("stream open") - .expect("a batch"); - if batch.sequence_number != 0 { - break batch; - } + .expect("publisher binds") + .to_string(); + let relay = KvEventRelay::new(RelayConfig { + endpoint, + replay_endpoint: None, + topic: "kv".to_string(), + history_batches: 100, + history_bytes: 64 << 20, + replay_timeout: Duration::from_secs(2), + load_tick: tick, + heartbeat_interval: heartbeat, + heartbeat_backoff: backoff, + }); + relay.set_load_source(loads); + relay.start(); + let mut lab = Lab { + publisher, + router: None, + relay, }; - assert_eq!(next.sequence_number, 4); - assert_eq!(next.dp_rank, Some(1)); - assert_eq!(stored(&next.events[0]).blocks[0].block_hash, 42); + if prime { + lab.prime().await; + } + lab + } - drop(stream); - let disconnected = timeout(Duration::from_secs(5), async { - while let Some(event) = monitor.next().await { - if matches!(event, SocketEvent::Disconnected(_)) { - return true; - } - } - false - }) - .await - .expect("the publisher notices in time"); - assert!(disconnected); + fn record_of(batch: &common::KvEventBatch) -> &common::EngineLoad { + batch.load.as_ref().expect("a load record on the batch") + } + + /// Every batch a subscriber receives, whether from the history, live or + /// a snapshot chunk, carries the servicer's load record, numbered per + /// stream; a relay without a source sends none. + #[tokio::test] + async fn every_batch_a_subscriber_receives_carries_the_load_record() { + let loads = StubLoads::new(3, 1); + let mut lab = start_lab_with_loads( + loads, + Duration::from_secs(10), + Duration::from_secs(10), + Duration::from_secs(10), + ) + .await; + let batch2 = golden::bytes(golden::BATCH2); + lab.publish(1, &batch2).await; + lab.wait_relayed(2).await; + // History first (0 and 1), then live (2). + let mut stream = lab.subscribe(0).expect("history"); + let first = read(&mut stream).await; + assert_eq!(first.sequence_number, 0); + let record = record_of(&first); + assert_eq!( + ( + record.running_requests, + record.waiting_requests, + record.waiting_uncached_tokens, + record.sample, + record.load_only + ), + (3, 1, Some(4_096), 1, false) + ); + // The first record carries the telemetry; the next event batch's + // record is the core only. + assert_eq!(record.cache_hit_rate, Some(0.5)); + assert_eq!(record.max_total_num_tokens, Some(4_096)); + let second = read(&mut stream).await; + assert_eq!(record_of(&second).sample, 2); + assert_eq!(record_of(&second).cache_hit_rate, None); + assert_eq!(record_of(&second).max_total_num_tokens, None); + lab.publish(2, &batch2).await; + let live = read(&mut stream).await; + assert_eq!((live.sequence_number, record_of(&live).sample), (2, 3)); + // A snapshot chunk (the window of this second lab has no start). + let mut rolled = start_lab_with_loads( + StubLoads::new(1, 0), + Duration::from_secs(10), + Duration::from_secs(10), + Duration::from_secs(10), + ) + .await; + rolled.publish(1, &batch2).await; + rolled.wait_relayed(2).await; + let mut whole = rolled.subscribe(0).expect("history"); + assert_eq!(record_of(&read(&mut whole).await).running_requests, 1); + // No source: no record. + let mut plain = start_lab(10, false).await; + plain.prime().await; + let mut bare = plain.subscribe(0).expect("history"); + assert!(read(&mut bare).await.load.is_none()); + } + + /// While the publisher is quiet the stream still speaks: a `load_only` + /// batch (no events, the last sequence repeated) when the record moved, + /// a heartbeat after the interval, and after two unchanged heartbeats + /// only every backoff interval. + #[tokio::test] + async fn a_quiet_publishers_stream_sends_load_only_batches_and_heartbeats() { + let loads = StubLoads::new(2, 0); + let mut lab = start_lab_with_loads( + Arc::clone(&loads), + Duration::from_millis(20), + Duration::from_millis(200), + Duration::from_millis(600), + ) + .await; + let mut stream = lab.subscribe(0).expect("history"); + let first = read(&mut stream).await; + assert_eq!(first.sequence_number, 0); + // A change with no KV event: a load-only batch within a tick or two. + let started = Instant::now(); + loads.set(5, 2); + let change = read(&mut stream).await; + let record = record_of(&change); + assert!(load_only_marker(&change), "{change:?}"); + assert_eq!((change.sequence_number, change.events.len()), (0, 0)); + assert_eq!((record.running_requests, record.waiting_requests), (5, 2)); + // A heartbeat carries the telemetry. + assert_eq!(record.cache_hit_rate, Some(0.5)); + assert!( + started.elapsed() < Duration::from_millis(150), + "the change took {:?}", + started.elapsed() + ); + // Nothing changes: heartbeats at the interval, then at the backoff. + let mut gaps = Vec::new(); + let mut last = Instant::now(); + for _ in 0..4 { + let beat = read(&mut stream).await; + assert!(load_only_marker(&beat)); + assert_eq!(record_of(&beat).running_requests, 5); + gaps.push(last.elapsed()); + last = Instant::now(); + } + assert!( + gaps[0] >= Duration::from_millis(150) && gaps[0] < Duration::from_millis(500), + "first heartbeat after {:?}", + gaps[0] + ); + assert!( + gaps[1] < Duration::from_millis(500), + "second heartbeat after {:?}", + gaps[1] + ); + assert!( + gaps[2] >= Duration::from_millis(500) && gaps[3] >= Duration::from_millis(500), + "backed off: {gaps:?}" + ); + // A real batch resets the backoff and carries the record itself. + lab.publish(1, &golden::bytes(golden::BATCH2)).await; + let live = read(&mut stream).await; + assert_eq!(live.sequence_number, 1); + assert!(!load_only_marker(&live)); + let beat = read(&mut stream).await; + assert!(load_only_marker(&beat)); + assert_eq!(beat.sequence_number, 1, "repeats the last sequence sent"); + } + + fn load_only_marker(batch: &common::KvEventBatch) -> bool { + batch.load.as_ref().is_some_and(|load| load.load_only) + } + + /// A relay whose replay socket is not bound yet: the relay keeps asking + /// for the publisher's start until it is. + async fn start_lab_without_replay_socket_yet() -> (Lab, String) { + let mut publisher = PubSocket::new(); + let endpoint = publisher + .bind("tcp://127.0.0.1:0") + .await + .expect("publisher binds") + .to_string(); + let replay_endpoint = format!( + "tcp://127.0.0.1:{}", + portpicker::pick_unused_port().expect("a free replay port") + ); + let relay = KvEventRelay::new(RelayConfig { + endpoint, + replay_endpoint: Some(replay_endpoint.clone()), + topic: "kv".to_string(), + history_batches: 100, + history_bytes: 64 << 20, + replay_timeout: Duration::from_secs(2), + load_tick: DEFAULT_LOAD_TICK, + heartbeat_interval: DEFAULT_HEARTBEAT_INTERVAL, + heartbeat_backoff: DEFAULT_HEARTBEAT_BACKOFF, + }); + relay.start(); + ( + Lab { + publisher, + router: None, + relay, + }, + replay_endpoint, + ) + } + + impl Lab { + async fn publish(&mut self, sequence: u64, payload: &[u8]) { + self.publisher + .send(golden::frame(b"kv", sequence, payload)) + .await + .expect("publish"); + } + + /// The subscription reaches the publisher a moment after the + /// connect; publish sequence 0 until the relay has it. + async fn prime(&mut self) { + let batch1 = golden::bytes(golden::BATCH1); + for _ in 0..250 { + self.publish(0, &batch1).await; + if self.relay.counts().relayed >= 1 { + return; + } + tokio::time::sleep(Duration::from_millis(20)).await; + } + panic!("the relay never saw sequence 0"); + } + + async fn wait_relayed(&self, count: u64) { + for _ in 0..250 { + if self.relay.counts().relayed >= count { + return; + } + tokio::time::sleep(Duration::from_millis(20)).await; + } + panic!( + "the relay relayed {} batches, not {count}", + self.relay.counts().relayed + ); + } + + fn subscribe(&self, cursor: u64) -> Result, Status> { + self.relay.subscribe(common::SubscribeKvEventsRequest { + start_sequence_number: cursor, + }) + } + + /// Answer the next replay request as vLLM's ROUTER does, with + /// `replies` and the end marker; the request must ask from `start`. + async fn answer_replay(&mut self, start: u64, replies: &[(u64, Vec)]) { + let router = self.router.as_mut().expect("a replay socket"); + let request = timeout(Duration::from_secs(5), router.recv()) + .await + .expect("a replay request in time") + .expect("request"); + let frames = replay_request_frames(&request, start); + self.send_replay(&frames, replies).await; + } + + /// Publish `sequence` (the subscription reaches the publisher a + /// moment after the connect) until the relay asks the replay socket, + /// which it must do from `start`; the request's frames. + async fn publish_until_replay_requested( + &mut self, + sequence: u64, + payload: &[u8], + start: u64, + ) -> Vec> { + let router = self.router.as_mut().expect("a replay socket"); + for _ in 0..250 { + self.publisher + .send(golden::frame(b"kv", sequence, payload)) + .await + .expect("publish"); + if let Ok(request) = timeout(Duration::from_millis(20), router.recv()).await { + return replay_request_frames(&request.expect("request"), start); + } + } + panic!("the relay never asked the replay socket"); + } + + /// Reply to the request `frames` with `replies` and the end marker. + async fn send_replay(&mut self, frames: &[Vec], replies: &[(u64, Vec)]) { + let router = self.router.as_mut().expect("a replay socket"); + for (sequence, payload) in replies { + let mut reply = ZmqMessage::from(frames[0].clone()); + reply.push_back(Vec::new().into()); + reply.push_back(b"kv".to_vec().into()); + reply.push_back(sequence.to_be_bytes().to_vec().into()); + reply.push_back(payload.clone().into()); + router.send(reply).await.expect("reply"); + } + let mut end = ZmqMessage::from(frames[0].clone()); + end.push_back(Vec::new().into()); + end.push_back(Vec::new().into()); + end.push_back(END_SEQUENCE.to_vec().into()); + end.push_back(Vec::new().into()); + router.send(end).await.expect("end marker"); + } + } + + /// A replay request's frames (`[identity, empty, start]`), checked to + /// ask from `start`. + fn replay_request_frames(request: &ZmqMessage, start: u64) -> Vec> { + let frames: Vec> = request.iter().map(|frame| frame.to_vec()).collect(); + assert_eq!(frames.len(), 3, "[identity, empty, start]"); + assert_eq!(frames[1], b""); + assert_eq!(frames[2], start.to_be_bytes()); + frames + } + + fn refused(result: Result, Status>) -> Status { + match result { + Err(status) => status, + Ok(_) => panic!("the subscription was accepted"), + } + } + + async fn read(stream: &mut BoxStream) -> common::KvEventBatch { + timeout(Duration::from_secs(5), stream.next()) + .await + .expect("a batch in time") + .expect("stream open") + .expect("a batch") + } + + async fn read_error(stream: &mut BoxStream) -> Status { + timeout(Duration::from_secs(5), stream.next()) + .await + .expect("an item in time") + .expect("stream open") + .expect_err("an error") + } + + async fn read_end(stream: &mut BoxStream) { + assert!(timeout(Duration::from_secs(5), stream.next()) + .await + .expect("the end in time") + .is_none()); + } + + #[tokio::test] + async fn history_serves_a_cursor_inside_the_window_then_live() { + let mut lab = start_lab(100, false).await; + lab.prime().await; + let batch2 = golden::bytes(golden::BATCH2); + for sequence in 1..=5 { + lab.publish(sequence, &batch2).await; + } + lab.wait_relayed(6).await; + + let mut stream = lab.subscribe(2).expect("inside the window"); + for expected in 3..=5 { + assert_eq!(read(&mut stream).await.sequence_number, expected); + } + lab.publish(6, &batch2).await; + assert_eq!(read(&mut stream).await.sequence_number, 6); + + // The newest sequence as the cursor: nothing owed, live from here. + let mut caught_up = lab.subscribe(6).expect("at the newest"); + lab.publish(7, &batch2).await; + assert_eq!(read(&mut caught_up).await.sequence_number, 7); + assert_eq!(read(&mut stream).await.sequence_number, 7); + assert_eq!(lab.relay.counts().served_from_history, 2); + } + + #[tokio::test] + async fn cursors_outside_the_window_are_out_of_range() { + let fresh = start_lab(10, false).await; + let status = refused(fresh.subscribe(1)); + assert_eq!(status.code(), tonic::Code::OutOfRange); + assert!(status.message().contains("holds no history yet")); + + let mut lab = start_lab(3, false).await; + lab.prime().await; + let batch2 = golden::bytes(golden::BATCH2); + for sequence in 1..=9 { + lab.publish(sequence, &batch2).await; + } + lab.wait_relayed(10).await; + // The window is 7..=9: cursor 6 wants 7, served; cursor 5 wants 6, gone. + let mut stream = lab.subscribe(6).expect("the oldest batch is wanted"); + for expected in 7..=9 { + assert_eq!(read(&mut stream).await.sequence_number, expected); + } + let behind = refused(lab.subscribe(5)); + assert_eq!(behind.code(), tonic::Code::OutOfRange); + assert!(behind.message().contains("from sequence 7")); + let ahead = refused(lab.subscribe(20)); + assert_eq!(ahead.code(), tonic::Code::OutOfRange); + assert!(ahead.message().contains("publisher restarted")); + assert_eq!(lab.relay.counts().out_of_range, 2); + } + + #[tokio::test] + async fn a_subscriber_without_a_cursor_gets_the_whole_history_or_a_snapshot() { + let mut lab = start_lab(100, false).await; + lab.prime().await; + let batch2 = golden::bytes(golden::BATCH2); + lab.publish(1, &batch2).await; + lab.publish(2, &batch2).await; + lab.wait_relayed(3).await; + // Everything since the publisher's first batch is here: hand it over. + let mut stream = lab.subscribe(0).expect("live"); + for expected in 0..=2 { + assert_eq!(read(&mut stream).await.sequence_number, expected); + } + lab.publish(3, &batch2).await; + assert_eq!(read(&mut stream).await.sequence_number, 3); + + // A relay that joined after the publisher's first batches has an + // incomplete window: what it knows of the state, as a snapshot cut + // at its newest sequence, then live. + let mut late = start_lab(100, false).await; + for _ in 0..250 { + late.publish(5, &batch2).await; + if late.relay.counts().relayed >= 1 { + break; + } + tokio::time::sleep(Duration::from_millis(20)).await; + } + assert_eq!(late.relay.counts().relayed, 1, "the relay saw sequence 5"); + let mut stream = late.subscribe(0).expect("a snapshot"); + let chunk = read(&mut stream).await; + assert_eq!(chunk.sequence_number, 5, "cut at the newest sequence"); + assert_eq!( + chunk.snapshot, + Some(common::KvSnapshotChunk { + index: 0, + count: 1, + blocks: 1, + unknown_before: 5, + }), + "the five sequences before the relay joined are unknown" + ); + assert_eq!(chunk.dp_rank, Some(1)); + assert!(is_cleared(&chunk.events[0])); + assert_eq!(stored_hashes(&chunk), vec![(Some(41), vec![42])]); + late.publish(6, &batch2).await; + assert_eq!( + read(&mut stream).await.sequence_number, + 6, + "live from the next sequence" + ); + assert_eq!(late.relay.counts().served_snapshots, 1); + } + + #[tokio::test] + async fn a_publisher_gap_is_filled_from_the_engines_replay() { + let mut lab = start_lab(100, true).await; + lab.prime().await; + let batch2 = golden::bytes(golden::BATCH2); + lab.publish(1, &batch2).await; + lab.wait_relayed(2).await; + let mut stream = lab.subscribe(1).expect("caught up"); + + // 2 and 3 never reach the SUB; 4 reveals the gap. + lab.publish(4, &batch2).await; + lab.answer_replay( + 2, + &[ + (2, batch2.clone()), + (3, batch2.clone()), + (4, batch2.clone()), + ], + ) + .await; + for expected in 2..=4 { + assert_eq!(read(&mut stream).await.sequence_number, expected); + } + lab.publish(5, &batch2).await; + assert_eq!(read(&mut stream).await.sequence_number, 5); + let counts = lab.relay.counts(); + assert_eq!( + ( + counts.publisher_gaps, + counts.gap_batches_recovered, + counts.gap_batches_lost + ), + (1, 3, 0) + ); + assert_eq!( + counts.relayed, 6, + "the replay's copy of 4 was relayed, the live one skipped" + ); + // The window is complete: a cursor inside it is served across the + // recovered stretch. + let mut resumed = lab.subscribe(1).expect("served"); + for expected in 2..=5 { + assert_eq!(read(&mut resumed).await.sequence_number, expected); + } + } + + #[tokio::test] + async fn a_gap_the_engine_cannot_fill_leaves_a_hole_not_a_dead_end() { + let mut lab = start_lab(100, false).await; + lab.prime().await; + let batch2 = golden::bytes(golden::BATCH2); + lab.publish(1, &batch2).await; + lab.wait_relayed(2).await; + let mut stream = lab.subscribe(1).expect("caught up"); + lab.publish(4, &batch2).await; + assert_eq!( + read(&mut stream).await.sequence_number, + 4, + "the live stream jumps" + ); + let counts = lab.relay.counts(); + assert_eq!( + ( + counts.publisher_gaps, + counts.gap_batches_recovered, + counts.gap_batches_lost + ), + (1, 0, 2) + ); + // A resume from inside the hole or before it skips it, as the gateway + // then settles the gap itself instead of looping on OUT_OF_RANGE. + let mut from_before = lab.subscribe(1).expect("served"); + assert_eq!(read(&mut from_before).await.sequence_number, 4); + let mut from_inside = lab.subscribe(2).expect("served"); + assert_eq!(read(&mut from_inside).await.sequence_number, 4); + // The window has a hole, so it is not the publisher's whole state: a + // subscriber without a cursor gets the snapshot (both copies of block + // 42, from sequences 1 and 4) cut at 4, then live. + let mut no_cursor = lab.subscribe(0).expect("a snapshot"); + let chunk = read(&mut no_cursor).await; + assert_eq!(chunk.sequence_number, 4); + assert_eq!(chunk.snapshot.as_ref().map(|chunk| chunk.blocks), Some(2)); + assert_eq!( + stored_hashes(&chunk), + vec![(Some(41), vec![42]), (Some(41), vec![42])] + ); + lab.publish(5, &batch2).await; + assert_eq!(read(&mut no_cursor).await.sequence_number, 5); + } + + #[tokio::test] + async fn a_partial_replay_recovers_what_it_can_and_loses_the_rest() { + let mut lab = start_lab(100, true).await; + lab.prime().await; + let batch2 = golden::bytes(golden::BATCH2); + let mut stream = lab.subscribe(0).expect("live"); + assert_eq!(read(&mut stream).await.sequence_number, 0); + lab.publish(5, &batch2).await; + // The engine's buffer starts at 3: 1 and 2 are gone for good. + lab.answer_replay(1, &[(3, batch2.clone()), (4, batch2.clone())]) + .await; + for expected in [3, 4, 5] { + assert_eq!(read(&mut stream).await.sequence_number, expected); + } + let counts = lab.relay.counts(); + assert_eq!( + (counts.gap_batches_recovered, counts.gap_batches_lost), + (2, 2) + ); + } + + #[tokio::test] + async fn undecodable_payloads_relay_as_empty_batches() { + let mut lab = start_lab(100, false).await; + lab.prime().await; + let mut stream = lab.subscribe(0).expect("live"); + assert_eq!(read(&mut stream).await.events.len(), 4); + lab.publish(1, b"not msgpack").await; + let empty = read(&mut stream).await; + assert_eq!((empty.sequence_number, empty.events.len()), (1, 0)); + let mut short = ZmqMessage::from(b"kv".to_vec()); + short.push_back(2u64.to_be_bytes().to_vec().into()); + lab.publisher.send(short).await.expect("publish"); + lab.publish(2, &golden::bytes(golden::BATCH2)).await; + assert_eq!( + read(&mut stream).await.sequence_number, + 2, + "a short frame is nothing" + ); + assert_eq!(lab.relay.counts().undecodable_batches, 1); + } + + #[tokio::test] + async fn a_sequence_regression_ends_live_streams_and_clears_the_history() { + let mut lab = start_lab(100, false).await; + lab.prime().await; + let batch2 = golden::bytes(golden::BATCH2); + for sequence in 1..=3 { + lab.publish(sequence, &batch2).await; + } + lab.wait_relayed(4).await; + let mut stream = lab.subscribe(0).expect("whole history"); + for expected in 0..=3 { + assert_eq!(read(&mut stream).await.sequence_number, expected); + } + + // The publisher restarts and counts from 0 again. + lab.publish(0, &batch2).await; + let status = read_error(&mut stream).await; + assert_eq!(status.code(), tonic::Code::DataLoss); + assert!(status.message().contains("restarted")); + read_end(&mut stream).await; + assert_eq!(lab.relay.counts().publisher_restarts, 1); + + // The old incarnation's cursor is refused; the new one's complete + // history is handed to a fresh subscriber. + let stale = refused(lab.subscribe(3)); + assert_eq!(stale.code(), tonic::Code::OutOfRange); + let mut fresh = lab + .subscribe(0) + .expect("the new incarnation from its start"); + assert_eq!(read(&mut fresh).await.sequence_number, 0); + lab.publish(1, &batch2).await; + assert_eq!(read(&mut fresh).await.sequence_number, 1); + } + + fn cleared_payload() -> Vec { + let batch = serde_json::json!([1700000002.0, [{"type": "AllBlocksCleared"}], 0]); + rmp_serde::to_vec_named(&batch).expect("encodes") + } + + #[test] + fn restart_rules_read_the_same_on_every_wire() { + let mut shared = Shared { + history: History::new(10, usize::MAX), + state: LiveState::new(), + cursor: Some(500), + started_at: Some(0), + unknown_before: 0, + generation: 0, + counts: RelayCounts::default(), + wire: WireCounts::default(), + failed: None, + }; + let restart = |reason| Admission::Restart { reason, last: 500 }; + assert_eq!(shared.admit(501, false), Admission::Accept); + assert_eq!( + shared.admit(501, true), + Admission::Accept, + "a flush continues the sequence" + ); + assert_eq!( + shared.admit(503, false), + Admission::Gap { from: 501, to: 502 } + ); + assert_eq!(shared.admit(500, false), Admission::Duplicate); + assert_eq!( + shared.admit(500, true), + restart(RestartReason::StartupClear) + ); + assert_eq!( + shared.admit(499, false), + restart(RestartReason::SequenceRegression) + ); + assert_eq!( + shared.admit(0, false), + restart(RestartReason::CounterRestarted) + ); + assert_eq!( + shared.admit(1, true), + restart(RestartReason::CounterRestarted) + ); + // Under a cursor of 1, a repeated 1 is a duplicate unless it clears. + shared.cursor = Some(1); + assert_eq!(shared.admit(1, false), Admission::Duplicate); + assert_eq!( + shared.admit(1, true), + Admission::Restart { + reason: RestartReason::StartupClear, + last: 1 + } + ); + assert_eq!( + shared.admit(0, false), + Admission::Restart { + reason: RestartReason::SequenceRegression, + last: 1 + } + ); + shared.cursor = None; + assert_eq!( + shared.admit(7, true), + Admission::Accept, + "no cursor yet: anything goes" + ); + } + + /// SGLang's first batch after a start carries `AllBlocksCleared`; under a + /// sequence the relay already passed it is a restart, not a duplicate. + #[tokio::test] + async fn the_engines_startup_clear_under_a_passed_cursor_is_a_restart() { + let mut lab = start_lab(100, false).await; + lab.prime().await; + let batch2 = golden::bytes(golden::BATCH2); + for sequence in 1..=3 { + lab.publish(sequence, &batch2).await; + } + lab.wait_relayed(4).await; + let mut stream = lab.subscribe(3).expect("caught up"); + // A repeated sequence without a clear is a duplicate. + lab.publish(3, &batch2).await; + lab.publish(4, &batch2).await; + assert_eq!(read(&mut stream).await.sequence_number, 4); + + // The engine comes back and its startup clear lands on a sequence + // the relay already passed. + lab.publish(2, &cleared_payload()).await; + let status = read_error(&mut stream).await; + assert_eq!(status.code(), tonic::Code::DataLoss); + assert_eq!(lab.relay.counts().publisher_restarts, 1); + // The new incarnation began at 2 with the clear: nothing is live, and + // a fresh subscriber is told so by a snapshot of one clear, cut at 2. + let mut fresh = lab.subscribe(0).expect("a snapshot"); + let chunk = read(&mut fresh).await; + assert_eq!((chunk.sequence_number, chunk.events.len()), (2, 1)); + assert!(is_cleared(&chunk.events[0])); + lab.publish(3, &batch2).await; + assert_eq!(read(&mut fresh).await.sequence_number, 3); + } + + /// A counter back at 0 or 1 after a cursor above them is a restart even + /// when nothing else says so (the mock engine's restart-publisher hook, + /// a vLLM process restart: no clear on that wire). + #[tokio::test] + async fn a_counter_back_at_its_start_is_a_restart() { + let mut lab = start_lab(100, false).await; + lab.prime().await; + let batch2 = golden::bytes(golden::BATCH2); + for sequence in 1..=5 { + lab.publish(sequence, &batch2).await; + } + lab.wait_relayed(6).await; + let mut stream = lab.subscribe(5).expect("caught up"); + lab.publish(1, &batch2).await; + let status = read_error(&mut stream).await; + assert_eq!(status.code(), tonic::Code::DataLoss); + read_end(&mut stream).await; + // The new incarnation started at 1, not 0: its window is not complete + // from the publisher's first batch, so a fresh subscriber gets the + // new incarnation's state as a snapshot (the restart emptied the old + // one), then live. + let mut fresh = lab.subscribe(0).expect("a snapshot"); + let chunk = read(&mut fresh).await; + assert_eq!(chunk.sequence_number, 1); + assert_eq!(stored_hashes(&chunk), vec![(Some(41), vec![42])]); + lab.publish(2, &batch2).await; + assert_eq!(read(&mut fresh).await.sequence_number, 2); + let mut resumed = lab.subscribe(1).expect("inside the new window"); + assert_eq!(read(&mut resumed).await.sequence_number, 2); + } + + /// A publisher already counting when the relay's subscription reaches + /// it: the relay asks the engine's replay for everything from 0 before + /// relaying what it saw, and the window is the publisher's whole life. + #[tokio::test] + async fn a_publisher_already_counting_when_the_relay_joins_is_replayed_from_its_start() { + let mut lab = start_lab(100, true).await; + let batch2 = golden::bytes(golden::BATCH2); + // Sequences 0..=2 went out before the subscription landed; 3 is the + // first the relay sees, and it must ask for 0 before relaying it. + let request = lab.publish_until_replay_requested(3, &batch2, 0).await; + lab.send_replay( + &request, + &[ + (0, batch2.clone()), + (1, batch2.clone()), + (2, batch2.clone()), + (3, batch2.clone()), + ], + ) + .await; + lab.wait_relayed(4).await; + let counts = lab.relay.counts(); + assert_eq!( + ( + counts.gap_batches_recovered, + counts.gap_batches_lost, + counts.unknown_before_start, + counts.publisher_gaps, + ), + (4, 0, 0, 0), + "{counts:?}" + ); + let mut stream = lab.subscribe(0).expect("the whole history"); + for expected in 0..=3 { + assert_eq!(read(&mut stream).await.sequence_number, expected); + } + lab.publish(4, &batch2).await; + assert_eq!(read(&mut stream).await.sequence_number, 4); + let counts = lab.relay.counts(); + assert_eq!( + (counts.served_from_history, counts.served_snapshots), + (1, 0) + ); + } + + /// The replay that answers a late join may itself start past 0 (the + /// engine's buffer rolled): what it gives is relayed, the sequences + /// before it are holes and counted as unknown, a subscriber from zero + /// gets a snapshot that says so, and cursors inside the holes are served + /// past them. + #[tokio::test] + async fn a_replay_that_starts_past_zero_leaves_the_earlier_sequences_unknown() { + let mut lab = start_lab(100, true).await; + let batch2 = golden::bytes(golden::BATCH2); + let request = lab.publish_until_replay_requested(7, &batch2, 0).await; + lab.send_replay( + &request, + &[ + (5, batch2.clone()), + (6, batch2.clone()), + (7, batch2.clone()), + ], + ) + .await; + lab.wait_relayed(3).await; + let counts = lab.relay.counts(); + assert_eq!( + ( + counts.gap_batches_recovered, + counts.gap_batches_lost, + counts.unknown_before_start, + ), + (3, 5, 5), + "{counts:?}" + ); + let mut fresh = lab.subscribe(0).expect("a snapshot"); + let chunk = read(&mut fresh).await; + assert_eq!(chunk.sequence_number, 7); + assert_eq!( + chunk.snapshot, + Some(common::KvSnapshotChunk { + index: 0, + count: 1, + blocks: 3, + unknown_before: 5, + }) + ); + // A cursor inside the unknown stretch resumes at the first batch + // after it; the gateway settles the jump as a gap of its own. + let mut resumed = lab.subscribe(2).expect("served past the holes"); + for expected in [5, 6, 7] { + assert_eq!(read(&mut resumed).await.sequence_number, expected); + } + lab.publish(8, &batch2).await; + assert_eq!(read(&mut fresh).await.sequence_number, 8); + assert_eq!(read(&mut resumed).await.sequence_number, 8); + } + + /// Without a replay socket a late join cannot be filled: the sequences + /// before the first seen are holes, counted as unknown, and the snapshot + /// a fresh subscriber gets carries the count instead of passing as whole. + #[tokio::test] + async fn a_relay_without_a_replay_socket_marks_the_batches_before_its_start_unknown() { + let mut late = start_lab(100, false).await; + let batch2 = golden::bytes(golden::BATCH2); + for _ in 0..250 { + late.publish(5, &batch2).await; + if late.relay.counts().relayed >= 1 { + break; + } + tokio::time::sleep(Duration::from_millis(20)).await; + } + let counts = late.relay.counts(); + assert_eq!( + ( + counts.relayed, + counts.gap_batches_lost, + counts.unknown_before_start, + counts.publisher_gaps, + ), + (1, 5, 5, 0), + "{counts:?}" + ); + let mut fresh = late.subscribe(0).expect("a snapshot"); + let chunk = read(&mut fresh).await; + assert_eq!((chunk.sequence_number, chunk.dp_rank), (5, Some(1))); + assert_eq!( + chunk.snapshot.as_ref().map(|chunk| chunk.unknown_before), + Some(5) + ); + assert_eq!(stored_hashes(&chunk), vec![(Some(41), vec![42])]); + // Cursors in the unknown stretch are served from the first batch. + let mut resumed = late.subscribe(3).expect("served past the holes"); + assert_eq!(read(&mut resumed).await.sequence_number, 5); + late.publish(6, &batch2).await; + assert_eq!(read(&mut resumed).await.sequence_number, 6); + assert_eq!(read(&mut fresh).await.sequence_number, 6); + assert_eq!(late.relay.counts().out_of_range, 0); + } + + /// A publisher that counts from 1 (the mock engine) or whose first batch + /// is its startup clear (SGLang) has nothing before it to ask for. + #[test] + fn a_first_batch_at_the_counters_start_or_with_a_clear_is_not_a_late_join() { + assert!(!joined_late(0, false)); + assert!(!joined_late(1, false)); + assert!(joined_late(2, false)); + assert!(!joined_late(2, true)); + assert!(joined_late(500, false)); + } + + /// A publisher that published at registration and nothing since: the + /// relay never gets a live batch to notice it by, so it takes the + /// replay's buffer at start, before any live batch, and a subscriber + /// from zero gets those blocks. + #[tokio::test] + async fn the_relay_takes_the_publishers_replay_at_start_before_any_live_batch() { + let mut lab = start_lab_without_replay_socket_yet().await; + let (lab, replay_endpoint) = (&mut lab.0, lab.1); + let batch2 = golden::bytes(golden::BATCH2); + // The engine comes up a moment after the servicer: the first attempt + // found no replay socket, the retry finds it. + tokio::time::sleep(Duration::from_millis(300)).await; + let mut router = RouterSocket::new(); + router.bind(&replay_endpoint).await.expect("replay binds"); + lab.router = Some(router); + // The mock counts from 1: sequences 1..=3 went out at registration. + lab.answer_replay( + 0, + &[ + (1, batch2.clone()), + (2, batch2.clone()), + (3, batch2.clone()), + ], + ) + .await; + lab.wait_relayed(3).await; + let counts = lab.relay.counts(); + assert_eq!( + ( + counts.primed_batches, + counts.unknown_before_start, + counts.gap_batches_lost, + counts.publisher_gaps, + ), + (3, 0, 0, 0), + "{counts:?}" + ); + // Nothing live ever came; the state is served from zero as a snapshot + // cut at 3 (a window from 1 is not complete from 0). + let mut fresh = lab.subscribe(0).expect("a snapshot"); + let chunk = read(&mut fresh).await; + assert_eq!(chunk.sequence_number, 3); + assert_eq!(chunk.snapshot.as_ref().map(|chunk| chunk.blocks), Some(3)); + assert_eq!( + chunk.snapshot.as_ref().map(|chunk| chunk.unknown_before), + Some(0) + ); + // Live continues from the publisher's next sequence. + lab.publish(4, &batch2).await; + assert_eq!(read(&mut fresh).await.sequence_number, 4); + let mut resumed = lab.subscribe(2).expect("inside the window"); + for expected in [3, 4] { + assert_eq!(read(&mut resumed).await.sequence_number, expected); + } + } + + /// The start replay answered empty, then the publisher spoke (the mock's + /// registration batches) without the subscription catching them, and + /// never again: the first subscriber from zero finds the relay empty, + /// which makes it ask the replay once more, and receives the batches. + #[tokio::test] + async fn a_first_subscriber_makes_an_empty_relay_ask_the_replay_again() { + let mut lab = start_lab(100, true).await; + let batch2 = golden::bytes(golden::BATCH2); + let mut stream = lab.subscribe(0).expect("live, nothing held"); + lab.answer_replay( + 0, + &[ + (1, batch2.clone()), + (2, batch2.clone()), + (3, batch2.clone()), + ], + ) + .await; + for expected in [1, 2, 3] { + assert_eq!(read(&mut stream).await.sequence_number, expected); + } + let counts = lab.relay.counts(); + assert_eq!( + ( + counts.primed_batches, + counts.unknown_before_start, + counts.relayed + ), + (3, 0, 3), + "{counts:?}" + ); + lab.publish(4, &batch2).await; + assert_eq!(read(&mut stream).await.sequence_number, 4); + // The next subscriber from zero gets the state the relay now holds + // (a snapshot: the window starts at 1), and asks nothing more. + let mut next = lab.subscribe(0).expect("a snapshot"); + assert_eq!(read(&mut next).await.sequence_number, 4); + assert_eq!(lab.relay.counts().served_snapshots, 1); + } + + /// A replay whose buffer rolled before the relay started marks what it + /// no longer holds as unknown, as the late-join path does. + #[tokio::test] + async fn a_start_replay_that_begins_past_the_publishers_start_marks_the_rest_unknown() { + let mut lab = start_lab_without_replay_socket_yet().await; + let (lab, replay_endpoint) = (&mut lab.0, lab.1); + let batch2 = golden::bytes(golden::BATCH2); + let mut router = RouterSocket::new(); + router.bind(&replay_endpoint).await.expect("replay binds"); + lab.router = Some(router); + lab.answer_replay(0, &[(6, batch2.clone()), (7, batch2.clone())]) + .await; + lab.wait_relayed(2).await; + let counts = lab.relay.counts(); + assert_eq!( + ( + counts.primed_batches, + counts.unknown_before_start, + counts.gap_batches_lost, + ), + (2, 6, 6), + "{counts:?}" + ); + let mut fresh = lab.subscribe(0).expect("a snapshot"); + let chunk = read(&mut fresh).await; + assert_eq!(chunk.sequence_number, 7); + assert_eq!( + chunk.snapshot.as_ref().map(|chunk| chunk.unknown_before), + Some(6) + ); + } + + #[tokio::test] + async fn dropping_the_relay_closes_the_publisher_subscription() { + let mut lab = start_lab(10, false).await; + let mut monitor = lab.publisher.monitor(); + lab.prime().await; + let mut stream = lab.subscribe(0).expect("live"); + assert_eq!(read(&mut stream).await.sequence_number, 0); + drop(lab.relay); + read_end(&mut stream).await; + let disconnected = timeout(Duration::from_secs(5), async { + while let Some(event) = monitor.next().await { + if matches!(event, SocketEvent::Disconnected(_)) { + return true; + } + } + false + }) + .await + .expect("the publisher notices in time"); + assert!(disconnected); + } + /// A publisher batch on rank 0 with one `BlockStored`: `hashes` chained + /// from `parent`, two tokens per block, stamped `ts`. + fn store_payload(ts: f64, hashes: &[i64], parent: Option) -> Vec { + let tokens: Vec = (0..hashes.len() as u32 * 2).collect(); + let batch = serde_json::json!([ts, [{ + "type": "BlockStored", + "block_hashes": hashes, + "parent_block_hash": parent, + "token_ids": tokens, + "block_size": 2, + "lora_id": null, + "medium": "GPU", + }], 0]); + rmp_serde::to_vec_named(&batch).expect("encodes") + } + + fn remove_payload(hashes: &[i64]) -> Vec { + let batch = serde_json::json!([1700000003.0, [{ + "type": "BlockRemoved", + "block_hashes": hashes, + "medium": "GPU", + }], 0]); + rmp_serde::to_vec_named(&batch).expect("encodes") + } + + fn is_cleared(event: &common::KvCacheEvent) -> bool { + matches!(event.data, Some(kv_cache_event::Data::Cleared(_))) + } + + /// `(parent, hashes)` of every stored event of a batch, in order. + fn stored_hashes(batch: &common::KvEventBatch) -> Vec<(Option, Vec)> { + batch + .events + .iter() + .filter_map(|event| match &event.data { + Some(kv_cache_event::Data::Stored(stored)) => Some(( + stored.parent_block_hash, + stored.blocks.iter().map(|block| block.block_hash).collect(), + )), + _ => None, + }) + .collect() + } + + /// Once the window has rolled, a subscriber without a cursor gets the + /// live set as a snapshot cut at the newest sequence, then live events + /// from the next one; a batch published between the subscription and + /// the first poll follows the snapshot, once. + #[tokio::test] + async fn after_the_window_rolled_a_subscriber_without_a_cursor_gets_a_snapshot_then_live() { + let mut lab = start_lab(3, false).await; + lab.prime().await; // sequence 0: BATCH1, whose clear leaves nothing live + lab.publish(1, &store_payload(1.0, &[10, 11], None)).await; + lab.publish(2, &store_payload(2.0, &[12], Some(11))).await; + lab.publish(3, &store_payload(3.0, &[20], None)).await; + lab.publish(4, &remove_payload(&[20])).await; + lab.publish(5, &store_payload(5.0, &[21, 22], None)).await; + lab.wait_relayed(6).await; + let mut stream = lab.subscribe(0).expect("a snapshot"); + // Published before the first poll: must follow the snapshot. + lab.publish(6, &store_payload(6.0, &[30], Some(22))).await; + lab.wait_relayed(7).await; + let chunk = read(&mut stream).await; + assert_eq!(chunk.sequence_number, 5, "stamped at the cut"); + assert_eq!( + chunk.snapshot, + Some(common::KvSnapshotChunk { + index: 0, + count: 1, + blocks: 5, + unknown_before: 0, + }) + ); + assert_eq!(chunk.dp_rank, Some(0)); + assert!(is_cleared(&chunk.events[0])); + assert_eq!( + stored_hashes(&chunk), + vec![(None, vec![10, 11, 12]), (None, vec![21, 22])], + "the live set as the engine stored it, chains merged" + ); + let live = read(&mut stream).await; + assert_eq!((live.sequence_number, live.snapshot), (6, None)); + assert_eq!(stored_hashes(&live), vec![(Some(22), vec![30])]); + lab.publish(7, &store_payload(7.0, &[31], Some(30))).await; + assert_eq!(read(&mut stream).await.sequence_number, 7); + let counts = lab.relay.counts(); + assert_eq!( + (counts.served_snapshots, counts.served_from_history), + (1, 0) + ); + } + + /// A cursor below the window is still refused, and the resubscription + /// from zero the gateway answers with gets the snapshot. + #[tokio::test] + async fn a_stale_cursor_below_the_window_is_refused_and_zero_gets_the_snapshot() { + let mut lab = start_lab(2, false).await; + lab.prime().await; + for seq in 1..=4 { + lab.publish(seq, &store_payload(seq as f64, &[seq as i64 * 10], None)) + .await; + } + lab.wait_relayed(5).await; + // The window is 3..=4; cursor 1 wants 2, gone. + let status = refused(lab.subscribe(1)); + assert_eq!(status.code(), tonic::Code::OutOfRange); + assert!( + status + .message() + .contains("resubscribe from zero for a state snapshot"), + "{}", + status.message() + ); + let mut stream = lab.subscribe(0).expect("a snapshot"); + let chunk = read(&mut stream).await; + assert_eq!(chunk.sequence_number, 4); + assert_eq!(chunk.snapshot.as_ref().map(|chunk| chunk.blocks), Some(4)); + assert_eq!( + stored_hashes(&chunk), + vec![ + (None, vec![10]), + (None, vec![20]), + (None, vec![30]), + (None, vec![40]) + ] + ); + lab.publish(5, &store_payload(5.0, &[50], None)).await; + assert_eq!(read(&mut stream).await.sequence_number, 5); + let counts = lab.relay.counts(); + assert_eq!((counts.out_of_range, counts.served_snapshots), (1, 1)); + } + + /// A relayed batch of `blocks` chained device blocks of 16 tokens on rank 0. + fn synthetic_batch(seq: u64, blocks: i64) -> common::KvEventBatch { + let first = seq as i64 * blocks + 1; + common::KvEventBatch { + sequence_number: seq, + timestamp: 1.0, + events: vec![common::KvCacheEvent { + event_id: seq, + data: Some(kv_cache_event::Data::Stored(common::KvBlocksStored { + blocks: (first..first + blocks) + .map(|hash| common::KvBlock { + block_hash: hash, + token_ids: (0..16).map(|i| hash as u32 ^ i).collect(), + block_size: 16, + ..Default::default() + }) + .collect(), + parent_block_hash: None, + tier: Some(KvCacheTier::Device as i32), + medium: Some("GPU".to_string()), + ..Default::default() + })), + }], + dp_rank: Some(0), + snapshot: None, + load: None, + } + } + + fn unix_now() -> f64 { + SystemTime::now() + .duration_since(UNIX_EPOCH) + .expect("after the epoch") + .as_secs_f64() + } + + /// Publish `sequences` one at a time, each stamped with its publish time + /// and read back on `live` before the next: the relay's latency per + /// batch in milliseconds. + async fn publish_timed( + lab: &mut Lab, + sequences: std::ops::Range, + live: &mut BoxStream, + ) -> Vec { + let mut latencies = Vec::with_capacity((sequences.end - sequences.start) as usize); + for seq in sequences { + lab.publish(seq, &store_payload(unix_now(), &[-(seq as i64)], None)) + .await; + let batch = read(live).await; + assert_eq!(batch.sequence_number, seq); + latencies.push((unix_now() - batch.timestamp) * 1e3); + } + latencies + } + + /// A large worker's pool (676k blocks) behind a rolled window: the + /// snapshot is cut atomically at the relay's cursor while live batches + /// stream through to another subscriber, which sees no stall beyond the + /// one pass under the lock; the snapshot subscriber, not read until the + /// cut is long past, gets every chunk and then the live stream from the + /// cut, nothing twice and nothing missing. + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + #[expect( + clippy::print_stderr, + reason = "the measured pause is this test's report; read it with --nocapture" + )] + async fn a_676k_block_snapshot_is_cut_atomically_and_does_not_stall_the_live_stream() { + const BATCHES: u64 = 10_564; + const PER_BATCH: i64 = 64; + let mut lab = start_lab(50, false).await; + lab.relay + .preload((0..BATCHES).map(|seq| synthetic_batch(seq, PER_BATCH))); + let preloaded = BATCHES * PER_BATCH as u64; + let last = BATCHES - 1; + let mut live = lab.subscribe(last).expect("at the newest sequence"); + // The SUB connects asynchronously: publish the next sequence until it lands. + let next = last + 1; + for _ in 0..250 { + lab.publish(next, &store_payload(unix_now(), &[-1], None)) + .await; + if lab.relay.counts().relayed > BATCHES { + break; + } + tokio::time::sleep(Duration::from_millis(20)).await; + } + assert_eq!(read(&mut live).await.sequence_number, next); + let control = publish_timed(&mut lab, next + 1..next + 201, &mut live).await; + + // The snapshot is taken while live batches keep flowing to `live`. + let relay = Arc::clone(&lab.relay); + let taking = tokio::task::spawn_blocking(move || { + let started = Instant::now(); + let stream = relay.subscribe(common::SubscribeKvEventsRequest { + start_sequence_number: 0, + }); + (stream, started.elapsed()) + }); + let during = publish_timed(&mut lab, next + 201..next + 1_201, &mut live).await; + let (snapshot, subscribe_took) = taking.await.expect("the subscribe call ran"); + let mut snapshot = snapshot.expect("a snapshot"); + let max = |values: &[f64]| values.iter().copied().fold(0.0f64, f64::max); + let median = |values: &[f64]| { + let mut sorted = values.to_vec(); + sorted.sort_by(f64::total_cmp); + sorted[sorted.len() / 2] + }; + eprintln!( + "676k snapshot: subscribe call (lock held for the collection) {subscribe_took:?}; \ + live latency ms: control median {:.3} max {:.3}, during median {:.3} max {:.3}", + median(&control), + max(&control), + median(&during), + max(&during) + ); + assert!( + subscribe_took < Duration::from_secs(2), + "{subscribe_took:?}" + ); + assert!( + max(&during) < 2_000.0, + "a live batch waited {:.1} ms on the snapshot", + max(&during) + ); + + // Not read until now: the chunks come whole, then live from the cut. + let mut chunks = Vec::new(); + let first_live = loop { + let batch = read(&mut snapshot).await; + if batch.snapshot.is_none() { + break batch; + } + chunks.push(batch); + }; + let through = chunks.last().expect("chunks").sequence_number; + assert!( + (next..next + 1_201).contains(&through), + "the cut fell inside the timed publishes: {through}" + ); + let blocks_at_cut = preloaded + (through - last); + let count = u32::try_from(chunks.len()).unwrap(); + assert_eq!( + count as usize, + (blocks_at_cut as usize).div_ceil(crate::kv_state::CHUNK_BLOCKS) + ); + for (index, chunk) in chunks.iter().enumerate() { + assert_eq!( + chunk.sequence_number, + through + 1 - u64::from(count) + index as u64, + "stamps are consecutive up to the cut" + ); + let marker = chunk.snapshot.as_ref().unwrap(); + assert_eq!((marker.index, marker.count), (index as u32, count)); + // Every live batch up to the cut stored one block: the state and + // the cut agree, so the collection and the cursor were one guard. + assert_eq!(marker.blocks, blocks_at_cut); + } + assert!(is_cleared(&chunks[0].events[0])); + let emitted: usize = chunks + .iter() + .map(|chunk| { + stored_hashes(chunk) + .iter() + .map(|(_, hashes)| hashes.len()) + .sum::() + }) + .sum(); + assert_eq!(emitted as u64, blocks_at_cut); + assert_eq!( + first_live.sequence_number, + through + 1, + "live continues right after the cut" + ); + let mut expected = through + 2; + while expected < next + 1_201 { + assert_eq!(read(&mut snapshot).await.sequence_number, expected); + expected += 1; + } + } + + /// The normalizer's counters ride on the relay: readable after every + /// relayed batch and logged with the relay's own counts, so a live run + /// shows what was forwarded, dropped by reason and, with the engine-hash + /// check on, verified. + #[tokio::test] + async fn the_relays_wire_counts_follow_the_normalizer() { + let mut lab = start_lab(8, false).await; + assert_eq!(lab.relay.wire_counts(), WireCounts::default()); + lab.prime().await; + lab.publish(1, &golden::bytes(golden::BATCH2)).await; + lab.wait_relayed(2).await; + let wire = lab.relay.wire_counts(); + assert!( + wire.forwarded_stored + wire.forwarded_removed + wire.forwarded_cleared > 0, + "{wire:?}" + ); + assert_eq!( + wire.hash_checked, 0, + "the check is off unless the environment asks" + ); + assert_eq!(wire.window_only_stores, 0); + } + + /// Before the stream has sent a batch there is no sequence a heartbeat + /// could repeat, so a quiet publisher at the start gets none: a gateway + /// that predates the load field would take the publisher's first batch + /// for a duplicate of a heartbeat numbered 0. The first real batch + /// opens the heartbeats. + #[tokio::test] + async fn no_heartbeat_goes_out_before_the_streams_first_batch() { + let loads = StubLoads::new(2, 0); + let mut lab = start_lab_with_loads_at( + Arc::clone(&loads), + Duration::from_millis(20), + Duration::from_millis(60), + Duration::from_millis(100), + false, + ) + .await; + let mut stream = lab.subscribe(0).expect("live from the start"); + loads.set(5, 1); + assert!( + timeout(Duration::from_millis(300), stream.next()) + .await + .is_err(), + "no batch may precede the publisher's first" + ); + lab.prime().await; + let first = read(&mut stream).await; + assert_eq!(first.sequence_number, 0); + assert!(!load_only_marker(&first)); + let beat = read(&mut stream).await; + assert!(load_only_marker(&beat)); + assert_eq!(beat.sequence_number, 0); } } diff --git a/crates/engine_servicer/src/kv_history.rs b/crates/engine_servicer/src/kv_history.rs new file mode 100644 index 0000000000..ccdd22ee1f --- /dev/null +++ b/crates/engine_servicer/src/kv_history.rs @@ -0,0 +1,352 @@ +//! A bounded, contiguous history of relayed KV-event batches: what a +//! subscriber's `start_sequence_number` is served from, and what a lagging +//! subscriber catches up from. +//! +//! The history covers one unbroken range of the publisher's sequence numbers. +//! A sequence the relay could not obtain (the publisher dropped it and its +//! replay did not have it) occupies its slot as a hole, so the window stays +//! contiguous and a resume past it skips it the way the live stream did. +//! Two caps bound it: a batch count (the engines keep `buffer_steps`, 10,000, +//! for their own replay) and a byte budget over the encoded batches. + +use std::{collections::VecDeque, sync::Arc}; + +use prost::Message; +use smg_grpc_client::common_proto::KvEventBatch; + +/// What a slot of the window holds. +enum Entry { + Batch(Arc), + /// The publisher's sequence passed with no batch to show for it. + Lost, +} + +/// The per-entry bookkeeping charged against the byte budget on top of the +/// encoded batch: the queue slot, the `Arc`, the proto's own allocations. +const ENTRY_OVERHEAD: usize = 96; + +/// Why a cursor cannot be served from the history. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub(crate) enum Window { + /// Nothing relayed yet. + Empty, + /// The batch after the cursor is older than the oldest one kept. + Behind { oldest: u64 }, + /// The cursor is past the newest sequence relayed: it belongs to another + /// publisher incarnation. + Ahead { newest: u64 }, +} + +pub struct History { + /// The sequence number of `entries[0]`. + first: u64, + entries: VecDeque, + bytes: usize, + holes: usize, + max_batches: usize, + max_bytes: usize, +} + +impl History { + pub fn new(max_batches: usize, max_bytes: usize) -> Self { + Self { + first: 0, + entries: VecDeque::new(), + bytes: 0, + holes: 0, + max_batches: max_batches.max(1), + max_bytes, + } + } + + pub fn is_empty(&self) -> bool { + self.entries.is_empty() + } + + pub fn len(&self) -> usize { + self.entries.len() + } + + /// Slots occupied by holes. + pub(crate) fn holes(&self) -> usize { + self.holes + } + + pub fn bytes(&self) -> usize { + self.bytes + } + + pub(crate) fn oldest(&self) -> Option { + (!self.is_empty()).then_some(self.first) + } + + pub(crate) fn newest(&self) -> Option { + self.oldest() + .map(|first| first + self.entries.len() as u64 - 1) + } + + /// The window holds every batch the publisher ever numbered: it starts at + /// 0 and has no holes. A subscriber without a cursor can take it as the + /// publisher's whole state. + pub(crate) fn complete_from_start(&self) -> bool { + !self.entries.is_empty() && self.first == 0 && self.holes == 0 + } + + /// Append the batch for `seq`, which must be the next sequence (any + /// sequence when the history is empty). A non-contiguous append is a + /// caller bug and is ignored. + pub fn push(&mut self, seq: u64, batch: Arc) { + let bytes = batch.encoded_len() + ENTRY_OVERHEAD; + self.append(seq, Entry::Batch(batch), bytes); + } + + /// Record that `seq` passed without a batch. + pub(crate) fn push_lost(&mut self, seq: u64) { + self.append(seq, Entry::Lost, ENTRY_OVERHEAD); + } + + /// Record that `from..=to` passed without batches (`from` the next + /// sequence, any when the history is empty). Only the newest + /// `max_batches` of them could stay in the window, so a longer run + /// starts the window over at the slots that would have survived. + pub(crate) fn push_lost_range(&mut self, from: u64, to: u64) { + if to < from { + return; + } + let mut from = from; + if to - from + 1 > self.max_batches as u64 { + self.clear(); + from = to + 1 - self.max_batches as u64; + } + for seq in from..=to { + self.push_lost(seq); + } + } + + fn append(&mut self, seq: u64, entry: Entry, bytes: usize) { + if let Some(newest) = self.newest() { + if seq != newest + 1 { + return; + } + } else { + self.first = seq; + } + if matches!(entry, Entry::Lost) { + self.holes += 1; + } + self.entries.push_back(entry); + self.bytes += bytes; + self.evict(); + } + + fn evict(&mut self) { + while self.entries.len() > self.max_batches + || (self.bytes > self.max_bytes && self.entries.len() > 1) + { + let Some(entry) = self.entries.pop_front() else { + break; + }; + self.first += 1; + match entry { + Entry::Batch(batch) => { + self.bytes = self + .bytes + .saturating_sub(batch.encoded_len() + ENTRY_OVERHEAD); + } + Entry::Lost => { + self.holes -= 1; + self.bytes = self.bytes.saturating_sub(ENTRY_OVERHEAD); + } + } + } + } + + pub(crate) fn clear(&mut self) { + self.entries.clear(); + self.bytes = 0; + self.holes = 0; + self.first = 0; + } + + /// The batches after `cursor`, in order, holes skipped; or why the window + /// cannot serve that cursor. A cursor equal to the newest sequence yields + /// nothing and is fine. + pub(crate) fn after(&self, cursor: u64) -> Result>, Window> { + let (Some(oldest), Some(newest)) = (self.oldest(), self.newest()) else { + return Err(Window::Empty); + }; + if cursor > newest { + return Err(Window::Ahead { newest }); + } + let want = cursor + 1; + if want < oldest { + return Err(Window::Behind { oldest }); + } + let skip = (want - oldest) as usize; + Ok(self.batches(skip)) + } + + /// Every batch in the window, holes skipped. + pub(crate) fn all(&self) -> Vec> { + self.batches(0) + } + + fn batches(&self, skip: usize) -> Vec> { + self.entries + .iter() + .skip(skip) + .filter_map(|entry| match entry { + Entry::Batch(batch) => Some(Arc::clone(batch)), + Entry::Lost => None, + }) + .collect() + } +} + +#[cfg(test)] +mod tests { + use smg_grpc_client::common_proto::{kv_cache_event, KvCacheCleared, KvCacheEvent}; + + use super::*; + + fn batch(seq: u64, events: usize) -> Arc { + Arc::new(KvEventBatch { + sequence_number: seq, + timestamp: 1.0, + events: (0..events) + .map(|id| KvCacheEvent { + event_id: id as u64, + data: Some(kv_cache_event::Data::Cleared(KvCacheCleared::default())), + }) + .collect(), + dp_rank: None, + snapshot: None, + load: None, + }) + } + + fn seqs(batches: &[Arc]) -> Vec { + batches.iter().map(|b| b.sequence_number).collect() + } + + #[test] + fn serves_cursors_inside_the_window_and_names_why_not_otherwise() { + let mut history = History::new(10, usize::MAX); + assert_eq!(history.after(0), Err(Window::Empty)); + assert!(!history.complete_from_start()); + for seq in 3..=7 { + history.push(seq, batch(seq, 1)); + } + assert_eq!((history.oldest(), history.newest()), (Some(3), Some(7))); + assert_eq!(seqs(&history.after(2).unwrap()), vec![3, 4, 5, 6, 7]); + assert_eq!(seqs(&history.after(5).unwrap()), vec![6, 7]); + assert!( + history.after(7).unwrap().is_empty(), + "nothing after the newest" + ); + assert_eq!(history.after(1), Err(Window::Behind { oldest: 3 })); + assert_eq!(history.after(8), Err(Window::Ahead { newest: 7 })); + assert!( + !history.complete_from_start(), + "the publisher's first batches were never seen" + ); + } + + #[test] + fn evicts_by_count_and_by_bytes_keeping_the_window_contiguous() { + let mut history = History::new(3, usize::MAX); + for seq in 0..=5 { + history.push(seq, batch(seq, 1)); + } + assert_eq!( + (history.oldest(), history.newest(), history.len()), + (Some(3), Some(5), 3) + ); + // Cursor 2 wants 3, the oldest kept; cursor 1 wants 2, gone. + assert_eq!(seqs(&history.after(2).unwrap()), vec![3, 4, 5]); + assert_eq!(history.after(1), Err(Window::Behind { oldest: 3 })); + + // Sequences 1..=5 encode to the same size (sequence 0 is elided). + let per_batch = batch(1, 4).encoded_len() + ENTRY_OVERHEAD; + let mut bounded = History::new(1_000, per_batch * 2 + 1); + for seq in 1..=5 { + bounded.push(seq, batch(seq, 4)); + } + assert_eq!(bounded.len(), 2, "two batches fit the byte budget"); + assert_eq!(bounded.oldest(), Some(4)); + assert!(bounded.bytes() <= per_batch * 2 + 1); + } + + #[test] + fn holes_keep_their_slot_and_are_skipped_on_resume() { + let mut history = History::new(10, usize::MAX); + history.push(0, batch(0, 1)); + history.push_lost(1); + history.push_lost(2); + history.push(3, batch(3, 1)); + assert_eq!(history.holes(), 2); + assert_eq!(history.newest(), Some(3)); + assert_eq!(seqs(&history.after(0).unwrap()), vec![3]); + assert_eq!(seqs(&history.after(1).unwrap()), vec![3]); + assert!( + !history.complete_from_start(), + "a hole means a batch is missing" + ); + // Evicting the holes restores completeness-from-start? No: the start moved. + let mut small = History::new(2, usize::MAX); + small.push(0, batch(0, 1)); + small.push_lost(1); + small.push(2, batch(2, 1)); + assert_eq!(small.oldest(), Some(1)); + assert_eq!(small.holes(), 1); + small.push(3, batch(3, 1)); + assert_eq!((small.oldest(), small.holes()), (Some(2), 0)); + } + + #[test] + fn a_complete_window_from_zero_is_the_publishers_whole_state() { + let mut history = History::new(10, usize::MAX); + history.push(0, batch(0, 1)); + history.push(1, batch(1, 1)); + assert!(history.complete_from_start()); + assert_eq!(seqs(&history.all()), vec![0, 1]); + history.clear(); + assert!(history.is_empty()); + assert_eq!(history.after(0), Err(Window::Empty)); + history.push(0, batch(0, 1)); + assert!( + history.complete_from_start(), + "a new incarnation starts over" + ); + } + + #[test] + fn a_run_of_lost_sequences_keeps_only_what_the_window_would() { + let mut history = History::new(4, usize::MAX); + history.push_lost_range(0, 9); + assert_eq!((history.oldest(), history.newest()), (Some(6), Some(9))); + assert_eq!((history.len(), history.holes()), (4, 4)); + history.push(10, batch(10, 1)); + // The batch evicts the oldest hole: the window is 7..=10. + assert_eq!(history.after(5), Err(Window::Behind { oldest: 7 })); + assert_eq!(seqs(&history.after(6).unwrap()), vec![10]); + assert!(!history.complete_from_start()); + // A short run after batches stays contiguous with them. + let mut short = History::new(10, usize::MAX); + short.push(0, batch(0, 1)); + short.push_lost_range(1, 3); + short.push(4, batch(4, 1)); + assert_eq!((history.len() + short.len(), short.holes()), (9, 3)); + assert_eq!(seqs(&short.after(0).unwrap()), vec![4]); + short.push_lost_range(5, 4); + assert_eq!(short.newest(), Some(4), "an empty run is nothing"); + } + + #[test] + fn a_non_contiguous_append_is_ignored() { + let mut history = History::new(10, usize::MAX); + history.push(0, batch(0, 1)); + history.push(5, batch(5, 1)); + assert_eq!(history.newest(), Some(0)); + } +} diff --git a/crates/engine_servicer/src/kv_state.rs b/crates/engine_servicer/src/kv_state.rs new file mode 100644 index 0000000000..4a99bce80d --- /dev/null +++ b/crates/engine_servicer/src/kv_state.rs @@ -0,0 +1,1022 @@ +//! The relay's record of the engine's live blocks, and the state snapshot it +//! serves a subscriber from that record. +//! +//! [`LiveState`] is a function of the normalized stream the relay forwards, +//! kept per data-parallel rank, tier and engine hash: a `Stored` adds one +//! physical copy of each of its blocks together with what emitting the block +//! again takes (parent, tokens, LoRA id, cache level, extra keys, and the +//! event's tier, medium, group, locality, ownership, session and namespace); +//! a `Removed` takes one copy away and drops the block at zero; an +//! `AllBlocksCleared` empties the rank. Copies are capped per block the way +//! the gateway caps them ([`COPIES_CAP`]), so a snapshot rebuilds the copy +//! counts the gateway would hold after the same stream. Memory is one entry +//! per live `(rank, tier, hash)`, about 230 bytes at block size 16; nothing +//! is kept per relayed batch. +//! +//! A snapshot ([`LiveState::snapshot`]) is one pass over the entries that +//! clones the per-block `Arc`s. The relay takes it under its lock together +//! with the sequence the state is current through, so the cut is atomic: +//! every batch up to that sequence is in it and every later one reaches the +//! subscriber through the live channel afterwards. Everything else, the +//! parent-first order, the run merging and the batch protos, happens outside +//! the lock in [`SnapshotChunks`] as the subscriber's stream is polled, so a +//! slow snapshot reader never holds up the publisher task or any other +//! subscriber. Each rank keeps its entries in a dense slot vector in store +//! order under a hash index, so the pass is a sequential scan whose `Arc` +//! clones follow the records' allocation order: tens of milliseconds for a +//! pool of several hundred thousand blocks in a release build, which is the +//! worst live latency another subscriber sees while a snapshot is taken; +//! the ordering pass and the chunk encoding that follow run off the lock. +//! +//! The chunks of one snapshot carry consecutive sequence numbers ending at +//! the sequence the state was taken at, so the gateway's per-rank cursor +//! continues into the live stream with no gap and no duplicate; chunk 0 +//! begins with an `AllBlocksCleared`, and every chunk is marked with +//! [`KvSnapshotChunk`](common::KvSnapshotChunk). The chunk size grows from +//! [`CHUNK_BLOCKS`] only when the stamping needs fewer chunks than the +//! default would make (a state far larger than the number of sequences +//! behind it), and then never beyond twice the largest store any single +//! relayed batch could have carried. + +use std::{ + collections::{BTreeMap, HashMap}, + sync::Arc, +}; + +use smg_grpc_client::common_proto::{self as common, kv_cache_event, KvCacheTier}; + +/// The most physical copies of one block counted per tier: the gateway's +/// `COPIES_CAP` in `kv_event_monitor.rs`, mirrored so a snapshot rebuilds +/// the counts the gateway would hold. +pub const COPIES_CAP: u32 = 8; + +/// Blocks per snapshot chunk unless the stamping rule needs larger chunks. +pub const CHUNK_BLOCKS: usize = 2_048; + +/// What a `Stored` event says about all of its blocks: shared by the blocks +/// of one event, and by consecutive events that repeat it. +#[derive(Clone, Debug, Default, PartialEq, Eq)] +pub(crate) struct StoredTail { + pub tier: Option, + pub medium: Option, + pub group_idx: Option, + pub kv_cache_spec_kind: Option, + pub kv_cache_spec_sliding_window: Option, + pub locality: Option, + pub ownership: Option, + pub session_id: Option, + pub lora_name: Option, + pub cache_salt: Option, +} + +impl StoredTail { + fn of(stored: &common::KvBlocksStored) -> Self { + Self { + tier: stored.tier, + medium: stored.medium.clone(), + group_idx: stored.group_idx, + kv_cache_spec_kind: stored.kv_cache_spec_kind.clone(), + kv_cache_spec_sliding_window: stored.kv_cache_spec_sliding_window, + locality: stored.locality, + ownership: stored.ownership.clone(), + session_id: stored.session_id.clone(), + lora_name: stored.lora_name.clone(), + cache_salt: stored.cache_salt.clone(), + } + } + + /// Whether `stored` carries exactly this tail (without allocating one). + fn matches(&self, stored: &common::KvBlocksStored) -> bool { + self.tier == stored.tier + && self.medium == stored.medium + && self.group_idx == stored.group_idx + && self.kv_cache_spec_kind == stored.kv_cache_spec_kind + && self.kv_cache_spec_sliding_window == stored.kv_cache_spec_sliding_window + && self.locality == stored.locality + && self.ownership == stored.ownership + && self.session_id == stored.session_id + && self.lora_name == stored.lora_name + && self.cache_salt == stored.cache_salt + } +} + +/// One live block, as the relay can emit it again. +#[derive(Debug)] +pub(crate) struct LiveBlock { + pub hash: i64, + /// The hash the block chains from: the previous block of its store, or + /// the store's `parent_block_hash` for the first. + pub parent: Option, + /// The tier it is live on (a `KvCacheTier` value). + pub tier: i32, + pub tokens: Box<[u32]>, + pub block_size: i32, + pub lora_id: Option, + pub cache_level: Option, + pub extra_keys: Box<[common::KvBlockExtraKey]>, + pub tail: Arc, + /// Store order across the relay's lifetime; a snapshot keeps it. + pub order: u64, +} + +struct Entry { + block: Arc, + copies: u32, +} + +/// One rank's live blocks: a dense slot vector in store order (freed slots +/// are reused) under a `(tier, hash)` index, so a snapshot is one +/// sequential scan whose `Arc` clones follow the records' allocation order +/// instead of a hash map's. +#[derive(Default)] +struct RankBlocks { + slots: Vec>, + index: HashMap<(i32, i64), u32>, + free: Vec, + live: usize, +} + +impl RankBlocks { + fn len(&self) -> usize { + self.live + } + + fn is_empty(&self) -> bool { + self.live == 0 + } + + fn get(&self, key: (i32, i64)) -> Option<&Entry> { + let slot = *self.index.get(&key)?; + self.slots[slot as usize].as_ref() + } + + fn get_mut(&mut self, key: (i32, i64)) -> Option<&mut Entry> { + let slot = *self.index.get(&key)?; + self.slots[slot as usize].as_mut() + } + + /// Store a block the rank does not hold on that tier. + fn insert(&mut self, key: (i32, i64), entry: Entry) { + let slot = match self.free.pop() { + Some(slot) => { + self.slots[slot as usize] = Some(entry); + slot + } + None => { + self.slots.push(Some(entry)); + u32::try_from(self.slots.len() - 1).unwrap_or(u32::MAX) + } + }; + self.index.insert(key, slot); + self.live += 1; + } + + fn remove(&mut self, key: (i32, i64)) -> Option { + let slot = self.index.remove(&key)?; + let entry = self.slots[slot as usize].take(); + if entry.is_some() { + self.free.push(slot); + self.live -= 1; + } + entry + } + + fn iter(&self) -> impl Iterator { + self.slots.iter().flatten() + } +} + +/// What the state has done since the relay started. +#[derive(Clone, Debug, Default, PartialEq, Eq)] +pub(crate) struct StateCounts { + /// Physical copies added by stores. + pub stored: u64, + /// Copies taken away by removals. + pub removed: u64, + /// Removals of a block the state did not hold on that tier. + pub removed_unknown: u64, + /// Stores of a block already at [`COPIES_CAP`] copies on its tier. + pub capped: u64, + /// Ranks emptied by `AllBlocksCleared`. + pub cleared: u64, +} + +/// The engine's live blocks per rank, as the relayed stream describes them. +#[derive(Default)] +pub struct LiveState { + ranks: BTreeMap, RankBlocks>, + /// The last store's tail, shared with the next store that repeats it. + last_tail: Option>, + order: u64, + /// Live physical copies over all ranks and tiers. + blocks: u64, + counts: StateCounts, +} + +/// The tier a store or removal names, as the gateway reads it: `tier` when +/// set, else the blocks' `cache_level` (absent means the device). +fn tier_key(tier: Option, cache_level: Option) -> i32 { + match tier { + Some(tier) if tier != KvCacheTier::Unspecified as i32 => tier, + _ => match cache_level.unwrap_or(0) { + 0 => KvCacheTier::Device as i32, + 1 => KvCacheTier::Host as i32, + 2 => KvCacheTier::Disk as i32, + _ => KvCacheTier::External as i32, + }, + } +} + +impl LiveState { + pub fn new() -> Self { + Self::default() + } + + /// Fold one normalized batch into the state. + pub fn apply(&mut self, batch: &common::KvEventBatch) { + for event in &batch.events { + match &event.data { + Some(kv_cache_event::Data::Stored(stored)) => self.store(batch.dp_rank, stored), + Some(kv_cache_event::Data::Removed(removed)) => { + self.remove(batch.dp_rank, removed); + } + Some(kv_cache_event::Data::Cleared(_)) => self.clear_rank(batch.dp_rank), + None => {} + } + } + } + + fn tail_for(&mut self, stored: &common::KvBlocksStored) -> Arc { + if let Some(tail) = &self.last_tail { + if tail.matches(stored) { + return Arc::clone(tail); + } + } + let tail = Arc::new(StoredTail::of(stored)); + self.last_tail = Some(Arc::clone(&tail)); + tail + } + + fn store(&mut self, rank: Option, stored: &common::KvBlocksStored) { + if stored.blocks.is_empty() { + return; + } + let tier = tier_key( + stored.tier, + stored.blocks.first().and_then(|block| block.cache_level), + ); + let tail = self.tail_for(stored); + let blocks = self.ranks.entry(rank).or_default(); + let mut parent = stored.parent_block_hash; + for block in &stored.blocks { + let key = (tier, block.block_hash); + if let Some(entry) = blocks.get_mut(key) { + if entry.copies < COPIES_CAP { + entry.copies += 1; + self.blocks += 1; + } else { + self.counts.capped += 1; + } + } else { + self.order += 1; + blocks.insert( + key, + Entry { + block: Arc::new(LiveBlock { + hash: block.block_hash, + parent, + tier, + tokens: block.token_ids.clone().into_boxed_slice(), + block_size: block.block_size, + lora_id: block.lora_id, + cache_level: block.cache_level, + extra_keys: block.extra_keys.clone().into_boxed_slice(), + tail: Arc::clone(&tail), + order: self.order, + }), + copies: 1, + }, + ); + self.blocks += 1; + } + self.counts.stored += 1; + parent = Some(block.block_hash); + } + } + + fn remove(&mut self, rank: Option, removed: &common::KvBlocksRemoved) { + let tier = tier_key(removed.tier, removed.cache_level); + let Some(blocks) = self.ranks.get_mut(&rank) else { + self.counts.removed_unknown += removed.block_hashes.len() as u64; + return; + }; + for &hash in &removed.block_hashes { + let key = (tier, hash); + let Some(entry) = blocks.get_mut(key) else { + self.counts.removed_unknown += 1; + continue; + }; + entry.copies -= 1; + self.blocks -= 1; + self.counts.removed += 1; + if entry.copies == 0 { + blocks.remove(key); + } + } + } + + fn clear_rank(&mut self, rank: Option) { + if let Some(blocks) = self.ranks.remove(&rank) { + let copies: u64 = blocks.iter().map(|entry| u64::from(entry.copies)).sum(); + self.blocks -= copies; + } + self.counts.cleared += 1; + } + + /// Forget everything: the publisher started over. + pub fn clear(&mut self) { + self.ranks.clear(); + self.last_tail = None; + self.blocks = 0; + } + + /// Live `(rank, tier, hash)` entries. + pub fn entries(&self) -> usize { + self.ranks.values().map(RankBlocks::len).sum() + } + + /// Live physical copies. + pub fn blocks(&self) -> u64 { + self.blocks + } + + /// Ranks holding at least one block. + pub fn ranks(&self) -> usize { + self.ranks.values().filter(|rank| !rank.is_empty()).count() + } + + pub(crate) fn counts(&self) -> &StateCounts { + &self.counts + } + + /// Live copies of `hash` on `tier` for `rank`. + pub fn copies(&self, rank: Option, tier: KvCacheTier, hash: i64) -> u32 { + self.ranks + .get(&rank) + .and_then(|blocks| blocks.get((tier as i32, hash))) + .map_or(0, |entry| entry.copies) + } + + /// The live set as it stands: the per-block records shared, nothing + /// copied. One pass over the entries; the caller orders and encodes it + /// outside the lock with [`SnapshotChunks::new`]. + pub fn snapshot(&self) -> Snapshot { + let ranks = self + .ranks + .iter() + .filter(|(_, blocks)| !blocks.is_empty()) + .map(|(rank, blocks)| { + let mut entries = Vec::with_capacity(blocks.len()); + entries.extend(blocks.iter().map(|entry| SnapshotEntry { + block: Arc::clone(&entry.block), + copies: entry.copies, + })); + (*rank, entries) + }) + .collect(); + Snapshot { + ranks, + blocks: self.blocks, + } + } +} + +/// One live block of a snapshot and how many copies the engine holds. +pub(crate) struct SnapshotEntry { + pub(crate) block: Arc, + pub(crate) copies: u32, +} + +/// The live set at one point of the stream, unordered. +pub struct Snapshot { + /// Per rank, in rank order; ranks without blocks are left out. + pub(crate) ranks: Vec<(Option, Vec)>, + /// Physical copies over all ranks. + pub(crate) blocks: u64, +} + +impl Snapshot { + pub fn entries(&self) -> usize { + self.ranks.iter().map(|(_, entries)| entries.len()).sum() + } +} + +/// A snapshot ordered and cut into the batches a subscriber receives, built +/// lazily chunk by chunk. +pub struct SnapshotChunks { + /// Per rank, one unit per physical copy in emission order: parents + /// before children, store order otherwise, copies adjacent. + ranks: Vec<(Option, Vec>)>, + chunk_blocks: usize, + count: u32, + next: u32, + /// `(rank, offset)` of the next chunk's first unit. + position: (usize, usize), + through: u64, + timestamp: f64, + blocks: u64, + unknown_before: u64, +} + +impl SnapshotChunks { + /// Order `snapshot`, taken at sequence `through`, and size its chunks so + /// that they can be stamped `through - count + 1 ..= through`; `timestamp` + /// is what the chunks carry (the gateway reads it as the publish time), + /// `unknown_before` how many publisher sequences before the relay's + /// record the snapshot cannot cover (0 for a record from the start). + pub fn new(snapshot: Snapshot, through: u64, timestamp: f64, unknown_before: u64) -> Self { + let blocks = snapshot.blocks; + let ranks: Vec<(Option, Vec>)> = snapshot + .ranks + .into_iter() + .map(|(rank, entries)| (rank, order_rank(entries))) + .collect(); + let largest = ranks + .iter() + .map(|(_, units)| units.len()) + .max() + .unwrap_or(0); + let chunks_at = |chunk_blocks: usize| -> usize { + ranks + .iter() + .map(|(_, units)| units.len().div_ceil(chunk_blocks)) + .sum::() + .max(1) + }; + // Every rank in the state came from at least one relayed batch, so + // one chunk per rank always fits under `through + 1`; the loop only + // runs when the default chunk would need more sequences than exist. + let mut chunk_blocks = CHUNK_BLOCKS; + while chunks_at(chunk_blocks) as u64 > through + 1 && chunk_blocks < largest { + chunk_blocks *= 2; + } + let count = u32::try_from(chunks_at(chunk_blocks)).unwrap_or(u32::MAX); + Self { + ranks, + chunk_blocks, + count, + next: 0, + position: (0, 0), + through, + timestamp, + blocks, + unknown_before, + } + } + + /// Chunks in the snapshot (at least one: the clear). + pub fn chunk_count(&self) -> u32 { + self.count + } + + /// Physical copies in the snapshot. + pub fn blocks(&self) -> u64 { + self.blocks + } + + /// The sequence the state was taken at: the last chunk's stamp. + pub fn through(&self) -> u64 { + self.through + } + + /// The next chunk, or `None` once all `count` have been produced. + pub(crate) fn next_chunk(&mut self) -> Option { + if self.next >= self.count { + return None; + } + let index = self.next; + self.next += 1; + let sequence_number = self + .through + .saturating_sub(u64::from(self.count) - 1) + .saturating_add(u64::from(index)); + let mut events = Vec::new(); + if index == 0 { + events.push(common::KvCacheEvent { + event_id: 0, + data: Some(kv_cache_event::Data::Cleared( + common::KvCacheCleared::default(), + )), + }); + } + let mut dp_rank = None; + while self.position.0 < self.ranks.len() { + let (rank_index, offset) = self.position; + let (rank, units) = &self.ranks[rank_index]; + if offset >= units.len() { + self.position = (rank_index + 1, 0); + continue; + } + let end = (offset + self.chunk_blocks).min(units.len()); + self.position = if end == units.len() { + (rank_index + 1, 0) + } else { + (rank_index, end) + }; + dp_rank = *rank; + stored_events(&units[offset..end], &mut events); + break; + } + Some(common::KvEventBatch { + sequence_number, + timestamp: self.timestamp, + events, + dp_rank, + snapshot: Some(common::KvSnapshotChunk { + index, + count: self.count, + blocks: self.blocks, + unknown_before: self.unknown_before, + }), + load: None, + }) + } +} + +impl Iterator for SnapshotChunks { + type Item = common::KvEventBatch; + + fn next(&mut self) -> Option { + self.next_chunk() + } +} + +/// One rank's entries in emission order, one unit per physical copy: store +/// order, with the live ancestors of a block hoisted ahead of it when they +/// were stored later (a parent evicted and stored again under a surviving +/// child), so the gateway's index derives every position from a parent it +/// already holds. +fn order_rank(mut entries: Vec) -> Vec> { + entries.sort_unstable_by_key(|entry| entry.block.order); + let mut by_hash: HashMap> = HashMap::with_capacity(entries.len()); + for (index, entry) in entries.iter().enumerate() { + by_hash.entry(entry.block.hash).or_default().push(index); + } + let units: usize = entries.iter().map(|entry| entry.copies as usize).sum(); + let mut out = Vec::with_capacity(units); + let mut done = vec![false; entries.len()]; + let mut path: Vec = Vec::new(); + let push_units = |index: usize, out: &mut Vec>| { + let entry = &entries[index]; + for _ in 0..entry.copies { + out.push(Arc::clone(&entry.block)); + } + }; + for index in 0..entries.len() { + if done[index] { + continue; + } + path.clear(); + let mut cursor = entries[index].block.parent; + while let Some(hash) = cursor { + let Some(indices) = by_hash.get(&hash) else { + // Not live: the gateway starts a chain there, as it did + // when the parent left the engine before its child. + break; + }; + if done[indices[0]] { + break; + } + // Every tier's entry of a hash is marked and emitted together. + for &sibling in indices { + done[sibling] = true; + } + path.push(hash); + cursor = entries[indices[0]].block.parent; + } + for hash in path.iter().rev() { + for &ancestor in &by_hash[hash] { + push_units(ancestor, &mut out); + } + } + for &sibling in &by_hash[&entries[index].block.hash] { + if !done[sibling] { + done[sibling] = true; + push_units(sibling, &mut out); + } + } + } + out +} + +/// Whether unit `at` is one of several copies of the same block (copies are +/// adjacent in the emission order). +fn is_copy(units: &[Arc], at: usize) -> bool { + (at > 0 && Arc::ptr_eq(&units[at - 1], &units[at])) + || (at + 1 < units.len() && Arc::ptr_eq(&units[at], &units[at + 1])) +} + +/// `Stored` events for `units`: a run of blocks chaining one from the other +/// under one tail and tier becomes one event, as the engine stored them; +/// each extra copy of a block is its own single-block event so the gateway +/// counts it. +fn stored_events(units: &[Arc], events: &mut Vec) { + let mut start = 0; + while start < units.len() { + let mut end = start + 1; + if !is_copy(units, start) { + while end < units.len() + && !is_copy(units, end) + && units[end].parent == Some(units[end - 1].hash) + && units[end].tier == units[end - 1].tier + && Arc::ptr_eq(&units[end].tail, &units[end - 1].tail) + { + end += 1; + } + } + events.push(stored_event(&units[start..end])); + start = end; + } +} + +fn stored_event(run: &[Arc]) -> common::KvCacheEvent { + let tail = &run[0].tail; + common::KvCacheEvent { + event_id: 0, + data: Some(kv_cache_event::Data::Stored(common::KvBlocksStored { + blocks: run + .iter() + .map(|block| common::KvBlock { + block_hash: block.hash, + token_ids: block.tokens.to_vec(), + block_size: block.block_size, + lora_id: block.lora_id, + cache_level: block.cache_level, + extra_keys: block.extra_keys.to_vec(), + }) + .collect(), + parent_block_hash: run[0].parent, + tier: tail.tier, + medium: tail.medium.clone(), + group_idx: tail.group_idx, + kv_cache_spec_kind: tail.kv_cache_spec_kind.clone(), + kv_cache_spec_sliding_window: tail.kv_cache_spec_sliding_window, + locality: tail.locality, + ownership: tail.ownership.clone(), + session_id: tail.session_id.clone(), + lora_name: tail.lora_name.clone(), + cache_salt: tail.cache_salt.clone(), + })), + } +} + +#[cfg(test)] +mod tests { + use std::time::Instant; + + use smg_grpc_client::common_proto::{KvBlock, KvBlocksRemoved, KvBlocksStored, KvCacheEvent}; + + use super::*; + + fn block(hash: i64) -> KvBlock { + KvBlock { + block_hash: hash, + token_ids: (0..4u32) + .map(|i| (hash as u32).wrapping_mul(16).wrapping_add(i)) + .collect(), + block_size: 4, + ..Default::default() + } + } + + fn stored(parent: Option, hashes: &[i64], tier: KvCacheTier) -> KvCacheEvent { + KvCacheEvent { + event_id: 1, + data: Some(kv_cache_event::Data::Stored(KvBlocksStored { + blocks: hashes.iter().map(|&hash| block(hash)).collect(), + parent_block_hash: parent, + tier: Some(tier as i32), + medium: Some("GPU".to_string()), + group_idx: Some(0), + kv_cache_spec_kind: Some("full_attention".to_string()), + ..Default::default() + })), + } + } + + fn removed(hashes: &[i64], tier: KvCacheTier) -> KvCacheEvent { + KvCacheEvent { + event_id: 1, + data: Some(kv_cache_event::Data::Removed(KvBlocksRemoved { + block_hashes: hashes.to_vec(), + tier: Some(tier as i32), + ..Default::default() + })), + } + } + + fn cleared() -> KvCacheEvent { + KvCacheEvent { + event_id: 1, + data: Some(kv_cache_event::Data::Cleared( + common::KvCacheCleared::default(), + )), + } + } + + fn batch(seq: u64, rank: Option, events: Vec) -> common::KvEventBatch { + common::KvEventBatch { + sequence_number: seq, + timestamp: 1.0, + events, + dp_rank: rank, + snapshot: None, + load: None, + } + } + + /// `(rank, tier, hash, parent)` per block of every stored event, in + /// emission order. + fn emitted(chunks: &[common::KvEventBatch]) -> Vec<(Option, i32, i64, Option)> { + let mut out = Vec::new(); + for chunk in chunks { + for event in &chunk.events { + if let Some(kv_cache_event::Data::Stored(stored)) = &event.data { + let mut parent = stored.parent_block_hash; + for block in &stored.blocks { + out.push(( + chunk.dp_rank, + stored.tier.unwrap(), + block.block_hash, + parent, + )); + parent = Some(block.block_hash); + } + } + } + } + out + } + + fn chunks_of(state: &LiveState, through: u64) -> Vec { + SnapshotChunks::new(state.snapshot(), through, 2.0, 0).collect() + } + + #[test] + fn stores_add_copies_removals_take_them_and_a_clear_empties_the_rank() { + let device = KvCacheTier::Device; + let mut state = LiveState::new(); + state.apply(&batch(1, Some(0), vec![stored(None, &[1, 2], device)])); + state.apply(&batch(2, Some(0), vec![stored(Some(2), &[3], device)])); + state.apply(&batch(3, Some(1), vec![stored(None, &[1], device)])); + assert_eq!((state.entries(), state.blocks(), state.ranks()), (4, 4, 2)); + // A second physical copy of 2 (vLLM's resend) counts without a new entry. + state.apply(&batch(4, Some(0), vec![stored(Some(1), &[2], device)])); + assert_eq!((state.entries(), state.blocks()), (4, 5)); + assert_eq!(state.copies(Some(0), device, 2), 2); + // One removal per copy; the block stays until the last one. + state.apply(&batch(5, Some(0), vec![removed(&[2], device)])); + assert_eq!(state.copies(Some(0), device, 2), 1); + state.apply(&batch(6, Some(0), vec![removed(&[2, 99], device)])); + assert_eq!(state.copies(Some(0), device, 2), 0); + assert_eq!(state.counts().removed_unknown, 1); + assert_eq!((state.entries(), state.blocks()), (3, 3)); + // The host tier is a copy of its own. + state.apply(&batch( + 7, + Some(0), + vec![stored(None, &[1], KvCacheTier::Host)], + )); + assert_eq!(state.copies(Some(0), KvCacheTier::Host, 1), 1); + assert_eq!(state.copies(Some(0), device, 1), 1); + state.apply(&batch(8, Some(0), vec![removed(&[1], device)])); + assert_eq!(state.copies(Some(0), device, 1), 0); + assert_eq!(state.copies(Some(0), KvCacheTier::Host, 1), 1); + // Clearing rank 0 leaves rank 1 alone. + state.apply(&batch(9, Some(0), vec![cleared()])); + assert_eq!((state.entries(), state.blocks(), state.ranks()), (1, 1, 1)); + assert_eq!(state.copies(Some(1), device, 1), 1); + assert_eq!(state.counts().cleared, 1); + state.clear(); + assert_eq!((state.entries(), state.blocks(), state.ranks()), (0, 0, 0)); + } + + #[test] + fn copies_are_capped_like_the_gateway_counts_them() { + let mut state = LiveState::new(); + for seq in 0..20 { + state.apply(&batch( + seq, + None, + vec![stored(None, &[7], KvCacheTier::Device)], + )); + } + assert_eq!(state.copies(None, KvCacheTier::Device, 7), COPIES_CAP); + assert_eq!(state.blocks(), u64::from(COPIES_CAP)); + assert_eq!(state.counts().capped, 20 - u64::from(COPIES_CAP)); + let chunks = chunks_of(&state, 19); + assert_eq!(emitted(&chunks).len(), COPIES_CAP as usize); + } + + #[test] + fn a_snapshot_lists_the_live_set_parents_first_with_the_store_fields() { + let device = KvCacheTier::Device; + let mut state = LiveState::new(); + // A child stored before its (re-stored) parent must still follow it. + state.apply(&batch(1, Some(0), vec![stored(None, &[1, 2, 3], device)])); + state.apply(&batch(2, Some(0), vec![removed(&[1], device)])); + state.apply(&batch(3, Some(0), vec![stored(Some(3), &[4], device)])); + state.apply(&batch(4, Some(0), vec![stored(None, &[1], device)])); + state.apply(&batch(5, Some(0), vec![stored(Some(4), &[5], device)])); + let chunks = chunks_of(&state, 5); + assert_eq!(chunks.len(), 1); + let chunk = &chunks[0]; + assert_eq!(chunk.sequence_number, 5); + assert_eq!(chunk.dp_rank, Some(0)); + assert_eq!( + chunk.snapshot, + Some(common::KvSnapshotChunk { + index: 0, + count: 1, + blocks: 5, + unknown_before: 0, + }) + ); + assert!( + matches!(chunk.events[0].data, Some(kv_cache_event::Data::Cleared(_))), + "the clear comes first" + ); + let order = emitted(&chunks); + assert_eq!(order.len(), 5); + let position = |hash: i64| order.iter().position(|entry| entry.2 == hash).unwrap(); + for (_, _, hash, parent) in &order { + if let Some(parent) = parent { + assert!( + position(*parent) < position(*hash), + "{parent} before {hash}: {order:?}" + ); + } + } + // Block 1's chain was broken by the removal: 2 keeps naming 1 as its + // parent, 1 is live again, so 1 comes first and the run 2, 3 follows. + assert_eq!(order[0].2, 1); + assert_eq!( + &order[1..3], + &[(Some(0), 1, 2, Some(1)), (Some(0), 1, 3, Some(2))] + ); + let stores: Vec<&KvBlocksStored> = chunk + .events + .iter() + .filter_map(|event| match &event.data { + Some(kv_cache_event::Data::Stored(stored)) => Some(stored), + _ => None, + }) + .collect(); + // Every block chains from the one emitted before it under one tail: + // the five come as one run, the way an engine stores a prompt. + assert_eq!(stores.len(), 1, "{stores:#?}"); + let first = &stores[0]; + assert_eq!(first.parent_block_hash, None); + assert_eq!(first.blocks.len(), 5); + assert_eq!(first.medium.as_deref(), Some("GPU")); + assert_eq!(first.kv_cache_spec_kind.as_deref(), Some("full_attention")); + assert_eq!(first.group_idx, Some(0)); + assert_eq!(first.blocks[0].token_ids, block(1).token_ids); + assert_eq!(first.blocks[0].block_size, 4); + } + + #[test] + fn chunks_are_cut_per_rank_and_stamped_up_to_the_cut() { + let device = KvCacheTier::Device; + let mut state = LiveState::new(); + let hashes: Vec = (1..=CHUNK_BLOCKS as i64 + 10).collect(); + state.apply(&batch(1, Some(0), vec![stored(None, &hashes, device)])); + state.apply(&batch(2, Some(1), vec![stored(None, &[1, 2], device)])); + let chunks = chunks_of(&state, 40); + assert_eq!(chunks.len(), 3, "rank 0 takes two chunks, rank 1 one"); + assert_eq!( + chunks + .iter() + .map(|chunk| (chunk.sequence_number, chunk.dp_rank)) + .collect::>(), + vec![(38, Some(0)), (39, Some(0)), (40, Some(1))] + ); + for (index, chunk) in chunks.iter().enumerate() { + let marker = chunk.snapshot.as_ref().unwrap(); + assert_eq!((marker.index as usize, marker.count), (index, 3)); + assert_eq!(marker.blocks, CHUNK_BLOCKS as u64 + 12); + assert_eq!( + matches!(chunk.events[0].data, Some(kv_cache_event::Data::Cleared(_))), + index == 0 + ); + } + // The split run continues from the first chunk's last block. + let second = emitted(&chunks[1..2]); + assert_eq!( + second[0], + ( + Some(0), + 1, + CHUNK_BLOCKS as i64 + 1, + Some(CHUNK_BLOCKS as i64) + ) + ); + assert_eq!(emitted(&chunks).len(), CHUNK_BLOCKS + 12); + } + + #[test] + fn few_sequences_behind_a_large_state_mean_larger_chunks() { + let mut state = LiveState::new(); + let hashes: Vec = (1..=4 * CHUNK_BLOCKS as i64).collect(); + state.apply(&batch( + 0, + None, + vec![stored(None, &hashes, KvCacheTier::Device)], + )); + state.apply(&batch( + 1, + None, + vec![stored(None, &[-1], KvCacheTier::Device)], + )); + // Two sequences exist (0 and 1): at most two chunks, stamped 0 and 1. + let chunks = chunks_of(&state, 1); + assert_eq!(chunks.len(), 2); + assert_eq!( + chunks + .iter() + .map(|chunk| chunk.sequence_number) + .collect::>(), + vec![0, 1] + ); + assert_eq!(emitted(&chunks).len(), 4 * CHUNK_BLOCKS + 1); + } + + #[test] + fn an_empty_state_snapshots_as_one_clear() { + let state = LiveState::new(); + let chunks = chunks_of(&state, 12); + assert_eq!(chunks.len(), 1); + assert_eq!(chunks[0].sequence_number, 12); + assert_eq!(chunks[0].events.len(), 1); + assert_eq!( + chunks[0].snapshot, + Some(common::KvSnapshotChunk { + index: 0, + count: 1, + blocks: 0, + unknown_before: 0, + }) + ); + assert_eq!(chunks[0].dp_rank, None); + // A record that starts late says so on every chunk. + let late: Vec = + SnapshotChunks::new(state.snapshot(), 12, 2.0, 7).collect(); + assert_eq!( + late[0].snapshot.as_ref().map(|chunk| chunk.unknown_before), + Some(7) + ); + } + + #[test] + #[expect( + clippy::print_stderr, + reason = "the measured pass is this test's report; read it with --nocapture" + )] + fn a_snapshot_of_a_large_state_is_one_brief_pass() { + // A 676k-block pool (a large worker's) in chains of 64. + let mut state = LiveState::new(); + let mut hash = 1i64; + for seq in 0..(676_144 / 64) { + let hashes: Vec = (hash..hash + 64).collect(); + hash += 64; + state.apply(&batch( + seq, + Some(0), + vec![stored(None, &hashes, KvCacheTier::Device)], + )); + } + assert_eq!(state.blocks(), 676_096); + let started = Instant::now(); + let snapshot = state.snapshot(); + let collected = started.elapsed(); + let started = Instant::now(); + let mut chunks = SnapshotChunks::new(snapshot, 20_000, 3.0, 0); + let ordered = started.elapsed(); + let started = Instant::now(); + let mut blocks = 0; + let count = chunks.chunk_count(); + while let Some(chunk) = chunks.next_chunk() { + blocks += emitted(std::slice::from_ref(&chunk)).len(); + } + let encoded = started.elapsed(); + eprintln!( + "676k snapshot: collect {collected:?}, order {ordered:?}, {count} chunks built in {encoded:?}" + ); + assert_eq!(blocks, 676_096); + assert_eq!(count as usize, 676_096_usize.div_ceil(CHUNK_BLOCKS)); + assert!( + collected.as_millis() < 1_000, + "collecting under the lock took {collected:?}" + ); + } +} diff --git a/crates/engine_servicer/src/kv_wire.rs b/crates/engine_servicer/src/kv_wire.rs new file mode 100644 index 0000000000..b91ab57ad5 --- /dev/null +++ b/crates/engine_servicer/src/kv_wire.rs @@ -0,0 +1,2276 @@ +//! The engines' KV-cache event wire format and its normalization into the +//! proto the gateway indexes. +//! +//! Both vLLM (`vllm/distributed/kv_events.py`) and SGLang +//! (`sglang/srt/disaggregation/kv_events.py`) publish one msgspec batch per +//! scheduler step: a positional array `[ts, events, dp_rank]` whose events +//! are tagged maps (`{"type": "BlockStored", ...}`, defaulted fields +//! omitted). Older publishers emit events as positional arrays with the tag +//! first; both layouts decode here. +//! +//! Normalization applies, per stream, what a per-worker prefix index needs +//! and nothing else: +//! - only local blocks (`locality` absent or `LOCAL`); remote ones describe a +//! shared pool and are dropped; +//! - only blocks the engine manages itself (`ownership` other than a +//! residency agent); +//! - the storage tier from `medium`, unknown media dropped; +//! - the main-attention KV cache group where a rank has one (hybrid models +//! publish one group per attention kind under the same hashes, and the +//! sliding-window and state-space groups' events, stores and removals, are +//! dropped); a rank that publishes only sliding-window groups forwards +//! their stores, the hashes aligned to the tail of the tokens, as its only +//! signal (see [`Normalizer`]); +//! - whole blocks only (hash count × block size == token count, or the +//! tail of a longer span), no self-referencing hash chains, no offload +//! placeholders (a chunk key with no tokens); +//! - speculative-decoding bigram pages folded to their tokens, so an Eagle +//! engine's blocks hash like a plain engine's; +//! - the cache namespace (LoRA name, cache salt) carried on every store and +//! inherited down the parent chain when a child store omits it. +//! +//! Stores and removals are forwarded one for one. vLLM keeps up to two +//! physical copies of one hash and removes them one at a time, so a removal +//! can arrive while another copy is still cached: every copy's store and +//! every removal go through, and the gateway counts copies per worker, tier +//! and hash (capped), as the relay's own live-block record and the hash +//! check's parent memory do. +//! +//! Hash identity: an integer hash is used as is (vLLM sends the low 64 bits +//! of the digest as an unsigned integer, SGLang the high 64 bits as a signed +//! one; both are 64-bit patterns the proto carries as `int64`); a raw digest +//! folds to its last eight bytes big-endian, the integer vLLM would have sent +//! for it. + +use std::{ + collections::{HashMap, HashSet}, + fmt, +}; + +use serde::{ + de::{self, IgnoredAny, MapAccess, SeqAccess, Visitor}, + Deserialize, Deserializer, +}; +use smg_grpc_client::common_proto::{ + self as common, kv_block_extra_key, kv_cache_event, KvCacheLocality, KvCacheTier, +}; +use tracing::{debug, warn}; + +use crate::{ + engine_hash::{self, Digest32, EngineHash, VllmExtraKey}, + kv_state::COPIES_CAP, +}; + +/// `int.from_bytes(bytes, "big")` kept to 64 bits: the whole value for the +/// publisher's eight-byte sequence frame, the low 64 bits of a longer hash. +pub(crate) fn low64_big_endian(bytes: &[u8]) -> u64 { + bytes + .iter() + .fold(0, |value, &byte| (value << 8) | u64::from(byte)) +} + +/// A publisher batch: msgspec `array_like`, `[ts, events, dp_rank]`, the +/// rank named `data_parallel_rank` by vLLM and `attn_dp_rank` by SGLang and +/// omittable by both; later fields are tolerated by the caller's codec. +#[derive(Deserialize)] +pub struct WireBatch { + pub ts: f64, + pub events: Vec, + #[serde(default)] + pub dp_rank: Option, +} + +/// A block hash as the proto's signed 64-bit identity: sha256 bytes keep +/// their low 64 bits read big-endian; an int is already 64 bits wide. +#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)] +pub struct BlockHash(pub i64); + +impl<'de> Deserialize<'de> for BlockHash { + fn deserialize>(deserializer: D) -> Result { + struct HashVisitor; + + impl Visitor<'_> for HashVisitor { + type Value = BlockHash; + + fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter.write_str("a block hash as bytes or an integer") + } + + fn visit_bytes(self, bytes: &[u8]) -> Result { + Ok(BlockHash(low64_big_endian(bytes) as i64)) + } + + fn visit_u64(self, value: u64) -> Result { + Ok(BlockHash(value as i64)) + } + + fn visit_i64(self, value: i64) -> Result { + Ok(BlockHash(value)) + } + } + + deserializer.deserialize_any(HashVisitor) + } +} + +/// The fields a store and a removal share beyond hashes and tokens. +#[derive(Clone, Debug, Default, PartialEq, Eq)] +pub struct EventTail { + pub medium: Option, + pub group_idx: Option, + pub kv_cache_spec_kind: Option, + pub kv_cache_spec_sliding_window: Option, + pub locality: Option, + pub ownership: Option, + pub session_id: Option, +} + +/// Token ids of a store: plain ids, or the (token, next token) bigrams a +/// speculative-decoding (Eagle) publisher emits. +#[derive(Clone, Debug, PartialEq, Eq)] +pub enum WireTokens { + Ids(Vec), + Bigrams(Vec<(u32, u32)>), +} + +/// One entry of vLLM's untagged per-block `extra_keys`. +#[derive(Clone, Debug, PartialEq, Eq)] +pub enum ExtraKey { + Text(String), + Number(i64), + Blob(Vec), + Multimodal { + identifier: String, + offset: i64, + }, + /// An item of a shape this relay does not model; kept as a marker so the + /// per-block key count stays truthful. + Opaque, +} + +#[derive(Clone, Debug, PartialEq)] +pub struct WireStored { + pub block_hashes: Vec, + pub parent_block_hash: Option, + pub token_ids: WireTokens, + pub block_size: i64, + pub lora_id: Option, + pub lora_name: Option, + pub cache_salt: Option, + /// One entry per block, `None` for a block without extra keys. + pub extra_keys: Option>>>, + pub tail: EventTail, +} + +#[derive(Clone, Debug, PartialEq)] +pub struct WireRemoved { + pub block_hashes: Vec, + pub tail: EventTail, +} + +#[derive(Clone, Debug, PartialEq)] +pub enum WireEvent { + BlockStored(WireStored), + BlockRemoved(WireRemoved), + AllBlocksCleared { + ownership: Option, + }, + /// An event type this relay does not convert (one a newer engine added): + /// skipped on its own so the batch's other events still go through. + Unknown, + /// A known event whose named field is missing or of a shape this relay + /// cannot read: skipped on its own, the rest of the batch still goes + /// through. + Malformed(&'static str), +} + +// --------------------------------------------------------------------------- +// Decoding: tagged maps and tag-first arrays +// --------------------------------------------------------------------------- + +/// A loosely typed scalar for the optional tail slots of the array layout and +/// for `extra_keys` items. +#[derive(Clone, Debug, PartialEq, Eq)] +enum Loose { + Nil, + Unsigned(u64), + Signed(i64), + Text(String), + Bytes(Vec), + Seq(Vec), + Other, +} + +impl<'de> Deserialize<'de> for Loose { + fn deserialize>(deserializer: D) -> Result { + struct LooseVisitor; + + impl<'de> Visitor<'de> for LooseVisitor { + type Value = Loose; + + fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter.write_str("a scalar, bytes, string, array or nil") + } + + fn visit_unit(self) -> Result { + Ok(Loose::Nil) + } + + fn visit_none(self) -> Result { + Ok(Loose::Nil) + } + + fn visit_some>( + self, + deserializer: D2, + ) -> Result { + Loose::deserialize(deserializer) + } + + fn visit_bool(self, value: bool) -> Result { + Ok(Loose::Unsigned(u64::from(value))) + } + + fn visit_u64(self, value: u64) -> Result { + Ok(Loose::Unsigned(value)) + } + + fn visit_i64(self, value: i64) -> Result { + Ok(if value >= 0 { + Loose::Unsigned(value as u64) + } else { + Loose::Signed(value) + }) + } + + fn visit_f64(self, _value: f64) -> Result { + Ok(Loose::Other) + } + + fn visit_str(self, value: &str) -> Result { + Ok(Loose::Text(value.to_owned())) + } + + fn visit_string(self, value: String) -> Result { + Ok(Loose::Text(value)) + } + + fn visit_bytes(self, value: &[u8]) -> Result { + Ok(Loose::Bytes(value.to_vec())) + } + + fn visit_byte_buf(self, value: Vec) -> Result { + Ok(Loose::Bytes(value)) + } + + fn visit_seq>(self, mut seq: A) -> Result { + let mut items = Vec::new(); + while let Some(item) = seq.next_element::()? { + items.push(item); + } + Ok(Loose::Seq(items)) + } + + fn visit_map>(self, mut map: A) -> Result { + while map.next_entry::()?.is_some() {} + Ok(Loose::Other) + } + } + + deserializer.deserialize_any(LooseVisitor) + } +} + +impl Loose { + fn text(self) -> Option { + match self { + Loose::Text(text) => Some(text), + _ => None, + } + } + + fn unsigned32(self) -> Option { + match self { + Loose::Unsigned(value) => u32::try_from(value).ok(), + _ => None, + } + } + + fn signed(self) -> Option { + match self { + Loose::Unsigned(value) => i64::try_from(value).ok(), + Loose::Signed(value) => Some(value), + _ => None, + } + } + + fn block_hash(&self) -> Option { + match self { + Loose::Unsigned(value) => Some(BlockHash(*value as i64)), + Loose::Signed(value) => Some(BlockHash(*value)), + Loose::Bytes(bytes) => Some(BlockHash(low64_big_endian(bytes) as i64)), + _ => None, + } + } + + fn block_hashes(self) -> Option> { + match self { + Loose::Seq(items) => items.iter().map(Loose::block_hash).collect(), + _ => None, + } + } + + fn tokens(self) -> Option { + let Loose::Seq(items) = self else { + return None; + }; + let mut ids = Vec::with_capacity(items.len()); + let mut pairs = Vec::new(); + for item in items { + match item { + Loose::Unsigned(value) => ids.push(u32::try_from(value).ok()?), + Loose::Seq(pair) => match pair.as_slice() { + [Loose::Unsigned(first), Loose::Unsigned(second)] => { + pairs.push((u32::try_from(*first).ok()?, u32::try_from(*second).ok()?)); + } + _ => return None, + }, + _ => return None, + } + } + if pairs.is_empty() { + Some(WireTokens::Ids(ids)) + } else if ids.is_empty() { + Some(WireTokens::Bigrams(pairs)) + } else { + None + } + } + + fn extra_key(self) -> ExtraKey { + match self { + Loose::Text(text) => ExtraKey::Text(text), + Loose::Unsigned(value) => { + i64::try_from(value).map_or(ExtraKey::Opaque, ExtraKey::Number) + } + Loose::Signed(value) => ExtraKey::Number(value), + Loose::Bytes(bytes) => ExtraKey::Blob(bytes), + Loose::Seq(items) => match items.as_slice() { + [Loose::Text(identifier), Loose::Unsigned(offset)] => ExtraKey::Multimodal { + identifier: identifier.clone(), + offset: i64::try_from(*offset).unwrap_or(i64::MAX), + }, + [Loose::Text(identifier), Loose::Signed(offset)] => ExtraKey::Multimodal { + identifier: identifier.clone(), + offset: *offset, + }, + _ => ExtraKey::Opaque, + }, + Loose::Nil | Loose::Other => ExtraKey::Opaque, + } + } + + /// `extra_keys`: one list (or nil) per block. + fn extra_keys(self) -> Option>>> { + match self { + Loose::Seq(per_block) => Some( + per_block + .into_iter() + .map(|keys| match keys { + Loose::Seq(items) => { + Some(items.into_iter().map(Loose::extra_key).collect()) + } + _ => None, + }) + .collect(), + ), + _ => None, + } + } + + /// The array layout's seventh store slot: vLLM's `lora_name`. + fn namespace_slot(self) -> (Option, Option) { + match self { + Loose::Text(name) => (Some(name), None), + _ => (None, None), + } + } +} + +/// Everything a map-layout event may carry, collected before dispatch on +/// `type` so key order does not matter. +#[derive(Default)] +struct Fields { + event_type: Option, + block_hashes: Option, + parent_block_hash: Option, + token_ids: Option, + block_size: Option, + lora_id: Option, + lora_name: Option, + cache_salt: Option, + extra_keys: Option, + tail: [Option; 7], +} + +const TAIL_KEYS: [&str; 7] = [ + "medium", + "group_idx", + "kv_cache_spec_kind", + "kv_cache_spec_sliding_window", + "locality", + "ownership", + "session_id", +]; + +fn tail_from(slots: [Option; 7]) -> EventTail { + let [medium, group_idx, kind, sliding, locality, ownership, session_id] = slots; + EventTail { + medium: medium.and_then(Loose::text), + group_idx: group_idx.and_then(Loose::unsigned32), + kv_cache_spec_kind: kind.and_then(Loose::text), + kv_cache_spec_sliding_window: sliding.and_then(Loose::unsigned32), + locality: locality.and_then(Loose::text), + ownership: ownership.and_then(Loose::text), + session_id: session_id.and_then(Loose::text), + } +} + +impl Fields { + fn into_event(self) -> WireEvent { + let Some(event_type) = self.event_type else { + return WireEvent::Malformed("type"); + }; + match event_type.as_str() { + "BlockStored" => { + let Some(block_hashes) = self.block_hashes.and_then(Loose::block_hashes) else { + return WireEvent::Malformed("block_hashes"); + }; + let Some(token_ids) = self.token_ids.and_then(Loose::tokens) else { + return WireEvent::Malformed("token_ids"); + }; + let Some(block_size) = self.block_size.and_then(Loose::signed) else { + return WireEvent::Malformed("block_size"); + }; + WireEvent::BlockStored(WireStored { + block_hashes, + parent_block_hash: self.parent_block_hash.and_then(|hash| hash.block_hash()), + token_ids, + block_size, + lora_id: self.lora_id.and_then(Loose::signed), + lora_name: self.lora_name.and_then(Loose::text), + cache_salt: self.cache_salt.and_then(Loose::text), + extra_keys: self.extra_keys.and_then(Loose::extra_keys), + tail: tail_from(self.tail), + }) + } + "BlockRemoved" => match self.block_hashes.and_then(Loose::block_hashes) { + Some(block_hashes) => WireEvent::BlockRemoved(WireRemoved { + block_hashes, + tail: tail_from(self.tail), + }), + None => WireEvent::Malformed("block_hashes"), + }, + "AllBlocksCleared" => { + let [_, _, _, _, _, ownership, _] = self.tail; + WireEvent::AllBlocksCleared { + ownership: ownership.and_then(Loose::text), + } + } + _ => WireEvent::Unknown, + } + } +} + +impl<'de> Deserialize<'de> for WireEvent { + fn deserialize>(deserializer: D) -> Result { + struct EventVisitor; + + impl<'de> Visitor<'de> for EventVisitor { + type Value = WireEvent; + + fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter.write_str("a KV cache event as a tagged map or a tag-first array") + } + + fn visit_map>(self, mut map: A) -> Result { + let mut fields = Fields::default(); + while let Some(key) = map.next_key::()? { + match key.as_str() { + "type" => fields.event_type = map.next_value::()?.text(), + "block_hashes" => fields.block_hashes = Some(map.next_value()?), + "parent_block_hash" => fields.parent_block_hash = Some(map.next_value()?), + "token_ids" => fields.token_ids = Some(map.next_value()?), + "block_size" => fields.block_size = Some(map.next_value()?), + "lora_id" => fields.lora_id = Some(map.next_value()?), + "lora_name" => fields.lora_name = Some(map.next_value()?), + "cache_salt" => fields.cache_salt = Some(map.next_value()?), + "extra_keys" => fields.extra_keys = Some(map.next_value()?), + other => match TAIL_KEYS.iter().position(|name| *name == other) { + Some(slot) => fields.tail[slot] = Some(map.next_value()?), + None => { + map.next_value::()?; + } + }, + } + } + Ok(fields.into_event()) + } + + /// The tag-first array layout, slots in the order vLLM's msgspec + /// structs declare their fields: a store is `[tag, block_hashes, + /// parent, token_ids, block_size, lora_id, medium, lora_name, + /// extra_keys, group_idx, kind, sliding_window, locality, + /// ownership, session_id]`, a removal `[tag, block_hashes, + /// medium, group_idx, locality, ownership]`, a clear `[tag]`; + /// trailing defaults are omitted. SGLang's legacy arrays put + /// `cache_salt` where vLLM has `lora_name`; the two are not + /// distinguishable there, and SGLang has published maps since it + /// added the field. + fn visit_seq>(self, mut seq: A) -> Result { + let tag: String = seq + .next_element()? + .ok_or_else(|| de::Error::invalid_length(0, &"an event tag"))?; + let mut slots = Vec::new(); + while let Some(slot) = seq.next_element::()? { + slots.push(slot); + } + let mut slots = slots.into_iter(); + let mut next = || slots.next().unwrap_or(Loose::Nil); + let mut fields = Fields { + event_type: Some(tag.clone()), + ..Fields::default() + }; + match tag.as_str() { + "BlockStored" => { + fields.block_hashes = Some(next()); + fields.parent_block_hash = Some(next()); + fields.token_ids = Some(next()); + fields.block_size = Some(next()); + fields.lora_id = Some(next()); + fields.tail[0] = Some(next()); + let (lora_name, cache_salt) = next().namespace_slot(); + fields.lora_name = lora_name.map(Loose::Text); + fields.cache_salt = cache_salt.map(Loose::Text); + fields.extra_keys = Some(next()); + for slot in 1..=6 { + fields.tail[slot] = Some(next()); + } + } + "BlockRemoved" => { + fields.block_hashes = Some(next()); + for slot in [0, 1, 4, 5] { + fields.tail[slot] = Some(next()); + } + } + "AllBlocksCleared" => { + fields.tail[5] = Some(next()); + } + _ => {} + } + Ok(fields.into_event()) + } + } + + deserializer.deserialize_any(EventVisitor) + } +} + +// --------------------------------------------------------------------------- +// Normalization +// --------------------------------------------------------------------------- + +/// Why the relay dropped an event instead of forwarding it. +#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)] +pub enum DropReason { + UnknownType, + UnsupportedOwnership, + NonLocalLocality, + UnknownMedium, + NonMainAttentionGroup, + UnalignedBlocks, + SelfReferencingHashes, + /// A store with nothing to index: an offload chunk placeholder (a chunk + /// key with no tokens) or an empty hash list. + Placeholder, + /// A known event with a field missing or unreadable. + Malformed, +} + +impl DropReason { + pub fn as_str(self) -> &'static str { + match self { + Self::UnknownType => "unknown_type", + Self::UnsupportedOwnership => "unsupported_ownership", + Self::NonLocalLocality => "non_local_locality", + Self::UnknownMedium => "unknown_medium", + Self::NonMainAttentionGroup => "non_main_attention_group", + Self::UnalignedBlocks => "unaligned_blocks", + Self::SelfReferencingHashes => "self_referencing_hashes", + Self::Placeholder => "placeholder", + Self::Malformed => "malformed", + } + } +} + +/// What one stream has forwarded and dropped, by reason. +#[derive(Clone, Debug, Default, PartialEq, Eq)] +pub struct Counts { + pub forwarded_stored: u64, + pub forwarded_removed: u64, + pub forwarded_cleared: u64, + /// Forwarded stores whose every hash this stream had already seen on that + /// rank and tier: vLLM's second physical copy, or a replayed batch. + pub duplicate_stores: u64, + /// Forwarded stores whose tokens arrived as speculative-decoding bigrams. + pub bigram_stores: u64, + /// Blocks rehashed the worker's way under the opt-in engine-hash check. + pub hash_checked: u64, + /// Checked blocks whose published hash differs from the recomputed one. + pub hash_mismatch: u64, + /// Blocks the check could not rehash: an unknown parent, a LoRA block + /// (the adapter path is not in the event), or a key shape it does not + /// model. + pub hash_unverifiable: u64, + /// Forwarded stores of a sliding-window group on a rank that showed no + /// main-attention group (a window-only model): the rank's only signal, + /// see the group policy on [`Normalizer`]. + pub window_only_stores: u64, + /// Forwarded stores that named fewer blocks than their tokens spanned, + /// the hashes aligned to the tail of the tokens. + pub tail_aligned_stores: u64, + pub dropped: HashMap, +} + +impl Counts { + pub fn dropped(&self, reason: DropReason) -> u64 { + self.dropped.get(&reason).copied().unwrap_or(0) + } +} + +/// The cache namespace block hashes were computed under. +#[derive(Clone, Debug, Default, PartialEq, Eq)] +pub struct Namespace { + pub lora_name: Option, + pub cache_salt: Option, +} + +impl Namespace { + fn is_empty(&self) -> bool { + self.lora_name.is_none() && self.cache_salt.is_none() + } +} + +/// What this stream remembers about a stored engine hash, kept until the +/// last of its physical copies is removed: the engines keep up to two +/// copies of a block and remove them one at a time, and the gateway counts +/// them the same way, so a child stored after one copy went still has its +/// parent here. +#[derive(Clone, Debug, Default)] +struct BlockRecord { + /// The namespace it was stored under, for children that omit theirs. + namespace: Option, + /// Physical copies stored and not yet removed, capped as the gateway + /// caps them. + copies: u32, + /// Its recomputed full digest when the hash check ran: what a child + /// chains on. + digest: Option, +} + +#[derive(Default)] +struct RankState { + /// Per tier: the engine hashes this stream has seen stored and not yet + /// removed, with what it remembers about each. + tiers: HashMap>, + /// KV cache groups seen on stores: whether each is a main-attention group. + groups: HashMap, + /// Hashes forwarded from a sliding-window group while the rank had no + /// main-attention group, so their removals still go through once it has. + window_forwarded: HashSet, +} + +impl RankState { + fn main_seen(&self) -> bool { + self.groups.values().any(|&main| main) + } +} + +/// What the shared gates let through, and under which group policy. +struct Admitted { + tier: KvCacheTier, + locality: KvCacheLocality, + /// A sliding-window group's event on a rank without a main-attention + /// group, forwarded as the rank's only signal. + window_only: bool, +} + +/// Per-stream normalization state (one engine endpoint, all its DP ranks). +/// +/// KV-cache groups: an event names its group (`group_idx`) and, on stores, +/// the group's attention kind. A model with full attention alongside sliding +/// windows (vLLM's hybrid KV-cache manager) publishes every block twice under +/// the same hash, once per group: the full-attention group holds the whole +/// prefix, which is what a prefix hit needs, while the window group names +/// only the window's blocks against the whole computed span and removes them +/// as the window moves on. So on a rank with a main-attention group every +/// non-main event, store and removal alike, is dropped and counted +/// ([`DropReason::NonMainAttentionGroup`]); forwarding the window group's +/// removals would take the full group's copies out of the gateway's count. +/// A rank that shows only sliding-window groups (a window-only model) would +/// otherwise give the gateway nothing: its window stores are forwarded with +/// their group identity, the hashes aligned to the tail of the tokens, and +/// counted ([`Counts::window_only_stores`]). +#[derive(Default)] +pub struct Normalizer { + ranks: HashMap, + counts: Counts, + /// The engine hash to recompute per store, when verification is on. + hash_check: Option, +} + +/// The environment variable that turns the engine-hash check on for every +/// relay in the process: `sglang` or `vllm-sha256-cbor`. +pub const HASH_CHECK_ENV: &str = "SMG_KV_EVENT_HASH_CHECK"; + +const MAIN_ATTENTION_KINDS: [&str; 3] = ["full_attention", "mla_attention", "sink_full_attention"]; + +/// The tier an engine medium names: the device when the medium is absent, +/// `None` for a medium this relay does not know. +pub fn tier_of(medium: Option<&str>) -> Option { + let Some(medium) = medium else { + return Some(KvCacheTier::Device); + }; + let upper = medium.to_ascii_uppercase(); + match upper.as_str() { + "GPU" | "DEVICE" => Some(KvCacheTier::Device), + "CPU" | "CPU_PINNED" | "CPU_TIER1" => Some(KvCacheTier::Host), + "CPU_TIER2" | "DISK" | "NVME" | "STORAGE" => Some(KvCacheTier::Disk), + "EXTERNAL" | "NETWORK" | "REMOTE" | "SHARED" => Some(KvCacheTier::External), + _ => None, + } +} + +/// `KvBlock.cache_level` for a tier: `None` on the device (older consumers +/// read an absent level as the device), the tier's rank otherwise. +pub(crate) fn cache_level_of(tier: KvCacheTier) -> Option { + match tier { + KvCacheTier::Unspecified | KvCacheTier::Device => None, + KvCacheTier::Host => Some(1), + KvCacheTier::Disk => Some(2), + KvCacheTier::External => Some(3), + } +} + +fn locality_of(locality: Option<&str>) -> Result { + match locality.map(str::to_ascii_uppercase).as_deref() { + None | Some("LOCAL") => Ok(KvCacheLocality::Local), + Some("REMOTE") => Err(()), + Some(_) => Err(()), + } +} + +fn is_residency_agent(ownership: Option<&str>) -> bool { + ownership.is_some_and(|owner| owner.eq_ignore_ascii_case("kvcr")) +} + +fn extra_key_proto(key: ExtraKey) -> Option { + let key = match key { + ExtraKey::Text(text) => kv_block_extra_key::Key::Text(text), + ExtraKey::Number(number) => kv_block_extra_key::Key::Number(number), + ExtraKey::Blob(blob) => kv_block_extra_key::Key::Blob(blob), + ExtraKey::Multimodal { identifier, offset } => { + kv_block_extra_key::Key::Multimodal(common::KvMultimodalKey { identifier, offset }) + } + ExtraKey::Opaque => return None, + }; + Some(common::KvBlockExtraKey { key: Some(key) }) +} + +/// vLLM's cache salt rides inside `extra_keys`: the first text item of the +/// first block's keys that is not the LoRA name. +fn salt_from_extra_keys( + extra_keys: Option<&[Option>]>, + lora_name: Option<&str>, +) -> Option { + let first = extra_keys?.first()?.as_ref()?; + first.iter().find_map(|key| match key { + ExtraKey::Text(text) if Some(text.as_str()) != lora_name && !text.is_empty() => { + Some(text.clone()) + } + _ => None, + }) +} + +impl Normalizer { + pub fn new() -> Self { + Self::default() + } + + /// A normalizer that also rehashes every verifiable store with `check` + /// and counts mismatches; nothing is dropped for it. See + /// [`crate::engine_hash`]. + pub fn with_hash_check(check: EngineHash) -> Self { + Self { + hash_check: Some(check), + ..Self::default() + } + } + + /// [`Self::new`], or [`Self::with_hash_check`] when [`HASH_CHECK_ENV`] + /// names an algorithm; an unknown value is logged and the check stays off. + pub fn from_env() -> Self { + match std::env::var(HASH_CHECK_ENV) { + Ok(value) if !value.trim().is_empty() => match EngineHash::parse(&value) { + Some(check) => Self::with_hash_check(check), + None => { + warn!(%value, "{HASH_CHECK_ENV} names no known engine hash; check off"); + Self::new() + } + }, + _ => Self::new(), + } + } + + pub fn counts(&self) -> &Counts { + &self.counts + } + + pub fn hash_check(&self) -> Option { + self.hash_check + } + + fn drop(&mut self, reason: DropReason, event_id: u64) -> DropReason { + let count = self.counts.dropped.entry(reason).or_insert(0); + *count += 1; + if *count <= 3 { + debug!(event_id, reason = reason.as_str(), "KV event not forwarded"); + } + reason + } + + /// A whole publisher batch as its proto, `event_id` advancing once per + /// event whether or not it is forwarded, so ids stay monotonic. + pub fn normalize_batch( + &mut self, + batch: WireBatch, + sequence_number: u64, + event_id: &mut u64, + ) -> common::KvEventBatch { + // Learn the batch's groups before normalizing any of its events: the + // engine lists a sliding-window group's store before the + // full-attention group's in the same step, and what to do with the + // window group depends on whether the rank has a main group at all. + let rank = batch.dp_rank.unwrap_or(-1); + for event in &batch.events { + if let WireEvent::BlockStored(stored) = event { + if let (Some(group), Some(kind)) = ( + stored.tail.group_idx, + stored.tail.kv_cache_spec_kind.as_deref(), + ) { + self.ranks + .entry(rank) + .or_default() + .groups + .insert(group, MAIN_ATTENTION_KINDS.contains(&kind)); + } + } + } + let mut events = Vec::with_capacity(batch.events.len()); + for event in batch.events { + *event_id += 1; + if let Ok(converted) = self.normalize(event, batch.dp_rank, *event_id) { + events.push(converted); + } + } + common::KvEventBatch { + sequence_number, + timestamp: batch.ts, + events, + dp_rank: batch.dp_rank, + snapshot: None, + load: None, + } + } + + /// One event as the proto the gateway should index, or why not. + pub fn normalize( + &mut self, + event: WireEvent, + dp_rank: Option, + event_id: u64, + ) -> Result { + let rank = dp_rank.unwrap_or(-1); + let data = match event { + WireEvent::Unknown => return Err(self.drop(DropReason::UnknownType, event_id)), + WireEvent::Malformed(field) => { + debug!(event_id, field, "KV event field unreadable"); + return Err(self.drop(DropReason::Malformed, event_id)); + } + WireEvent::AllBlocksCleared { ownership } => { + if is_residency_agent(ownership.as_deref()) { + return Err(self.drop(DropReason::UnsupportedOwnership, event_id)); + } + self.ranks.remove(&rank); + self.counts.forwarded_cleared += 1; + kv_cache_event::Data::Cleared(common::KvCacheCleared { ownership }) + } + WireEvent::BlockStored(stored) => self.normalize_stored(stored, rank, event_id)?, + WireEvent::BlockRemoved(removed) => self.normalize_removed(removed, rank, event_id)?, + }; + Ok(common::KvCacheEvent { + event_id, + data: Some(data), + }) + } + + /// The shared gates: ownership, locality, medium, cache group. `hashes` + /// are the event's, for the one non-main event a rank with a main group + /// still forwards: the removal of blocks it forwarded while window-only. + fn admit( + &mut self, + tail: &EventTail, + rank: i32, + event_id: u64, + learn_group: bool, + hashes: &[BlockHash], + ) -> Result { + if is_residency_agent(tail.ownership.as_deref()) { + return Err(self.drop(DropReason::UnsupportedOwnership, event_id)); + } + let locality = locality_of(tail.locality.as_deref()) + .map_err(|()| self.drop(DropReason::NonLocalLocality, event_id))?; + let tier = tier_of(tail.medium.as_deref()) + .ok_or_else(|| self.drop(DropReason::UnknownMedium, event_id))?; + let mut window_only = false; + if let Some(group) = tail.group_idx { + let state = self.ranks.entry(rank).or_default(); + let main = match tail.kv_cache_spec_kind.as_deref() { + Some(kind) => { + let main = MAIN_ATTENTION_KINDS.contains(&kind); + if learn_group { + state.groups.insert(group, main); + } + main + } + // A kind-less event on a group we learned follows the group; + // an unknown group is treated as main, as legacy publishers + // with a single group are. + None => state.groups.get(&group).copied().unwrap_or(true), + }; + if !main { + let forwarded_before = !learn_group + && hashes + .iter() + .any(|hash| state.window_forwarded.contains(&hash.0)); + if state.main_seen() && !forwarded_before { + return Err(self.drop(DropReason::NonMainAttentionGroup, event_id)); + } + window_only = true; + } + } + Ok(Admitted { + tier, + locality, + window_only, + }) + } + + fn normalize_stored( + &mut self, + stored: WireStored, + rank: i32, + event_id: u64, + ) -> Result { + let Admitted { + tier, + locality, + window_only, + } = self.admit(&stored.tail, rank, event_id, true, &stored.block_hashes)?; + // A bigram page lists (token, next token) per position; its tokens + // are the first elements and the page grid is unchanged. The pairs + // stay around for the engine-hash check, which hashes both words. + let (mut token_ids, mut bigram_words): (Vec, Option>) = match stored.token_ids + { + WireTokens::Ids(ids) => (ids, None), + WireTokens::Bigrams(pairs) => ( + pairs.iter().map(|&(token, _)| token).collect(), + Some( + pairs + .iter() + .flat_map(|&(token, next)| [token, next]) + .collect(), + ), + ), + }; + let bigrams = bigram_words.is_some(); + if stored.block_hashes.is_empty() || token_ids.is_empty() { + return Err(self.drop(DropReason::Placeholder, event_id)); + } + // The hashes name block_size tokens each, normally the whole span. A + // sliding-window group's store spans the whole range the step + // computed while naming only the window's blocks, the newest ones, + // so fewer hashes than blocks take the tail of the tokens, never the + // head; a span that is not whole blocks is unreadable. + let width = usize::try_from(stored.block_size) + .ok() + .filter(|&width| width > 0 && i32::try_from(width).is_ok()); + let Some(width) = width else { + return Err(self.drop(DropReason::UnalignedBlocks, event_id)); + }; + let span = token_ids.len(); + let tail_aligned = match stored.block_hashes.len().checked_mul(width) { + Some(need) if need == span => false, + Some(need) if need < span && span.is_multiple_of(width) => { + let start = span - need; + token_ids = token_ids.split_off(start); + bigram_words = bigram_words.map(|mut words| words.split_off(start * 2)); + true + } + _ => return Err(self.drop(DropReason::UnalignedBlocks, event_id)), + }; + { + let mut seen = HashSet::with_capacity(stored.block_hashes.len() + 1); + if let Some(parent) = stored.parent_block_hash { + seen.insert(parent.0); + } + if stored.block_hashes.iter().any(|hash| !seen.insert(hash.0)) { + return Err(self.drop(DropReason::SelfReferencingHashes, event_id)); + } + } + + // Namespace: what the event says, else what the parent was stored under. + let lora_name = stored.lora_name.filter(|name| !name.is_empty()); + let cache_salt = stored + .cache_salt + .filter(|salt| !salt.is_empty()) + .or_else(|| salt_from_extra_keys(stored.extra_keys.as_deref(), lora_name.as_deref())); + let mut namespace = Namespace { + lora_name, + cache_salt, + }; + let blocks_state = self + .ranks + .entry(rank) + .or_default() + .tiers + .entry(tier as i32) + .or_default(); + // vLLM names the salt on block 0 only and SGLang repeats it; a chain + // hashed under a namespace stays in it, so a child fills what it + // omits from its parent. + if namespace.lora_name.is_none() || namespace.cache_salt.is_none() { + if let Some(parent) = stored + .parent_block_hash + .and_then(|parent| blocks_state.get(&parent.0)) + .and_then(|record| record.namespace.clone()) + { + if namespace.lora_name.is_none() { + namespace.lora_name = parent.lora_name; + } + if namespace.cache_salt.is_none() { + namespace.cache_salt = parent.cache_salt; + } + } + } + let stored_namespace = (!namespace.is_empty()).then(|| namespace.clone()); + let mut all_seen = true; + for hash in &stored.block_hashes { + match blocks_state.get_mut(&hash.0) { + // A second physical copy: the record stays, its digest with + // it, with one more copy counted. + Some(record) => { + record.copies = record.copies.saturating_add(1).min(COPIES_CAP); + record.namespace.clone_from(&stored_namespace); + } + None => { + all_seen = false; + blocks_state.insert( + hash.0, + BlockRecord { + namespace: stored_namespace.clone(), + copies: 1, + digest: None, + }, + ); + } + } + } + if all_seen { + self.counts.duplicate_stores += 1; + } + if bigrams { + self.counts.bigram_stores += 1; + } + if tail_aligned { + self.counts.tail_aligned_stores += 1; + } + if let Some(check) = self.hash_check { + verify_hashes( + check, + blocks_state, + &mut self.counts, + HashInput { + hashes: &stored.block_hashes, + parent: stored.parent_block_hash, + tokens: &token_ids, + width, + bigram_words: bigram_words.as_deref(), + extra_keys: stored.extra_keys.as_deref(), + lora: stored.lora_id.is_some() || namespace.lora_name.is_some(), + cache_salt: namespace.cache_salt.as_deref(), + }, + ); + } + + if window_only { + self.counts.window_only_stores += 1; + self.ranks + .entry(rank) + .or_default() + .window_forwarded + .extend(stored.block_hashes.iter().map(|hash| hash.0)); + } + + let cache_level = cache_level_of(tier); + let mut extra_keys = stored.extra_keys.unwrap_or_default().into_iter(); + let blocks = stored + .block_hashes + .iter() + .zip(token_ids.chunks_exact(width)) + .map(|(hash, tokens)| common::KvBlock { + block_hash: hash.0, + token_ids: tokens.to_vec(), + block_size: i32::try_from(width).unwrap_or(i32::MAX), + lora_id: stored.lora_id, + cache_level, + extra_keys: extra_keys + .next() + .flatten() + .map(|keys| keys.into_iter().filter_map(extra_key_proto).collect()) + .unwrap_or_default(), + }) + .collect(); + self.counts.forwarded_stored += 1; + let EventTail { + medium, + group_idx, + kv_cache_spec_kind, + kv_cache_spec_sliding_window, + ownership, + session_id, + .. + } = stored.tail; + Ok(kv_cache_event::Data::Stored(common::KvBlocksStored { + blocks, + parent_block_hash: stored.parent_block_hash.map(|hash| hash.0), + tier: Some(tier as i32), + medium, + group_idx, + kv_cache_spec_kind, + kv_cache_spec_sliding_window, + locality: Some(locality as i32), + ownership, + session_id, + lora_name: namespace.lora_name, + cache_salt: namespace.cache_salt, + })) + } + + fn normalize_removed( + &mut self, + removed: WireRemoved, + rank: i32, + event_id: u64, + ) -> Result { + let Admitted { + tier, + locality, + window_only, + } = self.admit(&removed.tail, rank, event_id, false, &removed.block_hashes)?; + let mut block_hashes: Vec = removed.block_hashes.iter().map(|hash| hash.0).collect(); + if let Some(state) = self.ranks.get_mut(&rank) { + if window_only { + // Once the rank has a main group, a window group's removal + // reaches only the blocks forwarded while it had none. + if state.main_seen() { + block_hashes.retain(|hash| state.window_forwarded.contains(hash)); + } + for hash in &block_hashes { + state.window_forwarded.remove(hash); + } + } + if let Some(blocks_state) = state.tiers.get_mut(&(tier as i32)) { + // One physical copy goes; the record goes with the last. + for hash in &block_hashes { + if let Some(record) = blocks_state.get_mut(hash) { + record.copies = record.copies.saturating_sub(1); + if record.copies == 0 { + blocks_state.remove(hash); + } + } + } + } + } + self.counts.forwarded_removed += 1; + let EventTail { + medium, + group_idx, + ownership, + .. + } = removed.tail; + Ok(kv_cache_event::Data::Removed(common::KvBlocksRemoved { + block_hashes, + cache_level: cache_level_of(tier), + tier: Some(tier as i32), + medium, + group_idx, + locality: Some(locality as i32), + ownership, + })) + } +} + +/// The engine-hash check's view of one admitted store. +struct HashInput<'a> { + hashes: &'a [BlockHash], + parent: Option, + tokens: &'a [u32], + width: usize, + /// Both words of every bigram, when the page came as bigrams. + bigram_words: Option<&'a [u32]>, + extra_keys: Option<&'a [Option>]>, + /// The store belongs to a LoRA request (vLLM folds the adapter path into + /// the hash and does not publish it). + lora: bool, + cache_salt: Option<&'a str>, +} + +/// Rehash a store's blocks the way `check` says the worker did and count the +/// outcome; what is forwarded never changes. Digests are kept on the records +/// so children can chain on them. +fn verify_hashes( + check: EngineHash, + records: &mut HashMap, + counts: &mut Counts, + input: HashInput<'_>, +) { + let blocks = input.hashes.len() as u64; + let mut prior: Option = match input.parent { + Some(parent) => match records.get(&parent.0).and_then(|record| record.digest) { + Some(digest) => Some(digest), + None => { + counts.hash_unverifiable += blocks; + return; + } + }, + None => match check { + EngineHash::Sglang => input.cache_salt.map(engine_hash::sglang_salt_seed), + // `vllm_block` applies NONE_HASH itself. + EngineHash::VllmSha256Cbor => None, + }, + }; + if check == EngineHash::VllmSha256Cbor && (input.lora || input.bigram_words.is_some()) { + counts.hash_unverifiable += blocks; + return; + } + for (index, hash) in input.hashes.iter().enumerate() { + let tokens = &input.tokens[index * input.width..(index + 1) * input.width]; + let digest = match check { + EngineHash::Sglang => { + let words = match input.bigram_words { + Some(words) => &words[index * 2 * input.width..(index + 1) * 2 * input.width], + None => tokens, + }; + engine_hash::sglang_page(prior.as_ref(), words) + } + EngineHash::VllmSha256Cbor => { + let keys = input + .extra_keys + .and_then(|keys| keys.get(index)) + .and_then(Option::as_deref); + let Ok(keys) = vllm_keys(keys, index) else { + counts.hash_unverifiable += blocks - index as u64; + return; + }; + engine_hash::vllm_block(prior.as_ref(), tokens, keys.as_deref()) + } + }; + let expected = match check { + EngineHash::Sglang => engine_hash::sglang_event_int(&digest), + EngineHash::VllmSha256Cbor => engine_hash::vllm_event_int(&digest), + }; + counts.hash_checked += 1; + if expected != hash.0 { + counts.hash_mismatch += 1; + } + if let Some(record) = records.get_mut(&hash.0) { + record.digest = Some(digest); + } + prior = Some(digest); + } +} + +/// vLLM's untagged event keys as the tagged keys inside the hash: block 0's +/// text is the cache salt (a LoRA request was excluded before), a pair is a +/// multimodal item, bytes are a prompt-embeddings digest. `Err` for a shape +/// the hash input cannot be rebuilt from. +fn vllm_keys(keys: Option<&[ExtraKey]>, index: usize) -> Result>, ()> { + let Some(keys) = keys.filter(|keys| !keys.is_empty()) else { + return Ok(None); + }; + keys.iter() + .map(|key| match key { + ExtraKey::Text(text) if index == 0 => Ok(VllmExtraKey::CacheSalt(text.clone())), + ExtraKey::Multimodal { identifier, offset } => Ok(VllmExtraKey::Mm { + identifier: identifier.clone(), + offset: *offset, + }), + ExtraKey::Blob(blob) => Ok(VllmExtraKey::PromptEmbeds(blob.clone())), + ExtraKey::Text(_) | ExtraKey::Number(_) | ExtraKey::Opaque => Err(()), + }) + .collect::, ()>>() + .map(Some) +} + +#[cfg(test)] +mod shapes_tests; + +#[cfg(test)] +mod tests { + use serde_json::{json, Value}; + + use super::*; + + /// A publisher batch from JSON: objects become msgpack maps (the engines' + /// tagged-map events), arrays become arrays (the batch envelope and the + /// legacy event layout). + fn batch_from(value: Value) -> WireBatch { + let bytes = rmp_serde::to_vec_named(&value).expect("encodes"); + rmp_serde::from_slice(&bytes).expect("decodes") + } + + fn normalize_all(batches: Vec) -> (Vec, Normalizer) { + let mut normalizer = Normalizer::new(); + let mut event_id = 0; + let out = batches + .into_iter() + .enumerate() + .map(|(seq, value)| normalize(&mut normalizer, value, seq as u64, &mut event_id)) + .collect(); + (out, normalizer) + } + + fn normalize( + normalizer: &mut Normalizer, + value: Value, + seq: u64, + event_id: &mut u64, + ) -> common::KvEventBatch { + normalizer.normalize_batch(batch_from(value), seq, event_id) + } + + fn one(events: Vec) -> Value { + json!([1700000000.5, events, 0]) + } + + fn store(hashes: &[i64], parent: Option, tokens: &[u32]) -> Value { + json!({ + "type": "BlockStored", + "block_hashes": hashes, + "parent_block_hash": parent, + "token_ids": tokens, + "block_size": 4, + "lora_id": null, + "medium": "GPU", + "lora_name": null, + "group_idx": 0, + "kv_cache_spec_kind": "full_attention", + }) + } + + fn remove(hashes: &[i64]) -> Value { + json!({"type": "BlockRemoved", "block_hashes": hashes, "medium": "GPU", "group_idx": 0}) + } + + fn with(mut value: Value, key: &str, item: Value) -> Value { + value[key] = item; + value + } + + fn stored(event: &common::KvCacheEvent) -> &common::KvBlocksStored { + match event.data { + Some(kv_cache_event::Data::Stored(ref stored)) => stored, + ref other => panic!("not a store: {other:?}"), + } + } + + fn removed(event: &common::KvCacheEvent) -> &common::KvBlocksRemoved { + match event.data { + Some(kv_cache_event::Data::Removed(ref removed)) => removed, + ref other => panic!("not a removal: {other:?}"), + } + } + + #[test] + fn hash_identity_folds_like_vllm() { + let mut digest = [0u8; 32]; + digest[23] = 0xaa; + digest[24..].copy_from_slice(&0x8000_0000_0000_0001u64.to_be_bytes()); + assert_eq!(low64_big_endian(&digest), 0x8000_0000_0000_0001); + assert_eq!(low64_big_endian(&[0, 0, 0, 0, 0, 0, 0, 7]), 7); + + let (batches, _) = normalize_all(vec![one(vec![store( + &[i64::MIN + 1, -3], + None, + &[1, 2, 3, 4, 5, 6, 7, 8], + )])]); + let blocks = &stored(&batches[0].events[0]).blocks; + assert_eq!( + blocks[0].block_hash, + i64::MIN + 1, + "a u64 above i64::MAX keeps its bits" + ); + assert_eq!( + blocks[1].block_hash, -3, + "SGLang's signed form passes through" + ); + } + + #[test] + fn unknown_and_malformed_events_cost_themselves_not_the_batch() { + let (batches, normalizer) = normalize_all(vec![one(vec![ + json!({"type": "BlockMigrated", "block_hashes": [1], "destination": "peer"}), + json!({"type": "BlockStored", "block_hashes": "nope", "token_ids": [1], "block_size": 1}), + json!({"block_hashes": [1]}), + json!(["BlockStored", "nope"]), + remove(&[9]), + ])]); + assert_eq!(batches[0].events.len(), 1); + assert_eq!(removed(&batches[0].events[0]).block_hashes, vec![9]); + assert_eq!( + batches[0].events[0].event_id, 5, + "ids advance for dropped events too" + ); + let counts = normalizer.counts(); + assert_eq!(counts.dropped(DropReason::UnknownType), 1); + assert_eq!(counts.dropped(DropReason::Malformed), 3); + assert_eq!(counts.forwarded_removed, 1); + } + + #[test] + fn residency_agent_events_are_dropped() { + let (batches, normalizer) = normalize_all(vec![one(vec![ + with(store(&[1], None, &[1, 2, 3, 4]), "ownership", json!("kvcr")), + with(remove(&[1]), "ownership", json!("KVCR")), + json!({"type": "AllBlocksCleared", "ownership": "kvcr"}), + json!({"type": "AllBlocksCleared"}), + ])]); + assert_eq!(batches[0].events.len(), 1); + assert!(matches!( + batches[0].events[0].data, + Some(kv_cache_event::Data::Cleared(_)) + )); + assert_eq!( + normalizer + .counts() + .dropped(DropReason::UnsupportedOwnership), + 3 + ); + assert_eq!(normalizer.counts().forwarded_cleared, 1); + } + + #[test] + fn remote_and_unknown_localities_are_dropped() { + let (batches, normalizer) = normalize_all(vec![one(vec![ + with( + store(&[1], None, &[1, 2, 3, 4]), + "locality", + json!("REMOTE"), + ), + with(remove(&[1]), "locality", json!("elsewhere")), + with(store(&[2], None, &[1, 2, 3, 4]), "locality", json!("local")), + ])]); + assert_eq!(batches[0].events.len(), 1); + assert_eq!( + stored(&batches[0].events[0]).locality, + Some(KvCacheLocality::Local as i32) + ); + assert_eq!(normalizer.counts().dropped(DropReason::NonLocalLocality), 2); + } + + #[test] + fn media_map_to_tiers_and_cache_levels() { + let table = [ + (None, KvCacheTier::Device, None), + (Some("GPU"), KvCacheTier::Device, None), + (Some("device"), KvCacheTier::Device, None), + (Some("CPU"), KvCacheTier::Host, Some(1)), + (Some("CPU_PINNED"), KvCacheTier::Host, Some(1)), + (Some("CPU_TIER1"), KvCacheTier::Host, Some(1)), + (Some("CPU_TIER2"), KvCacheTier::Disk, Some(2)), + (Some("DISK"), KvCacheTier::Disk, Some(2)), + (Some("NVME"), KvCacheTier::Disk, Some(2)), + (Some("STORAGE"), KvCacheTier::Disk, Some(2)), + (Some("EXTERNAL"), KvCacheTier::External, Some(3)), + (Some("NETWORK"), KvCacheTier::External, Some(3)), + (Some("REMOTE"), KvCacheTier::External, Some(3)), + (Some("SHARED"), KvCacheTier::External, Some(3)), + ]; + for (medium, tier, level) in table { + assert_eq!(tier_of(medium), Some(tier), "{medium:?}"); + assert_eq!(cache_level_of(tier), level, "{medium:?}"); + } + assert_eq!(tier_of(Some("MARS")), None); + + let (batches, normalizer) = normalize_all(vec![one(vec![ + with( + store(&[1], None, &[1, 2, 3, 4]), + "medium", + json!("CPU_PINNED"), + ), + with(remove(&[1]), "medium", json!("STORAGE")), + with(store(&[2], None, &[1, 2, 3, 4]), "medium", json!("MARS")), + with(store(&[3], None, &[1, 2, 3, 4]), "medium", Value::Null), + ])]); + let events = &batches[0].events; + assert_eq!(events.len(), 3); + let host = stored(&events[0]); + assert_eq!(host.tier, Some(KvCacheTier::Host as i32)); + assert_eq!(host.medium.as_deref(), Some("CPU_PINNED")); + assert_eq!(host.blocks[0].cache_level, Some(1)); + let disk = removed(&events[1]); + assert_eq!(disk.tier, Some(KvCacheTier::Disk as i32)); + assert_eq!(disk.cache_level, Some(2)); + let device = stored(&events[2]); + assert_eq!(device.tier, Some(KvCacheTier::Device as i32)); + assert_eq!(device.medium, None); + assert_eq!(device.blocks[0].cache_level, None); + assert_eq!(normalizer.counts().dropped(DropReason::UnknownMedium), 1); + } + + #[test] + fn non_main_attention_groups_are_dropped_and_remembered() { + let sliding = |hash: i64| { + let event = with(store(&[hash], None, &[1, 2, 3, 4]), "group_idx", json!(1)); + let event = with(event, "kv_cache_spec_kind", json!("sliding_window")); + with(event, "kv_cache_spec_sliding_window", json!(128)) + }; + let (batches, normalizer) = normalize_all(vec![ + one(vec![ + sliding(1), + store(&[2], None, &[1, 2, 3, 4]), + with( + with( + store(&[3], None, &[1, 2, 3, 4]), + "kv_cache_spec_kind", + json!("mla_attention"), + ), + "group_idx", + json!(2), + ), + with( + with( + store(&[4], None, &[1, 2, 3, 4]), + "kv_cache_spec_kind", + json!("sink_full_attention"), + ), + "group_idx", + json!(3), + ), + with( + with( + store(&[5], None, &[1, 2, 3, 4]), + "kv_cache_spec_kind", + json!("mamba"), + ), + "group_idx", + json!(4), + ), + ]), + one(vec![ + // Removals carry no kind: the learned groups decide. + with(remove(&[1]), "group_idx", json!(1)), + with(remove(&[2]), "group_idx", json!(0)), + // An unlearned group without a kind counts as main. + with(remove(&[7]), "group_idx", json!(9)), + // A kind-less store on a learned non-main group is dropped too. + with( + with(store(&[8], None, &[1, 2, 3, 4]), "group_idx", json!(4)), + "kv_cache_spec_kind", + Value::Null, + ), + ]), + ]); + assert_eq!(batches[0].events.len(), 3); + assert_eq!(batches[1].events.len(), 2); + assert_eq!(removed(&batches[1].events[0]).block_hashes, vec![2]); + assert_eq!(removed(&batches[1].events[1]).block_hashes, vec![7]); + assert_eq!( + normalizer + .counts() + .dropped(DropReason::NonMainAttentionGroup), + 4 + ); + assert_eq!( + stored(&batches[0].events[1]).kv_cache_spec_kind.as_deref(), + Some("mla_attention") + ); + } + + #[test] + fn placeholders_unaligned_and_self_referencing_stores_are_dropped() { + let (batches, normalizer) = normalize_all(vec![one(vec![ + // vLLM's CPU offload placeholder: a chunk key, no tokens, block_size 0. + with( + with(store(&[1], None, &[]), "block_size", json!(0)), + "medium", + json!("CPU"), + ), + store(&[], None, &[1, 2, 3, 4]), + store(&[2], None, &[1, 2, 3, 4, 5, 6]), + with(store(&[3], None, &[1, 2, 3, 4]), "block_size", json!(0)), + store(&[4], Some(4), &[1, 2, 3, 4]), + store(&[5, 5], None, &[1, 2, 3, 4, 5, 6, 7, 8]), + store(&[6], None, &[1, 2, 3, 4]), + ])]); + assert_eq!(batches[0].events.len(), 1); + assert_eq!(stored(&batches[0].events[0]).blocks[0].block_hash, 6); + let counts = normalizer.counts(); + assert_eq!(counts.dropped(DropReason::Placeholder), 2); + assert_eq!(counts.dropped(DropReason::UnalignedBlocks), 2); + assert_eq!(counts.dropped(DropReason::SelfReferencingHashes), 2); + } + + #[test] + fn bigram_pages_fold_to_their_tokens() { + let (batches, normalizer) = normalize_all(vec![one(vec![with( + store(&[1], None, &[]), + "token_ids", + json!([[1, 2], [2, 3], [3, 4], [4, 5]]), + )])]); + let block = &stored(&batches[0].events[0]).blocks[0]; + assert_eq!(block.token_ids, vec![1, 2, 3, 4]); + assert_eq!(block.block_size, 4); + assert_eq!(normalizer.counts().bigram_stores, 1); + assert_eq!(normalizer.counts().forwarded_stored, 1); + } + + #[test] + fn stores_and_removals_are_forwarded_one_for_one() { + let (batches, normalizer) = normalize_all(vec![ + one(vec![ + store(&[1, 2], None, &[1, 2, 3, 4, 5, 6, 7, 8]), + store(&[1, 2], None, &[1, 2, 3, 4, 5, 6, 7, 8]), + store(&[2, 3], Some(1), &[5, 6, 7, 8, 9, 10, 11, 12]), + ]), + one(vec![ + remove(&[1]), + remove(&[1]), + remove(&[1, 2]), + store(&[1], None, &[1, 2, 3, 4]), + ]), + ]); + assert_eq!(batches[0].events.len(), 3); + assert_eq!(batches[1].events.len(), 4); + for event in &batches[1].events[..2] { + assert_eq!(removed(event).block_hashes, vec![1]); + } + assert_eq!(removed(&batches[1].events[2]).block_hashes, vec![1, 2]); + let counts = normalizer.counts(); + assert_eq!(counts.forwarded_stored, 4); + assert_eq!(counts.forwarded_removed, 3); + assert_eq!( + counts.duplicate_stores, 1, + "only the exact resend; a re-store after removal is new" + ); + assert!(counts.dropped.is_empty()); + } + + #[test] + fn namespaces_come_from_the_event_its_extra_keys_or_its_parent() { + // (A prompt-embeddings digest is msgpack bin, which JSON cannot + // express; the generated fixtures cover it.) + let lora = |value: Value| with(value, "lora_name", json!("adapter")); + let (batches, _) = normalize_all(vec![ + one(vec![ + // vLLM: the salt rides in block 0's extra keys, after the LoRA + // name and the multimodal (identifier, offset) pairs. + with( + lora(store(&[1], None, &[1, 2, 3, 4])), + "extra_keys", + json!([["adapter", ["mm-abc", 0], "salt-1"]]), + ), + with( + lora(store(&[2], Some(1), &[5, 6, 7, 8])), + "extra_keys", + json!([["adapter"]]), + ), + store(&[3], Some(2), &[9, 10, 11, 12]), + // SGLang: the salt is a field; empty strings count as absent. + with( + store(&[4], None, &[1, 2, 3, 4]), + "cache_salt", + json!("tenant-a"), + ), + with( + with(store(&[5], None, &[1, 2, 3, 4]), "cache_salt", json!("")), + "lora_name", + json!(""), + ), + ]), + // Another rank does not inherit from this one. + json!([1700000001.0, [store(&[6], Some(3), &[13, 14, 15, 16])], 1]), + ]); + let events = &batches[0].events; + let first = stored(&events[0]); + assert_eq!(first.lora_name.as_deref(), Some("adapter")); + assert_eq!(first.cache_salt.as_deref(), Some("salt-1")); + let keys: Vec<_> = first.blocks[0] + .extra_keys + .iter() + .map(|key| key.key.clone().expect("a key")) + .collect(); + assert_eq!( + keys, + vec![ + kv_block_extra_key::Key::Text("adapter".into()), + kv_block_extra_key::Key::Multimodal(common::KvMultimodalKey { + identifier: "mm-abc".into(), + offset: 0, + }), + kv_block_extra_key::Key::Text("salt-1".into()), + ] + ); + let child = stored(&events[1]); + assert_eq!(child.lora_name.as_deref(), Some("adapter")); + assert_eq!( + child.cache_salt.as_deref(), + Some("salt-1"), + "inherited from block 0" + ); + let grandchild = stored(&events[2]); + assert_eq!(grandchild.lora_name.as_deref(), Some("adapter")); + assert_eq!(grandchild.cache_salt.as_deref(), Some("salt-1")); + assert_eq!(stored(&events[3]).cache_salt.as_deref(), Some("tenant-a")); + assert_eq!(stored(&events[3]).lora_name, None); + assert_eq!(stored(&events[4]).cache_salt, None); + assert_eq!(stored(&events[4]).lora_name, None); + let other_rank = stored(&batches[1].events[0]); + assert_eq!(other_rank.lora_name, None); + assert_eq!(other_rank.cache_salt, None); + } + + #[test] + fn a_clear_resets_its_rank_only() { + let (batches, normalizer) = normalize_all(vec![ + one(vec![store(&[1], None, &[1, 2, 3, 4])]), + json!([1700000001.0, [store(&[1], None, &[1, 2, 3, 4])], 1]), + one(vec![ + json!({"type": "AllBlocksCleared"}), + store(&[1], None, &[1, 2, 3, 4]), + ]), + json!([1700000003.0, [store(&[1], None, &[1, 2, 3, 4])], 1]), + ]); + assert!(matches!( + batches[2].events[0].data, + Some(kv_cache_event::Data::Cleared(_)) + )); + assert_eq!(batches[2].events.len(), 2); + assert_eq!(normalizer.counts().forwarded_cleared, 1); + assert_eq!( + normalizer.counts().duplicate_stores, + 1, + "rank 1 kept its seen set" + ); + assert_eq!(batches[3].dp_rank, Some(1)); + } + + #[test] + fn array_and_map_layouts_decode_alike() { + let map = one(vec![ + with( + with( + store(&[1, 2], Some(7), &[1, 2, 3, 4, 5, 6, 7, 8]), + "session_id", + json!("req-1"), + ), + "lora_name", + json!("adapter"), + ), + with(remove(&[1]), "locality", json!("LOCAL")), + json!({"type": "AllBlocksCleared"}), + ]); + let array = json!([ + 1700000000.5, + [ + [ + "BlockStored", + [1, 2], + 7, + [1, 2, 3, 4, 5, 6, 7, 8], + 4, + null, + "GPU", + "adapter", + null, + 0, + "full_attention", + null, + null, + null, + "req-1" + ], + ["BlockRemoved", [1], "GPU", 0, "LOCAL"], + ["AllBlocksCleared"], + ], + 0 + ]); + let (from_map, _) = normalize_all(vec![map]); + let (from_array, _) = normalize_all(vec![array]); + assert_eq!(from_map, from_array); + let first = stored(&from_map[0].events[0]); + assert_eq!(first.session_id.as_deref(), Some("req-1")); + assert_eq!(first.lora_name.as_deref(), Some("adapter")); + assert_eq!(first.parent_block_hash, Some(7)); + assert_eq!(first.group_idx, Some(0)); + assert_eq!( + removed(&from_map[0].events[1]).locality, + Some(KvCacheLocality::Local as i32) + ); + } + + #[test] + fn legacy_arrays_keep_their_trailing_slots_in_order() { + // ownership is the removal's sixth slot and must gate the event even + // when locality (the fifth) is set. + let (batches, normalizer) = normalize_all(vec![json!([ + 1700000000.5, + [ + ["BlockRemoved", [1], "STORAGE", 0, "LOCAL", "kvcr"], + ["BlockRemoved", [2], "STORAGE", 0, "REMOTE"], + ["BlockRemoved", [3], "GPU"], + ], + 0 + ])]); + assert_eq!(batches[0].events.len(), 1); + assert_eq!(removed(&batches[0].events[0]).block_hashes, vec![3]); + assert_eq!( + normalizer + .counts() + .dropped(DropReason::UnsupportedOwnership), + 1 + ); + assert_eq!(normalizer.counts().dropped(DropReason::NonLocalLocality), 1); + } + + #[test] + fn hash_check_verifies_sglang_chains_and_counts_mismatches() { + use crate::engine_hash::{sglang_chain, sglang_salt_seed}; + + let chain = sglang_chain(&[1, 2, 3, 4, 5, 6, 7, 8], 4, None); + let (first, second) = (chain[0].1, chain[1].1); + assert_eq!(first, -3488128144981237669); + let mut normalizer = Normalizer::with_hash_check(EngineHash::Sglang); + let mut event_id = 0; + let batch = normalize( + &mut normalizer, + one(vec![ + store(&[first], None, &[1, 2, 3, 4]), + store(&[second], Some(first), &[5, 6, 7, 8]), + // Tampered: still forwarded, counted. + store(&[second + 1], Some(first), &[5, 6, 7, 8]), + // A parent this stream never saw: nothing to chain on. + store(&[99], Some(12345), &[9, 10, 11, 12]), + ]), + 1, + &mut event_id, + ); + assert_eq!(batch.events.len(), 4, "a mismatch never drops"); + let counts = normalizer.counts(); + assert_eq!( + ( + counts.hash_checked, + counts.hash_mismatch, + counts.hash_unverifiable + ), + (3, 1, 1) + ); + + // A salted request seeds its chain; a coalesced two-page store + // chains page to page. + let seed = sglang_salt_seed("tenant-a"); + let salted = sglang_chain(&[1, 2, 3, 4, 5, 6, 7, 8], 4, Some(&seed)); + let mut normalizer = Normalizer::with_hash_check(EngineHash::Sglang); + normalize( + &mut normalizer, + one(vec![with( + store(&[salted[0].1, salted[1].1], None, &[1, 2, 3, 4, 5, 6, 7, 8]), + "cache_salt", + json!("tenant-a"), + )]), + 1, + &mut 0, + ); + let counts = normalizer.counts(); + assert_eq!((counts.hash_checked, counts.hash_mismatch), (2, 0)); + + // An Eagle bigram page hashes both words of every pair. + let mut normalizer = Normalizer::with_hash_check(EngineHash::Sglang); + normalize( + &mut normalizer, + one(vec![with( + store(&[-638950109823820341], None, &[]), + "token_ids", + json!([[1, 2], [2, 3], [3, 4], [4, 5]]), + )]), + 1, + &mut 0, + ); + let counts = normalizer.counts(); + assert_eq!((counts.hash_checked, counts.hash_mismatch), (1, 0)); + } + + #[test] + fn hash_check_verifies_vllm_sha256_cbor_chains() { + use crate::engine_hash::{vllm_block, vllm_event_int}; + + // Vectors A and B from the reference run. + let (a, b) = (-8885242862429187823i64, -3153830497298837583i64); + let mut normalizer = Normalizer::with_hash_check(EngineHash::VllmSha256Cbor); + let batch = normalize( + &mut normalizer, + one(vec![ + store(&[a, b], None, &[1, 2, 3, 4, 5, 6, 7, 8]), + // A LoRA block: the adapter path is not in the event. + with( + store(&[7], Some(b), &[9, 10, 11, 12]), + "lora_name", + json!("adapter"), + ), + // Unaligned: dropped before the check runs. + store(&[8], Some(b), &[9, 10, 11, 12, 13]), + ]), + 1, + &mut 0, + ); + assert_eq!(batch.events.len(), 2); + let counts = normalizer.counts(); + assert_eq!( + ( + counts.hash_checked, + counts.hash_mismatch, + counts.hash_unverifiable + ), + (2, 0, 1) + ); + + // Multimodal, salt and prompt-embeddings keys rebuild the tagged hash + // input (vectors E, F, G); the events carry the bytes JSON cannot. + let embeds: Vec = (0..32).collect(); + let e = vllm_block( + None, + &[1, 2, 3, 4], + Some(&[ + VllmExtraKey::Mm { + identifier: "mm-abc".into(), + offset: 0, + }, + VllmExtraKey::CacheSalt("salt-1".into()), + VllmExtraKey::PromptEmbeds(embeds.clone()), + ]), + ); + let f = vllm_block( + Some(&e), + &[5, 6, 7, 8], + Some(&[VllmExtraKey::Mm { + identifier: "mm-abc".into(), + offset: -4, + }]), + ); + let g = vllm_block(Some(&f), &[9, 10, 11, 12], None); + let stored = |hashes: Vec, parent: Option, tokens: Vec, keys| { + WireEvent::BlockStored(WireStored { + block_hashes: hashes.into_iter().map(BlockHash).collect(), + parent_block_hash: parent.map(BlockHash), + token_ids: WireTokens::Ids(tokens), + block_size: 4, + lora_id: None, + lora_name: None, + cache_salt: None, + extra_keys: keys, + tail: EventTail { + medium: Some("GPU".into()), + group_idx: Some(0), + kv_cache_spec_kind: Some("full_attention".into()), + ..EventTail::default() + }, + }) + }; + let mut normalizer = Normalizer::with_hash_check(EngineHash::VllmSha256Cbor); + let events = [ + stored( + vec![vllm_event_int(&e)], + None, + vec![1, 2, 3, 4], + Some(vec![Some(vec![ + ExtraKey::Multimodal { + identifier: "mm-abc".into(), + offset: 0, + }, + ExtraKey::Text("salt-1".into()), + ExtraKey::Blob(embeds), + ])]), + ), + stored( + vec![vllm_event_int(&f)], + Some(vllm_event_int(&e)), + vec![5, 6, 7, 8], + Some(vec![Some(vec![ExtraKey::Multimodal { + identifier: "mm-abc".into(), + offset: -4, + }])]), + ), + stored( + vec![vllm_event_int(&g)], + Some(vllm_event_int(&f)), + vec![9, 10, 11, 12], + None, + ), + // A key shape the hash input cannot be rebuilt from. + stored( + vec![5], + Some(vllm_event_int(&g)), + vec![13, 14, 15, 16], + Some(vec![Some(vec![ExtraKey::Number(3)])]), + ), + ]; + for (index, event) in events.into_iter().enumerate() { + assert!(normalizer + .normalize(event, Some(0), index as u64 + 1) + .is_ok()); + } + let counts = normalizer.counts(); + assert_eq!( + ( + counts.hash_checked, + counts.hash_mismatch, + counts.hash_unverifiable + ), + (3, 0, 1) + ); + assert_eq!(stored_salt(&normalizer), Some("salt-1".to_string())); + } + + fn stored_salt(normalizer: &Normalizer) -> Option { + normalizer + .ranks + .get(&0) + .and_then(|rank| rank.tiers.get(&(KvCacheTier::Device as i32))) + .and_then(|tier| tier.values().find_map(|record| record.namespace.clone())) + .and_then(|namespace| namespace.cache_salt) + } + + #[test] + fn hash_check_is_off_unless_asked() { + assert_eq!(Normalizer::new().hash_check(), None); + assert_eq!( + Normalizer::with_hash_check(EngineHash::Sglang).hash_check(), + Some(EngineHash::Sglang) + ); + // Without the check, a wrong hash is nobody's business here. + let (_, normalizer) = normalize_all(vec![one(vec![store(&[1], None, &[1, 2, 3, 4])])]); + assert_eq!(normalizer.counts().hash_checked, 0); + } + + /// A sliding-window group's store: the window's blocks against the whole + /// span the step computed. + fn window(hashes: &[i64], parent: Option, tokens: &[u32]) -> Value { + let event = with(store(hashes, parent, tokens), "group_idx", json!(1)); + let event = with(event, "kv_cache_spec_kind", json!("sliding_window")); + with(event, "kv_cache_spec_sliding_window", json!(128)) + } + + fn remove_in(hashes: &[i64], group: u32) -> Value { + with(remove(hashes), "group_idx", json!(group)) + } + + #[test] + fn two_groups_with_identical_hashes_index_the_full_attention_group_once() { + // vLLM's hybrid manager publishes every block in both groups under + // the same hash, the window group's store first in the step and its + // removals as the window moves on; only the full-attention group's + // events are forwarded, stores and removals alike, so the gateway + // counts one copy per physical block. + let (batches, normalizer) = normalize_all(vec![ + one(vec![ + window(&[2], None, &[1, 2, 3, 4, 5, 6, 7, 8]), + store(&[1, 2], None, &[1, 2, 3, 4, 5, 6, 7, 8]), + ]), + one(vec![ + remove_in(&[2], 1), + store(&[1, 2], None, &[1, 2, 3, 4, 5, 6, 7, 8]), // the second copy + remove_in(&[1], 0), + remove_in(&[1], 0), + ]), + ]); + assert_eq!(batches[0].events.len(), 1); + let full = stored(&batches[0].events[0]); + assert_eq!( + full.blocks.iter().map(|b| b.block_hash).collect::>(), + vec![1, 2] + ); + assert_eq!(full.group_idx, Some(0)); + assert_eq!(batches[1].events.len(), 3); + assert_eq!(removed(&batches[1].events[1]).block_hashes, vec![1]); + assert_eq!(removed(&batches[1].events[2]).block_hashes, vec![1]); + let counts = normalizer.counts(); + assert_eq!(counts.dropped(DropReason::NonMainAttentionGroup), 2); + assert_eq!((counts.forwarded_stored, counts.forwarded_removed), (2, 2)); + assert_eq!( + ( + counts.duplicate_stores, + counts.window_only_stores, + counts.tail_aligned_stores + ), + (1, 0, 0) + ); + } + + #[test] + fn a_window_groups_hashes_take_the_tail_of_its_tokens() { + // A rank with only a sliding-window group: a store spans the whole + // range the step computed and names the window's last blocks, so the + // hashes get the tail of the tokens, and the group rides along. + let span: Vec = (1..=12).collect(); + let (batches, normalizer) = normalize_all(vec![one(vec![ + window(&[5], None, &span), + window(&[6, 7], Some(5), &span), + // A span that is not whole blocks is unreadable. + window(&[8], Some(7), &(1..=13).collect::>()), + // The window moved on without caching anything new. + window(&[], Some(7), &(1..=16).collect::>()), + ])]); + assert_eq!(batches[0].events.len(), 2); + let first = stored(&batches[0].events[0]); + assert_eq!(first.blocks[0].token_ids, vec![9, 10, 11, 12]); + assert_eq!(first.group_idx, Some(1)); + assert_eq!(first.kv_cache_spec_kind.as_deref(), Some("sliding_window")); + assert_eq!(first.kv_cache_spec_sliding_window, Some(128)); + let second = stored(&batches[0].events[1]); + assert_eq!( + second + .blocks + .iter() + .map(|b| b.token_ids.clone()) + .collect::>(), + vec![vec![5, 6, 7, 8], vec![9, 10, 11, 12]] + ); + assert_eq!(second.parent_block_hash, Some(5)); + let counts = normalizer.counts(); + assert_eq!( + (counts.window_only_stores, counts.tail_aligned_stores), + (2, 2) + ); + assert_eq!(counts.dropped(DropReason::UnalignedBlocks), 1); + assert_eq!(counts.dropped(DropReason::Placeholder), 1); + } + + #[test] + fn a_window_only_rank_is_forwarded_until_a_main_group_appears() { + let (batches, normalizer) = normalize_all(vec![ + one(vec![ + window(&[1], None, &[1, 2, 3, 4]), + remove_in(&[1], 1), + window(&[2], None, &[1, 2, 3, 4, 5, 6, 7, 8]), + ]), + one(vec![ + store(&[3], None, &[1, 2, 3, 4]), // a main-attention group shows up + window(&[4], Some(3), &[1, 2, 3, 4, 5, 6, 7, 8]), + remove_in(&[2], 1), // forwarded while window-only: still removable + remove_in(&[9], 1), // never forwarded: dropped + remove_in(&[2, 9], 1), // nothing left of it + ]), + ]); + assert_eq!(batches[0].events.len(), 3); + assert_eq!(batches[1].events.len(), 2); + assert_eq!(removed(&batches[1].events[1]).block_hashes, vec![2]); + let counts = normalizer.counts(); + assert_eq!( + (counts.window_only_stores, counts.tail_aligned_stores), + (2, 1) + ); + assert_eq!(counts.dropped(DropReason::NonMainAttentionGroup), 3); + } + + #[test] + fn one_copy_removals_pass_through_uncollapsed() { + // vLLM keeps up to two physical blocks per hash and publishes a + // removal when one copy goes; the gateway counts copies, so both + // stores and both removals must reach it as they are. + let (batches, normalizer) = normalize_all(vec![one(vec![ + store(&[1], None, &[1, 2, 3, 4]), + store(&[1], None, &[1, 2, 3, 4]), + remove(&[1]), + remove(&[1]), + ])]); + assert_eq!(batches[0].events.len(), 4); + let counts = normalizer.counts(); + assert_eq!( + ( + counts.forwarded_stored, + counts.forwarded_removed, + counts.duplicate_stores + ), + (2, 2, 1) + ); + } + + /// The engine keeps up to two physical copies of a block and removes + /// them one at a time; the check's parent memory counts the copies the + /// way the gateway does, so a child stored after the first copy went + /// still chains on its parent's digest, and only after the last copy is + /// the parent gone and the child unverifiable. Both removals are + /// forwarded either way. + #[test] + fn the_hash_checks_parent_memory_counts_physical_copies() { + let root = engine_hash::sglang_chain(&[1, 2, 3, 4], 4, None)[0].1; + let child_b = engine_hash::sglang_chain(&[1, 2, 3, 4, 5, 6, 7, 8], 4, None)[1].1; + let child_c = engine_hash::sglang_chain(&[1, 2, 3, 4, 13, 14, 15, 16], 4, None)[1].1; + let mut normalizer = Normalizer::with_hash_check(EngineHash::Sglang); + let mut event_id = 0; + normalize( + &mut normalizer, + one(vec![ + store(&[root], None, &[1, 2, 3, 4]), + store(&[root], None, &[1, 2, 3, 4]), // the second physical copy + remove(&[root]), // one copy goes + store(&[child_b], Some(root), &[5, 6, 7, 8]), + ]), + 1, + &mut event_id, + ); + let counts = normalizer.counts(); + assert_eq!( + ( + counts.hash_checked, + counts.hash_mismatch, + counts.hash_unverifiable, + counts.duplicate_stores + ), + (3, 0, 0, 1) + ); + normalize( + &mut normalizer, + one(vec![ + remove(&[root]), // the last copy + store(&[child_c], Some(root), &[13, 14, 15, 16]), + ]), + 2, + &mut event_id, + ); + let counts = normalizer.counts(); + assert_eq!((counts.hash_checked, counts.hash_unverifiable), (3, 1)); + assert_eq!((counts.forwarded_stored, counts.forwarded_removed), (4, 2)); + } +} diff --git a/crates/engine_servicer/src/kv_wire/shapes_tests.rs b/crates/engine_servicer/src/kv_wire/shapes_tests.rs new file mode 100644 index 0000000000..bda38712c2 --- /dev/null +++ b/crates/engine_servicer/src/kv_wire/shapes_tests.rs @@ -0,0 +1,1296 @@ +//! The engines' KV-event wire shapes, built in code and normalized: vLLM and +//! SGLang, each in the current tagged-map layout and the legacy tag-first +//! array layout, with every field variant the relay reads (both hash forms, +//! the tiers and their cache levels, parents, LoRA names and cache salts, +//! extra keys, bigram pages, a clear, a second DP rank and every drop rule), +//! the normalizer's output asserted event by event and in its counters, and +//! the normalized streams' round trip into the gateway's reference index. +//! +//! The bytes follow msgspec's encoding of the engines' own structs (vLLM +//! `vllm/distributed/kv_events.py`, SGLang +//! `python/sglang/srt/disaggregation/kv_events.py`): a batch is the array +//! `[ts, events, dp_rank]`; an event is a map with its `type` tag in which, +//! under `omit_defaults`, a field with a default appears only when set, or +//! the legacy array with the tag first and every field in declaration order, +//! nil when unset (`omit_defaults` does not thin an array). The +//! Python servicer's `tests/test_kv_relay.py` builds the same scenarios with +//! msgspec itself and expects the same output, which keeps the two relays +//! in step. + +use std::collections::BTreeMap; + +use engine_zmq_client::codec::TrailingTolerant; +use kv_index::{ + compute_request_content_hashes, + salt::{content_hash_with_seed, namespace_seed, namespaced_request_content_hashes}, + ApplyError, ContentHash, ReferenceIndexer, SequenceHash, StoredBlock, +}; +use rmpv::Value; +use sha2::{Digest, Sha256}; +use smg_grpc_client::common_proto::{ + kv_block_extra_key, kv_cache_event, KvBlocksStored, KvCacheEvent, KvCacheTier, KvEventBatch, +}; + +use super::{Counts, Normalizer, WireBatch}; + +const BLOCK_SIZE: u64 = 4; + +// --------------------------------------------------------------------------- +// Encoding: msgspec's two layouts +// --------------------------------------------------------------------------- + +/// How msgspec lays an event out: a map carrying its `type` tag, or the +/// legacy tag-first array. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +enum Layout { + Map, + Array, +} + +impl Layout { + /// An event from its tag, its required fields and its defaulted fields in + /// declaration order (`None` is a field left at its default). The map + /// keeps every required field, nil included, and only the set defaulted + /// ones; the array keeps every slot, nil when unset. + fn event( + self, + tag: &str, + required: Vec<(&str, Value)>, + defaulted: Vec<(&str, Option)>, + ) -> Value { + match self { + Layout::Map => { + let mut entries = vec![(Value::from("type"), Value::from(tag))]; + entries.extend( + required + .into_iter() + .map(|(key, value)| (Value::from(key), value)), + ); + entries.extend( + defaulted + .into_iter() + .filter_map(|(key, value)| Some((Value::from(key), value?))), + ); + Value::Map(entries) + } + Layout::Array => { + let mut items = vec![Value::from(tag)]; + items.extend(required.into_iter().map(|(_, value)| value)); + items.extend( + defaulted + .into_iter() + .map(|(_, value)| value.unwrap_or(Value::Nil)), + ); + Value::Array(items) + } + } + } + + fn name(self) -> &'static str { + match self { + Layout::Map => "map", + Layout::Array => "array", + } + } +} + +/// A block hash as an engine publishes it. +#[derive(Clone, Copy, Debug)] +enum Hash { + /// vLLM's integer form: the low 64 bits of the digest, unsigned. + Unsigned(u64), + /// SGLang's integer form: the high 64 bits of the digest, signed. + Signed(i64), + /// vLLM's raw digest (`VLLM_KV_EVENTS_USE_INT_BLOCK_HASHES=0`). + Digest([u8; 32]), +} + +impl Hash { + fn wire(self) -> Value { + match self { + Hash::Unsigned(value) => Value::from(value), + Hash::Signed(value) => Value::from(value), + Hash::Digest(digest) => Value::from(digest.to_vec()), + } + } + + /// The identity the relay forwards: the same 64 bits as an `i64`, a + /// digest's low 64 bits read big-endian. + fn forwarded(self) -> i64 { + match self { + Hash::Unsigned(value) => value as i64, + Hash::Signed(value) => value, + Hash::Digest(digest) => i64::from_be_bytes(low64(&digest)), + } + } +} + +fn digest(label: &str) -> [u8; 32] { + Sha256::digest(label.as_bytes()).into() +} + +fn low64(digest: &[u8; 32]) -> [u8; 8] { + let mut low = [0u8; 8]; + low.copy_from_slice(&digest[24..]); + low +} + +fn vllm_hash(label: &str) -> Hash { + Hash::Unsigned(u64::from_be_bytes(low64(&digest(label)))) +} + +fn sglang_hash(label: &str) -> Hash { + let mut high = [0u8; 8]; + high.copy_from_slice(&digest(label)[..8]); + Hash::Signed(i64::from_be_bytes(high)) +} + +fn hashes(hashes: &[Hash]) -> Value { + Value::Array(hashes.iter().map(|hash| hash.wire()).collect()) +} + +fn ints(tokens: &[u32]) -> Value { + Value::Array( + tokens + .iter() + .map(|&token| Value::from(u64::from(token))) + .collect(), + ) +} + +/// An Eagle bigram page: `[token, next token]` per position. +fn bigrams(pairs: &[(u32, u32)]) -> Value { + Value::Array( + pairs + .iter() + .map(|&(token, next)| { + Value::Array(vec![ + Value::from(u64::from(token)), + Value::from(u64::from(next)), + ]) + }) + .collect(), + ) +} + +fn opt_text(text: Option<&str>) -> Value { + text.map_or(Value::Nil, Value::from) +} + +fn opt_hash(hash: Option) -> Value { + hash.map_or(Value::Nil, Hash::wire) +} + +/// vLLM's `extra_keys`: one list per block, nil for a block without any. +fn extra_keys(per_block: &[Option>]) -> Value { + Value::Array( + per_block + .iter() + .map(|keys| { + keys.as_ref() + .map_or(Value::Nil, |items| Value::Array(items.clone())) + }) + .collect(), + ) +} + +/// A publisher batch, `[ts, events, dp_rank]`, as bytes. Both engines' +/// batches are `array_like`; vLLM's omits a rank left at its default and +/// SGLang's always writes the slot, which is the same thing once a rank is +/// set, as it is in every scenario here. +fn batch(ts: f64, rank: i64, events: Vec) -> Vec { + let value = Value::Array(vec![ + Value::from(ts), + Value::Array(events), + Value::from(rank), + ]); + let mut bytes = Vec::new(); + rmpv::encode::write_value(&mut bytes, &value).expect("msgpack encodes"); + bytes +} + +/// vLLM's `BlockStored` with the struct's defaults: `None` in a defaulted +/// field is a field left at its default. +#[derive(Clone)] +struct VllmStored { + hashes: Vec, + parent: Option, + tokens: Vec, + block_size: u64, + lora_id: Option, + medium: Option<&'static str>, + lora_name: Option<&'static str>, + extra_keys: Option>>>, + group_idx: Option, + spec_kind: Option<&'static str>, + sliding_window: Option, + locality: Option<&'static str>, + ownership: Option<&'static str>, + session_id: Option<&'static str>, +} + +impl VllmStored { + /// A store on the device from the main attention group. + fn gpu(hashes: Vec, parent: Option, tokens: Vec) -> Self { + Self { + hashes, + parent, + tokens, + block_size: BLOCK_SIZE, + lora_id: None, + medium: Some("GPU"), + lora_name: None, + extra_keys: None, + group_idx: Some(0), + spec_kind: Some("full_attention"), + sliding_window: None, + locality: None, + ownership: None, + session_id: None, + } + } + + fn wire(&self, layout: Layout) -> Value { + layout.event( + "BlockStored", + vec![ + ("block_hashes", hashes(&self.hashes)), + ("parent_block_hash", opt_hash(self.parent)), + ("token_ids", ints(&self.tokens)), + ("block_size", Value::from(self.block_size)), + ("lora_id", self.lora_id.map_or(Value::Nil, Value::from)), + ("medium", opt_text(self.medium)), + ("lora_name", opt_text(self.lora_name)), + ], + vec![ + ("extra_keys", self.extra_keys.as_deref().map(extra_keys)), + ("group_idx", self.group_idx.map(Value::from)), + ("kv_cache_spec_kind", self.spec_kind.map(Value::from)), + ( + "kv_cache_spec_sliding_window", + self.sliding_window.map(Value::from), + ), + ("locality", self.locality.map(Value::from)), + ("ownership", self.ownership.map(Value::from)), + ("session_id", self.session_id.map(Value::from)), + ], + ) + } +} + +/// vLLM's `BlockRemoved`: the medium is required, the rest defaulted. +fn vllm_removed(layout: Layout, removed: &[Hash], medium: &str, group_idx: Option) -> Value { + layout.event( + "BlockRemoved", + vec![ + ("block_hashes", hashes(removed)), + ("medium", Value::from(medium)), + ], + vec![ + ("group_idx", group_idx.map(Value::from)), + ("locality", None), + ("ownership", None), + ], + ) +} + +/// SGLang's `BlockStored`: `lora_id` is required (always nil), the medium, +/// salt and session are defaulted. +#[derive(Clone)] +struct SglangStored { + hashes: Vec, + parent: Option, + tokens: Value, + medium: Option<&'static str>, + cache_salt: Option<&'static str>, + session_id: Option<&'static str>, +} + +impl SglangStored { + fn gpu(hashes: Vec, parent: Option, tokens: &[u32]) -> Self { + Self { + hashes, + parent, + tokens: ints(tokens), + medium: Some("GPU"), + cache_salt: None, + session_id: None, + } + } + + fn wire(&self, layout: Layout) -> Value { + layout.event( + "BlockStored", + vec![ + ("block_hashes", hashes(&self.hashes)), + ("parent_block_hash", opt_hash(self.parent)), + ("token_ids", self.tokens.clone()), + ("block_size", Value::from(BLOCK_SIZE)), + ("lora_id", Value::Nil), + ], + vec![ + ("medium", self.medium.map(Value::from)), + ("cache_salt", self.cache_salt.map(Value::from)), + ("session_id", self.session_id.map(Value::from)), + ], + ) + } +} + +/// SGLang's `BlockRemoved`: only the hashes are required. +fn sglang_removed(layout: Layout, removed: &[Hash], medium: &str) -> Value { + layout.event( + "BlockRemoved", + vec![("block_hashes", hashes(removed))], + vec![("medium", Some(Value::from(medium)))], + ) +} + +fn cleared(layout: Layout) -> Value { + layout.event("AllBlocksCleared", Vec::new(), Vec::new()) +} + +/// An event type the relay does not know (stands in for a future one). +fn migrated(layout: Layout, moved: &[Hash]) -> Value { + layout.event( + "BlockMigrated", + vec![ + ("block_hashes", hashes(moved)), + ("destination", Value::from("peer")), + ], + Vec::new(), + ) +} + +/// A store whose hashes are not a list: no struct produces it, so it is +/// written out by hand in each layout. +fn malformed_store(layout: Layout) -> Value { + match layout { + Layout::Map => Value::Map(vec![ + (Value::from("type"), Value::from("BlockStored")), + (Value::from("block_hashes"), Value::from("nope")), + (Value::from("parent_block_hash"), Value::Nil), + (Value::from("token_ids"), ints(&[1, 2, 3, 4])), + (Value::from("block_size"), Value::from(BLOCK_SIZE)), + ]), + Layout::Array => Value::Array(vec![ + Value::from("BlockStored"), + Value::from("nope"), + Value::Nil, + ints(&[1, 2, 3, 4]), + Value::from(BLOCK_SIZE), + ]), + } +} + +// --------------------------------------------------------------------------- +// Expectations +// --------------------------------------------------------------------------- + +/// A forwarded store as the gateway reads it. +#[derive(Debug)] +struct WantStored { + rank: i32, + hashes: Vec, + parent: Option, + tier: KvCacheTier, + cache_level: Option, + tokens: Vec>, + lora_name: Option<&'static str>, + cache_salt: Option<&'static str>, + group_idx: Option, + session_id: Option<&'static str>, + /// Per block, when the scenario pins them. + extra_keys: Option>>, +} + +impl WantStored { + /// A plain device store on rank 0. + fn device(hashes: Vec, tokens: Vec>) -> Self { + Self { + rank: 0, + hashes, + parent: None, + tier: KvCacheTier::Device, + cache_level: None, + tokens, + lora_name: None, + cache_salt: None, + group_idx: None, + session_id: None, + extra_keys: None, + } + } +} + +/// One of a block's extra keys, in the shape the relay forwards it. +#[derive(Debug, PartialEq, Eq)] +enum Key { + Text(&'static str), + Multimodal(&'static str, i64), + BlobLen(usize), +} + +#[derive(Debug)] +enum Want { + Stored(WantStored), + Removed { + rank: i32, + hashes: Vec, + tier: KvCacheTier, + cache_level: Option, + }, + Cleared { + rank: i32, + }, +} + +fn removed(hashes: Vec) -> Want { + Want::Removed { + rank: 0, + hashes, + tier: KvCacheTier::Device, + cache_level: None, + } +} + +fn removed_on(hashes: Vec, tier: KvCacheTier, cache_level: i32) -> Want { + Want::Removed { + rank: 0, + hashes, + tier, + cache_level: Some(cache_level), + } +} + +#[derive(Debug, Default)] +struct WantCounts { + stored: u64, + removed: u64, + cleared: u64, + duplicate_stores: u64, + bigram_stores: u64, + dropped: BTreeMap<&'static str, u64>, +} + +/// One engine's stream in one layout: its batches as the publisher's bytes +/// and what the relay must make of them. +struct Scenario { + name: String, + engine: &'static str, + layout: Layout, + batches: Vec>, + want: Vec, + counts: WantCounts, +} + +/// The scenario decoded and normalized the way the relay forwards it. +fn normalized(scenario: &Scenario) -> (Vec, Counts) { + let mut normalizer = Normalizer::new(); + let mut event_id = 0; + let batches = scenario + .batches + .iter() + .enumerate() + .map(|(seq, bytes)| { + let batch = rmp_serde::from_slice::>(bytes) + .unwrap_or_else(|error| panic!("{} batch {seq}: {error}", scenario.name)) + .0; + normalizer.normalize_batch(batch, seq as u64, &mut event_id) + }) + .collect(); + (batches, normalizer.counts().clone()) +} + +fn key_shape(key: &kv_block_extra_key::Key) -> Key { + match key { + kv_block_extra_key::Key::Text(text) => Key::Text(Box::leak(text.clone().into_boxed_str())), + kv_block_extra_key::Key::Number(number) => { + panic!("no scenario forwards a numeric key, got {number}") + } + kv_block_extra_key::Key::Blob(blob) => Key::BlobLen(blob.len()), + kv_block_extra_key::Key::Multimodal(mm) => { + Key::Multimodal(Box::leak(mm.identifier.clone().into_boxed_str()), mm.offset) + } + } +} + +/// Every forwarded event against its expectation, then the counters. +fn check(scenario: &Scenario, batches: &[KvEventBatch], counts: &Counts) { + let name = &scenario.name; + let forwarded: Vec<(Option, &KvCacheEvent)> = batches + .iter() + .flat_map(|batch| batch.events.iter().map(move |event| (batch.dp_rank, event))) + .collect(); + assert_eq!( + forwarded.len(), + scenario.want.len(), + "{name}: forwarded event count; got {:#?}", + forwarded.iter().map(|(_, event)| event).collect::>() + ); + for (index, ((rank, event), want)) in forwarded.iter().zip(&scenario.want).enumerate() { + let at = format!("{name} forwarded event {index} (id {})", event.event_id); + match (&event.data, want) { + (Some(kv_cache_event::Data::Stored(stored)), Want::Stored(want)) => { + assert_eq!(*rank, Some(want.rank), "{at}: dp_rank"); + let got_hashes: Vec = stored.blocks.iter().map(|b| b.block_hash).collect(); + assert_eq!(got_hashes, want.hashes, "{at}: hashes"); + assert_eq!(stored.parent_block_hash, want.parent, "{at}: parent"); + assert_eq!(stored.tier, Some(want.tier as i32), "{at}: tier"); + let got_tokens: Vec> = + stored.blocks.iter().map(|b| b.token_ids.clone()).collect(); + assert_eq!(got_tokens, want.tokens, "{at}: tokens"); + for block in &stored.blocks { + assert_eq!(block.cache_level, want.cache_level, "{at}: cache_level"); + assert_eq!( + block.block_size as usize, + block.token_ids.len(), + "{at}: block_size" + ); + } + assert_eq!( + stored.lora_name.as_deref(), + want.lora_name, + "{at}: lora_name" + ); + assert_eq!( + stored.cache_salt.as_deref(), + want.cache_salt, + "{at}: cache_salt" + ); + assert_eq!(stored.group_idx, want.group_idx, "{at}: group_idx"); + assert_eq!( + stored.session_id.as_deref(), + want.session_id, + "{at}: session_id" + ); + if let Some(expected_keys) = &want.extra_keys { + let got_keys: Vec> = stored + .blocks + .iter() + .map(|block| { + block + .extra_keys + .iter() + .map(|key| key_shape(key.key.as_ref().expect("a key"))) + .collect() + }) + .collect(); + assert_eq!(&got_keys, expected_keys, "{at}: extra_keys"); + } + } + ( + Some(kv_cache_event::Data::Removed(got)), + Want::Removed { + rank: want_rank, + hashes, + tier, + cache_level, + }, + ) => { + assert_eq!(*rank, Some(*want_rank), "{at}: dp_rank"); + assert_eq!(&got.block_hashes, hashes, "{at}: hashes"); + assert_eq!(got.tier, Some(*tier as i32), "{at}: tier"); + assert_eq!(got.cache_level, *cache_level, "{at}: cache_level"); + } + (Some(kv_cache_event::Data::Cleared(_)), Want::Cleared { rank: want_rank }) => { + assert_eq!(*rank, Some(*want_rank), "{at}: dp_rank"); + } + (got, want) => panic!("{at}: got {got:?}, wanted {want:?}"), + } + } + + let want = &scenario.counts; + assert_eq!( + counts.forwarded_stored, want.stored, + "{name}: forwarded_stored" + ); + assert_eq!( + counts.forwarded_removed, want.removed, + "{name}: forwarded_removed" + ); + assert_eq!( + counts.forwarded_cleared, want.cleared, + "{name}: forwarded_cleared" + ); + assert_eq!( + counts.duplicate_stores, want.duplicate_stores, + "{name}: duplicate_stores" + ); + assert_eq!( + counts.bigram_stores, want.bigram_stores, + "{name}: bigram_stores" + ); + let dropped: BTreeMap<&'static str, u64> = counts + .dropped + .iter() + .map(|(reason, count)| (reason.as_str(), *count)) + .collect(); + assert_eq!(dropped, want.dropped, "{name}: dropped"); +} + +// --------------------------------------------------------------------------- +// Scenarios +// --------------------------------------------------------------------------- + +/// vLLM: both hash forms, a sliding-window group, a second physical copy +/// with per-copy removals, the offload tiers and every drop rule, a pool +/// reset, a LoRA request with a multimodal item, a cache salt and a prompt +/// embeddings digest whose child inherits the salt, and a second DP rank. +fn vllm(layout: Layout) -> Scenario { + let h = |name: &str| vllm_hash(&format!("vllm-{name}")); + let e = |name: &str| h(name).forwarded(); + let (d1, d2) = ( + Hash::Digest(digest("vllm-digest-1")), + Hash::Digest(digest("vllm-digest-2")), + ); + let embeds = digest("vllm-prompt-embeds"); + // The low 64 bits of a digest can exceed i64::MAX; make sure one does. + assert!( + "abcdefghijk" + .chars() + .any(|name| matches!(h(&name.to_string()), Hash::Unsigned(value) if value >= 1 << 63)), + "pick labels with a high bit set" + ); + let gpu = VllmStored::gpu; + let mut batches = Vec::new(); + let mut want = Vec::new(); + + // Batch 0: a plain chain in both hash forms; a sliding-window group store. + batches.push(batch( + 1_700_000_000.0, + 0, + vec![ + VllmStored { + session_id: Some("req-1"), + ..gpu(vec![h("a"), h("b")], None, (1..=8).collect()) + } + .wire(layout), + // Sliding-window group: more tokens than hashes x block size, no + // hashes at all. Dropped by the group gate. + VllmStored { + group_idx: Some(1), + spec_kind: Some("sliding_window"), + sliding_window: Some(128), + ..gpu(Vec::new(), None, (1..=8).collect()) + } + .wire(layout), + // The same group's usual shape: token_ids span the whole computed + // range and block_hashes name only the window's last blocks, so + // the hashes belong to the tail of the tokens. Dropped whole by + // the group gate, never sliced from the head. + VllmStored { + group_idx: Some(1), + spec_kind: Some("sliding_window"), + sliding_window: Some(128), + ..gpu(vec![h("k")], None, (1..=24).collect()) + } + .wire(layout), + // Raw digests, the parent given as an int, no extra keys on + // either block. + VllmStored { + extra_keys: Some(vec![None, None]), + ..gpu(vec![d1, d2], Some(h("b")), (9..=16).collect()) + } + .wire(layout), + ], + )); + want.extend([ + Want::Stored(WantStored { + group_idx: Some(0), + session_id: Some("req-1"), + ..WantStored::device( + vec![e("a"), e("b")], + vec![vec![1, 2, 3, 4], vec![5, 6, 7, 8]], + ) + }), + Want::Stored(WantStored { + parent: Some(e("b")), + group_idx: Some(0), + ..WantStored::device( + vec![d1.forwarded(), d2.forwarded()], + vec![vec![9, 10, 11, 12], vec![13, 14, 15, 16]], + ) + }), + ]); + + // Batch 1: a second physical copy, per-copy removals, offload tiers and + // every other drop rule. + batches.push(batch( + 1_700_000_001.0, + 0, + vec![ + gpu(vec![h("a"), h("b")], None, (1..=8).collect()).wire(layout), // duplicate copy + vllm_removed(layout, &[h("a")], "GPU", Some(0)), + vllm_removed(layout, &[h("a")], "GPU", Some(0)), // the other copy + // CPU offload placeholder: a chunk key, no tokens, block_size 0. + VllmStored { + medium: Some("CPU"), + block_size: 0, + spec_kind: None, + ..gpu(vec![h("c")], None, Vec::new()) + } + .wire(layout), + vllm_removed(layout, &[h("c")], "CPU", Some(0)), + VllmStored { + medium: Some("STORAGE"), + locality: Some("REMOTE"), + ..gpu(vec![h("d")], None, vec![1, 2, 3, 4]) + } + .wire(layout), + VllmStored { + medium: Some("STORAGE"), + locality: Some("LOCAL"), + ownership: Some("kvcr"), + ..gpu(vec![h("d")], None, vec![1, 2, 3, 4]) + } + .wire(layout), + VllmStored { + medium: Some("STORAGE"), + locality: Some("LOCAL"), + ..gpu(vec![h("d")], None, vec![1, 2, 3, 4]) + } + .wire(layout), + VllmStored { + medium: Some("MARS"), + ..gpu(vec![h("f")], None, vec![1, 2, 3, 4]) + } + .wire(layout), + gpu(vec![h("g")], None, vec![1, 2, 3, 4, 5, 6]).wire(layout), // unaligned + gpu(vec![h("i")], Some(h("i")), vec![1, 2, 3, 4]).wire(layout), // parent is itself + migrated(layout, &[h("j")]), + malformed_store(layout), + ], + )); + want.extend([ + Want::Stored(WantStored { + group_idx: Some(0), + ..WantStored::device( + vec![e("a"), e("b")], + vec![vec![1, 2, 3, 4], vec![5, 6, 7, 8]], + ) + }), + removed(vec![e("a")]), + removed(vec![e("a")]), + removed_on(vec![e("c")], KvCacheTier::Host, 1), + Want::Stored(WantStored { + tier: KvCacheTier::Disk, + cache_level: Some(2), + group_idx: Some(0), + ..WantStored::device(vec![e("d")], vec![vec![1, 2, 3, 4]]) + }), + ]); + + // Batch 2: the pool reset, then the chain again (not a duplicate any more). + batches.push(batch( + 1_700_000_002.0, + 0, + vec![ + cleared(layout), + gpu(vec![h("a"), h("b")], None, (1..=8).collect()).wire(layout), + ], + )); + want.extend([ + Want::Cleared { rank: 0 }, + Want::Stored(WantStored { + group_idx: Some(0), + ..WantStored::device( + vec![e("a"), e("b")], + vec![vec![1, 2, 3, 4], vec![5, 6, 7, 8]], + ) + }), + ]); + + // Batch 3: a LoRA request with a multimodal item, a cache salt and prompt + // embeddings; the salt rides in block 0's extra keys only and the child + // inherits it. + batches.push(batch( + 1_700_000_003.0, + 0, + vec![ + VllmStored { + lora_id: Some(7), + lora_name: Some("adapter"), + extra_keys: Some(vec![Some(vec![ + Value::from("adapter"), + Value::Array(vec![Value::from("mm-abc"), Value::from(0u64)]), + Value::from("salt-1"), + Value::from(embeds.to_vec()), + ])]), + session_id: Some("req-2"), + ..gpu(vec![h("e")], None, vec![1, 2, 3, 4]) + } + .wire(layout), + VllmStored { + lora_id: Some(7), + lora_name: Some("adapter"), + extra_keys: Some(vec![Some(vec![Value::from("adapter")])]), + session_id: Some("req-2"), + ..gpu(vec![h("h")], Some(h("e")), vec![5, 6, 7, 8]) + } + .wire(layout), + ], + )); + want.extend([ + Want::Stored(WantStored { + lora_name: Some("adapter"), + cache_salt: Some("salt-1"), + group_idx: Some(0), + session_id: Some("req-2"), + extra_keys: Some(vec![vec![ + Key::Text("adapter"), + Key::Multimodal("mm-abc", 0), + Key::Text("salt-1"), + Key::BlobLen(32), + ]]), + ..WantStored::device(vec![e("e")], vec![vec![1, 2, 3, 4]]) + }), + Want::Stored(WantStored { + parent: Some(e("e")), + lora_name: Some("adapter"), + cache_salt: Some("salt-1"), + group_idx: Some(0), + session_id: Some("req-2"), + extra_keys: Some(vec![vec![Key::Text("adapter")]]), + ..WantStored::device(vec![e("h")], vec![vec![5, 6, 7, 8]]) + }), + ]); + + // Batch 4: another DP rank stores the same hashes; seen-sets are per rank. + batches.push(batch( + 1_700_000_004.0, + 1, + vec![gpu(vec![h("a"), h("b")], None, (1..=8).collect()).wire(layout)], + )); + want.push(Want::Stored(WantStored { + rank: 1, + group_idx: Some(0), + ..WantStored::device( + vec![e("a"), e("b")], + vec![vec![1, 2, 3, 4], vec![5, 6, 7, 8]], + ) + })); + + Scenario { + name: format!("vllm-{}", layout.name()), + engine: "vllm", + layout, + batches, + want, + counts: WantCounts { + stored: 8, + removed: 3, + cleared: 1, + duplicate_stores: 1, + bigram_stores: 0, + dropped: [ + ("non_main_attention_group", 2), + ("placeholder", 1), + ("non_local_locality", 1), + ("unsupported_ownership", 1), + ("unknown_medium", 1), + ("unaligned_blocks", 1), + ("self_referencing_hashes", 1), + ("unknown_type", 1), + ("malformed", 1), + ] + .into_iter() + .collect(), + }, + } +} + +/// SGLang: the startup clear, a chain with a coalesced two-page store, +/// HiCache write-through (store on the host, demote, load back, evict the +/// host copy), a salted chain, an Eagle bigram page, the DISK and EXTERNAL +/// media, an unknown event type and a second attention DP rank. +fn sglang(layout: Layout) -> Scenario { + let s = |name: &str| sglang_hash(&format!("sglang-{name}")); + let e = |name: &str| s(name).forwarded(); + assert!( + "abcdefgh".chars().any(|name| e(&name.to_string()) < 0), + "pick labels with a negative i64" + ); + let gpu = SglangStored::gpu; + // The legacy array has no readable salt slot (it is vLLM's `lora_name` + // there), so that variant stores the chain unsalted; and it puts + // `session_id` where vLLM has `extra_keys`, where the relay cannot read + // it. + let (salt, session) = match layout { + Layout::Map => (Some("tenant-a"), Some("req-1")), + Layout::Array => (None, None), + }; + let mut batches = Vec::new(); + let mut want = Vec::new(); + + // Batch 0: the first batch after startup clears. + batches.push(batch(1_700_000_000.0, 0, vec![cleared(layout)])); + want.push(Want::Cleared { rank: 0 }); + + // Batch 1: a chain; the second store is coalesced over two pages. + batches.push(batch( + 1_700_000_001.0, + 0, + vec![ + SglangStored { + session_id: Some("req-1"), + ..gpu(vec![s("a")], None, &[1, 2, 3, 4]) + } + .wire(layout), + SglangStored { + session_id: Some("req-1"), + ..gpu( + vec![s("b"), s("c")], + Some(s("a")), + &(5..=12).collect::>(), + ) + } + .wire(layout), + ], + )); + want.extend([ + Want::Stored(WantStored { + session_id: session, + ..WantStored::device(vec![e("a")], vec![vec![1, 2, 3, 4]]) + }), + Want::Stored(WantStored { + parent: Some(e("a")), + session_id: session, + ..WantStored::device( + vec![e("b"), e("c")], + vec![vec![5, 6, 7, 8], vec![9, 10, 11, 12]], + ) + }), + ]); + + // Batch 2: HiCache write-through: back up to host, demote (device copy + // goes, host stays), load back, evict the host copy. + batches.push(batch( + 1_700_000_002.0, + 0, + vec![ + SglangStored { + medium: Some("CPU_PINNED"), + ..gpu(vec![s("a")], None, &[1, 2, 3, 4]) + } + .wire(layout), + sglang_removed(layout, &[s("a")], "GPU"), + gpu(vec![s("a")], None, &[1, 2, 3, 4]).wire(layout), + sglang_removed(layout, &[s("a")], "CPU_PINNED"), + ], + )); + want.extend([ + Want::Stored(WantStored { + tier: KvCacheTier::Host, + cache_level: Some(1), + ..WantStored::device(vec![e("a")], vec![vec![1, 2, 3, 4]]) + }), + removed(vec![e("a")]), + Want::Stored(WantStored::device(vec![e("a")], vec![vec![1, 2, 3, 4]])), + removed_on(vec![e("a")], KvCacheTier::Host, 1), + ]); + + // Batch 3: a salted request's chain. + batches.push(batch( + 1_700_000_003.0, + 0, + vec![ + SglangStored { + cache_salt: salt, + ..gpu(vec![s("d")], None, &[1, 2, 3, 4]) + } + .wire(layout), + SglangStored { + cache_salt: salt, + ..gpu(vec![s("e")], Some(s("d")), &[5, 6, 7, 8]) + } + .wire(layout), + ], + )); + want.extend([ + Want::Stored(WantStored { + cache_salt: salt, + ..WantStored::device(vec![e("d")], vec![vec![1, 2, 3, 4]]) + }), + Want::Stored(WantStored { + parent: Some(e("d")), + cache_salt: salt, + ..WantStored::device(vec![e("e")], vec![vec![5, 6, 7, 8]]) + }), + ]); + + // Batch 4: an Eagle bigram page, removed again in the same batch. (Its + // tokens differ from the plain chain's: two engine hashes with the same + // tokens at the same position share one index membership per worker.) + batches.push(batch( + 1_700_000_004.0, + 0, + vec![ + SglangStored { + tokens: bigrams(&[(21, 22), (22, 23), (23, 24), (24, 25)]), + ..gpu(vec![s("f")], None, &[]) + } + .wire(layout), + sglang_removed(layout, &[s("f")], "GPU"), + ], + )); + want.extend([ + Want::Stored(WantStored::device(vec![e("f")], vec![vec![21, 22, 23, 24]])), + removed(vec![e("f")]), + ]); + + // Batch 5: the tiers the default core never emits but defines; an event + // type the relay does not know. + batches.push(batch( + 1_700_000_005.0, + 0, + vec![ + SglangStored { + medium: Some("DISK"), + ..gpu(vec![s("g")], None, &[1, 2, 3, 4]) + } + .wire(layout), + SglangStored { + medium: Some("EXTERNAL"), + ..gpu(vec![s("h")], None, &[1, 2, 3, 4]) + } + .wire(layout), + migrated(layout, &[s("h")]), + ], + )); + want.extend([ + Want::Stored(WantStored { + tier: KvCacheTier::Disk, + cache_level: Some(2), + ..WantStored::device(vec![e("g")], vec![vec![1, 2, 3, 4]]) + }), + Want::Stored(WantStored { + tier: KvCacheTier::External, + cache_level: Some(3), + ..WantStored::device(vec![e("h")], vec![vec![1, 2, 3, 4]]) + }), + ]); + + // Batch 6: another attention DP rank; its batch carries its rank. + batches.push(batch( + 1_700_000_006.0, + 1, + vec![gpu(vec![s("a")], None, &[1, 2, 3, 4]).wire(layout)], + )); + want.push(Want::Stored(WantStored { + rank: 1, + ..WantStored::device(vec![e("a")], vec![vec![1, 2, 3, 4]]) + })); + + Scenario { + name: format!("sglang-{}", layout.name()), + engine: "sglang", + layout, + batches, + want, + counts: WantCounts { + stored: 10, + removed: 3, + cleared: 1, + duplicate_stores: 0, + bigram_stores: 1, + dropped: [("unknown_type", 1)].into_iter().collect(), + }, + } +} + +fn scenarios() -> [Scenario; 4] { + [ + vllm(Layout::Map), + vllm(Layout::Array), + sglang(Layout::Map), + sglang(Layout::Array), + ] +} + +fn run(scenario: &Scenario) { + let (batches, counts) = normalized(scenario); + check(scenario, &batches, &counts); +} + +#[test] +fn vllm_tagged_maps_normalize_as_expected() { + run(&vllm(Layout::Map)); +} + +#[test] +fn vllm_legacy_arrays_normalize_as_expected() { + run(&vllm(Layout::Array)); +} + +#[test] +fn sglang_tagged_maps_normalize_as_expected() { + run(&sglang(Layout::Map)); +} + +#[test] +fn sglang_legacy_arrays_normalize_as_expected() { + run(&sglang(Layout::Array)); +} + +/// The two layouts of one engine carry the same events, so the relay must +/// forward the same stream from either, down to the event ids; the one +/// difference is what the legacy array cannot say (SGLang's salt and +/// session, which the array scenario leaves out on both sides). +#[test] +fn both_layouts_forward_the_same_stream() { + for engine in [vllm, sglang] { + let (from_maps, map_counts) = normalized(&engine(Layout::Map)); + let (from_arrays, array_counts) = normalized(&engine(Layout::Array)); + assert_eq!(map_counts, array_counts); + let strip = |batches: Vec| -> Vec { + batches + .into_iter() + .map(|mut batch| { + for event in &mut batch.events { + if let Some(kv_cache_event::Data::Stored(stored)) = &mut event.data { + stored.session_id = None; + stored.cache_salt = None; + } + } + batch + }) + .collect() + }; + assert_eq!(strip(from_maps), strip(from_arrays)); + } +} + +// --------------------------------------------------------------------------- +// Round trip: wire -> proto -> index +// --------------------------------------------------------------------------- + +/// Apply normalized events to the index the way the gateway's monitor does +/// for the device tier: stores hashed under the event's namespace, removals +/// by engine hash, clears per worker, each rank its own worker. Host +/// residency is the monitor's own bookkeeping and has its tests there; here +/// host events are skipped. +fn index(batches: &[KvEventBatch]) -> ReferenceIndexer { + let mut indexer = ReferenceIndexer::new(); + for batch in batches { + let worker = u32::try_from(batch.dp_rank.unwrap_or(0)).expect("a rank"); + for event in &batch.events { + match &event.data { + Some(kv_cache_event::Data::Stored(stored)) => { + apply_stored(&mut indexer, worker, stored); + } + Some(kv_cache_event::Data::Removed(removed)) => { + if device_tier(removed.tier) { + let hashes: Vec = removed + .block_hashes + .iter() + .map(|&hash| SequenceHash::from(hash)) + .collect(); + indexer.apply_removed(worker, &hashes); + } + } + Some(kv_cache_event::Data::Cleared(_)) => indexer.apply_cleared(worker), + None => {} + } + } + } + indexer +} + +fn apply_stored(indexer: &mut ReferenceIndexer, worker: u32, stored: &KvBlocksStored) { + if !device_tier(stored.tier) { + return; + } + let seed = namespace_seed(stored.lora_name.as_deref(), stored.cache_salt.as_deref()); + let converted: Vec = stored + .blocks + .iter() + .map(|block| StoredBlock { + seq_hash: SequenceHash::from(block.block_hash), + content_hash: content_hash_with_seed(&block.token_ids, seed), + }) + .collect(); + let parent = stored.parent_block_hash.map(SequenceHash::from); + if let Err(ApplyError::WorkerNotTracked | ApplyError::ParentBlockNotFound) = + indexer.apply_stored(worker, &converted, parent) + { + indexer + .apply_stored(worker, &converted, None) + .expect("a parentless store applies"); + } +} + +fn device_tier(tier: Option) -> bool { + matches!( + KvCacheTier::try_from(tier.unwrap_or_default()), + Ok(KvCacheTier::Device | KvCacheTier::Unspecified) + ) +} + +/// The longest prefix of the request `worker` holds. +fn depth(indexer: &ReferenceIndexer, worker: u32, hashes: &[ContentHash]) -> u32 { + indexer + .find_matches(hashes) + .get(&worker) + .copied() + .unwrap_or(0) +} + +#[test] +fn the_normalized_streams_round_trip_into_the_reference_index() { + let tokens: Vec = (1..=16).collect(); + let plain = |n: usize| compute_request_content_hashes(&tokens[..n], 4); + for scenario in &scenarios() { + let (batches, _) = normalized(scenario); + let indexer = index(&batches); + let name = &scenario.name; + match scenario.engine { + "vllm" => { + // The chain of two device blocks survives the duplicate copy, + // its per-copy removals and the reset (it is stored again). + assert_eq!(depth(&indexer, 0, &plain(8)), 2, "{name}: plain chain"); + // The digest-hashed continuation was cleared and not restored. + assert_eq!(depth(&indexer, 0, &plain(16)), 2, "{name}: cleared tail"); + // The LoRA + salt chain matches only under its namespace, the + // child included (it inherited the salt from block 0). + let salted = namespaced_request_content_hashes( + &tokens[..8], + 4, + Some("adapter"), + Some("salt-1"), + ); + assert_eq!(depth(&indexer, 0, &salted), 2, "{name}: salted chain"); + let lora_only = + namespaced_request_content_hashes(&tokens[..8], 4, Some("adapter"), None); + assert_eq!( + depth(&indexer, 0, &lora_only), + 0, + "{name}: lora without salt" + ); + let salt_only = + namespaced_request_content_hashes(&tokens[..8], 4, None, Some("salt-1")); + assert_eq!( + depth(&indexer, 0, &salt_only), + 0, + "{name}: salt without lora" + ); + assert_eq!(depth(&indexer, 1, &plain(8)), 2, "{name}: rank 1"); + } + "sglang" => { + // Device chain of three pages; the HiCache demote and host + // eviction left the load-back copy in place. + assert_eq!(depth(&indexer, 0, &plain(12)), 3, "{name}: plain chain"); + if scenario.layout == Layout::Map { + let salted = + namespaced_request_content_hashes(&tokens[..8], 4, None, Some("tenant-a")); + assert_eq!(depth(&indexer, 0, &salted), 2, "{name}: salted chain"); + let other = + namespaced_request_content_hashes(&tokens[..8], 4, None, Some("tenant-b")); + assert_eq!(depth(&indexer, 0, &other), 0, "{name}: other salt"); + } + assert_eq!(depth(&indexer, 1, &plain(4)), 1, "{name}: rank 1"); + } + other => panic!("{name}: unknown engine {other}"), + } + } +} diff --git a/crates/engine_servicer/src/lib.rs b/crates/engine_servicer/src/lib.rs index 70c145571b..e0f47c020c 100644 --- a/crates/engine_servicer/src/lib.rs +++ b/crates/engine_servicer/src/lib.rs @@ -10,16 +10,29 @@ //! — it launches the headless engine and drives the server through the PyO3 //! binding — which is why the server runs on its own thread and reports back //! through plain flags instead of a Python-visible runtime. +//! +//! The engines' ZMQ KV-cache event publishers are relayed by [`kv_events`] +//! after [`kv_wire`] normalizes them; that module documents the per-engine +//! hash folding rule and the one-for-one forwarding of stores and removals. +//! [`kv_state`] keeps the engine's live blocks from that stream, the state +//! snapshot a subscriber receives once the relay's history has rolled. +pub mod engine_hash; mod engine_link; mod error; mod health; -mod kv_events; +pub mod kv_events; +pub mod kv_history; +pub mod kv_state; +pub mod kv_wire; +mod load_tracker; mod proto_json; mod requests; mod server; pub mod sglang; mod stop_match; +#[cfg(test)] +mod testing; mod tokenizer_bundle; pub mod tokenspeed; pub mod vllm; diff --git a/crates/engine_servicer/src/load_tracker.rs b/crates/engine_servicer/src/load_tracker.rs new file mode 100644 index 0000000000..155b152b7a --- /dev/null +++ b/crates/engine_servicer/src/load_tracker.rs @@ -0,0 +1,283 @@ +//! What a servicer can tell about its engine's load from the requests it +//! forwards: queued token-work, generation throughput and the prefix-cache +//! hit rate. The gateway's expected-wait score reads them from `GetLoads`; +//! an engine whose stats do not carry them (vLLM reports running, waiting +//! and KV usage) would otherwise be scored on defaults and routed unlike its +//! peers. +//! +//! The servicer sees every request it submits, the engine's first output for +//! it (with the prompt and cached token counts) and every token streamed, so: +//! +//! - queued token-work is the uncached prompt tokens of the requests the +//! engine has not started: of the submitted requests without a first +//! output, the youngest `num_waiting_reqs` (the engine admits FCFS, so the +//! older ones are the ones in prefill), each prompt discounted by the +//! recent hit rate since its own cached count is not known until it starts; +//! - generation throughput is the tokens streamed over the last +//! [`THROUGHPUT_WINDOW`]; +//! - the hit rate is cached over prompt tokens across the last +//! [`HIT_RATE_SAMPLES`] first outputs. +//! +//! These are the field semantics the SGLang servicer reports from its +//! scheduler and the mock engine from its queue. + +use std::{ + collections::VecDeque, + sync::{Mutex, MutexGuard, PoisonError}, + time::{Duration, Instant}, +}; + +/// How far back streamed tokens count toward the throughput. +pub(crate) const THROUGHPUT_WINDOW: Duration = Duration::from_secs(2); +/// How many first outputs the hit rate averages over. +pub(crate) const HIT_RATE_SAMPLES: usize = 64; + +/// The three fields, as `GetLoads` reports them. +#[derive(Clone, Copy, Debug, Default, PartialEq)] +pub(crate) struct LoadEstimate { + /// Uncached prompt tokens of the requests the engine has not started. + pub queued_token_work: i32, + /// Tokens streamed per second over the window. + pub gen_throughput: f64, + /// Cached over prompt tokens of recent first outputs, in `[0, 1]`. + pub cache_hit_rate: f64, +} + +struct Pending { + request_id: String, + prompt_tokens: u32, +} + +#[derive(Default)] +struct Inner { + /// Submitted requests without a first output yet, oldest first. + pending: VecDeque, + /// Recent first outputs: (prompt tokens, cached tokens). + hits: VecDeque<(u64, u64)>, + prompt_sum: u64, + cached_sum: u64, + /// Tokens streamed, by time, within the window. + generated: VecDeque<(Instant, u64)>, + generated_sum: u64, +} + +impl Inner { + fn trim(&mut self, now: Instant) { + while let Some(&(at, tokens)) = self.generated.front() { + if now.duration_since(at) <= THROUGHPUT_WINDOW { + break; + } + self.generated.pop_front(); + self.generated_sum -= tokens; + } + } + + fn hit_rate(&self) -> f64 { + if self.prompt_sum == 0 { + 0.0 + } else { + (self.cached_sum as f64 / self.prompt_sum as f64).clamp(0.0, 1.0) + } + } +} + +/// Per-servicer load bookkeeping; every method is cheap and lock-scoped. +#[derive(Default)] +pub(crate) struct LoadTracker { + inner: Mutex, +} + +impl LoadTracker { + fn lock(&self) -> MutexGuard<'_, Inner> { + self.inner.lock().unwrap_or_else(PoisonError::into_inner) + } + + /// A request with `prompt_tokens` was handed to the engine. + pub(crate) fn submitted(&self, request_id: &str, prompt_tokens: u32) { + self.lock().pending.push_back(Pending { + request_id: request_id.to_string(), + prompt_tokens, + }); + } + + /// The engine's first output for a request: it has started, and its + /// prompt had `cached_tokens` of `prompt_tokens` in the prefix cache. + pub(crate) fn first_output(&self, request_id: &str, prompt_tokens: u32, cached_tokens: u32) { + let mut inner = self.lock(); + if let Some(index) = inner + .pending + .iter() + .position(|pending| pending.request_id == request_id) + { + inner.pending.remove(index); + } + if prompt_tokens == 0 { + return; + } + let sample = ( + u64::from(prompt_tokens), + u64::from(cached_tokens.min(prompt_tokens)), + ); + inner.hits.push_back(sample); + inner.prompt_sum += sample.0; + inner.cached_sum += sample.1; + if inner.hits.len() > HIT_RATE_SAMPLES { + if let Some((prompt, cached)) = inner.hits.pop_front() { + inner.prompt_sum -= prompt; + inner.cached_sum -= cached; + } + } + } + + /// `tokens` were streamed to the client just now. + pub(crate) fn generated(&self, tokens: u32) { + self.generated_at(tokens, Instant::now()); + } + + pub(crate) fn generated_at(&self, tokens: u32, at: Instant) { + if tokens == 0 { + return; + } + let mut inner = self.lock(); + inner.generated.push_back((at, u64::from(tokens))); + inner.generated_sum += u64::from(tokens); + inner.trim(at); + } + + /// A request ended (or was aborted) without the engine starting it. + pub(crate) fn finished(&self, request_id: &str) { + let mut inner = self.lock(); + if let Some(index) = inner + .pending + .iter() + .position(|pending| pending.request_id == request_id) + { + inner.pending.remove(index); + } + } + + /// The estimate now, given the engine's own count of waiting requests. + pub(crate) fn estimate(&self, num_waiting_reqs: i32) -> LoadEstimate { + self.estimate_at(num_waiting_reqs, Instant::now()) + } + + pub(crate) fn estimate_at(&self, num_waiting_reqs: i32, now: Instant) -> LoadEstimate { + let mut inner = self.lock(); + inner.trim(now); + let hit_rate = inner.hit_rate(); + let waiting = usize::try_from(num_waiting_reqs) + .unwrap_or(0) + .min(inner.pending.len()); + let queued: f64 = inner + .pending + .iter() + .rev() + .take(waiting) + .map(|pending| f64::from(pending.prompt_tokens) * (1.0 - hit_rate)) + .sum(); + LoadEstimate { + queued_token_work: queued.round().clamp(0.0, f64::from(i32::MAX)) as i32, + gen_throughput: inner.generated_sum as f64 / THROUGHPUT_WINDOW.as_secs_f64(), + cache_hit_rate: hit_rate, + } + } + + /// Submitted requests the engine has not started. + #[cfg(test)] + pub(crate) fn pending(&self) -> usize { + self.lock().pending.len() + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn queued_token_work_is_the_youngest_waiting_prompts_discounted_by_the_hit_rate() { + let tracker = LoadTracker::default(); + let now = Instant::now(); + assert_eq!(tracker.estimate_at(5, now), LoadEstimate::default()); + + // Three in flight, none started: the engine says two are waiting, so + // the oldest is in prefill and the two youngest are queued work. + tracker.submitted("a", 1_000); + tracker.submitted("b", 300); + tracker.submitted("c", 200); + assert_eq!(tracker.estimate_at(2, now).queued_token_work, 500); + assert_eq!(tracker.estimate_at(0, now).queued_token_work, 0); + assert_eq!( + tracker.estimate_at(10, now).queued_token_work, + 1_500, + "never more than we hold" + ); + + // The first output of `a`: half its prompt was cached. Queued prompts + // are discounted by that rate. + tracker.first_output("a", 1_000, 500); + assert_eq!(tracker.pending(), 2); + let estimate = tracker.estimate_at(2, now); + assert_eq!(estimate.cache_hit_rate, 0.5); + assert_eq!(estimate.queued_token_work, 250); + + // A request that ends before starting leaves the queue. + tracker.finished("b"); + assert_eq!(tracker.estimate_at(2, now).queued_token_work, 100); + tracker.first_output("c", 200, 200); + assert_eq!(tracker.pending(), 0); + assert_eq!(tracker.estimate_at(3, now).queued_token_work, 0); + // 700 cached of 1_200 prompt tokens. + assert!((tracker.estimate_at(0, now).cache_hit_rate - 700.0 / 1_200.0).abs() < 1e-9); + } + + #[test] + fn throughput_counts_tokens_inside_the_window_only() { + let tracker = LoadTracker::default(); + let start = Instant::now(); + tracker.generated_at(1_000, start); + tracker.generated_at(3_000, start + Duration::from_secs(1)); + let per_second = |estimate: LoadEstimate| estimate.gen_throughput; + assert_eq!( + per_second(tracker.estimate_at(0, start + Duration::from_secs(1))), + 2_000.0 + ); + // Two seconds after the first sample it falls out of the window. + assert_eq!( + per_second(tracker.estimate_at(0, start + Duration::from_millis(2_500))), + 1_500.0 + ); + assert_eq!( + per_second(tracker.estimate_at(0, start + Duration::from_secs(4))), + 0.0 + ); + tracker.generated_at(0, start + Duration::from_secs(4)); + assert_eq!( + per_second(tracker.estimate_at(0, start + Duration::from_secs(4))), + 0.0 + ); + } + + #[test] + fn the_hit_rate_averages_recent_first_outputs_and_ignores_empty_prompts() { + let tracker = LoadTracker::default(); + let now = Instant::now(); + tracker.submitted("x", 0); + tracker.first_output("x", 0, 0); + assert_eq!(tracker.estimate_at(0, now).cache_hit_rate, 0.0); + for index in 0..HIT_RATE_SAMPLES { + tracker.first_output(&format!("old-{index}"), 100, 0); + } + assert_eq!(tracker.estimate_at(0, now).cache_hit_rate, 0.0); + for index in 0..HIT_RATE_SAMPLES { + tracker.first_output(&format!("new-{index}"), 100, 100); + } + assert_eq!( + tracker.estimate_at(0, now).cache_hit_rate, + 1.0, + "the old samples aged out" + ); + // Cached can never exceed the prompt. + tracker.first_output("odd", 10, 50); + assert!(tracker.estimate_at(0, now).cache_hit_rate <= 1.0); + } +} diff --git a/crates/engine_servicer/src/sglang/engine.rs b/crates/engine_servicer/src/sglang/engine.rs index 4dd10b7111..402a57a55b 100644 --- a/crates/engine_servicer/src/sglang/engine.rs +++ b/crates/engine_servicer/src/sglang/engine.rs @@ -4,7 +4,7 @@ use std::{sync::Arc, time::Duration}; -use engine_zmq_adapter::{connect_with_eos, EosTokenIds}; +use engine_zmq_adapter::{connect_with_eos, EosTokenIds, Handshake}; use openai_protocol::worker::RuntimeType; use tracing::{error, info}; @@ -29,7 +29,7 @@ pub(super) async fn connect_engine( &ipc_base_url, state.model.model_path.clone(), RuntimeType::Sglang, - Some(&handshake_address), + Handshake::TcpOrIpc(&handshake_address), engine_count, eos, startup_timeout, diff --git a/crates/engine_servicer/src/sglang/mod.rs b/crates/engine_servicer/src/sglang/mod.rs index 2b929242f3..e38a8f28ea 100644 --- a/crates/engine_servicer/src/sglang/mod.rs +++ b/crates/engine_servicer/src/sglang/mod.rs @@ -83,6 +83,16 @@ pub struct SglangModelInfo { /// The data-parallel size the launcher configured; the handshake's figure /// wins once the engines are up. pub data_parallel_size: i32, + /// SGLang's `--kv-events-config` ZMQ publisher endpoint as configured + /// (`tcp://*:5557` style, rank 0's port), or empty when events are off; + /// `SubscribeKvEvents` relays it through [`crate::kv_wire`]'s rules. + pub kv_events_endpoint: String, + /// The same config's `replay_endpoint` (SGLang's replay ROUTER), or + /// empty when it runs none: what the relay asks for gaps in flight and + /// for the batches published before its subscription joined. + pub kv_events_replay_endpoint: String, + /// The publisher's topic prefix (empty by default). + pub kv_events_topic: String, } /// How to bind, where the engines dial in, and what to advertise. @@ -95,6 +105,8 @@ pub struct SglangServicerConfig { pub ipc_base_url: String, /// `tcp://host:port` the headless scheduler dials for the handshake (the /// launcher's `--zmq-handshake-address`). + /// An `ipc://` endpoint is accepted too: one per test, so parallel + /// tests never share a probed TCP port. pub handshake_address: String, /// Engines that will dial in (the data-parallel size). pub engine_count: usize, @@ -107,8 +119,26 @@ pub struct SglangServicerConfig { pub engine_startup_timeout: Duration, } +/// The servicer's `GetLoads` figures, read for the load record the relay +/// attaches to every batch it streams. +struct LoadFromState(std::sync::Weak); + +impl crate::kv_events::LoadSource for LoadFromState { + fn load(&self, dp_rank: Option) -> Option { + let state = self.0.upgrade()?; + let response = info::loads(&state, dp_rank).ok()?; + let load = response.loads.first()?; + // The engine reports no queued token-work on this wire yet: the + // record's `waiting_uncached_tokens` stays unset. + Some(smg_grpc_client::common_proto::EngineLoad::from(load)) + } +} + pub(super) struct State { pub(super) model: SglangModelInfo, + /// The KV-event relay for the engine's ZMQ publisher; `None` when events + /// are off (`SubscribeKvEvents` is then UNIMPLEMENTED). + pub(super) kv_relay: Option>, /// The local tokenizer directory `GetTokenizer` bundles; `None` when none /// resolved. pub(super) tokenizer_dir: Option, @@ -173,8 +203,12 @@ impl SglangServicerServer { if !config.ipc_base_url.starts_with("ipc://") { return Err(invalid("ipc_base_url must be ipc://")); } - if !config.handshake_address.starts_with("tcp://") { - return Err(invalid("handshake_address must be tcp://host:port")); + if !config.handshake_address.starts_with("tcp://") + && !config.handshake_address.starts_with("ipc://") + { + return Err(invalid( + "handshake_address must be tcp://host:port or ipc://", + )); } if config.engine_count == 0 { return Err(invalid("engine_count must be positive")); @@ -195,15 +229,24 @@ impl SglangServicerServer { return Err(invalid(&format!("{name} must be a JSON object"))); } } + let kv_relay = crate::kv_events::KvEventRelay::for_publisher( + &config.model.kv_events_endpoint, + Some(&config.model.kv_events_replay_endpoint), + &config.model.kv_events_topic, + ); let state = Arc::new(State { tokenizer_dir: config.tokenizer_dir.clone(), model: config.model, + kv_relay, engine: EngineLink::default(), registry: Arc::new(RequestRegistry::default()), serving: AtomicBool::new(true), started: Instant::now(), started_at: SystemTime::now(), }); + if let Some(relay) = &state.kv_relay { + relay.set_load_source(Arc::new(LoadFromState(Arc::downgrade(&state)))); + } let service = service::SglangService { state: Arc::clone(&state), }; @@ -213,6 +256,7 @@ impl SglangServicerServer { Arc::new(move || health_state.is_serving()), ); let connect_state = Arc::clone(&state); + let kv_relay = state.kv_relay.clone(); let SglangServicerConfig { bind_address, ipc_base_url, @@ -225,6 +269,12 @@ impl SglangServicerServer { "smg-sglang-servicer", &bind_address, move |listener: TcpListener, shutdown: Shutdown, last_error| async move { + // The KV-event relay follows the publisher from the start, + // before any gateway asks, so its history and live-block + // record cover the engine's whole life. + if let Some(relay) = &kv_relay { + relay.start_at_boot(); + } #[expect( clippy::disallowed_methods, reason = "engine connect is fire-and-forget; the runtime drop cancels it" diff --git a/crates/engine_servicer/src/sglang/service.rs b/crates/engine_servicer/src/sglang/service.rs index 561fb95722..a05dffe374 100644 --- a/crates/engine_servicer/src/sglang/service.rs +++ b/crates/engine_servicer/src/sglang/service.rs @@ -1,6 +1,7 @@ //! `sglang.grpc.scheduler.SglangScheduler` over the shared state: each RPC //! delegates to its handler module. What the Python servicer does not serve -//! either (LoRA loading, the KV-event relay) answers UNIMPLEMENTED here. +//! either (LoRA loading) answers UNIMPLEMENTED here; the scheduler's KV-event +//! publisher is relayed as for the other engines. use std::sync::Arc; @@ -11,7 +12,7 @@ use smg_grpc_client::{ use tonic::{Request, Response, Status}; use super::{admin, embed, generate, info, State}; -use crate::{tokenizer_bundle, BoxStream}; +use crate::{kv_events, tokenizer_bundle, BoxStream}; /// What neither servicer serves: the msgpack wire has no adapter-loading /// message, and the Python servicer does not implement these RPCs either. @@ -141,12 +142,12 @@ impl SglangScheduler for SglangService { async fn subscribe_kv_events( &self, - _request: Request, + request: Request, ) -> Result, Status> { - Err(Status::unimplemented( - "SubscribeKvEvents is not available through the Rust SGLang servicer yet: the \ - scheduler's KV-event publisher is not relayed on this path", - )) + let Some(relay) = &self.state.kv_relay else { + return Err(Status::unimplemented(kv_events::SGLANG_DISABLED_MESSAGE)); + }; + relay.subscribe(request.into_inner()).map(Response::new) } async fn load_lo_ra_adapter( diff --git a/crates/engine_servicer/src/sglang/tests.rs b/crates/engine_servicer/src/sglang/tests.rs index 1aad0e3df8..dd62882cdf 100644 --- a/crates/engine_servicer/src/sglang/tests.rs +++ b/crates/engine_servicer/src/sglang/tests.rs @@ -8,14 +8,16 @@ use std::time::Duration; use bytes::Bytes; use engine_zmq_client::{ codec::{decode_msgpack, encode_msgpack, OpaqueValue}, - mock_engine::{connect_to_frontend, default_ready_response, MockEngineInput, MockEngineOutput}, + mock_engine::{ + connect_to_frontend, default_ready_response, MockEngineInput, MockEngineOutput, + MOCK_DEADLINE, + }, protocol::sglang::{ output::{BatchEmbeddingSlimOutput, BatchTokenIDSlimOutput, ControlReplySlim, MatchedStop}, request::{SglangRequestType, TokenizedEmbeddingReqInput, TokenizedGenerateReqInput}, }, EngineId, }; -use portpicker::pick_unused_port; use prost_types::value::Kind; use smg_grpc_client::{ common_proto as common, @@ -27,7 +29,7 @@ use tonic_health::pb::{ }; use super::*; -use crate::ServicerError; +use crate::{testing::Bounded, ServicerError}; fn model_info() -> SglangModelInfo { SglangModelInfo { @@ -64,11 +66,21 @@ fn config(dir: &std::path::Path, handshake: &str, model: SglangModelInfo) -> Sgl } } -fn handshake_address() -> String { - format!( - "tcp://127.0.0.1:{}", - pick_unused_port().expect("a free handshake port") - ) +/// The handshake endpoint of one test: an IPC socket under its own +/// directory. A probed TCP port is not reserved, so two tests running in +/// parallel could pick the same one and a mock engine would handshake with +/// the other test's servicer and wait forever for its INIT. +fn handshake_address(dir: &std::path::Path) -> String { + format!("ipc://{}", dir.join("handshake").display()) +} + +/// The port a socket bound to `tcp://127.0.0.1:0` was given. +fn bound_port(endpoint: &str) -> u16 { + endpoint + .rsplit(':') + .next() + .and_then(|port| port.parse().ok()) + .expect("a bound tcp endpoint ends with its port") } async fn wait_until(mut condition: impl FnMut() -> bool) { @@ -81,6 +93,17 @@ async fn wait_until(mut condition: impl FnMut() -> bool) { panic!("condition not met within 10s"); } +/// A channel to the servicer whose every request fails after [`MOCK_DEADLINE`] +/// instead of waiting on a servicer that never answers. +async fn grpc_channel(address: impl std::fmt::Display) -> Channel { + Channel::from_shared(format!("http://{address}")) + .expect("grpc address") + .timeout(MOCK_DEADLINE) + .connect() + .await + .expect("grpc client") +} + /// A bound servicer, a handshaken mock scheduler, and a gRPC client. struct Harness { server: SglangServicerServer, @@ -92,7 +115,7 @@ struct Harness { async fn harness(model: SglangModelInfo) -> Harness { let dir = tempfile::tempdir().unwrap(); - let handshake = handshake_address(); + let handshake = handshake_address(dir.path()); let server = SglangServicerServer::start(config(dir.path(), &handshake, model)) .expect("servicer starts"); let engine = connect_to_frontend( @@ -103,9 +126,7 @@ async fn harness(model: SglangModelInfo) -> Harness { .await .expect("mock scheduler handshake"); wait_until(|| server.engine_ready()).await; - let client = SglangSchedulerClient::connect(format!("http://{}", server.address())) - .await - .expect("grpc client"); + let client = SglangSchedulerClient::new(grpc_channel(server.address()).await); let (engine_in, engine_out) = engine.split(); Harness { server, @@ -249,7 +270,7 @@ fn config_is_validated_before_binding() { #[tokio::test] async fn health_follows_the_engine_link_and_the_drain_flag() { let dir = tempfile::tempdir().unwrap(); - let handshake = handshake_address(); + let handshake = handshake_address(dir.path()); let server = SglangServicerServer::start(config(dir.path(), &handshake, model_info())) .expect("servicer starts"); let channel = Channel::from_shared(format!("http://{}", server.address())) @@ -329,22 +350,22 @@ async fn streaming_generate_maps_steps_to_chunks_and_a_complete() { ) .await; - let first = stream.message().await.unwrap().unwrap(); + let first = stream.message().bounded().await.unwrap().unwrap(); assert_eq!(first.request_id, "r1"); let first = chunk(first); assert_eq!(first.token_ids, vec![10]); assert_eq!(first.prompt_tokens, 3); assert_eq!( - chunk(stream.message().await.unwrap().unwrap()).token_ids, + chunk(stream.message().bounded().await.unwrap().unwrap()).token_ids, vec![11] ); - let done = complete(stream.message().await.unwrap().unwrap()); + let done = complete(stream.message().bounded().await.unwrap().unwrap()); assert_eq!(done.finish_reason, "length"); assert_eq!(done.output_ids, vec![10, 11]); assert_eq!(done.prompt_tokens, 3); assert_eq!(done.completion_tokens, 2); assert_eq!(done.index, 0); - assert!(stream.message().await.unwrap().is_none()); + assert!(stream.message().bounded().await.unwrap().is_none()); h.server.stop(Duration::from_secs(5)).expect("clean stop"); } @@ -365,10 +386,10 @@ async fn non_streaming_generate_delivers_only_the_complete() { &batch("r2", vec![11, 12], 3, Some("stop"), None), ) .await; - let done = complete(stream.message().await.unwrap().unwrap()); + let done = complete(stream.message().bounded().await.unwrap().unwrap()); assert_eq!(done.finish_reason, "stop"); assert_eq!(done.output_ids, vec![10, 11, 12]); - assert!(stream.message().await.unwrap().is_none()); + assert!(stream.message().bounded().await.unwrap().is_none()); h.server.stop(Duration::from_secs(5)).expect("clean stop"); } @@ -391,7 +412,7 @@ async fn string_stops_reach_the_scheduler_and_its_match_comes_back() { let mut done = batch("r3", vec![10, 11], 2, Some("stop"), None); done.finished_matched = vec![Some(MatchedStop::Text("###".to_string()))]; send(&mut h.engine_out, &done).await; - let done = complete(stream.message().await.unwrap().unwrap()); + let done = complete(stream.message().bounded().await.unwrap().unwrap()); assert_eq!(done.finish_reason, "stop"); assert_eq!( done.matched_stop, @@ -416,7 +437,7 @@ async fn abort_rpc_ends_the_stream_and_reaches_the_scheduler() { recv_add(&mut h.engine_in).await; send(&mut h.engine_out, &batch("r4", vec![10], 1, None, None)).await; assert_eq!( - chunk(stream.message().await.unwrap().unwrap()).token_ids, + chunk(stream.message().bounded().await.unwrap().unwrap()).token_ids, vec![10] ); @@ -430,10 +451,10 @@ async fn abort_rpc_ends_the_stream_and_reaches_the_scheduler() { .unwrap() .into_inner(); assert!(response.success); - let done = complete(stream.message().await.unwrap().unwrap()); + let done = complete(stream.message().bounded().await.unwrap().unwrap()); assert_eq!(done.finish_reason, "abort"); assert_eq!(done.output_ids, vec![10]); - assert!(stream.message().await.unwrap().is_none()); + assert!(stream.message().bounded().await.unwrap().is_none()); assert_eq!(recv_abort(&mut h.engine_in).await, vec!["r4".to_string()]); let unknown = h @@ -507,8 +528,8 @@ async fn n2_fans_out_under_the_parent_id() { }; send(&mut h.engine_out, &both).await; let mut completes = [ - complete(stream.message().await.unwrap().unwrap()), - complete(stream.message().await.unwrap().unwrap()), + complete(stream.message().bounded().await.unwrap().unwrap()), + complete(stream.message().bounded().await.unwrap().unwrap()), ]; completes.sort_by_key(|done| done.index); assert_eq!( @@ -519,7 +540,7 @@ async fn n2_fans_out_under_the_parent_id() { (completes[1].index, &completes[1].output_ids), (1, &vec![11]) ); - assert!(stream.message().await.unwrap().is_none()); + assert!(stream.message().bounded().await.unwrap().is_none()); h.server.stop(Duration::from_secs(5)).expect("clean stop"); } @@ -541,11 +562,11 @@ async fn logprobs_and_ranked_candidates_pass_through() { step.output_top_logprobs_val = vec![vec![vec![-0.5, -1.5]]]; step.output_top_logprobs_idx = vec![vec![vec![10, 12]]]; send(&mut h.engine_out, &step).await; - let first = chunk(stream.message().await.unwrap().unwrap()); + let first = chunk(stream.message().bounded().await.unwrap().unwrap()); let logprobs = first.output_logprobs.expect("chunk logprobs"); assert_eq!(logprobs.token_logprobs, vec![-0.5]); assert_eq!(logprobs.top_logprobs[0].token_ids, vec![10, 12]); - let done = complete(stream.message().await.unwrap().unwrap()); + let done = complete(stream.message().bounded().await.unwrap().unwrap()); let logprobs = done.output_logprobs.expect("complete logprobs"); assert_eq!(logprobs.token_ids, vec![10]); assert_eq!(logprobs.top_logprobs.len(), 1); @@ -568,7 +589,7 @@ async fn a_scheduler_refusal_is_the_callers_error() { refusal.finished_messages = vec![Some("n=2 is not served on this wire".to_string())]; refusal.finished_status = vec![Some(400)]; send(&mut h.engine_out, &refusal).await; - let error = stream.message().await.expect_err("a status"); + let error = stream.message().bounded().await.expect_err("a status"); assert_eq!(error.code(), Code::InvalidArgument); assert!(error.message().contains("n=2 is not served"), "{error}"); h.server.stop(Duration::from_secs(5)).expect("clean stop"); @@ -657,7 +678,7 @@ async fn info_rpcs_report_launcher_facts_and_handshake_figures() { &batch("r8", vec![10], 1, None, Some((1, 2, 50, 1000))), ) .await; - chunk(stream.message().await.unwrap().unwrap()); + chunk(stream.message().bounded().await.unwrap().unwrap()); let busy = h .client .get_loads(sg::GetLoadsRequest { @@ -887,7 +908,7 @@ async fn prompt_logprobs_and_reasoning_tokens_pass_through() { ..batch("p1", vec![10], 1, None, None) }; send(&mut h.engine_out, &first).await; - let chunk1 = chunk(stream.message().await.unwrap().unwrap()); + let chunk1 = chunk(stream.message().bounded().await.unwrap().unwrap()); assert_eq!(chunk1.reasoning_tokens, 1); let input = chunk1 .input_logprobs @@ -902,19 +923,278 @@ async fn prompt_logprobs_and_reasoning_tokens_pass_through() { ..batch("p1", vec![11], 2, Some("stop"), None) }; send(&mut h.engine_out, &last).await; - let chunk2 = chunk(stream.message().await.unwrap().unwrap()); + let chunk2 = chunk(stream.message().bounded().await.unwrap().unwrap()); assert!( chunk2.input_logprobs.is_none(), "prompt logprobs go out once" ); assert_eq!(chunk2.reasoning_tokens, 2); - let done = complete(stream.message().await.unwrap().unwrap()); + let done = complete(stream.message().bounded().await.unwrap().unwrap()); assert_eq!(done.reasoning_tokens, 2); assert_eq!( done.input_logprobs.map(|input| input.token_ids), Some(vec![1, 2, 3]) ); - assert!(stream.message().await.unwrap().is_none()); + assert!(stream.message().bounded().await.unwrap().is_none()); + h.server.stop(Duration::from_secs(5)).expect("clean stop"); +} + +/// With a publisher configured, `SubscribeKvEvents` relays it: the call +/// resolves before any event, batches arrive under the publisher's sequence +/// numbers, and stopping the servicer closes the subscription. +#[tokio::test] +async fn subscribe_kv_events_relays_a_configured_publisher() { + use zeromq::{prelude::*, PubSocket, SocketEvent}; + + use crate::kv_events::golden; + + let mut publisher = PubSocket::new(); + let mut monitor = publisher.monitor(); + let port = bound_port( + &publisher + .bind("tcp://127.0.0.1:0") + .await + .expect("publisher binds") + .to_string(), + ); + let mut model = model_info(); + // A bind wildcard, as SGLang's config spells it; the relay resolves it. + model.kv_events_endpoint = format!("tcp://*:{port}"); + model.kv_events_topic = "kv".to_string(); + let mut h = harness(model).await; + let mut stream = tokio::time::timeout( + Duration::from_secs(5), + h.client + .subscribe_kv_events(common::SubscribeKvEventsRequest::default()), + ) + .await + .expect("the call resolves before any event is published") + .expect("subscribe") + .into_inner(); + + // The subscription reaches the publisher a moment after the connect; + // probe with sequence 0 until a batch comes through. + let batch1 = golden::bytes(golden::BATCH1); + let mut first = None; + for _ in 0..200 { + publisher + .send(golden::frame(b"kv", 0, &batch1)) + .await + .expect("publish"); + if let Ok(item) = tokio::time::timeout(Duration::from_millis(50), stream.message()).await { + first = Some(item.expect("stream open").expect("a batch")); + break; + } + } + let first = first.expect("the subscription went live"); + assert_eq!(first.sequence_number, 0); + assert_eq!(first.events.len(), 4); + + // The relay keeps its publisher subscription for the servicer's + // lifetime (its history outlives any one stream); stopping the servicer + // closes it. + drop(stream); + h.server.stop(Duration::from_secs(5)).expect("clean stop"); + let disconnected = tokio::time::timeout(Duration::from_secs(5), async { + while let Some(event) = monitor.next().bounded().await { + if matches!(event, SocketEvent::Disconnected(_)) { + return true; + } + } + false + }) + .await + .expect("the publisher notices the stopped servicer in time"); + assert!(disconnected); +} + +/// The relay follows the publisher from the servicer's start, before any +/// gateway subscribes: batches published with nobody listening are in its +/// history, and the first subscription gets them as the engine's whole state. +#[tokio::test] +async fn the_relay_subscribes_to_the_publisher_at_boot_before_any_gateway() { + use zeromq::{prelude::*, PubSocket}; + + use crate::kv_events::golden; + + let mut publisher = PubSocket::new(); + let port = bound_port( + &publisher + .bind("tcp://127.0.0.1:0") + .await + .expect("publisher binds") + .to_string(), + ); + let mut model = model_info(); + model.kv_events_endpoint = format!("tcp://*:{port}"); + model.kv_events_topic = "kv".to_string(); + let mut h = harness(model).await; + let relay = h + .server + .state + .kv_relay + .clone() + .expect("a relay for the publisher"); + // The SUB connect is asynchronous: publish sequence 0 until the relay, + // with no subscriber of its own yet, has taken it (repeats are duplicates). + let batch1 = golden::bytes(golden::BATCH1); + for _ in 0..250 { + publisher + .send(golden::frame(b"kv", 0, &batch1)) + .await + .expect("publish"); + if relay.counts().relayed >= 1 { + break; + } + tokio::time::sleep(Duration::from_millis(20)).await; + } + assert_eq!( + relay.counts().relayed, + 1, + "the relay took sequence 0 before any gateway subscribed" + ); + publisher + .send(golden::frame(b"kv", 1, &golden::bytes(golden::BATCH2))) + .await + .expect("publish"); + for _ in 0..250 { + if relay.counts().relayed >= 2 { + break; + } + tokio::time::sleep(Duration::from_millis(20)).await; + } + assert_eq!(relay.counts().relayed, 2); + + // The first gateway gets both from the history: the engine's whole state. + let mut stream = h + .client + .subscribe_kv_events(common::SubscribeKvEventsRequest::default()) + .await + .expect("subscribe") + .into_inner(); + for expected in [0, 1] { + let batch = tokio::time::timeout(Duration::from_secs(5), stream.message()) + .await + .expect("a batch in time") + .expect("stream open") + .expect("a batch"); + assert_eq!(batch.sequence_number, expected); + } + assert_eq!(relay.counts().served_from_history, 1); + drop(stream); + h.server.stop(Duration::from_secs(5)).expect("clean stop"); +} + +/// The relay asks the engine's replay for the batches it missed before its +/// subscription joined: a publisher already at sequence 3 when the servicer +/// starts, whose replay covers 0..=3, leaves the window whole from the +/// publisher's first batch, and the first gateway gets all four from it. +#[tokio::test] +async fn a_publisher_already_counting_when_the_servicer_starts_is_replayed_from_its_start() { + use zeromq::{prelude::*, PubSocket, RouterSocket, ZmqMessage}; + + use crate::kv_events::golden; + + let mut publisher = PubSocket::new(); + let pub_port = bound_port( + &publisher + .bind("tcp://127.0.0.1:0") + .await + .expect("publisher binds") + .to_string(), + ); + let mut router = RouterSocket::new(); + let replay_port = bound_port( + &router + .bind("tcp://127.0.0.1:0") + .await + .expect("replay socket binds") + .to_string(), + ); + let mut model = model_info(); + model.kv_events_endpoint = format!("tcp://*:{pub_port}"); + model.kv_events_replay_endpoint = format!("tcp://*:{replay_port}"); + model.kv_events_topic = "kv".to_string(); + let mut h = harness(model).await; + let relay = h + .server + .state + .kv_relay + .clone() + .expect("a relay for the publisher"); + // Sequences 0..=2 went out before the subscription landed: publish 3 + // until the relay asks the replay socket, which it must do from 0. + let batch1 = golden::bytes(golden::BATCH1); + let batch2 = golden::bytes(golden::BATCH2); + let mut request = None; + for _ in 0..250 { + publisher + .send(golden::frame(b"kv", 3, &batch2)) + .await + .expect("publish"); + if let Ok(message) = tokio::time::timeout(Duration::from_millis(20), router.recv()).await { + request = Some(message.expect("a replay request")); + break; + } + } + let request = request.expect("the relay asked the replay socket"); + let frames: Vec> = request.iter().map(|frame| frame.to_vec()).collect(); + assert_eq!(frames.len(), 3, "[identity, empty, start]"); + assert_eq!( + frames[2], + 0u64.to_be_bytes(), + "asked from the publisher's start" + ); + for (sequence, payload) in [(0u64, &batch1), (1, &batch2), (2, &batch2), (3, &batch2)] { + let mut reply = ZmqMessage::from(frames[0].clone()); + reply.push_back(Vec::new().into()); + reply.push_back(b"kv".to_vec().into()); + reply.push_back(sequence.to_be_bytes().to_vec().into()); + reply.push_back(payload.clone().into()); + router.send(reply).await.expect("reply"); + } + let mut end = ZmqMessage::from(frames[0].clone()); + end.push_back(Vec::new().into()); + end.push_back(Vec::new().into()); + end.push_back([0xff; 8].to_vec().into()); + end.push_back(Vec::new().into()); + router.send(end).await.expect("end marker"); + for _ in 0..250 { + if relay.counts().relayed >= 4 { + break; + } + tokio::time::sleep(Duration::from_millis(20)).await; + } + // Whichever asked first, the start replay or the late join: all four came + // from the replay socket and nothing is unknown. + let counts = relay.counts(); + assert_eq!( + ( + counts.relayed, + counts.gap_batches_recovered + counts.primed_batches, + counts.unknown_before_start + ), + (4, 4, 0), + "{counts:?}" + ); + + // The window is the publisher's whole life: the first gateway gets it. + let mut stream = h + .client + .subscribe_kv_events(common::SubscribeKvEventsRequest::default()) + .await + .expect("subscribe") + .into_inner(); + for expected in 0..=3 { + let batch = tokio::time::timeout(Duration::from_secs(5), stream.message()) + .await + .expect("a batch in time") + .expect("stream open") + .expect("a batch"); + assert_eq!(batch.sequence_number, expected); + } + assert_eq!(relay.counts().served_from_history, 1); + drop(stream); h.server.stop(Duration::from_secs(5)).expect("clean stop"); } @@ -949,7 +1229,7 @@ async fn unserved_rpcs_report_the_gap() { #[tokio::test] async fn startup_timeout_keeps_the_server_up_and_reports_the_error() { let dir = tempfile::tempdir().unwrap(); - let mut config = config(dir.path(), &handshake_address(), model_info()); + let mut config = config(dir.path(), &handshake_address(dir.path()), model_info()); config.engine_startup_timeout = Duration::from_millis(300); let server = SglangServicerServer::start(config).expect("servicer starts"); wait_until(|| matches!(server.last_error(), Ok(Some(_)))).await; diff --git a/crates/engine_servicer/src/testing.rs b/crates/engine_servicer/src/testing.rs new file mode 100644 index 0000000000..c500e5c3bb --- /dev/null +++ b/crates/engine_servicer/src/testing.rs @@ -0,0 +1,24 @@ +//! The deadline every servicer test's wait on the servicer runs under. + +use std::{future::Future, panic::Location}; + +use engine_zmq_client::mock_engine::MOCK_DEADLINE; +use tokio::time::timeout; + +/// A wait on the servicer (a response stream's next message, a watch's next +/// event) under the mock engine's deadline: a servicer that stalls fails the +/// test at the waiting line after thirty seconds instead of hanging the gate. +pub(crate) trait Bounded: Future + Sized { + #[track_caller] + fn bounded(self) -> impl Future { + let caller = Location::caller(); + async move { + match timeout(MOCK_DEADLINE, self).await { + Ok(output) => output, + Err(_) => panic!("{caller}: nothing from the servicer within {MOCK_DEADLINE:?}"), + } + } + } +} + +impl Bounded for F {} diff --git a/crates/engine_servicer/src/tokenspeed/engine.rs b/crates/engine_servicer/src/tokenspeed/engine.rs index 16b5703267..af2d6fc337 100644 --- a/crates/engine_servicer/src/tokenspeed/engine.rs +++ b/crates/engine_servicer/src/tokenspeed/engine.rs @@ -5,7 +5,7 @@ use std::{sync::Arc, time::Duration}; -use engine_zmq_adapter::{connect_with_eos, EosTokenIds}; +use engine_zmq_adapter::{connect_with_eos, EosTokenIds, Handshake}; use openai_protocol::worker::RuntimeType; use tracing::{error, info, warn}; @@ -50,7 +50,7 @@ pub(super) async fn connect_engine( &ipc_base_url, state.model.model_path.clone(), RuntimeType::TokenSpeed, - Some(&handshake_address), + Handshake::TcpOrIpc(&handshake_address), engine_count, eos, startup_timeout, diff --git a/crates/engine_servicer/src/tokenspeed/mod.rs b/crates/engine_servicer/src/tokenspeed/mod.rs index 327edc4137..83155a419d 100644 --- a/crates/engine_servicer/src/tokenspeed/mod.rs +++ b/crates/engine_servicer/src/tokenspeed/mod.rs @@ -80,6 +80,10 @@ pub struct TokenSpeedModelInfo { /// TokenSpeed's ZMQ KV-event publisher endpoint and topic; an empty /// endpoint means `SubscribeKvEvents` is UNIMPLEMENTED. pub kv_events_endpoint: String, + /// The same config's `replay_endpoint` (the publisher's replay ROUTER), + /// or empty when it runs none: what the relay asks for gaps in flight + /// and for the batches published before its subscription joined. + pub kv_events_replay_endpoint: String, pub kv_events_topic: String, } @@ -93,6 +97,8 @@ pub struct TokenSpeedServicerConfig { pub ipc_base_url: String, /// `tcp://host:port` the headless scheduler dials for the handshake (its /// `--data-parallel-address`/`--data-parallel-rpc-port`). + /// An `ipc://` endpoint is accepted too: one per test, so parallel + /// tests never share a probed TCP port. pub handshake_address: String, /// Engines that will dial in (the attention data-parallel size). pub engine_count: usize, @@ -105,8 +111,26 @@ pub struct TokenSpeedServicerConfig { pub engine_startup_timeout: Duration, } +/// The servicer's `GetLoads` figures, read for the load record the relay +/// attaches to every batch it streams. +struct LoadFromState(std::sync::Weak); + +impl crate::kv_events::LoadSource for LoadFromState { + fn load(&self, dp_rank: Option) -> Option { + let state = self.0.upgrade()?; + let response = info::loads(&state, dp_rank).ok()?; + let load = response.loads.first()?; + // The engine reports no queued token-work on this wire yet: the + // record's `waiting_uncached_tokens` stays unset. + Some(smg_grpc_client::common_proto::EngineLoad::from(load)) + } +} + pub(super) struct State { pub(super) model: TokenSpeedModelInfo, + /// The KV-event relay for the engine's ZMQ publisher; `None` when events + /// are off (`SubscribeKvEvents` is then UNIMPLEMENTED). + pub(super) kv_relay: Option>, /// The local tokenizer directory the servicer loaded (`GetTokenizer` /// bundles it); `None` when none resolved. pub(super) tokenizer_dir: Option, @@ -194,8 +218,12 @@ impl TokenSpeedServicerServer { if !config.ipc_base_url.starts_with("ipc://") { return Err(invalid("ipc_base_url must be ipc://")); } - if !config.handshake_address.starts_with("tcp://") { - return Err(invalid("handshake_address must be tcp://host:port")); + if !config.handshake_address.starts_with("tcp://") + && !config.handshake_address.starts_with("ipc://") + { + return Err(invalid( + "handshake_address must be tcp://host:port or ipc://", + )); } if config.engine_count == 0 { return Err(invalid("engine_count must be positive")); @@ -206,9 +234,15 @@ impl TokenSpeedServicerServer { if config.model.model_path.trim().is_empty() { return Err(invalid("model_path must not be empty")); } + let kv_relay = crate::kv_events::KvEventRelay::for_publisher( + &config.model.kv_events_endpoint, + Some(&config.model.kv_events_replay_endpoint), + &config.model.kv_events_topic, + ); let state = Arc::new(State { tokenizer_dir: config.tokenizer_dir.clone(), model: config.model, + kv_relay, engine: EngineLink::default(), tokenizer: OnceLock::new(), registry: Arc::new(RequestRegistry::default()), @@ -216,6 +250,9 @@ impl TokenSpeedServicerServer { started: Instant::now(), started_at: SystemTime::now(), }); + if let Some(relay) = &state.kv_relay { + relay.set_load_source(Arc::new(LoadFromState(Arc::downgrade(&state)))); + } if let Some(tokenizer) = tokenizer { let _ = state.tokenizer.set(Some(tokenizer)); } @@ -228,6 +265,7 @@ impl TokenSpeedServicerServer { Arc::new(move || health_state.is_serving()), ); let connect_state = Arc::clone(&state); + let kv_relay = state.kv_relay.clone(); let TokenSpeedServicerConfig { bind_address, ipc_base_url, @@ -241,6 +279,12 @@ impl TokenSpeedServicerServer { "smg-tokenspeed-servicer", &bind_address, move |listener: TcpListener, shutdown: Shutdown, last_error| async move { + // The KV-event relay follows the publisher from the start, + // before any gateway asks, so its history and live-block + // record cover the engine's whole life. + if let Some(relay) = &kv_relay { + relay.start_at_boot(); + } #[expect( clippy::disallowed_methods, reason = "engine connect is fire-and-forget; the runtime drop cancels it" diff --git a/crates/engine_servicer/src/tokenspeed/service.rs b/crates/engine_servicer/src/tokenspeed/service.rs index 9e30de0061..1e2bfb838f 100644 --- a/crates/engine_servicer/src/tokenspeed/service.rs +++ b/crates/engine_servicer/src/tokenspeed/service.rs @@ -122,16 +122,11 @@ impl TokenSpeedScheduler for TokenSpeedService { &self, request: Request, ) -> Result, Status> { - let model = &self.state.model; - if model.kv_events_endpoint.is_empty() { + let Some(relay) = &self.state.kv_relay else { return Err(Status::unimplemented( kv_events::TOKENSPEED_DISABLED_MESSAGE, )); - } - Ok(Response::new(kv_events::subscribe( - &model.kv_events_endpoint, - model.kv_events_topic.clone(), - request.into_inner(), - ))) + }; + relay.subscribe(request.into_inner()).map(Response::new) } } diff --git a/crates/engine_servicer/src/tokenspeed/tests.rs b/crates/engine_servicer/src/tokenspeed/tests.rs index 2cc60f00e9..d1bb0ff59f 100644 --- a/crates/engine_servicer/src/tokenspeed/tests.rs +++ b/crates/engine_servicer/src/tokenspeed/tests.rs @@ -11,7 +11,10 @@ use std::{ use bytes::Bytes; use engine_zmq_client::{ codec::{decode_msgpack, encode_msgpack}, - mock_engine::{connect_to_frontend, default_ready_response, MockEngineInput, MockEngineOutput}, + mock_engine::{ + connect_to_frontend, default_ready_response, MockEngineInput, MockEngineOutput, + MOCK_DEADLINE, + }, protocol::tokenspeed::{ output::BatchTokenIDOutSlim, request::{TokenSpeedRequestType, TokenizedGenerateReqInput}, @@ -19,7 +22,6 @@ use engine_zmq_client::{ EngineId, }; use llm_tokenizer::{mock::MockTokenizer, traits::Tokenizer}; -use portpicker::pick_unused_port; use prost_types::value::Kind; use smg_grpc_client::{ common_proto as common, @@ -31,7 +33,7 @@ use tonic_health::pb::{ }; use super::*; -use crate::{kv_events, ServicerError}; +use crate::{kv_events, testing::Bounded, ServicerError}; fn model_info() -> TokenSpeedModelInfo { TokenSpeedModelInfo { @@ -73,11 +75,33 @@ fn config( } } -fn handshake_address() -> String { - format!( - "tcp://127.0.0.1:{}", - pick_unused_port().expect("a free handshake port") - ) +/// The handshake endpoint of one test: an IPC socket under its own +/// directory. A probed TCP port is not reserved, so two tests running in +/// parallel could pick the same one and a mock engine would handshake with +/// the other test's servicer and wait forever for its INIT. +fn handshake_address(dir: &std::path::Path) -> String { + format!("ipc://{}", dir.join("handshake").display()) +} + +/// The handshake endpoint is free again: ZMQ unlinks the ipc socket file once +/// the bound socket is dropped, so a fresh listener can take the path. +async fn assert_handshake_released(handshake: &str) { + use std::os::unix::net::UnixListener; + let path = handshake.trim_start_matches("ipc://"); + let deadline = Instant::now() + Duration::from_secs(5); + while std::fs::metadata(path).is_ok() && Instant::now() < deadline { + tokio::time::sleep(Duration::from_millis(20)).await; + } + UnixListener::bind(path).expect("handshake endpoint released"); +} + +/// The port a socket bound to `tcp://127.0.0.1:0` was given. +fn bound_port(endpoint: &str) -> u16 { + endpoint + .rsplit(':') + .next() + .and_then(|port| port.parse().ok()) + .expect("a bound tcp endpoint ends with its port") } async fn wait_until(mut condition: impl FnMut() -> bool) { @@ -90,6 +114,17 @@ async fn wait_until(mut condition: impl FnMut() -> bool) { panic!("condition not met within 10s"); } +/// A channel to the servicer whose every request fails after [`MOCK_DEADLINE`] +/// instead of waiting on a servicer that never answers. +async fn grpc_channel(address: impl std::fmt::Display) -> Channel { + Channel::from_shared(format!("http://{address}")) + .expect("grpc address") + .timeout(MOCK_DEADLINE) + .connect() + .await + .expect("grpc client") +} + /// A bound servicer, a handshaken mock scheduler, and a gRPC client. struct Harness { server: TokenSpeedServicerServer, @@ -101,7 +136,7 @@ struct Harness { async fn harness(model: TokenSpeedModelInfo, tokenizer: Option>) -> Harness { let dir = tempfile::tempdir().unwrap(); - let handshake = handshake_address(); + let handshake = handshake_address(dir.path()); let config = config(dir.path(), &handshake, model); let server = match tokenizer { Some(tokenizer) => TokenSpeedServicerServer::start_with_tokenizer(config, tokenizer), @@ -116,9 +151,7 @@ async fn harness(model: TokenSpeedModelInfo, tokenizer: Option { @@ -613,7 +646,7 @@ async fn choices_fan_out_under_the_parent_id() { (1, vec![21], true), ] ); - assert!(stream.message().await.unwrap().is_none()); + assert!(stream.message().bounded().await.unwrap().is_none()); h.server.stop(Duration::from_secs(5)).expect("clean stop"); } @@ -636,7 +669,15 @@ async fn logprobs_pass_through_and_the_unsupported_kinds_are_refused() { step.output_token_logprobs_val = vec![vec![-0.5]]; step.output_token_logprobs_idx = vec![vec![10]]; send(&mut h.engine_out, &step).await; - let chunk = match stream.message().await.unwrap().unwrap().response.unwrap() { + let chunk = match stream + .message() + .bounded() + .await + .unwrap() + .unwrap() + .response + .unwrap() + { ts::generate_response::Response::Chunk(chunk) => chunk, other @ ts::generate_response::Response::Complete(_) => { panic!("expected a chunk, got {other:?}") @@ -645,7 +686,7 @@ async fn logprobs_pass_through_and_the_unsupported_kinds_are_refused() { let logprobs = chunk.output_logprobs.expect("chunk logprobs"); assert_eq!(logprobs.token_ids, vec![10]); assert!((logprobs.token_logprobs[0] + 0.5).abs() < 1e-6); - let done = complete(stream.message().await.unwrap().unwrap()); + let done = complete(stream.message().bounded().await.unwrap().unwrap()); assert_eq!(done.output_logprobs.unwrap().token_ids, vec![10]); let mut request = generate_request("lp2", true, Vec::new()); @@ -771,7 +812,7 @@ async fn info_rpcs_report_the_launcher_facts_and_the_handshake() { &batch("ld", vec![11], 2, Some("length"), Some((2, 1, 10, 100))), ) .await; - let done = complete(stream.message().await.unwrap().unwrap()); + let done = complete(stream.message().bounded().await.unwrap().unwrap()); assert_eq!(done.output_ids, vec![10, 11]); let idle = h .client @@ -824,13 +865,203 @@ async fn control_rpcs_report_what_the_wire_cannot_carry() { h.server.stop(Duration::from_secs(5)).expect("clean stop"); } +/// The relay follows the publisher from the servicer's start, before any +/// gateway subscribes: batches published with nobody listening are in its +/// history, and the first subscription gets them as the engine's whole state. +#[tokio::test] +async fn the_relay_subscribes_to_the_publisher_at_boot_before_any_gateway() { + use zeromq::{prelude::*, PubSocket}; + + use crate::kv_events::golden; + + let mut publisher = PubSocket::new(); + let port = bound_port( + &publisher + .bind("tcp://127.0.0.1:0") + .await + .expect("publisher binds") + .to_string(), + ); + let mut model = model_info(); + model.kv_events_endpoint = format!("tcp://*:{port}"); + model.kv_events_topic = "kv".to_string(); + let mut h = harness(model, None).await; + let relay = h + .server + .state + .kv_relay + .clone() + .expect("a relay for the publisher"); + // The SUB connect is asynchronous: publish sequence 0 until the relay, + // with no subscriber of its own yet, has taken it (repeats are duplicates). + let batch1 = golden::bytes(golden::BATCH1); + for _ in 0..250 { + publisher + .send(golden::frame(b"kv", 0, &batch1)) + .await + .expect("publish"); + if relay.counts().relayed >= 1 { + break; + } + tokio::time::sleep(Duration::from_millis(20)).await; + } + assert_eq!( + relay.counts().relayed, + 1, + "the relay took sequence 0 before any gateway subscribed" + ); + publisher + .send(golden::frame(b"kv", 1, &golden::bytes(golden::BATCH2))) + .await + .expect("publish"); + for _ in 0..250 { + if relay.counts().relayed >= 2 { + break; + } + tokio::time::sleep(Duration::from_millis(20)).await; + } + assert_eq!(relay.counts().relayed, 2); + + // The first gateway gets both from the history: the engine's whole state. + let mut stream = h + .client + .subscribe_kv_events(common::SubscribeKvEventsRequest::default()) + .await + .expect("subscribe") + .into_inner(); + for expected in [0, 1] { + let batch = tokio::time::timeout(Duration::from_secs(5), stream.message()) + .await + .expect("a batch in time") + .expect("stream open") + .expect("a batch"); + assert_eq!(batch.sequence_number, expected); + } + assert_eq!(relay.counts().served_from_history, 1); + drop(stream); + h.server.stop(Duration::from_secs(5)).expect("clean stop"); +} + +/// The relay asks the engine's replay for the batches it missed before its +/// subscription joined: a publisher already at sequence 3 when the servicer +/// starts, whose replay covers 0..=3, leaves the window whole from the +/// publisher's first batch, and the first gateway gets all four from it. +#[tokio::test] +async fn a_publisher_already_counting_when_the_servicer_starts_is_replayed_from_its_start() { + use zeromq::{prelude::*, PubSocket, RouterSocket, ZmqMessage}; + + use crate::kv_events::golden; + + let mut publisher = PubSocket::new(); + let pub_port = bound_port( + &publisher + .bind("tcp://127.0.0.1:0") + .await + .expect("publisher binds") + .to_string(), + ); + let mut router = RouterSocket::new(); + let replay_port = bound_port( + &router + .bind("tcp://127.0.0.1:0") + .await + .expect("replay socket binds") + .to_string(), + ); + let mut model = model_info(); + model.kv_events_endpoint = format!("tcp://*:{pub_port}"); + model.kv_events_replay_endpoint = format!("tcp://*:{replay_port}"); + model.kv_events_topic = "kv".to_string(); + let mut h = harness(model, None).await; + let relay = h + .server + .state + .kv_relay + .clone() + .expect("a relay for the publisher"); + // Sequences 0..=2 went out before the subscription landed: publish 3 + // until the relay asks the replay socket, which it must do from 0. + let batch1 = golden::bytes(golden::BATCH1); + let batch2 = golden::bytes(golden::BATCH2); + let mut request = None; + for _ in 0..250 { + publisher + .send(golden::frame(b"kv", 3, &batch2)) + .await + .expect("publish"); + if let Ok(message) = tokio::time::timeout(Duration::from_millis(20), router.recv()).await { + request = Some(message.expect("a replay request")); + break; + } + } + let request = request.expect("the relay asked the replay socket"); + let frames: Vec> = request.iter().map(|frame| frame.to_vec()).collect(); + assert_eq!(frames.len(), 3, "[identity, empty, start]"); + assert_eq!( + frames[2], + 0u64.to_be_bytes(), + "asked from the publisher's start" + ); + for (sequence, payload) in [(0u64, &batch1), (1, &batch2), (2, &batch2), (3, &batch2)] { + let mut reply = ZmqMessage::from(frames[0].clone()); + reply.push_back(Vec::new().into()); + reply.push_back(b"kv".to_vec().into()); + reply.push_back(sequence.to_be_bytes().to_vec().into()); + reply.push_back(payload.clone().into()); + router.send(reply).await.expect("reply"); + } + let mut end = ZmqMessage::from(frames[0].clone()); + end.push_back(Vec::new().into()); + end.push_back(Vec::new().into()); + end.push_back([0xff; 8].to_vec().into()); + end.push_back(Vec::new().into()); + router.send(end).await.expect("end marker"); + for _ in 0..250 { + if relay.counts().relayed >= 4 { + break; + } + tokio::time::sleep(Duration::from_millis(20)).await; + } + // Whichever asked first, the start replay or the late join: all four came + // from the replay socket and nothing is unknown. + let counts = relay.counts(); + assert_eq!( + ( + counts.relayed, + counts.gap_batches_recovered + counts.primed_batches, + counts.unknown_before_start + ), + (4, 4, 0), + "{counts:?}" + ); + + // The window is the publisher's whole life: the first gateway gets it. + let mut stream = h + .client + .subscribe_kv_events(common::SubscribeKvEventsRequest::default()) + .await + .expect("subscribe") + .into_inner(); + for expected in 0..=3 { + let batch = tokio::time::timeout(Duration::from_secs(5), stream.message()) + .await + .expect("a batch in time") + .expect("stream open") + .expect("a batch"); + assert_eq!(batch.sequence_number, expected); + } + assert_eq!(relay.counts().served_from_history, 1); + drop(stream); + h.server.stop(Duration::from_secs(5)).expect("clean stop"); +} + /// A scheduler that never dials in fails the link at the configured bound (not /// a fixed one: a cold kernel cache makes a real start exceed ten minutes), the -/// server stays up to report it, and `stop` releases the handshake port. +/// server stays up to report it, and `stop` releases the handshake endpoint. #[tokio::test] async fn a_scheduler_that_never_dials_in_fails_the_link_at_the_startup_bound() { let dir = tempfile::tempdir().unwrap(); - let handshake = handshake_address(); + let handshake = handshake_address(dir.path()); let server = TokenSpeedServicerServer::start(TokenSpeedServicerConfig { engine_startup_timeout: Duration::from_millis(300), ..config(dir.path(), &handshake, model_info()) @@ -852,6 +1083,5 @@ async fn a_scheduler_that_never_dials_in_fails_the_link_at_the_startup_bound() { ); server.stop(Duration::from_secs(5)).unwrap(); - let port = handshake.rsplit(':').next().unwrap(); - std::net::TcpListener::bind(format!("127.0.0.1:{port}")).expect("handshake port released"); + assert_handshake_released(&handshake).await; } diff --git a/crates/engine_servicer/src/vllm/engine.rs b/crates/engine_servicer/src/vllm/engine.rs index 2c99bb1ff7..d19cc3f1d4 100644 --- a/crates/engine_servicer/src/vllm/engine.rs +++ b/crates/engine_servicer/src/vllm/engine.rs @@ -3,7 +3,9 @@ use std::{sync::Arc, time::Duration}; -use engine_zmq_adapter::{connect_with_eos, structured_outputs_backend_from_config, EosTokenIds}; +use engine_zmq_adapter::{ + connect_with_eos, structured_outputs_backend_from_config, EosTokenIds, Handshake, +}; use openai_protocol::worker::RuntimeType; use tracing::{error, info, warn}; @@ -52,7 +54,7 @@ pub(super) async fn connect_engine( &ipc_base_url, state.model.model_path.clone(), RuntimeType::Vllm, - Some(&handshake_address), + Handshake::TcpOrIpc(&handshake_address), engine_count, eos, startup_timeout, diff --git a/crates/engine_servicer/src/vllm/generate.rs b/crates/engine_servicer/src/vllm/generate.rs index 6169005264..0a58321d2a 100644 --- a/crates/engine_servicer/src/vllm/generate.rs +++ b/crates/engine_servicer/src/vllm/generate.rs @@ -162,6 +162,12 @@ async fn submit( // comes off; a refusal below leaves the engine without the request (a // fan-out's earlier subs are aborted with the error) and still owes one. let owed = notice.disarm(); + let prompt_tokens = match req.input.as_ref() { + Some(vllm::generate_request::Input::Tokenized(tokenized)) => { + u32::try_from(tokenized.input_ids.len()).unwrap_or(u32::MAX) + } + _ => 0, + }; let subs = match client .generate_vllm_streams_with_media(req, processed_media) .await @@ -174,6 +180,13 @@ async fn submit( return Err(status); } }; + // Queued token-work for `GetLoads`: in flight from here until the + // engine's first output (or the stream's end). + state.loads.submitted(&request_id, prompt_tokens); + let loads = LoadsGuard { + state: Arc::clone(state), + request_id: request_id.clone(), + }; let mut choices = SelectAll::new(); for sub in subs { // Each choice decodes its own text: the matcher is per sequence. @@ -188,6 +201,7 @@ async fn submit( streaming, min_tokens, media_identity.clone(), + request_id.clone(), )); } Ok(Box::pin(GenerateStream { @@ -195,9 +209,23 @@ async fn submit( cancel, aborted: None, _registration: registration, + _loads: loads, })) } +/// Takes the request out of the load tracker's queue when the stream ends, +/// in case the engine never started it. +struct LoadsGuard { + state: Arc, + request_id: String, +} + +impl Drop for LoadsGuard { + fn drop(&mut self) { + self.state.loads.finished(&self.request_id); + } +} + /// One choice of a generate request: the ZMQ-mapped stream plus the string /// stop matcher the engine cannot run itself. struct ChoiceStream { @@ -220,6 +248,8 @@ struct ChoiceStream { /// A PD prefill leg's account of the media it processed, stamped on /// every `Complete` of the request for the decode leg. media_identity: Option>, + /// The request this choice belongs to, for the load tracker. + request_id: String, } impl ChoiceStream { @@ -230,9 +260,11 @@ impl ChoiceStream { streaming: bool, min_tokens: u32, media_identity: Option>, + request_id: String, ) -> Self { Self { media_identity, + request_id, state, inner: Some(inner), stops: StopMatcher::new(decoder, min_tokens), @@ -243,8 +275,9 @@ impl ChoiceStream { } /// Count the prompt once per request (every choice of an `n > 1` fan-out - /// reports the same prompt) for the stats line. - fn count_prompt(&mut self, prompt_tokens: u32) { + /// reports the same prompt) for the stats line, and tell the load tracker + /// the engine started it and how much of the prompt was cached. + fn count_prompt(&mut self, prompt_tokens: u32, cached_tokens: u32) { if self.prompt_counted || prompt_tokens == 0 { return; } @@ -253,6 +286,9 @@ impl ChoiceStream { .stats .prompt_tokens .fetch_add(u64::from(prompt_tokens), Ordering::Relaxed); + self.state + .loads + .first_output(&self.request_id, prompt_tokens, cached_tokens); } self.prompt_counted = true; } @@ -318,15 +354,18 @@ impl Stream for ChoiceStream { if let Some(vllm::generate_response::Response::Complete(complete)) = item.response.as_ref() { - this.count_prompt(complete.prompt_tokens); + this.count_prompt(complete.prompt_tokens, complete.cached_tokens); } return Poll::Ready(Some(Ok(this.stamped(item)))); }; - this.count_prompt(chunk.prompt_tokens); + this.count_prompt(chunk.prompt_tokens, chunk.cached_tokens); this.state .stats .generation_tokens .fetch_add(chunk.token_ids.len() as u64, Ordering::Relaxed); + this.state + .loads + .generated(u32::try_from(chunk.token_ids.len()).unwrap_or(u32::MAX)); let matched = match this.stops.feed(&chunk.token_ids) { Ok(matched) => matched, Err(status) => { @@ -372,6 +411,7 @@ struct GenerateStream { /// at that point, yielded before the stream ends. aborted: Option>, _registration: Registration, + _loads: LoadsGuard, } impl Stream for GenerateStream { diff --git a/crates/engine_servicer/src/vllm/info.rs b/crates/engine_servicer/src/vllm/info.rs index 532a3ca298..a194478abe 100644 --- a/crates/engine_servicer/src/vllm/info.rs +++ b/crates/engine_servicer/src/vllm/info.rs @@ -176,7 +176,11 @@ fn server_facts(state: &State) -> vllm::GetServerInfoResponse { } /// `GetLoads`: the per-rank load piggybacked on engine output, in the gRPC -/// response shape, with the handshake's capacity figures. +/// response shape, with the handshake's capacity figures and, from this +/// servicer's own bookkeeping ([`crate::load_tracker`]), the queued +/// token-work, generation throughput and hit rate that vLLM's stats do not +/// carry. The servicer forwards to one engine process, so those three go on +/// the first rank's entry. pub(super) fn loads(state: &State) -> Result { let client = state.engine()?; let ready = client.ready_response(); @@ -192,17 +196,28 @@ pub(super) fn loads(state: &State) -> Result { let mut loads: Vec = snapshot .loads .into_iter() - .map(|load| vllm::SchedulerLoad { - dp_rank: load.dp_rank, - num_running_reqs: load.num_running_reqs, - num_waiting_reqs: load.num_waiting_reqs, - num_total_reqs: load.num_running_reqs.saturating_add(load.num_waiting_reqs), - token_usage: load.token_usage, - max_running_requests, - max_total_num_tokens, - ..Default::default() + .map(|load| { + let num_used_tokens = (load.token_usage * f64::from(max_total_num_tokens)).round(); + vllm::SchedulerLoad { + dp_rank: load.dp_rank, + num_running_reqs: load.num_running_reqs, + num_waiting_reqs: load.num_waiting_reqs, + num_total_reqs: load.num_running_reqs.saturating_add(load.num_waiting_reqs), + num_used_tokens: num_used_tokens.clamp(0.0, f64::from(i32::MAX)) as i32, + token_usage: load.token_usage, + utilization: load.token_usage, + max_running_requests, + max_total_num_tokens, + ..Default::default() + } }) .collect(); + if let Some(first) = loads.first_mut() { + let estimate = state.loads.estimate(first.num_waiting_reqs); + first.num_waiting_uncached_tokens = estimate.queued_token_work; + first.gen_throughput = estimate.gen_throughput; + first.cache_hit_rate = estimate.cache_hit_rate; + } if loads.is_empty() { // No output batch has carried a snapshot yet. The Python servicer // reports a zero-filled entry per rank, and the Router reads an diff --git a/crates/engine_servicer/src/vllm/mod.rs b/crates/engine_servicer/src/vllm/mod.rs index dd15ea1c7d..4012bbc13e 100644 --- a/crates/engine_servicer/src/vllm/mod.rs +++ b/crates/engine_servicer/src/vllm/mod.rs @@ -137,6 +137,8 @@ pub struct VllmServicerConfig { pub ipc_base_url: String, /// `tcp://host:port` the headless engine dials for the handshake (its /// `--data-parallel-address`/`--data-parallel-rpc-port`). + /// An `ipc://` endpoint is accepted too: one per test, so parallel + /// tests never share a probed TCP port. pub handshake_address: String, /// Engines that will dial in (the engine-level data-parallel size). pub engine_count: usize, @@ -179,9 +181,39 @@ pub(super) struct Stats { pub(super) generation_tokens: AtomicU64, } +/// The servicer's `GetLoads` figures, read for the load record the relay +/// attaches to every batch it streams. +struct LoadFromState(std::sync::Weak); + +impl crate::kv_events::LoadSource for LoadFromState { + fn load(&self, dp_rank: Option) -> Option { + let state = self.0.upgrade()?; + let response = info::loads(&state).ok()?; + let rank = dp_rank.unwrap_or(0); + let load = response + .loads + .iter() + .find(|load| load.dp_rank == rank) + .or_else(|| response.loads.first())?; + // The queued token-work is this servicer's estimate for the first + // rank (see `info::loads`); the other ranks do not report it. + let estimated = response.loads.first().map(|first| first.dp_rank) == Some(load.dp_rank); + let mut record = smg_grpc_client::common_proto::EngineLoad::from(load); + record.waiting_uncached_tokens = + estimated.then(|| u32::try_from(load.num_waiting_uncached_tokens).unwrap_or(0)); + Some(record) + } +} + pub(super) struct State { pub(super) model: VllmModelInfo, + /// The KV-event relay for the engine's ZMQ publisher; `None` when events + /// are off (`SubscribeKvEvents` is then UNIMPLEMENTED). + pub(super) kv_relay: Option>, pub(super) stats: Stats, + /// Queued token-work, generation throughput and hit rate from the + /// requests this servicer forwards, for `GetLoads`. + pub(super) loads: crate::load_tracker::LoadTracker, /// The local tokenizer directory the servicer loaded (`GetTokenizer` /// bundles it); `None` when none resolved. pub(super) tokenizer_dir: Option, @@ -276,8 +308,12 @@ impl VllmServicerServer { if !config.ipc_base_url.starts_with("ipc://") { return Err(invalid("ipc_base_url must be ipc://")); } - if !config.handshake_address.starts_with("tcp://") { - return Err(invalid("handshake_address must be tcp://host:port")); + if !config.handshake_address.starts_with("tcp://") + && !config.handshake_address.starts_with("ipc://") + { + return Err(invalid( + "handshake_address must be tcp://host:port or ipc://", + )); } if config.engine_count == 0 { return Err(invalid("engine_count must be positive")); @@ -288,10 +324,17 @@ impl VllmServicerServer { if config.model.model_path.trim().is_empty() { return Err(invalid("model_path must not be empty")); } + let kv_relay = crate::kv_events::KvEventRelay::for_publisher( + &config.model.kv_events_endpoint, + Some(&config.model.kv_events_replay_endpoint), + &config.model.kv_events_topic, + ); let state = Arc::new(State { tokenizer_dir: config.tokenizer_dir.clone(), model: config.model, + kv_relay, stats: Stats::default(), + loads: crate::load_tracker::LoadTracker::default(), engine: EngineLink::default(), tokenizer: OnceLock::new(), registry: Arc::new(RequestRegistry::default()), @@ -299,6 +342,9 @@ impl VllmServicerServer { started: Instant::now(), media: config.media_processor.map(MediaGate::new), }); + if let Some(relay) = &state.kv_relay { + relay.set_load_source(Arc::new(LoadFromState(Arc::downgrade(&state)))); + } if let Some(tokenizer) = tokenizer { let _ = state.tokenizer.set(Some(tokenizer)); } @@ -312,6 +358,7 @@ impl VllmServicerServer { ); let connect_state = Arc::clone(&state); let stats_state = Arc::clone(&state); + let kv_relay = state.kv_relay.clone(); let VllmServicerConfig { bind_address, ipc_base_url, @@ -325,6 +372,12 @@ impl VllmServicerServer { "smg-vllm-servicer", &bind_address, move |listener: TcpListener, shutdown: Shutdown, last_error| async move { + // The KV-event relay follows the publisher from the start, + // before any gateway asks, so its history and live-block + // record cover the engine's whole life. + if let Some(relay) = &kv_relay { + relay.start_at_boot(); + } // Fire-and-forget on the server runtime: dropping the runtime // with the server thread cancels a still-running connect. #[expect( diff --git a/crates/engine_servicer/src/vllm/service.rs b/crates/engine_servicer/src/vllm/service.rs index 09d62c47d9..bc370dd83c 100644 --- a/crates/engine_servicer/src/vllm/service.rs +++ b/crates/engine_servicer/src/vllm/service.rs @@ -103,14 +103,9 @@ impl VllmEngine for VllmEngineService { &self, request: Request, ) -> Result, Status> { - let model = &self.state.model; - if model.kv_events_endpoint.is_empty() { + let Some(relay) = &self.state.kv_relay else { return Err(Status::unimplemented(kv_events::VLLM_DISABLED_MESSAGE)); - } - Ok(Response::new(kv_events::subscribe( - &model.kv_events_endpoint, - model.kv_events_topic.clone(), - request.into_inner(), - ))) + }; + relay.subscribe(request.into_inner()).map(Response::new) } } diff --git a/crates/engine_servicer/src/vllm/tests.rs b/crates/engine_servicer/src/vllm/tests.rs index 061f17c372..c331b5bb1b 100644 --- a/crates/engine_servicer/src/vllm/tests.rs +++ b/crates/engine_servicer/src/vllm/tests.rs @@ -14,7 +14,7 @@ use engine_zmq_client::{ codec::{decode_msgpack, tensor::WireTensor, OpaqueValue}, mock_engine::{ connect_to_frontend, default_ready_response, EngineInbound, MockEngineInput, - MockEngineOutput, + MockEngineOutput, MOCK_DEADLINE, }, protocol::vllm::{ multimodal::MmKwargValue, @@ -30,7 +30,6 @@ use engine_zmq_client::{ EngineId, }; use llm_tokenizer::{mock::MockTokenizer, traits::Tokenizer}; -use portpicker::pick_unused_port; use smg_grpc_client::{ common_proto as common, tokenizer_bundle::{validate_bundle_sha256, with_extracted_bundle, StreamBundle}, @@ -44,7 +43,7 @@ use tonic_health::pb::{ use zip::{CompressionMethod, ZipArchive}; use super::*; -use crate::{kv_events, tokenizer_bundle, ServicerError}; +use crate::{kv_events, testing::Bounded, tokenizer_bundle, ServicerError}; fn model_info() -> VllmModelInfo { VllmModelInfo { @@ -77,11 +76,33 @@ fn config(dir: &std::path::Path, handshake: &str, model: VllmModelInfo) -> VllmS } } -fn handshake_address() -> String { - format!( - "tcp://127.0.0.1:{}", - pick_unused_port().expect("a free handshake port") - ) +/// The handshake endpoint of one test: an IPC socket under its own +/// directory. A probed TCP port is not reserved, so two tests running in +/// parallel could pick the same one and a mock engine would handshake with +/// the other test's servicer and wait forever for its INIT. +fn handshake_address(dir: &std::path::Path) -> String { + format!("ipc://{}", dir.join("handshake").display()) +} + +/// The handshake endpoint is free again: ZMQ unlinks the ipc socket file once +/// the bound socket is dropped, so a fresh listener can take the path. +async fn assert_handshake_released(handshake: &str) { + use std::os::unix::net::UnixListener; + let path = handshake.trim_start_matches("ipc://"); + let deadline = Instant::now() + Duration::from_secs(5); + while fs::metadata(path).is_ok() && Instant::now() < deadline { + tokio::time::sleep(Duration::from_millis(20)).await; + } + UnixListener::bind(path).expect("handshake endpoint released"); +} + +/// The port a socket bound to `tcp://127.0.0.1:0` was given. +fn bound_port(endpoint: &str) -> u16 { + endpoint + .rsplit(':') + .next() + .and_then(|port| port.parse().ok()) + .expect("a bound tcp endpoint ends with its port") } async fn wait_until(mut condition: impl FnMut() -> bool) { @@ -94,6 +115,17 @@ async fn wait_until(mut condition: impl FnMut() -> bool) { panic!("condition not met within 10s"); } +/// A channel to the servicer whose every request fails after [`MOCK_DEADLINE`] +/// instead of waiting on a servicer that never answers. +async fn grpc_channel(address: impl std::fmt::Display) -> Channel { + Channel::from_shared(format!("http://{address}")) + .expect("grpc address") + .timeout(MOCK_DEADLINE) + .connect() + .await + .expect("grpc client") +} + /// A bound servicer, a handshaken mock engine, and a gRPC client. struct Harness { server: VllmServicerServer, @@ -113,7 +145,7 @@ async fn harness_with( media_processor: Option>, ) -> Harness { let dir = tempfile::tempdir().unwrap(); - let handshake = handshake_address(); + let handshake = handshake_address(dir.path()); let mut config = config(dir.path(), &handshake, model); config.media_processor = media_processor; let server = match tokenizer { @@ -129,9 +161,7 @@ async fn harness_with( .await .expect("mock engine handshake"); wait_until(|| server.engine_ready()).await; - let client = VllmEngineClient::connect(format!("http://{}", server.address())) - .await - .expect("grpc client"); + let client = VllmEngineClient::new(grpc_channel(server.address()).await); let (engine_in, engine_out) = engine.split(); Harness { server, @@ -247,7 +277,7 @@ fn start_rejects_malformed_config() { ..good.clone() }; let bad_handshake = VllmServicerConfig { - handshake_address: "ipc:///tmp/hs".to_string(), + handshake_address: "udp://127.0.0.1:1".to_string(), ..good.clone() }; let no_engines = VllmServicerConfig { @@ -275,7 +305,7 @@ fn start_rejects_malformed_config() { #[tokio::test] async fn health_gates_on_the_engine_link_and_the_drain_flag() { let dir = tempfile::tempdir().unwrap(); - let handshake = handshake_address(); + let handshake = handshake_address(dir.path()); let server = VllmServicerServer::start(config(dir.path(), &handshake, model_info())) .expect("servicer starts"); let address = format!("http://{}", server.address()); @@ -383,7 +413,7 @@ async fn streams_chunks_then_a_cumulative_complete() { .await .unwrap(); assert_eq!( - chunk_tokens(stream.message().await.unwrap().unwrap()), + chunk_tokens(stream.message().bounded().await.unwrap().unwrap()), vec![10] ); @@ -414,14 +444,14 @@ async fn streams_chunks_then_a_cumulative_complete() { .await .unwrap(); assert_eq!( - chunk_tokens(stream.message().await.unwrap().unwrap()), + chunk_tokens(stream.message().bounded().await.unwrap().unwrap()), vec![11] ); - let done = complete(stream.message().await.unwrap().unwrap()); + let done = complete(stream.message().bounded().await.unwrap().unwrap()); assert_eq!(done.output_ids, vec![10, 11]); assert_eq!(done.finish_reason, "length"); assert_eq!(done.completion_tokens, 2); - assert!(stream.message().await.unwrap().is_none()); + assert!(stream.message().bounded().await.unwrap().is_none()); } /// A non-streaming request gets the terminal `Complete` only, as from the @@ -449,10 +479,10 @@ async fn non_streaming_yields_only_the_complete() { )) .await .unwrap(); - let done = complete(stream.message().await.unwrap().unwrap()); + let done = complete(stream.message().bounded().await.unwrap().unwrap()); assert_eq!(done.output_ids, vec![10, 11]); assert_eq!(done.finish_reason, "stop"); - assert!(stream.message().await.unwrap().is_none()); + assert!(stream.message().bounded().await.unwrap().is_none()); } /// EngineCore cannot match string stops: the servicer strips them from the @@ -491,14 +521,14 @@ async fn string_stops_are_matched_by_the_servicer() { .unwrap(); assert_eq!( - chunk_tokens(stream.message().await.unwrap().unwrap()), + chunk_tokens(stream.message().bounded().await.unwrap().unwrap()), vec![1] ); assert_eq!( - chunk_tokens(stream.message().await.unwrap().unwrap()), + chunk_tokens(stream.message().bounded().await.unwrap().unwrap()), vec![2] ); - let done = complete(stream.message().await.unwrap().unwrap()); + let done = complete(stream.message().bounded().await.unwrap().unwrap()); assert_eq!(done.finish_reason, "stop"); assert_eq!(done.output_ids, vec![1, 2]); assert_eq!( @@ -507,7 +537,7 @@ async fn string_stops_are_matched_by_the_servicer() { "Hello world".to_string() )) ); - assert!(stream.message().await.unwrap().is_none()); + assert!(stream.message().bounded().await.unwrap().is_none()); // The engine is still generating from its point of view: it gets the // abort for the choice the servicer ended. assert_eq!(recv_abort(&mut h.engine_in).await, vec!["r3".to_string()]); @@ -550,15 +580,15 @@ async fn an_engine_finish_on_the_matching_tick_keeps_the_engine_complete() { if streaming { assert_eq!( - chunk_tokens(stream.message().await.unwrap().unwrap()), + chunk_tokens(stream.message().bounded().await.unwrap().unwrap()), vec![1] ); assert_eq!( - chunk_tokens(stream.message().await.unwrap().unwrap()), + chunk_tokens(stream.message().bounded().await.unwrap().unwrap()), vec![2] ); } - let done = complete(stream.message().await.unwrap().unwrap()); + let done = complete(stream.message().bounded().await.unwrap().unwrap()); assert_eq!(done.output_ids, vec![1, 2], "streaming={streaming}"); assert_eq!(done.finish_reason, "stop"); assert_eq!(done.completion_tokens, 2); @@ -566,7 +596,7 @@ async fn an_engine_finish_on_the_matching_tick_keeps_the_engine_complete() { done.matched_stop, Some(vllm::generate_complete::MatchedStop::MatchedTokenId(2)) ); - assert!(stream.message().await.unwrap().is_none()); + assert!(stream.message().bounded().await.unwrap().is_none()); } } @@ -603,7 +633,7 @@ async fn abort_rpc_cancels_an_in_flight_stream() { .await .unwrap(); assert_eq!( - chunk_tokens(stream.message().await.unwrap().unwrap()), + chunk_tokens(stream.message().bounded().await.unwrap().unwrap()), vec![10] ); @@ -615,10 +645,10 @@ async fn abort_rpc_cancels_an_in_flight_stream() { .expect("abort"); // Ends as on the Python servicer: a terminal `abort` Complete with the // output so far, then the stream closes; the engine side is aborted. - let aborted = complete(stream.message().await.unwrap().unwrap()); + let aborted = complete(stream.message().bounded().await.unwrap().unwrap()); assert_eq!(aborted.finish_reason, "abort"); assert_eq!(aborted.output_ids, vec![10]); - assert!(stream.message().await.unwrap().is_none()); + assert!(stream.message().bounded().await.unwrap().is_none()); assert_eq!(recv_abort(&mut h.engine_in).await, vec!["r4".to_string()]); // An unknown id is a no-op, not an error (idempotent cleanup). h.client @@ -665,7 +695,7 @@ async fn kv_transfer_params_pass_through_both_ways() { })); } h.engine_out.send_outputs(&outputs).await.unwrap(); - let finished = complete(stream.message().await.unwrap().unwrap()); + let finished = complete(stream.message().bounded().await.unwrap().unwrap()); assert_eq!(finished.finish_reason, "length"); let returned: serde_json::Value = serde_json::from_str( finished @@ -679,7 +709,7 @@ async fn kv_transfer_params_pass_through_both_ways() { let legacy = finished.kv_transfer_params.expect("legacy mirror"); assert_eq!(legacy.remote_host, "10.0.0.1"); assert_eq!(legacy.remote_port, 5600); - assert!(stream.message().await.unwrap().is_none()); + assert!(stream.message().bounded().await.unwrap().is_none()); h.server.stop(Duration::from_secs(5)).expect("clean stop"); } @@ -866,11 +896,11 @@ async fn string_stops_wait_for_min_tokens() { } for expected in [1, 2, 1, 2] { assert_eq!( - chunk_tokens(stream.message().await.unwrap().unwrap()), + chunk_tokens(stream.message().bounded().await.unwrap().unwrap()), vec![expected] ); } - let done = complete(stream.message().await.unwrap().unwrap()); + let done = complete(stream.message().bounded().await.unwrap().unwrap()); assert_eq!(done.finish_reason, "stop"); assert_eq!(done.output_ids, vec![1, 2, 1, 2]); assert_eq!( @@ -879,7 +909,7 @@ async fn string_stops_wait_for_min_tokens() { "Hello world".to_string() )) ); - assert!(stream.message().await.unwrap().is_none()); + assert!(stream.message().bounded().await.unwrap().is_none()); assert_eq!(recv_abort(&mut h.engine_in).await, vec!["mt1".to_string()]); } @@ -907,11 +937,11 @@ async fn string_stops_may_span_the_min_tokens_boundary() { } for _ in 0..3 { assert_eq!( - chunk_tokens(stream.message().await.unwrap().unwrap()), + chunk_tokens(stream.message().bounded().await.unwrap().unwrap()), vec![2] ); } - let done = complete(stream.message().await.unwrap().unwrap()); + let done = complete(stream.message().bounded().await.unwrap().unwrap()); assert_eq!(done.finish_reason, "stop"); assert_eq!(done.output_ids, vec![2, 2, 2]); assert_eq!( @@ -920,7 +950,7 @@ async fn string_stops_may_span_the_min_tokens_boundary() { "world world".to_string() )) ); - assert!(stream.message().await.unwrap().is_none()); + assert!(stream.message().bounded().await.unwrap().is_none()); assert_eq!(recv_abort(&mut h.engine_in).await, vec!["mt2".to_string()]); } @@ -951,7 +981,7 @@ async fn spec_decode_counts_reach_the_complete() { }); } h.engine_out.send_outputs(&outputs).await.unwrap(); - let done = complete(stream.message().await.unwrap().unwrap()); + let done = complete(stream.message().bounded().await.unwrap().unwrap()); assert_eq!(done.output_ids, vec![5, 6]); assert_eq!(done.spec_accepted_tokens, 5); assert_eq!(done.spec_draft_tokens, 9); @@ -1143,6 +1173,263 @@ async fn info_rpcs_report_config_and_handshake_facts() { h.server.stop(Duration::from_secs(5)).expect("clean stop"); } +/// The relay follows the publisher from the servicer's start, before any +/// gateway subscribes: batches published with nobody listening are in its +/// history, and the first subscription gets them as the engine's whole state. +#[tokio::test] +async fn the_relay_subscribes_to_the_publisher_at_boot_before_any_gateway() { + use kv_events::golden; + use zeromq::{prelude::*, PubSocket}; + + let mut publisher = PubSocket::new(); + let port = bound_port( + &publisher + .bind("tcp://127.0.0.1:0") + .await + .expect("publisher binds") + .to_string(), + ); + let mut model = model_info(); + model.kv_events_endpoint = format!("tcp://*:{port}"); + model.kv_events_topic = "kv".to_string(); + let mut h = harness(model, None).await; + let relay = h + .server + .state + .kv_relay + .clone() + .expect("a relay for the publisher"); + // The SUB connect is asynchronous: publish sequence 0 until the relay, + // with no subscriber of its own yet, has taken it (repeats are duplicates). + let batch1 = golden::bytes(golden::BATCH1); + for _ in 0..250 { + publisher + .send(golden::frame(b"kv", 0, &batch1)) + .await + .expect("publish"); + if relay.counts().relayed >= 1 { + break; + } + tokio::time::sleep(Duration::from_millis(20)).await; + } + assert_eq!( + relay.counts().relayed, + 1, + "the relay took sequence 0 before any gateway subscribed" + ); + publisher + .send(golden::frame(b"kv", 1, &golden::bytes(golden::BATCH2))) + .await + .expect("publish"); + for _ in 0..250 { + if relay.counts().relayed >= 2 { + break; + } + tokio::time::sleep(Duration::from_millis(20)).await; + } + assert_eq!(relay.counts().relayed, 2); + + // The first gateway gets both from the history: the engine's whole state. + let mut stream = h + .client + .subscribe_kv_events(common::SubscribeKvEventsRequest::default()) + .await + .expect("subscribe") + .into_inner(); + for expected in [0, 1] { + let batch = tokio::time::timeout(Duration::from_secs(5), stream.message()) + .await + .expect("a batch in time") + .expect("stream open") + .expect("a batch"); + assert_eq!(batch.sequence_number, expected); + } + assert_eq!(relay.counts().served_from_history, 1); + drop(stream); + h.server.stop(Duration::from_secs(5)).expect("clean stop"); +} + +/// Every relayed batch carries the servicer's load record: the figures +/// `GetLoads` answers with, so the gateway reads the queue, running set, KV +/// usage and window at every scheduler step. +#[tokio::test] +async fn relayed_batches_carry_the_servicers_load_record() { + use zeromq::{prelude::*, PubSocket}; + + let mut publisher = PubSocket::new(); + let port = bound_port( + &publisher + .bind("tcp://127.0.0.1:0") + .await + .expect("publisher binds") + .to_string(), + ); + let mut model = model_info(); + model.kv_events_endpoint = format!("tcp://*:{port}"); + model.kv_events_topic = "kv".to_string(); + let mut h = harness(model, None).await; + let mut stream = h + .client + .subscribe_kv_events(common::SubscribeKvEventsRequest::default()) + .await + .expect("subscribe") + .into_inner(); + let batch1 = kv_events::golden::bytes(kv_events::golden::BATCH1); + let mut first = None; + for _ in 0..200 { + publisher + .send(kv_events::golden::frame(b"kv", 0, &batch1)) + .await + .expect("publish"); + if let Ok(item) = tokio::time::timeout(Duration::from_millis(50), stream.message()).await { + first = Some(item.expect("stream open").expect("a batch")); + break; + } + } + let first = first.expect("the subscription went live"); + let record = first.load.expect("the batch carries the load record"); + let loads = h + .client + .get_loads(vllm::GetLoadsRequest::default()) + .await + .expect("loads") + .into_inner(); + let rank0 = &loads.loads[0]; + assert_eq!( + ( + record.running_requests, + record.waiting_requests, + record.max_running_requests + ), + ( + u32::try_from(rank0.num_running_reqs).unwrap(), + u32::try_from(rank0.num_waiting_reqs).unwrap(), + u32::try_from(rank0.max_running_requests).unwrap() + ) + ); + assert_eq!( + record.waiting_uncached_tokens, + Some(u32::try_from(rank0.num_waiting_uncached_tokens).unwrap()), + "the vLLM servicer estimates the queued token-work" + ); + assert!((record.token_usage - rank0.token_usage).abs() < f64::EPSILON); + assert!(record.sample >= 1 && !record.load_only); + drop(stream); + h.server.stop(Duration::from_secs(5)).expect("clean stop"); +} + +/// The relay asks the engine's replay for the batches it missed before its +/// subscription joined: a publisher already at sequence 3 when the servicer +/// starts, whose replay covers 0..=3, leaves the window whole from the +/// publisher's first batch, and the first gateway gets all four from it. +#[tokio::test] +async fn a_publisher_already_counting_when_the_servicer_starts_is_replayed_from_its_start() { + use kv_events::golden; + use zeromq::{prelude::*, PubSocket, RouterSocket, ZmqMessage}; + + let mut publisher = PubSocket::new(); + let pub_port = bound_port( + &publisher + .bind("tcp://127.0.0.1:0") + .await + .expect("publisher binds") + .to_string(), + ); + let mut router = RouterSocket::new(); + let replay_port = bound_port( + &router + .bind("tcp://127.0.0.1:0") + .await + .expect("replay socket binds") + .to_string(), + ); + let mut model = model_info(); + model.kv_events_endpoint = format!("tcp://*:{pub_port}"); + model.kv_events_replay_endpoint = format!("tcp://*:{replay_port}"); + model.kv_events_topic = "kv".to_string(); + let mut h = harness(model, None).await; + let relay = h + .server + .state + .kv_relay + .clone() + .expect("a relay for the publisher"); + // Sequences 0..=2 went out before the subscription landed: publish 3 + // until the relay asks the replay socket, which it must do from 0. + let batch1 = golden::bytes(golden::BATCH1); + let batch2 = golden::bytes(golden::BATCH2); + let mut request = None; + for _ in 0..250 { + publisher + .send(golden::frame(b"kv", 3, &batch2)) + .await + .expect("publish"); + if let Ok(message) = tokio::time::timeout(Duration::from_millis(20), router.recv()).await { + request = Some(message.expect("a replay request")); + break; + } + } + let request = request.expect("the relay asked the replay socket"); + let frames: Vec> = request.iter().map(|frame| frame.to_vec()).collect(); + assert_eq!(frames.len(), 3, "[identity, empty, start]"); + assert_eq!( + frames[2], + 0u64.to_be_bytes(), + "asked from the publisher's start" + ); + for (sequence, payload) in [(0u64, &batch1), (1, &batch2), (2, &batch2), (3, &batch2)] { + let mut reply = ZmqMessage::from(frames[0].clone()); + reply.push_back(Vec::new().into()); + reply.push_back(b"kv".to_vec().into()); + reply.push_back(sequence.to_be_bytes().to_vec().into()); + reply.push_back(payload.clone().into()); + router.send(reply).await.expect("reply"); + } + let mut end = ZmqMessage::from(frames[0].clone()); + end.push_back(Vec::new().into()); + end.push_back(Vec::new().into()); + end.push_back([0xff; 8].to_vec().into()); + end.push_back(Vec::new().into()); + router.send(end).await.expect("end marker"); + for _ in 0..250 { + if relay.counts().relayed >= 4 { + break; + } + tokio::time::sleep(Duration::from_millis(20)).await; + } + // Whichever asked first, the start replay or the late join: all four came + // from the replay socket and nothing is unknown. + let counts = relay.counts(); + assert_eq!( + ( + counts.relayed, + counts.gap_batches_recovered + counts.primed_batches, + counts.unknown_before_start + ), + (4, 4, 0), + "{counts:?}" + ); + + // The window is the publisher's whole life: the first gateway gets it. + let mut stream = h + .client + .subscribe_kv_events(common::SubscribeKvEventsRequest::default()) + .await + .expect("subscribe") + .into_inner(); + for expected in 0..=3 { + let batch = tokio::time::timeout(Duration::from_secs(5), stream.message()) + .await + .expect("a batch in time") + .expect("stream open") + .expect("a batch"); + assert_eq!(batch.sequence_number, expected); + } + assert_eq!(relay.counts().served_from_history, 1); + drop(stream); + h.server.stop(Duration::from_secs(5)).expect("clean stop"); +} + /// Representative tokenizer files plus ones the Python builder excludes. /// `tokenizer.json` is incompressible and larger than a chunk, so the bundle /// streams as several. @@ -1185,7 +1472,7 @@ async fn get_tokenizer_streams_a_bundle_the_router_loader_accepts() { let tokenizer_dir = dir.path().join("tokenizer"); fs::create_dir(&tokenizer_dir).unwrap(); let files = write_tokenizer_dir(&tokenizer_dir); - let mut config = config(dir.path(), &handshake_address(), model_info()); + let mut config = config(dir.path(), &handshake_address(dir.path()), model_info()); config.tokenizer_dir = Some(tokenizer_dir.to_string_lossy().into_owned()); // The bundle comes off the configured directory; no engine is needed. let server = VllmServicerServer::start(config).expect("servicer starts"); @@ -1199,7 +1486,7 @@ async fn get_tokenizer_streams_a_bundle_the_router_loader_accepts() { .expect("get_tokenizer") .into_inner(); let mut chunks = Vec::new(); - while let Some(chunk) = stream.message().await.unwrap() { + while let Some(chunk) = stream.message().bounded().await.unwrap() { chunks.push(chunk); } let (last, full) = chunks.split_last().expect("at least one chunk"); @@ -1258,8 +1545,12 @@ async fn get_tokenizer_streams_a_bundle_the_router_loader_accepts() { #[tokio::test] async fn get_tokenizer_without_a_tokenizer_dir_is_refused() { let dir = tempfile::tempdir().unwrap(); - let server = VllmServicerServer::start(config(dir.path(), &handshake_address(), model_info())) - .expect("servicer starts"); + let server = VllmServicerServer::start(config( + dir.path(), + &handshake_address(dir.path()), + model_info(), + )) + .expect("servicer starts"); let mut client = VllmEngineClient::connect(format!("http://{}", server.address())) .await .unwrap(); @@ -1277,7 +1568,7 @@ async fn get_tokenizer_without_a_tokenizer_dir_is_refused() { let dir = tempfile::tempdir().unwrap(); let missing = dir.path().join("missing"); - let mut config = config(dir.path(), &handshake_address(), model_info()); + let mut config = config(dir.path(), &handshake_address(dir.path()), model_info()); config.tokenizer_dir = Some(missing.to_string_lossy().into_owned()); let server = VllmServicerServer::start(config).expect("servicer starts"); let mut client = VllmEngineClient::connect(format!("http://{}", server.address())) @@ -1299,8 +1590,8 @@ async fn get_tokenizer_without_a_tokenizer_dir_is_refused() { /// `SubscribeKvEvents` without a publisher is UNIMPLEMENTED with the Python /// servicer's message. With one configured, the call resolves before any /// event (headers go out eagerly, as the Python relay's initial metadata), -/// batches arrive under the publisher's sequence numbers, and dropping the -/// stream closes the subscription on the publisher's side. +/// batches arrive under the publisher's sequence numbers, and stopping the +/// servicer closes the subscription on the publisher's side. #[tokio::test] async fn subscribe_kv_events_relays_a_publisher_or_is_unimplemented() { use zeromq::{prelude::*, PubSocket, SocketEvent}; @@ -1316,13 +1607,15 @@ async fn subscribe_kv_events_relays_a_publisher_or_is_unimplemented() { assert_eq!(status.message(), kv_events::VLLM_DISABLED_MESSAGE); h.server.stop(Duration::from_secs(5)).expect("clean stop"); - let port = pick_unused_port().expect("a free publisher port"); let mut publisher = PubSocket::new(); let mut monitor = publisher.monitor(); - publisher - .bind(&format!("tcp://127.0.0.1:{port}")) - .await - .expect("publisher binds"); + let port = bound_port( + &publisher + .bind("tcp://127.0.0.1:0") + .await + .expect("publisher binds") + .to_string(), + ); let mut model = model_info(); // A bind wildcard, as vLLM's config spells it; the relay resolves it. model.kv_events_endpoint = format!("tcp://*:{port}"); @@ -1385,9 +1678,13 @@ async fn subscribe_kv_events_relays_a_publisher_or_is_unimplemented() { assert_eq!(stored.blocks[0].block_hash, 42); assert_eq!(stored.blocks[0].token_ids, vec![100, 101]); + // The relay keeps its publisher subscription for the servicer's + // lifetime (its history outlives any one stream); stopping the servicer + // closes it. drop(stream); + h.server.stop(Duration::from_secs(5)).expect("clean stop"); let disconnected = tokio::time::timeout(Duration::from_secs(5), async { - while let Some(event) = monitor.next().await { + while let Some(event) = monitor.next().bounded().await { if matches!(event, SocketEvent::Disconnected(_)) { return true; } @@ -1395,9 +1692,8 @@ async fn subscribe_kv_events_relays_a_publisher_or_is_unimplemented() { false }) .await - .expect("the publisher notices the dropped stream in time"); + .expect("the publisher notices the stopped servicer in time"); assert!(disconnected); - h.server.stop(Duration::from_secs(5)).expect("clean stop"); } /// `FlushCache`, under the Python servicer's `admin.flush_cache` contract. @@ -2212,8 +2508,8 @@ async fn a_pd_prefill_leg_returns_the_media_identity() { )) .await .unwrap(); - let _chunk = stream.message().await.unwrap().unwrap(); - let done = complete(stream.message().await.unwrap().unwrap()); + let _chunk = stream.message().bounded().await.unwrap().unwrap(); + let done = complete(stream.message().bounded().await.unwrap().unwrap()); assert_eq!(done.media_identity, Some(identity)); h.server.stop(Duration::from_secs(5)).expect("clean stop"); } @@ -2447,11 +2743,11 @@ async fn a_caller_leaving_an_admitted_decode_leg_sends_no_notice() { } /// An engine that never dials in fails the link at the configured bound, the -/// server stays up to report it, and `stop` releases the handshake port. +/// server stays up to report it, and `stop` releases the handshake endpoint. #[tokio::test] async fn an_engine_that_never_dials_in_fails_the_link_at_the_startup_bound() { let dir = tempfile::tempdir().unwrap(); - let handshake = handshake_address(); + let handshake = handshake_address(dir.path()); let server = VllmServicerServer::start(VllmServicerConfig { engine_startup_timeout: Duration::from_millis(300), ..config(dir.path(), &handshake, model_info()) @@ -2473,6 +2769,5 @@ async fn an_engine_that_never_dials_in_fails_the_link_at_the_startup_bound() { ); server.stop(Duration::from_secs(5)).unwrap(); - let port = handshake.rsplit(':').next().unwrap(); - std::net::TcpListener::bind(format!("127.0.0.1:{port}")).expect("handshake port released"); + assert_handshake_released(&handshake).await; } diff --git a/crates/engine_servicer/tests/kv_event_snapshot.rs b/crates/engine_servicer/tests/kv_event_snapshot.rs new file mode 100644 index 0000000000..df18086b0a --- /dev/null +++ b/crates/engine_servicer/tests/kv_event_snapshot.rs @@ -0,0 +1,676 @@ +//! The relay's state snapshot against generated publisher streams in the +//! engines' own hashes: the live set the snapshot emits equals a reference +//! replay of the normalized stream, copies and tiers included; a live parent +//! always precedes its children; the emitted chains re-verify under the +//! engine's own hash, which needs every parent's digest before its child and +//! so fails on any ordering slip; and the chunks are stamped the way the +//! gateway's cursor needs them. +//! +//! The streams come from a seeded generator, not from recordings: two DP +//! ranks storing chains block by block after live parents, second physical +//! copies, removals of live blocks, host-tier copies and their evictions, +//! and one clear of a rank midway, every hash computed the way the engine +//! computes it so the relay's hash check verifies the whole stream. + +#![allow( + clippy::expect_used, + clippy::unwrap_used, + clippy::panic, + clippy::print_stderr +)] + +use std::collections::{BTreeMap, HashSet}; + +use engine_servicer::{ + engine_hash::{self, Digest32, EngineHash}, + kv_state::{LiveState, SnapshotChunks, CHUNK_BLOCKS}, + kv_wire::{ + BlockHash, Counts, EventTail, ExtraKey, Normalizer, WireBatch, WireEvent, WireRemoved, + WireStored, WireTokens, + }, +}; +use smg_grpc_client::common_proto::{ + kv_block_extra_key, kv_cache_event, KvBlocksStored, KvEventBatch, +}; + +/// `(rank, tier, hash)` with its physical copies. +type LiveSet = BTreeMap<(Option, i32, i64), u32>; + +const BLOCK_SIZE: usize = 16; +const RANKS: usize = 2; + +// --------------------------------------------------------------------------- +// The generated publisher +// --------------------------------------------------------------------------- + +/// xorshift64*: deterministic and dependency-free. +struct Rng(u64); + +impl Rng { + fn next(&mut self) -> u64 { + let mut x = self.0; + x ^= x >> 12; + x ^= x << 25; + x ^= x >> 27; + self.0 = x; + x.wrapping_mul(0x2545_F491_4F6C_DD1D) + } + + fn below(&mut self, n: usize) -> usize { + (self.next() % n as u64) as usize + } + + fn chance(&mut self, percent: u64) -> bool { + self.next() % 100 < percent + } +} + +/// A block the publisher has stored: what a child chains on, what a second +/// copy or a host backup repeats, and how many physical copies it has. +#[derive(Clone)] +struct Published { + hash: i64, + digest: Digest32, + tokens: Vec, + parent: Option, + copies: u32, +} + +#[derive(Default)] +struct RankModel { + /// Device blocks with their record intact at the relay (stored, never + /// removed since): parents, repeats and host backups come from here. + live: Vec, + /// Blocks a copy of which was removed: no longer parents at the relay, + /// still holding copies the engine frees one at a time. + fading: Vec, + /// Blocks backed up to the host tier and not yet evicted there. + host: Vec, + /// Children holding copies, per block: only a block without any is + /// freed, as the engines free the tail of a chain first, so a live child + /// always has its parent in the state. + children: BTreeMap, +} + +impl RankModel { + fn held(&mut self, hash: i64) { + *self.children.entry(hash).or_default() += 1; + } + + /// The last copy of a block with `parent` went. + fn released(&mut self, parent: Option) { + if let Some(parent) = parent { + if let Some(count) = self.children.get_mut(&parent) { + *count -= 1; + if *count == 0 { + self.children.remove(&parent); + } + } + } + } + + fn is_leaf(&self, hash: i64) -> bool { + !self.children.contains_key(&hash) + } +} + +struct Generator { + rng: Rng, + engine: EngineHash, + next_token: u32, + ranks: Vec, + /// The batch that clears rank 0. + clear_at: usize, +} + +/// One block's digest and published integer in `engine`'s algorithm. +fn block(engine: EngineHash, prior: Option<&Digest32>, tokens: &[u32]) -> (Digest32, i64) { + match engine { + EngineHash::Sglang => { + let digest = engine_hash::sglang_page(prior, tokens); + (digest, engine_hash::sglang_event_int(&digest)) + } + EngineHash::VllmSha256Cbor => { + let digest = engine_hash::vllm_block(prior, tokens, None); + (digest, engine_hash::vllm_event_int(&digest)) + } + } +} + +impl Generator { + fn new(engine: EngineHash, seed: u64, batches: usize) -> Self { + Self { + rng: Rng(seed.max(1)), + engine, + next_token: 0, + ranks: (0..RANKS).map(|_| RankModel::default()).collect(), + clear_at: batches / 3, + } + } + + /// Tokens no earlier block had, so every new block is a new hash and the + /// only repeated hashes are the deliberate second copies. + fn fresh_tokens(&mut self) -> Vec { + (0..BLOCK_SIZE) + .map(|_| { + self.next_token += 1; + self.next_token + }) + .collect() + } + + fn medium(&self, host: bool) -> &'static str { + match (self.engine, host) { + (EngineHash::Sglang, true) => "CPU_PINNED", + (EngineHash::VllmSha256Cbor, true) => "CPU", + (_, false) => "GPU", + } + } + + fn tail(&self, host: bool) -> EventTail { + let vllm = self.engine == EngineHash::VllmSha256Cbor; + EventTail { + medium: Some(self.medium(host).to_string()), + group_idx: vllm.then_some(0), + kv_cache_spec_kind: vllm.then(|| "full_attention".to_string()), + ..EventTail::default() + } + } + + fn stored( + &self, + hashes: Vec, + parent: Option, + tokens: Vec, + host: bool, + ) -> WireEvent { + WireEvent::BlockStored(WireStored { + block_hashes: hashes.into_iter().map(BlockHash).collect(), + parent_block_hash: parent.map(BlockHash), + token_ids: WireTokens::Ids(tokens), + block_size: BLOCK_SIZE as i64, + lora_id: None, + lora_name: None, + cache_salt: None, + extra_keys: None, + tail: self.tail(host), + }) + } + + fn removed(&self, hashes: Vec, host: bool) -> WireEvent { + WireEvent::BlockRemoved(WireRemoved { + block_hashes: hashes.into_iter().map(BlockHash).collect(), + tail: self.tail(host), + }) + } + + /// A chain of one to three new blocks after a live parent, or a new root. + fn chain(&mut self, rank: usize) -> WireEvent { + let parent = { + let live = &self.ranks[rank].live; + (!live.is_empty() && self.rng.chance(75)) + .then(|| live[self.rng.below(live.len())].clone()) + }; + let mut prior = parent.as_ref().map(|parent| parent.digest); + let mut prior_hash = parent.as_ref().map(|parent| parent.hash); + let mut hashes = Vec::new(); + let mut tokens = Vec::new(); + for _ in 0..=self.rng.below(3) { + let block_tokens = self.fresh_tokens(); + let (digest, hash) = block(self.engine, prior.as_ref(), &block_tokens); + let model = &mut self.ranks[rank]; + model.live.push(Published { + hash, + digest, + tokens: block_tokens.clone(), + parent: prior_hash, + copies: 1, + }); + if let Some(prior_hash) = prior_hash { + model.held(prior_hash); + } + hashes.push(hash); + tokens.extend(block_tokens); + prior = Some(digest); + prior_hash = Some(hash); + } + self.stored(hashes, parent.map(|parent| parent.hash), tokens, false) + } + + /// A second physical copy of a live block whose parent is still live + /// (so the repeat verifies), capped like the gateway counts copies. + fn second_copy(&mut self, rank: usize) -> Option { + let model = &self.ranks[rank]; + let live_hashes: HashSet = model.live.iter().map(|block| block.hash).collect(); + let candidates: Vec = (0..model.live.len()) + .filter(|&index| { + let block = &model.live[index]; + block.copies < 8 + && block + .parent + .is_none_or(|parent| live_hashes.contains(&parent)) + }) + .collect(); + let index = *candidates.get(self.rng.below(candidates.len().max(1)))?; + let block = &mut self.ranks[rank].live[index]; + block.copies += 1; + let (hash, parent, tokens) = (block.hash, block.parent, block.tokens.clone()); + Some(self.stored(vec![hash], parent, tokens, false)) + } + + /// One copy of a leaf freed: a never-removed block moves to `fading`, a + /// fading block loses one more copy, the last copy releases its parent. + fn free_leaf(&mut self, rank: usize) -> Option { + let model = &self.ranks[rank]; + let mut candidates: Vec<(bool, usize)> = (0..model.live.len()) + .filter(|&index| model.is_leaf(model.live[index].hash)) + .map(|index| (true, index)) + .collect(); + candidates.extend( + (0..model.fading.len()) + .filter(|&index| model.is_leaf(model.fading[index].hash)) + .map(|index| (false, index)), + ); + let (from_live, index) = *candidates.get(self.rng.below(candidates.len().max(1)))?; + let model = &mut self.ranks[rank]; + if from_live { + let mut block = model.live.swap_remove(index); + block.copies -= 1; + let (hash, parent) = (block.hash, block.parent); + if block.copies == 0 { + model.released(parent); + } else { + model.fading.push(block); + } + return Some(hash); + } + model.fading[index].copies -= 1; + let hash = model.fading[index].hash; + if model.fading[index].copies == 0 { + let block = model.fading.swap_remove(index); + model.released(block.parent); + } + Some(hash) + } + + /// One event of `rank`: a chain, a second copy, a removal of one or two + /// leaf copies, a host backup of a live root, or a host eviction. + fn event(&mut self, rank: usize) -> WireEvent { + let roll = self.rng.below(100); + if roll < 55 || self.ranks[rank].live.is_empty() { + return self.chain(rank); + } + if roll < 70 { + if let Some(event) = self.second_copy(rank) { + return event; + } + return self.chain(rank); + } + if roll < 85 { + let mut hashes = Vec::new(); + for _ in 0..=self.rng.below(2) { + if let Some(hash) = self.free_leaf(rank) { + if !hashes.contains(&hash) { + hashes.push(hash); + } + } + } + if !hashes.is_empty() { + return self.removed(hashes, false); + } + return self.chain(rank); + } + if roll < 93 { + let roots: Vec = self.ranks[rank] + .live + .iter() + .filter(|block| block.parent.is_none()) + .cloned() + .collect(); + if !roots.is_empty() { + let root = roots[self.rng.below(roots.len())].clone(); + let event = self.stored(vec![root.hash], None, root.tokens.clone(), true); + self.ranks[rank].host.push(root); + return event; + } + } + let host = &mut self.ranks[rank].host; + if host.is_empty() { + return self.chain(rank); + } + let index = self.rng.below(host.len()); + let evicted = host.swap_remove(index).hash; + self.removed(vec![evicted], true) + } + + fn batch(&mut self, index: usize) -> WireBatch { + let rank = self.rng.below(RANKS); + let mut events = Vec::new(); + if index == self.clear_at { + events.push(WireEvent::AllBlocksCleared { ownership: None }); + self.ranks[0] = RankModel::default(); + } + let rank = if index == self.clear_at { 0 } else { rank }; + for _ in 0..=self.rng.below(4) { + events.push(self.event(rank)); + } + WireBatch { + ts: 1_700_000_000.0 + index as f64, + events, + dp_rank: Some(rank as i32), + } + } +} + +/// `batches` publisher batches of `engine`'s stream, normalized with the +/// engine-hash check on, as the relay forwards them. +fn generated(engine: EngineHash, seed: u64, batches: usize) -> (Vec, Counts) { + let mut generator = Generator::new(engine, seed, batches); + let mut normalizer = Normalizer::with_hash_check(engine); + let mut event_id = 0; + let stream = (0..batches) + .map(|index| { + let batch = generator.batch(index); + normalizer.normalize_batch(batch, index as u64, &mut event_id) + }) + .collect(); + (stream, normalizer.counts().clone()) +} + +// --------------------------------------------------------------------------- +// The reference and the snapshot's reading +// --------------------------------------------------------------------------- + +/// One batch into the live set a naive replay of the normalized stream +/// leaves: stores add a copy (capped like the gateway counts them), removals +/// take one, a clear empties the rank. +fn replay(live: &mut LiveSet, batch: &KvEventBatch) { + for event in &batch.events { + match &event.data { + Some(kv_cache_event::Data::Stored(stored)) => { + let tier = stored.tier.expect("the relay sets the tier"); + for block in &stored.blocks { + let copies = live + .entry((batch.dp_rank, tier, block.block_hash)) + .or_insert(0); + if *copies < 8 { + *copies += 1; + } + } + } + Some(kv_cache_event::Data::Removed(removed)) => { + let tier = removed.tier.expect("the relay sets the tier"); + for &hash in &removed.block_hashes { + let key = (batch.dp_rank, tier, hash); + if let Some(copies) = live.get_mut(&key) { + *copies -= 1; + if *copies == 0 { + live.remove(&key); + } + } + } + } + Some(kv_cache_event::Data::Cleared(_)) => { + live.retain(|(rank, _, _), _| *rank != batch.dp_rank); + } + None => {} + } + } +} + +fn reference(batches: &[KvEventBatch]) -> LiveSet { + let mut live = LiveSet::new(); + for batch in batches { + replay(&mut live, batch); + } + live +} + +/// Every stored event of the chunks with the rank it came under. +fn stores(chunks: &[KvEventBatch]) -> Vec<(Option, &KvBlocksStored)> { + chunks + .iter() + .flat_map(|chunk| { + chunk + .events + .iter() + .filter_map(move |event| match &event.data { + Some(kv_cache_event::Data::Stored(stored)) => Some((chunk.dp_rank, stored)), + _ => None, + }) + }) + .collect() +} + +fn emitted(chunks: &[KvEventBatch]) -> LiveSet { + let mut live = LiveSet::new(); + for (rank, stored) in stores(chunks) { + for block in &stored.blocks { + *live + .entry((rank, stored.tier.unwrap(), block.block_hash)) + .or_insert(0) += 1; + } + } + live +} + +/// A snapshot store as the publisher would have sent it, so the relay's +/// hash check can rehash the emitted chain. +fn wire_stored(stored: &KvBlocksStored) -> WireEvent { + let first = &stored.blocks[0]; + WireEvent::BlockStored(WireStored { + block_hashes: stored + .blocks + .iter() + .map(|block| BlockHash(block.block_hash)) + .collect(), + parent_block_hash: stored.parent_block_hash.map(BlockHash), + token_ids: WireTokens::Ids( + stored + .blocks + .iter() + .flat_map(|block| block.token_ids.iter().copied()) + .collect(), + ), + block_size: i64::from(first.block_size), + lora_id: first.lora_id, + lora_name: stored.lora_name.clone(), + cache_salt: stored.cache_salt.clone(), + extra_keys: Some( + stored + .blocks + .iter() + .map(|block| { + (!block.extra_keys.is_empty()).then(|| { + block + .extra_keys + .iter() + .map(|key| match key.key.clone().expect("a key") { + kv_block_extra_key::Key::Text(text) => ExtraKey::Text(text), + kv_block_extra_key::Key::Number(number) => ExtraKey::Number(number), + kv_block_extra_key::Key::Blob(blob) => ExtraKey::Blob(blob), + kv_block_extra_key::Key::Multimodal(mm) => ExtraKey::Multimodal { + identifier: mm.identifier, + offset: mm.offset, + }, + }) + .collect() + }) + }) + .collect(), + ), + tail: EventTail { + medium: stored.medium.clone(), + group_idx: stored.group_idx, + kv_cache_spec_kind: stored.kv_cache_spec_kind.clone(), + kv_cache_spec_sliding_window: stored.kv_cache_spec_sliding_window, + locality: None, + ownership: stored.ownership.clone(), + session_id: stored.session_id.clone(), + }, + }) +} + +/// Rehash the emitted stores in order with `engine`'s algorithm. +fn rehash(chunks: &[KvEventBatch], engine: EngineHash) -> Counts { + let mut checker = Normalizer::with_hash_check(engine); + let mut event_id = 0; + for (rank, stored) in stores(chunks) { + event_id += 1; + checker + .normalize(wire_stored(stored), rank, event_id) + .expect("a snapshot store is forwardable"); + } + checker.counts().clone() +} + +struct Case { + name: &'static str, + engine: EngineHash, + seed: u64, +} + +const CASES: &[Case] = &[ + Case { + name: "vllm", + engine: EngineHash::VllmSha256Cbor, + seed: 0x5eed_0001, + }, + Case { + name: "sglang", + engine: EngineHash::Sglang, + seed: 0x5eed_0002, + }, +]; + +/// Enough batches for the live set to span several chunks. +const BATCHES: usize = 4_000; + +#[test] +fn snapshots_of_the_generated_streams_equal_their_live_sets() { + for case in CASES { + let name = case.name; + let (batches, stream_counts) = generated(case.engine, case.seed, BATCHES); + // The generator speaks the engine's hash: everything verified. + assert!(stream_counts.hash_checked > 0, "{name}: blocks checked"); + assert_eq!( + (stream_counts.hash_mismatch, stream_counts.hash_unverifiable), + (0, 0), + "{name}: the generated stream verifies" + ); + assert!(stream_counts.duplicate_stores > 0, "{name}: second copies"); + assert!(stream_counts.forwarded_removed > 0, "{name}: removals"); + assert_eq!(stream_counts.forwarded_cleared, 1, "{name}: the clear"); + + let mut state = LiveState::new(); + for batch in &batches { + state.apply(batch); + } + let through = batches.iter().map(|b| b.sequence_number).max().unwrap(); + let chunks: Vec = + SnapshotChunks::new(state.snapshot(), through, 1.0, 0).collect(); + + let want = reference(&batches); + let got = emitted(&chunks); + assert_eq!(got, want, "{name}: the emitted live set"); + assert_eq!( + state.blocks(), + want.values().map(|&copies| u64::from(copies)).sum::(), + "{name}: live copies" + ); + assert_eq!(state.entries(), want.len(), "{name}: live entries"); + assert!( + want.keys().any(|(_, tier, _)| *tier != 1), + "{name}: the stream has host-tier entries" + ); + assert!(chunks.len() > 1, "{name}: the live set spans chunks"); + + // Framing: the clear first, every chunk marked, stamps up to `through`. + assert!(matches!( + chunks[0].events[0].data, + Some(kv_cache_event::Data::Cleared(_)) + )); + let count = chunks.len() as u32; + for (index, chunk) in chunks.iter().enumerate() { + let marker = chunk.snapshot.as_ref().expect("marked"); + assert_eq!( + (marker.index, marker.count), + (index as u32, count), + "{name}" + ); + assert_eq!(marker.blocks, state.blocks(), "{name}"); + assert_eq!( + chunk.sequence_number, + through + 1 - u64::from(count) + index as u64, + "{name}: stamp of chunk {index}" + ); + let blocks: usize = stores(std::slice::from_ref(chunk)) + .iter() + .map(|(_, stored)| stored.blocks.len()) + .sum(); + assert!( + blocks <= CHUNK_BLOCKS, + "{name}: chunk {index} holds {blocks}" + ); + } + + // Order: a live parent precedes its children, per rank. + let live_hashes: HashSet<(Option, i64)> = + want.keys().map(|(rank, _, hash)| (*rank, *hash)).collect(); + let mut seen: HashSet<(Option, i64)> = HashSet::new(); + for (rank, stored) in stores(&chunks) { + if let Some(parent) = stored.parent_block_hash { + assert!( + seen.contains(&(rank, parent)) || !live_hashes.contains(&(rank, parent)), + "{name}: parent {parent} of {} emitted after its child", + stored.blocks[0].block_hash + ); + } + for block in &stored.blocks { + seen.insert((rank, block.block_hash)); + } + } + + // The engine's own hash over the emitted chains: every block checked + // with a known parent, every digest reproduced. + let counts = rehash(&chunks, case.engine); + let blocks: u64 = got.values().map(|&copies| u64::from(copies)).sum(); + assert_eq!(counts.hash_checked, blocks, "{name}: every block rehashed"); + assert_eq!(counts.hash_unverifiable, 0, "{name}: a parent was missing"); + assert_eq!(counts.hash_mismatch, 0, "{name}: the chain rehashes"); + eprintln!( + "{name}: {} batches -> {} live entries, {} copies, {} chunks; stream hash check \ + {}/{} mismatched, snapshot {}/{}", + batches.len(), + state.entries(), + state.blocks(), + chunks.len(), + stream_counts.hash_mismatch, + stream_counts.hash_checked, + counts.hash_mismatch, + counts.hash_checked, + ); + } +} + +/// Cutting the stream anywhere and snapshotting there equals the reference +/// at that point: the state follows removals and clears, not just stores. +#[test] +fn snapshots_at_every_cut_of_a_stream_follow_the_reference() { + let case = &CASES[0]; + let (batches, _) = generated(case.engine, case.seed, 600); + let mut state = LiveState::new(); + let mut want = LiveSet::new(); + for batch in &batches { + state.apply(batch); + replay(&mut want, batch); + let chunks: Vec = + SnapshotChunks::new(state.snapshot(), batch.sequence_number, 1.0, 0).collect(); + assert_eq!( + emitted(&chunks), + want, + "after batch {}", + batch.sequence_number + ); + } +} diff --git a/crates/engine_zmq_adapter/src/client.rs b/crates/engine_zmq_adapter/src/client.rs index b6c9b93abe..5365756e9e 100644 --- a/crates/engine_zmq_adapter/src/client.rs +++ b/crates/engine_zmq_adapter/src/client.rs @@ -34,7 +34,8 @@ use crate::{ SglangProfileStart, }, sockets::{ - ensure_ipc_socket_dir, unlink_stale_socket, zmq_socket_addresses, ZMQ_CONNECT_TIMEOUT, + ensure_ipc_socket_dir, unlink_stale_socket, zmq_socket_addresses, Handshake, + ZMQ_CONNECT_TIMEOUT, }, stream::ZmqGenerateStream, tokenspeed::{ @@ -145,7 +146,7 @@ pub async fn connect_for_worker( base_url, model_id, runtime, - handshake_override, + Handshake::registered(handshake_override), engine_count, eos, ZMQ_CONNECT_TIMEOUT, @@ -159,17 +160,20 @@ pub async fn connect_for_worker( /// the engine's own config (the Rust gRPC servicer) and has no model dir to /// read them from. `startup_timeout` bounds the handshake: the gateway's /// connector passes [`ZMQ_CONNECT_TIMEOUT`]; a servicer that launches its own -/// engine passes what that engine's start may take. +/// engine passes what that engine's start may take. `handshake` is the +/// endpoint the engine dials: a registered worker's is tcp-only +/// ([`Handshake::registered`]); a servicer's own link may bind an `ipc://` +/// socket ([`Handshake::TcpOrIpc`]). pub async fn connect_with_eos( base_url: &str, model_id: String, runtime: RuntimeType, - handshake_override: Option<&str>, + handshake: Handshake<'_>, engine_count: usize, eos: EosTokenIds, startup_timeout: Duration, ) -> Result { - let (handshake, input, output) = zmq_socket_addresses(base_url, handshake_override)?; + let (handshake, input, output) = zmq_socket_addresses(base_url, handshake)?; ensure_ipc_socket_dir(base_url).await?; // ZMQ refuses to bind over an existing ipc socket file, so leftovers from // a dead gateway would fail every reconnect with a bare transport error. diff --git a/crates/engine_zmq_adapter/src/lib.rs b/crates/engine_zmq_adapter/src/lib.rs index 9f2263bc60..51b4175d3b 100644 --- a/crates/engine_zmq_adapter/src/lib.rs +++ b/crates/engine_zmq_adapter/src/lib.rs @@ -40,7 +40,7 @@ pub use embed::translate_embed_request; pub use engine_zmq_client::protocol::vllm::pooling::{PoolerDefaults, PoolingParams}; pub use eos::{fold_tokenizer_eos_backstop, EosTokenIds}; pub use sglang::{to_sglang_response, SglangGenerateStream, SglangProfileStart}; -pub use sockets::zmq_handshake_address; +pub use sockets::{zmq_handshake_address, Handshake}; pub use stream::ZmqGenerateStream; pub use tokenspeed::{to_tokenspeed_response, TokenSpeedGenerateStream}; pub use vllm::{ diff --git a/crates/engine_zmq_adapter/src/sockets.rs b/crates/engine_zmq_adapter/src/sockets.rs index 9362a5ac11..89ecd0a886 100644 --- a/crates/engine_zmq_adapter/src/sockets.rs +++ b/crates/engine_zmq_adapter/src/sockets.rs @@ -34,6 +34,27 @@ pub(crate) fn derive_handshake_port(path: &str) -> u16 { 20000 + (hash % 10000) as u16 } +/// The handshake endpoint a ZMQ link binds, as its caller names it. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum Handshake<'a> { + /// The loopback `tcp://` address derived from the worker's `ipc://` base. + Derived, + /// A registered worker's `zmq_handshake_address` override: `tcp://` only, + /// since the engine dials a TCP handshake and can reach nothing else. + Tcp(&'a str), + /// A servicer's own engine link: `tcp://`, or an `ipc://` socket for a + /// loopback engine (the test harness binds one per test instead of a + /// probed port two tests can share). + TcpOrIpc(&'a str), +} + +impl<'a> Handshake<'a> { + /// A registered worker's handshake: its override, or the derived address. + pub fn registered(handshake_override: Option<&'a str>) -> Self { + handshake_override.map_or(Self::Derived, Self::Tcp) + } +} + /// Derive the ZMQ socket addresses for a worker from its base URL. /// /// Mirrors vLLM's headless topology: the **handshake is TCP** (the engine dials @@ -48,20 +69,23 @@ pub(crate) fn derive_handshake_port(path: &str) -> u16 { /// `tcp://127.0.0.1:30500` (its `--data-parallel-address`/ /// `--data-parallel-rpc-port` defaults, outside the derived 20000..=29999 /// band), so setting the override to that value pairs a bare -/// `ts serve --headless` with a manually registered worker. +/// `ts serve --headless` with a manually registered worker. A servicer's own +/// link may name an `ipc://` handshake instead ([`Handshake::TcpOrIpc`]): its +/// engine is on the same host, and the test harness binds one socket per test. /// Returns `(handshake, input, output)`. /// /// [`zmq_handshake_address`] exposes just the handshake half, for the /// registration-time validation of that address. pub(crate) fn zmq_socket_addresses( base_url: &str, - handshake_override: Option<&str>, + handshake: Handshake<'_>, ) -> Result<(String, String, String), String> { let path = base_url .strip_prefix("ipc://") .ok_or_else(|| format!("ZMQ worker URL must be ipc://, got '{base_url}'"))?; - let handshake = match handshake_override { - Some(address) => { + let handshake = match handshake { + Handshake::Derived => format!("tcp://{ZMQ_LOOPBACK_HOST}:{}", derive_handshake_port(path)), + Handshake::Tcp(address) => { if !address.starts_with("tcp://") { return Err(format!( "zmq_handshake_address must be a tcp:// address \ @@ -70,7 +94,14 @@ pub(crate) fn zmq_socket_addresses( } address.to_string() } - None => format!("tcp://{ZMQ_LOOPBACK_HOST}:{}", derive_handshake_port(path)), + Handshake::TcpOrIpc(address) => { + if !address.starts_with("tcp://") && !address.starts_with("ipc://") { + return Err(format!( + "handshake address must be tcp://host:port or ipc://, got '{address}'" + )); + } + address.to_string() + } }; let input = format!("ipc://{path}-in.sock"); let output = format!("ipc://{path}-out.sock"); @@ -87,7 +118,8 @@ pub fn zmq_handshake_address( base_url: &str, handshake_override: Option<&str>, ) -> Result { - zmq_socket_addresses(base_url, handshake_override).map(|(handshake, _, _)| handshake) + zmq_socket_addresses(base_url, Handshake::registered(handshake_override)) + .map(|(handshake, _, _)| handshake) } /// Create the parent directory for a worker's `ipc://` sockets. Kept off the @@ -221,34 +253,64 @@ mod tests { #[test] fn zmq_socket_addresses_derive_handshake_by_default() { let (handshake, input, output) = - zmq_socket_addresses("ipc:///tmp/smg-zmq/ts0.ipc", None).unwrap(); + zmq_socket_addresses("ipc:///tmp/smg-zmq/ts0.ipc", Handshake::Derived).unwrap(); assert_eq!(handshake, "tcp://127.0.0.1:25152"); assert_eq!(input, "ipc:///tmp/smg-zmq/ts0.ipc-in.sock"); assert_eq!(output, "ipc:///tmp/smg-zmq/ts0.ipc-out.sock"); } #[test] - fn zmq_socket_addresses_honor_handshake_override() { - // TokenSpeed's default dial target — outside the derived band; the - // override must be bound verbatim while the data plane stays derived. - let (handshake, input, output) = - zmq_socket_addresses("ipc:///tmp/smg-zmq/ts0.ipc", Some("tcp://127.0.0.1:30500")) - .unwrap(); - assert_eq!(handshake, "tcp://127.0.0.1:30500"); + fn a_servicers_handshake_may_be_ipc_and_nothing_else_beyond_tcp() { + // The servicer's own link binds an ipc:// socket for a loopback engine + // (one per test); any other scheme is still a misconfiguration. + let (handshake, input, _) = zmq_socket_addresses( + "ipc:///tmp/smg-zmq/ts0.ipc", + Handshake::TcpOrIpc("ipc:///tmp/t/handshake"), + ) + .unwrap(); + assert_eq!(handshake, "ipc:///tmp/t/handshake"); assert_eq!(input, "ipc:///tmp/smg-zmq/ts0.ipc-in.sock"); - assert_eq!(output, "ipc:///tmp/smg-zmq/ts0.ipc-out.sock"); + let refused = zmq_socket_addresses( + "ipc:///tmp/smg-zmq/ts0.ipc", + Handshake::TcpOrIpc("udp://x:1"), + ) + .unwrap_err(); + assert!(refused.contains("tcp://"), "{refused}"); } #[test] fn zmq_socket_addresses_reject_non_tcp_override() { - // The engine dials a TCP handshake; a non-tcp override is a config - // error and must fail loudly rather than bind something unexpected. - let err = zmq_socket_addresses("ipc:///tmp/smg-zmq/ts0.ipc", Some("ipc:///tmp/hs.sock")) - .unwrap_err(); + // The engine dials a TCP handshake; a non-tcp override on a registered + // worker is a config error and must fail loudly rather than bind + // something no engine can reach. + let err = zmq_socket_addresses( + "ipc:///tmp/smg-zmq/ts0.ipc", + Handshake::Tcp("ipc:///tmp/hs.sock"), + ) + .unwrap_err(); assert!( err.contains("tcp://"), "error must name the required scheme: {err}" ); + assert!( + zmq_handshake_address("ipc:///tmp/smg-zmq/ts0.ipc", Some("ipc:///tmp/hs.sock")) + .is_err(), + "registration validation keeps the tcp-only rule" + ); + } + + #[test] + fn zmq_socket_addresses_honor_handshake_override() { + // TokenSpeed's default dial target — outside the derived band; the + // override must be bound verbatim while the data plane stays derived. + let (handshake, input, output) = zmq_socket_addresses( + "ipc:///tmp/smg-zmq/ts0.ipc", + Handshake::Tcp("tcp://127.0.0.1:30500"), + ) + .unwrap(); + assert_eq!(handshake, "tcp://127.0.0.1:30500"); + assert_eq!(input, "ipc:///tmp/smg-zmq/ts0.ipc-in.sock"); + assert_eq!(output, "ipc:///tmp/smg-zmq/ts0.ipc-out.sock"); } #[tokio::test] diff --git a/crates/engine_zmq_client/src/mock_engine.rs b/crates/engine_zmq_client/src/mock_engine.rs index 1ed9f26173..fb83a7e530 100644 --- a/crates/engine_zmq_client/src/mock_engine.rs +++ b/crates/engine_zmq_client/src/mock_engine.rs @@ -29,7 +29,12 @@ use crate::{ Error, Result, }; -const CONNECT_TIMEOUT: Duration = Duration::from_secs(5); +/// Every wait of the mock engine, the handshake's stages and each receive, +/// ends in `Error::HandshakeTimeout` naming the step after this long: a +/// servicer that never answers (or answers another test's engine) costs a +/// test thirty seconds and a message instead of a hung gate. +pub const MOCK_DEADLINE: Duration = Duration::from_secs(30); +const CONNECT_TIMEOUT: Duration = MOCK_DEADLINE; /// Per-test IPC endpoint namespace backed by a unique temporary directory, so /// concurrent tests never collide on socket paths. @@ -114,7 +119,12 @@ impl MockEngineInput { /// Receive one request's raw frames (`[request_type, payload, aux..]`); the /// DEALER identity is already stripped. pub async fn recv_frames(&mut self) -> Result> { - Ok(self.socket.recv().await?.into_vec()) + let message = within( + "mock engine: a request from the servicer", + self.socket.recv(), + ) + .await?; + Ok(message?.into_vec()) } /// Receive and classify the next request. @@ -257,6 +267,16 @@ fn peer_identity(engine_id: &EngineId) -> Result { } /// Wait for a ZMQ endpoint to become connectable before dialing it. +/// `step` under the mock's deadline; the error names the step that hung. +async fn within(step: &'static str, future: impl std::future::Future) -> Result { + timeout(MOCK_DEADLINE, future) + .await + .map_err(|_| Error::HandshakeTimeout { + stage: step, + timeout: MOCK_DEADLINE, + }) +} + async fn wait_for_endpoint(endpoint: &str) -> Result<()> { let Some(socket_path) = endpoint.strip_prefix("ipc://") else { return Ok(()); @@ -288,13 +308,19 @@ pub async fn connect_to_frontend( let mut handshake_options = SocketOptions::default(); handshake_options.peer_identity(identity.clone()); let mut handshake = DealerSocket::with_options(handshake_options); - handshake.connect(handshake_address).await?; + within( + "mock engine: connecting to the handshake endpoint", + handshake.connect(handshake_address), + ) + .await??; // HELLO -> (INIT) -> READY. handshake .send(ZmqMessage::from(encode_msgpack(&ready_message("HELLO"))?)) .await?; - let init_frames = handshake.recv().await?.into_vec(); + let init_frames = within("mock engine: INIT from the servicer", handshake.recv()) + .await?? + .into_vec(); let [init_frame] = init_frames.as_slice() else { return Err(Error::UnexpectedHandshakeMessage { message: format!("expected one INIT frame, got {}", init_frames.len()), @@ -327,14 +353,22 @@ pub async fn connect_to_frontend( let mut input_options = SocketOptions::default(); input_options.peer_identity(identity); let mut input = DealerSocket::with_options(input_options); - input.connect(input_address).await?; + within( + "mock engine: connecting to the input endpoint", + input.connect(input_address), + ) + .await??; input .send(ZmqMessage::from(encode_msgpack(&ready_response)?)) .await?; wait_for_endpoint(output_address).await?; let mut output = PushSocket::new(); - output.connect(output_address).await?; + within( + "mock engine: connecting to the output endpoint", + output.connect(output_address), + ) + .await??; Ok(MockEngine { init, diff --git a/crates/grpc_client/proto/common.proto b/crates/grpc_client/proto/common.proto index 1c01886f05..5ee3dcb9b4 100644 --- a/crates/grpc_client/proto/common.proto +++ b/crates/grpc_client/proto/common.proto @@ -26,6 +26,123 @@ message KvEventBatch { double timestamp = 2; repeated KvCacheEvent events = 3; optional int32 dp_rank = 4; + reserved 5, 6; + // Set when this batch is one chunk of a relay state snapshot rather than a + // batch the publisher sent (see KvSnapshotChunk). + KvSnapshotChunk snapshot = 7; + // The engine's load as the servicer knows it when it sends this batch (see + // EngineLoad); absent from servicers that predate the field. + EngineLoad load = 8; +} + +// The servicer's load record, attached to every KvEventBatch it streams so a +// subscriber learns the engine's queue, running set, KV usage and rate at +// every scheduler step instead of at its GetLoads poll: the same numbers +// GetLoads answers with, for the batch's dp_rank. A batch with `load_only` +// set carries no events and repeats the last sequence sent on its stream; +// the servicer sends one when the record changed without a KV event (an +// admission served from cache, a completion), coalesced per stream, and as +// a heartbeat while the engine is idle; the subscriber reads the load and +// admits nothing. +message EngineLoad { + uint32 running_requests = 1; + uint32 waiting_requests = 2; + // Queued token-work: prompt tokens of waiting requests not served from the + // prefix cache; absent when the servicer cannot tell (distinct from 0). + optional uint32 waiting_uncached_tokens = 3; + // KV-cache utilization in [0, 1], as GetLoads reports it. + double token_usage = 4; + // Tokens generated per second, as GetLoads reports it (0 when not known). + double gen_throughput = 5; + // The engine's admission window (max running requests); 0 when unknown. + uint32 max_running_requests = 6; + // Age of the sample behind this record in milliseconds (0 when the + // servicer reads the engine's current stats). + uint32 age_ms = 7; + // Monotonic per stream: the subscriber can order records. + uint64 sample = 8; + bool load_only = 9; + // The engine's telemetry beyond the routing core, filled by each servicer + // with what its engine reports and left unset otherwise (absent is + // distinct from zero): what the gateway's smg_engine_* gauges and GET + // /loads show under pushes exactly as under a GetLoads poll. Sent on the + // heartbeat cadence (load_only batches and a stream's first record), not + // on every event batch; no routing input reads these fields. + optional double cache_hit_rate = 10; + optional int32 num_used_tokens = 11; + optional int32 max_total_num_tokens = 12; + optional EngineMemory memory = 13; + optional EngineQueues queues = 14; + optional EngineSpeculative speculative = 15; + optional EngineLora lora = 16; + optional EngineDisaggregation disaggregation = 17; +} + +// The sections below mirror the SGLang scheduler's GetLoads messages field +// for field (TokenSpeed reports memory and queues in the same shape; vLLM +// reports none), so a servicer copies them through. +message EngineMemory { + double weight_gb = 1; + double kv_cache_gb = 2; + double graph_gb = 3; + int32 token_capacity = 4; +} + +message EngineQueues { + int32 waiting = 1; + int32 grammar = 2; + int32 paused = 3; + int32 retracted = 4; +} + +message EngineSpeculative { + double accept_length = 1; + double accept_rate = 2; +} + +message EngineLora { + int32 slots_used = 1; + int32 slots_total = 2; + double utilization = 3; +} + +message EngineDisaggregation { + string mode = 1; // "prefill", "decode", or "null" + int32 prefill_prealloc_queue_reqs = 2; + int32 prefill_inflight_queue_reqs = 3; + int32 decode_prealloc_queue_reqs = 4; + int32 decode_transfer_queue_reqs = 5; + int32 decode_retracted_queue_reqs = 6; + double kv_transfer_speed_gb_s = 7; + double kv_transfer_latency_ms = 8; +} + +// One chunk of a state snapshot. A relay that keeps a bounded history serves a +// subscriber without a usable cursor (start_sequence_number 0 once the history +// no longer reaches back to the publisher's first batch) the engine's live +// blocks as it recorded them from the stream, instead of live events only. The +// `count` chunks of one snapshot carry consecutive sequence numbers ending at +// the sequence the state was taken at, so live events continue after the last +// chunk with no gap and no duplicate. Chunk 0 begins with an AllBlocksCleared +// for the rank; the Stored events that follow list every live block, one event +// per physical copy, parents before children, with the fields the original +// store carried. Synthesized events have event_id 0. A subscriber applies the +// chunks as a resync of the worker's state and sets its cursor to each chunk's +// sequence number. +message KvSnapshotChunk { + // 0-based position of this chunk in the snapshot. + uint32 index = 1; + // Chunks in the snapshot. + uint32 count = 2; + // Live blocks in the whole snapshot, physical copies counted. + uint64 blocks = 3; + // Publisher sequences before the relay's record that it never saw: 0 when + // the record covers the engine's whole life (the relay joined before the + // publisher's second batch, or the engine's replay reached back to the + // start); otherwise the snapshot lacks whatever those batches stored and + // did not remove since, and the subscriber should treat the worker as + // partially known. + uint64 unknown_before = 4; } message KvCacheEvent { @@ -37,9 +154,65 @@ message KvCacheEvent { } } +// Storage tier of a KV block, from the engine's `medium` (vLLM: GPU, CPU, +// STORAGE; SGLang: GPU, CPU_PINNED, DISK, EXTERNAL). `KvBlock.cache_level` +// mirrors it for older consumers: 0 device, 1 host, 2 disk, 3 external; a +// block with neither set is on the device. +enum KvCacheTier { + KV_CACHE_TIER_UNSPECIFIED = 0; + KV_CACHE_TIER_DEVICE = 1; + KV_CACHE_TIER_HOST = 2; + KV_CACHE_TIER_DISK = 3; + KV_CACHE_TIER_EXTERNAL = 4; +} + +// Where the publishing engine says the blocks live relative to itself. Only +// local blocks belong in a per-worker prefix index; remote ones describe a +// shared pool. +enum KvCacheLocality { + KV_CACHE_LOCALITY_UNSPECIFIED = 0; + KV_CACHE_LOCALITY_LOCAL = 1; + KV_CACHE_LOCALITY_REMOTE = 2; +} + message KvBlocksStored { repeated KvBlock blocks = 1; optional int64 parent_block_hash = 2; + // Tier the blocks were stored in (see KvCacheTier) and the engine's raw + // medium string it was derived from. + optional KvCacheTier tier = 3; + optional string medium = 4; + // KV cache group the blocks belong to (hybrid models publish one group per + // attention kind; only full-attention groups feed a prefix index). + optional uint32 group_idx = 5; + optional string kv_cache_spec_kind = 6; + optional uint32 kv_cache_spec_sliding_window = 7; + optional KvCacheLocality locality = 8; + // Who manages the blocks: the engine itself, or a residency agent. + optional string ownership = 9; + // Request context that stored the blocks (attribution only). + optional string session_id = 10; + // Cache namespace: LoRA adapter and cache salt the block hashes were + // computed under. Consumers fold them into their own content hashes. + optional string lora_name = 11; + optional string cache_salt = 12; +} + +// One entry of vLLM's per-block `extra_keys`: the inputs beyond token ids +// that went into the block hash, published untagged (a LoRA name or cache +// salt as text, multimodal items as (identifier, offset), digests as bytes). +message KvBlockExtraKey { + oneof key { + string text = 1; + bytes blob = 2; + int64 number = 3; + KvMultimodalKey multimodal = 4; + } +} + +message KvMultimodalKey { + string identifier = 1; + int64 offset = 2; } message KvBlock { @@ -48,14 +221,22 @@ message KvBlock { int32 block_size = 3; optional int64 lora_id = 4; optional int32 cache_level = 5; + repeated KvBlockExtraKey extra_keys = 6; } message KvBlocksRemoved { repeated int64 block_hashes = 1; optional int32 cache_level = 2; + optional KvCacheTier tier = 3; + optional string medium = 4; + optional uint32 group_idx = 5; + optional KvCacheLocality locality = 6; + optional string ownership = 7; } -message KvCacheCleared {} +message KvCacheCleared { + optional string ownership = 1; +} // ===================== // Admin Operations diff --git a/crates/grpc_client/proto/sglang_scheduler.proto b/crates/grpc_client/proto/sglang_scheduler.proto index 31c17f53c6..57645c4ee7 100644 --- a/crates/grpc_client/proto/sglang_scheduler.proto +++ b/crates/grpc_client/proto/sglang_scheduler.proto @@ -27,7 +27,9 @@ service SglangScheduler { // Get server information rpc GetServerInfo(GetServerInfoRequest) returns (GetServerInfoResponse); - // Get comprehensive load metrics + // Get comprehensive load metrics. The gateway's fallback: a servicer that + // pushes its load on the KV-event stream (KvEventBatch.load) is not polled + // while those records flow. rpc GetLoads(GetLoadsRequest) returns (GetLoadsResponse); // Flush the KV cache on all scheduler processes diff --git a/crates/grpc_client/proto/tokenspeed_scheduler.proto b/crates/grpc_client/proto/tokenspeed_scheduler.proto index e7b9503ca7..7ccc2980e8 100644 --- a/crates/grpc_client/proto/tokenspeed_scheduler.proto +++ b/crates/grpc_client/proto/tokenspeed_scheduler.proto @@ -19,6 +19,8 @@ service TokenSpeedScheduler { rpc Abort(AbortRequest) returns (AbortResponse); rpc GetModelInfo(GetModelInfoRequest) returns (GetModelInfoResponse); rpc GetServerInfo(GetServerInfoRequest) returns (GetServerInfoResponse); + // The gateway's fallback: a servicer that pushes its load on the KV-event + // stream (KvEventBatch.load) is not polled while those records flow. rpc GetLoads(GetLoadsRequest) returns (GetLoadsResponse); // Flush the KV cache on the scheduler diff --git a/crates/grpc_client/proto/vllm_engine.proto b/crates/grpc_client/proto/vllm_engine.proto index 00d0a1c06a..3ac4989d90 100644 --- a/crates/grpc_client/proto/vllm_engine.proto +++ b/crates/grpc_client/proto/vllm_engine.proto @@ -29,7 +29,9 @@ service VllmEngine { // Get server information rpc GetServerInfo(GetServerInfoRequest) returns (GetServerInfoResponse); - // Get scheduler load metrics + // Get scheduler load metrics. The gateway's fallback: a servicer that + // pushes its load on the KV-event stream (KvEventBatch.load) is not polled + // while those records flow. rpc GetLoads(GetLoadsRequest) returns (GetLoadsResponse); // Get tokenizer artifacts for remote construction @@ -488,4 +490,9 @@ message SchedulerLoad { double cache_hit_rate = 9; double utilization = 10; int32 max_running_requests = 11; + // Queued token-work: prompt tokens of waiting requests not served from the + // prefix cache, what the gateway's expected-wait score drains at + // gen_throughput. The servicer estimates it from the requests it forwarded + // that the engine has not started (same semantics as the SGLang field). + int32 num_waiting_uncached_tokens = 12; } diff --git a/crates/grpc_client/src/channel.rs b/crates/grpc_client/src/channel.rs index 51a8b41d9e..18e427b046 100644 --- a/crates/grpc_client/src/channel.rs +++ b/crates/grpc_client/src/channel.rs @@ -33,6 +33,22 @@ pub fn normalize_grpc_endpoint(endpoint: &str) -> String { /// Matches the upstream HTTP client's connect timeout in `AppContext`. pub const DEFAULT_CONNECT_TIMEOUT: Duration = Duration::from_secs(10); +/// HTTP/2 keepalive: a ping every 30 seconds, answered within 10. The +/// interval is bounded from below by the engines' gRPC servers, not by how +/// fast a dead peer should be noticed: a grpc-core server (vLLM's and +/// SGLang's native gRPC servers) answers a ping that arrives within its +/// `min_ping_interval_without_data` (five minutes by default) of the previous +/// one with a strike, and the second strike is a GOAWAY "too_many_pings" +/// that fails every in-flight stream with `UNAVAILABLE`. A ping every second +/// did exactly that whenever a prefill batch left the connection without +/// data for two seconds. A peer that drops its connection still fails its +/// streams at once; one that falls silent is noticed at the keepalive +/// timeout or at the next poll's deadline, whichever comes first, and the +/// gateway's liveness veto turns the failure into exclusion from routing. +const KEEPALIVE_INTERVAL: Duration = Duration::from_secs(30); +const KEEPALIVE_TIMEOUT: Duration = Duration::from_secs(10); +const TCP_KEEPALIVE: Duration = Duration::from_secs(60); + /// Connect a `tonic::Channel` to the given endpoint with the SMG-standard /// keep-alive and HTTP/2 window profile applied, using /// [`DEFAULT_CONNECT_TIMEOUT`]. @@ -68,10 +84,10 @@ fn configured_endpoint( let http_endpoint = normalize_grpc_endpoint(endpoint); Ok(Endpoint::from_shared(http_endpoint)? .connect_timeout(connect_timeout) - .http2_keep_alive_interval(Duration::from_secs(30)) - .keep_alive_timeout(Duration::from_secs(10)) + .http2_keep_alive_interval(KEEPALIVE_INTERVAL) + .keep_alive_timeout(KEEPALIVE_TIMEOUT) .keep_alive_while_idle(true) - .tcp_keepalive(Some(Duration::from_secs(60))) + .tcp_keepalive(Some(TCP_KEEPALIVE)) .tcp_nodelay(true) .http2_adaptive_window(false) // 16MB stream window, 32MB connection window — sized for the diff --git a/crates/grpc_client/src/engine_load.rs b/crates/grpc_client/src/engine_load.rs new file mode 100644 index 0000000000..cb786a40fa --- /dev/null +++ b/crates/grpc_client/src/engine_load.rs @@ -0,0 +1,380 @@ +//! The pushed load record (`common::EngineLoad`, on `KvEventBatch.load`) +//! against the engines' `GetLoads` shapes: what a servicer fills from its +//! engine's per-rank report, and what the gateway reads back into its +//! per-rank snapshot, so the `smg_engine_*` gauges and `GET /loads` show the +//! same under pushes as under a poll. +//! +//! The routing core (running and waiting requests, the queued token-work, +//! KV usage, generation rate, the admission window) rides every record; the +//! telemetry (cache hit rate, used and total tokens, and SGLang's memory, +//! queue, speculative, LoRA and disaggregation sections, which TokenSpeed +//! shares for memory and queues and vLLM has none of) is filled when the +//! engine reports it and left unset otherwise. `waiting_uncached_tokens` is +//! the servicer's to set: whether the figure is an estimate for one rank or +//! the engine's own is a per-servicer rule. + +use openai_protocol::worker::{ + EngineMemoryMetricsSnapshot, EngineQueueMetricsSnapshot, SchedulerLoadSnapshot, +}; + +use crate::{common_proto as common, sglang_proto, tokenspeed_proto, vllm_proto}; + +fn unsigned(value: i32) -> u32 { + u32::try_from(value).unwrap_or(0) +} + +/// The record's core from the figures every engine reports. +fn core( + running: i32, + waiting: i32, + token_usage: f64, + gen_throughput: f64, + max_running_requests: i32, +) -> common::EngineLoad { + common::EngineLoad { + running_requests: unsigned(running), + waiting_requests: unsigned(waiting), + token_usage, + gen_throughput, + max_running_requests: unsigned(max_running_requests), + ..Default::default() + } +} + +/// vLLM's report: the core, the cache hit rate and the token counts; no +/// sections on its wire. +impl From<&vllm_proto::SchedulerLoad> for common::EngineLoad { + fn from(load: &vllm_proto::SchedulerLoad) -> Self { + Self { + cache_hit_rate: Some(load.cache_hit_rate), + num_used_tokens: Some(load.num_used_tokens), + max_total_num_tokens: Some(load.max_total_num_tokens), + ..core( + load.num_running_reqs, + load.num_waiting_reqs, + load.token_usage, + load.gen_throughput, + load.max_running_requests, + ) + } + } +} + +/// SGLang's report: the core, the cache hit rate, the token counts and +/// every optional section it carries. +impl From<&sglang_proto::SchedulerLoad> for common::EngineLoad { + fn from(load: &sglang_proto::SchedulerLoad) -> Self { + Self { + cache_hit_rate: Some(load.cache_hit_rate), + num_used_tokens: Some(load.num_used_tokens), + max_total_num_tokens: Some(load.max_total_num_tokens), + memory: load.memory.as_ref().map(|memory| common::EngineMemory { + weight_gb: memory.weight_gb, + kv_cache_gb: memory.kv_cache_gb, + graph_gb: memory.graph_gb, + token_capacity: memory.token_capacity, + }), + queues: load.queues.as_ref().map(|queues| common::EngineQueues { + waiting: queues.waiting, + grammar: queues.grammar, + paused: queues.paused, + retracted: queues.retracted, + }), + speculative: load + .speculative + .as_ref() + .map(|speculative| common::EngineSpeculative { + accept_length: speculative.accept_length, + accept_rate: speculative.accept_rate, + }), + lora: load.lora.as_ref().map(|lora| common::EngineLora { + slots_used: lora.slots_used, + slots_total: lora.slots_total, + utilization: lora.utilization, + }), + disaggregation: load.disaggregation.as_ref().map(|disagg| { + common::EngineDisaggregation { + mode: disagg.mode.clone(), + prefill_prealloc_queue_reqs: disagg.prefill_prealloc_queue_reqs, + prefill_inflight_queue_reqs: disagg.prefill_inflight_queue_reqs, + decode_prealloc_queue_reqs: disagg.decode_prealloc_queue_reqs, + decode_transfer_queue_reqs: disagg.decode_transfer_queue_reqs, + decode_retracted_queue_reqs: disagg.decode_retracted_queue_reqs, + kv_transfer_speed_gb_s: disagg.kv_transfer_speed_gb_s, + kv_transfer_latency_ms: disagg.kv_transfer_latency_ms, + } + }), + ..core( + load.num_running_reqs, + load.num_waiting_reqs, + load.token_usage, + load.gen_throughput, + load.max_running_requests, + ) + } + } +} + +/// TokenSpeed's report: the core, the cache hit rate, the token counts, and +/// the memory and queue sections. +impl From<&tokenspeed_proto::SchedulerLoad> for common::EngineLoad { + fn from(load: &tokenspeed_proto::SchedulerLoad) -> Self { + Self { + cache_hit_rate: Some(load.cache_hit_rate), + num_used_tokens: Some(load.num_used_tokens), + max_total_num_tokens: Some(load.max_total_num_tokens), + memory: load.memory.as_ref().map(|memory| common::EngineMemory { + weight_gb: memory.weight_gb, + kv_cache_gb: memory.kv_cache_gb, + graph_gb: memory.graph_gb, + token_capacity: memory.token_capacity, + }), + queues: load.queues.as_ref().map(|queues| common::EngineQueues { + waiting: queues.waiting, + grammar: queues.grammar, + paused: queues.paused, + retracted: queues.retracted, + }), + ..core( + load.num_running_reqs, + load.num_waiting_reqs, + load.token_usage, + load.gen_throughput, + load.max_running_requests, + ) + } + } +} + +/// The per-rank snapshot a poll would have produced from what the record +/// carries, the way the engines' `GetLoads` responses convert: the two +/// canonical disaggregation queue depths roll the per-stage counters up +/// (prefill = prealloc + inflight, decode = prealloc + transfer + +/// retracted), `utilization` is the KV usage, `num_total_reqs` the running +/// and waiting sum. The rank is the caller's; telemetry the record leaves +/// unset stays at the type's default, which the gateway merges with what it +/// knew from the last record that carried it. +impl From<&common::EngineLoad> for SchedulerLoadSnapshot { + fn from(record: &common::EngineLoad) -> Self { + let signed = |value: u32| i32::try_from(value).unwrap_or(i32::MAX); + let disagg = record.disaggregation.as_ref(); + Self { + num_running_reqs: signed(record.running_requests), + num_waiting_reqs: signed(record.waiting_requests), + num_waiting_uncached_tokens: record.waiting_uncached_tokens.map_or(0, signed), + num_total_reqs: signed( + record + .running_requests + .saturating_add(record.waiting_requests), + ), + num_used_tokens: record.num_used_tokens.unwrap_or(0), + max_total_num_tokens: record.max_total_num_tokens.unwrap_or(0), + token_usage: record.token_usage, + gen_throughput: record.gen_throughput, + cache_hit_rate: record.cache_hit_rate.unwrap_or(0.0), + utilization: record.token_usage, + max_running_requests: signed(record.max_running_requests), + memory: record + .memory + .as_ref() + .map(|memory| EngineMemoryMetricsSnapshot { + weight_gb: memory.weight_gb, + kv_cache_gb: memory.kv_cache_gb, + graph_gb: memory.graph_gb, + token_capacity: memory.token_capacity, + }), + queues: record + .queues + .as_ref() + .map(|queues| EngineQueueMetricsSnapshot { + waiting: queues.waiting, + grammar: queues.grammar, + paused: queues.paused, + retracted: queues.retracted, + }), + kv_transfer_latency_ms: disagg.map(|d| d.kv_transfer_latency_ms), + kv_transfer_speed_gb_s: disagg.map(|d| d.kv_transfer_speed_gb_s), + prefill_queue_reqs: disagg.map(|d| { + d.prefill_prealloc_queue_reqs + .saturating_add(d.prefill_inflight_queue_reqs) + }), + decode_queue_reqs: disagg.map(|d| { + d.decode_prealloc_queue_reqs + .saturating_add(d.decode_transfer_queue_reqs) + .saturating_add(d.decode_retracted_queue_reqs) + }), + disagg_mode: disagg.map(|d| d.mode.clone()), + ..Default::default() + } + } +} + +/// Telemetry fields of the record cleared: what an event batch carries +/// between heartbeats. +pub fn core_only(record: &mut common::EngineLoad) { + record.cache_hit_rate = None; + record.num_used_tokens = None; + record.max_total_num_tokens = None; + record.memory = None; + record.queues = None; + record.speculative = None; + record.lora = None; + record.disaggregation = None; +} + +#[cfg(test)] +mod tests { + use super::*; + + /// A SGLang rank with every section set, as the Python servicer reports + /// one. `utilization`, `num_total_reqs` and the queued token-work are + /// what the record derives (KV usage, the sum, the servicer's own), so + /// the poll's values are set to those for the comparison. + fn sglang_rank() -> sglang_proto::SchedulerLoad { + sglang_proto::SchedulerLoad { + dp_rank: 2, + num_running_reqs: 7, + num_waiting_reqs: 3, + num_total_reqs: 10, + num_used_tokens: 12_000, + max_total_num_tokens: 64_000, + token_usage: 0.1875, + gen_throughput: 1_500.5, + cache_hit_rate: 0.42, + utilization: 0.1875, + max_running_requests: 128, + num_waiting_uncached_tokens: 0, + memory: Some(sglang_proto::MemoryMetrics { + weight_gb: 15.0, + kv_cache_gb: 40.0, + graph_gb: 1.5, + token_capacity: 64_000, + }), + speculative: Some(sglang_proto::SpeculativeMetrics { + accept_length: 2.5, + accept_rate: 0.8, + }), + lora: Some(sglang_proto::LoRaMetrics { + slots_used: 2, + slots_total: 8, + utilization: 0.25, + }), + disaggregation: Some(sglang_proto::DisaggregationMetrics { + mode: "prefill".to_string(), + prefill_prealloc_queue_reqs: 4, + prefill_inflight_queue_reqs: 5, + decode_prealloc_queue_reqs: 1, + decode_transfer_queue_reqs: 2, + decode_retracted_queue_reqs: 3, + kv_transfer_speed_gb_s: 12.0, + kv_transfer_latency_ms: 3.5, + }), + queues: Some(sglang_proto::QueueMetrics { + waiting: 3, + grammar: 1, + paused: 0, + retracted: 2, + }), + } + } + + fn pushed(record: &common::EngineLoad, dp_rank: i32) -> SchedulerLoadSnapshot { + let mut snapshot = SchedulerLoadSnapshot::from(record); + snapshot.dp_rank = dp_rank; + snapshot + } + + #[test] + fn a_sglang_rank_reads_the_same_from_the_record_as_from_the_poll() { + let rank = sglang_rank(); + let polled = SchedulerLoadSnapshot::from(rank.clone()); + let record = common::EngineLoad::from(&rank); + assert_eq!(record.cache_hit_rate, Some(0.42)); + assert_eq!( + record.disaggregation.as_ref().map(|d| d.mode.as_str()), + Some("prefill") + ); + assert_eq!( + format!("{polled:?}"), + format!("{:?}", pushed(&record, rank.dp_rank)) + ); + } + + #[test] + fn a_tokenspeed_rank_reads_the_same_from_the_record_as_from_the_poll() { + let rank = tokenspeed_proto::SchedulerLoad { + dp_rank: 1, + num_running_reqs: 4, + num_waiting_reqs: 2, + num_total_reqs: 6, + num_used_tokens: 8_192, + max_total_num_tokens: 32_768, + max_running_requests: 64, + num_waiting_uncached_tokens: 0, + token_usage: 0.25, + gen_throughput: 900.0, + cache_hit_rate: 0.6, + utilization: 0.25, + memory: Some(tokenspeed_proto::MemoryMetrics { + weight_gb: 7.0, + kv_cache_gb: 20.0, + graph_gb: 0.5, + token_capacity: 32_768, + }), + queues: Some(tokenspeed_proto::QueueMetrics { + waiting: 2, + grammar: 0, + paused: 0, + retracted: 1, + }), + }; + let polled = SchedulerLoadSnapshot::from(rank); + let record = common::EngineLoad::from(&rank); + assert!(record.memory.is_some() && record.disaggregation.is_none()); + assert_eq!( + format!("{polled:?}"), + format!("{:?}", pushed(&record, rank.dp_rank)) + ); + } + + #[test] + fn a_vllm_rank_reads_the_same_from_the_record_as_from_the_poll() { + let rank = vllm_proto::SchedulerLoad { + dp_rank: 0, + num_running_reqs: 9, + num_waiting_reqs: 0, + num_total_reqs: 9, + num_used_tokens: 100_000, + max_total_num_tokens: 676_144, + token_usage: 0.1479, + gen_throughput: 2_000.0, + cache_hit_rate: 0.33, + utilization: 0.1479, + max_running_requests: 256, + num_waiting_uncached_tokens: 0, + }; + let polled = SchedulerLoadSnapshot::from(rank); + let record = common::EngineLoad::from(&rank); + assert!(record.memory.is_none() && record.queues.is_none()); + assert_eq!( + format!("{polled:?}"), + format!("{:?}", pushed(&record, rank.dp_rank)) + ); + } + + /// Absent telemetry is absent, not zero: a core-only record converts + /// with the type's defaults for the caller to merge, never with a + /// section it did not carry. + #[test] + fn a_core_only_record_carries_no_sections() { + let mut record = common::EngineLoad::from(&sglang_rank()); + core_only(&mut record); + let snapshot = SchedulerLoadSnapshot::from(&record); + assert!(record.cache_hit_rate.is_none() && record.memory.is_none()); + assert!(snapshot.memory.is_none() && snapshot.disagg_mode.is_none()); + assert_eq!( + (snapshot.num_running_reqs, snapshot.max_total_num_tokens), + (7, 0) + ); + } +} diff --git a/crates/grpc_client/src/lib.rs b/crates/grpc_client/src/lib.rs index adb6b9ccd0..40b1dd0f9d 100644 --- a/crates/grpc_client/src/lib.rs +++ b/crates/grpc_client/src/lib.rs @@ -15,6 +15,7 @@ pub mod common_proto { } pub mod abort_on_drop; pub mod channel; +pub mod engine_load; pub mod mlx_engine; pub mod sglang_scheduler; pub mod tokenizer_bundle; diff --git a/crates/grpc_client/src/sglang_scheduler.rs b/crates/grpc_client/src/sglang_scheduler.rs index 7e5f729dff..94d6caef10 100644 --- a/crates/grpc_client/src/sglang_scheduler.rs +++ b/crates/grpc_client/src/sglang_scheduler.rs @@ -869,6 +869,7 @@ impl From for openai_protocol::worker::WorkerLoadRespon dp_rank_count: resp.dp_rank_count, loads: resp.loads.into_iter().map(Into::into).collect(), aggregate, + sampled_at: None, } } } diff --git a/crates/grpc_client/src/tokenspeed_scheduler.rs b/crates/grpc_client/src/tokenspeed_scheduler.rs index f3c7cd28fb..dccdded90e 100644 --- a/crates/grpc_client/src/tokenspeed_scheduler.rs +++ b/crates/grpc_client/src/tokenspeed_scheduler.rs @@ -745,6 +745,7 @@ impl From for openai_protocol::worker::Worke dp_rank_count: resp.dp_rank_count, loads: resp.loads.into_iter().map(Into::into).collect(), aggregate, + sampled_at: None, } } } diff --git a/crates/grpc_client/src/vllm_engine.rs b/crates/grpc_client/src/vllm_engine.rs index 9496e7c1c3..255847a25c 100644 --- a/crates/grpc_client/src/vllm_engine.rs +++ b/crates/grpc_client/src/vllm_engine.rs @@ -706,8 +706,7 @@ impl From for openai_protocol::worker::SchedulerLoadSnapsh dp_rank: load.dp_rank, num_running_reqs: load.num_running_reqs, num_waiting_reqs: load.num_waiting_reqs, - // vLLM does not report queued token-work; degrade to 0. - num_waiting_uncached_tokens: 0, + num_waiting_uncached_tokens: load.num_waiting_uncached_tokens, num_total_reqs: load.num_total_reqs, num_used_tokens: load.num_used_tokens, max_total_num_tokens: load.max_total_num_tokens, @@ -731,6 +730,7 @@ impl From for openai_protocol::worker::WorkerLoadRespon dp_rank_count: resp.dp_rank_count, loads: resp.loads.into_iter().map(Into::into).collect(), aggregate: None, + sampled_at: None, } } } diff --git a/crates/kv_index/Cargo.toml b/crates/kv_index/Cargo.toml index 218e899f15..6aff1db947 100644 --- a/crates/kv_index/Cargo.toml +++ b/crates/kv_index/Cargo.toml @@ -16,6 +16,8 @@ categories = ["data-structures", "caching"] [dependencies] bincode = "1.3" blake3 = { workspace = true } +crossbeam-queue = "0.3" +crossbeam-utils = "0.8" dashmap = { workspace = true } once_cell = "1.21.4" parking_lot = { workspace = true } @@ -24,12 +26,19 @@ serde = { version = "1", features = ["derive"] } tracing = { workspace = true } xxhash-rust = { version = "0.8", features = ["xxh3"] } +[features] + [dev-dependencies] +anyhow = { workspace = true } clap = { version = "4", features = ["derive"] } criterion = { version = "0.8", features = ["html_reports"] } rand = { workspace = true } serde_json = "1" tokio = { workspace = true, features = ["rt-multi-thread", "macros", "time"] } +# mooncake_replay: CPU pinning and absolute monotonic sleeps, and one allocator for every +# backend the harness measures. +libc = "0.2" +mimalloc = { version = "0.1", default-features = false } [[bench]] name = "throughput_bench" @@ -39,5 +48,13 @@ harness = false name = "match_insert" harness = false +[[bench]] +name = "mooncake_replay" +harness = false + +[[bench]] +name = "churn" +harness = false + [lints] workspace = true diff --git a/crates/kv_index/README.md b/crates/kv_index/README.md new file mode 100644 index 0000000000..8454d3ab46 --- /dev/null +++ b/crates/kv_index/README.md @@ -0,0 +1,232 @@ +# kv_index: the KV-event index behind cache-aware routing + +The gateway's cache-aware routing keeps an index of which worker holds which prefix of which +prompt. The engines publish it: every scheduler step, vLLM and SGLang emit the block hashes they +stored, the hashes they evicted and the occasional full clear, over a ZMQ publisher that the +servicer relays into a gRPC stream the gateway subscribes to. A request is hashed into the same +per-block chain, the index answers how many leading blocks of that chain each worker holds, and +the policy credits the overlap when it picks a worker. + +This document describes the index this crate provides, the relay that feeds it and the routing +changes around it. Benchmarks and how to run them: `benches/README.md`. + +## The problem with one entry per block + +The positional indexer (`src/event_tree.rs`, the default) keys one entry per (position, content +hash) with the set of holding workers, and a lookup probes one entry per request block. It is +exact (its jump search was not: the first commits of this series fixed it against the reference +indexer and made it verify every position), but exactness costs it one dependent memory access per +block per lookup, the per-entry worker set grows with every holder, and every store or removal is +a hash-map write per block under a shard lock. A fleet of engines sharing long prefixes makes all +three worse at once: the same chain is stored by many workers, requests are hundreds of blocks +long, and the engines evict and re-store at the rate they serve. + +## The chain index + +`ChainIndex` (`src/chain_index.rs`) stores the engines' block-hash chains run-length compressed. + +**Chains and runs.** Every stored block names its parent, so each worker's blocks form chains and +the chains of a fleet form a trie keyed by block hash. The index stores that trie as runs: a run is +a maximal stretch of consecutive positions on one chain whose set of holding workers is the same at +every position. It stores one content hash and one engine hash per position (shared by every +worker that holds it) and, per run, a coverage bitset with one bit per worker slot plus a table of +partial holders, `(worker, cutoff)` entries for workers that hold a prefix of the run. A worker is +in the bitset or in the table, never both. A tail eviction lowers a cutoff or turns a full holder +into a partial one; a prefix join adds a partial entry; a decode that extends a prefix to the end +of the run promotes the worker back to the bitset. None of those split the run, so a popular +prompt prefix stays one run however many workers hold different lengths of it. + +**Children at any offset.** A chain that diverges from a run becomes a child keyed by the offset +and the next block hash in the run's child table (an open-addressing table in the arena, so the +root with tens of thousands of children inserts in constant time). The run itself is not split by +the divergence. This is what keeps the index's shape stable under the engines' churn: a request's +decode blocks are private content stored under the prompt's last block, they diverge from the +prompt chain at the prompt's end, and they die when the engine evicts them; as children they come +and go without leaving a boundary in the prompt chain. Splits remain for the two cases that need +one: a hole (a worker evicted a block in the middle of a chain it otherwise keeps, so the coverage +genuinely changes at that position) and a store whose parent sits inside a run (a stale parent). + +**Holes are exact.** Because coverage is uniform within a run, a worker that evicts a middle block +stops covering the piece that holds it and nothing else; a lookup stops there for that worker and +the chain matches up to the hole and no further, which is what the engine would serve. A later +store from the hole on heals it. + +**Lookups write nothing.** A lookup walks from the root, compares the request's content hashes +against each run on the path (one compare loop per run, not one probe per block) and ANDs the alive +set with the run's coverage; a worker scores the position where it drops out, and its partial +cutoff is applied when it has one. Readers take no locks, allocate nothing but the result map (or +nothing at all through `score_into`) and store to no shared memory: every run header, hash array +and child table lives in an arena addressed by integer ids, and a run's window (hash array, base, +length, coverage, partial table, child entry) is read under a seqlock version that is checked +again after the reads, so a split, a growth, an unlink or a reuse of the run is atomic to a +reader. Stale reads before the version check are clamped and probed with bounds checks, never +indexed. + +**Writers lock one run at a time.** A store walk descends through runs under their versions and +locks a run only when it has to change it: a join, a promotion, a split, an in-place append of the +worker's own leaf (one lock, no allocation) or a new child, which is claimed in the child table by +compare-and-swap and published with a release store, so a new chain costs the root no lock. A +walk whose plan is invalidated by a concurrent split gives up and restarts from the parent block. +A run is matched by the engine hash at the end of a window, with a bisection to the first +differing position on a mismatch; the content hash at the landing is checked, and a store that +fails it is placed by content, block by block. Dead run headers, hash arrays (reference-counted +across the runs a split leaves sharing one) and child tables return to lock-free free lists; a run +id carries a generation so a stale pointer to a reused id is recognised. Hash arrays come in +size classes a quarter of an octave apart with a best-fit search, so an array wastes at most a +quarter of its words to its class. + +**Lanes.** Events of one worker apply in the order the engine sent them, so the unit of +scheduling is the worker. Each event lane keeps, for the workers it owns, a lane map +(`src/lane_map.rs`): an open-addressing table from engine hash to (run, offset) with 16-byte slots +and a tag byte per slot, Fibonacci home slot, linear probing over the tags, backward-shift deletion +so steady churn never accumulates tombstones, and the home slots of the next keys prefetched a few +keys ahead so one event's cache misses overlap. The lane pool (`src/lane_pool.rs`) gives every +worker a bounded FIFO queue, puts a worker with queued events on exactly one lane's ready list, and +lets any lane steal a whole ready worker from a lane busy with another one, or from a lane that has +stopped running once a worker's backlog behind it is `steal_after` deep; `enqueue` never blocks +and never drops (at the cap the event comes back to the producer, which holds its cursor). The +lanes are the caller's threads, so a harness keeps its pinning and accounting. + +**Shards.** `ShardedChainIndex` (`src/sharded.rs`) holds one complete chain index per shard, each +written only by the lanes assigned to it; a worker id carries its shard, stores and removals go to +the worker's shard, and a lookup probes every shard's root and walks the shards that hold the +request's first block, reporting each worker once. The worker sets are disjoint, so the union of +the shards' scores is the single index's answer whatever the assignment. One shard is the chain +index behind one indirection, byte for byte. On a machine with several memory nodes, one shard per +node removes every cache line the event lanes of different nodes would otherwise write in common. + +**Memory.** 16 bytes per distinct block on a chain (its content hash and its engine hash), shared +by every worker that holds it, plus a run header, the coverage words and a child table per +branching run, against the 16-byte lane-map slot each lane keeps per held block for removals. The +index holds exactly what the engines report and shrinks with their removals; nothing in it ages or +is pruned. + +**Exactness.** `src/reference.rs` is a single-threaded reference indexer kept as literal as +possible: per worker, blocks by engine hash and by (position, content hash) with the prefix hashes +that reach them, stores placed after the parent, removals and clears by forgetting, a lookup that +scores the longest prefix matched position by position. Every indexer in this crate is a function +of the event stream and must equal it after any replay, on content and on every score; the tests +and benches in `tests/` and `benches/` check that on seeded corpora with holes, under concurrent +lanes with worker replacement, under churn, and, in the gateway, on recorded engine streams through +the gateway's own apply path. Guardrail tests also check that a burst of lookups leaves the index's +counters and content as they were and that a clear and a removal return every arena word but the +root's child table to a free list. + +## The relay + +The servicer relays the engine's publisher into `SubscribeKvEvents` through one relay per engine +(`crates/engine_servicer/src/kv_events.rs`, `kv_wire.rs`, `kv_history.rs`, `kv_state.rs`, +`engine_hash.rs`, `load_tracker.rs`; the Python servicers mirror the decoding and normalization in +`smg_grpc_servicer/kv_relay.py`). + +- **Wire and normalization.** Both engines' layouts (tagged maps and the older tag-first arrays) + decode into one model; a field that cannot be read costs that event, not the batch. Per stream, + the normalizer keeps only local, engine-managed blocks of the main-attention cache group, maps + the medium to a tier (device, host, disk, external), drops placeholders, unaligned and + self-referencing stores, folds speculative-decoding bigram pages to their tokens, carries the + cache namespace (LoRA name, cache salt) on every store and down the parent chain, and counts + every drop by reason. Stores and removals are forwarded one for one; the gateway counts physical + copies per worker, rank and tier, so vLLM's duplicate copies and SGLang's host copies evict a + block only when no copy remains. +- **Hash check.** The relay can rehash every admitted store the engine's own way (vLLM's + `sha256_cbor` chain and SGLang's per-page SHA-256 chain are reproduced) and count matches, + mismatches and unverifiable blocks, so a worker with another algorithm, seed or page size shows + as a rate instead of silent misses (`SMG_KV_EVENT_HASH_CHECK`). +- **History and resume.** The relay subscribes once, at the servicer's start, for the servicer's + lifetime (`SMG_KV_EVENT_RELAY_START=lazy` defers it), keeps the last `SMG_KV_EVENT_HISTORY_BATCHES` + batches within `SMG_KV_EVENT_HISTORY_BYTES` under the publisher's own sequence numbers, serves a + subscriber's cursor from it, asks the engine's replay socket for sequences its own socket missed, + and answers a cursor below the window, beyond the newest sequence or before its own start with + `OUT_OF_RANGE`, which makes the gateway clear and resubscribe from zero. +- **Snapshot.** The relay keeps the engine's live blocks per rank, tier and hash, folded from the + stream it relays. A subscription from zero whose history no longer starts at the publisher's + first batch receives a state snapshot: chunks marked `KvSnapshotChunk`, chunk 0 beginning with + a clear, stores parent-first with the original fields, consecutive sequence numbers ending at + the relay's cursor, cut atomically under the lock that admits batches; live events continue from + the next sequence with no gap and no duplicate. A fresh gateway in front of warm engines starts + with their whole resident cache. +- **Slow joiner.** A ZMQ subscriber sees nothing published before it joined, so the relay asks the + engine's replay for everything from zero at start, after a detected publisher restart, and when + a subscription from zero finds it holding nothing; what the replay no longer has is marked in the + snapshot and the gateway keeps the rank degraded until the next resync. +- **Publisher restarts** are read by rules that hold on every wire (the sequence goes backwards on + the socket, the engine's startup clear arrives under a passed sequence, the counter is back at 0 + or 1 under a larger cursor); each ends live streams with `DATA_LOSS` and starts a new incarnation. +- **Pushed loads.** Each servicer installs a load source over its own `GetLoads` figures; every + batch a subscriber receives carries the engine's load, and while the publisher is quiet the + stream sends load-only batches on change and as a heartbeat (1 s, backing off to 5 s). The vLLM + servicers derive queued uncached token-work, generation throughput and hit rate from the requests + they forward, the fields the gateway's expected-wait score drains. + +## Routing and recovery in the gateway + +- **Two backends, one type.** `KvIndex` (`model_gateway/src/worker/kv_index_backend.rs`) holds the + positional indexer or the chain index, selected by `--kv-index {positional,chain}`; the monitor + and the policy see the operations they used before. An in-gateway exactness test replays the + recorded engine streams and a synthetic hole corpus through the monitor's own apply path into + both backends and the reference indexer. +- **Per-rank admission.** One cursor per (worker, dp_rank): contiguous batches apply, duplicates + skip, a gap gets one replay request and is otherwise settled (a small one keeps the blocks and + marks the rank degraded, a large one clears the worker), a publisher restart clears the worker's + state, `OUT_OF_RANGE` and `DATA_LOSS` reset everything, a snapshot chunk applies as a resync. +- **Liveness beside the health check.** A worker that dies or restarts fails every stream on its + connection at once, and one that falls silent fails them at the keepalive timeout (the channel + profile pings every 30 s, answered within 10 s: faster pings draw a `too_many_pings` GOAWAY from + the engines' grpc-core servers) or at the next poll's deadline; a liveness tracker turns a + connection failure on the KV stream or the load poll into an `unreachable` veto after + `--worker-stall-secs`, vetoes a worker that holds requests and a growing queue without producing a + token for `--worker-wedge-secs` as `wedged` (the clock starts at the first dispatch and the bound + stretches to the prefill still in flight), re-admits a worker on its first contact with a closed + circuit breaker, and never touches the health status. A veto steers, it never refuses: when every + candidate is vetoed, the request goes to the least-loaded ready worker among them. +- **Warm-up slice.** A worker that just became routable, or whose index is thin against the + fleet's (`--worker-warmup-thin-ratio`), receives a share of cache-miss traffic until its index has + grown, so an empty index does not starve it while every prompt has a holder elsewhere. +- **Completion reporting and the selection layer.** Every request's end reaches the policy that + placed it through the worker's load guard, on every path; booked state is reconciled with the + live in-flight count each poll. Worker selection runs through a cost-function selection layer + (`model_gateway/src/policies/cost/`) whose default reproduces the existing cache-aware decision; + event-driven requests are hashed under their cache namespace. +- **Protection by default.** Worker overload protection is on as steering: a worker at or above + `--worker-overload-waiting-requests` or `--worker-overload-token-usage` is left out of selection + while another is under them, and a fleet uniformly over them is routed to its least-loaded worker; + shedding stays opt-in (`--worker-overload-shed`). + +## Flags and settings + +Gateway: `--kv-index {positional,chain}` (default `positional`; `run` is accepted as a deprecated +alias of `chain`); `--kv-indexer-ttl-secs`, `--kv-indexer-max-entries` (positional only); +`--worker-stall-secs` (2), `--worker-wedge-secs` (3); `--worker-warmup-secs` (60), +`--worker-warmup-share` (0.25), `--worker-warmup-blocks` (1024), `--worker-warmup-thin-ratio` +(0.5); `--worker-overload-protection` (on), `--disable-worker-overload-protection`, +`--worker-overload-waiting-requests` (8), `--worker-overload-token-usage` (0.8), +`--worker-overload-shed` (off); `--selection-policy` (`cache-aware-default`), +`--selection-accounting-ttl-ms` (0); +`--load-monitor-interval` (10; a `GetLoads` poll goes only to a worker whose KV-event stream pushed no load +record within the interval, the poll being the fallback for servicers that do not push). + +Servicers: `SMG_KV_EVENT_HISTORY_BATCHES` (10,000), `SMG_KV_EVENT_HISTORY_BYTES` (256 MiB), +`SMG_KV_EVENT_RELAY_START` (`lazy` to subscribe at the first gateway), +`SMG_KV_EVENT_HASH_CHECK` (`sglang` or `vllm-sha256-cbor`). + +Metrics added: `smg_kv_index_lookup_seconds{index}`, `smg_kv_event_apply_seconds{worker}`, +`smg_kv_event_blocks_total{worker,op}`, `smg_kv_event_parentless_stores_total{worker}`, +`smg_kv_index_blocks{worker}`, `smg_kv_index_memberships` and `smg_kv_index_entries` per model, +`smg_kv_index_runs_live`, `smg_kv_index_blocks_live`, `smg_kv_index_arena_bytes`, +`smg_kv_index_arena_free_bytes`, `smg_kv_index_slab_bytes`, `smg_kv_index_moved_hashes`, `smg_kv_index_engine_conflicts`; +`smg_kv_event_batches_total{disposition}`, `smg_kv_event_gaps_total{outcome}`, +`smg_kv_event_resyncs_total{reason}`, `smg_kv_event_lag_seconds`, `smg_kv_event_degraded_ranks`, +`smg_kv_event_subscriptions_total`; `smg_worker_stalled{reason}`, +`smg_worker_stall_transitions_total`, `smg_worker_overload_fallback_total`, +`smg_policy_inflight_reconciled_total{policy}`, `smg_cache_aware_policy_branch_total{branch}`. + +## Testing without engines + +`crates/mock_worker --engine realistic` is a vLLM-style engine: a pass scheduler with a token +budget, a block-level KV pool with reference counts, LRU and LIFO preemption, admission only when +the whole prompt fits, tail-first frees, KV events per pass on the gRPC stream and on the engines' +own ZMQ wires (`--kv-events-zmq-base-port`, `--kv-events-wire vllm|sglang`), an admin API with +fault hooks (drop, delay, publisher restart, pause) and engine truth, and timing from a +calibration file. The mock worker's replay binary (`replay`) replays a Mooncake-style trace through the +gateway and scores every decision against the fleet's arrival-time oracle and the engines' own +cached-token counts. The mock worker's README describes both binaries' flags. diff --git a/crates/kv_index/benches/README.md b/crates/kv_index/benches/README.md new file mode 100644 index 0000000000..4ae79a14e5 --- /dev/null +++ b/crates/kv_index/benches/README.md @@ -0,0 +1,140 @@ +# kv_index benchmarks + +| File | What it is | +| --- | --- | +| `throughput_bench.rs` | Criterion micro-benchmarks of the indexers (see its module doc). | +| `match_insert.rs` | Criterion benchmarks of single operations: a store that extends a shared chain, a store that diverges from it, a lookup over a chain many workers hold. | +| `churn.rs` | The churn bench: workers with block-LRU caches over a shared chain pool, as engines behave, with decode tails and restarts; records the index's shape over time and checks exactness against the reference indexer at every sample. | +| `mooncake_replay.rs` | Open-loop replay of an indexer corpus (an exported Mooncake-style trace schedule) against this crate's indexers: `cargo bench -p kv-index --bench mooncake_replay -- --help`. | + +Build the replay once and run the binary directly, one process per trial: + +``` +cargo bench -p kv-index --bench mooncake_replay --no-run +target/release/deps/mooncake_replay- --help +``` + +## The replay's definition + +The replay follows the open-loop indexer benchmark definition that KV routers are compared on: a +trace of requests and KV-cache events is replayed with its deadlines scaled into a window; lookups +run on query lanes and events on event lanes, sharded by worker; throughput is counted in block +operations (requested, stored and removed block hashes) over the time from the start of issue to +the last completion, drain included; latency is the lookup service time measured at each lane. + +- Query lane = `worker_id % query_lanes`. Event lane = round robin over `(worker_id, dp_rank)` in + order of first appearance. Event issuers own contiguous worker ranges (or, with + `--issuer-by-lane`, the workers of a lane range); query lanes are sharded over the query issuers + in contiguous ranges. +- Before the trial: `malloc_trim`, a quiescence sleep (`--pre-run-quiescence-ms`, default 5000), + a pass over the corpus pages, thread pinning (`--issuer-cpus`, `--query-issuer-cpus`, + `--backend-cpus`). +- Issue: each issuer sleeps to the absolute monotonic deadline and spins the last + `--issuer-spin-us`; the queries of a deadline are published before its events. +- Validity: the generator is valid when nothing failed and the issue span is at most 1.01 × the + window; `kept_up` additionally needs the last completion within 1.10 × the window. A sustained + point (see Sustained throughput below) asks for more: a valid generator and at least 99% of the offered + block ops achieved. +- Rates: `achieved_block_ops_per_sec = total_block_ops / (last_completion − start)`, drain + included; `offered` divides by the window; `actual_issue` by the issue span. Percentiles are + nearest-rank (p50, p99, p99.9, max). +- `--owned-payloads` (default on) charges the lanes what a harness that hands each lane an + owned event charges every backend: each event arrives as an owned payload (40 bytes per block, + allocated before the trial) that the lane converts into this crate's blocks and frees after the + apply, and each lookup copies its hashes into this crate's hash type. With it off, lanes read the + corpus slabs and copy nothing: the cheapest way to drive a backend, not a comparable number. + `--payload-home issuer` builds those payloads on the issuing thread instead of the main thread, + so on a multi-socket machine the payloads live on the issuer's socket. + +Backends (`--backend`): `positional` (`PositionalIndexer`), `chain` (`ChainIndex`, or +`ShardedChainIndex` with `--shards N`; `run` is its deprecated spelling), `reference` (the +single-threaded exactness reference; small corpora only) and `null` (no indexer: the harness's own +ceiling on a layout). A new index plugs in by implementing `ReplayBackend`. + +### Corpus format `SMGMCK01`, version 1 + +All integers are little-endian. One file holds one prepared schedule for one window; deadlines are +already scaled to that window and are rescaled linearly for another (`--benchmark-duration-ms`) or +for an offered rate (`--offered-block-ops-per-sec`, which sets the window from the corpus's +block-op total). + +| Section | Layout | +| --- | --- | +| Magic | 8 bytes `SMGMCK01` | +| Header | u32 version (1), u32 block_size, u64 reference_window_ns, u64 trace_duplication_factor, u64 trace_length_factor, u64 inference_worker_duplication_factor, u64 logical_workers (max worker id + 1) | +| Totals | 7 × u64: requests, stored_events, removed_events, cleared_events, request_blocks, stored_blocks, removed_blocks | +| Trace path | u64 length, UTF-8 bytes (informational) | +| Query hashes | u64 count, count × u64 local block hash | +| Stored blocks | u64 count, count × (u64 block_hash, u64 tokens_hash) | +| Removed hashes | u64 count, count × u64 block_hash | +| Operations | u64 count, then per operation: u32 id, u64 deadline_ns, u64 worker_id, u8 kind, kind-specific fields | + +| kind | Fields | +| --- | --- | +| 0 query | u64 start into the query hash slab, u32 length | +| 1 stored | u32 dp_rank, u64 event_id, u8 has_parent, u64 parent (0 when absent), u8 has_start_position, u32 start_position (0 when absent), u64 start into the stored-block slab, u32 length | +| 2 removed | u32 dp_rank, u64 event_id, u64 start into the removed-hash slab, u32 length | +| 3 cleared | u32 dp_rank, u64 event_id | + +Operations appear in issue order (deadline, queries before events at equal deadlines, worker id, +trace order) and ids are their dense positions; the loader verifies both, and the totals. The +export is deterministic: two exports of the same arguments hash identically. + +### Flags + +| Flag | Meaning | +| --- | --- | +| `--backend`, `--jump-size`, `--max-workers`, `--shards`, `--lane-memory inherit\|local` | The index under test; `--shards` places each pinned event lane's workers on the shard of the lane's NUMA node; `--lane-memory local` makes a pinned lane prefer its own node for what it allocates. | +| `--benchmark-duration-ms` or `--offered-block-ops-per-sec` | The window, or the offered rate that sets it. | +| `--query-lanes`, `--event-lanes`, `--issuer-threads`, `--query-issuer-threads` | Lane and issuer counts. | +| `--issuer-cpus`, `--query-issuer-cpus`, `--backend-cpus`, `--pin-event-lanes`, `--issuer-by-lane` | Placement. Issuer CPUs inside the lane set are refused. | +| `--owned-payloads`, `--payload-home main\|issuer` | The owned-payload cost model above. | +| `--lane-scheduling owned\|stealing`, `--steal-after` | How event lanes share work: each lane applies its own workers only, or the event lanes of each shard form one lane pool that serves whole workers and takes a worker from a lane that is not running once its backlog is `--steal-after` events deep; the pools' counters go into the result. | +| `--queries on\|off`, `--lookups all-shards\|per-shard`, `--count-shard-heads` | Diagnostics: events only; one lane per shard merged by the last to finish; a count of lookups by how many shards hold the request's first block. | +| `--result-json-output` | The result record: rates, lookup percentiles, queue depths, per-lane CPU and completion times, the layout, and a provenance object (argv, binary and corpus hashes, trace parameters). | + +The first line of every log states the issuer CPUs, the query-issuer CPUs and the lane CPU set. + +## Sustained throughput + +Sustained throughput is the highest offered rate at which a trial keeps up: generator valid and +at least 99% of the offered block operations achieved. A window-driven replay only says "keeps up +at this window", so the threshold is bracketed: start from a rate that keeps up and one that fails, +run several fresh processes at the geometric midpoint (every one of them must keep up for the point +to pass), and move the bracket until it is within 10%. `--offered-block-ops-per-sec` is the knob; +one process per trial. A published point is then a series of fresh-process trials at the kept-up +rate, each followed by a control trial of the same binary and configuration (so the pair shows the +noise floor an A/A comparison would show), summarised as medians with bootstrap confidence +intervals of achieved block ops/s and lookup p50/p99, with the trials that overlapped foreign load +on the measurement cores discarded and replaced. Capacity (the achieved rate when overloaded) is +compared only within one harness and layout. Every point is reported with the lookup percentiles +measured at that load. + +## Exactness + +Every indexer answers the same question as the single-threaded `ReferenceIndexer` +(`src/reference.rs`): after any replay, the set of (worker, position, block) and every lookup +score equal the reference's. That is checked three ways: + +- `tests/exactness_positional.rs` and `tests/exactness_chain.rs` replay seeded corpora (new + conversations, extensions, siblings diverging at any position, tail and middle evictions, + clears, worker removal and arrival) into the positional and the chain index beside the + reference, compare every lookup kind after every 256 events and the full block set at the end; + `KV_INDEX_EXACTNESS_EVENTS`, `KV_INDEX_EXACTNESS_SEED` and `KV_INDEX_EXACTNESS_SHARDS` scale, + reseed and shard them; +- `tests/concurrency_chain.rs` runs 16 event lanes and 4 readers with worker replacement and + replays every lane's log into the reference at the end; +- the churn bench and `tests/churn_gate.rs` check the chain index against the reference at every + sample while runs split and die. + +In the replay, `--backend reference` runs the reference itself on a small corpus so a backend's +result record can be compared with it. + +## Churn + +`cargo bench -p kv-index --bench churn -- --help`: workers, chains, cache blocks, decode blocks +per request, eviction order (tail-first or by hash), restart schedule, request count and sample +interval. The series it prints per sample: lookup p50/p99, runs walked per lookup, runs and blocks +live, mean run length, splits by cause, mergeable adjacent pairs, memory, and the exactness +verdict. `--json ` writes the series and the run's figures as JSON; nothing is written +without it. diff --git a/crates/kv_index/benches/churn.rs b/crates/kv_index/benches/churn.rs new file mode 100644 index 0000000000..4d81d08c40 --- /dev/null +++ b/crates/kv_index/benches/churn.rs @@ -0,0 +1,320 @@ +//! Churn bench: hours of a soak's block-LRU churn on shared chains, compressed to the index's +//! own speed, with the figures that tell fragmentation from everything else. +//! +//! The generator (`kv_index::churn`) runs workers with block-LRU caches over a shared chain +//! pool: every request is a lookup over a random prefix of a random chain, routed to the worker +//! with the longest prefix, which stores what it lacks and evicts past its capacity in the +//! configured order (the mock's tail-first order frees suffixes; hash order frees middle +//! stretches). One thread drives it: fragmentation is a property of the event sequence, not of +//! concurrency, and the lookups' service time is what is measured. +//! +//! Every `--sample-every-requests` requests a row is recorded: lookup p50/p99/p999, runs walked +//! per lookup (mean and p99), runs live, blocks live and the mean run length, the split counters +//! by cause, the mergeable adjacent pairs found by a debug walk (a run with one child, equal +//! coverage, no prefix holder of its own on the child), arena and header bytes, and the events +//! applied; and the index is checked against the reference indexer on a sample of queries. The +//! series goes to the file named by `--json`, when one is given. + +#![expect(clippy::print_stdout)] + +use std::{ + io::Write, + path::PathBuf, + time::{Duration, Instant}, +}; + +use clap::{Parser, ValueEnum}; +use kv_index::{ + churn::{Churn, ChurnConfig, FreeOrder, StepReport}, + ReferenceIndexer, ShardedChainIndex, +}; +use serde_json::json; + +#[derive(Clone, Copy, Debug, ValueEnum)] +enum FreeOrderArg { + /// The mock's order: equal-age blocks freed from the chain's tail (suffix evictions). + TailFirst, + /// Equal-age blocks freed by hash: stretches anywhere, holes included. + Hash, +} + +#[derive(Parser, Debug)] +#[command( + about = "Block-LRU churn on shared chains against the chain index, with the fragmentation series" +)] +struct Args { + #[arg(long, default_value = "128")] + workers: usize, + #[arg(long, default_value = "4000")] + chains: usize, + #[arg(long, default_value = "16")] + min_chain_len: usize, + #[arg(long, default_value = "256")] + max_chain_len: usize, + /// Probability that a chain shares a prefix with an earlier one (where runs branch). + #[arg(long, default_value = "0.7")] + share: f64, + /// Blocks a worker keeps before evicting. + #[arg(long, default_value = "20000")] + cache_blocks: usize, + #[arg(long, value_enum, default_value = "tail-first")] + free_order: FreeOrderArg, + /// Blocks a request decodes after its prompt (private tails stored block by block). + #[arg(long, default_value = "4")] + decode_blocks: usize, + /// Probability that a request decodes at all. + #[arg(long, default_value = "1.0")] + decode_probability: f64, + /// Restart (clear) one worker every this many requests and refill it cold; 0 = never. + #[arg(long, default_value = "0")] + restart_every: u64, + /// Requests routed to a restarted worker whatever the holders. + #[arg(long, default_value = "2000")] + refill_requests: u64, + #[arg(long, default_value = "1")] + shards: usize, + /// Worker slots per shard. + #[arg(long, default_value = "256")] + max_workers: usize, + #[arg(long, default_value = "2000000")] + requests: u64, + /// Stretch the run over this many minutes of wall time (sleeping between requests) instead + /// of running at the index's speed. + #[arg(long)] + minutes: Option, + #[arg(long, default_value = "100000")] + sample_every_requests: u64, + /// Queries checked against the reference indexer at every sample. + #[arg(long, default_value = "200")] + exact_samples: usize, + /// Queries checked against the reference indexer at the end. + #[arg(long, default_value = "5000")] + final_exact_samples: usize, + #[arg(long, default_value = "20261006")] + seed: u64, + /// Write the sample series and the run's figures to this file; without it nothing is + /// written (the bench also runs under `cargo test --all-targets`, in the crate directory). + #[arg(long)] + json: Option, + /// Passed by `cargo bench`; ignored. + #[arg(long, hide = true)] + bench: bool, +} + +fn percentile(sorted: &[u64], num: usize, den: usize) -> u64 { + if sorted.is_empty() { + return 0; + } + let rank = sorted.len().saturating_mul(num).div_ceil(den).max(1); + sorted[rank.saturating_sub(1).min(sorted.len() - 1)] +} + +#[derive(Default)] +struct Window { + lookup_ns: Vec, + walked: Vec, + stored: u64, + removed: u64, + healed: u64, + holders: u64, + decoded: u64, +} + +impl Window { + fn add(&mut self, r: &StepReport) { + self.lookup_ns.push(r.lookup_ns); + self.walked.push(r.runs_walked as u64); + self.stored += r.stored_blocks as u64; + self.removed += r.removed_blocks as u64; + self.healed += r.healed_blocks as u64; + self.decoded += r.decoded_blocks as u64; + self.holders += r.holders as u64; + } +} + +fn main() -> anyhow::Result<()> { + let args = Args::parse(); + let cfg = ChurnConfig { + workers: args.workers, + chains: args.chains, + min_chain_len: args.min_chain_len, + max_chain_len: args.max_chain_len, + share: args.share, + cache_blocks: args.cache_blocks, + free_order: match args.free_order { + FreeOrderArg::TailFirst => FreeOrder::TailFirst, + FreeOrderArg::Hash => FreeOrder::Hash, + }, + decode_blocks: args.decode_blocks, + decode_probability: args.decode_probability, + restart_every: args.restart_every, + refill_requests: args.refill_requests, + seed: args.seed, + }; + let index = ShardedChainIndex::new(args.shards.max(1), args.max_workers); + let mut reference = ReferenceIndexer::new(); + let mut churn = Churn::new(cfg.clone(), &index); + let divergence_points = churn.pool.divergence_points(); + let pool_blocks: usize = churn.pool.chains.iter().map(|c| c.contents.len()).sum(); + println!( + "churn: {} workers, {} chains ({} blocks, {} divergence points), cache {} blocks per worker, free order {:?}, {} shards, {} requests{}", + args.workers, + args.chains, + pool_blocks, + divergence_points, + args.cache_blocks, + cfg.free_order, + args.shards.max(1), + args.requests, + args.minutes.map_or(String::new(), |m| format!(" over {m} minutes")), + ); + let pace = args + .minutes + .map(|m| Duration::from_secs_f64(m * 60.0 / args.requests.max(1) as f64)); + let started = Instant::now(); + let mut window = Window::default(); + let mut series = Vec::new(); + let mut exact_failures = Vec::new(); + for request in 1..=args.requests { + let report = churn.step(&index, Some(&mut reference)); + window.add(&report); + if let Some(pace) = pace { + if request % 1000 == 0 { + let due = started + pace * request as u32; + let now = Instant::now(); + if due > now { + std::thread::sleep(due - now); + } + } + } + if request % args.sample_every_requests == 0 || request == args.requests { + let stats = index.stats(); + let (mergeable_pairs, mergeable_blocks) = index.debug_mergeable(); + let exact = churn.check_exact(&index, &reference, args.exact_samples); + if let Err(message) = &exact { + exact_failures.push(format!("request {request}: {message}")); + } + let mut lookups = std::mem::take(&mut window.lookup_ns); + lookups.sort_unstable(); + let mut walked = std::mem::take(&mut window.walked); + walked.sort_unstable(); + let n = lookups.len().max(1) as f64; + let row = json!({ + "requests": request, + "elapsed_s": started.elapsed().as_secs_f64(), + "lookup_p50_us": percentile(&lookups, 50, 100) as f64 / 1e3, + "lookup_p99_us": percentile(&lookups, 99, 100) as f64 / 1e3, + "lookup_p999_us": percentile(&lookups, 999, 1000) as f64 / 1e3, + "lookup_mean_us": lookups.iter().sum::() as f64 / n / 1e3, + "runs_walked_mean": walked.iter().sum::() as f64 / n, + "runs_walked_p99": percentile(&walked, 99, 100), + "holders_mean": window.holders as f64 / n, + "runs_live": stats.runs_live, + "blocks_live": stats.blocks_live, + "mean_run_len": stats.blocks_live as f64 / stats.runs_live.max(1) as f64, + "memberships": index.current_size(), + "distinct_blocks": index.entry_count(), + "splits_by_branch": stats.splits_by_branch, + "splits_by_hole": stats.splits_by_hole, + "splits_by_mid_run_store": stats.splits_by_mid_run_store, + "splits_by_prefix_holders": stats.splits_by_prefix_holders, + "runs_died": stats.runs_died, + "mergeable_pairs": mergeable_pairs, + "mergeable_blocks": mergeable_blocks, + "partial_entries": stats.partial_entries, + "max_partials": stats.max_partials, + "child_entries": stats.child_entries, + "child_tombstones": stats.child_tombstones, + "arena_bytes": stats.arena_bytes, + "header_bytes": stats.header_bytes, + "slab_bytes": stats.slab_bytes, + "stored_blocks": window.stored, + "removed_blocks": window.removed, + "healed_blocks": window.healed, + "decoded_blocks": window.decoded, + "restarts": churn.restarts, + "exact": exact.is_ok(), + }); + println!( + "{:>9} req {:7.1}s lookup p50 {:6.2} p99 {:7.2} us walked mean {:5.2} p99 {:3} | runs live {:7} blocks live {:9} mean len {:6.1} | splits branch {} hole {} mid-run {} holders {} died {} mergeable {} ({} blocks) | stored {} removed {} healed {} | exact {}", + request, + row["elapsed_s"].as_f64().unwrap_or(0.0), + row["lookup_p50_us"].as_f64().unwrap_or(0.0), + row["lookup_p99_us"].as_f64().unwrap_or(0.0), + row["runs_walked_mean"].as_f64().unwrap_or(0.0), + row["runs_walked_p99"].as_u64().unwrap_or(0), + stats.runs_live, + stats.blocks_live, + row["mean_run_len"].as_f64().unwrap_or(0.0), + stats.splits_by_branch, + stats.splits_by_hole, + stats.splits_by_mid_run_store, + stats.splits_by_prefix_holders, + stats.runs_died, + mergeable_pairs, + mergeable_blocks, + window.stored, + window.removed, + window.healed, + exact.is_ok(), + ); + std::io::stdout().flush().ok(); + series.push(row); + window = Window::default(); + } + } + let final_exact = churn.check_exact(&index, &reference, args.final_exact_samples); + if let Err(message) = &final_exact { + exact_failures.push(format!("final: {message}")); + } + let result = json!({ + "harness": "smg-churn", + "config": { + "workers": args.workers, + "chains": args.chains, + "min_chain_len": args.min_chain_len, + "max_chain_len": args.max_chain_len, + "share": args.share, + "cache_blocks": args.cache_blocks, + "free_order": format!("{:?}", cfg.free_order).to_lowercase(), + "decode_blocks": args.decode_blocks, + "decode_probability": args.decode_probability, + "restart_every": args.restart_every, + "refill_requests": args.refill_requests, + "shards": args.shards.max(1), + "requests": args.requests, + "minutes": args.minutes, + "seed": args.seed, + "pool_blocks": pool_blocks, + "divergence_points": divergence_points, + }, + "series": series, + "exact_failures": exact_failures, + "final_exact": final_exact.is_ok(), + "elapsed_s": started.elapsed().as_secs_f64(), + }); + if let Some(out) = &args.json { + std::fs::write(out, serde_json::to_vec_pretty(&result)?)?; + } + println!( + "churn done in {:.1} s; final exactness {}{}", + started.elapsed().as_secs_f64(), + if final_exact.is_ok() { "ok" } else { "FAILED" }, + if exact_failures.is_empty() { + String::new() + } else { + format!( + "; {} checkpoint failures, first: {}", + exact_failures.len(), + exact_failures[0] + ) + } + ); + anyhow::ensure!( + exact_failures.is_empty(), + "{} exactness failures, first: {}", + exact_failures.len(), + exact_failures[0] + ); + Ok(()) +} diff --git a/crates/kv_index/benches/match_insert.rs b/crates/kv_index/benches/match_insert.rs index b8c5a56abc..def18d90a6 100644 --- a/crates/kv_index/benches/match_insert.rs +++ b/crates/kv_index/benches/match_insert.rs @@ -8,7 +8,7 @@ //! warm tree; "miss" appends a never-seen turn each time; "hit_8_threads" //! runs the hit path on eight threads against one tree and reports wall time //! divided by operations. -#![allow(clippy::expect_used, clippy::cast_possible_truncation)] +#![expect(clippy::expect_used)] use std::{ cell::RefCell, @@ -23,8 +23,9 @@ use std::{ use criterion::{criterion_group, criterion_main, BatchSize, Criterion, Throughput}; use kv_index::{ - compute_content_hash, compute_request_content_hashes, PositionalIndexer, SequenceHash, - StoredBlock, TokenTree, Tree, WorkerBlockMap, + compute_content_hash, compute_request_content_hashes, request_prefix_hashes, ChainBlockMap, + ChainIndex, ContentHash, PositionalIndexer, SequenceHash, ShardedChainIndex, StoredBlock, + TokenTree, Tree, WorkerBlockMap, }; const TENANTS: usize = 64; @@ -280,10 +281,159 @@ fn bench_event_path(c: &mut Criterion) { group.finish(); } +/// Blocks of a content chain as an engine hashes them (the engine hash is the chain hash). +fn chain_blocks(stream: u64, len: usize) -> Vec { + let contents: Vec = (0..len) + .map(|p| compute_content_hash(&[stream as u32, (stream >> 32) as u32, p as u32])) + .collect(); + contents + .iter() + .zip(request_prefix_hashes(&contents)) + .map(|(&content_hash, seq_hash)| StoredBlock { + seq_hash, + content_hash, + }) + .collect() +} + +/// A chain index in which `holders` workers hold one 88-block chain: the shape of the Mooncake +/// replay, where nineteen of twenty stores land on a run another worker built. +fn shared_run(holders: usize) -> (ChainIndex, Vec) { + let index = ChainIndex::with_max_workers(64); + let blocks = chain_blocks(1, 88); + for w in 0..holders { + let worker = index.intern_worker(&tenant(w)).expect("worker id"); + let mut map = ChainBlockMap::default(); + index + .apply_stored(worker, &blocks, None, &mut map) + .expect("store"); + } + (index, blocks) +} + +/// The store walk of the chain index per event block: a worker storing a chain other workers +/// already hold (the walk matches the run and writes the lane map), a store that diverges from +/// the shared chain halfway (the walk finds the divergence and splits the run), and a fresh +/// chain under the root (a new run, allocation included) as the control. +fn bench_chain_index_store(c: &mut Criterion) { + let mut group = c.benchmark_group("chain_index_store"); + group.throughput(Throughput::Elements(88)); + + let (index, blocks) = shared_run(19); + let worker = index.intern_worker("newcomer").expect("worker id"); + let hashes: Vec = blocks.iter().map(|block| block.seq_hash).collect(); + group.bench_function("shared_chain/88_blocks/19_holders", |b| { + b.iter_batched( + || { + let mut map = ChainBlockMap::default(); + index + .apply_stored(worker, &blocks, None, &mut map) + .expect("store"); + index.apply_removed(worker, &hashes, &mut map); + map + }, + |mut map| { + index + .apply_stored(worker, &blocks, None, &mut map) + .expect("store"); + map + }, + BatchSize::SmallInput, + ); + }); + + let mut forked = blocks[..44].to_vec(); + forked.extend(chain_blocks(2, 88).into_iter().skip(44)); + let forked = { + // Re-chain the fork's engine hashes from the shared prefix. + let contents: Vec = forked.iter().map(|block| block.content_hash).collect(); + contents + .iter() + .zip(request_prefix_hashes(&contents)) + .map(|(&content_hash, seq_hash)| StoredBlock { + seq_hash, + content_hash, + }) + .collect::>() + }; + group.bench_function("divergent_chain/88_blocks/split_at_44", |b| { + b.iter_batched( + || { + let (index, _) = shared_run(19); + let worker = index.intern_worker("forker").expect("worker id"); + (index, worker) + }, + |(index, worker)| { + let mut map = ChainBlockMap::default(); + index + .apply_stored(worker, &forked, None, &mut map) + .expect("store"); + (index, map) + }, + BatchSize::SmallInput, + ); + }); + + group.bench_function("fresh_chain/88_blocks", |b| { + b.iter_batched( + || { + let index = ChainIndex::with_max_workers(64); + let worker = index.intern_worker("first").expect("worker id"); + (index, worker) + }, + |(index, worker)| { + let mut map = ChainBlockMap::default(); + index + .apply_stored(worker, &blocks, None, &mut map) + .expect("store"); + (index, map) + }, + BatchSize::SmallInput, + ); + }); + group.finish(); +} + +/// The lookup of the chain index: a request of 88 blocks that twenty workers hold whole (the +/// Mooncake shape, one run walked), and the same request against two shards holding it on both. +fn bench_chain_index_lookup(c: &mut Criterion) { + let mut group = c.benchmark_group("chain_index_lookup"); + group.throughput(Throughput::Elements(1)); + let (index, blocks) = shared_run(20); + let request: Vec = blocks.iter().map(|block| block.content_hash).collect(); + group.bench_function("88_blocks/20_holders/1_shard", |b| { + b.iter(|| { + let mut scored = 0usize; + index.score_into(&request, |content| content.0, false, |_, _| scored += 1); + scored + }); + }); + let sharded = ShardedChainIndex::new(2, 64); + for w in 0..20 { + let worker = sharded + .intern_worker_in(w % 2, &tenant(w)) + .expect("worker id"); + let mut map = ChainBlockMap::default(); + sharded + .apply_stored(worker, &blocks, None, &mut map) + .expect("store"); + } + group.bench_function("88_blocks/20_holders/2_shards", |b| { + b.iter(|| { + let mut scored = 0usize; + sharded.score_into(&request, |content| content.0, false, |_, _| scored += 1); + scored + }); + }); + group.finish(); +} + criterion_group!( benches, bench_token_tree, bench_string_tree, - bench_event_path + bench_event_path, + bench_chain_index_store, + bench_chain_index_lookup ); criterion_main!(benches); diff --git a/crates/kv_index/benches/mooncake_replay.rs b/crates/kv_index/benches/mooncake_replay.rs new file mode 100644 index 0000000000..dc672ce9a6 --- /dev/null +++ b/crates/kv_index/benches/mooncake_replay.rs @@ -0,0 +1,2783 @@ +//! Open-loop replay of an exported Mooncake indexer corpus against this crate's indexers, with +//! drain-inclusive accounting. +//! +//! The corpus is a prepared schedule of the public Mooncake trace (after trace duplication, +//! deadline sort and id assignment) in the export format described in `README.md` next to this +//! file. Replaying it here measures the indexer alone, with the harness's own threads and queues: +//! +//! - queries go to `worker_id % query_lanes` lanes, each an OS thread that services lookups +//! inline; events go to `(worker, dp_rank)`-pinned event lanes assigned round-robin on first +//! sight, each an OS thread draining an unbounded queue; +//! - one query issuer and `--issuer-threads` event issuers (events sharded by contiguous worker +//! ranges) issue at absolute deadlines (`clock_nanosleep` + a spin) and never wait for earlier +//! operations; at equal deadlines queries are published before events; +//! - timing ends at the last completion (drain included); a trial is `generator_valid` when every +//! operation was issued within 1.01× the window and `kept_up` when replay plus drain fit in +//! 1.10×; a block op is a requested, stored or removed block hash; +//! - lookup `query_service` is the time inside the indexer, `query_scheduled_to_finished` includes +//! queueing; percentiles use the nearest-rank method. +//! +//! What the harness charges every backend, stated so numbers can be read correctly: lanes are OS +//! threads parked on `std::thread::park`; event queues are `std::sync::mpsc` (unbounded, one +//! consumer). With `--owned-payloads` (default) the lanes pay for owning their events: each event +//! arrives as an owned payload in the engine's wire layout (40 bytes per block, allocated before +//! the trial) that the lane converts into this crate's 16-byte blocks and frees after the apply, +//! and each lookup copies its hashes into this crate's hash type, as an adapter in front of the +//! index would. The binary links mimalloc, so every backend is measured under one allocator. +#![expect(clippy::expect_used, clippy::print_stdout, clippy::print_stderr)] +// The harness pins threads and sleeps to absolute monotonic deadlines through libc, which the +// standard library does not expose; the five calls are wrapped in small checked helpers below. +#![expect(unsafe_code)] +#![recursion_limit = "256"] + +use std::{ + collections::BTreeMap, + hint::black_box, + sync::{ + atomic::{AtomicBool, AtomicU32, AtomicU64, AtomicUsize, Ordering}, + mpsc, Arc, Barrier, Mutex, + }, + thread, + time::{Duration, Instant}, +}; + +use clap::{Parser, ValueEnum}; +use kv_index::{ + ChainBlockMap, Claimed, ContentHash, Control, LaneHooks, LanePool, LanePoolConfig, + PositionalIndexer, QueueFull, ReferenceIndexer, SequenceHash, ShardedChainIndex, StoredBlock, + WorkerBlockMap, +}; +use rustc_hash::FxHashMap; +use serde_json::json; + +#[global_allocator] +static GLOBAL: mimalloc::MiMalloc = mimalloc::MiMalloc; + +const MAGIC: &[u8; 8] = b"SMGMCK01"; +const EMPTY_OPERATION_ID: u32 = u32::MAX; +const WARMUP_QUERIES: usize = 128; + +// --------------------------------------------------------------------------------------------- +// Corpus +// --------------------------------------------------------------------------------------------- + +#[derive(Clone, Copy, Debug, Default)] +struct Totals { + requests: u64, + stored_events: u64, + removed_events: u64, + cleared_events: u64, + request_blocks: u64, + stored_blocks: u64, + removed_blocks: u64, +} + +impl Totals { + fn events(&self) -> u64 { + self.stored_events + self.removed_events + self.cleared_events + } + fn block_ops(&self) -> u64 { + self.request_blocks + self.stored_blocks + self.removed_blocks + } + /// The block ops the dispatch schedules: all of them, or the events' alone under + /// `--queries off`. + fn scheduled_block_ops(&self, queries: Queries) -> u64 { + match queries { + Queries::On => self.block_ops(), + Queries::Off => self.stored_blocks + self.removed_blocks, + } + } +} + +#[derive(Clone, Copy, Debug)] +enum OpKind { + Query { + start: u64, + len: u32, + }, + Stored { + dp_rank: u32, + parent: Option, + start: u64, + len: u32, + }, + Removed { + dp_rank: u32, + start: u64, + len: u32, + }, + Cleared { + dp_rank: u32, + }, +} + +#[derive(Clone, Copy, Debug)] +struct Op { + id: u32, + deadline_ns: u64, + worker: u64, + kind: OpKind, +} + +impl Op { + fn is_query(&self) -> bool { + matches!(self.kind, OpKind::Query { .. }) + } + fn dp_rank(&self) -> u32 { + match self.kind { + OpKind::Query { .. } => 0, + OpKind::Stored { dp_rank, .. } + | OpKind::Removed { dp_rank, .. } + | OpKind::Cleared { dp_rank } => dp_rank, + } + } +} + +struct Corpus { + block_size: u32, + reference_window_ns: u64, + trace_duplication_factor: u64, + trace_length_factor: u64, + inference_worker_duplication_factor: u64, + logical_workers: u64, + totals: Totals, + trace_path: String, + /// blake3 of the corpus file, recorded in the result's provenance. + file_blake3: String, + hashes: Box<[ContentHash]>, + blocks: Box<[StoredBlock]>, + removed: Box<[SequenceHash]>, + ops: Vec, +} + +struct Cursor<'a> { + bytes: &'a [u8], + at: usize, +} + +impl<'a> Cursor<'a> { + fn take(&mut self, n: usize) -> anyhow::Result<&'a [u8]> { + let end = self + .at + .checked_add(n) + .filter(|&end| end <= self.bytes.len()) + .ok_or_else(|| anyhow::anyhow!("corpus truncated at byte {}", self.at))?; + let out = &self.bytes[self.at..end]; + self.at = end; + Ok(out) + } + fn u8(&mut self) -> anyhow::Result { + Ok(self.take(1)?[0]) + } + fn u32(&mut self) -> anyhow::Result { + let b = self.take(4)?; + Ok(u32::from_le_bytes([b[0], b[1], b[2], b[3]])) + } + fn u64(&mut self) -> anyhow::Result { + let b = self.take(8)?; + Ok(u64::from_le_bytes([ + b[0], b[1], b[2], b[3], b[4], b[5], b[6], b[7], + ])) + } +} + +fn load_corpus(path: &str) -> anyhow::Result { + let bytes = std::fs::read(path)?; + let file_blake3 = blake3::hash(&bytes).to_hex().to_string(); + let mut c = Cursor { + bytes: &bytes, + at: 0, + }; + anyhow::ensure!(c.take(8)? == MAGIC, "not an SMG Mooncake corpus: {path}"); + let version = c.u32()?; + anyhow::ensure!(version == 1, "unsupported corpus format version {version}"); + let block_size = c.u32()?; + let reference_window_ns = c.u64()?; + let trace_duplication_factor = c.u64()?; + let trace_length_factor = c.u64()?; + let inference_worker_duplication_factor = c.u64()?; + let logical_workers = c.u64()?; + let totals = Totals { + requests: c.u64()?, + stored_events: c.u64()?, + removed_events: c.u64()?, + cleared_events: c.u64()?, + request_blocks: c.u64()?, + stored_blocks: c.u64()?, + removed_blocks: c.u64()?, + }; + let path_len = c.u64()? as usize; + let trace_path = String::from_utf8_lossy(c.take(path_len)?).into_owned(); + let n = c.u64()? as usize; + let mut hashes = Vec::with_capacity(n); + for _ in 0..n { + hashes.push(ContentHash(c.u64()?)); + } + let n = c.u64()? as usize; + let mut blocks = Vec::with_capacity(n); + for _ in 0..n { + let seq = c.u64()?; + let content = c.u64()?; + blocks.push(StoredBlock { + seq_hash: SequenceHash(seq), + content_hash: ContentHash(content), + }); + } + let n = c.u64()? as usize; + let mut removed = Vec::with_capacity(n); + for _ in 0..n { + removed.push(SequenceHash(c.u64()?)); + } + let n = c.u64()? as usize; + let mut ops = Vec::with_capacity(n); + for _ in 0..n { + let id = c.u32()?; + let deadline_ns = c.u64()?; + let worker = c.u64()?; + let kind = match c.u8()? { + 0 => OpKind::Query { + start: c.u64()?, + len: c.u32()?, + }, + 1 => { + let dp_rank = c.u32()?; + let _event_id = c.u64()?; + let has_parent = c.u8()? != 0; + let parent = c.u64()?; + let _has_start = c.u8()?; + let _start_position = c.u32()?; + OpKind::Stored { + dp_rank, + parent: has_parent.then_some(parent), + start: c.u64()?, + len: c.u32()?, + } + } + 2 => { + let dp_rank = c.u32()?; + let _event_id = c.u64()?; + OpKind::Removed { + dp_rank, + start: c.u64()?, + len: c.u32()?, + } + } + 3 => { + let dp_rank = c.u32()?; + let _event_id = c.u64()?; + OpKind::Cleared { dp_rank } + } + other => anyhow::bail!("unknown operation kind {other}"), + }; + ops.push(Op { + id, + deadline_ns, + worker, + kind, + }); + } + anyhow::ensure!(c.at == bytes.len(), "trailing bytes in corpus"); + // The ids are the sorted positions; verify the invariants the replay relies on. + for (expected, op) in ops.iter().enumerate() { + anyhow::ensure!(op.id as usize == expected, "operation ids are not dense"); + } + anyhow::ensure!( + ops.windows(2).all(|w| w[0].deadline_ns <= w[1].deadline_ns), + "operations are not deadline-sorted" + ); + let mut recount = Totals::default(); + for op in &ops { + match op.kind { + OpKind::Query { len, .. } => { + recount.requests += 1; + recount.request_blocks += u64::from(len); + } + OpKind::Stored { len, .. } => { + recount.stored_events += 1; + recount.stored_blocks += u64::from(len); + } + OpKind::Removed { len, .. } => { + recount.removed_events += 1; + recount.removed_blocks += u64::from(len); + } + OpKind::Cleared { .. } => recount.cleared_events += 1, + } + } + anyhow::ensure!( + recount.block_ops() == totals.block_ops() && recount.events() == totals.events(), + "corpus totals disagree with its operations" + ); + Ok(Corpus { + block_size, + reference_window_ns, + trace_duplication_factor, + trace_length_factor, + inference_worker_duplication_factor, + logical_workers, + totals, + trace_path, + file_blake3, + hashes: hashes.into_boxed_slice(), + blocks: blocks.into_boxed_slice(), + removed: removed.into_boxed_slice(), + ops, + }) +} + +// --------------------------------------------------------------------------------------------- +// Backends +// --------------------------------------------------------------------------------------------- + +/// What a backend must offer the replay. `Lane` is the per-event-lane state (SMG keeps a block +/// map per worker in the lane that owns the worker, like the gateway's event monitor). Workers +/// are `(worker_id, dp_rank)` pairs as the engine streams key them; a backend interns them as it +/// likes. +trait ReplayBackend: Send + Sync + 'static { + type Lane: Send; + fn name(&self) -> &'static str; + fn new_lane(&self) -> Self::Lane; + /// A lane with its index and the CPUs it is pinned to (one CPU under `--pin-event-lanes`); + /// backends that place state by lane override it. + fn new_lane_for(&self, _lane: usize, _cpus: &[usize]) -> Self::Lane { + self.new_lane() + } + /// Figures to print and record at the end of a run, if the backend has any. + fn report(&self) -> Option { + None + } + /// Apply one stored event; `false` when the backend rejected it (counted). + fn apply_stored( + &self, + lane: &mut Self::Lane, + worker: (u64, u32), + blocks: &[StoredBlock], + parent: Option, + ) -> bool; + fn apply_removed( + &self, + lane: &mut Self::Lane, + worker: (u64, u32), + hashes: &[SequenceHash], + ) -> bool; + fn apply_cleared(&self, lane: &mut Self::Lane, worker: (u64, u32)) -> bool; + /// The shard an event lane's workers live in (its socket under `--shards` with pinned + /// lanes); the lane pools of `--lane-scheduling stealing` are one per shard, so a stolen + /// worker is always applied by a lane of its own shard. + fn lane_shard(&self, _lane: usize, _cpus: &[usize]) -> usize { + 0 + } + /// Answer one lookup; the return value is only consumed by `black_box`. + fn lookup(&self, hashes: &[ContentHash]) -> usize; + /// Shards a lookup can be fanned out over under `--lookups per-shard`; one when the backend + /// has no shards (the fan-out then is `all-shards`). + fn lookup_shards(&self) -> usize { + 1 + } + /// The CPUs of the event lanes placed on `shard`, when the backend places lanes by shard; + /// the query lanes of that shard float over them. + fn shard_cpus(&self, _shard: usize) -> Option> { + None + } + /// Answer one lookup over `shard` alone, appending its `(worker, score)` pairs to `out`: + /// the per-shard half of a fanned-out lookup. A backend with one shard never sees it. + fn lookup_shard(&self, _shard: usize, hashes: &[ContentHash], _out: &mut Vec<(u32, u32)>) { + black_box(self.lookup(hashes)); + } +} + +/// One stored block in the engine's wire layout: two hashes and an always-empty multimodal slot. +struct WireBlock { + block_hash: u64, + tokens_hash: u64, + /// Never set by the Mooncake trace; present so the record has the wire layout's size. + #[expect(dead_code)] + mm_extra_info: Option>, +} +const _: () = assert!(size_of::() == 40); + +/// The owned event payload an issuer moves to a lane under `--owned-payloads`; `None` when +/// the lane reads the event from the corpus slabs instead. +enum Payload { + None, + Stored(Vec), + Removed(Vec), +} + +struct EventMsg { + id: u32, + /// The worker's slot (its rank of first appearance in the corpus), the lane pool's key. + slot: u32, + payload: Payload, +} + +struct Positional { + inner: PositionalIndexer, +} + +struct PositionalLane { + workers: FxHashMap<(u64, u32), (u32, WorkerBlockMap)>, +} + +impl ReplayBackend for Positional { + type Lane = PositionalLane; + + fn name(&self) -> &'static str { + "smg-positional" + } + + fn new_lane(&self) -> Self::Lane { + PositionalLane { + workers: FxHashMap::default(), + } + } + + fn apply_stored( + &self, + lane: &mut Self::Lane, + worker: (u64, u32), + blocks: &[StoredBlock], + parent: Option, + ) -> bool { + let (smg_id, blocks_map) = lane.workers.entry(worker).or_insert_with(|| { + let id = self + .inner + .intern_worker(&format!("{}:{}", worker.0, worker.1)) + .expect("worker id space"); + (id, WorkerBlockMap::default()) + }); + self.inner + .apply_stored(*smg_id, blocks, parent, blocks_map) + .is_ok() + } + + fn apply_removed( + &self, + lane: &mut Self::Lane, + worker: (u64, u32), + hashes: &[SequenceHash], + ) -> bool { + let Some((smg_id, blocks_map)) = lane.workers.get_mut(&worker) else { + return false; + }; + self.inner.apply_removed(*smg_id, hashes, blocks_map); + true + } + + fn apply_cleared(&self, lane: &mut Self::Lane, worker: (u64, u32)) -> bool { + if let Some((smg_id, blocks_map)) = lane.workers.get_mut(&worker) { + self.inner.apply_cleared(*smg_id, blocks_map); + } + true + } + + fn lookup(&self, hashes: &[ContentHash]) -> usize { + self.inner.find_matches(hashes, false).scores.len() + } +} + +struct Chain { + inner: ShardedChainIndex, + /// The shard of every backend CPU (by position in the backend CPU list): the NUMA node's + /// index when the list spans exactly `shards` nodes, else contiguous groups of the list. + shard_of_cpu: FxHashMap, + event_lanes: usize, + /// Under `--count-shard-heads`: lookups by how many shards held the request's first block + /// (index = that count); a shared counter the query lanes write, so a diagnostic only. + heads_held: Option>, +} + +impl Chain { + fn shard_for(&self, lane: usize, cpus: &[usize]) -> usize { + let shards = self.inner.shards(); + if shards == 1 { + return 0; + } + match cpus { + [cpu] => self.shard_of_cpu.get(cpu).copied().unwrap_or(0), + // A floating lane has no socket: shards by lane index, exact but without the + // placement the split exists for. + _ => lane * shards / self.event_lanes.max(1), + } + } +} + +struct ChainLane { + shard: usize, + workers: FxHashMap<(u64, u32), (u32, ChainBlockMap)>, +} + +impl ReplayBackend for Chain { + type Lane = ChainLane; + + fn name(&self) -> &'static str { + "smg-chain" + } + + fn new_lane(&self) -> Self::Lane { + ChainLane { + shard: 0, + workers: FxHashMap::default(), + } + } + + fn new_lane_for(&self, lane: usize, cpus: &[usize]) -> Self::Lane { + ChainLane { + shard: self.shard_for(lane, cpus), + workers: FxHashMap::default(), + } + } + + fn lane_shard(&self, lane: usize, cpus: &[usize]) -> usize { + self.shard_for(lane, cpus) + } + + fn report(&self) -> Option { + let mut out = format!("shards = {}", self.inner.shards()); + for (shard, stats) in self.inner.shard_stats().iter().enumerate() { + out.push_str(&format!( + "\n shard {shard}: distinct blocks {} runs live {} blocks in live runs {} arena bytes {} (chunks {}, free-listed {}) slab bytes {} partial entries {} (max {} per run) child entries {} (tombstones {})", + self.inner.shard(shard).entry_count(), + stats.runs_live, + stats.blocks_live, + stats.arena_bytes, + stats.arena_chunk_bytes, + stats.arena_free_bytes, + stats.slab_bytes, + stats.partial_entries, + stats.max_partials, + stats.child_entries, + stats.child_tombstones + )); + } + let total = self.inner.stats(); + out.push_str(&format!( + "\n total: memberships {} distinct blocks summed over shards {} arena bytes {} chunks {} slab bytes {}", + self.inner.current_size(), + self.inner.entry_count(), + total.arena_bytes, + total.arena_chunk_bytes, + total.slab_bytes + )); + if let Some(counts) = &self.heads_held { + let counts: Vec = counts.iter().map(|c| c.load(Ordering::Relaxed)).collect(); + out.push_str(&format!( + "\n lookups by shards holding the first block (0, 1, 2, ..): {counts:?}" + )); + } + Some(out) + } + + fn apply_stored( + &self, + lane: &mut Self::Lane, + worker: (u64, u32), + blocks: &[StoredBlock], + parent: Option, + ) -> bool { + let (smg_id, blocks_map) = lane.workers.entry(worker).or_insert_with(|| { + let id = self + .inner + .intern_worker_in(lane.shard, &format!("{}:{}", worker.0, worker.1)) + .expect("worker slots; raise --max-workers"); + (id, ChainBlockMap::default()) + }); + self.inner + .apply_stored(*smg_id, blocks, parent, blocks_map) + .is_ok() + } + + fn apply_removed( + &self, + lane: &mut Self::Lane, + worker: (u64, u32), + hashes: &[SequenceHash], + ) -> bool { + let Some((smg_id, blocks_map)) = lane.workers.get_mut(&worker) else { + return false; + }; + self.inner.apply_removed(*smg_id, hashes, blocks_map); + true + } + + fn apply_cleared(&self, lane: &mut Self::Lane, worker: (u64, u32)) -> bool { + if let Some((smg_id, blocks_map)) = lane.workers.get_mut(&worker) { + self.inner.apply_cleared(*smg_id, blocks_map); + } + true + } + + fn lookup(&self, hashes: &[ContentHash]) -> usize { + if let (Some(counts), Some(first)) = (&self.heads_held, hashes.first()) { + let held = self.inner.shards_holding_head(first.0); + counts[held.min(counts.len() - 1)].fetch_add(1, Ordering::Relaxed); + } + let mut scored = 0usize; + self.inner + .score_into(hashes, |content| content.0, false, |_, _| scored += 1); + scored + } + + fn lookup_shards(&self) -> usize { + self.inner.shards() + } + + fn shard_cpus(&self, shard: usize) -> Option> { + let mut cpus: Vec = self + .shard_of_cpu + .iter() + .filter(|(_, &s)| s == shard) + .map(|(&cpu, _)| cpu) + .collect(); + cpus.sort_unstable(); + (!cpus.is_empty()).then_some(cpus) + } + + fn lookup_shard(&self, shard: usize, hashes: &[ContentHash], out: &mut Vec<(u32, u32)>) { + self.inner.score_shard_into( + shard, + hashes, + |content| content.0, + false, + |worker, score| out.push((worker, score)), + ); + } +} + +struct Reference { + inner: Mutex, + ids: Mutex>, +} + +impl Reference { + fn id(&self, key: (u64, u32)) -> u32 { + let mut ids = self.ids.lock().unwrap_or_else(|e| e.into_inner()); + let next = ids.len() as u32; + *ids.entry(key).or_insert(next) + } +} + +impl ReplayBackend for Reference { + type Lane = (); + + fn name(&self) -> &'static str { + "smg-reference" + } + + fn new_lane(&self) -> Self::Lane {} + + fn apply_stored( + &self, + _lane: &mut Self::Lane, + worker: (u64, u32), + blocks: &[StoredBlock], + parent: Option, + ) -> bool { + let worker = self.id(worker); + self.inner + .lock() + .unwrap_or_else(|e| e.into_inner()) + .apply_stored(worker, blocks, parent) + .is_ok() + } + + fn apply_removed( + &self, + _lane: &mut Self::Lane, + worker: (u64, u32), + hashes: &[SequenceHash], + ) -> bool { + let worker = self.id(worker); + self.inner + .lock() + .unwrap_or_else(|e| e.into_inner()) + .apply_removed(worker, hashes); + true + } + + fn apply_cleared(&self, _lane: &mut Self::Lane, worker: (u64, u32)) -> bool { + let worker = self.id(worker); + self.inner + .lock() + .unwrap_or_else(|e| e.into_inner()) + .apply_cleared(worker); + true + } + + fn lookup(&self, hashes: &[ContentHash]) -> usize { + self.inner + .lock() + .unwrap_or_else(|e| e.into_inner()) + .find_matches(hashes) + .len() + } +} + +/// Discards events and answers lookups with the sequence length: measures the harness alone +/// (issuers, queues, lanes, timing) with no indexer behind it, which bounds what any backend can +/// be measured at on a given layout. +struct Null; + +impl ReplayBackend for Null { + type Lane = (); + + fn name(&self) -> &'static str { + "null" + } + + fn new_lane(&self) -> Self::Lane {} + + fn apply_stored( + &self, + _lane: &mut Self::Lane, + _worker: (u64, u32), + _blocks: &[StoredBlock], + _parent: Option, + ) -> bool { + true + } + + fn apply_removed( + &self, + _lane: &mut Self::Lane, + _worker: (u64, u32), + _hashes: &[SequenceHash], + ) -> bool { + true + } + + fn apply_cleared(&self, _lane: &mut Self::Lane, _worker: (u64, u32)) -> bool { + true + } + + fn lookup(&self, hashes: &[ContentHash]) -> usize { + hashes.len() + } +} + +// --------------------------------------------------------------------------------------------- +// Lanes +// --------------------------------------------------------------------------------------------- + +struct QueryLane { + slots: Box<[AtomicU32]>, + published: AtomicUsize, + closed: AtomicBool, + consumer: Mutex>, +} + +impl QueryLane { + fn new(capacity: usize) -> Self { + Self { + slots: (0..capacity) + .map(|_| AtomicU32::new(EMPTY_OPERATION_ID)) + .collect::>() + .into_boxed_slice(), + published: AtomicUsize::new(0), + closed: AtomicBool::new(false), + consumer: Mutex::new(None), + } + } + fn wake(&self) { + if let Some(consumer) = self + .consumer + .lock() + .unwrap_or_else(|e| e.into_inner()) + .as_ref() + { + consumer.unpark(); + } + } + fn publish(&self, count: usize) { + self.published.store(count, Ordering::Release); + self.wake(); + } + fn close(&self) { + self.closed.store(true, Ordering::Release); + self.wake(); + } +} + +#[derive(Clone, Copy, Default)] +struct QueryCompletion { + id: u32, + started_ns: u64, + finished_ns: u64, +} + +/// The hand-over point of one fanned-out lookup (`--lookups per-shard`): the member lanes of a +/// lookup group share one slot per published position, since every member receives the same +/// ids in the same order. A lane that finishes before the others leaves its partial answer and +/// its start time here; the last one merges and records the completion. +#[derive(Default)] +struct FanSlot { + parked: Mutex>, +} + +struct FanPartial { + /// The earliest start among the lanes finished so far. + started_ns: u64, + /// Their partial answers, concatenated (worker sets are disjoint across shards). + scores: Vec<(u32, u32)>, + /// How many lanes have finished. + finished: usize, +} + +/// What a query lane walks: every shard, or its own shard of a fanned-out lookup. +enum LaneLookup { + AllShards, + Shard { + shard: usize, + fan: usize, + slots: Arc<[FanSlot]>, + }, +} + +/// Per-lane counters of the fan-out, summed per shard at the end (nothing shared on the lookup +/// path). +#[derive(Clone, Copy, Debug, Default)] +struct FanStats { + /// Lookups this lane walked its shard for. + lookups: u64, + /// Of those, lookups whose shard held something of the request. + with_holders: u64, + /// Holders this lane's shard reported, summed over its lookups. + holders: u64, + /// Lookups this lane completed: it finished last and merged. + merged: u64, +} + +impl FanStats { + fn add(&mut self, other: &FanStats) { + self.lookups += other.lookups; + self.with_holders += other.with_holders; + self.holders += other.holders; + self.merged += other.merged; + } +} + +struct QueryLaneOutput { + completions: Vec, + failure: Option<&'static str>, + cpu_ns: u64, + fan: FanStats, +} + +fn query_lane_worker( + backend: Arc, + lane: Arc, + corpus: Arc, + epoch: Instant, + cpus: Arc<[usize]>, + mirror: bool, + lookup: LaneLookup, +) -> QueryLaneOutput { + let _ = pin_current_thread(&cpus); + let cpu_started = thread_cpu_time_ns(); + *lane.consumer.lock().unwrap_or_else(|e| e.into_inner()) = Some(thread::current()); + let mut completions = Vec::with_capacity(lane.slots.len()); + let mut consumed = 0usize; + let mut failure = None; + let mut partial: Vec<(u32, u32)> = Vec::new(); + let mut stats = FanStats::default(); + loop { + let published = lane.published.load(Ordering::Acquire); + while consumed < published { + let id = lane.slots[consumed].load(Ordering::Relaxed); + if id == EMPTY_OPERATION_ID { + failure = Some("query_lane_missing_published_id"); + break; + } + let OpKind::Query { start, len } = corpus.ops[id as usize].kind else { + failure = Some("query_lane_non_query_id"); + break; + }; + let hashes = &corpus.hashes[start as usize..start as usize + len as usize]; + let started_ns = elapsed_ns(epoch); + match &lookup { + LaneLookup::AllShards => { + if mirror { + // The adapter's copy into this crate's hash type, inside the timed lookup. + let owned: Vec = hashes.to_vec(); + black_box(backend.lookup(&owned)); + } else { + black_box(backend.lookup(hashes)); + } + let finished_ns = elapsed_ns(epoch); + completions.push(QueryCompletion { + id, + started_ns, + finished_ns, + }); + } + LaneLookup::Shard { shard, fan, slots } => { + partial.clear(); + if mirror { + let owned: Vec = hashes.to_vec(); + backend.lookup_shard(*shard, &owned, &mut partial); + } else { + backend.lookup_shard(*shard, hashes, &mut partial); + } + stats.lookups += 1; + stats.with_holders += u64::from(!partial.is_empty()); + stats.holders += partial.len() as u64; + let Some(slot) = slots.get(consumed) else { + failure = Some("query_lane_fan_slot_overflow"); + break; + }; + let mut parked = slot.parked.lock().unwrap_or_else(|e| e.into_inner()); + match parked.as_mut() { + Some(earlier) if earlier.finished + 1 == *fan => { + // Last to finish: merge the answers and complete the lookup. + let mut merged = std::mem::take(&mut earlier.scores); + let started_ns = started_ns.min(earlier.started_ns); + *parked = None; + drop(parked); + merged.extend_from_slice(&partial); + black_box(&merged); + let finished_ns = elapsed_ns(epoch); + completions.push(QueryCompletion { + id, + started_ns, + finished_ns, + }); + stats.merged += 1; + } + Some(earlier) => { + earlier.finished += 1; + earlier.started_ns = earlier.started_ns.min(started_ns); + earlier.scores.extend_from_slice(&partial); + } + None => { + // First to finish: hand the partial answer over (one small vector + // allocated on this lane's socket, freed by the merging lane). + *parked = Some(FanPartial { + started_ns, + scores: partial.clone(), + finished: 1, + }); + } + } + } + } + consumed += 1; + } + if failure.is_some() { + break; + } + if lane.closed.load(Ordering::Acquire) && consumed == lane.published.load(Ordering::Acquire) + { + break; + } + if consumed == lane.published.load(Ordering::Acquire) + && !lane.closed.load(Ordering::Acquire) + { + thread::park(); + } + } + QueryLaneOutput { + completions, + failure, + cpu_ns: thread_cpu_time_ns().saturating_sub(cpu_started), + fan: stats, + } +} + +#[derive(Clone, Copy, Default)] +struct EventCompletion { + id: u32, + finished_ns: u64, + ok: bool, +} + +/// Where an event lane runs: its index, the CPUs it is pinned to (one under `--pin-event-lanes`) +/// and whether it prefers its own NUMA node for memory. +struct LanePlacement { + index: usize, + cpus: Arc<[usize]>, + local_memory: bool, +} + +/// What an event lane thread returns: its completions, its thread CPU time, the part of it +/// spent inside the apply calls (the rest is the lane's own loop, or the pool's overhead under +/// `--lane-scheduling stealing`), and the minor page faults it took during the run. +struct LaneOutcome { + completions: Vec, + cpu_ns: u64, + apply_ns: u64, + minor_faults: u64, +} + +fn event_lane_worker( + backend: Arc, + receiver: mpsc::Receiver, + corpus: Arc, + epoch: Instant, + placement: LanePlacement, + expected: usize, +) -> LaneOutcome { + place_lane(&placement); + let LanePlacement { + index: lane_index, + cpus, + .. + } = placement; + let cpu_started = thread_cpu_time_ns(); + let faults_started = thread_minor_faults(); + let mut lane = backend.new_lane_for(lane_index, &cpus); + let mut completions = Vec::with_capacity(expected); + let mut apply_ns = 0u64; + while let Ok(EventMsg { id, payload, .. }) = receiver.recv() { + let started = Instant::now(); + let ok = apply_event(&*backend, &mut lane, &corpus, id, payload); + apply_ns += started.elapsed().as_nanos() as u64; + completions.push(EventCompletion { + id, + finished_ns: elapsed_ns(epoch), + ok, + }); + } + LaneOutcome { + completions, + cpu_ns: thread_cpu_time_ns().saturating_sub(cpu_started), + apply_ns, + minor_faults: thread_minor_faults().saturating_sub(faults_started), + } +} + +/// Pin an event lane's thread and set its memory policy as its placement says. +fn place_lane(placement: &LanePlacement) { + let LanePlacement { + index: lane_index, + cpus, + local_memory, + } = placement; + let _ = pin_current_thread(cpus); + if *local_memory { + match cpus.first().and_then(|&cpu| node_of_cpu(cpu)) { + Some(node) => { + if let Err(err) = prefer_node(node) { + println!( + "event lane {lane_index}: set_mempolicy for node {node} failed: {err}" + ); + } + } + None => println!( + "event lane {lane_index}: no NUMA node for CPUs {cpus:?}, memory policy inherited" + ), + } + } +} + +/// Apply one event to the backend through a lane's (or a worker's) state; `false` when the +/// backend rejected it. +fn apply_event( + backend: &B, + lane: &mut B::Lane, + corpus: &Corpus, + id: u32, + payload: Payload, +) -> bool { + { + let op = &corpus.ops[id as usize]; + let worker = (op.worker, op.dp_rank()); + match (&op.kind, payload) { + (&OpKind::Stored { parent, .. }, Payload::Stored(wire)) => { + // The adapter's copy, wire records into this crate's blocks; the payload is + // freed after the apply, as a lane frees the event it was handed. + let owned: Vec = wire + .iter() + .map(|block| StoredBlock { + seq_hash: SequenceHash(block.block_hash), + content_hash: ContentHash(block.tokens_hash), + }) + .collect(); + let ok = backend.apply_stored(lane, worker, &owned, parent.map(SequenceHash)); + drop(wire); + ok + } + ( + &OpKind::Stored { + parent, start, len, .. + }, + Payload::None, + ) => backend.apply_stored( + lane, + worker, + &corpus.blocks[start as usize..start as usize + len as usize], + parent.map(SequenceHash), + ), + (&OpKind::Removed { .. }, Payload::Removed(wire)) => { + let owned: Vec = + wire.iter().map(|&hash| SequenceHash(hash)).collect(); + let ok = backend.apply_removed(lane, worker, &owned); + drop(wire); + ok + } + (&OpKind::Removed { start, len, .. }, Payload::None) => backend.apply_removed( + lane, + worker, + &corpus.removed[start as usize..start as usize + len as usize], + ), + (&OpKind::Cleared { .. }, _) => backend.apply_cleared(lane, worker), + _ => false, + } + } +} + +/// How long a pooled lane blocks on its channel when nothing is ready anywhere in its pool. +const POOL_WAIT_STEALABLE: Duration = Duration::from_micros(100); +const POOL_WAIT_QUIET: Duration = Duration::from_millis(1); + +/// An event lane under `--lane-scheduling stealing`: a lane of its shard's `LanePool`. It drains +/// its own channel into the pool's per-worker queues (its ingress is unchanged: the issuer still +/// sends a worker's events to the worker's home lane) and serves ready workers from any lane of +/// the pool, whole workers at a time, so a worker's events keep their order while a quiet lane +/// works off a stalled lane's backlog. The per-worker state is a backend lane state holding that +/// one worker, created by whichever lane of the pool first serves it, which under pinned lanes is +/// on the worker's own shard. +struct PoolLane<'a, B: ReplayBackend> { + backend: &'a B, + corpus: &'a Corpus, + receiver: mpsc::Receiver, + pool: &'a LanePool, + /// This lane's index within its pool. + lane: usize, + /// This lane's index among all event lanes (what `new_lane_for` places by). + lane_index: usize, + cpus: Arc<[usize]>, + epoch: Instant, + /// Lanes of this pool whose channel is still open; the pool is done when it is zero and + /// nothing is queued or held. + open_channels: &'a AtomicUsize, + held_total: &'a AtomicUsize, + closed: bool, + /// An event the pool refused (the worker's queue at its cap), re-offered before reading on. + held: Option, + completions: Vec, + enqueue_ns: u64, + stealable_idle_turns: u64, + apply_ns: u64, +} + +impl PoolLane<'_, B> { + fn offer(&mut self, msg: EventMsg) { + let started = Instant::now(); + let refused = self.pool.enqueue(self.lane, msg.slot, msg); + self.enqueue_ns += started.elapsed().as_nanos() as u64; + if let Err(QueueFull(msg)) = refused { + self.held = Some(msg); + self.held_total.fetch_add(1, Ordering::AcqRel); + } + } + + fn close(&mut self) { + if !self.closed { + self.closed = true; + self.open_channels.fetch_sub(1, Ordering::AcqRel); + } + } + + fn control(&self) -> Control { + if self.open_channels.load(Ordering::Acquire) == 0 + && self.held_total.load(Ordering::Acquire) == 0 + && self.pool.queued() == 0 + { + Control::Stop + } else { + Control::Continue + } + } +} + +impl LaneHooks for PoolLane<'_, B> { + fn apply(&mut self, claimed: Claimed<'_, B::Lane>, msg: EventMsg) { + let (backend, lane_index, cpus) = (self.backend, self.lane_index, &self.cpus); + let state = claimed + .state + .get_or_insert_with(|| backend.new_lane_for(lane_index, cpus)); + let started = Instant::now(); + let ok = apply_event(backend, state, self.corpus, msg.id, msg.payload); + self.apply_ns += started.elapsed().as_nanos() as u64; + self.completions.push(EventCompletion { + id: msg.id, + finished_ns: elapsed_ns(self.epoch), + ok, + }); + } + + fn pump(&mut self) -> Control { + if let Some(msg) = self.held.take() { + self.held_total.fetch_sub(1, Ordering::AcqRel); + self.offer(msg); + if self.held.is_some() { + return Control::Continue; + } + } + if self.closed { + return self.control(); + } + loop { + match self.receiver.try_recv() { + Ok(msg) => { + self.offer(msg); + if self.held.is_some() { + return Control::Continue; + } + } + Err(mpsc::TryRecvError::Empty) => return Control::Continue, + Err(mpsc::TryRecvError::Disconnected) => { + self.close(); + return self.control(); + } + } + } + } + + fn wait(&mut self) -> Control { + let stealable = self.pool.has_stealable(); + if stealable { + // Nothing was ready or stealable for this lane this turn, yet some lane is flagged. + self.stealable_idle_turns += 1; + } + let timeout = if stealable { + POOL_WAIT_STEALABLE + } else { + POOL_WAIT_QUIET + }; + if self.closed || self.held.is_some() { + self.pool.park_lane(self.lane, timeout); + return self.control(); + } + match self.receiver.recv_timeout(timeout) { + Ok(msg) => { + self.offer(msg); + Control::Continue + } + Err(mpsc::RecvTimeoutError::Timeout) => Control::Continue, + Err(mpsc::RecvTimeoutError::Disconnected) => { + self.close(); + self.control() + } + } + } +} + +/// The pool of one shard's event lanes and the counters its lanes share. +struct ShardPool { + pool: LanePool, + open_channels: AtomicUsize, + held_total: AtomicUsize, + /// Time the lanes spent inside `enqueue` (the worker queue's lock and the ready list's), + /// summed: a producer-side convoy on a stolen worker's queue shows up here. + enqueue_ns: AtomicU64, + /// Turns on which a lane found nothing to serve or steal while the pool still reported a + /// stealable lane: a stealable flag left set on a lane that is not running shows up here. + stealable_idle_turns: AtomicU64, +} + +#[expect(clippy::too_many_arguments)] +fn event_lane_worker_pooled( + backend: Arc, + receiver: mpsc::Receiver, + corpus: Arc, + epoch: Instant, + placement: LanePlacement, + expected: usize, + shard_pool: Arc>, + pool_lane: usize, +) -> LaneOutcome { + place_lane(&placement); + let LanePlacement { + index: lane_index, + cpus, + .. + } = placement; + let cpu_started = thread_cpu_time_ns(); + let faults_started = thread_minor_faults(); + let mut hooks = PoolLane { + backend: &*backend, + corpus: &corpus, + receiver, + pool: &shard_pool.pool, + lane: pool_lane, + lane_index, + cpus, + epoch, + open_channels: &shard_pool.open_channels, + held_total: &shard_pool.held_total, + closed: false, + held: None, + completions: Vec::with_capacity(expected), + enqueue_ns: 0, + stealable_idle_turns: 0, + apply_ns: 0, + }; + shard_pool.pool.run_lane(pool_lane, &mut hooks); + shard_pool + .enqueue_ns + .fetch_add(hooks.enqueue_ns, Ordering::Relaxed); + shard_pool + .stealable_idle_turns + .fetch_add(hooks.stealable_idle_turns, Ordering::Relaxed); + LaneOutcome { + completions: hooks.completions, + cpu_ns: thread_cpu_time_ns().saturating_sub(cpu_started), + apply_ns: hooks.apply_ns, + minor_faults: thread_minor_faults().saturating_sub(faults_started), + } +} + +// --------------------------------------------------------------------------------------------- +// Issuers +// --------------------------------------------------------------------------------------------- + +#[derive(Clone, Copy, Default)] +struct IssueRecord { + scheduled_ns: u64, + accepted_ns: u64, + is_query: bool, + accepted: bool, +} + +struct Clock { + epoch: Instant, + monotonic_epoch_ns: u64, + spin_ns: u64, +} + +impl Clock { + fn new(spin_ns: u64) -> anyhow::Result { + Ok(Self { + epoch: Instant::now(), + monotonic_epoch_ns: monotonic_now_ns()?, + spin_ns, + }) + } + fn now_ns(&self) -> u64 { + elapsed_ns(self.epoch) + } + fn wait_until(&self, target_ns: u64) { + let sleep_target = target_ns.saturating_sub(self.spin_ns); + if sleep_target > self.now_ns() { + sleep_until_monotonic(self.monotonic_epoch_ns.saturating_add(sleep_target)); + } + while self.now_ns() < target_ns { + std::hint::spin_loop(); + } + } +} + +fn elapsed_ns(epoch: Instant) -> u64 { + epoch.elapsed().as_nanos().min(u64::MAX as u128) as u64 +} + +fn monotonic_now_ns() -> anyhow::Result { + let mut ts = libc::timespec { + tv_sec: 0, + tv_nsec: 0, + }; + // SAFETY: clock_gettime writes a timespec it is given a valid pointer to. + let rc = unsafe { libc::clock_gettime(libc::CLOCK_MONOTONIC, &mut ts) }; + anyhow::ensure!(rc == 0, "clock_gettime failed"); + Ok((ts.tv_sec as u64) + .saturating_mul(1_000_000_000) + .saturating_add(ts.tv_nsec as u64)) +} + +fn sleep_until_monotonic(target_ns: u64) { + let request = libc::timespec { + tv_sec: (target_ns / 1_000_000_000) as libc::time_t, + tv_nsec: (target_ns % 1_000_000_000) as libc::c_long, + }; + loop { + // SAFETY: an absolute sleep on CLOCK_MONOTONIC with a valid timespec and no remainder. + let rc = unsafe { + libc::clock_nanosleep( + libc::CLOCK_MONOTONIC, + libc::TIMER_ABSTIME, + &request, + std::ptr::null_mut(), + ) + }; + if rc != libc::EINTR { + return; + } + } +} + +fn pin_current_thread(cpus: &[usize]) -> std::io::Result<()> { + if cpus.is_empty() { + return Ok(()); + } + // SAFETY: a zeroed cpu_set_t is a valid empty set; CPU_SET/CPU_ZERO only touch that set; + // sched_setaffinity(0) applies it to the calling thread. + unsafe { + let mut set = std::mem::zeroed::(); + libc::CPU_ZERO(&mut set); + for &cpu in cpus { + libc::CPU_SET(cpu, &mut set); + } + let rc = libc::sched_setaffinity(0, size_of::(), &set); + if rc != 0 { + return Err(std::io::Error::last_os_error()); + } + } + Ok(()) +} + +/// The NUMA node a CPU belongs to, from `/sys/devices/system/node/node*/cpulist`. +fn node_of_cpu(cpu: usize) -> Option { + let nodes = std::fs::read_dir("/sys/devices/system/node").ok()?; + for entry in nodes.flatten() { + let name = entry.file_name(); + let Some(number) = name.to_str().and_then(|n| n.strip_prefix("node")) else { + continue; + }; + let Ok(node) = number.parse::() else { + continue; + }; + let Ok(list) = std::fs::read_to_string(entry.path().join("cpulist")) else { + continue; + }; + if parse_cpu_list(list.trim()).is_ok_and(|cpus| cpus.contains(&cpu)) { + return Some(node); + } + } + None +} + +/// Prefer `node` for this thread's page allocations from now on (`set_mempolicy`, +/// `MPOL_PREFERRED`): the arena chunks, run slab chunks and lane maps a lane first touches land +/// on its own socket whatever the process policy (an interleaving policy set by the launcher, +/// for instance). +fn prefer_node(node: usize) -> std::io::Result<()> { + let mut mask = [0u64; 16]; + mask[node / 64] |= 1u64 << (node % 64); + // SAFETY: set_mempolicy reads `maxnode` bits from `mask`, which holds 16 * 64 of them, and + // changes only the calling thread's allocation policy. + let rc = unsafe { + libc::syscall( + libc::SYS_set_mempolicy, + libc::MPOL_PREFERRED as libc::c_long, + mask.as_ptr(), + (mask.len() * 64) as libc::c_ulong, + ) + }; + if rc != 0 { + return Err(std::io::Error::last_os_error()); + } + Ok(()) +} + +/// Minor page faults taken by the calling thread so far (`getrusage(RUSAGE_THREAD)`), so a lane's +/// record can say whether its time went into faulting in memory; 0 when the call fails. +fn thread_minor_faults() -> u64 { + // SAFETY: rusage is plain data and getrusage writes the whole struct on success. + let mut usage: libc::rusage = unsafe { std::mem::zeroed() }; + let rc = unsafe { libc::getrusage(libc::RUSAGE_THREAD, &mut usage) }; + if rc != 0 { + return 0; + } + u64::try_from(usage.ru_minflt).unwrap_or(0) +} + +fn thread_cpu_time_ns() -> u64 { + let mut ts = libc::timespec { + tv_sec: 0, + tv_nsec: 0, + }; + // SAFETY: as in monotonic_now_ns. + let rc = unsafe { libc::clock_gettime(libc::CLOCK_THREAD_CPUTIME_ID, &mut ts) }; + if rc != 0 { + return 0; + } + (ts.tv_sec as u64) + .saturating_mul(1_000_000_000) + .saturating_add(ts.tv_nsec as u64) +} + +fn parse_cpu_list(value: &str) -> anyhow::Result> { + let mut cpus = Vec::new(); + for part in value.split(',').map(str::trim).filter(|p| !p.is_empty()) { + if let Some((a, b)) = part.split_once('-') { + let (a, b): (usize, usize) = (a.parse()?, b.parse()?); + anyhow::ensure!(a <= b, "descending CPU range {part}"); + cpus.extend(a..=b); + } else { + cpus.push(part.parse()?); + } + } + Ok(cpus) +} + +/// Contiguous worker ranges per event issuer. +fn event_issuer_for(worker: u64, logical_workers: u64, issuers: usize) -> usize { + let per = (logical_workers as usize).div_ceil(issuers).max(1); + ((worker as usize) / per).min(issuers - 1) +} + +struct Shared { + clock: Clock, + start_ns: AtomicUsize, + /// Queries of each deadline group not yet published; its events wait for zero. + deadline_pending: Box<[AtomicU32]>, + peer_failed: AtomicBool, +} + +fn issue_queries( + shared: &Shared, + corpus: &Corpus, + dispatch: &[(u32, u32, u16)], // (id, deadline_group, lookup group) + lanes: &[Arc], + fan: usize, // lanes per lookup group: the lookup's shards under `--lookups per-shard` + records: &mut Vec<(u32, IssueRecord)>, +) -> (u64, Option<&'static str>) { + let cpu_started = thread_cpu_time_ns(); + let start_ns = shared.start_ns.load(Ordering::Acquire) as u64; + let mut cursors = vec![0usize; lanes.len()]; + let mut touched = Vec::with_capacity(lanes.len()); + let mut touched_flags = vec![false; lanes.len()]; + let mut i = 0usize; + let mut failure = None; + while i < dispatch.len() && failure.is_none() { + let deadline_ns = corpus.ops[dispatch[i].0 as usize].deadline_ns; + let group = dispatch[i].1 as usize; + shared + .clock + .wait_until(start_ns.saturating_add(deadline_ns)); + touched.clear(); + let group_start = i; + while i < dispatch.len() && corpus.ops[dispatch[i].0 as usize].deadline_ns == deadline_ns { + let (id, _, target) = dispatch[i]; + let first_lane = target as usize * fan; + for lane in first_lane..first_lane + fan { + let Some(slot) = lanes[lane].slots.get(cursors[lane]) else { + failure = Some("issuer_query_lane_overflow"); + break; + }; + slot.store(id, Ordering::Relaxed); + cursors[lane] += 1; + if !touched_flags[lane] { + touched_flags[lane] = true; + touched.push(lane); + } + } + if failure.is_some() { + break; + } + records.push(( + id, + IssueRecord { + scheduled_ns: start_ns.saturating_add(deadline_ns), + accepted_ns: shared.clock.now_ns(), + is_query: true, + accepted: true, + }, + )); + i += 1; + } + for &lane in &touched { + lanes[lane].publish(cursors[lane]); + touched_flags[lane] = false; + } + if failure.is_some() { + break; + } + shared.deadline_pending[group].fetch_sub((i - group_start) as u32, Ordering::Release); + } + if failure.is_some() { + shared.peer_failed.store(true, Ordering::Release); + } + (thread_cpu_time_ns().saturating_sub(cpu_started), failure) +} + +/// The owned payload of one event in the engine's wire layout, as its lane receives it under +/// `--owned-payloads`. +fn payload_for(corpus: &Corpus, op: &Op) -> Payload { + match op.kind { + OpKind::Stored { start, len, .. } => Payload::Stored( + corpus.blocks[start as usize..start as usize + len as usize] + .iter() + .map(|block| WireBlock { + block_hash: block.seq_hash.0, + tokens_hash: block.content_hash.0, + mm_extra_info: None, + }) + .collect(), + ), + OpKind::Removed { start, len, .. } => Payload::Removed( + corpus.removed[start as usize..start as usize + len as usize] + .iter() + .map(|hash| hash.0) + .collect(), + ), + _ => Payload::None, + } +} + +/// One event in an issuer's dispatch: its operation, deadline group, lane and (under +/// `--owned-payloads`) the owned payload the lane receives. +struct EventDispatch { + id: u32, + group: u32, + lane: u16, + slot: u32, + payload: Payload, +} + +fn issue_events( + shared: &Shared, + corpus: &Corpus, + dispatch: Vec, + senders: &[mpsc::Sender], + records: &mut Vec<(u32, IssueRecord)>, +) -> (u64, Option<&'static str>) { + let cpu_started = thread_cpu_time_ns(); + let start_ns = shared.start_ns.load(Ordering::Acquire) as u64; + let mut entries = dispatch.into_iter().peekable(); + let mut failure = None; + while let Some(head) = entries.peek() { + let deadline_ns = corpus.ops[head.id as usize].deadline_ns; + let group = head.group as usize; + shared + .clock + .wait_until(start_ns.saturating_add(deadline_ns)); + // Queries of this deadline are published before its events. + while shared.deadline_pending[group].load(Ordering::Acquire) != 0 { + if shared.peer_failed.load(Ordering::Acquire) { + failure = Some("issuer_peer_failed"); + break; + } + std::hint::spin_loop(); + } + if failure.is_some() { + break; + } + while entries + .peek() + .is_some_and(|entry| corpus.ops[entry.id as usize].deadline_ns == deadline_ns) + { + let Some(entry) = entries.next() else { + break; + }; + let message = EventMsg { + id: entry.id, + slot: entry.slot, + payload: entry.payload, + }; + if senders[entry.lane as usize].send(message).is_err() { + failure = Some("issuer_event_lane_offline"); + break; + } + records.push(( + entry.id, + IssueRecord { + scheduled_ns: start_ns.saturating_add(deadline_ns), + accepted_ns: shared.clock.now_ns(), + is_query: false, + accepted: true, + }, + )); + } + if failure.is_some() { + break; + } + } + if failure.is_some() { + shared.peer_failed.store(true, Ordering::Release); + } + (thread_cpu_time_ns().saturating_sub(cpu_started), failure) +} + +// --------------------------------------------------------------------------------------------- +// Chain +// --------------------------------------------------------------------------------------------- + +#[derive(Clone, Copy, Debug, PartialEq, Eq, ValueEnum)] +enum LaneMemory { + /// The process policy, whatever the launcher set. + Inherit, + /// Each pinned event lane prefers its own NUMA node (`set_mempolicy(MPOL_PREFERRED)`). + Local, +} + +/// Whether the corpus's lookups are issued. +#[derive(Clone, Copy, Debug, PartialEq, Eq, ValueEnum)] +enum Queries { + /// Queries and events, the corpus as recorded. + On, + /// Events only: queries are dropped at dispatch, the deadline schedule is kept and the query + /// lanes stay idle; offered, issued and achieved rates and the kept-up verdict count event + /// block ops alone, so the row is a lane-cost discriminator and not comparable with mixed + /// rows (the JSON says `"queries": "off"` and `total_requests` is 0). + Off, +} + +/// Which thread builds the owned event payloads of `--owned-payloads`, and so which +/// allocator arena and NUMA node they live on. +#[derive(Clone, Copy, Debug, PartialEq, Eq, ValueEnum)] +enum LaneScheduling { + /// Every event lane applies its own workers' events, in order, and nothing else (today). + Owned, + /// Each shard's event lanes form one lane pool: an idle lane serves any ready worker of + /// its shard, and takes one from a lane that is not running once `--steal-after` events + /// are queued behind it. + Stealing, +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq, ValueEnum)] +enum PayloadHome { + /// The main thread, before the trial, as a preparation step does: one arena for every + /// payload, which every lane frees into. On two sockets that one arena's lock and free lists + /// bounce between the sockets on every event, and the lanes' CPU per block doubles. + Main, + /// Each event issuer, in its pinned thread, for its own dispatch: payloads live on the + /// issuer's socket and the lanes free into the issuer's arena (under `--issuer-by-lane` the + /// same socket), so a two-socket row measures the index and not the allocator. + Issuer, +} + +/// How a lookup is spread over the index's shards. +#[derive(Clone, Copy, Debug, PartialEq, Eq, ValueEnum)] +enum Lookups { + /// One query lane walks every shard (the lane floats over all lane cores). + AllShards, + /// A lookup fans out to one query lane per shard, each floating over its shard's cores and + /// walking only its shard; the lane finishing last merges the partial answers. Worker sets + /// are disjoint across shards, so the merge is a concatenation and the answer is + /// `all-shards`'s. The remote half of a lookup becomes one hand-over of a small vector + /// instead of remote reads of the lines the local event lanes write. One shard: as + /// `all-shards`. + PerShard, +} + +#[derive(Clone, Copy, Debug, ValueEnum)] +enum BackendKind { + /// This crate's event-driven PositionalIndexer. + Positional, + /// The single-threaded reference indexer (small corpora only). + Reference, + /// This crate's run-compressed ChainIndex (worker slots from `--max-workers`); `run` is the + /// name it had and stays accepted until the positional indexer's removal. + #[value(alias = "run")] + Chain, + /// No indexer: the harness's own ceiling on this layout. + Null, +} + +#[derive(Parser, Debug)] +#[command( + about = "Open-loop replay of an exported Mooncake indexer corpus with drain-inclusive accounting" +)] +struct Args { + /// Corpus file in the export format described in benches/README.md: a prepared schedule of + /// the Mooncake trace. + corpus: String, + #[arg(long, value_enum, default_value = "positional")] + backend: BackendKind, + /// Jump size of the positional indexer's lookup. + #[arg(long, default_value = "8")] + jump_size: usize, + /// Worker slots of the chain index (one coverage bit per slot per run, at most 1024). + #[arg(long, default_value = "256")] + max_workers: usize, + /// Chain index shards, one per socket of the lane set: each shard is a complete chain index + /// written only by the lanes placed on it (a pinned lane goes to its CPU's NUMA node's + /// shard when the lane set spans exactly this many nodes, else to its contiguous group of + /// the backend CPU list), and a lookup walks every shard. One shard is the chain index itself. + #[arg(long, default_value = "1")] + shards: usize, + /// Memory policy of the event lanes: `inherit` the process policy, or `local`, which makes + /// each pinned lane prefer its own NUMA node for what it allocates (set_mempolicy). + #[arg(long, value_enum, default_value = "inherit")] + lane_memory: LaneMemory, + /// Diagnostic: count lookups by how many shards hold the request's first block (a shared + /// counter on the lookup path; not for timing rows). + #[arg(long)] + count_shard_heads: bool, + /// Issue the corpus's lookups (`on`) or events only (`off`, a lane-cost discriminator whose + /// rates count event block ops alone). + #[arg(long, value_enum, default_value = "on")] + queries: Queries, + /// Spread each lookup over the shards: `all-shards` (one lane walks every shard) or + /// `per-shard` (one lane per shard on the shard's cores, merged by the last to finish). + #[arg(long, value_enum, default_value = "all-shards")] + lookups: Lookups, + /// Who builds the owned event payloads: the `main` thread before the trial, or each `issuer` + /// in its own pinned thread (payloads on the issuer's socket). + #[arg(long, value_enum, default_value = "main")] + payload_home: PayloadHome, + /// How event lanes share work: `owned` (each lane applies its own workers only, as a + /// channel per lane) or `stealing` (the event lanes of each shard are the lanes of one + /// `LanePool`: any lane serves any ready worker of its shard, whole workers at a time, and a + /// lane that is not running loses a worker once its backlog is `--steal-after` events deep). + /// Recorded in the result with the pools' counters. + #[arg(long, value_enum, default_value = "owned")] + lane_scheduling: LaneScheduling, + /// Under `--lane-scheduling stealing`: events a ready worker may have queued on a lane that + /// is not serving before another lane takes it (the pool's `steal_after`); 0 keeps the + /// pool's serving-lane rule alone. + #[arg(long, default_value = "0")] + steal_after: usize, + /// Replay window in milliseconds; deadlines are rescaled linearly from the corpus's + /// reference window when they differ. + #[arg(long, conflicts_with = "offered_block_ops_per_sec")] + benchmark_duration_ms: Option, + /// Offered rate in block ops per second: sets the window from the corpus's block-op total + /// (window = total / rate), the knob a sustained-throughput threshold search moves. + #[arg(long)] + offered_block_ops_per_sec: Option, + #[arg(long, default_value = "128")] + query_lanes: usize, + /// Event lanes (OS threads applying events). + #[arg(long, default_value = "64")] + event_lanes: usize, + /// Event issuer threads; events are sharded by contiguous worker ranges. + #[arg(long, default_value = "4")] + issuer_threads: usize, + /// CPUs for the event issuers (one per issuer thread when given). + #[arg(long)] + issuer_cpus: Option, + /// Query issuer threads; query lanes are sharded over them in contiguous ranges (one issuer + /// caps the generator near 1.5B block ops/s here). + #[arg(long, default_value = "1")] + query_issuer_threads: usize, + /// CPUs for the query issuers (one per thread when given); `--query-issuer-cpu` is an alias. + #[arg(long, alias = "query-issuer-cpu")] + query_issuer_cpus: Option, + /// CPUs for query and event lanes. + #[arg(long)] + backend_cpus: Option, + #[arg(long, default_value = "100")] + issuer_spin_us: u64, + #[arg(long, default_value = "250")] + issue_lag_diagnostic_threshold_us: u64, + #[arg(long, default_value = "5000")] + pre_run_quiescence_ms: u64, + /// Pin each event lane to one backend CPU (round robin) instead of letting it float over + /// the set; a diagnostic for scheduler effects. + #[arg(long)] + pin_event_lanes: bool, + /// Give each worker's events to the issuer whose lane range holds the worker's lane (issuer k + /// feeds lanes [k * lanes / issuers, (k + 1) * lanes / issuers)) instead of contiguous + /// worker-id ranges; with `--pin-event-lanes` and the issuer CPUs listed in lane order every + /// issuer then sits on the socket of the lanes it feeds. A two-socket diagnostic. + #[arg(long)] + issuer_by_lane: bool, + /// Charge every backend for owning its events: each event arrives as an owned payload in the + /// engine's wire layout (40 bytes per block, allocated before the trial) that the lane converts + /// into this crate's blocks and frees after the apply, and each lookup copies its hashes into + /// this crate's hash type. Off: lanes read the corpus slabs and copy nothing. + #[arg(long, default_value = "true", action = clap::ArgAction::Set)] + owned_payloads: bool, + #[arg(long, default_value = "mooncake_replay_result.json")] + result_json_output: String, +} + +fn main() -> anyhow::Result<()> { + let args = Args::parse(); + let corpus = load_corpus(&args.corpus)?; + let window_ns = match (args.benchmark_duration_ms, args.offered_block_ops_per_sec) { + (Some(ms), _) => ms * 1_000_000, + (None, Some(rate)) => { + anyhow::ensure!(rate > 0.0, "--offered-block-ops-per-sec must be positive"); + (corpus.totals.scheduled_block_ops(args.queries) as f64 / rate * 1e9).round() as u64 + } + (None, None) => corpus.reference_window_ns, + }; + match args.backend { + BackendKind::Positional => run( + &args, + corpus, + window_ns, + Arc::new(Positional { + inner: PositionalIndexer::new(args.jump_size), + }), + ), + BackendKind::Reference => { + if corpus.ops.len() > 2_000_000 { + eprintln!( + "warning: the reference backend is single-threaded; this corpus is large" + ); + } + run( + &args, + corpus, + window_ns, + Arc::new(Reference { + inner: Mutex::new(ReferenceIndexer::new()), + ids: Mutex::new(FxHashMap::default()), + }), + ) + } + BackendKind::Chain => { + let backend_cpus = args + .backend_cpus + .as_deref() + .map(parse_cpu_list) + .transpose()? + .unwrap_or_default(); + let shard_of_cpu = shard_map(&backend_cpus, args.shards.max(1)); + run( + &args, + corpus, + window_ns, + Arc::new(Chain { + inner: ShardedChainIndex::new(args.shards.max(1), args.max_workers), + shard_of_cpu, + event_lanes: args.event_lanes, + heads_held: args.count_shard_heads.then(|| { + (0..=args.shards.max(1)) + .map(|_| AtomicUsize::new(0)) + .collect() + }), + }), + ) + } + BackendKind::Null => run(&args, corpus, window_ns, Arc::new(Null)), + } +} + +/// The shard of each backend CPU: the index of its NUMA node among the nodes the list spans when +/// there are exactly `shards` of them, else the CPU's contiguous group of the list. +fn shard_map(backend_cpus: &[usize], shards: usize) -> FxHashMap { + let mut map = FxHashMap::default(); + if backend_cpus.is_empty() || shards <= 1 { + return map; + } + let mut nodes: Vec = Vec::new(); + let mut node_of: Vec> = Vec::with_capacity(backend_cpus.len()); + for &cpu in backend_cpus { + let node = node_of_cpu(cpu); + if let Some(node) = node { + if !nodes.contains(&node) { + nodes.push(node); + } + } + node_of.push(node); + } + if nodes.len() == shards && node_of.iter().all(Option::is_some) { + for (&cpu, node) in backend_cpus.iter().zip(&node_of) { + let node = node.expect("checked"); + map.insert(cpu, nodes.iter().position(|&n| n == node).expect("listed")); + } + } else { + for (position, &cpu) in backend_cpus.iter().enumerate() { + map.insert(cpu, position * shards / backend_cpus.len()); + } + } + map +} + +fn run( + args: &Args, + mut corpus: Corpus, + window_ns: u64, + backend: Arc, +) -> anyhow::Result<()> { + anyhow::ensure!(args.query_lanes > 0 && args.event_lanes > 0 && args.issuer_threads > 0); + let backend_cpus: Arc<[usize]> = args + .backend_cpus + .as_deref() + .map(parse_cpu_list) + .transpose()? + .unwrap_or_default() + .into(); + let issuer_cpus = args + .issuer_cpus + .as_deref() + .map(parse_cpu_list) + .transpose()? + .unwrap_or_default(); + let issuer_threads = if issuer_cpus.is_empty() { + args.issuer_threads + } else { + issuer_cpus.len() + }; + let query_cpus = args + .query_issuer_cpus + .as_deref() + .map(parse_cpu_list) + .transpose()? + .unwrap_or_default(); + let query_issuers = if query_cpus.is_empty() { + args.query_issuer_threads.max(1) + } else { + query_cpus.len() + }; + // Issuers and lanes must not share cores: a lane on an issuer's core eats the issue schedule + // and the trial measures the layout mistake, not the indexer. + if !backend_cpus.is_empty() { + let overlap: Vec = issuer_cpus + .iter() + .chain(query_cpus.iter()) + .copied() + .filter(|cpu| backend_cpus.contains(cpu)) + .collect(); + anyhow::ensure!( + overlap.is_empty(), + "issuer CPUs {overlap:?} overlap the backend CPU set; give lanes their own cores" + ); + } + // The layout, first line of every run log, so the provenance shows it at a glance. + println!( + "layout: event issuers {} on {:?}, query issuers {} on {:?}, lanes {} event + {} query on {:?} ({} cores), shards {}, lane memory {:?}, event lanes {}, queries {:?}, lookups {:?}, payload home {:?}, lane scheduling {:?} (steal after {})", + issuer_threads, + issuer_cpus, + query_issuers, + query_cpus, + args.event_lanes, + args.query_lanes, + backend_cpus, + backend_cpus.len(), + args.shards.max(1), + args.lane_memory, + if args.pin_event_lanes { "pinned one per core" } else { "floating" }, + args.queries, + args.lookups, + args.payload_home, + args.lane_scheduling, + args.steal_after, + ); + if window_ns != corpus.reference_window_ns { + let reference = corpus.reference_window_ns.max(1) as u128; + for op in &mut corpus.ops { + op.deadline_ns = ((op.deadline_ns as u128 * window_ns as u128) / reference) as u64; + } + } + // Lane assignment and deadline groups. + let mirror = args.owned_payloads; + // Lanes per lookup: one, or one per shard under `--lookups per-shard` (query lane + // `group * fan + shard` walks `shard` for lookup group `group`). + let fan = match args.lookups { + Lookups::AllShards => 1, + Lookups::PerShard => backend.lookup_shards().max(1), + }; + anyhow::ensure!( + args.query_lanes.is_multiple_of(fan), + "--query-lanes {} is not a multiple of the {fan} shards a per-shard lookup fans out to", + args.query_lanes + ); + let lookup_groups = args.query_lanes / fan; + let mut lane_capacities = vec![0usize; args.query_lanes]; + let mut deadline_query_counts: Vec = Vec::new(); + let mut previous_deadline = None; + // A worker's slot is its rank of first appearance; its lane is the slot modulo the lane count. + let mut event_lane_of: FxHashMap<(u64, u32), u32> = FxHashMap::default(); + let mut event_lane_expected = vec![0usize; args.event_lanes]; + let mut query_dispatch: Vec> = vec![Vec::new(); query_issuers]; + let mut event_dispatch: Vec> = + (0..issuer_threads).map(|_| Vec::new()).collect(); + for op in &corpus.ops { + if previous_deadline != Some(op.deadline_ns) { + previous_deadline = Some(op.deadline_ns); + deadline_query_counts.push(0); + } + let group = (deadline_query_counts.len() - 1) as u32; + if op.is_query() { + if args.queries == Queries::Off { + continue; + } + let target = (op.worker as usize) % lookup_groups; + for member in 0..fan { + lane_capacities[target * fan + member] += 1; + } + deadline_query_counts[group as usize] += 1; + query_dispatch[target * query_issuers / lookup_groups].push(( + op.id, + group, + target as u16, + )); + } else { + let next = event_lane_of.len() as u32; + let slot = *event_lane_of + .entry((op.worker, op.dp_rank())) + .or_insert(next); + let lane = (slot as usize % args.event_lanes) as u16; + event_lane_expected[lane as usize] += 1; + let shard = if args.issuer_by_lane { + (lane as usize) * issuer_threads / args.event_lanes + } else { + event_issuer_for(op.worker, corpus.logical_workers, issuer_threads) + }; + // Under --owned-payloads the payload is built here, before the trial, or by the + // issuer itself under --payload-home issuer. + let payload = if mirror && args.payload_home == PayloadHome::Main { + payload_for(&corpus, op) + } else { + Payload::None + }; + event_dispatch[shard].push(EventDispatch { + id: op.id, + group, + lane, + slot, + payload, + }); + } + } + let deadline_pending: Box<[AtomicU32]> = deadline_query_counts + .iter() + .map(|&count| AtomicU32::new(count)) + .collect::>() + .into_boxed_slice(); + + // Quiescence: return preparation pages and let the allocator settle. + // SAFETY: malloc_trim takes an integer pad and has no other preconditions. + unsafe { + libc::malloc_trim(0); + } + if args.pre_run_quiescence_ms > 0 { + thread::sleep(Duration::from_millis(args.pre_run_quiescence_ms)); + } + let corpus = Arc::new(corpus); + // Page-touch the corpus once. + let mut checksum = 0u64; + for hash in &*corpus.hashes { + checksum ^= hash.0; + } + for block in &*corpus.blocks { + checksum ^= block.seq_hash.0 ^ block.content_hash.0; + } + for op in &corpus.ops { + checksum ^= op.deadline_ns ^ op.worker ^ u64::from(op.id); + } + black_box(checksum); + + pin_current_thread(&backend_cpus)?; + let clock = Clock::new(args.issuer_spin_us.saturating_mul(1_000))?; + let epoch = clock.epoch; + // Fixed lookup warm-up. + for op in corpus + .ops + .iter() + .filter(|op| op.is_query()) + .take(if args.queries == Queries::On { + WARMUP_QUERIES + } else { + 0 + }) + { + if let OpKind::Query { start, len } = op.kind { + black_box( + backend.lookup(&corpus.hashes[start as usize..start as usize + len as usize]), + ); + } + } + + // Lanes. + let lanes: Vec> = lane_capacities + .iter() + .map(|&capacity| Arc::new(QueryLane::new(capacity))) + .collect(); + // Under `--lookups per-shard`, the hand-over slots of every lookup group, one per published + // position (every member lane of a group receives the same ids in the same order). + let fan_slots: Vec> = (0..lookup_groups) + .map(|group| { + let slots = if fan > 1 { + lane_capacities[group * fan] + } else { + 0 + }; + (0..slots) + .map(|_| FanSlot::default()) + .collect::>() + .into() + }) + .collect(); + let mut query_threads = Vec::with_capacity(lanes.len()); + for (index, lane) in lanes.iter().enumerate() { + let (backend, lane, corpus) = (Arc::clone(&backend), Arc::clone(lane), Arc::clone(&corpus)); + let (lookup, cpus) = if fan > 1 { + let shard = index % fan; + let cpus: Arc<[usize]> = backend + .shard_cpus(shard) + .map_or_else(|| Arc::clone(&backend_cpus), Arc::from); + let lookup = LaneLookup::Shard { + shard, + fan, + slots: Arc::clone(&fan_slots[index / fan]), + }; + (lookup, cpus) + } else { + (LaneLookup::AllShards, Arc::clone(&backend_cpus)) + }; + query_threads.push(thread::spawn(move || { + query_lane_worker(backend, lane, corpus, epoch, cpus, mirror, lookup) + })); + } + // Wait until every query lane has registered its parker. + for lane in &lanes { + while lane + .consumer + .lock() + .unwrap_or_else(|e| e.into_inner()) + .is_none() + { + thread::yield_now(); + } + } + let lane_cpus: Vec> = (0..args.event_lanes) + .map(|idx| { + if args.pin_event_lanes && !backend_cpus.is_empty() { + Arc::from(vec![backend_cpus[idx % backend_cpus.len()]]) + } else { + Arc::clone(&backend_cpus) + } + }) + .collect(); + // Under `--lane-scheduling stealing` the event lanes of each shard are the lanes of one + // pool, so a worker is only ever applied by a lane of its own shard; the pool index of a lane + // is its rank among its shard's lanes. + let stealing = args.lane_scheduling == LaneScheduling::Stealing; + let lane_shard: Vec = (0..args.event_lanes) + .map(|idx| backend.lane_shard(idx, &lane_cpus[idx])) + .collect(); + let shard_count = lane_shard.iter().copied().max().map_or(1, |s| s + 1); + let mut pool_lane_index = vec![0usize; args.event_lanes]; + let mut lanes_per_shard = vec![0usize; shard_count]; + for (idx, &shard) in lane_shard.iter().enumerate() { + pool_lane_index[idx] = lanes_per_shard[shard]; + lanes_per_shard[shard] += 1; + } + let shard_pools: Vec>>> = lanes_per_shard + .iter() + .map(|&lanes| { + (stealing && lanes > 0).then(|| { + Arc::new(ShardPool { + pool: LanePool::new(LanePoolConfig { + lanes, + max_workers: event_lane_of.len().max(1), + // Never refuse: the harness's own queues are unbounded, and a refusal + // here would only move an event into the lane's held slot. + depth_cap: corpus.totals.events().try_into().unwrap_or(usize::MAX), + batch: 32, + steal_after: args.steal_after, + }), + open_channels: AtomicUsize::new(lanes), + held_total: AtomicUsize::new(0), + enqueue_ns: AtomicU64::new(0), + stealable_idle_turns: AtomicU64::new(0), + }) + }) + }) + .collect(); + let mut senders = Vec::with_capacity(args.event_lanes); + let mut event_threads = Vec::with_capacity(args.event_lanes); + for (idx, &expected) in event_lane_expected.iter().enumerate() { + let (tx, rx) = mpsc::channel::(); + senders.push(tx); + let (backend, corpus) = (Arc::clone(&backend), Arc::clone(&corpus)); + let placement = LanePlacement { + index: idx, + cpus: Arc::clone(&lane_cpus[idx]), + local_memory: matches!(args.lane_memory, LaneMemory::Local), + }; + match shard_pools[lane_shard[idx]].as_ref().map(Arc::clone) { + Some(shard_pool) => { + let pool_lane = pool_lane_index[idx]; + event_threads.push(thread::spawn(move || { + event_lane_worker_pooled( + backend, rx, corpus, epoch, placement, expected, shard_pool, pool_lane, + ) + })); + } + None => event_threads.push(thread::spawn(move || { + event_lane_worker(backend, rx, corpus, epoch, placement, expected) + })), + } + } + + let shared = Shared { + clock, + start_ns: AtomicUsize::new(0), + deadline_pending, + peer_failed: AtomicBool::new(false), + }; + let ready = Barrier::new(issuer_threads + query_issuers + 1); + let start = Barrier::new(issuer_threads + query_issuers + 1); + let (start_ns, outputs) = thread::scope(|scope| { + let mut handles = Vec::with_capacity(issuer_threads + query_issuers); + let (shared_ref, corpus_ref, lanes_ref, ready_ref, start_ref) = + (&shared, &corpus, &lanes, &ready, &start); + for (idx, dispatch) in query_dispatch.iter().enumerate() { + let cpu = query_cpus.get(idx).copied(); + handles.push(scope.spawn(move || { + let pin = cpu.map_or(Ok(()), |cpu| pin_current_thread(&[cpu])); + let mut failure = pin.err().map(|_| "issuer_affinity"); + if failure.is_some() { + shared_ref.peer_failed.store(true, Ordering::Release); + } + ready_ref.wait(); + start_ref.wait(); + let mut records = Vec::with_capacity(dispatch.len()); + let (cpu_ns, f) = if failure.is_none() { + issue_queries( + shared_ref, + corpus_ref, + dispatch, + lanes_ref, + fan, + &mut records, + ) + } else { + (0, None) + }; + failure = failure.or(f); + (records, cpu_ns, failure) + })); + } + let issuer_payloads = mirror && args.payload_home == PayloadHome::Issuer; + for (idx, dispatch) in event_dispatch.into_iter().enumerate() { + let cpu = issuer_cpus.get(idx).copied(); + let senders = &senders; + handles.push(scope.spawn(move || { + let mut dispatch = dispatch; + let pin = cpu.map_or(Ok(()), |cpu| pin_current_thread(&[cpu])); + let mut failure = pin.err().map(|_| "issuer_affinity"); + if failure.is_some() { + shared_ref.peer_failed.store(true, Ordering::Release); + } + if issuer_payloads { + // Built here, pinned: the payloads take this thread's arena and NUMA node. + for entry in &mut dispatch { + entry.payload = payload_for(corpus_ref, &corpus_ref.ops[entry.id as usize]); + } + } + ready_ref.wait(); + start_ref.wait(); + let mut records = Vec::with_capacity(dispatch.len()); + let (cpu_ns, f) = if failure.is_none() { + issue_events(shared_ref, corpus_ref, dispatch, senders, &mut records) + } else { + (0, None) + }; + failure = failure.or(f); + (records, cpu_ns, failure) + })); + } + ready.wait(); + let start_ns = shared.clock.now_ns().saturating_add(20_000_000); + shared.start_ns.store(start_ns as usize, Ordering::Release); + start.wait(); + let outputs: Vec<_> = handles + .into_iter() + .map(|h| h.join().expect("issuer thread panicked")) + .collect(); + (start_ns, outputs) + }); + let producer_stop_ns = shared.clock.now_ns(); + drop(senders); + for lane in &lanes { + lane.close(); + } + let mut failure_reasons: Vec = Vec::new(); + let mut query_results = Vec::with_capacity(lanes.len()); + let mut query_lane_cpu_ns = 0u64; + let mut fan_stats = vec![FanStats::default(); fan]; + for (index, handle) in query_threads.into_iter().enumerate() { + let output = handle.join().expect("query lane panicked"); + if let Some(failure) = output.failure { + failure_reasons.push(failure.to_string()); + } + query_lane_cpu_ns += output.cpu_ns; + fan_stats[index % fan].add(&output.fan); + query_results.push(output.completions); + } + let outcomes: Vec = event_threads + .into_iter() + .map(|h| h.join().expect("event lane panicked")) + .collect(); + let event_lane_cpu_ns: Vec = outcomes.iter().map(|o| o.cpu_ns).collect(); + let event_lane_apply_ms: Vec = outcomes.iter().map(|o| o.apply_ns as f64 / 1e6).collect(); + let event_lane_minor_faults: Vec = outcomes.iter().map(|o| o.minor_faults).collect(); + let event_results: Vec> = + outcomes.into_iter().map(|o| o.completions).collect(); + + // Merge issue records. + let n = corpus.ops.len(); + let mut records = vec![IssueRecord::default(); n]; + let mut issuer_cpu_ns = 0u64; + for (local, cpu_ns, failure) in outputs { + issuer_cpu_ns += cpu_ns; + if let Some(failure) = failure { + failure_reasons.push(failure.to_string()); + } + for (id, record) in local { + if records[id as usize].accepted { + failure_reasons.push("duplicate_issue_record".to_string()); + } + records[id as usize] = record; + } + } + // Completions by id, with order checks. + let mut query_done: Vec> = vec![None; n]; + for (lane_idx, completions) in query_results.iter().enumerate() { + let group = lane_idx / fan; + let expected: Vec = query_dispatch + .iter() + .flatten() + .filter(|(_, _, target)| *target as usize == group) + .map(|(id, _, _)| *id) + .collect(); + let actual: Vec = completions.iter().map(|c| c.id).collect(); + let in_order = if fan == 1 { + expected == actual + } else { + // A member lane completes the lookups it merged, in publish order: a subsequence. + let mut remaining = expected.iter(); + actual.iter().all(|id| remaining.any(|e| e == id)) + }; + if !in_order { + failure_reasons.push(format!("query_lane_order_{lane_idx}")); + } + for c in completions { + if query_done[c.id as usize].replace(*c).is_some() { + failure_reasons.push("duplicate_query_completion".to_string()); + } + } + } + let mut event_done: Vec> = vec![None; n]; + // A worker's events in the order they were applied: by finish time, since under + // `--lane-scheduling stealing` a worker's events are applied by more than one lane (one at + // a time, so finish order is apply order) and its completions sit in several lanes' lists. + let mut actual_by_worker: BTreeMap<(u64, u32), Vec<(u64, u32)>> = BTreeMap::new(); + let mut failed_events = 0usize; + for completions in &event_results { + for c in completions { + let op = &corpus.ops[c.id as usize]; + actual_by_worker + .entry((op.worker, op.dp_rank())) + .or_default() + .push((c.finished_ns, c.id)); + if !c.ok { + failed_events += 1; + } + if event_done[c.id as usize].replace(*c).is_some() { + failure_reasons.push("duplicate_event_completion".to_string()); + } + } + } + let actual_by_worker: BTreeMap<(u64, u32), Vec> = actual_by_worker + .into_iter() + .map(|(worker, mut done)| { + done.sort_unstable(); + (worker, done.into_iter().map(|(_, id)| id).collect()) + }) + .collect(); + let mut expected_by_worker: BTreeMap<(u64, u32), Vec> = BTreeMap::new(); + for op in corpus.ops.iter().filter(|op| !op.is_query()) { + expected_by_worker + .entry((op.worker, op.dp_rank())) + .or_default() + .push(op.id); + } + let mut fifo_violations = 0usize; + for (worker, expected) in &expected_by_worker { + if actual_by_worker.get(worker) != Some(expected) { + fifo_violations += 1; + } + } + if fifo_violations > 0 { + failure_reasons.push(format!("event_worker_fifo_{fifo_violations}")); + } + + let tolerance_ns = args.issue_lag_diagnostic_threshold_us.saturating_mul(1_000); + let (mut read_lag, mut update_lag, mut queue_wait, mut service, mut query_e2e) = + (Vec::new(), Vec::new(), Vec::new(), Vec::new(), Vec::new()); + let (mut update_acc_fin, mut update_e2e) = (Vec::new(), Vec::new()); + let (mut delayed_reads, mut delayed_updates, mut races) = (0usize, 0usize, 0usize); + let mut query_edges = Vec::new(); + let mut update_edges = Vec::new(); + let mut last_completion = 0u64; + let mut unissued = 0usize; + let (mut queued_queries_at_stop, mut outstanding_updates_at_stop) = (0usize, 0usize); + for (id, record) in records.iter().enumerate() { + if !record.accepted { + // Under `--queries off` the lookups are left out of the dispatch on purpose. + if args.queries == Queries::On || !corpus.ops[id].is_query() { + unissued += 1; + } + continue; + } + let lag = record.accepted_ns.saturating_sub(record.scheduled_ns); + if record.is_query { + read_lag.push(lag); + delayed_reads += usize::from(lag > tolerance_ns); + let Some(c) = query_done[id] else { + failure_reasons.push(format!("missing_query_completion_{id}")); + continue; + }; + last_completion = last_completion.max(c.finished_ns); + queued_queries_at_stop += usize::from( + record.accepted_ns <= producer_stop_ns && c.started_ns > producer_stop_ns, + ); + queue_wait.push(c.started_ns.saturating_sub(record.accepted_ns)); + service.push(c.finished_ns.saturating_sub(c.started_ns)); + query_e2e.push(c.finished_ns.saturating_sub(record.scheduled_ns)); + query_edges.push((record.accepted_ns, 1i8)); + query_edges.push((c.started_ns.max(record.accepted_ns), -1i8)); + } else { + update_lag.push(lag); + delayed_updates += usize::from(lag > tolerance_ns); + let Some(c) = event_done[id] else { + failure_reasons.push(format!("missing_event_completion_{id}")); + continue; + }; + last_completion = last_completion.max(c.finished_ns); + outstanding_updates_at_stop += usize::from(c.finished_ns > producer_stop_ns); + if c.finished_ns < record.accepted_ns { + races += 1; + } + update_acc_fin.push(c.finished_ns.saturating_sub(record.accepted_ns)); + update_e2e.push(c.finished_ns.saturating_sub(record.scheduled_ns)); + update_edges.push((record.accepted_ns, 1i8)); + update_edges.push((c.finished_ns.max(record.accepted_ns), -1i8)); + } + } + if unissued > 0 { + failure_reasons.push(format!("unissued_operations_{unissued}")); + } + if failure_reasons.len() > 32 { + failure_reasons.truncate(32); + failure_reasons.push("...".to_string()); + } + let end_ns = if last_completion > 0 { + last_completion + } else { + producer_stop_ns + }; + let issue_span_ns = records + .iter() + .filter(|r| r.accepted) + .map(|r| r.accepted_ns) + .max() + .unwrap_or(start_ns) + .saturating_sub(start_ns); + let drain_ns = end_ns.saturating_sub(producer_stop_ns); + let elapsed = end_ns.saturating_sub(start_ns).max(1); + let totals = corpus.totals; + let (issued_requests, issued_request_blocks) = match args.queries { + Queries::On => (totals.requests, totals.request_blocks), + Queries::Off => (0, 0), + }; + let total_logical_ops = issued_requests + totals.events(); + let total_block_ops = totals.scheduled_block_ops(args.queries); + let offered_s = window_ns.max(1) as f64 / 1e9; + let achieved_s = elapsed as f64 / 1e9; + let issue_s = issue_span_ns.max(1) as f64 / 1e9; + let issue_span_valid = issue_span_ns <= window_ns.saturating_mul(101) / 100; + let generator_valid = failure_reasons.is_empty() && issue_span_valid; + let kept_up = generator_valid && elapsed <= window_ns.saturating_mul(110) / 100; + + // Per-lane diagnostics: CPU time, event count and the time of the + // last completion of each event lane, which tell scheduler starvation from slower work. + let event_lane_cpu_ms: Vec = event_lane_cpu_ns + .iter() + .map(|&ns| ns as f64 / 1e6) + .collect(); + let event_lane_events: Vec = event_results.iter().map(Vec::len).collect(); + let event_lane_last_finished_ms: Vec = event_results + .iter() + .map(|c| { + c.last() + .map_or(0.0, |c| c.finished_ns.saturating_sub(start_ns) as f64 / 1e6) + }) + .collect(); + let result = json!({ + "schema_version": 3, + "harness": "smg-mooncake-replay", + "backend": backend.name(), + "corpus": args.corpus, + "owned_payloads": args.owned_payloads, + "shards": args.shards.max(1), + "lane_memory": format!("{:?}", args.lane_memory).to_lowercase(), + "queries": format!("{:?}", args.queries).to_lowercase(), + "lookups": match args.lookups { + Lookups::AllShards => "all-shards", + Lookups::PerShard => "per-shard", + }, + "payload_home": format!("{:?}", args.payload_home).to_lowercase(), + "lane_scheduling": format!("{:?}", args.lane_scheduling).to_lowercase(), + "steal_after": args.steal_after, + "lane_pool": stealing.then(|| { + shard_pools + .iter() + .enumerate() + .filter_map(|(shard, pool)| pool.as_ref().map(|pool| (shard, pool))) + .map(|(shard, shard_pool)| { + let m = shard_pool.pool.metrics(); + json!({ + "shard": shard, + "lanes": shard_pool.pool.config().lanes, + "steal_after": shard_pool.pool.config().steal_after, + "enqueued": m.enqueued, + "applied": m.applied, + "rejected": m.rejected, + "steals": m.steals, + "max_depth": m.max_depth, + "max_queued": m.max_queued, + "busy_ms": m.busy_ns as f64 / 1e6, + "idle_ms": m.idle_ns as f64 / 1e6, + "enqueue_wait_ms": shard_pool.enqueue_ns.load(Ordering::Relaxed) as f64 / 1e6, + "stealable_idle_turns": shard_pool.stealable_idle_turns.load(Ordering::Relaxed), + }) + }) + .collect::>() + }), + "lookup_fanout": (fan > 1).then(|| json!({ + "fan": fan, + "groups": lookup_groups, + "shards": fan_stats.iter().enumerate().map(|(shard, s)| json!({ + "shard": shard, + "cpus": backend.shard_cpus(shard), + "lookups": s.lookups, + "with_holders": s.with_holders, + "holders": s.holders, + "merged": s.merged, + })).collect::>(), + })), + "backend_report": backend.report(), + "provenance": { + "argv": std::env::args().collect::>(), + "binary": std::env::current_exe().ok().map(|p| p.display().to_string()), + "binary_blake3": std::env::current_exe() + .ok() + .and_then(|p| std::fs::read(p).ok()) + .map(|b| blake3::hash(&b).to_hex().to_string()), + "corpus_blake3": corpus.file_blake3, + "corpus_reference_window_ns": corpus.reference_window_ns, + "trace_path": corpus.trace_path, + "trace_block_size": corpus.block_size, + "trace_duplication_factor": corpus.trace_duplication_factor, + "trace_length_factor": corpus.trace_length_factor, + "inference_worker_duplication_factor": corpus.inference_worker_duplication_factor, + "num_unique_inference_workers": corpus.logical_workers, + "jump_size": args.jump_size, + "issuer_spin_us": args.issuer_spin_us, + "issue_lag_diagnostic_threshold_us": args.issue_lag_diagnostic_threshold_us, + }, + "timer": "clock_nanosleep_monotonic_absolute", + "benchmark_duration_ms": window_ns / 1_000_000, + "block_size": corpus.block_size, + "pre_run_quiescence_ms": args.pre_run_quiescence_ms, + "query_lanes": args.query_lanes, + "issuer_threads": issuer_threads, + "event_workers": args.event_lanes, + "issuer_cpus": issuer_cpus, + "query_issuer_threads": query_issuers, + "query_issuer_cpus": query_cpus, + "backend_cpus": backend_cpus.to_vec(), + "total_requests": issued_requests, + "total_events": totals.events(), + "total_stored_events": totals.stored_events, + "total_removed_events": totals.removed_events, + "total_cleared_events": totals.cleared_events, + "total_request_blocks": issued_request_blocks, + "total_stored_blocks": totals.stored_blocks, + "total_removed_blocks": totals.removed_blocks, + "total_logical_ops": total_logical_ops, + "total_block_ops": total_block_ops, + "offered_logical_ops_per_sec": total_logical_ops as f64 / offered_s, + "actual_issue_logical_ops_per_sec": total_logical_ops as f64 / issue_s, + "achieved_logical_ops_per_sec": total_logical_ops as f64 / achieved_s, + "offered_block_ops_per_sec": total_block_ops as f64 / offered_s, + "actual_issue_block_ops_per_sec": total_block_ops as f64 / issue_s, + "achieved_block_ops_per_sec": total_block_ops as f64 / achieved_s, + "read_issue_lag": distribution(read_lag), + "update_issue_lag": distribution(update_lag), + "generator_gate": "issue_span_exact_completion", + "query_queue_wait": distribution(queue_wait), + "query_service": distribution(service), + "query_scheduled_to_finished": distribution(query_e2e), + "update_accepted_to_finished": distribution(update_acc_fin), + "update_scheduled_to_finished": distribution(update_e2e), + "delayed_reads": delayed_reads, + "delayed_updates": delayed_updates, + "maximum_query_queue_depth": maximum_depth(&mut query_edges), + "maximum_outstanding_updates": maximum_depth(&mut update_edges), + "queued_queries_at_stop": queued_queries_at_stop, + "outstanding_updates_at_stop": outstanding_updates_at_stop, + "post_acceptance_completion_races": races, + "rejected_events": failed_events, + "issuer_cpu_ns": issuer_cpu_ns, + "pin_event_lanes": args.pin_event_lanes, + "issuer_by_lane": args.issuer_by_lane, + "event_lane_cpu_ms": event_lane_cpu_ms, + "event_lane_apply_ms": event_lane_apply_ms, + "event_lane_minor_faults": event_lane_minor_faults, + "event_lane_events": event_lane_events, + "event_lane_last_finished_ms": event_lane_last_finished_ms, + "query_lane_cpu_ms_total": query_lane_cpu_ns as f64 / 1e6, + "issue_span_ns": issue_span_ns, + "drain_ns": drain_ns, + "generator_valid": generator_valid, + "kept_up": kept_up, + "failure_reasons": failure_reasons, + }); + println!( + "{} window {} ms: offered {:.1}M achieved {:.1}M block ops/s, kept_up {}, valid {}, \ + lookup service p50 {:.2} us p99 {:.2} us, scheduled->finished p99 {:.1} us, drain {:.1} ms", + backend.name(), + window_ns / 1_000_000, + result["offered_block_ops_per_sec"].as_f64().unwrap_or(0.0) / 1e6, + result["achieved_block_ops_per_sec"].as_f64().unwrap_or(0.0) / 1e6, + kept_up, + generator_valid, + result["query_service"]["p50_ns"].as_u64().unwrap_or(0) as f64 / 1e3, + result["query_service"]["p99_ns"].as_u64().unwrap_or(0) as f64 / 1e3, + result["query_scheduled_to_finished"]["p99_ns"].as_u64().unwrap_or(0) as f64 / 1e3, + drain_ns as f64 / 1e6, + ); + if fan > 1 { + let shards: Vec = fan_stats + .iter() + .enumerate() + .map(|(shard, s)| { + format!( + "shard {shard}: lookups {} with holders {} holders {} merged {}", + s.lookups, s.with_holders, s.holders, s.merged + ) + }) + .collect(); + println!("lookups per shard (fan-out {fan}): {}", shards.join("; ")); + } + if !result["failure_reasons"] + .as_array() + .is_some_and(Vec::is_empty) + { + println!("failure reasons: {}", result["failure_reasons"]); + } + std::fs::write( + &args.result_json_output, + serde_json::to_vec_pretty(&result)?, + )?; + Ok(()) +} + +fn maximum_depth(edges: &mut [(u64, i8)]) -> usize { + edges.sort_unstable_by(|l, r| l.0.cmp(&r.0).then_with(|| r.1.cmp(&l.1))); + let (mut depth, mut maximum) = (0isize, 0isize); + for &(_, delta) in edges.iter() { + depth += delta as isize; + maximum = maximum.max(depth); + } + maximum.max(0) as usize +} + +fn distribution(mut values: Vec) -> serde_json::Value { + if values.is_empty() { + return json!({"p50_ns": 0, "p99_ns": 0, "p999_ns": 0, "max_ns": 0}); + } + values.sort_unstable(); + let rank = |num: usize, den: usize| { + let r = values.len().saturating_mul(num).div_ceil(den).max(1); + values[r.saturating_sub(1).min(values.len() - 1)] + }; + json!({ + "p50_ns": rank(50, 100), + "p99_ns": rank(99, 100), + "p999_ns": rank(999, 1000), + "max_ns": values[values.len() - 1], + }) +} diff --git a/crates/kv_index/src/chain_index.rs b/crates/kv_index/src/chain_index.rs new file mode 100644 index 0000000000..a24ac58270 --- /dev/null +++ b/crates/kv_index/src/chain_index.rs @@ -0,0 +1,2132 @@ +//! A run-compressed KV index: the event-driven prefix index as a tree of runs with per-run +//! worker coverage bitsets, lock-free and allocation-free for readers, bounded in memory. +//! +//! The index answers the same question as [`PositionalIndexer`](crate::PositionalIndexer): for a +//! request given as its per-block content hashes, how many leading blocks does each worker hold, +//! where "holds" means the worker stored, at every position up to there, the block that sits on +//! the request's chain. It is fed by the same engine events (stored / removed / cleared, keyed by +//! the engine's block hashes) and keeps the same per-worker block map the gateway's event monitor +//! owns, so it drops into the same call sites. Its results are checked against +//! [`ReferenceIndexer`](crate::ReferenceIndexer) by `tests/exactness_chain.rs` (including evictions +//! that leave holes in a chain) and under concurrent lanes by `tests/concurrency_chain.rs`. +//! +//! Shape: +//! - A **run** is a maximal stretch of consecutive positions on one chain whose set of holding +//! workers is the same at every position. It stores one content hash per position and one +//! coverage bit per worker. Runs form a tree: a run's children continue it with different next +//! blocks. A store appends to a run or adds a child; a divergence or a coverage change splits a +//! run into a prefix and a suffix; positions never move except by a split, which forwards them. +//! Because coverage is uniform within a run, a worker that evicts a block in the middle of a +//! chain simply stops covering the piece that holds it, and a lookup stops there for that +//! worker and nowhere else: holes are exact. +//! - A **lookup** walks from the root, comparing the request's content hashes against each run on +//! the path (one compare loop per run, not one probe per block) and ANDing the alive set with +//! the run's coverage; a worker that drops out scores the position where it dropped. Readers +//! take no locks, allocate nothing but the result map and write to no memory at all: every run +//! header, hash array and child table lives in an arena addressed by integer ids, a run's +//! window `(hash array, base, length, children)` is read under a seqlock version that is +//! checked again after the run's hashes, coverage and child entry have been read, so a split, +//! a growth, an unlink or a reuse of the run is atomic to a reader. +//! - **Writers** (the event lanes, one per engine worker) lock one run at a time, plus its parent +//! for the moment it takes to unlink an empty run. A decode extension of the worker's own leaf +//! appends in place: one lock, no allocation. A split shares the hash array between prefix and +//! suffix (no copy) and leaves a forwarding record so map entries written before it still +//! resolve, which keeps other lanes' maps untouched. The lane's own map takes an engine hash to +//! `(run, offset)`. +//! - **Memory is recycled.** Run headers, hash arrays (reference-counted across the runs a split +//! leaves sharing one) and child tables return to free lists when they die; a run id carries a +//! generation so a stale child entry or forwarding record to a reused id is recognised. +//! Children are an open-addressing table (linear probing, tombstones, rebuilt at 3/4 load), so +//! a node with many children, the root above all, inserts in constant time. +//! +//! Engine hashes: the index trusts the engine's parent pointers and block identities, as the +//! positional indexer does. Nothing is shared between workers through the maps, so one engine +//! reusing a hash cannot corrupt another worker's view. The index carries one engine hash per +//! distinct block, the one its first holder stored, in an array parallel to the content hashes. +//! The engine hash is a chain hash (a hash of the parent's hash and the block's content), which +//! is what lets a store walk match a run by one engine hash at the end of its window instead of +//! a content hash per block (`match_run`); the content hash at the landing is still checked, and +//! a store that fails it (`landing_mismatches`) is placed by its content, block by block. A +//! worker whose engine names the same content by other hashes (`engine_conflicts` in the stats) +//! is matched by content too; its own hashes key its lane map, so nothing else changes for it. +//! The index therefore assumes one engine hash per worker per (parent, content) position. A +//! second name for a position the worker holds (a twin) is filed onto that position and +//! counted, not stored twice: the worker still scores the position, but once the first name is +//! removed the position goes with it while the second name is still in the lane map, and a +//! store under that name cuts the run at the parent and goes on after it (`store_in_run`). The +//! reference indexer keeps a position until its last name goes, so under twins the index holds +//! at most what the reference holds and never scores a worker above it (the exactness suite's +//! twins test); the normalizer guarantees one name per position for vLLM, whose hash is the +//! chain hash, so the count stays zero in the gateway, and a reading of sliding-window events +//! against the wrong tokens is what produced thousands of them in an offline capture. +//! A lane-map slot carries its key: a probe is settled in the map itself, with no read into the +//! index per block (key-less 8-byte slots checked through the index cost 30-40% of lane CPU). +//! +//! Memory: 16 bytes per distinct block on a chain (its content hash and its engine hash, shared +//! by every worker that holds it, with about 12% slack for growth) plus a 64-byte run header, the +//! coverage words and a child table per branching run, against the 16-byte lane-map slot each +//! lane keeps per held block for removals. + +use std::{ + collections::BTreeSet, + sync::{ + atomic::{fence, AtomicIsize, AtomicU32, AtomicU64, AtomicUsize, Ordering}, + Arc, OnceLock, + }, +}; + +use crossbeam_queue::SegQueue; +use crossbeam_utils::CachePadded; +use dashmap::{mapref::entry::Entry, DashMap}; +use parking_lot::{Mutex, MutexGuard}; +use rustc_hash::{FxBuildHasher, FxHashMap, FxHashSet}; + +use crate::event_tree::{ + chain_prefix_hash, ApplyError, ContentHash, OverlapScores, SequenceHash, StoredBlock, + WorkerIdExhausted, +}; + +mod arena; +mod slab; +#[cfg(test)] +mod tests; +mod walk; + +use arena::*; +use slab::*; + +/// The virtual root: position 0's parent, holds no blocks, never dies. +const ROOT: u32 = 0; +/// "No table" / "no array": word 0 of the arena is never handed out. +const NONE: u32 = 0; +/// Forwarding target of blocks whose last holder evicted them: nowhere. +const GONE: u32 = u32::MAX; +/// A table slot whose child was unlinked. +const TOMB: u64 = u64::MAX; +/// Runs per slab chunk (1024) and chunks in the directory (64 Mi runs in all). +const RUN_CHUNK_BITS: u32 = 10; +const RUN_CHUNK: usize = 1 << RUN_CHUNK_BITS; +const RUN_DIR: usize = 1 << 16; +/// Words per arena chunk (1 Mi, 8 MiB) and chunks in the directory (4 Gi words in all). +const WORD_CHUNK_BITS: u32 = 20; +const WORD_CHUNK: usize = 1 << WORD_CHUNK_BITS; +const WORD_DIR: usize = 1 << 12; +/// Hash array capacities: multiples of 8 up to 128, then powers of two. +const SMALL_ARRAY_CLASSES: usize = 16; +const ARRAY_CLASSES: usize = SMALL_ARRAY_CLASSES + 4 * 13; +/// Classes above the wanted one an allocation may take a freed array from (one octave). +const ARRAY_FIT_SPAN: usize = 4; +/// Child table slot counts: powers of two from 2. +const MIN_TABLE_SLOTS: usize = 2; +const TABLE_CLASSES: usize = 27; +/// Coverage words per run at most: 1024 workers. +const MAX_WORDS: usize = 16; +/// Partial holders a lookup can buffer per run: every worker at most. +const MAX_PARTIAL: usize = MAX_WORDS * 64; +/// Prefix holders a run carries before it is split at their median cutoff: a lookup reads every +/// entry of a run it walks, so a hot chain held to a hundred different depths costs a hundred +/// entries per lookup as one run and a few bitset words as a handful; a split here turns the +/// holders at or past the median into whole holders of the prefix. Sixteen left the A2 pairs +/// unchanged at one and two shards where thirty-two still cost 0.1 us at two. Splits for this reason are +/// bounded by holders, not by requests, so they do not accumulate the way splits at every +/// divergence did. +const PARTIAL_CAP: usize = 16; + +thread_local! { + /// The lookup's buffer for one run's partial-holder entries, per thread: filled and read + /// within one walk, never cleared, so a lookup does not zero 8 KB of stack first. + static PARTIAL_BUFFER: std::cell::RefCell<[u64; MAX_PARTIAL]> = + const { std::cell::RefCell::new([0; MAX_PARTIAL]) }; +} + +/// One step of [`ChainIndex::hop`] along a block's forwarding chain. +enum Hop { + /// The place forwards on: the next place and the generation its run must have. + Next(BlockRef, u32), + /// The place stands; the run's generation. + Here(u32), + /// Nobody holds the block any more. + Gone, +} + +/// Where one of a worker's blocks lives: the run and the offset of the block within it. +#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, PartialOrd, Ord)] +pub struct BlockRef { + pub run: u32, + pub offset: u32, +} + +pub use crate::lane_map::ChainBlockMap; + +/// Memory and shape counters, as the replay harness reports them and the gateway exports them. +#[derive(Debug, Clone, Copy, Default)] +pub struct ChainIndexStats { + /// Run headers ever created (resident). + pub runs_allocated: usize, + /// Dead headers waiting for reuse. + pub runs_free: usize, + /// Runs linked in the tree. + pub runs_live: usize, + /// Content hashes held by live runs. + pub blocks_live: usize, + /// Bytes of the word arena handed out so far (hash arrays, child tables, free lists + /// included). + pub arena_bytes: usize, + /// Bytes of the word arena sitting in free lists. + pub arena_free_bytes: usize, + /// Bytes the word arena holds from the process allocator: `arena_bytes` rounded up to whole + /// chunks (8 MiB each). + pub arena_chunk_bytes: usize, + /// Bytes of run headers, coverage words included (all allocated runs). + pub header_bytes: usize, + /// Bytes the run slab holds from the process allocator: whole chunks of headers and coverage + /// words, used or not. + pub slab_bytes: usize, + /// Stored blocks whose engine hash differed from the one the index carries for the block + /// (an engine fleet that does not agree on hashes); such a block is matched by content and + /// keyed by the worker's own hash in its lane map. + pub engine_conflicts: usize, + /// Stores whose blocks carried the engine hashes the index holds for other content: what + /// the relay's hash check refuses at the engine, seen from the index. Such a store is + /// placed by its content, so the index stays exact; the count names a broken engine. + pub landing_mismatches: usize, + /// Blocks a store moved to another place for their worker (an engine hash stored again at + /// a different position, as after a store without its parent); the old membership is + /// released so a hash is held at one place per worker. + pub moved_hashes: usize, + /// Runs split because a stored chain diverged inside them: always zero, since a chain that + /// diverges inside a run hangs a child off that offset and leaves the run whole; the churn + /// gate and the split-counter test pin it. + pub splits_by_branch: usize, + /// Runs split because a removal left a worker holding blocks on both sides of a hole. + pub splits_by_hole: usize, + /// Runs split because a store entered them under a parent the worker did not hold up to + /// (a stale parent, or a worker re-entering a chain it had lost the head of). + pub splits_by_mid_run_store: usize, + /// Runs split at the median cutoff of their prefix holders once more than `PARTIAL_CAP` + /// held a prefix of them (bounded by holders, never by requests). + pub splits_by_prefix_holders: usize, + /// Runs unlinked from the tree since the index was created (their headers go back to the + /// slab's free list). + pub runs_died: usize, + /// Prefix-holder entries over the live runs (a lookup reads a run's entries when a holder it + /// follows is not a whole holder of the run), and the most one run carries. + pub partial_entries: usize, + pub max_partials: usize, + /// Child-table entries over the live runs, live and tombstoned. + pub child_entries: usize, + pub child_tombstones: usize, +} + +/// Why a run was split: a hole a worker opens in a run it holds, or the cut at a parent entry +/// that points past what the worker holds. +#[derive(Clone, Copy)] +enum SplitCause { + Hole, + MidRunStore, +} + +/// Worker slots: a slot is in use from `intern_worker` until `remove_worker`, after which it is +/// handed out again (every coverage bit of a removed worker is clear by then). +#[derive(Default)] +struct WorkerRegistry { + names: Vec>>, + free: Vec, +} + +/// The run-compressed index. Worker ids are interned `u32`s, as in the positional indexer. +pub struct ChainIndex { + slab: RunSlab, + arena: WordArena, + words: usize, + max_workers: usize, + worker_to_id: DashMap, u32, FxBuildHasher>, + registry: Mutex, + /// Blocks held per worker, one cache line each: the lanes write these on every event. + worker_blocks: Box<[CachePadded]>, + /// Per-worker signed contributions to the distinct-block count, summed on demand. + distinct_blocks: Box<[CachePadded]>, + /// Stored blocks whose engine hash was not the one the index carries. + engine_conflicts: AtomicUsize, + /// Stores whose blocks carried the engine hashes the index holds for other content. + landing_mismatches: AtomicUsize, + /// Blocks a store moved to another place for their worker (the old membership released). + /// Coverage words a lookup has to read: enough for the highest worker id interned so far + /// (ids are handed out densely and reused), at most `words`. An index sized for a thousand + /// workers that serves eight reads one word per run instead of sixteen. + live_words: AtomicUsize, + moved_hashes: AtomicUsize, + /// Splits by cause and runs unlinked, always counted (one relaxed increment per split or + /// death): the fragmentation figures a churn harness and the gateway's gauges read. + splits_branch: AtomicUsize, + splits_hole: AtomicUsize, + splits_mid_run: AtomicUsize, + splits_prefix_holders: AtomicUsize, + runs_died: AtomicUsize, +} + +impl Default for ChainIndex { + fn default() -> Self { + Self::new() + } +} + +/// How a worker holds a run. +enum Holding { + Full, + /// The first `cutoff` blocks only. + Partial(usize), + None, +} + +/// Outcome of one attempt to walk a store into the tree. +enum Walk { + Done, + /// The walk met a run another lane unlinked meanwhile; start over from the parent block. + Restart, + /// The parent block is not held after all (its run died since the map was written). + NoParent, +} + +/// What the lock-free look at a run during a store walk decided. +enum Plan { + /// Nothing changes here: `matched` blocks are already held; carry on after them. With them, + /// how many of the matched blocks carry an engine hash that is not the one the index holds. + Skip(usize, usize), + /// Continue in the child `(id, generation)` from its first block. + Descend(u32, u32), + /// Open a new child for the blocks in hand at this offset of the run (its end, or a + /// divergence inside it), without the run's lock if the table has room. + InsertAt(usize), + /// The run must change: take its lock and re-read it. + Lock, +} + +/// Outcome of claiming a child-table slot. +enum Claim { + Inserted, + /// Another writer already linked a child with this head. + Exists(u32, u32), + /// No free slot (or no table): grow under the lock. + Full, + /// The run changed since the plan was made (a split moved its end): plan again. + Changed, +} + +/// What [`ChainIndex::store_in_run`] found. +enum InRun<'b> { + /// Every block is placed. + Done, + /// The blocks still to place start right after the run's last block. + Continue(&'b [StoredBlock]), + /// The blocks still to place diverge from the run at this offset: they continue it from + /// there as a child. + ContinueAt(&'b [StoredBlock], usize), + /// Carry on in this run (id, generation) from its first block. + MoveTo(u32, u32), +} + +/// Blocks a store walk placed: `count` blocks from `start` in the event sit in `run` from +/// `offset`. Recorded under the run lock, written into the lane map after it. +struct Placed { + run: u32, + offset: u32, + start: usize, + count: usize, + /// Blocks of the range whose engine hash differs from the one the index carries for the + /// block (counted in the stats; the worker's own hash keys its lane map either way). + conflicts: usize, +} + +/// How many of the stored blocks carry an engine hash that is not the one the index holds +/// (`engines` runs parallel to `stored`); zero in a fleet whose engines agree. +fn engine_conflicts(stored: &[StoredBlock], engines: &[AtomicU64]) -> usize { + stored + .iter() + .zip(engines) + .filter(|(block, slot)| block.seq_hash.0 != slot.load(Ordering::Relaxed)) + .count() +} + +/// A batch of a worker's offsets in one run, with the generation the run must still have. +struct Removal { + run: u32, + generation: Option, + offsets: Vec, +} + +impl ChainIndex { + /// An index for up to 256 workers. + pub fn new() -> Self { + Self::with_max_workers(256) + } + + /// An index for up to `max_workers` interned workers (at most 1024); coverage costs one bit + /// per worker per run, rounded up to 64. + pub fn with_max_workers(max_workers: usize) -> Self { + let max_workers = max_workers.clamp(1, MAX_WORDS * 64); + let words = max_workers.div_ceil(64); + let slab = RunSlab::new(words); + // The root is run 0: position 0's parent, no blocks, no coverage. + let root = slab.alloc( + 0, + ROOT, + Window { + block: NONE, + base: 0, + engine: NONE, + len: 0, + children: NONE, + partials: NONE, + forwards: NONE, + }, + ); + debug_assert_eq!(root, ROOT); + Self { + slab, + arena: WordArena::new(), + words, + max_workers, + worker_to_id: DashMap::with_hasher(FxBuildHasher), + registry: Mutex::new(WorkerRegistry::default()), + worker_blocks: (0..max_workers) + .map(|_| CachePadded::new(AtomicUsize::new(0))) + .collect(), + distinct_blocks: (0..max_workers) + .map(|_| CachePadded::new(AtomicIsize::new(0))) + .collect(), + engine_conflicts: AtomicUsize::new(0), + landing_mismatches: AtomicUsize::new(0), + live_words: AtomicUsize::new(0), + moved_hashes: AtomicUsize::new(0), + splits_branch: AtomicUsize::new(0), + splits_hole: AtomicUsize::new(0), + splits_mid_run: AtomicUsize::new(0), + splits_prefix_holders: AtomicUsize::new(0), + runs_died: AtomicUsize::new(0), + } + } + + /// Lock a run's writer state. + #[inline] + fn lock_run(&self, run_id: u32) -> MutexGuard<'_, RunMeta> { + self.slab.run(run_id).meta.lock() + } + + /// Count a split by its cause: a hole a removal left (`Own`), or a store under a parent the + /// worker did not hold up to (`Shared`). + #[inline] + fn count_split(&self, cause: SplitCause) { + let counter = match cause { + SplitCause::Hole => &self.splits_hole, + SplitCause::MidRunStore => &self.splits_mid_run, + }; + counter.fetch_add(1, Ordering::Relaxed); + } + + /// Intern a worker name; the same name maps to the same id until the worker is removed. + /// + /// One writer per id: an id's holdings are the lane map of whoever interned it. A second + /// subscription that interns a name before the first has called + /// [`remove_worker`](Self::remove_worker) receives the first one's id, and the first one's + /// removal then takes the second one's blocks out of the runs they share and frees the slot + /// under it, where a third name may be interned. A caller replacing a worker under the same + /// name finishes the removal before interning again, or keeps the id and empties it with + /// [`apply_cleared`](Self::apply_cleared) instead; the index cannot tell two writers apart. + pub fn intern_worker(&self, worker: &str) -> Result { + if let Some(entry) = self.worker_to_id.get(worker) { + return Ok(*entry.value()); + } + let name: Arc = Arc::from(worker); + match self.worker_to_id.entry(name.clone()) { + Entry::Occupied(entry) => Ok(*entry.get()), + Entry::Vacant(entry) => { + let mut registry = self.registry.lock(); + let id = match registry.free.pop() { + Some(id) => id, + None if registry.names.len() < self.max_workers => { + registry.names.push(None); + let id = (registry.names.len() - 1) as u32; + // Published before the id can hold anything: the worker's first store + // comes after this returns. + self.live_words + .fetch_max(id as usize / 64 + 1, Ordering::Release); + id + } + None => return Err(WorkerIdExhausted), + }; + registry.names[id as usize] = Some(name); + entry.insert(id); + Ok(id) + } + } + } + + /// Hand a removed worker's slot back once nothing refers to it any more. + fn release_worker(&self, worker: u32) { + // Lock order is name map shard, then registry (as in `intern_worker`): never hold the + // registry while touching the name map, or an intern and a release deadlock each other. + let name = { + let mut registry = self.registry.lock(); + registry + .names + .get_mut(worker as usize) + .and_then(Option::take) + }; + let Some(name) = name else { + return; + }; + self.worker_to_id.remove(&*name); + self.registry.lock().free.push(worker); + } + + pub fn worker_id(&self, worker: &str) -> Option { + self.worker_to_id.get(worker).map(|entry| *entry.value()) + } + + /// Whether no worker holds any block: the root has no live child, read under the root's + /// version. O(1), for a caller that asks before every lookup (`current_size` reads a line + /// per worker slot). + pub fn is_empty(&self) -> bool { + let root = self.slab.run(ROOT); + loop { + let (window, version) = root.snapshot(); + let empty = window.children == NONE || self.arena.table_live(window.children) == 0; + if root.confirm(version) { + return empty; + } + } + } + + /// Blocks held across all workers (a block two workers hold counts twice). + pub fn current_size(&self) -> usize { + self.worker_blocks + .iter() + .map(|count| count.load(Ordering::Relaxed)) + .sum() + } + + /// Distinct blocks held by at least one worker. + pub fn entry_count(&self) -> usize { + let total: isize = self + .distinct_blocks + .iter() + .map(|count| count.load(Ordering::Relaxed)) + .sum(); + total.max(0) as usize + } + + pub fn worker_block_count(&self, worker: u32) -> usize { + self.worker_blocks + .get(worker as usize) + .map_or(0, |count| count.load(Ordering::Relaxed)) + } + + fn credit(&self, worker: u32, blocks: usize) { + if blocks == 0 { + return; + } + self.worker_blocks[worker as usize].fetch_add(blocks, Ordering::Relaxed); + } + + fn debit(&self, worker: u32, blocks: usize) { + if blocks == 0 { + return; + } + self.worker_blocks[worker as usize].fetch_sub(blocks, Ordering::Relaxed); + } + + /// `worker` became the first holder of `blocks` distinct blocks. + fn distinct_add(&self, worker: u32, blocks: usize) { + if blocks != 0 { + self.distinct_blocks[worker as usize].fetch_add(blocks as isize, Ordering::Relaxed); + } + } + + /// `worker` was the last holder of `blocks` distinct blocks. + fn distinct_sub(&self, worker: u32, blocks: usize) { + if blocks != 0 { + self.distinct_blocks[worker as usize].fetch_sub(blocks as isize, Ordering::Relaxed); + } + } + + /// The hash at `offset` of a run, from the writer's side (the run is locked). + #[inline] + fn hash_at(&self, run: &Run, offset: usize) -> u64 { + let data = run.block.load(Ordering::Relaxed) + run.base.load(Ordering::Relaxed); + self.arena + .word(data + offset as u32) + .load(Ordering::Relaxed) + } + + /// The child-table key of a child continuing a run from `offset` with the content hash + /// `head`: a child may hang off any offset of its parent (a divergence inside a run does not + /// split the run), so the key carries the offset beside the hash. Never zero, which a table + /// slot reads as empty. + #[inline] + fn child_key(offset: usize, head: u64) -> u64 { + let mixed = head + ^ (offset as u64 + 1) + .wrapping_mul(0x9E37_79B9_7F4A_7C15) + .rotate_left(23); + if mixed == 0 { + 1 + } else { + mixed + } + } + + /// The content hash of a run's first block, read under its version (the run is not locked). + fn child_head(&self, run_id: u32) -> u64 { + let run = self.slab.run(run_id); + loop { + let (window, version) = run.snapshot(); + let head = self + .arena + .word(window.block + window.base) + .load(Ordering::Relaxed); + if run.confirm(version) { + return head; + } + } + } + + /// One step along a block's forwarding chain, read under the run's version: the next place + /// and the generation it must have, the place standing (with the run's generation), or the + /// block gone (a dead or reused run, a `GONE` suffix). A split publishes its shorter length + /// and its forwarding record under one lock but in two steps, so a reader that meets an + /// offset at or past the length without a record waits for the lock once and reads again; + /// after that the record is there, or the reference was stale. + fn hop(&self, at: BlockRef, expected: Option) -> Hop { + let mut waited = false; + loop { + if at.run == GONE { + return Hop::Gone; + } + let run = self.slab.run(at.run); + let (window, version) = run.snapshot(); + let generation = (version >> 32) as u32; + if expected.is_some_and(|wanted| wanted != generation) + || (at.run != ROOT && window.len == 0) + { + return Hop::Gone; + } + let hop = self.arena.forwards_find(window.forwards, at.offset); + if !run.confirm(version) { + continue; + } + match hop { + Some((next, next_generation)) => return Hop::Next(next, next_generation), + None if at.run != ROOT && at.offset >= window.len => { + if waited { + return Hop::Gone; + } + drop(run.meta.lock()); + waited = true; + } + None => return Hop::Here(generation), + } + } + } + + /// Where a block recorded at `at` lives now, with the generation of the run: the end of its + /// forwarding chain. `None` when nobody holds it any more. + fn resolve(&self, mut at: BlockRef) -> Option<(BlockRef, u32)> { + let mut expected: Option = None; + loop { + match self.hop(at, expected) { + Hop::Next(next, generation) => { + at = next; + expected = Some(generation); + } + Hop::Here(generation) => return Some((at, generation)), + Hop::Gone => return None, + } + } + } + + /// Whether `to` is `from` or a place `from` has forwarded to since: the places one block has + /// had. Forwarding records only accumulate while their runs live (a prefix outlives its + /// suffix, being its parent), so the answer does not depend on when it is read, unlike the + /// ends of two chains resolved one after the other. + fn forwards_to(&self, from: BlockRef, to: BlockRef) -> bool { + let mut at = from; + let mut expected: Option = None; + loop { + if at == to { + return true; + } + match self.hop(at, expected) { + Hop::Next(next, generation) => { + at = next; + expected = Some(generation); + } + Hop::Here(_) | Hop::Gone => return false, + } + } + } + + /// Whether `worker`'s lane map holds the block with engine hash `key`. + pub fn is_held(&self, map: &ChainBlockMap, key: SequenceHash) -> bool { + let _ = self; + map.contains_key(key) + } + + /// Add `child` to the run's table (the run is locked): claims a slot like a lock-free + /// inserter, growing the table under a version step that first drains inserters in flight. + /// `Some` when another writer linked a child with this head meanwhile: the caller descends + /// into that one instead. + fn link_child(&self, run: &Run, offset: usize, head: u64, child: u32) -> Option<(u32, u32)> { + let generation = self.slab.run(child).generation(); + let key = Self::child_key(offset, head); + loop { + let table = run.children.load(Ordering::Acquire); + match self.arena.table_claim(table, key, child, generation) { + Claim::Inserted => return None, + Claim::Exists(other, other_generation) => return Some((other, other_generation)), + Claim::Full | Claim::Changed => {} + } + run.begin_update(); + run.wait_inflight(); + let table = run.children.load(Ordering::Relaxed); + let grown = if table == NONE { + self.arena.alloc_table(MIN_TABLE_SLOTS) + } else { + self.arena.table_grown(table) + }; + run.children.store(grown, Ordering::Relaxed); + run.end_update(); + self.arena.free_table(table); + } + } + + /// Take `child` out of the run's table (the run is locked); an emptied table goes away once + /// no insert is in flight on it. + fn unlink_child(&self, run: &Run, child: u32) { + let table = run.children.load(Ordering::Relaxed); + if table == NONE { + return; + } + let live = self.arena.table_take(table, child); + if live == 0 { + run.begin_update(); + run.wait_inflight(); + if self.arena.word(table + 1).load(Ordering::Relaxed) == 0 { + run.children.store(NONE, Ordering::Relaxed); + run.end_update(); + self.arena.free_table(table); + } else { + run.end_update(); + } + } else if self.arena.table_dead(table) > self.arena.table_slots(table) / 2 { + // Children come and go at every offset of a long-lived run (a decode tail per prompt + // end): a table more than half tombstones is rebuilt from its live entries, so a + // probe never walks the dead keys of a thousand finished requests. + run.begin_update(); + run.wait_inflight(); + let fresh = self.arena.table_from(&self.arena.table_entries(table)); + run.children.store(fresh, Ordering::Relaxed); + run.end_update(); + self.arena.free_table(table); + } + } + + /// A lock-free attempt to link `child` (prepared, unpublished) under `run` for the blocks + /// after its last one, as seen in the snapshot with `planned` version: holds `inflight` + /// across the claim so a split or table growth cannot move the table under the insert, and + /// gives up with `Claim::Changed` if the run moved on since the plan (its end is elsewhere + /// now). `Claim::Full` means the locked path must do it. + fn insert_child( + &self, + run_id: u32, + child: u32, + offset: usize, + head: u64, + planned: u64, + ) -> Claim { + let run = self.slab.run(run_id); + loop { + run.inflight.fetch_add(1, Ordering::SeqCst); + let version = run.version.load(Ordering::SeqCst); + if version & 1 == 1 { + run.inflight.fetch_sub(1, Ordering::SeqCst); + std::hint::spin_loop(); + continue; + } + if version != planned { + run.inflight.fetch_sub(1, Ordering::SeqCst); + return Claim::Changed; + } + let table = run.children.load(Ordering::Acquire); + let len = run.len(); + if run.version.load(Ordering::SeqCst) != version { + run.inflight.fetch_sub(1, Ordering::SeqCst); + continue; + } + if offset > len { + run.inflight.fetch_sub(1, Ordering::SeqCst); + return Claim::Changed; + } + // The child continues this run from `offset`: its end, or a divergence inside it. + let new_run = self.slab.run(child); + new_run + .start + .store((run.start() + offset) as u32, Ordering::Relaxed); + new_run.parent.store(run_id, Ordering::Release); + let claim = self.arena.table_claim( + table, + Self::child_key(offset, head), + child, + new_run.generation(), + ); + run.inflight.fetch_sub(1, Ordering::SeqCst); + return claim; + } + } + + /// A run prepared for a store that another lane beat to the slot: back to the slab. + fn discard_run(&self, run_id: u32, worker: u32, freed: &mut Vec) { + let mut meta = self.slab.run(run_id).meta.lock(); + clear(self.slab.coverage(run_id), worker); + self.kill(run_id, &mut meta, freed); + } + + /// Retire a run that is unlinked (locked by the caller): its array reference goes, its + /// window empties, and its id is queued for reuse once the caller has dropped the lock. + fn kill(&self, run_id: u32, meta: &mut RunMeta, freed: &mut Vec) { + let run = self.slab.run(run_id); + let block = run.block.load(Ordering::Relaxed); + let engine = run.engine.load(Ordering::Relaxed); + let partials = run.partials.load(Ordering::Relaxed); + let forwards = run.forwards.load(Ordering::Relaxed); + run.begin_update(); + run.block.store(NONE, Ordering::Relaxed); + run.base.store(0, Ordering::Relaxed); + run.engine.store(NONE, Ordering::Relaxed); + run.len.store(0, Ordering::Relaxed); + run.children.store(NONE, Ordering::Relaxed); + run.partials.store(NONE, Ordering::Relaxed); + run.forwards.store(NONE, Ordering::Relaxed); + run.end_update(); + self.arena.array_release(block); + self.arena.array_release(engine); + self.arena.free_partials(partials); + self.arena.free_table(forwards); + meta.dead = true; + self.runs_died.fetch_add(1, Ordering::Relaxed); + freed.push(run_id); + } + + /// Dead ids go back to the slab only after their locks are released, so a thread holding a + /// live run's lock and reviving a dead id never waits on a thread that holds the dead id's + /// lock and wants the live run. + fn recycle(&self, freed: &mut Vec) { + for id in freed.drain(..) { + self.slab.free.push(id); + } + } + + /// How `worker` holds a run: all of it, a prefix of it, or nothing. + fn holding(&self, run_id: u32, worker: u32) -> Holding { + if has(self.slab.coverage(run_id), worker) { + return Holding::Full; + } + let table = self.slab.run(run_id).partials.load(Ordering::Relaxed); + match self.arena.partial_find(table, worker) { + Some((_, cutoff)) => Holding::Partial(cutoff as usize), + None => Holding::None, + } + } + + /// Blocks of a run that `worker` holds. + fn held_by(&self, run_id: u32, worker: u32) -> usize { + match self.holding(run_id, worker) { + Holding::Full => self.slab.run(run_id).len(), + Holding::Partial(cutoff) => cutoff, + Holding::None => 0, + } + } + + /// Blocks of a run held by at least one worker: all of them while anybody holds the whole + /// run, otherwise the longest partial prefix. + fn held_len(&self, run_id: u32) -> usize { + let run = self.slab.run(run_id); + if coverage_is_empty(self.slab.coverage(run_id)) { + self.arena.partial_max(run.partials.load(Ordering::Relaxed)) + } else { + run.len() + } + } + + fn has_holders(&self, run_id: u32) -> bool { + !coverage_is_empty(self.slab.coverage(run_id)) + || self + .arena + .partials_live(self.slab.run(run_id).partials.load(Ordering::Relaxed)) + > 0 + } + + /// Distinct-block accounting around a change to `run_id` (locked by the caller): the delta + /// of blocks held by anybody, attributed to `worker`. A split changes nothing in total (the + /// prefix and the suffix hold between them what the run held), so callers take `before` + /// after any split; the suffix is published by then and other lanes account for their own + /// changes to it. + fn settle_distinct(&self, worker: u32, before: usize, run_id: u32) { + let after = self.held_len(run_id); + if after > before { + self.distinct_add(worker, after - before); + } else { + self.distinct_sub(worker, before - after); + } + } + + /// Make `worker` hold exactly `[0, cutoff)` of the run (the run is locked): the whole run when + /// `cutoff` reaches its length, nothing when 0. Readers treat the coverage bit as the truth + /// when both forms are visible, so a worker gains its bit before its partial entry goes and + /// gains a partial entry before its bit goes. + fn set_holding(&self, run_id: u32, worker: u32, cutoff: usize) { + let run = self.slab.run(run_id); + let coverage = self.slab.coverage(run_id); + let len = run.len(); + let was_full = has(coverage, worker); + let table = run.partials.load(Ordering::Relaxed); + let entry = if was_full { + None + } else { + self.arena.partial_find(table, worker) + }; + if cutoff >= len { + if !was_full { + set(coverage, worker); + } + if let Some((slot, _)) = entry { + self.drop_partial(run, table, slot); + } + } else if cutoff == 0 { + if was_full { + clear(coverage, worker); + } + if let Some((slot, _)) = entry { + self.drop_partial(run, table, slot); + } + } else { + match entry { + Some((slot, old)) => { + if old as usize != cutoff { + self.arena.partial_set(table, slot, worker, cutoff as u32); + } + } + None if was_full => { + // A reader that saw neither the entry nor the bit would score nothing for a + // worker that holds a prefix: keep the two writes inside one version step. + run.begin_update(); + self.add_partial(run, table, worker, cutoff as u32); + clear(coverage, worker); + run.end_update(); + } + None => self.add_partial(run, table, worker, cutoff as u32), + } + } + } + + /// Append a partial entry, growing (and republishing) the table when it is full. + fn add_partial(&self, run: &Run, table: u32, worker: u32, cutoff: u32) { + if self.arena.partial_put(table, worker, cutoff) { + return; + } + let grown = if table == NONE { + self.arena.alloc_partials(MIN_TABLE_SLOTS) + } else { + self.arena.partials_grown(table, 1) + }; + let placed = self.arena.partial_put(grown, worker, cutoff); + debug_assert!(placed); + let nested = run.version.load(Ordering::Relaxed) & 1 == 1; + if !nested { + run.begin_update(); + } + run.partials.store(grown, Ordering::Relaxed); + if !nested { + run.end_update(); + } + self.arena.free_partials(table); + } + + /// Tombstone a partial entry; an emptied table goes away. + fn drop_partial(&self, run: &Run, table: u32, slot: usize) { + if self.arena.partial_remove(table, slot) == 0 { + run.begin_update(); + run.partials.store(NONE, Ordering::Relaxed); + run.end_update(); + self.arena.free_partials(table); + } + } + + /// Split `run` (locked by the caller, whose guard is `_meta`) at `at`: the run keeps `[0, at)`; a new suffix run takes + /// `[at, len)` on the same hash array, with the run's children, its full holders and the + /// partial holders reaching past `at`; partial holders reaching `at` become full holders of + /// the prefix. A suffix nobody would hold is not created when the run has no children: the + /// forwarding record says those blocks are gone. No worker's holdings change in total. + /// Split the run (locked by the caller) at the median cutoff of its prefix holders when more + /// than `PARTIAL_CAP` of them hold a prefix of it: the holders at or past the median become + /// whole holders of the prefix, the rest keep their entries on one side or the other. + fn cap_prefix_holders(&self, run_id: u32, meta: &mut RunMeta) { + let run = self.slab.run(run_id); + let table = run.partials.load(Ordering::Relaxed); + if table == NONE || self.arena.partials_live(table) <= PARTIAL_CAP { + return; + } + let mut cutoffs: Vec = self + .arena + .partial_entries(table) + .iter() + .map(|&(_, cutoff)| cutoff) + .collect(); + if cutoffs.len() <= PARTIAL_CAP { + return; + } + cutoffs.sort_unstable(); + let at = cutoffs[cutoffs.len() / 2] as usize; + if at == 0 || at >= run.len() { + return; + } + self.splits_prefix_holders.fetch_add(1, Ordering::Relaxed); + self.split_locked(run_id, meta, at); + } + + fn split_locked(&self, run_id: u32, _meta: &mut RunMeta, at: usize) -> u32 { + let run = self.slab.run(run_id); + let coverage = self.slab.coverage(run_id); + let len = run.len(); + debug_assert!(at > 0 && at < len, "split inside the run: 0 < {at} < {len}"); + let block = run.block.load(Ordering::Relaxed); + let engine = run.engine.load(Ordering::Relaxed); + let base = run.base.load(Ordering::Relaxed); + let children = run.children.load(Ordering::Relaxed); + let partials = self + .arena + .partial_entries(run.partials.load(Ordering::Relaxed)); + let beyond: Vec<(u32, u32)> = partials + .iter() + .filter(|(_, cutoff)| *cutoff as usize > at) + .map(|&(worker, cutoff)| (worker, cutoff - at as u32)) + .collect(); + let suffix_held = !coverage_is_empty(coverage) || !beyond.is_empty(); + // Children hang off any offset of the run: those past `at` move to the suffix, keyed by + // their offset within it; the rest stay. Decided under the version step, after + // the inserts in flight have landed, so none is missed. + run.begin_update(); + run.wait_inflight(); + let start = run.start(); + let mut kept: Vec<(u64, u32, u32)> = Vec::new(); + let mut moved: Vec<(u64, u32, u32)> = Vec::new(); + for (key, child, generation) in self.arena.table_entries(children) { + let child_offset = self.slab.run(child).start().saturating_sub(start); + // A child at `at` itself stays an end child of the prefix (a walk never descends + // from offset zero of a run, so the suffix must not carry one there). + if child_offset <= at { + kept.push((key, child, generation)); + } else { + moved.push(( + Self::child_key(child_offset - at, self.child_head(child)), + child, + generation, + )); + } + } + let suffix_id = if !suffix_held && moved.is_empty() { + run.len.store(at as u32, Ordering::Relaxed); + run.end_update(); + GONE + } else { + self.arena.array_retain(block); + self.arena.array_retain(engine); + let suffix_partials = self.arena.partials_from(&beyond); + let suffix_children = self.arena.table_from(&moved); + let suffix_id = self.slab.alloc( + start + at, + run_id, + Window { + block, + base: base + at as u32, + engine, + len: (len - at) as u32, + children: suffix_children, + partials: suffix_partials, + forwards: NONE, + }, + ); + let suffix = self.slab.run(suffix_id); + for (slot, word) in self.slab.coverage(suffix_id).iter().zip(coverage) { + slot.store(word.load(Ordering::Relaxed), Ordering::Relaxed); + } + for &(_, child, _) in &moved { + self.slab + .run(child) + .parent + .store(suffix_id, Ordering::Release); + } + kept.push(( + Self::child_key(at, self.hash_at(run, at)), + suffix_id, + suffix.generation(), + )); + let table = self.arena.table_from(&kept); + run.len.store(at as u32, Ordering::Relaxed); + run.children.store(table, Ordering::Relaxed); + run.end_update(); + self.arena.free_table(children); + suffix_id + }; + // The prefix is `[0, at)` now: a partial holder that reached it holds all of it. Done + // after the truncation so no reader sees a bit for the whole old run. + for (worker, cutoff) in partials { + if cutoff as usize >= at { + self.set_holding(run_id, worker, at); + } + } + let generation = if suffix_id == GONE { + 0 + } else { + self.slab.run(suffix_id).generation() + }; + self.add_forward(run, at as u32, suffix_id, generation); + suffix_id + } + + /// Record a split on the run (locked by the caller); a full table is replaced under a version + /// step so a lock-free reader never follows a recycled one. + fn add_forward(&self, run: &Run, at: u32, suffix: u32, generation: u32) { + let table = run.forwards.load(Ordering::Relaxed); + if self.arena.forwards_push(table, at, suffix, generation) { + return; + } + let grown = if table == NONE { + self.arena.alloc_forwards(MIN_TABLE_SLOTS) + } else { + self.arena.forwards_grown(table) + }; + let placed = self.arena.forwards_push(grown, at, suffix, generation); + debug_assert!(placed); + run.begin_update(); + run.forwards.store(grown, Ordering::Relaxed); + run.end_update(); + self.arena.free_table(table); + } + + /// Unlink `run` (locked by the caller, known to be an uncovered leaf) from its parent and + /// retire it, then the parent if that leaves it an uncovered leaf too. + fn unlink_locked(&self, run_id: u32, meta: &mut RunMeta, freed: &mut Vec) { + if run_id == ROOT || meta.dead { + return; + } + let run = self.slab.run(run_id); + loop { + let parent_id = run.parent.load(Ordering::Acquire); + let parent = self.slab.run(parent_id); + // Child-then-parent is the only order in which two run locks are ever held. A split + // of the parent may have re-parented this run while we waited; check and retry. + let mut parent_meta = parent.meta.lock(); + if run.parent.load(Ordering::Acquire) != parent_id { + continue; + } + self.unlink_child(parent, run_id); + self.kill(run_id, meta, freed); + if parent_id != ROOT + && parent.children.load(Ordering::Relaxed) == NONE + && !self.has_holders(parent_id) + { + self.unlink_locked(parent_id, &mut parent_meta, freed); + } + return; + } + } + + /// Store `blocks` for `worker` after `parent` (position 0 when `None`). + pub fn apply_stored( + &self, + worker: u32, + blocks: &[StoredBlock], + parent: Option, + map: &mut ChainBlockMap, + ) -> Result<(), ApplyError> { + if blocks.is_empty() { + return Ok(()); + } + let origin = match parent { + None => None, + Some(hash) => { + if map.is_empty() { + return Err(ApplyError::WorkerNotTracked); + } + match map.get(hash) { + Some(at) => Some((hash, at)), + None => return Err(ApplyError::ParentBlockNotFound), + } + } + }; + loop { + match self.store_walk(worker, blocks, origin, map) { + Walk::Done => return Ok(()), + Walk::Restart => {} + Walk::NoParent => return Err(ApplyError::ParentBlockNotFound), + } + } + } + + /// One attempt to place a store, from the parent block (or the root) down the tree. Map + /// entries for the placed blocks are written after the run locks are released: the lane map + /// is private, and a split meanwhile is covered by the forwarding records. + fn store_walk( + &self, + worker: u32, + blocks: &[StoredBlock], + origin: Option<(SequenceHash, BlockRef)>, + map: &mut ChainBlockMap, + ) -> Walk { + let mut pending: Vec = Vec::new(); + let outcome = self.store_walk_locked(worker, blocks, origin, map, &mut pending); + self.write_placements(worker, blocks, pending, map); + outcome + } + + /// The lane map writes of one store, after the run locks are released. A block the engine + /// names by a hash this worker already holds elsewhere (a store that moves the hash, as a + /// store without its parent followed by the whole chain does) keeps one place per hash: the + /// old membership is released. A re-store of the same block at the same place is not a + /// move, and "the same place" is read through the forwarding records: the placement was + /// recorded under the run's lock, and splits by other lanes land freely between that and + /// this write, so the recorded place and the map's old entry may both be behind the block's + /// current place by any number of splits. A block that did not move has its new place on + /// the forwarding chain from its old one, and that stays true however many splits follow + /// (`forwards_to`); comparing the two places resolved to their ends instead raced with a + /// split between the two reads, took the block for a moved hash, and released the worker's + /// holding at its real place while the map kept the entry (lookups then scored the worker + /// short, and a store under the block as parent found no parent). + fn write_placements( + &self, + worker: u32, + blocks: &[StoredBlock], + pending: Vec, + map: &mut ChainBlockMap, + ) { + let mut moved: Vec = Vec::new(); + for placed in pending { + let first = BlockRef { + run: placed.run, + offset: placed.offset, + }; + let range = &blocks[placed.start..placed.start + placed.count]; + if placed.conflicts > 0 { + self.engine_conflicts + .fetch_add(placed.conflicts, Ordering::Relaxed); + } + map.insert_run( + range.iter().map(|stored| stored.seq_hash), + first, + |old, new| { + if !self.forwards_to(old, new) { + moved.push(old); + } + }, + ); + } + if !moved.is_empty() { + self.moved_hashes.fetch_add(moved.len(), Ordering::Relaxed); + let work = group_by_run(moved); + self.apply_removals(worker, work); + } + } + + fn store_walk_locked( + &self, + worker: u32, + blocks: &[StoredBlock], + origin: Option<(SequenceHash, BlockRef)>, + map: &mut ChainBlockMap, + pending: &mut Vec, + ) -> Walk { + let (mut run_id, mut offset, mut expected) = match origin { + None => (ROOT, 0usize, self.slab.run(ROOT).generation()), + Some((hash, at)) => { + let Some((at, generation)) = self.resolve(at) else { + map.remove(hash); + return Walk::NoParent; + }; + map.insert(hash, at); + (at.run, at.offset as usize + 1, generation) + } + }; + let mut remaining = blocks; + // Set when a lock-free insert found the child table full: the next look at the same run + // takes its lock, whose path grows the table. + let mut force_lock = false; + loop { + // Look at the run under its version first: a run this worker already holds up to + // the blocks in hand, or a run whose child continues them, is passed without its + // lock. Only a run that must change is locked, and re-read under the lock. + let run = self.slab.run(run_id); + let (window, version) = run.snapshot(); + if (version >> 32) as u32 != expected || (run_id != ROOT && window.len == 0) { + return Walk::Restart; + } + let len = window.len as usize; + if offset > len { + // A split moved the blocks after the parent into a suffix: resolve again. + return Walk::Restart; + } + let block_start = blocks.len() - remaining.len(); + let plan = if force_lock { + force_lock = false; + Plan::Lock + } else if offset < len { + let held = self.held_in(run_id, &window, worker); + if held < offset { + Plan::Lock + } else { + let data = window.block + window.base + offset as u32; + let hashes = self.arena.words(data, len - offset); + let engines = self + .arena + .words(window.engine + window.base + offset as u32, len - offset); + let (matched, conflicts) = self.match_run(hashes, engines, remaining, false); + if matched > 0 && offset + matched <= held { + // Blocks already held: step past them. A divergence after them is the + // next round's business, at the offset where it starts. + Plan::Skip(matched, conflicts) + } else if matched == 0 && offset > 0 { + // The blocks in hand leave the run's content right here: they continue + // it as a child hanging off this offset, found or opened without the + // lock. The run itself does not change. + self.plan_child(&window, offset, remaining[0].content_hash.0) + } else { + Plan::Lock + } + } + } else { + self.plan_child(&window, len, remaining[0].content_hash.0) + }; + if !run.confirm(version) { + continue; + } + match plan { + Plan::InsertAt(at) => { + let head = remaining[0].content_hash.0; + let contents: Vec = remaining + .iter() + .map(|stored| stored.content_hash.0) + .collect(); + let engines: Vec = + remaining.iter().map(|stored| stored.seq_hash.0).collect(); + let block = self + .arena + .alloc_array(&contents, capacity_for(contents.len())); + let engine = self + .arena + .alloc_array(&engines, self.arena.array_capacity(block)); + let new_id = self.slab.alloc( + 0, + run_id, + Window { + block, + base: 0, + engine, + len: contents.len() as u32, + children: NONE, + partials: NONE, + forwards: NONE, + }, + ); + set(self.slab.coverage(new_id), worker); + match self.insert_child(run_id, new_id, at, head, version) { + Claim::Inserted => { + pending.push(Placed { + run: new_id, + offset: 0, + start: block_start, + count: remaining.len(), + conflicts: 0, + }); + self.credit(worker, contents.len()); + self.distinct_add(worker, contents.len()); + return Walk::Done; + } + Claim::Exists(child, generation) => { + let mut freed = Vec::new(); + self.discard_run(new_id, worker, &mut freed); + self.recycle(&mut freed); + run_id = child; + expected = generation; + offset = 0; + } + Claim::Full => { + let mut freed = Vec::new(); + self.discard_run(new_id, worker, &mut freed); + self.recycle(&mut freed); + force_lock = true; + } + Claim::Changed => { + let mut freed = Vec::new(); + self.discard_run(new_id, worker, &mut freed); + self.recycle(&mut freed); + return Walk::Restart; + } + } + } + Plan::Skip(matched, conflicts) => { + pending.push(Placed { + run: run_id, + offset: offset as u32, + start: block_start, + count: matched, + conflicts, + }); + if matched == remaining.len() { + return Walk::Done; + } + remaining = &remaining[matched..]; + offset += matched; + } + Plan::Descend(child, generation) => { + run_id = child; + expected = generation; + offset = 0; + } + Plan::Lock => { + let mut meta = self.lock_run(run_id); + if meta.dead || run.generation() != expected || offset > run.len() { + return Walk::Restart; + } + let next = match self.store_in_run( + worker, + run_id, + &mut meta, + offset, + remaining, + block_start, + pending, + ) { + InRun::Done => return Walk::Done, + InRun::Continue(rest) => { + remaining = rest; + let end = run.len(); + self.store_at( + worker, + run_id, + &meta, + end, + remaining, + blocks.len() - remaining.len(), + pending, + ) + } + InRun::ContinueAt(rest, at) => { + remaining = rest; + self.store_at( + worker, + run_id, + &meta, + at, + remaining, + blocks.len() - remaining.len(), + pending, + ) + } + InRun::MoveTo(child, generation) => Some((child, generation)), + }; + drop(meta); + match next { + Some((child, generation)) => { + run_id = child; + expected = generation; + offset = 0; + } + None => return Walk::Done, + } + } + } + } + } + + /// How far `remaining` continues a run whose hashes from the position in hand are `contents` + /// and `engines`: the matched count and how many of the matched blocks carry an engine hash + /// that is not the one the index holds. + /// + /// An engine hash is a chain hash: an engine names a block by a hash of its parent's hash and + /// the block's own content, so two chains that agree at a position agree at every position + /// before it. The match is therefore decided by the engine hash at the end of the window, + /// and on a mismatch by bisection to the first differing position, instead of a content + /// compare per block. Two checks keep it exact for a worker whose engine breaks the + /// assumption, both by falling back to the compare per block: the content hash at the + /// landing (a block named by the hash the index carries but holding other content, which + /// is what the relay's hash check exists to catch; `landing_mismatches` counts it), and the + /// content hash right after the match (an engine that hashes the same content differently + /// matches by content where it cannot by engine hash, and is counted). + /// + /// `count` says whether a landing mismatch is counted: the lock-free look at a run finds it + /// first and then takes the lock, where the locked look finds it again. + fn match_run( + &self, + contents: &[AtomicU64], + engines: &[AtomicU64], + remaining: &[StoredBlock], + count: bool, + ) -> (usize, usize) { + let window = remaining.len().min(contents.len()).min(engines.len()); + if window == 0 { + return (0, 0); + } + let engine_eq = |at: usize| remaining[at].seq_hash.0 == engines[at].load(Ordering::Relaxed); + let content_eq = + |at: usize| remaining[at].content_hash.0 == contents[at].load(Ordering::Relaxed); + let matched = if engine_eq(window - 1) { + window + } else { + // The positions that agree are a prefix: find the first that does not. + let (mut low, mut high) = (0usize, window - 1); + while low < high { + let mid = low + (high - low) / 2; + if engine_eq(mid) { + low = mid + 1; + } else { + high = mid; + } + } + low + }; + let landing_holds = matched == 0 || content_eq(matched - 1); + let same_content_after = matched < window && content_eq(matched); + if landing_holds && !same_content_after { + return (matched, 0); + } + if !landing_holds && count { + self.landing_mismatches.fetch_add(1, Ordering::Relaxed); + } + let matched = remaining + .iter() + .zip(contents) + .take_while(|(stored, slot)| stored.content_hash.0 == slot.load(Ordering::Relaxed)) + .count(); + ( + matched, + engine_conflicts(&remaining[..matched], &engines[..matched]), + ) + } + + /// Blocks of a run `worker` holds, read from a snapshot (no lock). + #[inline] + fn held_in(&self, run_id: u32, window: &Window, worker: u32) -> usize { + if has(self.slab.coverage(run_id), worker) { + window.len as usize + } else { + self.arena + .partial_find(window.partials, worker) + .map_or(0, |(_, cutoff)| cutoff as usize) + } + } + + /// The plan where the blocks in hand leave the run's content, at its end or at a divergence + /// at `offset`: descend into the child that continues the run there, open one without the + /// lock when the run has a table, or lock a leaf (the worker's own to extend, or a shared one + /// without a table yet). + fn plan_child(&self, window: &Window, offset: usize, head: u64) -> Plan { + match self + .arena + .table_find(window.children, Self::child_key(offset, head)) + { + Some((child, generation)) => { + // The child's header is the next line the walk reads: ask for it while the + // version is confirmed. + crate::prefetch::prefetch_read(std::ptr::from_ref(self.slab.run(child))); + Plan::Descend(child, generation) + } + None if window.children != NONE => Plan::InsertAt(offset), + None => Plan::Lock, + } + } + + /// Match `remaining` against the run from `offset`: join the run as far as it matches, record + /// the matched blocks, and say how to go on (a divergence inside the run continues as a + /// child hanging off it; the run is never split for one). + #[expect(clippy::too_many_arguments)] + fn store_in_run<'b>( + &self, + worker: u32, + run_id: u32, + meta: &mut RunMeta, + offset: usize, + remaining: &'b [StoredBlock], + block_start: usize, + pending: &mut Vec, + ) -> InRun<'b> { + let run = self.slab.run(run_id); + let len = run.len(); + if offset >= len { + return InRun::Continue(remaining); + } + let held = self.held_by(run_id, worker); + if held < offset { + // The parent entry pointed past what this worker holds: the engine re-stored under + // a stale parent, or named one position by two engine hashes (the same content under + // the same parent) and removed the one the index filed the position under. Cut here + // and join the suffix; when nobody holds anything past the cut and no child hangs + // there, the split leaves nothing to join (`GONE`) and the run simply ends at the + // cut: the blocks go after it, like any store past a run's end. + self.count_split(SplitCause::MidRunStore); + let suffix = self.split_locked(run_id, meta, offset); + if suffix == GONE { + return InRun::Continue(remaining); + } + return InRun::MoveTo(suffix, self.slab.run(suffix).generation()); + } + let base = run.base.load(Ordering::Relaxed) + offset as u32; + let data = run.block.load(Ordering::Relaxed) + base; + let hashes = self.arena.words(data, len - offset); + let engines = self + .arena + .words(run.engine.load(Ordering::Relaxed) + base, len - offset); + let (matched, conflicts) = self.match_run(hashes, engines, remaining, true); + let available = len - offset; + let reach = offset + matched; + // A divergence inside the run leaves the run whole: the blocks from the divergence on + // continue it as a child hanging off `reach` (see `store_at`). + let diverges = matched < available && matched < remaining.len(); + let before = self.held_len(run_id); + if reach > held { + self.set_holding(run_id, worker, reach); + self.credit(worker, reach - held); + } + self.settle_distinct(worker, before, run_id); + pending.push(Placed { + run: run_id, + offset: offset as u32, + start: block_start, + count: matched, + conflicts, + }); + if matched == remaining.len() { + // The store ends in this run: a safe point to cap its prefix holders (a split here + // moves nothing this store still has to place). + self.cap_prefix_holders(run_id, meta); + return InRun::Done; + } + if diverges { + return InRun::ContinueAt(&remaining[matched..], reach); + } + InRun::Continue(&remaining[matched..]) + } + + /// Place `remaining` after `offset` blocks of the run (locked by the caller): in the child + /// that already continues the run there, in place when the run is this worker's own leaf and + /// `offset` is its end, or in a new child linked at `offset`. `Some` names a child that + /// already existed, for the caller to carry on in. + #[expect(clippy::too_many_arguments)] + fn store_at( + &self, + worker: u32, + run_id: u32, + _meta: &RunMeta, + offset: usize, + remaining: &[StoredBlock], + block_start: usize, + pending: &mut Vec, + ) -> Option<(u32, u32)> { + let run = self.slab.run(run_id); + let coverage = self.slab.coverage(run_id); + let head = remaining[0].content_hash.0; + let children = run.children.load(Ordering::Relaxed); + if let Some(found) = self + .arena + .table_find(children, Self::child_key(offset, head)) + { + return Some(found); + } + let len = run.len(); + debug_assert!( + offset <= len, + "a child hangs off the run: {offset} <= {len}" + ); + let contents: Vec = remaining + .iter() + .map(|stored| stored.content_hash.0) + .collect(); + let engines: Vec = remaining.iter().map(|stored| stored.seq_hash.0).collect(); + let own_leaf = offset == len + && run_id != ROOT + && children == NONE + && run.forwards.load(Ordering::Relaxed) == NONE + && covered_only_by(coverage, worker); + let (target, first) = if own_leaf { + self.append(run, &contents, &engines); + (run_id, len) + } else { + let block = self + .arena + .alloc_array(&contents, capacity_for(contents.len())); + let engine = self + .arena + .alloc_array(&engines, self.arena.array_capacity(block)); + let new_id = self.slab.alloc( + run.start() + offset, + run_id, + Window { + block, + base: 0, + engine, + len: contents.len() as u32, + children: NONE, + partials: NONE, + forwards: NONE, + }, + ); + set(self.slab.coverage(new_id), worker); + if let Some(existing) = self.link_child(run, offset, head, new_id) { + let mut freed = Vec::new(); + self.discard_run(new_id, worker, &mut freed); + self.recycle(&mut freed); + return Some(existing); + } + (new_id, 0) + }; + pending.push(Placed { + run: target, + offset: first as u32, + start: block_start, + count: remaining.len(), + conflicts: 0, + }); + self.credit(worker, contents.len()); + self.distinct_add(worker, contents.len()); + None + } + + /// Extend a leaf in place when its window ends the hash array and the array has room; + /// otherwise move it to a larger array. + fn append(&self, run: &Run, contents: &[u64], engines: &[u64]) { + let block = run.block.load(Ordering::Relaxed); + let engine = run.engine.load(Ordering::Relaxed); + let base = run.base.load(Ordering::Relaxed) as usize; + let len = run.len(); + let header = self.arena.array_header(block); + let (used, capacity) = unpack(header.load(Ordering::Acquire)); + let end = base + len; + let claimed = end == used as usize + && end + contents.len() <= capacity as usize + && header + .compare_exchange( + pack(used, capacity), + pack((end + contents.len()) as u32, capacity), + Ordering::AcqRel, + Ordering::Acquire, + ) + .is_ok(); + if claimed { + let slots = self.arena.words(block + end as u32, contents.len()); + for (slot, &hash) in slots.iter().zip(contents) { + slot.store(hash, Ordering::Relaxed); + } + let slots = self.arena.words(engine + end as u32, engines.len()); + for (slot, &hash) in slots.iter().zip(engines) { + slot.store(hash, Ordering::Relaxed); + } + run.len + .store((len + contents.len()) as u32, Ordering::Release); + return; + } + let mut grown: Vec = self + .arena + .words(block + base as u32, len) + .iter() + .map(|slot| slot.load(Ordering::Relaxed)) + .collect(); + grown.extend_from_slice(contents); + let mut grown_engines: Vec = self + .arena + .words(engine + base as u32, len) + .iter() + .map(|slot| slot.load(Ordering::Relaxed)) + .collect(); + grown_engines.extend_from_slice(engines); + let new_block = self.arena.alloc_array(&grown, capacity_for(grown.len())); + // The engine array must hold at least what the content array can: an in-place append + // claims room from the content array's header and writes both. + let new_engine = self + .arena + .alloc_array(&grown_engines, self.arena.array_capacity(new_block)); + run.begin_update(); + run.block.store(new_block, Ordering::Relaxed); + run.base.store(0, Ordering::Relaxed); + run.engine.store(new_engine, Ordering::Relaxed); + run.len.store(grown.len() as u32, Ordering::Release); + run.end_update(); + self.arena.array_release(block); + self.arena.array_release(engine); + } + + /// Forget the named blocks of `worker`; unknown hashes are ignored. + pub fn apply_removed(&self, worker: u32, hashes: &[SequenceHash], map: &mut ChainBlockMap) { + let mut refs: Vec = Vec::with_capacity(hashes.len()); + map.remove_all(hashes, |at| refs.push(at)); + let work = group_by_run(refs); + self.apply_removals(worker, work); + } + + /// Drop `worker` from the grouped places, one run at a time, forwarding what a split moved. + fn apply_removals(&self, worker: u32, mut work: Vec) { + // The runs' headers are the lines the locked work reads first: ask for all of them now + // so their misses overlap instead of serialising one run after another. + for removal in &work { + if removal.run != GONE { + crate::prefetch::prefetch_read(std::ptr::from_ref(self.slab.run(removal.run))); + } + } + let mut freed = Vec::new(); + while let Some(removal) = work.pop() { + self.remove_from_run(worker, removal, &mut work, &mut freed); + self.recycle(&mut freed); + } + } + + /// Drop `worker` from the offsets of one run, forwarding offsets a split moved on. + fn remove_from_run( + &self, + worker: u32, + removal: Removal, + work: &mut Vec, + freed: &mut Vec, + ) { + if removal.run == GONE { + return; + } + let run = self.slab.run(removal.run); + let mut meta = self.lock_run(removal.run); + if meta.dead + || removal + .generation + .is_some_and(|generation| generation != run.generation()) + { + return; + } + let offsets = reforward( + &self.arena, + run.forwards.load(Ordering::Relaxed), + removal.offsets, + work, + ); + if offsets.is_empty() || self.held_by(removal.run, worker) == 0 { + return; + } + self.remove_ranges(worker, removal.run, &mut meta, offsets); + if !self.has_holders(removal.run) && run.children.load(Ordering::Relaxed) == NONE { + self.unlink_locked(removal.run, &mut meta, freed); + } + } + + /// Drop `worker` from the given offsets of a run it covers, splitting the run so that the + /// pieces it still covers keep their offsets. + fn remove_ranges(&self, worker: u32, run_id: u32, meta: &mut RunMeta, mut offsets: Vec) { + offsets.sort_unstable(); + offsets.dedup(); + let mut held = self.held_by(run_id, worker); + offsets.retain(|&offset| (offset as usize) < held); + // Contiguous ranges, highest first, so earlier ranges keep their offsets after the splits + // a later range causes. + let mut ranges: Vec<(usize, usize)> = Vec::new(); + for &offset in &offsets { + let offset = offset as usize; + match ranges.last_mut() { + Some((_, high)) if *high + 1 == offset => *high = offset, + _ => ranges.push((offset, offset)), + } + } + for (low, high) in ranges.into_iter().rev() { + // Blocks after the range stay held by this worker too: a hole, so the tail becomes + // its own run. A range reaching the worker's end only lowers its cutoff. + if high + 1 < held { + self.count_split(SplitCause::Hole); + self.split_locked(run_id, meta, high + 1); + } + let before = self.held_len(run_id); + self.set_holding(run_id, worker, low); + self.debit(worker, high + 1 - low); + self.settle_distinct(worker, before, run_id); + held = low; + } + // Every range is applied: cap the run's prefix holders (a tail eviction adds one). + self.cap_prefix_holders(run_id, meta); + } + + /// Forget every block of `worker` (the engine cleared its cache); the map is emptied. + pub fn apply_cleared(&self, worker: u32, map: &mut ChainBlockMap) { + let drained = std::mem::take(map); + self.drop_worker(worker, drained); + } + + /// Forget every block of `worker` (the worker left) and free its slot; the same name interns + /// afresh afterwards. Nothing may write under `worker` from here on: the slot goes to the + /// next name interned. + pub fn remove_worker(&self, worker: u32, map: ChainBlockMap) { + self.drop_worker(worker, map); + self.release_worker(worker); + } + + fn drop_worker(&self, worker: u32, map: ChainBlockMap) { + // Keyed by (run, generation): a forwarding record to a dead generation of an id must not + // shadow the live run that reused the id. + let mut seen: FxHashSet<(u32, Option)> = FxHashSet::default(); + let mut work: Vec<(u32, Option)> = Vec::new(); + for (_, at) in map { + if seen.insert((at.run, None)) { + work.push((at.run, None)); + } + } + let mut freed = Vec::new(); + while let Some((run_id, generation)) = work.pop() { + if run_id == GONE { + continue; + } + let run = self.slab.run(run_id); + let mut meta = self.lock_run(run_id); + if meta.dead || generation.is_some_and(|generation| generation != run.generation()) { + continue; + } + // Blocks of this worker may have moved into suffixes since the map was written. + for (_, suffix, suffix_generation) in self + .arena + .forward_records(run.forwards.load(Ordering::Relaxed)) + { + if suffix != GONE && seen.insert((suffix, Some(suffix_generation))) { + work.push((suffix, Some(suffix_generation))); + } + } + let held = self.held_by(run_id, worker); + if held > 0 { + let before = self.held_len(run_id); + self.set_holding(run_id, worker, 0); + self.debit(worker, held); + self.settle_distinct(worker, before, run_id); + if !self.has_holders(run_id) && run.children.load(Ordering::Relaxed) == NONE { + self.unlink_locked(run_id, &mut meta, &mut freed); + } + } + drop(meta); + self.recycle(&mut freed); + } + } + + /// Every block every worker holds, as `(worker, position, content hash, prefix hash)`; + /// for tests and for comparing against the reference indexer. Not consistent under + /// concurrent writes. + #[doc(hidden)] + pub fn debug_blocks(&self) -> BTreeSet<(u32, usize, ContentHash, SequenceHash)> { + let mut out = BTreeSet::new(); + let mut stack: Vec<(u32, Option)> = Vec::new(); + let (root, _) = self.slab.run(ROOT).snapshot(); + for (_, child, _) in self.arena.table_entries(root.children) { + stack.push((child, None)); + } + while let Some((run_id, mut prefix)) = stack.pop() { + let run = self.slab.run(run_id); + let (window, _) = run.snapshot(); + let start = run.start(); + let holders = workers(self.slab.coverage(run_id)); + let partials = self.arena.partial_entries(window.partials); + let hashes = self + .arena + .words(window.block + window.base, window.len as usize); + // The prefix hash after each count of the run's blocks: a child hangs off any offset + // and continues from the hash at its own. + let mut prefixes: Vec> = Vec::with_capacity(hashes.len() + 1); + prefixes.push(prefix); + for (offset, slot) in hashes.iter().enumerate() { + let content = ContentHash(slot.load(Ordering::Relaxed)); + let next = match prefix { + Some(previous) => chain_prefix_hash(previous, content), + None => SequenceHash(content.0), + }; + for &worker in &holders { + out.insert((worker, start + offset, content, next)); + } + for &(worker, cutoff) in &partials { + if offset < cutoff as usize { + out.insert((worker, start + offset, content, next)); + } + } + prefix = Some(next); + prefixes.push(prefix); + } + for (_, child, _) in self.arena.table_entries(window.children) { + let child_offset = self.slab.run(child).start().saturating_sub(start); + stack.push((child, prefixes[child_offset.min(prefixes.len() - 1)])); + } + } + out + } + + /// Adjacent runs a compaction could fold: a run with exactly one child whose coverage equals + /// its own and that has no prefix holders of its own (a prefix holder of the child that is + /// not a whole holder of the parent holds a suffix, which one run could not express). + /// Returns the pairs and the blocks the children hold. A diagnostic walk, not consistent + /// under concurrent writes. + #[doc(hidden)] + pub fn debug_mergeable(&self) -> (usize, usize) { + let (strict, _, blocks) = self.debug_mergeable_by_rule(); + (strict, blocks) + } + + /// As [`debug_mergeable`](Self::debug_mergeable), but also counting the pairs a generalised + /// merge could fold: the child's whole holders a subset of the parent's and every prefix + /// holder of the child a whole holder of the parent (the parent's extra holders would become + /// prefix holders of the merged run). Returns `(strict pairs, generalised pairs, blocks of the + /// generalised pairs' children)`. + #[doc(hidden)] + fn debug_mergeable_by_rule(&self) -> (usize, usize, usize) { + let (mut strict, mut general, mut blocks) = (0usize, 0usize, 0usize); + let mut stack: Vec = Vec::new(); + let (root, _) = self.slab.run(ROOT).snapshot(); + for (_, child, _) in self.arena.table_entries(root.children) { + stack.push(child); + } + while let Some(run_id) = stack.pop() { + let (window, _) = self.slab.run(run_id).snapshot(); + let children: Vec = self + .arena + .table_entries(window.children) + .into_iter() + .map(|(_, child, _)| child) + .collect(); + if let [only] = children[..] { + let (child_window, _) = self.slab.run(only).snapshot(); + let parent_words: Vec = self + .slab + .coverage(run_id) + .iter() + .map(|w| w.load(Ordering::Relaxed)) + .collect(); + let child_words: Vec = self + .slab + .coverage(only) + .iter() + .map(|w| w.load(Ordering::Relaxed)) + .collect(); + let same_coverage = parent_words == child_words; + let subset = parent_words + .iter() + .zip(&child_words) + .all(|(p, c)| c & !p == 0); + let child_partials = self.arena.partial_entries(child_window.partials); + let partials_held = child_partials + .iter() + .all(|&(worker, _)| has(self.slab.coverage(run_id), worker)); + if same_coverage && child_partials.is_empty() { + strict += 1; + } + if subset && partials_held { + general += 1; + blocks += child_window.len as usize; + } + } + stack.extend(children); + } + (strict, general, blocks) + } + + /// Shape and memory counters. + pub fn stats(&self) -> ChainIndexStats { + let allocated = self.slab.allocated(); + let mut stats = ChainIndexStats { + runs_allocated: allocated, + runs_free: self.slab.free.len(), + arena_bytes: self.arena.used() as usize * size_of::(), + arena_free_bytes: self.arena.free_words() * size_of::(), + arena_chunk_bytes: self.arena.chunk_bytes(), + header_bytes: allocated * (size_of::() + self.words * size_of::()), + slab_bytes: self.slab.chunk_bytes(), + engine_conflicts: self.engine_conflicts.load(Ordering::Relaxed), + landing_mismatches: self.landing_mismatches.load(Ordering::Relaxed), + moved_hashes: self.moved_hashes.load(Ordering::Relaxed), + splits_by_branch: self.splits_branch.load(Ordering::Relaxed), + splits_by_hole: self.splits_hole.load(Ordering::Relaxed), + splits_by_mid_run_store: self.splits_mid_run.load(Ordering::Relaxed), + splits_by_prefix_holders: self.splits_prefix_holders.load(Ordering::Relaxed), + runs_died: self.runs_died.load(Ordering::Relaxed), + ..ChainIndexStats::default() + }; + for id in 1..allocated as u32 { + let run = self.slab.run(id); + if run.meta.lock().dead { + continue; + } + stats.runs_live += 1; + stats.blocks_live += run.len(); + let (window, _) = run.snapshot(); + let partials = self.arena.partial_entries(window.partials).len(); + stats.partial_entries += partials; + stats.max_partials = stats.max_partials.max(partials); + if window.children != NONE { + let live = self.arena.table_live(window.children); + stats.child_entries += live; + stats.child_tombstones += self.arena.table_dead(window.children); + } + } + stats + } +} + +/// Offsets taken from the map may have moved into suffixes since they were written: send those +/// on, tagged with the generation the suffix had at the split. +/// The places of one worker's blocks grouped by run (stored coordinates; the per-run work +/// forwards what a split moved). +fn group_by_run(mut refs: Vec) -> Vec { + refs.sort_unstable_by_key(|at| at.run); + let mut work: Vec = Vec::new(); + let mut index = 0; + while index < refs.len() { + let run = refs[index].run; + let end = refs[index..] + .iter() + .position(|at| at.run != run) + .map_or(refs.len(), |count| index + count); + work.push(Removal { + run, + generation: None, + offsets: refs[index..end].iter().map(|at| at.offset).collect(), + }); + index = end; + } + work +} + +fn reforward( + arena: &WordArena, + forwards: u32, + mut offsets: Vec, + work: &mut Vec, +) -> Vec { + if forwards == NONE { + return offsets; + } + let mut forwarded: FxHashMap<(u32, u32), Vec> = FxHashMap::default(); + offsets.retain(|&offset| match arena.forwards_find(forwards, offset) { + Some((next, generation)) => { + if next.run != GONE { + forwarded + .entry((next.run, generation)) + .or_default() + .push(next.offset); + } + false + } + None => true, + }); + work.extend( + forwarded + .into_iter() + .map(|((run, generation), offsets)| Removal { + run, + generation: Some(generation), + offsets, + }), + ); + offsets +} diff --git a/crates/kv_index/src/chain_index/arena.rs b/crates/kv_index/src/chain_index/arena.rs new file mode 100644 index 0000000000..c286617aab --- /dev/null +++ b/crates/kv_index/src/chain_index/arena.rs @@ -0,0 +1,670 @@ +//! The word arena behind the chain index: hash and engine-hash arrays, child tables, +//! partial-holder tables and forwarding records in one address space of 64-bit words, with +//! per-class free lists. + +use super::*; + +/// Hash array capacity class: `8, 16, .., 128` in steps of 8, then four classes per octave +/// (`160, 192, 224, 256, 320, ..`), so an array wastes at most a quarter of its words to the +/// class, not half; a run of 140 blocks with its slack lands in 160, not 256. +pub(super) fn array_class(capacity: usize) -> usize { + if capacity <= 8 * SMALL_ARRAY_CLASSES { + capacity.div_ceil(8).max(1) - 1 + } else { + // `capacity` in `(low, 2 * low]` with `low` a power of two from 128 up. + let bits = usize::BITS - (capacity - 1).leading_zeros(); + let low = 1usize << (bits - 1); + let quarter = (capacity - low).div_ceil(low / 4).max(1) - 1; + SMALL_ARRAY_CLASSES + 4 * (bits as usize - 8) + quarter + } +} + +pub(super) fn class_capacity(class: usize) -> usize { + if class < SMALL_ARRAY_CLASSES { + 8 * (class + 1) + } else { + let above = class - SMALL_ARRAY_CLASSES; + let low = 1usize << (7 + above / 4); + low + (above % 4 + 1) * (low / 4) + } +} + +/// Room for `len` hashes plus a little for decode extensions. +pub(super) fn capacity_for(len: usize) -> usize { + len + (len / 8).max(2) +} + +pub(super) fn table_class(slots: usize) -> usize { + (slots.trailing_zeros() as usize).saturating_sub(MIN_TABLE_SLOTS.trailing_zeros() as usize) +} + +/// Append-only storage of 64-bit words in fixed chunks with free lists per size class. Holds +/// hash arrays (`used | capacity << 32`, `refs`, then the hashes) and child tables +/// (`slots | used << 32`, `live`, then `(head hash, run id | generation << 32)` slots). An +/// allocation never crosses a chunk, so any array is one slice. +pub(super) struct WordArena { + pub(super) dir: Box<[OnceLock>]>, + pub(super) next: AtomicU64, + pub(super) free_arrays: Vec>, + pub(super) free_tables: Vec>, + pub(super) free_partials: Vec>, +} + +impl WordArena { + pub(super) fn new() -> Self { + Self { + dir: (0..WORD_DIR).map(|_| OnceLock::new()).collect(), + next: AtomicU64::new(1), + free_arrays: (0..ARRAY_CLASSES).map(|_| SegQueue::new()).collect(), + free_tables: (0..TABLE_CLASSES).map(|_| SegQueue::new()).collect(), + free_partials: (0..TABLE_CLASSES).map(|_| SegQueue::new()).collect(), + } + } + + #[inline] + pub(super) fn chunk(&self, index: usize) -> &[AtomicU64] { + self.dir[index].get_or_init(|| (0..WORD_CHUNK).map(|_| AtomicU64::new(0)).collect()) + } + + /// `count` consecutive words starting at `start` (all within one chunk, as allocated). A + /// lock-free reader may compute `count` from a header that is being recycled under it; it + /// confirms the run's version afterwards and discards what it read, so the slice is clamped + /// to the chunk rather than trusted. + #[inline] + pub(super) fn words(&self, start: u32, count: usize) -> &[AtomicU64] { + let start = start as usize; + let offset = start & (WORD_CHUNK - 1); + let end = offset.saturating_add(count).min(WORD_CHUNK); + &self.chunk(start >> WORD_CHUNK_BITS)[offset..end] + } + + #[inline] + pub(super) fn word(&self, at: u32) -> &AtomicU64 { + &self.words(at, 1)[0] + } + + /// Fresh words inside one chunk. + pub(super) fn bump(&self, count: usize) -> u32 { + debug_assert!(0 < count && count <= WORD_CHUNK); + loop { + let current = self.next.load(Ordering::Relaxed); + let mut start = current as usize; + if start >> WORD_CHUNK_BITS != (start + count - 1) >> WORD_CHUNK_BITS { + start = ((start >> WORD_CHUNK_BITS) + 1) << WORD_CHUNK_BITS; + } + let end = start + count; + assert!( + end <= WORD_DIR * WORD_CHUNK, + "chain index arena exhausted: more than 2^32 hash words" + ); + if self + .next + .compare_exchange_weak(current, end as u64, Ordering::Relaxed, Ordering::Relaxed) + .is_ok() + { + self.chunk(start >> WORD_CHUNK_BITS); + return start as u32; + } + } + } + + pub(super) fn used(&self) -> u64 { + self.next.load(Ordering::Relaxed) + } + + /// Bytes of chunks taken from the process allocator: what the arena costs in memory. + pub(super) fn chunk_bytes(&self) -> usize { + let chunks = self + .dir + .iter() + .filter(|chunk| chunk.get().is_some()) + .count(); + chunks * WORD_CHUNK * size_of::() + } + + /// Words sitting in free lists. + pub(super) fn free_words(&self) -> usize { + let arrays: usize = self + .free_arrays + .iter() + .enumerate() + .map(|(class, list)| list.len() * (class_capacity(class) + 2)) + .sum(); + let tables: usize = self + .free_tables + .iter() + .enumerate() + .map(|(class, list)| list.len() * (2 + 2 * (MIN_TABLE_SLOTS << class))) + .sum(); + let partials: usize = self + .free_partials + .iter() + .enumerate() + .map(|(class, list)| list.len() * (2 + (MIN_TABLE_SLOTS << class))) + .sum(); + arrays + tables + partials + } + + // ---- hash arrays: [used | capacity << 32][refs][hash; capacity], data = start + 2 ---- + + /// A hash array holding `contents` with room for at least `capacity`; returns the data start. + pub(super) fn alloc_array(&self, contents: &[u64], capacity: usize) -> u32 { + let wanted = array_class(capacity.max(contents.len()).max(8)); + // A freed array of the wanted class, else of one up to an octave larger: fresh words + // are bumped only when nothing in that range is free, so the free lists of neighbouring + // classes do not fill while others bump (churn moves arrays between classes as runs + // grow and die). + let (class, recycled) = (wanted..ARRAY_CLASSES.min(wanted + ARRAY_FIT_SPAN + 1)) + .find_map(|class| self.free_arrays[class].pop().map(|start| (class, start))) + .unwrap_or((wanted, 0)); + let capacity = class_capacity(class); + let start = if recycled == 0 { + self.bump(capacity + 2) + } else { + recycled + }; + let data = start + 2; + for (slot, &hash) in self.words(data, contents.len()).iter().zip(contents) { + slot.store(hash, Ordering::Relaxed); + } + self.word(start + 1).store(1, Ordering::Relaxed); + self.word(start).store( + pack(contents.len() as u32, capacity as u32), + Ordering::Release, + ); + data + } + + #[inline] + pub(super) fn array_header(&self, data: u32) -> &AtomicU64 { + self.word(data - 2) + } + + /// Words an array can hold: the class it was given, which may be larger than asked for. + pub(super) fn array_capacity(&self, data: u32) -> usize { + unpack(self.array_header(data).load(Ordering::Relaxed)).1 as usize + } + + /// One more run shares this array. + pub(super) fn array_retain(&self, data: u32) { + self.word(data - 1).fetch_add(1, Ordering::Relaxed); + } + + /// One run fewer uses this array; the last one frees it. + pub(super) fn array_release(&self, data: u32) { + if data == NONE { + return; + } + if self.word(data - 1).fetch_sub(1, Ordering::AcqRel) == 1 { + let (_, capacity) = unpack(self.array_header(data).load(Ordering::Relaxed)); + self.free_arrays[array_class(capacity as usize)].push(data - 2); + } + } + + // ---- child tables: [slots | used << 32][live][(head, run | gen << 32); slots] ---- + + pub(super) fn alloc_table(&self, slots: usize) -> u32 { + let class = table_class(slots); + let recycled = self.free_tables[class].pop(); + let table = recycled.unwrap_or_else(|| self.bump(2 + 2 * slots)); + for word in self.words(table + 2, 2 * slots) { + word.store(0, Ordering::Relaxed); + } + self.word(table + 1).store(0, Ordering::Relaxed); + self.word(table).store(slots as u64, Ordering::Release); + table + } + + pub(super) fn free_table(&self, table: u32) { + if table == NONE { + return; + } + let slots = self.table_slots(table); + self.free_tables[table_class(slots)].push(table); + } + + #[inline] + pub(super) fn table_slots(&self, table: u32) -> usize { + self.word(table).load(Ordering::Relaxed) as u32 as usize + } + + /// The child continuing with `head`, as `(run, generation)`. + #[inline] + /// Live entries of a child table (`NONE` has none). + pub(super) fn table_live(&self, table: u32) -> usize { + if table == NONE { + 0 + } else { + self.word(table + 1).load(Ordering::Relaxed) as usize + } + } + + pub(super) fn table_find(&self, table: u32, head: u64) -> Option<(u32, u32)> { + if table == NONE { + return None; + } + let slots = self.table_slots(table); + if slots == 0 || !slots.is_power_of_two() { + // A header being recycled under a lock-free reader; the version check discards this. + return None; + } + let words = self.words(table + 2, 2 * slots); + let mask = slots - 1; + let mut index = head as usize & mask; + for _ in 0..slots { + let entry = words.get(2 * index + 1)?.load(Ordering::Acquire); + if entry == 0 { + return None; + } + if entry != TOMB && words.get(2 * index)?.load(Ordering::Relaxed) == head { + return Some((entry as u32, (entry >> 32) as u32)); + } + index = (index + 1) & mask; + } + None + } + + /// Claim a slot for `key` without a lock: the head word is taken by compare-and-swap, then the + /// run word is published. A reader stops at an empty run word, so an entry in flight is simply + /// not there yet for it; a writer that meets the same head in flight waits for it. + pub(super) fn table_claim(&self, table: u32, key: u64, child: u32, generation: u32) -> Claim { + if table == NONE { + return Claim::Full; + } + let slots = self.table_slots(table); + if slots == 0 || !slots.is_power_of_two() { + return Claim::Full; + } + let words = self.words(table + 2, 2 * slots); + let mask = slots - 1; + let mut index = key as usize & mask; + for _ in 0..slots { + let (Some(head_word), Some(run_word)) = + (words.get(2 * index), words.get(2 * index + 1)) + else { + return Claim::Full; + }; + let mut head = head_word.load(Ordering::Acquire); + if head == 0 { + match head_word.compare_exchange(0, key, Ordering::AcqRel, Ordering::Acquire) { + Ok(_) => { + run_word.store( + u64::from(child) | (u64::from(generation) << 32), + Ordering::Release, + ); + self.word(table).fetch_add(1 << 32, Ordering::Relaxed); + self.word(table + 1).fetch_add(1, Ordering::Relaxed); + return Claim::Inserted; + } + Err(taken) => head = taken, + } + } + if head == key { + loop { + let run = run_word.load(Ordering::Acquire); + if run == 0 { + std::hint::spin_loop(); + continue; + } + if run == TOMB { + break; + } + return Claim::Exists(run as u32, (run >> 32) as u32); + } + } + index = (index + 1) & mask; + } + Claim::Full + } + + /// Live `(head, run, generation)` entries of a table. + pub(super) fn table_entries(&self, table: u32) -> Vec<(u64, u32, u32)> { + if table == NONE { + return Vec::new(); + } + let slots = self.table_slots(table); + let words = self.words(table + 2, 2 * slots); + (0..slots) + .filter_map(|index| { + let entry = words[2 * index + 1].load(Ordering::Acquire); + (entry != 0 && entry != TOMB).then(|| { + ( + words[2 * index].load(Ordering::Relaxed), + entry as u32, + (entry >> 32) as u32, + ) + }) + }) + .collect() + } + + /// Write an entry into a free or tombstoned slot of a table that has room (writer side). + pub(super) fn table_put(&self, table: u32, head: u64, run: u32, generation: u32) { + let slots = self.table_slots(table); + let words = self.words(table + 2, 2 * slots); + let mask = slots - 1; + let mut index = head as usize & mask; + loop { + let entry = words[2 * index + 1].load(Ordering::Relaxed); + if entry == 0 || entry == TOMB { + words[2 * index].store(head, Ordering::Relaxed); + words[2 * index + 1].store( + u64::from(run) | (u64::from(generation) << 32), + Ordering::Release, + ); + let header = self.word(table); + let used = (header.load(Ordering::Relaxed) >> 32) + u64::from(entry == 0); + header.store(slots as u64 | (used << 32), Ordering::Relaxed); + self.word(table + 1).fetch_add(1, Ordering::Relaxed); + return; + } + index = (index + 1) & mask; + } + } + + /// Tombstone the entry of `run`; returns the live count left. + pub(super) fn table_take(&self, table: u32, run: u32) -> u64 { + let slots = self.table_slots(table); + let words = self.words(table + 2, 2 * slots); + for index in 0..slots { + let entry = words[2 * index + 1].load(Ordering::Relaxed); + if entry != 0 && entry != TOMB && entry as u32 == run { + words[2 * index + 1].store(TOMB, Ordering::Release); + return self.word(table + 1).fetch_sub(1, Ordering::Relaxed) - 1; + } + } + self.word(table + 1).load(Ordering::Relaxed) + } + + /// A new table holding the live entries of `table` plus room for one more. + pub(super) fn table_grown(&self, table: u32) -> u32 { + let entries = self.table_entries(table); + let needed = entries.len() + 1; + let mut slots = MIN_TABLE_SLOTS; + while needed * 4 > slots * 3 { + slots *= 2; + } + let grown = self.alloc_table(slots); + for (head, run, generation) in entries { + self.table_put(grown, head, run, generation); + } + grown + } + + /// A table holding `entries` (`(key, run, generation)`) with room for one more; `NONE` for + /// none. + pub(super) fn table_from(&self, entries: &[(u64, u32, u32)]) -> u32 { + if entries.is_empty() { + return NONE; + } + let needed = entries.len() + 1; + let mut slots = MIN_TABLE_SLOTS; + while needed * 4 > slots * 3 { + slots *= 2; + } + let table = self.alloc_table(slots); + for &(key, run, generation) in entries { + self.table_put(table, key, run, generation); + } + table + } + + /// Tombstoned slots of a table: entries taken out since it was built. + pub(super) fn table_dead(&self, table: u32) -> usize { + let header = self.word(table).load(Ordering::Relaxed); + ((header >> 32) as usize) + .saturating_sub(self.word(table + 1).load(Ordering::Relaxed) as usize) + } +} + +impl WordArena { + // ---- partial holders: [slots | used << 32][live][(worker | cutoff << 32); slots] ---- + // Entries are appended in slot order; a removed entry is a tombstone. A worker listed here + // holds the run's blocks `[0, cutoff)` with `0 < cutoff < len`. + + pub(super) fn alloc_partials(&self, slots: usize) -> u32 { + let class = table_class(slots); + let recycled = self.free_partials[class].pop(); + let table = recycled.unwrap_or_else(|| self.bump(2 + slots)); + for word in self.words(table + 2, slots) { + word.store(0, Ordering::Relaxed); + } + self.word(table + 1).store(0, Ordering::Relaxed); + self.word(table).store(slots as u64, Ordering::Release); + table + } + + pub(super) fn free_partials(&self, table: u32) { + if table == NONE { + return; + } + let slots = self.word(table).load(Ordering::Relaxed) as u32 as usize; + self.free_partials[table_class(slots)].push(table); + } + + /// `(slots, used)` of a partial table. + #[inline] + pub(super) fn partials_shape(&self, table: u32) -> (usize, usize) { + let header = self.word(table).load(Ordering::Relaxed); + (header as u32 as usize, (header >> 32) as usize) + } + + pub(super) fn partials_live(&self, table: u32) -> usize { + if table == NONE { + 0 + } else { + self.word(table + 1).load(Ordering::Relaxed) as usize + } + } + + /// Live `(worker, cutoff)` entries, in slot order. + pub(super) fn partial_entries(&self, table: u32) -> Vec<(u32, u32)> { + if table == NONE { + return Vec::new(); + } + let (_, used) = self.partials_shape(table); + self.words(table + 2, used) + .iter() + .map(|slot| slot.load(Ordering::Relaxed)) + .filter(|entry| *entry != 0 && *entry != TOMB) + .map(|entry| (entry as u32, (entry >> 32) as u32)) + .collect() + } + + /// The slot and cutoff of `worker`, if it is a partial holder. + pub(super) fn partial_find(&self, table: u32, worker: u32) -> Option<(usize, u32)> { + if table == NONE { + return None; + } + let (_, used) = self.partials_shape(table); + self.words(table + 2, used) + .iter() + .enumerate() + .find_map(|(index, slot)| { + let entry = slot.load(Ordering::Relaxed); + (entry != 0 && entry != TOMB && entry as u32 == worker) + .then_some((index, (entry >> 32) as u32)) + }) + } + + /// The largest cutoff among the partial holders (0 when there are none). + pub(super) fn partial_max(&self, table: u32) -> usize { + if table == NONE { + return 0; + } + let (_, used) = self.partials_shape(table); + self.words(table + 2, used) + .iter() + .map(|slot| slot.load(Ordering::Relaxed)) + .filter(|entry| *entry != 0 && *entry != TOMB) + .map(|entry| (entry >> 32) as usize) + .max() + .unwrap_or(0) + } + + pub(super) fn partial_set(&self, table: u32, slot: usize, worker: u32, cutoff: u32) { + self.word(table + 2 + slot as u32).store( + u64::from(worker) | (u64::from(cutoff) << 32), + Ordering::Release, + ); + } + + /// Append an entry; `false` when the table is full. + pub(super) fn partial_put(&self, table: u32, worker: u32, cutoff: u32) -> bool { + if table == NONE { + return false; + } + let (slots, used) = self.partials_shape(table); + if used >= slots { + return false; + } + self.partial_set(table, used, worker, cutoff); + self.word(table) + .store(slots as u64 | ((used as u64 + 1) << 32), Ordering::Relaxed); + self.word(table + 1).fetch_add(1, Ordering::Relaxed); + true + } + + /// Tombstone a slot; returns the live count left. + pub(super) fn partial_remove(&self, table: u32, slot: usize) -> usize { + self.word(table + 2 + slot as u32) + .store(TOMB, Ordering::Release); + self.word(table + 1).fetch_sub(1, Ordering::Relaxed) as usize - 1 + } + + // ---- forwarding records: [slots | count << 32][0][(at | suffix << 32), generation; slots] ---- + // Appended under the run's lock, read without it: `count` is published with a release store + // after the record, and a reader checks the run's version around the read. Same word count as + // a child table of the same slot count, so the two share free lists. + + pub(super) fn alloc_forwards(&self, slots: usize) -> u32 { + let class = table_class(slots); + let recycled = self.free_tables[class].pop(); + let table = recycled.unwrap_or_else(|| self.bump(2 + 2 * slots)); + self.word(table + 1).store(0, Ordering::Relaxed); + self.word(table).store(slots as u64, Ordering::Release); + table + } + + #[inline] + pub(super) fn forwards_shape(&self, table: u32) -> (usize, usize) { + let header = self.word(table).load(Ordering::Acquire); + (header as u32 as usize, (header >> 32) as usize) + } + + /// The records of a table, oldest first. + pub(super) fn forward_records(&self, table: u32) -> Vec<(u32, u32, u32)> { + if table == NONE { + return Vec::new(); + } + let (_, count) = self.forwards_shape(table); + let words = self.words(table + 2, 2 * count); + (0..count) + .map(|index| { + let first = words[2 * index].load(Ordering::Relaxed); + ( + first as u32, + (first >> 32) as u32, + words[2 * index + 1].load(Ordering::Relaxed) as u32, + ) + }) + .collect() + } + + /// Where a block at `offset` of the run went: the oldest record whose split point is at or + /// before it (later splits cut the shorter prefix). + #[inline] + pub(super) fn forwards_find(&self, table: u32, offset: u32) -> Option<(BlockRef, u32)> { + if table == NONE { + return None; + } + let (_, count) = self.forwards_shape(table); + let words = self.words(table + 2, 2 * count); + (0..count).find_map(|index| { + let first = words.get(2 * index)?.load(Ordering::Relaxed); + let at = first as u32; + if at > offset { + return None; + } + Some(( + BlockRef { + run: (first >> 32) as u32, + offset: offset - at, + }, + words.get(2 * index + 1)?.load(Ordering::Relaxed) as u32, + )) + }) + } + + /// Append a record; `false` when the table is full. + pub(super) fn forwards_push(&self, table: u32, at: u32, suffix: u32, generation: u32) -> bool { + if table == NONE { + return false; + } + let (slots, count) = self.forwards_shape(table); + if count >= slots { + return false; + } + let words = self.words(table + 2, 2 * slots); + words[2 * count].store(u64::from(at) | (u64::from(suffix) << 32), Ordering::Relaxed); + words[2 * count + 1].store(u64::from(generation), Ordering::Relaxed); + self.word(table) + .store(slots as u64 | ((count as u64 + 1) << 32), Ordering::Release); + true + } + + /// A new table with the records of `table` and room for more. + pub(super) fn forwards_grown(&self, table: u32) -> u32 { + let records = self.forward_records(table); + let mut slots = MIN_TABLE_SLOTS; + while slots < 2 * (records.len() + 1) { + slots *= 2; + } + let grown = self.alloc_forwards(slots); + for (at, suffix, generation) in records { + self.forwards_push(grown, at, suffix, generation); + } + grown + } + + /// A partial table holding `entries`; `NONE` for none. + pub(super) fn partials_from(&self, entries: &[(u32, u32)]) -> u32 { + if entries.is_empty() { + return NONE; + } + let mut slots = MIN_TABLE_SLOTS; + while slots < 2 * entries.len() { + slots *= 2; + } + let table = self.alloc_partials(slots); + for &(worker, cutoff) in entries { + self.partial_put(table, worker, cutoff); + } + table + } + + /// A new table with the live entries of `table` and room for `extra` more. + pub(super) fn partials_grown(&self, table: u32, extra: usize) -> u32 { + let entries = self.partial_entries(table); + let needed = entries.len() + extra; + let mut slots = MIN_TABLE_SLOTS; + while slots < 2 * needed { + slots *= 2; + } + let grown = self.alloc_partials(slots); + for (worker, cutoff) in entries { + self.partial_put(grown, worker, cutoff); + } + grown + } +} + +#[inline] +pub(super) fn pack(used: u32, capacity: u32) -> u64 { + (u64::from(capacity) << 32) | u64::from(used) +} + +#[inline] +pub(super) fn unpack(header: u64) -> (u32, u32) { + (header as u32, (header >> 32) as u32) +} diff --git a/crates/kv_index/src/chain_index/slab.rs b/crates/kv_index/src/chain_index/slab.rs new file mode 100644 index 0000000000..dd316c2a0b --- /dev/null +++ b/crates/kv_index/src/chain_index/slab.rs @@ -0,0 +1,298 @@ +//! Run headers and the slab that holds them: the seqlock window a reader snapshots, the +//! per-run metadata under its lock, the coverage words of whole holders, and the free list of +//! dead ids with their generations. + +use super::*; + +/// Writer-side bookkeeping of a run, under its lock. +#[derive(Default)] +pub(super) struct RunMeta { + /// Unlinked from the tree; its id may be reused (with the next generation). + pub(super) dead: bool, +} + +/// A reader's consistent view of a run's window. +#[derive(Clone, Copy)] +pub(super) struct Window { + /// Data start of the hash array, or `NONE`. + pub(super) block: u32, + /// Offset of the run's first hash within the array. + pub(super) base: u32, + /// Data start of the engine-hash array, or `NONE`. + pub(super) engine: u32, + pub(super) len: u32, + /// Child table, or `NONE`. + pub(super) children: u32, + /// Partial-holder table, or `NONE`. + pub(super) partials: u32, + /// Forwarding records of the splits this run has undergone, or `NONE`: a block that sat at + /// `offset >= at` before a split lives in that split's suffix at `offset - at` (and may have + /// been forwarded again from there); a `GONE` suffix means nobody holds it any more. Split + /// points decrease along the records: a run never grows after a split. + pub(super) forwards: u32, +} + +pub(super) struct Run { + /// Absolute position of the run's first block. + pub(super) start: AtomicU32, + /// Run id of the parent (the root's parent is itself). + pub(super) parent: AtomicU32, + /// Generation in the high half, seqlock in the low half: odd while an update is in flight. + pub(super) version: AtomicU64, + pub(super) block: AtomicU32, + pub(super) base: AtomicU32, + /// Data start of the engine-hash array, parallel to `block` (same base, same length, shared + /// and released with it), or `NONE`. One engine hash per distinct block, the first holder's. + pub(super) engine: AtomicU32, + pub(super) len: AtomicU32, + pub(super) children: AtomicU32, + /// Workers holding only a prefix of the run, with how much: `(worker, cutoff)` entries in the + /// arena. A worker is either in the coverage bitset (whole run) or here, never both. + pub(super) partials: AtomicU32, + pub(super) forwards: AtomicU32, + /// Lock-free child inserts in progress on this run; a writer that replaces the child table + /// or hands it to a suffix waits for this to drain inside its version step. + pub(super) inflight: AtomicU32, + pub(super) meta: Mutex, +} + +impl Run { + pub(super) fn blank() -> Self { + Self { + start: AtomicU32::new(0), + parent: AtomicU32::new(ROOT), + version: AtomicU64::new(0), + block: AtomicU32::new(NONE), + base: AtomicU32::new(0), + engine: AtomicU32::new(NONE), + len: AtomicU32::new(0), + children: AtomicU32::new(NONE), + partials: AtomicU32::new(NONE), + forwards: AtomicU32::new(NONE), + inflight: AtomicU32::new(0), + meta: Mutex::new(RunMeta::default()), + } + } + + #[inline] + pub(super) fn start(&self) -> usize { + self.start.load(Ordering::Relaxed) as usize + } + + #[inline] + pub(super) fn len(&self) -> usize { + self.len.load(Ordering::Acquire) as usize + } + + #[inline] + pub(super) fn generation(&self) -> u32 { + (self.version.load(Ordering::Relaxed) >> 32) as u32 + } + + /// The window as a reader sees it, with the version to confirm afterwards. An in-place + /// append only grows `len`, published with a release store after its hashes, and needs no + /// version; everything else that changes the window goes through `begin_update`. + #[inline] + pub(super) fn snapshot(&self) -> (Window, u64) { + loop { + let before = self.version.load(Ordering::Acquire); + if before & 1 == 1 { + std::hint::spin_loop(); + continue; + } + let window = Window { + block: self.block.load(Ordering::Relaxed), + base: self.base.load(Ordering::Relaxed), + engine: self.engine.load(Ordering::Relaxed), + len: self.len.load(Ordering::Acquire), + children: self.children.load(Ordering::Relaxed), + partials: self.partials.load(Ordering::Relaxed), + forwards: self.forwards.load(Ordering::Relaxed), + }; + fence(Ordering::Acquire); + if self.version.load(Ordering::Relaxed) == before { + return (window, before); + } + } + } + + /// Whether everything read since the snapshot belongs to it. + #[inline] + pub(super) fn confirm(&self, version: u64) -> bool { + fence(Ordering::Acquire); + self.version.load(Ordering::Relaxed) == version + } + + pub(super) fn begin_update(&self) { + // SeqCst against the inserters' `inflight` increment: either they see the odd version and + // back off, or the writer sees their count and waits (see `wait_inflight`). + self.version.fetch_add(1, Ordering::SeqCst); + } + + /// Wait for lock-free child inserts to finish; called inside a version step, so no new one + /// starts meanwhile. + pub(super) fn wait_inflight(&self) { + while self.inflight.load(Ordering::SeqCst) != 0 { + std::hint::spin_loop(); + } + } + + pub(super) fn end_update(&self) { + self.version.fetch_add(1, Ordering::Release); + } + + /// Start a new life of this header: next generation, fields reset. + pub(super) fn reincarnate(&self, start: usize, parent: u32, window: Window) { + self.version.fetch_add((1 << 32) | 1, Ordering::Acquire); + self.start.store(start as u32, Ordering::Relaxed); + self.parent.store(parent, Ordering::Relaxed); + self.block.store(window.block, Ordering::Relaxed); + self.base.store(window.base, Ordering::Relaxed); + self.engine.store(window.engine, Ordering::Relaxed); + self.len.store(window.len, Ordering::Relaxed); + self.children.store(window.children, Ordering::Relaxed); + self.partials.store(window.partials, Ordering::Relaxed); + self.forwards.store(window.forwards, Ordering::Relaxed); + self.end_update(); + } +} + +pub(super) struct RunChunk { + pub(super) runs: Box<[Run]>, + /// `words` coverage words per run, run-major. + pub(super) coverage: Box<[AtomicU64]>, +} + +/// Run storage: a directory of fixed-size chunks created on first use, with a free list of dead +/// ids. +pub(super) struct RunSlab { + pub(super) dir: Box<[OnceLock]>, + pub(super) next: AtomicU32, + pub(super) free: SegQueue, + pub(super) words: usize, +} + +impl RunSlab { + pub(super) fn new(words: usize) -> Self { + Self { + dir: (0..RUN_DIR).map(|_| OnceLock::new()).collect(), + next: AtomicU32::new(0), + free: SegQueue::new(), + words, + } + } + + #[inline] + pub(super) fn chunk(&self, index: usize) -> &RunChunk { + self.dir[index].get_or_init(|| RunChunk { + runs: (0..RUN_CHUNK).map(|_| Run::blank()).collect(), + coverage: (0..RUN_CHUNK * self.words) + .map(|_| AtomicU64::new(0)) + .collect(), + }) + } + + #[inline] + pub(super) fn run(&self, id: u32) -> &Run { + let id = id as usize; + &self.chunk(id >> RUN_CHUNK_BITS).runs[id & (RUN_CHUNK - 1)] + } + + #[inline] + pub(super) fn coverage(&self, id: u32) -> &[AtomicU64] { + let id = id as usize; + let chunk = self.chunk(id >> RUN_CHUNK_BITS); + let first = (id & (RUN_CHUNK - 1)) * self.words; + &chunk.coverage[first..first + self.words] + } + + /// A run that is not yet reachable from the tree: a dead header given its next life, or a + /// fresh one. + pub(super) fn alloc(&self, start: usize, parent: u32, window: Window) -> u32 { + if let Some(id) = self.free.pop() { + let run = self.run(id); + let mut meta = run.meta.lock(); + debug_assert!(meta.dead && coverage_is_empty(self.coverage(id))); + meta.dead = false; + run.reincarnate(start, parent, window); + return id; + } + let id = self.next.fetch_add(1, Ordering::Relaxed); + assert!( + (id as usize) < RUN_DIR * RUN_CHUNK, + "chain index slab exhausted: more than 2^26 runs" + ); + let run = self.run(id); + run.start.store(start as u32, Ordering::Relaxed); + run.parent.store(parent, Ordering::Relaxed); + run.block.store(window.block, Ordering::Relaxed); + run.base.store(window.base, Ordering::Relaxed); + run.engine.store(window.engine, Ordering::Relaxed); + run.len.store(window.len, Ordering::Relaxed); + run.children.store(window.children, Ordering::Relaxed); + run.partials.store(window.partials, Ordering::Relaxed); + run.forwards.store(window.forwards, Ordering::Relaxed); + id + } + + pub(super) fn allocated(&self) -> usize { + self.next.load(Ordering::Relaxed) as usize + } + + /// Bytes of chunks taken from the process allocator: headers and coverage words of every + /// slot, used or not. + pub(super) fn chunk_bytes(&self) -> usize { + let chunks = self + .dir + .iter() + .filter(|chunk| chunk.get().is_some()) + .count(); + chunks * RUN_CHUNK * (size_of::() + self.words * size_of::()) + } +} + +#[inline] +pub(super) fn has(coverage: &[AtomicU64], worker: u32) -> bool { + coverage[(worker / 64) as usize].load(Ordering::Relaxed) & (1u64 << (worker % 64)) != 0 +} + +pub(super) fn set(coverage: &[AtomicU64], worker: u32) { + coverage[(worker / 64) as usize].fetch_or(1u64 << (worker % 64), Ordering::Relaxed); +} + +pub(super) fn clear(coverage: &[AtomicU64], worker: u32) { + coverage[(worker / 64) as usize].fetch_and(!(1u64 << (worker % 64)), Ordering::Relaxed); +} + +pub(super) fn coverage_is_empty(coverage: &[AtomicU64]) -> bool { + coverage + .iter() + .all(|word| word.load(Ordering::Relaxed) == 0) +} + +/// Exactly `worker` and nobody else. +pub(super) fn covered_only_by(coverage: &[AtomicU64], worker: u32) -> bool { + let word = (worker / 64) as usize; + let bit = 1u64 << (worker % 64); + coverage.iter().enumerate().all(|(index, slot)| { + let value = slot.load(Ordering::Relaxed); + if index == word { + value == bit + } else { + value == 0 + } + }) +} + +pub(super) fn workers(coverage: &[AtomicU64]) -> Vec { + coverage + .iter() + .enumerate() + .flat_map(|(index, word)| { + let value = word.load(Ordering::Relaxed); + (0..64) + .filter(move |bit| value & (1u64 << bit) != 0) + .map(move |bit| (index * 64 + bit) as u32) + }) + .collect() +} diff --git a/crates/kv_index/src/chain_index/tests.rs b/crates/kv_index/src/chain_index/tests.rs new file mode 100644 index 0000000000..d810845ab4 --- /dev/null +++ b/crates/kv_index/src/chain_index/tests.rs @@ -0,0 +1,1043 @@ +use super::{arena::*, slab::*, *}; +use crate::reference::{request_prefix_hashes, ReferenceIndexer}; + +fn content(stream: u64, position: usize) -> ContentHash { + crate::compute_content_hash(&[stream as u32, (stream >> 32) as u32, position as u32]) +} + +fn blocks_of(contents: &[ContentHash]) -> Vec { + contents + .iter() + .zip(request_prefix_hashes(contents)) + .map(|(&content_hash, seq_hash)| StoredBlock { + seq_hash, + content_hash, + }) + .collect() +} + +fn scores(index: &ChainIndex, query: &[ContentHash]) -> Vec<(u32, u32)> { + let mut v: Vec<(u32, u32)> = index + .find_matches(query, false) + .scores + .into_iter() + .collect(); + v.sort_unstable(); + v +} + +#[test] +fn array_classes_round_up_and_back() { + for capacity in [1usize, 8, 9, 16, 100, 128, 129, 256, 257, 1000, 4096, 5000] { + let class = array_class(capacity); + assert!(class_capacity(class) >= capacity, "capacity {capacity}"); + assert_eq!(array_class(class_capacity(class)), class); + } + assert_eq!(table_class(2), 0); + assert_eq!(table_class(4), 1); + assert_eq!(table_class(1024), 9); +} + +#[test] +fn store_lookup_and_divergence() { + let index = ChainIndex::with_max_workers(8); + let w = index.intern_worker("w").expect("id"); + let mut map = ChainBlockMap::default(); + let held: Vec = (0..10).map(|p| content(1, p)).collect(); + index + .apply_stored(w, &blocks_of(&held), None, &mut map) + .expect("store"); + assert_eq!(scores(&index, &held), vec![(w, 10)]); + assert_eq!(scores(&index, &held[..4]), vec![(w, 4)]); + let mut diverged = held.clone(); + diverged[5] = content(2, 0); + assert_eq!(scores(&index, &diverged), vec![(w, 5)]); + let mut extended = held.clone(); + extended.push(content(3, 0)); + assert_eq!(scores(&index, &extended), vec![(w, 10)]); + assert_eq!(scores(&index, &[content(9, 0)]), vec![]); + assert_eq!(index.current_size(), 10); + assert_eq!(index.entry_count(), 10); +} + +#[test] +fn two_workers_share_a_prefix_and_split_at_the_fork() { + let index = ChainIndex::with_max_workers(8); + let a = index.intern_worker("a").expect("id"); + let b = index.intern_worker("b").expect("id"); + let (mut ma, mut mb) = (ChainBlockMap::default(), ChainBlockMap::default()); + let base: Vec = (0..6).map(|p| content(1, p)).collect(); + let mut fork = base[..3].to_vec(); + fork.extend((0..4).map(|p| content(2, p))); + index + .apply_stored(a, &blocks_of(&base), None, &mut ma) + .expect("store a"); + index + .apply_stored(b, &blocks_of(&fork), None, &mut mb) + .expect("store b"); + assert_eq!(scores(&index, &base), vec![(a, 6), (b, 3)]); + assert_eq!(scores(&index, &fork), vec![(a, 3), (b, 7)]); + let mut early: Vec<(u32, u32)> = index.find_matches(&base, true).scores.into_iter().collect(); + early.sort_unstable(); + assert_eq!(early, vec![(a, 1), (b, 1)]); + let mut reference = ReferenceIndexer::new(); + reference + .apply_stored(a, &blocks_of(&base), None) + .expect("ref a"); + reference + .apply_stored(b, &blocks_of(&fork), None) + .expect("ref b"); + assert_eq!(index.debug_blocks(), reference.blocks()); + assert_eq!(index.entry_count(), 10); +} + +#[test] +fn a_worker_holds_both_sides_of_its_own_divergence() { + let index = ChainIndex::with_max_workers(8); + let w = index.intern_worker("w").expect("id"); + let mut map = ChainBlockMap::default(); + let base: Vec = (0..6).map(|p| content(1, p)).collect(); + let mut fork = base[..3].to_vec(); + fork.extend((0..2).map(|p| content(2, p))); + let blocks = blocks_of(&base); + index + .apply_stored(w, &blocks, None, &mut map) + .expect("base"); + let fork_blocks = blocks_of(&fork); + index + .apply_stored(w, &fork_blocks[3..], Some(blocks[2].seq_hash), &mut map) + .expect("fork"); + assert_eq!(scores(&index, &base), vec![(w, 6)]); + assert_eq!(scores(&index, &fork), vec![(w, 5)]); + assert_eq!(index.current_size(), 8); +} + +#[test] +fn a_hole_stops_the_match_at_the_hole() { + let index = ChainIndex::with_max_workers(8); + let w = index.intern_worker("w").expect("id"); + let mut map = ChainBlockMap::default(); + let held: Vec = (0..8).map(|p| content(1, p)).collect(); + let blocks = blocks_of(&held); + index + .apply_stored(w, &blocks, None, &mut map) + .expect("store"); + index.apply_removed(w, &[blocks[3].seq_hash], &mut map); + assert_eq!(scores(&index, &held), vec![(w, 3)]); + assert_eq!(index.current_size(), 7); + // Re-storing the missing block after its parent heals the hole. + index + .apply_stored(w, &blocks[3..4], Some(blocks[2].seq_hash), &mut map) + .expect("heal"); + assert_eq!(scores(&index, &held), vec![(w, 8)]); + let mut reference = ReferenceIndexer::new(); + reference.apply_stored(w, &blocks, None).expect("ref"); + assert_eq!(index.debug_blocks(), reference.blocks()); +} + +#[test] +fn a_hole_in_a_shared_run_affects_only_the_evicting_worker() { + let index = ChainIndex::with_max_workers(8); + let v = index.intern_worker("v").expect("id"); + let w = index.intern_worker("w").expect("id"); + let (mut mv, mut mw) = (ChainBlockMap::default(), ChainBlockMap::default()); + let held: Vec = (0..10).map(|p| content(1, p)).collect(); + let blocks = blocks_of(&held); + index.apply_stored(v, &blocks, None, &mut mv).expect("v"); + index.apply_stored(w, &blocks, None, &mut mw).expect("w"); + index.apply_removed(w, &[blocks[5].seq_hash], &mut mw); + assert_eq!(scores(&index, &held), vec![(v, 10), (w, 5)]); + index.apply_removed(v, &[blocks[7].seq_hash, blocks[8].seq_hash], &mut mv); + assert_eq!(scores(&index, &held), vec![(v, 7), (w, 5)]); + index + .apply_stored(w, &blocks[5..6], Some(blocks[4].seq_hash), &mut mw) + .expect("heal"); + assert_eq!(scores(&index, &held), vec![(v, 7), (w, 10)]); + let mut reference = ReferenceIndexer::new(); + reference.apply_stored(v, &blocks, None).expect("ref v"); + reference.apply_stored(w, &blocks, None).expect("ref w"); + reference.apply_removed(v, &[blocks[7].seq_hash, blocks[8].seq_hash]); + assert_eq!(index.debug_blocks(), reference.blocks()); +} + +#[test] +fn tail_removal_truncates_and_a_clear_empties() { + let index = ChainIndex::with_max_workers(8); + let w = index.intern_worker("w").expect("id"); + let mut map = ChainBlockMap::default(); + let held: Vec = (0..8).map(|p| content(1, p)).collect(); + let blocks = blocks_of(&held); + index + .apply_stored(w, &blocks, None, &mut map) + .expect("store"); + let tail: Vec = blocks[5..].iter().map(|b| b.seq_hash).collect(); + index.apply_removed(w, &tail, &mut map); + assert_eq!(scores(&index, &held), vec![(w, 5)]); + assert_eq!(map.len(), 5); + assert_eq!(index.entry_count(), 5); + index.apply_cleared(w, &mut map); + assert!(map.is_empty()); + assert_eq!(scores(&index, &held), vec![]); + assert_eq!(index.current_size(), 0); + assert_eq!(index.entry_count(), 0); + assert!(index.debug_blocks().is_empty()); + let stats = index.stats(); + assert_eq!(stats.runs_live, 0); + assert_eq!(stats.runs_free, 1, "the dead run waits for reuse"); + assert_eq!( + stats.arena_free_bytes, + stats.arena_bytes - 8, + "every array and table is back in a free list" + ); +} + +#[test] +fn dead_runs_and_arrays_are_reused() { + let index = ChainIndex::with_max_workers(8); + let w = index.intern_worker("w").expect("id"); + let mut map = ChainBlockMap::default(); + for round in 0..200u64 { + let held: Vec = (0..12).map(|p| content(10 + round, p)).collect(); + let blocks = blocks_of(&held); + index + .apply_stored(w, &blocks, None, &mut map) + .expect("store"); + assert_eq!(scores(&index, &held), vec![(w, 12)]); + let hashes: Vec = blocks.iter().map(|b| b.seq_hash).collect(); + index.apply_removed(w, &hashes, &mut map); + assert_eq!(scores(&index, &held), vec![]); + } + let stats = index.stats(); + assert!( + stats.runs_allocated <= 3, + "runs were not recycled: {stats:?}" + ); + assert!( + stats.arena_bytes < 4096, + "arena words were not recycled: {stats:?}" + ); + assert_eq!(index.current_size(), 0); +} + +#[test] +fn appends_reuse_the_array_until_another_worker_joins() { + let index = ChainIndex::with_max_workers(8); + let w = index.intern_worker("w").expect("id"); + let v = index.intern_worker("v").expect("id"); + let (mut mw, mut mv) = (ChainBlockMap::default(), ChainBlockMap::default()); + let held: Vec = (0..40).map(|p| content(1, p)).collect(); + let blocks = blocks_of(&held); + index + .apply_stored(w, &blocks[..4], None, &mut mw) + .expect("first"); + for step in 1..10 { + let from = step * 4; + index + .apply_stored( + w, + &blocks[from..from + 4], + Some(blocks[from - 1].seq_hash), + &mut mw, + ) + .expect("extend"); + } + assert_eq!( + index.stats().runs_live, + 1, + "decode extensions stay in one run" + ); + assert_eq!(scores(&index, &held), vec![(w, 40)]); + index + .apply_stored(v, &blocks[..20], None, &mut mv) + .expect("join"); + assert_eq!(scores(&index, &held), vec![(w, 40), (v, 20)]); + assert_eq!( + index.stats().runs_live, + 1, + "a prefix holder joins as a partial holder, no split" + ); + let more: Vec = (0..3).map(|p| content(2, p)).collect(); + let mut long = held.clone(); + long.extend(more); + let long_blocks = blocks_of(&long); + index + .apply_stored(w, &long_blocks[40..], Some(blocks[39].seq_hash), &mut mw) + .expect("extend after join"); + assert_eq!(scores(&index, &long), vec![(w, 43), (v, 20)]); + assert_eq!(index.stats().runs_live, 1, "the run is still w's own leaf"); +} + +#[test] +fn many_children_grow_the_table_and_stay_findable() { + let index = ChainIndex::with_max_workers(8); + let w = index.intern_worker("w").expect("id"); + let mut map = ChainBlockMap::default(); + let prompt: Vec = (0..3).map(|p| content(1, p)).collect(); + index + .apply_stored(w, &blocks_of(&prompt), None, &mut map) + .expect("prompt"); + let anchor = blocks_of(&prompt)[2].seq_hash; + let mut chains = Vec::new(); + for branch in 0..300u64 { + let mut chain = prompt.clone(); + chain.extend((0..2).map(|p| content(100 + branch, p))); + index + .apply_stored(w, &blocks_of(&chain)[3..], Some(anchor), &mut map) + .expect("branch"); + chains.push(chain); + } + for chain in &chains { + assert_eq!(scores(&index, chain), vec![(w, 5)]); + } + let mut unknown = prompt.clone(); + unknown.push(content(999, 0)); + assert_eq!(scores(&index, &unknown), vec![(w, 3)]); + // Unlinking every other branch tombstones its slot; the rest stay findable. + for chain in chains.iter().step_by(2) { + let hashes: Vec = blocks_of(chain)[3..].iter().map(|b| b.seq_hash).collect(); + index.apply_removed(w, &hashes, &mut map); + } + for (branch, chain) in chains.iter().enumerate() { + let expected = if branch % 2 == 0 { 3 } else { 5 }; + assert_eq!( + scores(&index, chain), + vec![(w, expected)], + "branch {branch}" + ); + } + let mut reference = ReferenceIndexer::new(); + reference + .apply_stored(w, &blocks_of(&prompt), None) + .expect("ref prompt"); + for chain in chains.iter().skip(1).step_by(2) { + reference + .apply_stored(w, &blocks_of(chain)[3..], Some(anchor)) + .expect("ref branch"); + } + assert_eq!(index.debug_blocks(), reference.blocks()); +} + +#[test] +fn tail_evictions_and_regrowth_do_not_split() { + let index = ChainIndex::with_max_workers(8); + let w = index.intern_worker("w").expect("id"); + let v = index.intern_worker("v").expect("id"); + let (mut mw, mut mv) = (ChainBlockMap::default(), ChainBlockMap::default()); + let held: Vec = (0..40).map(|p| content(1, p)).collect(); + let blocks = blocks_of(&held); + index.apply_stored(w, &blocks, None, &mut mw).expect("w"); + index.apply_stored(v, &blocks, None, &mut mv).expect("v"); + let mut reference = ReferenceIndexer::new(); + reference.apply_stored(w, &blocks, None).expect("ref w"); + reference.apply_stored(v, &blocks, None).expect("ref v"); + // v evicts its tail twice: a cutoff, not a split. + for keep in [30usize, 12] { + let gone: Vec = blocks[keep..].iter().map(|b| b.seq_hash).collect(); + index.apply_removed(v, &gone, &mut mv); + reference.apply_removed(v, &gone); + assert_eq!(scores(&index, &held), vec![(w, 40), (v, keep as u32)]); + assert_eq!(index.stats().runs_live, 1, "tail eviction to {keep}"); + assert_eq!(index.worker_block_count(v), keep); + assert_eq!(index.entry_count(), 40); + } + // A decode extends v back to the end of the run: full holder again, still one run. + index + .apply_stored(v, &blocks[12..], Some(blocks[11].seq_hash), &mut mv) + .expect("regrow"); + reference + .apply_stored(v, &blocks[12..], Some(blocks[11].seq_hash)) + .expect("ref regrow"); + assert_eq!(scores(&index, &held), vec![(w, 40), (v, 40)]); + assert_eq!(index.stats().runs_live, 1); + assert_eq!(index.debug_blocks(), reference.blocks()); + // w evicts everything, v keeps a prefix: the run survives with one partial holder. + let all: Vec = blocks.iter().map(|b| b.seq_hash).collect(); + index.apply_removed(w, &all, &mut mw); + reference.apply_removed(w, &all); + index.apply_removed(v, &all[25..], &mut mv); + reference.apply_removed(v, &all[25..]); + assert_eq!(scores(&index, &held), vec![(v, 25)]); + assert_eq!(index.entry_count(), 25); + assert_eq!(index.debug_blocks(), reference.blocks()); + let walked = index.score_into(&held, |c| c.0, false, |_, _| {}); + assert_eq!(walked, 1); +} + +#[test] +fn a_staircase_of_prefix_holders_is_a_few_runs() { + let index = ChainIndex::with_max_workers(128); + let held: Vec = (0..64).map(|p| content(3, p)).collect(); + let blocks = blocks_of(&held); + let mut maps: Vec = (0..64).map(|_| ChainBlockMap::default()).collect(); + let mut reference = ReferenceIndexer::new(); + for step in 0..64usize { + assert_eq!(index.intern_worker(&format!("w{step}")), Ok(step as u32)); + } + // The longest holder stores first (a full prefill); every shorter prefix then joins as a + // partial holder, the way evictions and cache hits shape a shared prompt. + for step in (0..64usize).rev() { + index + .apply_stored(step as u32, &blocks[..=step], None, &mut maps[step]) + .expect("store"); + reference + .apply_stored(step as u32, &blocks[..=step], None) + .expect("ref"); + } + // A divergence never splits, but more than `PARTIAL_CAP` prefix holders on one run do, + // at their median cutoff: 64 prefixes of one chain end up in a few runs with the longer + // holders whole on the prefixes, never in one run per holder. + let runs = index.stats().runs_live; + assert!( + (2..=2 * 64 / PARTIAL_CAP).contains(&runs), + "64 prefixes of one chain under the prefix-holder cap of {PARTIAL_CAP}: {runs} runs" + ); + assert!(index.stats().splits_by_prefix_holders >= 1); + let expected: Vec<(u32, u32)> = (0..64u32).map(|w| (w, w + 1)).collect(); + assert_eq!(scores(&index, &held), expected); + assert_eq!(index.debug_blocks(), reference.blocks()); + let walked = index.score_into(&held, |c| c.0, false, |_, _| {}); + assert_eq!(walked, runs, "the walk visits each run of the chain once"); + let mut early: Vec<(u32, u32)> = index.find_matches(&held, true).scores.into_iter().collect(); + early.sort_unstable(); + assert_eq!(early, (0..64u32).map(|w| (w, 1)).collect::>()); + // A hole in the longest holder's prefix still splits, and only for it. + index.apply_removed(63, &[blocks[40].seq_hash], &mut maps[63]); + reference.apply_removed(63, &[blocks[40].seq_hash]); + assert_eq!(scores(&index, &held)[63], (63, 40)); + assert!(index.stats().runs_live >= runs); + assert_eq!(index.debug_blocks(), reference.blocks()); +} + +/// A lock-free reader may still hold a table id whose words were recycled and rewritten by +/// the time it reads them; the version check discards the read, so the read itself must only +/// return garbage, never panic. Rewrites the headers a recycled slot could carry: a child +/// table that became a forwards table of another size, and a partial table whose used count +/// outgrew its slots. +#[test] +fn stale_table_headers_are_read_without_panicking() { + let arena = WordArena::new(); + let child_table = arena.alloc_table(MIN_TABLE_SLOTS); + arena.table_put(child_table, 0xabcd, 7, 1); + let partial_table = arena.alloc_partials(MIN_TABLE_SLOTS); + arena.partial_put(partial_table, 3, 5); + let forwards_table = arena.alloc_forwards(MIN_TABLE_SLOTS); + arena.forwards_push(forwards_table, 4, 9, 2); + // Garbage of every shape a recycled header might show: zero, not a power of two, huge. + for garbage in [ + 0u64, + 3, + u64::MAX, + (u64::MAX << 32) | 5, + (7u64 << 32) | (1 << 31), + ] { + for table in [child_table, partial_table, forwards_table] { + arena.word(table).store(garbage, Ordering::Relaxed); + let _ = arena.table_find(table, 0xabcd); + let _ = arena.forwards_find(table, 4); + let _ = arena.partial_find(table, 3); + let _ = arena.partial_max(table); + let (_, used) = arena.partials_shape(table); + let _ = arena.words(table + 2, used).len(); + } + } + // A whole walk over an index keeps working after the headers it reads are rewritten. + let index = ChainIndex::with_max_workers(8); + let w = index.intern_worker("w").expect("id"); + let mut map = ChainBlockMap::default(); + let held: Vec = (0..12).map(|p| content(1, p)).collect(); + let blocks = blocks_of(&held); + index + .apply_stored(w, &blocks, None, &mut map) + .expect("store"); + assert_eq!(scores(&index, &held), vec![(w, 12)]); +} + +/// The gateway's parent-missing fallback: a worker holds b0..b3, evicts b2 and b3, the +/// engine extends after b3 but the parent is unknown, so b4 and b5 are stored without a +/// parent and land at positions 0 and 1 under the hashes of positions 4 and 5; when the +/// engine then announces the whole chain, each hash is held at one place only, the old +/// memberships are released, the counts read six blocks, a query for the mislaid pair +/// scores nothing, and a worker interned later into the freed id inherits nothing. +#[test] +fn a_hash_stored_again_at_another_position_releases_its_old_place() { + let index = ChainIndex::with_max_workers(8); + let mut reference = ReferenceIndexer::new(); + let w = index.intern_worker("w").expect("id"); + let mut map = ChainBlockMap::default(); + let chain: Vec = (0..6).map(|p| content(23, p)).collect(); + let blocks = blocks_of(&chain); + index + .apply_stored(w, &blocks[..4], None, &mut map) + .expect("b0..b3"); + reference.apply_stored(w, &blocks[..4], None).expect("ref"); + let evicted = [blocks[2].seq_hash, blocks[3].seq_hash]; + index.apply_removed(w, &evicted, &mut map); + reference.apply_removed(w, &evicted); + // The fallback: b4 and b5 with no parent, so at positions 0 and 1. + index + .apply_stored(w, &blocks[4..], None, &mut map) + .expect("fallback"); + reference + .apply_stored(w, &blocks[4..], None) + .expect("ref fallback"); + assert_eq!(scores(&index, &chain[4..]), vec![(w, 2)]); + assert_eq!(index.worker_block_count(w), 4); + // The engine announces the chain whole. + index + .apply_stored(w, &blocks, None, &mut map) + .expect("whole"); + reference.apply_stored(w, &blocks, None).expect("ref whole"); + assert_eq!(index.worker_block_count(w), 6); + assert_eq!(index.current_size(), 6); + assert_eq!(index.entry_count(), 6); + assert_eq!(index.stats().moved_hashes, 2); + assert_eq!(scores(&index, &chain), vec![(w, 6)]); + assert_eq!(scores(&index, &chain[4..]), vec![]); + assert_eq!(index.debug_blocks(), reference.blocks()); + assert_eq!(reference.find_matches(&chain[4..]).len(), 0); + // The id is freed and reused: nothing is inherited. + index.remove_worker(w, map); + assert!(index.is_empty()); + let v = index.intern_worker("v").expect("id"); + assert_eq!(v, w); + assert_eq!(scores(&index, &chain), vec![]); + assert_eq!(index.worker_block_count(v), 0); +} + +/// A worker that stores its chain again beside another worker's fork off it re-stores every +/// block at the same place: nothing moves, nothing is released. +#[test] +fn a_re_store_beside_a_fork_moves_nothing() { + let index = ChainIndex::with_max_workers(8); + let a = index.intern_worker("a").expect("id"); + let b = index.intern_worker("b").expect("id"); + let (mut ma, mut mb) = (ChainBlockMap::default(), ChainBlockMap::default()); + let chain: Vec = (0..30).map(|p| content(24, p)).collect(); + let blocks = blocks_of(&chain); + index.apply_stored(a, &blocks, None, &mut ma).expect("a"); + let mut fork = chain[..12].to_vec(); + fork.extend((12..20).map(|p| content(25, p))); + index + .apply_stored(b, &blocks_of(&fork), None, &mut mb) + .expect("b"); + // a re-stores the whole chain and then a tail after its own parent. + index + .apply_stored(a, &blocks, None, &mut ma) + .expect("a again"); + index + .apply_stored(a, &blocks[20..], Some(blocks[19].seq_hash), &mut ma) + .expect("a tail"); + assert_eq!(index.stats().moved_hashes, 0); + assert_eq!(index.worker_block_count(a), 30); + assert_eq!(scores(&index, &chain), vec![(a, 30), (b, 12)]); + let mut reference = ReferenceIndexer::new(); + reference.apply_stored(a, &blocks, None).expect("ref a"); + reference + .apply_stored(b, &blocks_of(&fork), None) + .expect("ref b"); + assert_eq!(index.debug_blocks(), reference.blocks()); +} + +/// A split by another lane can land between a store's placement, recorded under the run's +/// lock, and its lane-map write. The write then meets a recorded place that forwards to the +/// suffix, exactly as the map's old entry for the block does: not a move, nothing released. +/// Driven through the write phase directly, with the placement recorded before the split +/// (a hole another holder opens; a divergence hangs a child and splits nothing). +#[test] +fn a_placement_recorded_before_a_split_is_not_a_move() { + let index = ChainIndex::with_max_workers(8); + let a = index.intern_worker("a").expect("id"); + let b = index.intern_worker("b").expect("id"); + let (mut ma, mut mb) = (ChainBlockMap::default(), ChainBlockMap::default()); + let chain: Vec = (0..8).map(|p| content(26, p)).collect(); + let blocks = blocks_of(&chain); + index.apply_stored(a, &blocks, None, &mut ma).expect("a"); + index.apply_stored(b, &blocks, None, &mut mb).expect("b"); + let first = ma.get(blocks[0].seq_hash).expect("mapped"); + // b drops block 1: a hole, the run is split at 2. a's map keeps its entries in the + // original run's coordinates, forwarded to the suffix that holds blocks 2..8 now. + index.apply_removed(b, &[blocks[1].seq_hash], &mut mb); + let (suffix, _) = index + .resolve(BlockRef { + run: first.run, + offset: 2, + }) + .expect("forwarded"); + assert_ne!(suffix.run, first.run); + assert_eq!(suffix.offset, 0); + // What a's re-store of the chain records at this point: blocks 0..2 in the prefix, + // blocks 2..8 in the suffix. + let pending = vec![ + Placed { + run: first.run, + offset: 0, + start: 0, + count: 2, + conflicts: 0, + }, + Placed { + run: suffix.run, + offset: 0, + start: 2, + count: 6, + conflicts: 0, + }, + ]; + // Before the write lands, b drops block 4: the suffix is split at 3. + index.apply_removed(b, &[blocks[4].seq_hash], &mut mb); + assert_ne!( + index + .resolve(BlockRef { + run: suffix.run, + offset: 3, + }) + .expect("forwarded again") + .0 + .run, + suffix.run + ); + index.write_placements(a, &blocks, pending, &mut ma); + assert_eq!(index.stats().moved_hashes, 0); + assert_eq!(index.worker_block_count(a), 8); + assert_eq!(ma.len(), 8); + assert_eq!(scores(&index, &chain), vec![(a, 8), (b, 1)]); + // The entries resolve to places that credit a: a store under block 6 finds its parent. + index + .apply_stored(a, &blocks[7..], Some(blocks[6].seq_hash), &mut ma) + .expect("a tail"); + assert_eq!(index.worker_block_count(a), 8); + let mut reference = ReferenceIndexer::new(); + reference.apply_stored(a, &blocks, None).expect("ref a"); + reference.apply_stored(b, &blocks, None).expect("ref b"); + reference.apply_removed(b, &[blocks[1].seq_hash]); + reference.apply_removed(b, &[blocks[4].seq_hash]); + assert_eq!(index.debug_blocks(), reference.blocks()); +} + +/// An engine that names one position by two engine hashes (the same content under the same +/// parent: a twin) has the second filed onto the first by content, and the lane map carries +/// both names for one held position. Removing the first name takes the position with it; +/// a store under the second then points past what the worker holds, which cuts the run at +/// the parent and joins what lies beyond. When nothing lies beyond (nobody else holds the +/// tail, no child hangs there) there is nothing to join: the run ends at the cut and the +/// blocks go after it. +#[test] +fn a_store_under_a_twin_whose_first_name_was_removed_does_not_panic() { + let index = ChainIndex::with_max_workers(8); + let w = index.intern_worker("w").expect("id"); + let mut map = ChainBlockMap::default(); + let chain: Vec = (0..8).map(|p| content(29, p)).collect(); + let blocks = blocks_of(&chain); + index + .apply_stored(w, &blocks, None, &mut map) + .expect("chain"); + // Twins of positions 4..8: the same content under the same parent, other engine hashes. + let twins: Vec = blocks[4..] + .iter() + .enumerate() + .map(|(i, b)| StoredBlock { + seq_hash: SequenceHash(b.seq_hash.0 ^ (0x5151 << (i + 8))), + content_hash: b.content_hash, + }) + .collect(); + index + .apply_stored(w, &twins, Some(blocks[3].seq_hash), &mut map) + .expect("twins"); + assert_eq!(index.stats().engine_conflicts, 4); + assert_eq!(index.worker_block_count(w), 8); + assert_eq!(map.len(), 12); + // The first names of positions 4..8 go: the positions go with them. + let first_names: Vec = blocks[4..].iter().map(|b| b.seq_hash).collect(); + index.apply_removed(w, &first_names, &mut map); + assert_eq!(index.worker_block_count(w), 4); + assert_eq!(scores(&index, &chain), vec![(w, 4)]); + // A store under the twin at position 5: the parent entry points past the holding. + let tail = blocks_of(&[ + chain[0], + chain[1], + chain[2], + chain[3], + chain[4], + chain[5], + content(30, 6), + ]); + let stored = index.apply_stored(w, &tail[6..], Some(twins[1].seq_hash), &mut map); + assert!(stored.is_ok(), "{stored:?}"); + assert_eq!(index.worker_block_count(w), 5); + // Positions 0..4 and the new block at 6 are held; positions 4 and 5 are not (the + // documented twin gap: one engine hash per held position). + assert_eq!(scores(&index, &chain[..4]), vec![(w, 4)]); + let query: Vec = tail.iter().map(|b| b.content_hash).collect(); + assert_eq!(scores(&index, &query), vec![(w, 4)]); + assert!(index + .debug_blocks() + .iter() + .any(|b| b.0 == w && b.1 == 6 && b.2 == content(30, 6))); + // Everything the worker holds is still reachable and consistent. + assert_eq!(index.debug_blocks().len(), 5); +} +/// `is_empty` follows the blocks: false from the first store, true again once every block +/// is removed, cleared or taken with its worker. +#[test] +fn emptiness_follows_the_blocks() { + let index = ChainIndex::with_max_workers(8); + assert!(index.is_empty()); + let a = index.intern_worker("a").expect("id"); + let b = index.intern_worker("b").expect("id"); + let (mut ma, mut mb) = (ChainBlockMap::default(), ChainBlockMap::default()); + let chain: Vec = (0..12).map(|p| content(21, p)).collect(); + let blocks = blocks_of(&chain); + index.apply_stored(a, &blocks, None, &mut ma).expect("a"); + assert!(!index.is_empty()); + index + .apply_stored(b, &blocks[..6], None, &mut mb) + .expect("b"); + let hashes: Vec = blocks.iter().map(|block| block.seq_hash).collect(); + index.apply_removed(a, &hashes, &mut ma); + assert!(!index.is_empty(), "b still holds a prefix"); + index.apply_cleared(b, &mut mb); + assert!(index.is_empty()); + index + .apply_stored(a, &blocks[..3], None, &mut ma) + .expect("a again"); + assert!(!index.is_empty()); + index.remove_worker(a, ma); + assert!(index.is_empty()); +} + +/// A content array recycled from a larger class gets an engine-hash twin of at least that +/// capacity, so an in-place append (which claims room from the content header and writes +/// both arrays) never runs past the twin into the words behind it. +#[test] +fn the_engine_twin_is_as_large_as_a_recycled_content_array() { + let index = ChainIndex::with_max_workers(8); + let w = index.intern_worker("w").expect("id"); + let mut mw = ChainBlockMap::default(); + // One freed array two classes above what 140 blocks ask for: the content array takes + // it, the twin must not come out smaller. + let big = index.arena.alloc_array(&[0; 200], 200); + index.arena.array_release(big); + let chain: Vec = (0..200).map(|p| content(9, p)).collect(); + let blocks = blocks_of(&chain); + index + .apply_stored(w, &blocks[..140], None, &mut mw) + .expect("first store"); + let run_id = mw.get(blocks[0].seq_hash).expect("mapped").run; + let run = index.slab.run(run_id); + let (block, engine) = ( + run.block.load(Ordering::Relaxed), + run.engine.load(Ordering::Relaxed), + ); + assert_eq!(block, big, "the recycled array was taken"); + assert!( + index.arena.array_capacity(engine) >= index.arena.array_capacity(block), + "twin {} words, content {}", + index.arena.array_capacity(engine), + index.arena.array_capacity(block) + ); + // Words bumped right behind the twin: an overflowing append would land here. + let guard = index.arena.alloc_array(&[7; 8], 8); + index + .apply_stored(w, &blocks[140..], Some(blocks[139].seq_hash), &mut mw) + .expect("append"); + assert_eq!(run.len(), 200, "appended in place"); + assert_eq!( + run.block.load(Ordering::Relaxed), + block, + "same content array" + ); + for (slot, word) in index.arena.words(guard, 8).iter().enumerate() { + assert_eq!( + word.load(Ordering::Relaxed), + 7, + "guard word {slot} overwritten" + ); + } + for block in &blocks { + assert!(index.is_held(&mw, block.seq_hash), "{:?}", block.seq_hash); + } +} + +/// A run is matched by its engine chain: a store that agrees to the end of the window is +/// taken whole on one compare, one that diverges inside is split exactly where the bisection +/// lands, one that stops short is a prefix, one that runs past the end extends; the +/// reference agrees on every block and no counter moves. +#[test] +fn stores_match_a_run_by_its_engine_chain() { + let index = ChainIndex::with_max_workers(8); + let mut reference = ReferenceIndexer::new(); + let chain: Vec = (0..100).map(|p| content(5, p)).collect(); + let mut forked = chain[..57].to_vec(); + forked.extend((57..100).map(|p| content(6, p))); + let short = chain[..30].to_vec(); + let longer: Vec = (0..140).map(|p| content(5, p)).collect(); + let mut early_fork = chain[..1].to_vec(); + early_fork.extend((1..40).map(|p| content(7, p))); + let mut maps = Vec::new(); + for (name, contents) in [ + ("whole", &chain), + ("forked", &forked), + ("short", &short), + ("longer", &longer), + ("early", &early_fork), + ] { + let worker = index.intern_worker(name).expect("id"); + let mut map = ChainBlockMap::default(); + index + .apply_stored(worker, &blocks_of(contents), None, &mut map) + .expect("store"); + reference + .apply_stored(worker, &blocks_of(contents), None) + .expect("reference store"); + maps.push((worker, map)); + } + for query in [&chain, &forked, &short, &longer, &early_fork] { + let expected: Vec<(u32, u32)> = reference.find_matches(query).into_iter().collect(); + assert_eq!(scores(&index, query), expected); + } + assert_eq!(index.debug_blocks(), reference.blocks()); + // The run of the first 57 blocks is shared by four workers and the walk after the + // divergence took the fork's own run: a decode extension of the fork appends to it. + let (fork_worker, fork_map) = &mut maps[1]; + let more: Vec = (100..110).map(|p| content(6, p)).collect(); + let mut extended = forked.clone(); + extended.extend(more.iter().copied()); + let blocks = blocks_of(&extended); + index + .apply_stored( + *fork_worker, + &blocks[100..], + Some(blocks[99].seq_hash), + fork_map, + ) + .expect("extend"); + reference + .apply_stored(*fork_worker, &blocks[100..], Some(blocks[99].seq_hash)) + .expect("reference extend"); + let expected: Vec<(u32, u32)> = reference.find_matches(&extended).into_iter().collect(); + assert_eq!(scores(&index, &extended), expected); + let stats = index.stats(); + assert_eq!(stats.engine_conflicts, 0); + assert_eq!(stats.landing_mismatches, 0); +} + +/// A worker whose engine names the same content by other hashes cannot match by engine +/// hash; the walk sees the content continue past the mismatch, compares block by block and +/// counts the conflicts, and the worker's own hashes key its lane map: lookups, an extension +/// after its own parent hash and removals by its hashes all stay exact. +#[test] +fn a_worker_with_other_engine_hashes_matches_by_content() { + let index = ChainIndex::with_max_workers(8); + let a = index.intern_worker("a").expect("id"); + let b = index.intern_worker("b").expect("id"); + let (mut ma, mut mb) = (ChainBlockMap::default(), ChainBlockMap::default()); + let chain: Vec = (0..50).map(|p| content(8, p)).collect(); + let blocks = blocks_of(&chain); + let other: Vec = blocks + .iter() + .map(|block| StoredBlock { + seq_hash: SequenceHash(block.seq_hash.0 ^ 0x5bd1_e995), + content_hash: block.content_hash, + }) + .collect(); + index + .apply_stored(a, &blocks[..40], None, &mut ma) + .expect("a"); + index + .apply_stored(b, &other[..40], None, &mut mb) + .expect("b"); + assert_eq!(scores(&index, &chain), vec![(a, 40), (b, 40)]); + assert_eq!(index.stats().engine_conflicts, 40); + index + .apply_stored(b, &other[40..], Some(other[39].seq_hash), &mut mb) + .expect("b extends"); + assert_eq!(scores(&index, &chain), vec![(a, 40), (b, 50)]); + let tail: Vec = other[20..].iter().map(|block| block.seq_hash).collect(); + index.apply_removed(b, &tail, &mut mb); + assert_eq!(scores(&index, &chain), vec![(a, 40), (b, 20)]); + for block in &other[20..] { + assert!(!index.is_held(&mb, block.seq_hash)); + } + for block in &other[..20] { + assert!(index.is_held(&mb, block.seq_hash)); + } + assert_eq!(index.stats().landing_mismatches, 0); +} + +/// A store that carries the engine hashes the index holds for other content (an engine +/// whose hashes do not follow its content, which the relay's hash check refuses) is counted +/// and placed by its content: the match ends where the content does. +#[test] +fn other_content_under_known_engine_hashes_is_counted_and_placed_by_content() { + let index = ChainIndex::with_max_workers(8); + let a = index.intern_worker("a").expect("id"); + let b = index.intern_worker("b").expect("id"); + let (mut ma, mut mb) = (ChainBlockMap::default(), ChainBlockMap::default()); + let chain: Vec = (0..20).map(|p| content(9, p)).collect(); + let blocks = blocks_of(&chain); + let mut other = chain[..10].to_vec(); + other.extend((10..20).map(|p| content(10, p))); + let impostor: Vec = blocks_of(&other) + .into_iter() + .zip(&blocks) + .map(|(block, known)| StoredBlock { + seq_hash: known.seq_hash, + content_hash: block.content_hash, + }) + .collect(); + index.apply_stored(a, &blocks, None, &mut ma).expect("a"); + index.apply_stored(b, &impostor, None, &mut mb).expect("b"); + assert_eq!(index.stats().landing_mismatches, 1); + assert_eq!(scores(&index, &chain), vec![(a, 20), (b, 10)]); + assert_eq!(scores(&index, &other), vec![(a, 10), (b, 20)]); + let mut reference = ReferenceIndexer::new(); + reference.apply_stored(a, &blocks, None).expect("ref a"); + reference.apply_stored(b, &impostor, None).expect("ref b"); + assert_eq!(index.debug_blocks(), reference.blocks()); + let hashes: Vec = impostor.iter().map(|block| block.seq_hash).collect(); + index.apply_removed(b, &hashes, &mut mb); + assert_eq!(scores(&index, &other), vec![(a, 10)]); + assert!(mb.is_empty()); +} + +/// The race the release harness caught: a split of the parent between a store walk's plan +/// and its lock-free claim must make the insert give up, not link the child after the new +/// end. Replayed deterministically: plan (take the version), split, then try the insert. +#[test] +fn a_child_insert_planned_before_a_split_gives_up() { + let index = ChainIndex::with_max_workers(8); + let w = index.intern_worker("w").expect("id"); + let mut mw = ChainBlockMap::default(); + let held: Vec = (0..10).map(|p| content(1, p)).collect(); + let blocks = blocks_of(&held); + index.apply_stored(w, &blocks, None, &mut mw).expect("w"); + let parent = mw.get(blocks[9].seq_hash).expect("mapped").run; + let (_, planned) = index.slab.run(parent).snapshot(); + // A divergence inside the run leaves its end where it is (a fork hangs off its offset), and + // neither does a tail eviction (the cutoff drops, the array stays for regrowth); a hole + // does: w drops blocks 3 and 4, and the run is cut at 5. + let hole: Vec = blocks[3..5].iter().map(|block| block.seq_hash).collect(); + index.apply_removed(w, &hole, &mut mw); + assert_eq!(index.slab.run(parent).len(), 5); + // The child prepared for blocks after the old end must not be linked after the new one. + let contents = [content(3, 0).0, content(3, 1).0]; + let block = index.arena.alloc_array(&contents, capacity_for(2)); + let engine = index + .arena + .alloc_array(&[30, 31], index.arena.array_capacity(block)); + let child = index.slab.alloc( + 0, + parent, + Window { + block, + base: 0, + engine, + len: 2, + children: NONE, + partials: NONE, + forwards: NONE, + }, + ); + set(index.slab.coverage(child), w); + assert!(matches!( + index.insert_child(parent, child, 10, contents[0], planned), + Claim::Changed + )); + let mut freed = Vec::new(); + index.discard_run(child, w, &mut freed); + index.recycle(&mut freed); + // The real path restarts from the parent block and lands the blocks after block 9, in + // the suffix the hole split off; the lookup still stops at the hole. + let mut longer = held.clone(); + longer.extend([content(3, 0), content(3, 1)]); + let longer_blocks = blocks_of(&longer); + index + .apply_stored(w, &longer_blocks[10..], Some(blocks[9].seq_hash), &mut mw) + .expect("extend"); + let mut reference = ReferenceIndexer::new(); + reference.apply_stored(w, &blocks, None).expect("ref w"); + reference.apply_removed(w, &hole); + reference + .apply_stored(w, &longer_blocks[10..], Some(blocks[9].seq_hash)) + .expect("ref extend"); + assert_eq!(index.debug_blocks(), reference.blocks()); + assert_eq!(scores(&index, &longer), vec![(w, 3)]); +} + +/// Interning and releasing worker slots from many threads at once: an intern holds a name-map +/// shard and then the registry, a release must not hold the registry while it takes a shard, +/// or the two deadlock (this test hung within seconds before the order was fixed). +#[test] +fn worker_slots_churn_from_many_threads_without_deadlock() { + let index = ChainIndex::with_max_workers(64); + std::thread::scope(|scope| { + for thread in 0..8u32 { + let index = &index; + scope.spawn(move || { + for round in 0..2_000u32 { + let name = format!("t{thread}-r{}", round % 5); + let id = index.intern_worker(&name).expect("slot"); + let mut map = ChainBlockMap::default(); + let held: Vec = + (0..3).map(|p| content(u64::from(id) + 1, p)).collect(); + index + .apply_stored(id, &blocks_of(&held), None, &mut map) + .expect("store"); + index.remove_worker(id, map); + } + }); + } + }); + assert_eq!(index.current_size(), 0); + assert_eq!( + index.intern_worker("after"), + Ok(index.intern_worker("after").expect("slot")) + ); +} + +#[test] +fn parent_errors_match_the_positional_indexer() { + let index = ChainIndex::with_max_workers(8); + let w = index.intern_worker("w").expect("id"); + let mut map = ChainBlockMap::default(); + let held: Vec = (0..3).map(|p| content(1, p)).collect(); + let blocks = blocks_of(&held); + assert!(matches!( + index.apply_stored(w, &blocks[1..], Some(blocks[0].seq_hash), &mut map), + Err(ApplyError::WorkerNotTracked) + )); + index + .apply_stored(w, &blocks[..1], None, &mut map) + .expect("store"); + assert!(matches!( + index.apply_stored(w, &blocks[2..], Some(blocks[1].seq_hash), &mut map), + Err(ApplyError::ParentBlockNotFound) + )); +} + +#[test] +fn worker_slots_are_bounded_and_reused_after_removal() { + let index = ChainIndex::with_max_workers(2); + assert_eq!(index.intern_worker("a"), Ok(0)); + assert_eq!(index.intern_worker("b"), Ok(1)); + assert_eq!(index.intern_worker("a"), Ok(0)); + assert_eq!(index.intern_worker("c"), Err(WorkerIdExhausted)); + let mut map = ChainBlockMap::default(); + let held: Vec = (0..4).map(|p| content(1, p)).collect(); + index + .apply_stored(0, &blocks_of(&held), None, &mut map) + .expect("store"); + index.remove_worker(0, map); + assert_eq!(index.worker_id("a"), None); + assert_eq!( + index.intern_worker("c"), + Ok(0), + "the freed slot is handed out again" + ); + assert_eq!( + scores(&index, &held), + vec![], + "nothing of the old holder survives" + ); + assert_eq!(index.intern_worker("a"), Err(WorkerIdExhausted)); +} diff --git a/crates/kv_index/src/chain_index/walk.rs b/crates/kv_index/src/chain_index/walk.rs new file mode 100644 index 0000000000..06f7dcfa8b --- /dev/null +++ b/crates/kv_index/src/chain_index/walk.rs @@ -0,0 +1,225 @@ +//! The lookup: a walk from the root along the request's content hashes, scoring every worker +//! by how many leading blocks it holds, read without locks under the runs' versions. + +use super::*; + +impl ChainIndex { + /// Score every worker by how many leading blocks of the request it holds. With `early_exit`, + /// report the workers holding the first block, each scored 1. + pub fn find_matches(&self, content_hashes: &[ContentHash], early_exit: bool) -> OverlapScores { + let mut out = OverlapScores::default(); + self.score_into( + content_hashes, + |content| content.0, + early_exit, + |worker, score| { + out.scores.insert(worker, score); + }, + ); + out + } + + /// The lookup behind [`find_matches`](Self::find_matches), for callers that keep their own + /// hash type and result shape: `hash_of` reads a block's content hash, `report` receives + /// every `(worker, score)` with a non-empty prefix (once each). Nothing is allocated. + /// Returns the number of runs walked, a measure of how fragmented the matched path is. + pub fn score_into( + &self, + content_hashes: &[T], + hash_of: impl Fn(&T) -> u64, + early_exit: bool, + report: impl FnMut(u32, u32), + ) -> usize { + let Some(first) = content_hashes.first() else { + return 0; + }; + match self.head_entry(hash_of(first)) { + Some(entry) => self.score_from(entry, content_hashes, hash_of, early_exit, report), + None => 0, + } + } + + /// The run under the root that starts with the content hash `first`, with its generation: + /// whether any worker holds a chain starting there, read under the root's version. This is + /// the lookup's first step, and what lets a sharded lookup pass over a shard that cannot + /// hold the request at all. + pub(crate) fn head_entry(&self, first: u64) -> Option<(u32, u32)> { + let root = self.slab.run(ROOT); + loop { + let (window, version) = root.snapshot(); + let found = self + .arena + .table_find(window.children, Self::child_key(0, first)); + if root.confirm(version) { + return found; + } + } + } + + /// The lookup from the run under the root that [`head_entry`](Self::head_entry) found. + pub(crate) fn score_from( + &self, + entry: (u32, u32), + content_hashes: &[T], + hash_of: impl Fn(&T) -> u64, + early_exit: bool, + mut report: impl FnMut(u32, u32), + ) -> usize { + let (mut run_id, mut expected) = entry; + // Only the words that can carry a holder: reading all sixteen of a thousand-worker + // index per run visited, and sweeping them three times, was most of a lookup's cost + // in a fleet of eight. + let words = self.live_words.load(Ordering::Acquire).min(self.words); + let mut alive = [0u64; MAX_WORDS]; + // The partial-holder entries of the run in hand live in a per-thread buffer: a walk + // writes the prefix it then reads, so nothing is zeroed per lookup (the 8 KB this + // buffer would cost on the stack was a sixth of a lookup). `report` must not look up. + PARTIAL_BUFFER.with(|buffer| { + let mut partial = buffer.borrow_mut(); + let mut position = 0usize; + let mut walked = 0usize; + loop { + let run = self.slab.run(run_id); + let (window, version) = run.snapshot(); + if (version >> 32) as u32 != expected { + break; + } + let len = window.len as usize; + let available = len.min(content_hashes.len() - position); + let hashes = self.arena.words(window.block + window.base, available); + let matched = content_hashes[position..position + available] + .iter() + .zip(hashes) + .take_while(|(content, slot)| hash_of(content) == slot.load(Ordering::Relaxed)) + .count(); + let coverage = self.slab.coverage(run_id); + let mut held = [0u64; MAX_WORDS]; + for (word, slot) in held[..words].iter_mut().zip(coverage) { + *word = slot.load(Ordering::Relaxed); + } + // The partial-holder table is another line of the arena: read it only when it + // can change the answer, that is at the first run (every entry is a holder there) + // or when a worker still alive does not hold this whole run. A worker moving from + // a prefix to the whole run gains its bit before it loses its entry, so the + // coverage word is read again after the table and both reads count: a bit seen + // either time makes the worker whole, and an entry gone by the table read means + // its bit was set before the second read. + let mut partials = 0usize; + if window.partials != NONE + && (position == 0 + || alive[..words] + .iter() + .zip(&held) + .any(|(word, whole)| word & !whole != 0)) + { + let (_, used) = self.arena.partials_shape(window.partials); + for slot in self.arena.words(window.partials + 2, used) { + let entry = slot.load(Ordering::Relaxed); + if entry != 0 && entry != TOMB && partials < MAX_PARTIAL { + partial[partials] = entry; + partials += 1; + } + } + for (word, slot) in held[..words].iter_mut().zip(coverage) { + *word |= slot.load(Ordering::Relaxed); + } + } + let next = if position + matched < content_hashes.len() { + // The request goes on past what matched here: a child may continue the run + // from this offset, at its end or at a divergence inside it. + self.arena.table_find( + window.children, + Self::child_key(matched, hash_of(&content_hashes[position + matched])), + ) + } else { + None + }; + if !run.confirm(version) { + continue; + } + walked += 1; + if matched == 0 { + break; + } + if position == 0 && early_exit { + alive = held; + for &entry in &partial[..partials] { + let worker = entry as u32; + alive[(worker / 64) as usize] |= 1u64 << (worker % 64); + } + emit(&alive[..words], 1, &mut report); + return walked; + } + // Prefix holders, in one pass: at the first run every entry is a holder, later + // only those still alive. One that holds the whole matched stretch of a run the + // request leaves inside goes on (a child hanging off the divergence may continue + // it); the rest end here with what they hold. A whole holder is never in the + // table, and a worker gaining its bit while its entry lingers counts as whole. + let mut through = [0u64; MAX_WORDS]; + for &entry in &partial[..partials] { + let worker = entry as u32; + let (index, bit) = ((worker / 64) as usize, 1u64 << (worker % 64)); + if held[index] & bit != 0 || (position > 0 && alive[index] & bit == 0) { + continue; + } + let cutoff = (entry >> 32) as usize; + if matched < len && cutoff >= matched { + through[index] |= bit; + } else { + report(worker, (position + cutoff.min(matched)) as u32); + // Reported here: not among the holders dropped below. + alive[index] &= !bit; + } + } + if position > 0 { + for (index, word) in alive[..words].iter_mut().enumerate() { + let dropped = *word & !held[index] & !through[index]; + if dropped != 0 { + emit_word(index, dropped, position as u32, &mut report); + } + } + for (index, word) in alive[..words].iter_mut().enumerate() { + *word &= held[index] | through[index]; + } + } else { + for (index, word) in alive[..words].iter_mut().enumerate() { + *word = held[index] | through[index]; + } + } + if alive[..words].iter().all(|word| *word == 0) { + return walked; + } + position += matched; + match next { + Some((child, generation)) => { + run_id = child; + expected = generation; + } + None => break, + } + } + emit(&alive[..words], position as u32, &mut report); + walked + }) + } +} + +fn emit(alive: &[u64], score: u32, report: &mut impl FnMut(u32, u32)) { + if score == 0 { + return; + } + for (index, &word) in alive.iter().enumerate() { + if word != 0 { + emit_word(index, word, score, report); + } + } +} + +#[inline] +fn emit_word(index: usize, mut word: u64, score: u32, report: &mut impl FnMut(u32, u32)) { + while word != 0 { + let bit = word.trailing_zeros(); + report((index * 64 + bit as usize) as u32, score); + word &= word - 1; + } +} diff --git a/crates/kv_index/src/churn.rs b/crates/kv_index/src/churn.rs new file mode 100644 index 0000000000..bffa18a3ff --- /dev/null +++ b/crates/kv_index/src/churn.rs @@ -0,0 +1,487 @@ +//! Churn generator for the chain index: workers with block-LRU caches serving prefix requests +//! over a shared chain pool, the way the mock engines do in a soak, so the index sees the +//! stores, evictions and heals that fragment runs over hours. Shared by the churn bench +//! (timing, series, JSON) and the gate test (a few minutes of events, a bound on runs live). +//! +//! Every request is one lookup (timed by the caller through the report) routed to the worker +//! with the longest prefix, which then "computes" the blocks it lacks and stores them under the +//! last block it holds, touches the whole prefix in its LRU, and evicts past its capacity in +//! the configured order: least recently used first, equal-age blocks tail-first (the mock's +//! order, which frees suffixes) or by hash (which frees middle stretches). + +#![expect(clippy::expect_used)] + +use std::collections::{BTreeMap, BTreeSet, HashMap}; + +use crate::{ + chain_prefix_hash, compute_content_hash, request_prefix_hashes, ChainBlockMap, ContentHash, + ReferenceIndexer, SequenceHash, ShardedChainIndex, StoredBlock, +}; + +/// Order among equal-age blocks when a worker frees past its capacity. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum FreeOrder { + /// Highest position of the chain first: an eviction takes a suffix. + TailFirst, + /// By block hash: an eviction takes stretches anywhere, holes included. + Hash, +} + +#[derive(Clone, Debug)] +pub struct ChurnConfig { + pub workers: usize, + pub chains: usize, + /// Blocks per chain, drawn uniformly from `min_chain_len..=max_chain_len`. + pub min_chain_len: usize, + pub max_chain_len: usize, + /// Chains after the first share a prefix with an earlier chain with this probability; the + /// shared length is uniform over the parent's length, which is where runs branch. + pub share: f64, + /// Blocks a worker keeps before it evicts. + pub cache_blocks: usize, + pub free_order: FreeOrder, + /// Blocks a request generates after its prompt prefix, as the mock's decode does: private + /// content stored block by block under the prefix, released to the LRU with the rest. + pub decode_blocks: usize, + /// Probability that a request decodes at all (an output shorter than a block stores none). + pub decode_probability: f64, + /// Every this many requests one worker restarts: its cache is emptied (`apply_cleared`) and + /// the next `refill_requests` requests are routed to it whatever the holders, as a balancing + /// router does with a cold engine. Zero: never. + pub restart_every: u64, + pub refill_requests: u64, + pub seed: u64, +} + +/// A chain of the pool with the offset where it diverged from its parent (its own length when +/// it has none): the content's branch points, which bound how many runs the index needs. +pub struct Chain { + pub contents: Vec, + pub blocks: Vec, + pub parent: Option, + pub divergence: usize, +} + +pub struct Pool { + pub chains: Vec, +} + +impl Pool { + pub fn generate(cfg: &ChurnConfig, rng: &mut Rng) -> Self { + let mut chains: Vec = Vec::with_capacity(cfg.chains); + for stream in 0..cfg.chains { + let len = cfg.min_chain_len + rng.below(cfg.max_chain_len - cfg.min_chain_len + 1); + let (parent, shared) = if stream > 0 && rng.unit() < cfg.share { + let parent = rng.below(stream); + let shared = 1 + rng.below( + chains[parent] + .contents + .len() + .min(len.saturating_sub(1)) + .max(1), + ); + (Some(parent), shared.min(len)) + } else { + (None, 0) + }; + let mut contents = Vec::with_capacity(len); + if let Some(parent) = parent { + contents.extend_from_slice(&chains[parent].contents[..shared]); + } + for position in contents.len()..len { + contents.push(content(stream as u64, position)); + } + let blocks: Vec = contents + .iter() + .zip(request_prefix_hashes(&contents)) + .map(|(&content_hash, seq_hash)| StoredBlock { + seq_hash, + content_hash, + }) + .collect(); + chains.push(Chain { + contents, + blocks, + parent, + divergence: if parent.is_some() { shared } else { len }, + }); + } + Self { chains } + } + + /// Distinct branch points of the content: pairs (parent chain, offset) where a chain leaves + /// its parent, counted once each. With every chain stored whole, the index needs at most one + /// run per chain plus one per branch point; the gate test holds runs live to a multiple of + /// that. + pub fn divergence_points(&self) -> usize { + let mut points = BTreeSet::new(); + for chain in &self.chains { + if let Some(parent) = chain.parent { + points.insert((parent, chain.divergence)); + } + } + points.len() + } +} + +fn content(stream: u64, position: usize) -> ContentHash { + compute_content_hash(&[stream as u32, (stream >> 32) as u32, position as u32]) +} + +/// xorshift64*, enough for a workload generator. +pub struct Rng(u64); + +impl Rng { + pub fn new(seed: u64) -> Self { + Self(seed.wrapping_mul(0x9E37_79B9_7F4A_7C15) | 1) + } + pub fn next_u64(&mut self) -> u64 { + let mut x = self.0; + x ^= x >> 12; + x ^= x << 25; + x ^= x >> 27; + self.0 = x; + x.wrapping_mul(0x2545_F491_4F6C_DD1D) + } + pub fn below(&mut self, n: usize) -> usize { + (self.next_u64() % n.max(1) as u64) as usize + } + pub fn unit(&mut self) -> f64 { + (self.next_u64() >> 11) as f64 / (1u64 << 53) as f64 + } +} + +/// A block in a worker's cache: where it sits and when it was last touched. +#[derive(Clone, Copy)] +struct Held { + chain: usize, + position: usize, + age: u64, + /// A decode block (private tail) rather than pool content. + decode: bool, +} + +/// Eviction order key: age first, then the configured tie-break. +#[derive(Clone, Copy, PartialEq, Eq, PartialOrd, Ord)] +struct EvictKey { + age: u64, + tie: u64, + chain: usize, + position: usize, +} + +struct Worker { + id: u32, + map: ChainBlockMap, + held: HashMap, + order: BTreeMap, +} + +impl Worker { + fn tie(free_order: FreeOrder, _chain: usize, position: usize, hash: SequenceHash) -> u64 { + match free_order { + // Tail first: a higher position sorts earlier among equal ages. + FreeOrder::TailFirst => u64::MAX - position as u64, + FreeOrder::Hash => hash.0, + } + } + + fn touch( + &mut self, + free_order: FreeOrder, + chain: usize, + position: usize, + hash: SequenceHash, + age: u64, + decode: bool, + ) { + if let Some(old) = self.held.get(&hash).copied() { + let key = EvictKey { + age: old.age, + tie: Self::tie(free_order, old.chain, old.position, hash), + chain: old.chain, + position: old.position, + }; + self.order.remove(&key); + } + self.held.insert( + hash, + Held { + chain, + position, + age, + decode, + }, + ); + let key = EvictKey { + age, + tie: Self::tie(free_order, chain, position, hash), + chain, + position, + }; + self.order.insert(key, hash); + } + + /// Blocks past the capacity, least recently used first, equal ages in the configured order. + fn victims(&mut self, capacity: usize) -> Vec { + let excess = self.held.len().saturating_sub(capacity); + let mut out = Vec::with_capacity(excess); + while out.len() < excess { + let Some((&key, &hash)) = self.order.iter().next() else { + break; + }; + self.order.remove(&key); + self.held.remove(&hash); + out.push(hash); + } + out + } + + /// Longest held prefix of `chain`: the engine's own prefix match stops at the first block + /// it lacks, whatever it holds beyond. + fn held_prefix(&self, chain: &Chain) -> usize { + chain + .blocks + .iter() + .take_while(|block| self.held.contains_key(&block.seq_hash)) + .count() + } +} + +/// What one request did. +#[derive(Clone, Copy, Debug, Default)] +pub struct StepReport { + pub lookup_ns: u64, + pub runs_walked: usize, + pub holders: usize, + pub best_score: u32, + pub stored_blocks: usize, + pub removed_blocks: usize, + /// Stored blocks that lay before a block the worker still held: a hole healed. + pub healed_blocks: usize, + /// Decode blocks stored (private tails). + pub decoded_blocks: usize, +} + +pub struct Churn { + cfg: ChurnConfig, + pub pool: Pool, + workers: Vec, + rng: Rng, + clock: u64, + scores: Vec<(u32, u32)>, + /// A restarted worker still being refilled: its index and the requests left to send it. + refill: Option<(usize, u64)>, + pub restarts: u64, +} + +impl Churn { + /// Interns `cfg.workers` workers round robin over the shards. + pub fn new(cfg: ChurnConfig, index: &ShardedChainIndex) -> Self { + let mut rng = Rng::new(cfg.seed); + let pool = Pool::generate(&cfg, &mut rng); + let workers = (0..cfg.workers) + .map(|w| Worker { + id: index + .intern_worker_in(w % index.shards(), &format!("churn-{w}")) + .expect("worker slots"), + map: ChainBlockMap::default(), + held: HashMap::new(), + order: BTreeMap::new(), + }) + .collect(); + Self { + cfg, + pool, + workers, + rng, + clock: 0, + scores: Vec::new(), + refill: None, + restarts: 0, + } + } + + /// Decode blocks still held by some worker: an upper bound on the live private tails, each of + /// which is a run of its own. + pub fn live_decode_blocks(&self) -> usize { + self.workers + .iter() + .map(|w| w.held.values().filter(|h| h.decode).count()) + .sum() + } + + pub fn worker_ids(&self) -> Vec { + self.workers.iter().map(|w| w.id).collect() + } + + /// One request: a lookup over a random prefix of a random chain, routed to the best holder, + /// which stores what it lacks and evicts past capacity. `reference` receives the same events. + pub fn step( + &mut self, + index: &ShardedChainIndex, + mut reference: Option<&mut ReferenceIndexer>, + ) -> StepReport { + self.clock += 1; + if self.cfg.restart_every > 0 && self.clock.is_multiple_of(self.cfg.restart_every) { + // A worker restarts: everything it held is gone, and it is refilled cold. + let which = self.rng.below(self.workers.len()); + let worker = &mut self.workers[which]; + index.apply_cleared(worker.id, &mut worker.map); + if let Some(reference) = reference.as_mut() { + reference.apply_cleared(worker.id); + } + worker.held.clear(); + worker.order.clear(); + self.refill = Some((which, self.cfg.refill_requests)); + self.restarts += 1; + } + let chain_pick = self.rng.below(self.pool.chains.len()); + let chain = &self.pool.chains[chain_pick]; + let prefix = 1 + self.rng.below(chain.contents.len()); + let query = &chain.contents[..prefix]; + self.scores.clear(); + let started = std::time::Instant::now(); + let walked = index.score_into( + query, + |c| c.0, + false, + |worker, score| { + self.scores.push((worker, score)); + }, + ); + let lookup_ns = started.elapsed().as_nanos() as u64; + let best = self.scores.iter().max_by_key(|(_, s)| *s).copied(); + let routed = match best { + Some((worker, _)) => self + .workers + .iter() + .position(|w| w.id == worker) + .unwrap_or_else(|| self.rng.below(self.workers.len())), + None => self.rng.below(self.workers.len()), + }; + // A cold worker being refilled takes the request instead of the best holder. + let worker_index = match self.refill { + Some((which, left)) if left > 0 => { + self.refill = Some((which, left - 1)); + which + } + _ => routed, + }; + let mut report = StepReport { + lookup_ns, + runs_walked: walked, + holders: self.scores.len(), + best_score: best.map_or(0, |(_, s)| s), + ..StepReport::default() + }; + let free_order = self.cfg.free_order; + let capacity = self.cfg.cache_blocks; + let worker = &mut self.workers[worker_index]; + let known = worker.held_prefix(chain); + if known < prefix { + let blocks = &chain.blocks[known..prefix]; + let parent = (known > 0).then(|| chain.blocks[known - 1].seq_hash); + // A block stored before one the worker still holds further along is a heal. + report.healed_blocks = blocks + .iter() + .filter(|b| worker.held.contains_key(&b.seq_hash)) + .count(); + index + .apply_stored(worker.id, blocks, parent, &mut worker.map) + .expect("store after a held parent"); + if let Some(reference) = reference.as_mut() { + reference + .apply_stored(worker.id, blocks, parent) + .expect("reference store"); + } + report.stored_blocks = blocks.len(); + } + for (position, block) in chain.blocks[..prefix].iter().enumerate() { + worker.touch( + free_order, + chain_pick, + position, + block.seq_hash, + self.clock, + false, + ); + } + // Decode: private blocks appended under the prefix, one store per block as the mock + // emits them, touched with the request and released with it. + if self.cfg.decode_blocks > 0 && self.rng.unit() < self.cfg.decode_probability { + let mut previous = chain.blocks[prefix - 1].seq_hash; + for k in 0..self.cfg.decode_blocks { + let content = compute_content_hash(&[ + self.clock as u32, + (self.clock >> 32) as u32, + u32::MAX - k as u32, + ]); + let block = StoredBlock { + seq_hash: chain_prefix_hash(previous, content), + content_hash: content, + }; + index + .apply_stored( + worker.id, + std::slice::from_ref(&block), + Some(previous), + &mut worker.map, + ) + .expect("decode store under the block before"); + if let Some(reference) = reference.as_mut() { + reference + .apply_stored(worker.id, std::slice::from_ref(&block), Some(previous)) + .expect("reference decode store"); + } + worker.touch( + free_order, + chain_pick, + prefix + k, + block.seq_hash, + self.clock, + true, + ); + previous = block.seq_hash; + report.decoded_blocks += 1; + } + } + let victims = worker.victims(capacity); + if !victims.is_empty() { + index.apply_removed(worker.id, &victims, &mut worker.map); + if let Some(reference) = reference.as_mut() { + reference.apply_removed(worker.id, &victims); + } + report.removed_blocks = victims.len(); + } + report + } + + /// Score a sample of chains against the reference: the exactness check at a checkpoint. + pub fn check_exact( + &mut self, + index: &ShardedChainIndex, + reference: &ReferenceIndexer, + samples: usize, + ) -> Result<(), String> { + for _ in 0..samples { + let chain = &self.pool.chains[self.rng.below(self.pool.chains.len())]; + let prefix = 1 + self.rng.below(chain.contents.len()); + let query = &chain.contents[..prefix]; + let mut ours: Vec<(u32, u32)> = index + .find_matches(query, false) + .scores + .into_iter() + .collect(); + let mut theirs: Vec<(u32, u32)> = reference.find_matches(query).into_iter().collect(); + ours.sort_unstable(); + theirs.sort_unstable(); + if ours != theirs { + return Err(format!( + "prefix {prefix} of a chain: index {ours:?}, reference {theirs:?}" + )); + } + } + Ok(()) + } +} diff --git a/crates/kv_index/src/event_tree.rs b/crates/kv_index/src/event_tree.rs index 502c79c7f6..b9f0197679 100644 --- a/crates/kv_index/src/event_tree.rs +++ b/crates/kv_index/src/event_tree.rs @@ -3,13 +3,15 @@ //! Uses a single `DashMap<(usize, ContentHash), IndexEntry>` keyed by (position, content_hash). //! Unbounded by default; [`PositionalIndexer::prune`] optionally bounds it with a //! last-touch TTL and/or a capacity ceiling (oldest-first eviction). -//! Jump search skips positions in strides, yielding amortized O(D/J + W) complexity. +//! A lookup verifies every position of the request in order, so a worker's score is exactly +//! the length of the prefix it holds contiguously (O(D) probes for a request of D blocks). //! //! **Dual-hash scheme**: backends send a position-aware `block_hash` (SequenceHash) //! and raw `token_ids` per block. The router computes a position-independent //! ContentHash (XXH3) from token_ids, then a rolling prefix hash (also XXH3) from -//! the ContentHash sequence. SeqEntry is keyed by the router's prefix hash for -//! precise disambiguation at query time. The backend's SequenceHash is stored in +//! the ContentHash sequence. SeqEntry is keyed by the router's prefix hash, and every lookup probe +//! (every position) matches on it, so a +//! request is credited only for the exact chain a worker holds. The backend's SequenceHash is stored in //! worker_blocks only, used for `apply_removed` reverse lookup. //! //! **Performance**: Internal u32 worker IDs eliminate Arc hashing and atomic @@ -31,7 +33,7 @@ use std::{ }; use dashmap::{mapref::entry::Entry, DashMap}; -use rustc_hash::{FxBuildHasher, FxHashMap, FxHashSet}; +use rustc_hash::{FxBuildHasher, FxHashMap}; /// Seed for XXH3 hashing. pub const XXH3_SEED: u64 = 1337; @@ -145,13 +147,53 @@ impl std::error::Error for WorkerIdExhausted {} /// Overlap scores: how many consecutive blocks each worker has cached. /// /// Keys are internal `u32` worker IDs. Use [`PositionalIndexer::worker_id`] to -/// map a worker URL to its internal ID for lookups. +/// map a worker URL to its internal ID for lookups. A worker's total block count +/// is available separately through [`PositionalIndexer::worker_block_count`]; it +/// is not collected per lookup, since no routing decision reads it there. #[derive(Debug, Default)] pub struct OverlapScores { /// internal_worker_id → number of matching prefix blocks (depth in indexer) pub scores: FxHashMap, - /// internal_worker_id → total blocks cached by this worker - pub tree_sizes: FxHashMap, +} + +/// A request's block content hashes as the lookup reads them: by position, without copying. +/// +/// Implemented for `[ContentHash]` and `[u64]`. A caller with its own hash newtype implements +/// it on a wrapper around its slice and calls +/// [`PositionalIndexer::find_matches_in`], so no `Vec` is built per lookup. +pub trait ContentSeq { + /// Number of blocks in the request. + fn len(&self) -> usize; + /// Content hash of the block at `position` (`position < len()`). + fn at(&self, position: usize) -> ContentHash; + /// Whether the request has no blocks. + fn is_empty(&self) -> bool { + self.len() == 0 + } +} + +impl ContentSeq for [ContentHash] { + #[inline] + fn len(&self) -> usize { + <[ContentHash]>::len(self) + } + + #[inline] + fn at(&self, position: usize) -> ContentHash { + self[position] + } +} + +impl ContentSeq for [u64] { + #[inline] + fn len(&self) -> usize { + <[u64]>::len(self) + } + + #[inline] + fn at(&self, position: usize) -> ContentHash { + ContentHash(self[position]) + } } /// Compute content hash from token IDs (position-independent): XXH3-64 with @@ -233,6 +275,135 @@ pub fn compute_request_content_hashes(tokens: &[u32], block_size: usize) -> Vec< // SeqEntry: optimizes for the common case (one seq_hash per position+content) // --------------------------------------------------------------------------- +/// Ids below this many live in the inline words of a [`WorkerSet`]. +const INLINE_WORKER_IDS: u32 = 128; + +/// Dense set of interned worker ids. +/// +/// Two inline words cover ids below [`INLINE_WORKER_IDS`] with no heap +/// allocation, which is every worker in a typical fleet; larger ids spill +/// into a boxed vector that is grown on demand (rare: a fleet past 128 +/// workers). The spill sits behind one thin pointer so the set is 24 bytes and +/// an [`IndexEntry`] 48, which keeps a map bucket (16-byte key included) at 64 +/// bytes. Membership is a bit test, which is what the lookup path does for +/// every active worker at every position. +#[derive(Debug, Clone, Default)] +struct WorkerSet { + low: [u64; 2], + high: Option>, +} + +/// The spill words of a [`WorkerSet`] (ids at or above [`INLINE_WORKER_IDS`]), boxed as one +/// thin pointer: a `Box<[u64]>` is two words and would put the set back at 32 bytes. The extra +/// indirection is paid only on the rare spill path. +#[derive(Debug, Clone, Default)] +struct Spill(Vec); + +impl std::ops::Deref for Spill { + type Target = Vec; + + fn deref(&self) -> &Vec { + &self.0 + } +} + +impl std::ops::DerefMut for Spill { + fn deref_mut(&mut self) -> &mut Vec { + &mut self.0 + } +} + +impl WorkerSet { + fn single(id: u32) -> Self { + let mut set = Self::default(); + set.insert(id); + set + } + + #[inline] + fn slot(id: u32) -> (usize, u64) { + ((id / 64) as usize, 1u64 << (id % 64)) + } + + /// Insert `id`; returns whether it was newly added. + fn insert(&mut self, id: u32) -> bool { + let (word, bit) = Self::slot(id); + if id < INLINE_WORKER_IDS { + let was = self.low[word] & bit != 0; + self.low[word] |= bit; + return !was; + } + let index = word - 2; + let high = self.high.get_or_insert_with(Box::default); + if high.len() <= index { + high.resize(index + 1, 0); + } + let was = high[index] & bit != 0; + high[index] |= bit; + !was + } + + /// Remove `id`; returns whether it was present. + fn remove(&mut self, id: u32) -> bool { + let (word, bit) = Self::slot(id); + if id < INLINE_WORKER_IDS { + let was = self.low[word] & bit != 0; + self.low[word] &= !bit; + return was; + } + let Some(high) = self.high.as_mut() else { + return false; + }; + let index = word - 2; + let Some(slot) = high.get_mut(index) else { + return false; + }; + let was = *slot & bit != 0; + *slot &= !bit; + was + } + + #[inline] + fn contains(&self, id: u32) -> bool { + let (word, bit) = Self::slot(id); + if id < INLINE_WORKER_IDS { + return self.low[word] & bit != 0; + } + self.high + .as_ref() + .and_then(|high| high.get(word - 2)) + .is_some_and(|slot| slot & bit != 0) + } + + fn is_empty(&self) -> bool { + self.low == [0, 0] + && self + .high + .as_ref() + .is_none_or(|high| high.iter().all(|w| *w == 0)) + } + + /// Every id in the set, ascending. + fn iter(&self) -> impl Iterator + '_ { + let words = self + .low + .iter() + .copied() + .chain(self.high.iter().flat_map(|high| high.iter().copied())); + words.enumerate().flat_map(|(index, mut word)| { + let base = index as u32 * 64; + std::iter::from_fn(move || { + if word == 0 { + return None; + } + let bit = word.trailing_zeros(); + word &= word - 1; + Some(base + bit) + }) + }) + } +} + /// Entry for the innermost level of the index. /// /// Optimizes for the common case where there's only one sequence hash @@ -240,16 +411,14 @@ pub fn compute_request_content_hashes(tokens: &[u32], block_size: usize) -> Vec< #[derive(Debug, Clone)] enum SeqEntry { /// Single seq_hash → workers mapping (common case, no HashMap allocation). - Single(SequenceHash, FxHashSet), + Single(SequenceHash, WorkerSet), /// Multiple seq_hash → workers mappings (rare: different prefixes with same content). - Multi(FxHashMap>), + Multi(FxHashMap), } impl SeqEntry { fn new(seq_hash: SequenceHash, worker_id: u32) -> Self { - let mut workers = FxHashSet::default(); - workers.insert(worker_id); - Self::Single(seq_hash, workers) + Self::Single(seq_hash, WorkerSet::single(worker_id)) } /// Insert a worker for a given seq_hash, upgrading to Multi if needed. @@ -280,14 +449,14 @@ impl SeqEntry { fn remove(&mut self, seq_hash: SequenceHash, worker_id: u32) -> (bool, bool) { match self { Self::Single(existing_hash, workers) if *existing_hash == seq_hash => { - let removed = workers.remove(&worker_id); + let removed = workers.remove(worker_id); (removed, workers.is_empty()) } Self::Single(_, _) => (false, false), Self::Multi(map) => { let mut removed = false; if let Some(workers) = map.get_mut(&seq_hash) { - removed = workers.remove(&worker_id); + removed = workers.remove(worker_id); if workers.is_empty() { map.remove(&seq_hash); } @@ -304,13 +473,13 @@ impl SeqEntry { fn accumulate_worker_counts(&self, acc: &mut FxHashMap) { match self { Self::Single(_, workers) => { - for &w in workers { + for w in workers.iter() { *acc.entry(w).or_default() += 1; } } Self::Multi(map) => { for workers in map.values() { - for &w in workers { + for w in workers.iter() { *acc.entry(w).or_default() += 1; } } @@ -319,26 +488,13 @@ impl SeqEntry { } /// Get workers for a specific prefix hash (used in query path and event processing). - fn get(&self, seq_hash: SequenceHash) -> Option<&FxHashSet> { + fn get(&self, seq_hash: SequenceHash) -> Option<&WorkerSet> { match self { Self::Single(existing_hash, workers) if *existing_hash == seq_hash => Some(workers), Self::Single(_, _) => None, Self::Multi(map) => map.get(&seq_hash), } } - - /// For Single entries, return the worker set directly without prefix hash check. - /// Content hash collisions at 64-bit XXH3 are practically impossible (~2^-64), - /// so a matching content_hash at the same position is unambiguous — the rolling - /// hash computation can be skipped entirely. - /// Returns None for Multi entries — caller must compute prefix hash to disambiguate. - #[inline] - fn workers_if_single(&self) -> Option<&FxHashSet> { - match self { - Self::Single(_, workers) => Some(workers), - Self::Multi(_) => None, - } - } } // --------------------------------------------------------------------------- @@ -488,9 +644,19 @@ impl IndexEntry { } } + /// Record an access at coarse time `now`. + /// + /// The stamp has whole-second resolution, so a store is needed at most + /// once per second per entry: every other call finds the value already + /// equal and performs a plain load. That keeps the query path free of + /// shared-memory writes in steady state (a hot entry probed a million + /// times a second is written once), while prune keeps the exact "last + /// store or read" semantics it had when every probe wrote. #[inline] fn touch(&self, now: u32) { - self.last_touch.store(now, Ordering::Relaxed); + if self.last_touch.load(Ordering::Relaxed) != now { + self.last_touch.store(now, Ordering::Relaxed); + } } } @@ -515,7 +681,7 @@ pub struct PruneStats { /// Uses a single `DashMap<(usize, ContentHash), IndexEntry>` — keyed by /// (position, content_hash). Unbounded by default; [`prune`](Self::prune) /// optionally bounds it with a last-touch TTL and a capacity ceiling. -/// Jump search gives amortized O(D/J + W) matching complexity. +/// A lookup probes every position of the request (O(D)); see [`find_matches`](Self::find_matches). /// /// Write-path methods take a caller-owned `&mut WorkerBlockMap` (one per worker). /// This gives direct HashMap access (~5ns) instead of DashMap hash+shard locking @@ -534,7 +700,8 @@ pub struct PositionalIndexer { /// Monotonic counter for assigning new worker IDs. Never recycled; u64 so /// exhaustion of the u32 id space is detected instead of wrapping. next_worker_id: AtomicU64, - /// Jump size for search optimization (default 64). + /// Accepted for API compatibility; the lookup no longer strides (see + /// [`find_matches`](Self::find_matches)). jump_size: usize, /// Origin of the coarse `last_touch` clock (whole seconds since creation). epoch: Instant, @@ -543,12 +710,27 @@ pub struct PositionalIndexer { test_now: AtomicU32, } +/// Per-thread scratch for the lookup path: the request's chain hashes and the active worker +/// set, reused across lookups so a query allocates nothing before it builds its result. +struct LookupScratch { + seq_hashes: Vec, + active: Vec, +} + +thread_local! { + static LOOKUP_SCRATCH: std::cell::RefCell = const { + std::cell::RefCell::new(LookupScratch { + seq_hashes: Vec::new(), + active: Vec::new(), + }) + }; +} + impl PositionalIndexer { - /// Create a new PositionalIndexer with the given jump size. + /// Create a new PositionalIndexer. /// - /// `jump_size` controls how many positions the search algorithm skips at a time. - /// Larger values reduce lookups on long matching prefixes but increase scan range - /// when workers drain. Default: 64. + /// `jump_size` is accepted for API compatibility and does not change the lookup: every + /// position is verified (see [`find_matches`](Self::find_matches)). pub fn new(jump_size: usize) -> Self { assert!(jump_size > 0, "jump_size must be greater than 0"); Self { @@ -605,7 +787,28 @@ impl PositionalIndexer { parent_seq_hash: Option, worker_blocks: &mut WorkerBlockMap, ) -> Result<(), ApplyError> { - if blocks.is_empty() { + self.apply_stored_iter( + worker_id, + blocks.iter().copied(), + parent_seq_hash, + worker_blocks, + ) + } + + /// [`apply_stored`](Self::apply_stored) with the blocks taken from an iterator, so a + /// caller translating another event format does not build a `Vec` per event. + pub fn apply_stored_iter( + &self, + worker_id: u32, + blocks: I, + parent_seq_hash: Option, + worker_blocks: &mut WorkerBlockMap, + ) -> Result<(), ApplyError> + where + I: IntoIterator, + { + let mut blocks = blocks.into_iter().peekable(); + if blocks.peek().is_none() { return Ok(()); } @@ -625,10 +828,11 @@ impl PositionalIndexer { let mut prev_prefix = parent_prefix; let mut num_new_blocks = 0usize; + let mut num_moved = 0usize; // One coarse stamp per batch — cheaper than per-block clock reads and // precise enough for prune's second-granularity TTL. let now = self.now_secs(); - for (i, block) in blocks.iter().enumerate() { + for (i, block) in blocks.enumerate() { let position = start_pos + i; let content_hash = block.content_hash; @@ -659,7 +863,28 @@ impl PositionalIndexer { // pruned (its stale reverse mapping kept) restores its count — // mirroring apply_removed, which decrements only memberships // actually removed. - worker_blocks.insert(block.seq_hash, (position, content_hash, prefix_hash)); + // A hash this worker already held at another place (a store without its parent + // followed by the whole chain, as the gateway's fallback produces) is held at the + // new place only: the latest store wins, and the old membership goes. + if let Some(old) = + worker_blocks.insert(block.seq_hash, (position, content_hash, prefix_hash)) + { + if old != (position, content_hash, prefix_hash) { + let (old_position, old_content, old_prefix) = old; + if let Entry::Occupied(mut occupied) = + self.index.entry((old_position, old_content)) + { + let (removed, now_empty) = + occupied.get_mut().seq.remove(old_prefix, worker_id); + if now_empty { + occupied.remove(); + } + if removed { + num_moved += 1; + } + } + } + } if membership_added { num_new_blocks += 1; } @@ -667,8 +892,10 @@ impl PositionalIndexer { } // Atomically update tree_sizes — lock-free array index. - if num_new_blocks > 0 { - self.tree_sizes.add(worker_id, num_new_blocks); + if num_new_blocks > num_moved { + self.tree_sizes.add(worker_id, num_new_blocks - num_moved); + } else if num_moved > num_new_blocks { + self.tree_sizes.sub(worker_id, num_moved - num_new_blocks); } Ok(()) @@ -692,8 +919,21 @@ impl PositionalIndexer { seq_hashes: &[SequenceHash], worker_blocks: &mut WorkerBlockMap, ) { + self.apply_removed_iter(worker_id, seq_hashes.iter().copied(), worker_blocks); + } + + /// [`apply_removed`](Self::apply_removed) with the hashes taken from an iterator, so a + /// caller translating another event format does not build a `Vec` per event. + pub fn apply_removed_iter( + &self, + worker_id: u32, + seq_hashes: I, + worker_blocks: &mut WorkerBlockMap, + ) where + I: IntoIterator, + { let mut num_removed = 0usize; - for &seq_hash in seq_hashes { + for seq_hash in seq_hashes { let Some((position, content_hash, prefix_hash)) = worker_blocks.remove(&seq_hash) else { continue; @@ -752,11 +992,39 @@ impl PositionalIndexer { self.tree_sizes.reset(worker_id); } + /// Number of blocks the index currently holds for `worker_id` (a lock-free counter read). + pub fn worker_block_count(&self, worker_id: u32) -> usize { + self.tree_sizes.load(worker_id) + } + /// Get total number of blocks across all workers. pub fn current_size(&self) -> usize { self.tree_sizes.total() } + /// Every membership in the index as `(worker, position, content hash, prefix hash)`. + /// + /// A full walk under the shard read locks, for the exactness harness only; never call it on + /// a request path. + #[doc(hidden)] + pub fn debug_blocks(&self) -> Vec<(u32, usize, ContentHash, SequenceHash)> { + let mut out = Vec::new(); + for entry in &self.index { + let (position, content) = *entry.key(); + match &entry.value().seq { + SeqEntry::Single(prefix, workers) => { + out.extend(workers.iter().map(|w| (w, position, content, *prefix))); + } + SeqEntry::Multi(map) => { + for (prefix, workers) in map { + out.extend(workers.iter().map(|w| (w, position, content, *prefix))); + } + } + } + } + out + } + /// Number of `(position, content_hash)` entries currently in the index. /// O(shards); intended for prune decisions and observability, not hot paths. pub fn entry_count(&self) -> usize { @@ -774,12 +1042,9 @@ impl PositionalIndexer { /// `None` or `Some(0)` disables the capacity pass. /// /// Semantics and caveats: - /// * Queries touch exactly the entries they read (position 0, jump - /// landings, and linear-scan ranges), so entries a hot request stream - /// actually needs stay resident; interior positions the jump shortcut - /// skips may age out — harmless, the shortcut never reads them, and a - /// later drain across an evicted position under-counts that one score - /// (same tolerance as the documented `apply_removed` gap behavior). + /// * Queries touch every entry they read (each position of the request up + /// to where its last worker drains), so entries a hot request stream + /// actually needs stay resident. /// * Eviction leaves the per-worker reverse maps (`WorkerBlockMap`) /// untouched; a later `apply_removed` for a pruned block is a safe /// no-op (membership check), and stale reverse entries are dropped by @@ -871,25 +1136,39 @@ impl PositionalIndexer { /// Find overlap scores for a request's content hash sequence. /// - /// Uses jump search: strides by `jump_size` positions, only scanning - /// intermediate positions when workers drain (stop matching). - /// Complexity: amortized O(D/J + W) where D=depth, J=jump_size, W=workers. + /// Verifies every position of the request in order (O(D) probes, D = depth), draining + /// workers as they stop matching; a worker's score is the length of the prefix it holds + /// contiguously, so an evicted middle block ends it exactly as it does in the engine. /// /// When `early_exit` is true, returns immediately after finding any match /// at position 0 (score = 1 for all matching workers). Useful when the caller /// only needs to know whether any worker has cached data for this sequence. /// - /// **Assumption**: Block sequences are prefix-closed — if a worker has a block at - /// position N, it has blocks at all positions 0..N. This holds when backends evict - /// from the tail (LRU). If `apply_removed` creates a mid-sequence gap, the rolling - /// prefix hash detects it (the chain breaks at the gap), but the jump heuristic may - /// over-count if it lands past the gap. In practice, backends only evict tail blocks. + /// A worker is credited exactly the prefix it holds contiguously: a block evicted from + /// the middle of a chain ends the score there, as the engine's own prefix match does. pub fn find_matches(&self, content_hashes: &[ContentHash], early_exit: bool) -> OverlapScores { - self.jump_search_matches(content_hashes, early_exit) + self.find_matches_in(content_hashes, early_exit) + } + + /// [`find_matches`](Self::find_matches) over any [`ContentSeq`], so a caller whose + /// request hashes live in its own type (plain `u64`s, or a newtype wrapped by the caller) + /// does not build a `Vec` per lookup. + pub fn find_matches_in( + &self, + sequence: &S, + early_exit: bool, + ) -> OverlapScores { + LOOKUP_SCRATCH.with(|scratch| { + let mut scratch = scratch.borrow_mut(); + let LookupScratch { seq_hashes, active } = &mut *scratch; + seq_hashes.clear(); + active.clear(); + self.scan_with(sequence, early_exit, seq_hashes, active) + }) } // ----------------------------------------------------------------------- - // Internal: router prefix hash + jump search + // Internal: router prefix hash + lookup scan // // The router computes its own rolling hash from ContentHashes (XXH3). // This hash is stored in SeqEntry during apply_stored and recomputed @@ -909,18 +1188,18 @@ impl PositionalIndexer { /// Lazily compute prefix hashes up to `target_pos`. #[inline] - fn ensure_seq_hash_computed( + fn ensure_seq_hash_computed( seq_hashes: &mut Vec, target_pos: usize, - sequence: &[ContentHash], + sequence: &S, ) { while seq_hashes.len() <= target_pos { let pos = seq_hashes.len(); if pos == 0 { - seq_hashes.push(SequenceHash(sequence[0].0)); + seq_hashes.push(SequenceHash(sequence.at(0).0)); } else { let prev = seq_hashes[pos - 1].0; - let current = sequence[pos].0; + let current = sequence.at(pos).0; seq_hashes.push(SequenceHash(Self::compute_next_seq_hash(prev, current))); } } @@ -960,67 +1239,34 @@ impl PositionalIndexer { // Internal: query helpers // ----------------------------------------------------------------------- - /// Get workers at a position matching content_hash (and prefix_hash for Multi). - /// Copies worker IDs into a Vec — used only once at position 0 to initialize `active`. - /// Skips rolling hash computation for Single entries (unambiguous match). - fn get_workers_lazy( - index: &PosIndex, - position: usize, - content_hash: ContentHash, - seq_hashes: &mut Vec, - sequence: &[ContentHash], - now: u32, - ) -> Option> { - let entry = index.get(&(position, content_hash))?; - entry.value().touch(now); - if let Some(workers) = entry.value().seq.workers_if_single() { - return Some(workers.iter().copied().collect()); - } - // Multi: need rolling hash to disambiguate - Self::ensure_seq_hash_computed(seq_hashes, position, sequence); - entry - .value() - .seq - .get(seq_hashes[position]) - .map(|workers| workers.iter().copied().collect()) - } - - /// Count workers at a position matching the prefix_hash (no set materialization). - /// Skips rolling hash computation for Single entries (unambiguous match). - fn count_workers_at( + /// Workers holding the request's first block (position 0: same content hash and prefix + /// hash), appended to the caller's `active` buffer. False when no entry exists at all. + fn collect_workers_at_start( index: &PosIndex, - position: usize, - content_hash: ContentHash, seq_hashes: &mut Vec, - sequence: &[ContentHash], + sequence: &S, now: u32, - ) -> usize { - let Some(entry) = index.get(&(position, content_hash)) else { - return 0; + active: &mut Vec, + ) -> bool { + let Some(entry) = index.get(&(0, sequence.at(0))) else { + return false; }; entry.value().touch(now); - if let Some(workers) = entry.value().seq.workers_if_single() { - return workers.len(); + Self::ensure_seq_hash_computed(seq_hashes, 0, sequence); + if let Some(workers) = entry.value().seq.get(seq_hashes[0]) { + active.extend(workers.iter()); } - // Multi: need rolling hash to disambiguate - Self::ensure_seq_hash_computed(seq_hashes, position, sequence); - entry - .value() - .seq - .get(seq_hashes[position]) - .map(|workers| workers.len()) - .unwrap_or(0) + true } /// Scan positions sequentially, draining workers that stop matching. - /// Accesses DashMap entries directly — no set cloning. - /// Skips rolling hash computation for Single entries (unambiguous match). - /// Uses retain guard: skips retain when workers.len() >= active.len() - /// (all active workers are still present, no work to do). + /// Accesses DashMap entries directly — no set cloning. Every position is matched on + /// content hash and prefix hash, and every active worker is checked against the matching + /// set: a set at least as large as `active` may still lack one of its workers. #[expect(clippy::too_many_arguments)] - fn linear_scan_drain( + fn linear_scan_drain( index: &PosIndex, - sequence: &[ContentHash], + sequence: &S, seq_hashes: &mut Vec, active: &mut Vec, internal_scores: &mut FxHashMap, @@ -1029,11 +1275,11 @@ impl PositionalIndexer { early_exit: bool, now: u32, ) { - for (offset, &content_hash) in sequence[lo..hi].iter().enumerate() { + for pos in lo..hi { if active.is_empty() { break; } - let pos = lo + offset; + let content_hash = sequence.at(pos); let Some(entry) = index.get(&(pos, content_hash)) else { for &w in active.iter() { @@ -1043,30 +1289,6 @@ impl PositionalIndexer { break; }; entry.value().touch(now); - - // Fast path: Single entry — skip rolling hash, use workers directly. - if let Some(workers) = entry.value().seq.workers_if_single() { - // Retain guard: only retain when some workers - // have dropped off. When workers.len() >= active.len(), all active - // workers are still present — skip the O(active) iteration. - if workers.len() < active.len() { - let mut i = 0; - while i < active.len() { - if workers.contains(&active[i]) { - i += 1; - } else { - internal_scores.insert(active[i], pos as u32); - active.swap_remove(i); - } - } - } - if early_exit && !active.is_empty() { - break; - } - continue; - } - - // Multi: need rolling hash to disambiguate. Self::ensure_seq_hash_computed(seq_hashes, pos, sequence); let seq_hash = seq_hashes[pos]; @@ -1078,120 +1300,78 @@ impl PositionalIndexer { break; }; - // Retain guard: only iterate when some workers dropped off. - if workers.len() < active.len() { - let mut i = 0; - while i < active.len() { - if workers.contains(&active[i]) { - i += 1; - } else { - internal_scores.insert(active[i], pos as u32); - active.swap_remove(i); - } + let mut i = 0; + while i < active.len() { + if workers.contains(active[i]) { + i += 1; + } else { + internal_scores.insert(active[i], pos as u32); + active.swap_remove(i); } } - if early_exit && !active.is_empty() { break; } } } - fn jump_search_matches( + /// The lookup proper, over caller-provided scratch buffers (see [`LookupScratch`]). + fn scan_with( &self, - content_hashes: &[ContentHash], + sequence: &S, early_exit: bool, + seq_hashes: &mut Vec, + active: &mut Vec, ) -> OverlapScores { let mut scores = OverlapScores::default(); - if content_hashes.is_empty() { + if sequence.is_empty() { return scores; } - let mut seq_hashes = Vec::with_capacity(content_hashes.len()); let now = self.now_secs(); - let Some(initial_workers) = Self::get_workers_lazy( - &self.index, - 0, - content_hashes[0], - &mut seq_hashes, - content_hashes, - now, - ) else { - return scores; - }; - - let mut active = initial_workers; - if active.is_empty() { + if !Self::collect_workers_at_start(&self.index, seq_hashes, sequence, now, active) + || active.is_empty() + { return scores; } - let len = content_hashes.len(); + let len = sequence.len(); let mut internal_scores: FxHashMap = FxHashMap::default(); // Early exit: just record that workers matched at position 0. if early_exit { - for &w in &active { + for &w in active.iter() { internal_scores.insert(w, 1); } scores.scores = internal_scores; - for &int_id in scores.scores.keys() { - scores - .tree_sizes - .insert(int_id, self.tree_sizes.load(int_id)); - } return scores; } - let mut current_pos = 0; - - while current_pos < len - 1 && !active.is_empty() { - let next_pos = (current_pos + self.jump_size).min(len - 1); - - let count = Self::count_workers_at( - &self.index, - next_pos, - content_hashes[next_pos], - &mut seq_hashes, - content_hashes, - now, - ); - - // If the worker count at the jump destination matches the active set size, - // all active workers are still present — safe to skip intermediate positions. - if count == active.len() { - current_pos = next_pos; - } else { - Self::linear_scan_drain( - &self.index, - content_hashes, - &mut seq_hashes, - &mut active, - &mut internal_scores, - current_pos + 1, - next_pos + 1, - false, - now, - ); - current_pos = next_pos; - } - } + // Every position is verified in order. A worker's score is the length of the prefix it + // holds contiguously, so an evicted middle block (a hole) must stop it, and a landing + // check alone cannot see holes: the prefix hash stored at a later position was computed + // when the chain was intact. The jump shortcut is therefore not taken; exactness + // (guardrail 1) over the probe count. `jump_size` is kept for API compatibility. + Self::linear_scan_drain( + &self.index, + sequence, + seq_hashes, + active, + &mut internal_scores, + 1, + len, + false, + now, + ); let final_score = len as u32; - for &w in &active { + for &w in active.iter() { internal_scores.insert(w, final_score); } scores.scores = internal_scores; - - // Populate tree_sizes from atomic counters — lock-free array index. - for &int_id in scores.scores.keys() { - scores - .tree_sizes - .insert(int_id, self.tree_sizes.load(int_id)); - } - scores } } @@ -1214,6 +1394,67 @@ impl fmt::Debug for PositionalIndexer { #[cfg(test)] mod tests { + /// The gateway's parent-missing fallback: b4 and b5 stored without a parent land at + /// positions 0 and 1; when the engine announces the chain whole they move to 4 and 5 and + /// their old memberships go, so the count reads six and a query for the pair alone scores + /// nothing. + #[test] + fn a_hash_stored_again_at_another_position_releases_its_old_place() { + use super::{ + compute_content_hash, ContentHash, PositionalIndexer, SequenceHash, StoredBlock, + WorkerBlockMap, + }; + let indexer = PositionalIndexer::new(8); + let worker = indexer.intern_worker("w").expect("id"); + let mut map = WorkerBlockMap::default(); + let contents: Vec = (0..6) + .map(|p| compute_content_hash(&[31, p as u32])) + .collect(); + let mut prefix = None::; + let blocks: Vec = contents + .iter() + .map(|&content| { + let next = match prefix { + Some(prev) => PositionalIndexer::compute_next_seq_hash(prev, content.0), + None => content.0, + }; + prefix = Some(next); + StoredBlock { + seq_hash: SequenceHash(next), + content_hash: content, + } + }) + .collect(); + indexer + .apply_stored(worker, &blocks[..4], None, &mut map) + .expect("b0..b3"); + indexer.apply_removed(worker, &[blocks[2].seq_hash, blocks[3].seq_hash], &mut map); + indexer + .apply_stored(worker, &blocks[4..], None, &mut map) + .expect("fallback"); + assert_eq!(indexer.current_size(), 4); + assert_eq!( + indexer + .find_matches(&contents[4..], false) + .scores + .get(&worker), + Some(&2) + ); + indexer + .apply_stored(worker, &blocks, None, &mut map) + .expect("whole"); + assert_eq!(indexer.current_size(), 6); + assert_eq!(map.len(), 6); + assert_eq!( + indexer.find_matches(&contents, false).scores.get(&worker), + Some(&6) + ); + assert!(!indexer + .find_matches(&contents[4..], false) + .scores + .contains_key(&worker)); + } + /// The streaming hasher this function used before; the one-shot path must /// stay bit-identical because workers report hashes computed the old way. fn streaming_content_hash(token_ids: &[u32]) -> ContentHash { @@ -1278,6 +1519,36 @@ mod tests { values.iter().map(|&v| ContentHash(v)).collect() } + #[test] + fn worker_set_inline_and_spilled_ids() { + let mut set = WorkerSet::default(); + assert!(set.is_empty()); + for id in [0u32, 63, 64, 127, 128, 200, 10_000] { + assert!(set.insert(id), "{id} newly inserted"); + assert!(!set.insert(id), "{id} already present"); + assert!(set.contains(id)); + } + assert!(!set.contains(1)); + assert!(!set.contains(129)); + assert!(!set.contains(100_000)); + assert_eq!( + set.iter().collect::>(), + vec![0, 63, 64, 127, 128, 200, 10_000], + "iteration is ascending across the inline and spilled words" + ); + for id in [63u32, 128, 10_000] { + assert!(set.remove(id)); + assert!(!set.remove(id)); + assert!(!set.contains(id)); + } + assert_eq!(set.iter().collect::>(), vec![0, 64, 127, 200]); + for id in [0u32, 64, 127, 200] { + set.remove(id); + } + assert!(set.is_empty(), "all bits cleared, including spilled words"); + assert!(!WorkerSet::default().remove(5)); + } + #[test] fn test_new_indexer_is_empty() { let indexer = PositionalIndexer::default(); @@ -1296,7 +1567,7 @@ mod tests { let scores = indexer.find_matches(&hashes(&[10, 20, 30]), false); assert_eq!(scores.scores.get(&w1), Some(&3)); - assert_eq!(scores.tree_sizes.get(&w1), Some(&3)); + assert_eq!(indexer.worker_block_count(w1), 3); } #[test] @@ -1358,7 +1629,7 @@ mod tests { // After removing block at position 2, w1 should only match 2 blocks let scores = indexer.find_matches(&hashes(&[10, 20, 30]), false); assert_eq!(scores.scores.get(&w1), Some(&2)); - assert_eq!(scores.tree_sizes.get(&w1), Some(&2)); + assert_eq!(indexer.worker_block_count(w1), 2); } #[test] @@ -1400,9 +1671,8 @@ mod tests { .apply_stored(w2, &blocks_w2, None, &mut wb2) .unwrap(); - let scores = indexer.find_matches(&hashes(&[10]), false); - assert_eq!(scores.tree_sizes.get(&w1), Some(&3)); - assert_eq!(scores.tree_sizes.get(&w2), Some(&2)); + assert_eq!(indexer.worker_block_count(w1), 3); + assert_eq!(indexer.worker_block_count(w2), 2); } #[test] @@ -1432,7 +1702,7 @@ mod tests { let scores = indexer.find_matches(&hashes(&[10, 20, 30, 40]), false); assert_eq!(scores.scores.get(&w1), Some(&4)); - assert_eq!(scores.tree_sizes.get(&w1), Some(&4)); + assert_eq!(indexer.worker_block_count(w1), 4); } #[test] @@ -1578,6 +1848,49 @@ mod tests { } // ----------------------------------------------------------------------- + #[test] + fn find_matches_in_reads_plain_u64_hashes_without_copying() { + let indexer = PositionalIndexer::new(8); + let w1 = indexer.intern_worker("w1").unwrap(); + let mut wb1 = WorkerBlockMap::default(); + let blocks = make_blocks(&[10, 20, 30, 40]); + indexer + .apply_stored_iter(w1, blocks.iter().copied(), None, &mut wb1) + .unwrap(); + let typed = hashes(&[10, 20, 30, 40]); + let raw: Vec = typed.iter().map(|h| h.0).collect(); + assert_eq!( + indexer.find_matches_in(raw.as_slice(), false).scores, + indexer.find_matches(&typed, false).scores + ); + assert_eq!( + indexer + .find_matches_in(raw.as_slice(), false) + .scores + .get(&w1), + Some(&4) + ); + indexer.apply_removed_iter(w1, [blocks[3].seq_hash], &mut wb1); + assert_eq!( + indexer + .find_matches_in(raw.as_slice(), false) + .scores + .get(&w1), + Some(&3) + ); + assert!(indexer.find_matches_in(&raw[..0], false).scores.is_empty()); + } + + /// The index stores one `IndexEntry` per (position, content) key; its size sets the + /// bytes per unique block, so a regression here is a memory regression. + #[test] + #[cfg(target_pointer_width = "64")] + fn index_entry_is_48_bytes() { + assert_eq!(size_of::(), 24); + assert_eq!(size_of::(), 40); + assert_eq!(size_of::(), 48); + } + // Jump search edge cases // ----------------------------------------------------------------------- @@ -1759,8 +2072,8 @@ mod tests { let scores = indexer.find_matches(&hashes(&[10, 20, 30]), true); // early_exit: score is 1 (matched at position 0), not full depth assert_eq!(scores.scores.get(&w1), Some(&1)); - // tree_sizes still populated - assert_eq!(scores.tree_sizes.get(&w1), Some(&3)); + // the worker's block count does not depend on the lookup mode + assert_eq!(indexer.worker_block_count(w1), 3); } #[test] @@ -1813,9 +2126,7 @@ mod tests { indexer.apply_removed(w1, &[blocks[3].seq_hash, blocks[4].seq_hash], &mut wb1); assert_eq!(indexer.current_size(), 3); - // Verify tree_sizes in query results - let scores = indexer.find_matches(&hashes(&[10, 20, 30]), false); - assert_eq!(scores.tree_sizes.get(&w1), Some(&3)); + assert_eq!(indexer.worker_block_count(w1), 3); } #[test] @@ -1827,19 +2138,18 @@ mod tests { // First store: 3 new blocks indexer.apply_stored(w1, &blocks, None, &mut wb1).unwrap(); - let scores = indexer.find_matches(&hashes(&[10, 20, 30]), false); - assert_eq!(scores.tree_sizes.get(&w1), Some(&3)); + assert_eq!(indexer.worker_block_count(w1), 3); // Replay the same store event — tree_size must not change indexer.apply_stored(w1, &blocks, None, &mut wb1).unwrap(); - let scores = indexer.find_matches(&hashes(&[10, 20, 30]), false); assert_eq!( - scores.tree_sizes.get(&w1), - Some(&3), + indexer.worker_block_count(w1), + 3, "Duplicate store event must not inflate tree_size" ); // Overlap scores should also be unchanged + let scores = indexer.find_matches(&hashes(&[10, 20, 30]), false); assert_eq!(scores.scores.get(&w1), Some(&3)); } @@ -1939,7 +2249,7 @@ mod tests { let content_hashes = hashes(&content); let mut seq_hashes: Vec = Vec::new(); - PositionalIndexer::ensure_seq_hash_computed(&mut seq_hashes, 4, &content_hashes); + PositionalIndexer::ensure_seq_hash_computed(&mut seq_hashes, 4, content_hashes.as_slice()); for (i, block) in blocks.iter().enumerate() { assert_eq!( @@ -1959,7 +2269,7 @@ mod tests { let scores = indexer.find_matches(&hashes(&[10, 20]), false); assert_eq!(scores.scores.get(&w1), Some(&2)); - assert_eq!(scores.tree_sizes.get(&w1), Some(&5)); + assert_eq!(indexer.worker_block_count(w1), 5); } #[test] @@ -2129,7 +2439,7 @@ mod tests { let query_hashes = compute_request_content_hashes(&query_tokens, block_size); let scores = indexer.find_matches(&query_hashes, false); assert_eq!(scores.scores.get(&w1), Some(&2)); - assert_eq!(scores.tree_sizes.get(&w1), Some(&2)); + assert_eq!(indexer.worker_block_count(w1), 2); } #[test] @@ -2518,8 +2828,8 @@ mod tests { "worker at depth {depth} has wrong score" ); assert_eq!( - scores.tree_sizes.get(&wid), - Some(&depth), + indexer.worker_block_count(wid), + depth, "worker at depth {depth} has wrong tree_size" ); } @@ -2580,11 +2890,7 @@ mod tests { let scores = indexer.find_matches(&hashes(&shared), false); for &wid in &probe_ids { assert_eq!(scores.scores.get(&wid), Some(&3), "worker {wid} score"); - assert_eq!( - scores.tree_sizes.get(&wid), - Some(&4), - "worker {wid} tree_size" - ); + assert_eq!(indexer.worker_block_count(wid), 4, "worker {wid} tree_size"); } // Only worker 2049 has the tail block — the rest drain at depth 3. diff --git a/crates/kv_index/src/lane_map.rs b/crates/kv_index/src/lane_map.rs new file mode 100644 index 0000000000..cec4609d64 --- /dev/null +++ b/crates/kv_index/src/lane_map.rs @@ -0,0 +1,504 @@ +//! The lane map: one worker's blocks by engine hash, as the event lane that owns the worker +//! keeps them for the chain index. +//! +//! An open-addressing table of 16-byte slots (engine hash, run id, offset) with a one-byte tag +//! per slot beside them: Fibonacci home slot, linear probing over the tags (dense, so a probe +//! run of several slots reads one cache line and touches a slot only when its tag matches), at +//! most three quarters full so every probe run ends at an empty tag, doubling growth, and +//! backward-shift deletion so steady churn (an engine storing and evicting at the same rate for +//! hours) never accumulates tombstones. Removals and insertions come in batches of an event's +//! blocks: the home slots of the next keys are prefetched a few keys ahead so their cache misses +//! overlap instead of serialising. +//! +//! Semantics match a hash map: a key maps to at most one place, `insert` replaces, `remove` of an +//! absent key is a no-op, iteration yields every entry once. A differential test against a hash +//! map model under random operations keeps the backshift right. + +use crate::{chain_index::BlockRef, event_tree::SequenceHash}; + +/// Smallest table a non-empty map allocates. +const MIN_SLOTS: usize = 16; +/// How many keys ahead a batch touches. +const AHEAD: usize = 8; +/// Run id that marks an empty slot (never a live location: run ids are below 2^26). +const EMPTY: u32 = u32::MAX; +/// Tag of an empty slot; a full slot's tag is seven bits of its key with the top bit set. +const VACANT: u8 = 0; + +#[derive(Clone, Copy)] +struct Slot { + key: u64, + run: u32, + offset: u32, +} + +const EMPTY_SLOT: Slot = Slot { + key: 0, + run: EMPTY, + offset: 0, +}; + +/// One worker's blocks by engine hash: where each lives in the chain index. +pub struct ChainBlockMap { + /// Power-of-two length, or empty before the first insert. + slots: Box<[Slot]>, + /// One tag per slot: `VACANT`, or the key's fingerprint. + tags: Box<[u8]>, + len: usize, + shift: u32, +} + +impl Default for ChainBlockMap { + fn default() -> Self { + Self { + slots: Box::default(), + tags: Box::default(), + len: 0, + shift: 64, + } + } +} + +/// Seven bits of the key the home slot does not use, with the top bit set so it is never +/// `VACANT`. +#[inline] +fn fingerprint(key: u64) -> u8 { + ((key >> 25) as u8) | 0x80 +} + +impl ChainBlockMap { + pub fn new() -> Self { + Self::default() + } + + /// Entries held. + pub fn len(&self) -> usize { + self.len + } + + pub fn is_empty(&self) -> bool { + self.len == 0 + } + + /// Slots allocated (16 bytes each, plus a tag byte). + pub fn capacity(&self) -> usize { + self.slots.len() + } + + #[inline] + fn home(&self, key: u64) -> usize { + // Fibonacci hashing spreads structured keys; engine hashes are already uniform. + (key.wrapping_mul(0x9E37_79B9_7F4A_7C15) >> self.shift) as usize + } + + #[inline] + fn mask(&self) -> usize { + self.slots.len() - 1 + } + + /// The slot holding `key`, or else the empty slot that ends its probe run. Requires slots. + /// Walks the tags; a slot is read only when its tag matches. + #[inline] + fn probe(&self, key: u64) -> Result { + let mask = self.mask(); + let tag = fingerprint(key); + let mut index = self.home(key); + loop { + let found = self.tags[index]; + if found == VACANT { + return Err(index); + } + if found == tag && self.slots[index].key == key { + return Ok(index); + } + index = (index + 1) & mask; + } + } + + fn find(&self, key: u64) -> Option { + if self.len == 0 { + return None; + } + self.probe(key).ok() + } + + pub fn contains_key(&self, key: SequenceHash) -> bool { + self.find(key.0).is_some() + } + + pub fn get(&self, key: SequenceHash) -> Option { + self.find(key.0).map(|index| { + let slot = &self.slots[index]; + BlockRef { + run: slot.run, + offset: slot.offset, + } + }) + } + + /// Room for `additional` more entries under three-quarters load. + fn reserve(&mut self, additional: usize) { + let needed = self.len + additional; + if needed * 4 <= self.slots.len() * 3 { + return; + } + let mut capacity = self.slots.len().max(MIN_SLOTS); + while needed * 4 > capacity * 3 { + capacity *= 2; + } + let old = std::mem::replace(&mut self.slots, (0..capacity).map(|_| EMPTY_SLOT).collect()); + self.tags = (0..capacity).map(|_| VACANT).collect(); + self.shift = 64 - capacity.trailing_zeros(); + let mask = capacity - 1; + for slot in old.iter().filter(|slot| slot.run != EMPTY) { + let mut index = self.home(slot.key); + while self.tags[index] != VACANT { + index = (index + 1) & mask; + } + self.slots[index] = *slot; + self.tags[index] = fingerprint(slot.key); + } + } + + /// Insert or replace; the previous location when the key was present. + pub fn insert(&mut self, key: SequenceHash, at: BlockRef) -> Option { + self.reserve(1); + match self.probe(key.0) { + Ok(index) => { + let slot = &mut self.slots[index]; + let previous = BlockRef { + run: slot.run, + offset: slot.offset, + }; + slot.run = at.run; + slot.offset = at.offset; + Some(previous) + } + Err(empty) => { + self.slots[empty] = Slot { + key: key.0, + run: at.run, + offset: at.offset, + }; + self.tags[empty] = fingerprint(key.0); + self.len += 1; + None + } + } + } + + /// Remove `key`; its location when it was present. + pub fn remove(&mut self, key: SequenceHash) -> Option { + let mut hole = self.find(key.0)?; + let removed = BlockRef { + run: self.slots[hole].run, + offset: self.slots[hole].offset, + }; + self.slots[hole] = EMPTY_SLOT; + self.tags[hole] = VACANT; + self.len -= 1; + // Backward-shift deletion: pull later entries of the probe run into the hole whenever + // the hole lies between their home slot and their current slot. + let mask = self.mask(); + let mut next = hole; + loop { + next = (next + 1) & mask; + if self.tags[next] == VACANT { + break; + } + let slot = self.slots[next]; + let home = self.home(slot.key); + if (next.wrapping_sub(home) & mask) >= (next.wrapping_sub(hole) & mask) { + self.slots[hole] = slot; + self.tags[hole] = self.tags[next]; + self.slots[next] = EMPTY_SLOT; + self.tags[next] = VACANT; + hole = next; + } + } + Some(removed) + } + + /// Ask the cache for the home slot of `key` before the key is used: a prefetch hint, so the + /// misses of a batch overlap instead of serialising. + #[inline] + fn touch(&self, key: u64) { + if !self.slots.is_empty() { + let home = self.home(key); + crate::prefetch::prefetch_read(&self.tags[home]); + crate::prefetch::prefetch_read(&self.slots[home]); + } + } + + /// Remove every key of a batch, reporting each present one's location, with the home slots + /// touched `AHEAD` keys early. + pub(crate) fn remove_all( + &mut self, + keys: &[SequenceHash], + mut on_removed: impl FnMut(BlockRef), + ) { + if self.len == 0 { + return; + } + for key in &keys[..keys.len().min(AHEAD)] { + self.touch(key.0); + } + for (index, key) in keys.iter().enumerate() { + if let Some(next) = keys.get(index + AHEAD) { + self.touch(next.0); + } + if let Some(at) = self.remove(*key) { + on_removed(at); + } + } + } + + /// Insert `count` consecutive blocks starting at `first` for the keys given, with the home + /// slots touched `AHEAD` keys early: the lane map writes of one store. A key already present + /// moves to its new place and `on_moved` sees its old and new places (a block the engine + /// stored again at the same place passes through here unchanged). + pub(crate) fn insert_run( + &mut self, + keys: impl ExactSizeIterator + Clone, + first: BlockRef, + mut on_moved: impl FnMut(BlockRef, BlockRef), + ) { + self.reserve(keys.len()); + let mut ahead = keys.clone(); + for key in ahead.by_ref().take(AHEAD) { + self.touch(key.0); + } + for (index, key) in keys.enumerate() { + if let Some(next) = ahead.next() { + self.touch(next.0); + } + let at = BlockRef { + run: first.run, + offset: first.offset + index as u32, + }; + match self.probe(key.0) { + Ok(slot) => { + let old = BlockRef { + run: self.slots[slot].run, + offset: self.slots[slot].offset, + }; + if old != at { + on_moved(old, at); + } + self.slots[slot].run = at.run; + self.slots[slot].offset = at.offset; + } + Err(empty) => { + self.slots[empty] = Slot { + key: key.0, + run: at.run, + offset: at.offset, + }; + self.tags[empty] = fingerprint(key.0); + self.len += 1; + } + } + } + } + + /// Every entry, in table order. + pub fn iter(&self) -> impl Iterator + '_ { + self.slots + .iter() + .filter(|slot| slot.run != EMPTY) + .map(|slot| { + ( + SequenceHash(slot.key), + BlockRef { + run: slot.run, + offset: slot.offset, + }, + ) + }) + } +} + +impl IntoIterator for ChainBlockMap { + type Item = (SequenceHash, BlockRef); + type IntoIter = std::vec::IntoIter<(SequenceHash, BlockRef)>; + + fn into_iter(self) -> Self::IntoIter { + self.iter().collect::>().into_iter() + } +} + +impl std::fmt::Debug for ChainBlockMap { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("ChainBlockMap") + .field("len", &self.len) + .field("slots", &self.slots.len()) + .finish() + } +} + +#[cfg(test)] +mod tests { + use rustc_hash::FxHashMap; + + use super::*; + + struct Rng(u64); + + impl Rng { + fn next(&mut self) -> u64 { + let mut x = self.0; + x ^= x >> 12; + x ^= x << 25; + x ^= x >> 27; + self.0 = x; + x.wrapping_mul(0x2545_F491_4F6C_DD1D) + } + } + + fn at(n: u64) -> BlockRef { + BlockRef { + run: (n % 1000) as u32, + offset: (n % 97) as u32, + } + } + + /// Random inserts, replacements, single and batched removals (present and absent keys), and + /// batched runs of inserts against a hash map model; the table must agree after every step + /// and at the end, entry for entry. + #[test] + fn matches_a_hash_map_under_random_operations() { + let mut rng = Rng(20261005); + let mut table = ChainBlockMap::default(); + let mut model: FxHashMap = FxHashMap::default(); + // Keys from a small universe so the same ones come back after removal and probe runs + // overlap; a few structured keys (multiples of a power of two) stress the home hash. + let key = |r: &mut Rng| { + let n = r.next() % 6000; + SequenceHash(if n.is_multiple_of(7) { + n << 40 + } else { + n.wrapping_mul(0x9E37_79B9) + }) + }; + for step in 0..200_000u64 { + match rng.next() % 10 { + 0..=3 => { + let k = key(&mut rng); + let v = at(rng.next()); + assert_eq!( + table.insert(k, v), + model.insert(k, v), + "insert at step {step}" + ); + } + 4..=5 => { + let k = key(&mut rng); + assert_eq!(table.remove(k), model.remove(&k), "remove at step {step}"); + } + 6 => { + let keys: Vec = + (0..(rng.next() % 40)).map(|_| key(&mut rng)).collect(); + let mut removed = 0usize; + table.remove_all(&keys, |_| removed += 1); + let expected = keys.iter().filter(|k| model.remove(k).is_some()).count(); + assert_eq!(removed, expected, "batch removal at step {step}"); + } + 7 => { + let keys: Vec = + (0..(rng.next() % 60)).map(|_| key(&mut rng)).collect(); + let first = at(rng.next()); + table.insert_run(keys.iter().copied(), first, |_, _| {}); + for (index, k) in keys.iter().enumerate() { + model.insert( + *k, + BlockRef { + run: first.run, + offset: first.offset + index as u32, + }, + ); + } + } + _ => { + let k = key(&mut rng); + assert_eq!(table.get(k), model.get(&k).copied(), "get at step {step}"); + assert_eq!(table.contains_key(k), model.contains_key(&k)); + } + } + assert_eq!(table.len(), model.len(), "len at step {step}"); + assert!(table.len() * 4 <= table.capacity() * 3 || table.capacity() == 0); + } + let mut seen: Vec<(SequenceHash, BlockRef)> = table.iter().collect(); + seen.sort_by_key(|(k, _)| k.0); + let mut expected: Vec<(SequenceHash, BlockRef)> = + model.iter().map(|(k, v)| (*k, *v)).collect(); + expected.sort_by_key(|(k, _)| k.0); + assert_eq!(seen, expected); + let drained: FxHashMap = table.into_iter().collect(); + assert_eq!(drained, model); + } + + /// Probe lengths at the load the table runs at: how many slots a hit walks and how many an + /// insert walks to its empty slot, for uniform keys at just under three quarters full. + #[test] + fn probe_lengths_at_three_quarters_load() { + let mut rng = Rng(7); + let mut table = ChainBlockMap::default(); + let keys: Vec = (0..24_000).map(|_| SequenceHash(rng.next())).collect(); + for (index, key) in keys.iter().enumerate() { + table.insert(*key, at(index as u64)); + } + assert_eq!(table.capacity(), 32_768); + let mask = table.mask(); + let mut hit = [0usize; 65]; + let mut miss = [0usize; 65]; + for key in &keys { + let home = table.home(key.0); + let found = table.probe(key.0).expect("present"); + hit[(found.wrapping_sub(home) & mask).min(64)] += 1; + } + for _ in 0..24_000 { + let key = rng.next(); + let home = table.home(key); + let empty = table.probe(key).expect_err("absent"); + miss[(empty.wrapping_sub(home) & mask).min(64)] += 1; + } + let mean = |h: &[usize; 65]| { + h.iter().enumerate().map(|(d, n)| d * n).sum::() as f64 + / h.iter().sum::() as f64 + }; + let tail = |h: &[usize; 65], d: usize| { + h[d..].iter().sum::() as f64 / h.iter().sum::() as f64 + }; + // Measured on this layout: hits 1.37 slots, inserts 6.16 with 11% past 16 slots. + assert!( + mean(&hit) < 4.0, + "hit probe mean {:.2} (>4: {:.1}%, >16: {:.2}%)", + mean(&hit), + 100.0 * tail(&hit, 5), + 100.0 * tail(&hit, 17) + ); + assert!( + mean(&miss) < 16.0, + "insert probe mean {:.2} (>4: {:.1}%, >16: {:.2}%)", + mean(&miss), + 100.0 * tail(&miss, 5), + 100.0 * tail(&miss, 17) + ); + } + + #[test] + fn empty_map_answers_without_slots() { + let mut table = ChainBlockMap::default(); + assert_eq!(table.capacity(), 0); + assert!(table.get(SequenceHash(1)).is_none()); + assert!(table.remove(SequenceHash(1)).is_none()); + table.remove_all(&[SequenceHash(1), SequenceHash(2)], |_| { + panic!("nothing to remove") + }); + table.insert_run(std::iter::empty(), at(0), |_, _| {}); + assert!(table.is_empty()); + assert!(table.insert(SequenceHash(1), at(1)).is_none()); + assert_eq!(table.capacity(), MIN_SLOTS); + assert_eq!(table.len(), 1); + } +} diff --git a/crates/kv_index/src/lane_pool.rs b/crates/kv_index/src/lane_pool.rs new file mode 100644 index 0000000000..e5dc2906b6 --- /dev/null +++ b/crates/kv_index/src/lane_pool.rs @@ -0,0 +1,870 @@ +//! Lane pool: per-worker FIFO queues served by a fixed set of lanes that steal whole workers. +//! +//! Events of one worker (an engine rank) apply in the order the engine sent them, so the unit of +//! scheduling is the worker, never the event: a worker with queued events sits in exactly one +//! lane's ready list, and the lane that takes it from there applies its events alone until the +//! queue drains or the lane's batch is up. Any lane may take any ready worker, so a lane whose +//! own workers are quiet drains the backlog of a busy one, and no event waits behind another +//! worker's events. The per-worker state (SMG's interned id and the worker's block map) moves with +//! the worker: the claim transitions hand it from lane to lane. +//! +//! # Producers +//! +//! [`LanePool::enqueue`] never blocks and never drops. A worker holds at most `depth_cap` queued +//! events; past that the event comes back as [`QueueFull`] and the producer applies its rule: +//! +//! - the gateway's event monitor stops reading that rank's stream (its per-rank cursor does not +//! advance, so the engine sees flow control) and retries once [`LanePool::depth`] fell, or +//! declares the rank stale and resyncs through the recovery path. Dropping one event is never +//! an option: the index would stay wrong for as long as the block lives; +//! - the replay adapters hold the event in a per-lane backlog and leave the rest in the harness's +//! own queue, which is what the harness's queue-depth row measures. +//! +//! [`PoolMetrics::rejected`] counts the refusals, `max_depth` the deepest worker queue seen. +//! +//! # Lanes +//! +//! A lane is any thread that calls [`LanePool::run_lane`]: in the replay harnesses the harness's +//! own lane threads (thread counts, pinning and per-lane CPU accounting stay the harness's), in the +//! gateway threads of the pool's owner. Each turn a lane runs the caller's non-blocking `pump` +//! (ingress of the caller's choosing, usually draining a channel into `enqueue`), takes the oldest +//! ready worker of its own list or steals the oldest of a lane that is busy applying another +//! worker (a lane that is not serving empties its own list within microseconds, so taking from +//! it would only move a worker's cache state between cores), applies up to `batch` events +//! and releases the worker; with nothing ready anywhere it runs the caller's `wait` (blocking +//! ingress, or [`LanePool::park_lane`]); the three come as one [`LaneHooks`] value. Idle lanes +//! cost nothing: a parked lane is woken by the enqueue that gives it work, directly when the +//! worker lands on its own list and as the chosen thief when it lands on a busy lane's list, so +//! the gateway's `wait` is a park with a long safety timeout rather than a poll. +//! +//! # Gateway handoff +//! +//! `KvEventMonitor::process_stream` applies the events of a batch inline today. With the pool it +//! enqueues them under the rank's slot, advances its sequence cursor only once the whole batch is +//! queued, and the rank's `WorkerIndexState` becomes the pool's per-worker state; removing a worker +//! is one more queued event that the apply closure answers by taking the state out. Nothing else +//! in the monitor changes. + +use std::{ + collections::VecDeque, + fmt, + sync::{ + atomic::{AtomicBool, AtomicU64, AtomicU8, AtomicUsize, Ordering}, + Mutex, MutexGuard, PoisonError, + }, + thread::Thread, + time::{Duration, Instant}, +}; + +use crossbeam_utils::CachePadded; + +/// No queued events and not in any ready list. +const IDLE: u8 = 0; +/// In exactly one lane's ready list. +const READY: u8 = 1; +/// One lane is applying the worker's events. +const RUNNING: u8 = 2; + +/// Shape of a pool. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub struct LanePoolConfig { + /// Threads that will call [`LanePool::run_lane`], each with its own ready list. + pub lanes: usize, + /// Worker slots; `enqueue` and `depth` take a slot index below this. + pub max_workers: usize, + /// Queued events a worker may hold; `enqueue` refuses the next one. + pub depth_cap: usize, + /// Events a lane applies from one worker before releasing it, so a stolen backlog does not + /// starve the lane's own ready workers. + pub batch: usize, + /// Events a ready worker may have queued on a lane that is not serving before any lane may + /// take it: the case of a lane that is not running (preempted, or pinned to a core something + /// else holds), which today's rule never reaches because only a serving lane with more on its + /// list is stolen from. Zero, the default, keeps today's rule exactly. The bound is in events + /// queued, never in time, so a quiet pool never steals and a stalled lane's backlog is + /// taken once it is `steal_after` deep. + pub steal_after: usize, +} + +/// The event handed back by [`LanePool::enqueue`] when the worker's queue is at its cap. +pub struct QueueFull(pub E); + +impl fmt::Debug for QueueFull { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str("QueueFull") + } +} + +/// What a lane does after its `pump` or `wait` hook. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum Control { + Continue, + Stop, +} + +/// A worker handed to [`LaneHooks::apply`]: its slot index and its state, which is `None` until +/// the hook creates it and `None` again once the hook takes it out. +pub struct Claimed<'a, W> { + pub worker: u32, + pub state: &'a mut Option, +} + +/// What a lane thread brings to [`LanePool::run_lane`]: how to apply an event and where its +/// ingress comes from. +pub trait LaneHooks { + /// Apply one event of one worker; the pool holds the worker for this call alone. + fn apply(&mut self, claimed: Claimed<'_, W>, event: E); + + /// Non-blocking ingress at the top of every turn (drain a channel into + /// [`LanePool::enqueue`], for instance). + fn pump(&mut self) -> Control { + Control::Continue + } + + /// No worker is ready on any lane: block for a bounded time (on the ingress, or in + /// [`LanePool::park_lane`]). + fn wait(&mut self) -> Control; +} + +/// Closures as hooks: `(apply, pump, wait)`. +impl LaneHooks for (A, P, T) +where + A: FnMut(Claimed<'_, W>, E), + P: FnMut() -> Control, + T: FnMut() -> Control, +{ + fn apply(&mut self, claimed: Claimed<'_, W>, event: E) { + (self.0)(claimed, event); + } + + fn pump(&mut self) -> Control { + (self.1)() + } + + fn wait(&mut self) -> Control { + (self.2)() + } +} + +/// Counters of a pool, summed over lanes. +#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)] +pub struct PoolMetrics { + pub enqueued: u64, + pub applied: u64, + /// Events refused by the depth cap (handed back, never dropped). + pub rejected: u64, + /// Ready workers taken from another lane's list. + pub steals: u64, + /// Events waiting now (those being applied by a lane count until the batch ends). + pub queued: usize, + /// Deepest queue one worker reached. + pub max_depth: usize, + /// Most events waiting at once across all workers (sampled every few hundred batches per + /// lane). + pub max_queued: usize, + /// Time lanes spent inside the apply closure. + pub busy_ns: u64, + /// Time lanes spent inside their `wait` hook. + pub idle_ns: u64, +} + +/// Queue capacity a worker keeps once its burst has drained. +const KEEP_CAPACITY: usize = 256; + +struct WorkerSlot { + state: AtomicU8, + queue: Mutex>, + /// Only the lane that moved the slot from READY to RUNNING takes this lock, so it is never + /// contended; it is a lock rather than a cell so the exclusivity is the type system's. + held: Mutex>, +} + +/// One lane's list and counters. Every counter is the lane's own, so the hot paths touch no line +/// another lane writes; the pool-wide figures are sums. +struct Lane { + ready: Mutex>, + parked: AtomicBool, + /// Inside `serve`: the lane is applying one worker's batch, so the others on its list wait. + /// Thieves take from serving lanes only; a lane that is pumping or waiting empties its own + /// list within microseconds, and moving a worker's state to another core for that would + /// cost more than it saves. + serving: AtomicBool, + thread: Mutex>, + busy_ns: AtomicU64, + idle_ns: AtomicU64, + /// Events this lane enqueued, refused, and applied. + enqueued: AtomicU64, + rejected: AtomicU64, + applied: AtomicU64, + steals: AtomicU64, + max_depth: AtomicUsize, + max_queued: AtomicUsize, +} + +/// Batches between two samples of the pool-wide queue depth by one lane (a sum over every +/// lane's counters, so not something to take per event). +const SAMPLE_EVERY: u64 = 256; + +pub struct LanePool { + config: LanePoolConfig, + workers: Box<[CachePadded>]>, + lanes: Box<[CachePadded]>, + /// One bit per lane: set while the lane is serving a worker and has more on its list, which + /// is exactly where a thief can take from. A hint: a stale bit costs the thief one empty + /// look, a missed one delays a steal by a batch. + stealable: Box<[CachePadded]>, +} + +fn lock(mutex: &Mutex) -> MutexGuard<'_, T> { + mutex.lock().unwrap_or_else(PoisonError::into_inner) +} + +impl LanePool { + /// # Panics + /// + /// When the configuration has no lanes, no workers, a zero cap or a zero batch. + #[must_use] + pub fn new(config: LanePoolConfig) -> Self { + assert!(config.lanes > 0, "a lane pool needs at least one lane"); + assert!( + config.max_workers > 0, + "a lane pool needs at least one worker slot" + ); + assert!( + config.depth_cap > 0, + "a lane pool needs a positive depth cap" + ); + assert!(config.batch > 0, "a lane pool needs a positive batch"); + Self { + config, + workers: (0..config.max_workers) + .map(|_| { + CachePadded::new(WorkerSlot { + state: AtomicU8::new(IDLE), + queue: Mutex::new(VecDeque::new()), + held: Mutex::new(None), + }) + }) + .collect(), + stealable: (0..config.lanes.div_ceil(64)) + .map(|_| CachePadded::new(AtomicU64::new(0))) + .collect(), + lanes: (0..config.lanes) + .map(|_| { + CachePadded::new(Lane { + ready: Mutex::new(VecDeque::new()), + parked: AtomicBool::new(false), + serving: AtomicBool::new(false), + thread: Mutex::new(None), + busy_ns: AtomicU64::new(0), + idle_ns: AtomicU64::new(0), + enqueued: AtomicU64::new(0), + rejected: AtomicU64::new(0), + applied: AtomicU64::new(0), + steals: AtomicU64::new(0), + max_depth: AtomicUsize::new(0), + max_queued: AtomicUsize::new(0), + }) + }) + .collect(), + } + } + + #[must_use] + pub fn config(&self) -> &LanePoolConfig { + &self.config + } + + /// Queue `event` for `worker`; an idle worker becomes ready on `lane`'s list. + /// + /// Never blocks and never drops: at `depth_cap` queued events the event comes back and the + /// caller holds it (see the module documentation for the rule each producer applies). + /// + /// # Errors + /// + /// [`QueueFull`] with the event when the worker's queue is at the cap. + /// + /// # Panics + /// + /// When `lane` or `worker` is out of range. + pub fn enqueue(&self, lane: usize, worker: u32, event: E) -> Result<(), QueueFull> { + let slot = &self.workers[worker as usize]; + let me = &self.lanes[lane]; + let depth = { + let mut queue = lock(&slot.queue); + if queue.len() >= self.config.depth_cap { + drop(queue); + me.rejected.fetch_add(1, Ordering::Relaxed); + return Err(QueueFull(event)); + } + queue.push_back(event); + queue.len() + }; + me.enqueued.fetch_add(1, Ordering::Relaxed); + me.max_depth.fetch_max(depth, Ordering::Relaxed); + if slot + .state + .compare_exchange(IDLE, READY, Ordering::AcqRel, Ordering::Relaxed) + .is_ok() + { + self.make_ready(lane, worker); + } + Ok(()) + } + + /// Events queued for `worker` now. + #[must_use] + pub fn depth(&self, worker: u32) -> usize { + lock(&self.workers[worker as usize].queue).len() + } + + /// Events queued across all workers now (enqueued minus applied, summed over lanes). + #[must_use] + pub fn queued(&self) -> usize { + let (mut enqueued, mut applied) = (0u64, 0u64); + for lane in &self.lanes { + enqueued += lane.enqueued.load(Ordering::Relaxed); + applied += lane.applied.load(Ordering::Relaxed); + } + usize::try_from(enqueued.saturating_sub(applied)).unwrap_or(usize::MAX) + } + + #[must_use] + pub fn metrics(&self) -> PoolMetrics { + let mut metrics = PoolMetrics::default(); + for lane in &self.lanes { + metrics.enqueued += lane.enqueued.load(Ordering::Relaxed); + metrics.rejected += lane.rejected.load(Ordering::Relaxed); + metrics.applied += lane.applied.load(Ordering::Relaxed); + metrics.steals += lane.steals.load(Ordering::Relaxed); + metrics.busy_ns += lane.busy_ns.load(Ordering::Relaxed); + metrics.idle_ns += lane.idle_ns.load(Ordering::Relaxed); + metrics.max_depth = metrics + .max_depth + .max(lane.max_depth.load(Ordering::Relaxed)); + metrics.max_queued = metrics + .max_queued + .max(lane.max_queued.load(Ordering::Relaxed)); + } + metrics.queued = metrics + .enqueued + .saturating_sub(metrics.applied) + .try_into() + .unwrap_or(usize::MAX); + metrics + } + + fn make_ready(&self, lane: usize, worker: u32) { + let target = &self.lanes[lane]; + lock(&target.ready).push_back(worker); + if target.serving.load(Ordering::Relaxed) { + self.mark_stealable(lane, true); + } + if target.parked.swap(false, Ordering::SeqCst) { + Self::unpark(target); + } else { + // The target lane is busy: wake one parked lane, which finds its own list empty and + // steals the worker. Clearing the flag on the sleeper's behalf makes each wake-up + // reach a different lane. + for (index, other) in self.lanes.iter().enumerate() { + if index != lane && other.parked.swap(false, Ordering::SeqCst) { + Self::unpark(other); + break; + } + } + } + } + + fn unpark(lane: &Lane) { + if let Some(thread) = lock(&lane.thread).as_ref() { + thread.unpark(); + } + } + + /// Serve `lane` until the hooks' `pump` or `wait` returns [`Control::Stop`]. + /// + /// `pump` runs at the top of every turn and must not block; `wait` runs when no worker is + /// ready on any lane and may block for a bounded time; `apply` sees one worker at a time with + /// exclusive access to that worker's state. + /// + /// # Panics + /// + /// When `lane` is out of range. + pub fn run_lane(&self, lane: usize, hooks: &mut impl LaneHooks) { + let me = &self.lanes[lane]; + *lock(&me.thread) = Some(std::thread::current()); + let mut seed = (lane as u64).wrapping_mul(0x9E37_79B9_7F4A_7C15) | 1; + loop { + if hooks.pump() == Control::Stop { + break; + } + let next = lock(&me.ready) + .pop_front() + .or_else(|| self.steal(lane, &mut seed)); + let Some(worker) = next else { + let started = Instant::now(); + let control = hooks.wait(); + me.idle_ns + .fetch_add(started.elapsed().as_nanos() as u64, Ordering::Relaxed); + if control == Control::Stop { + break; + } + continue; + }; + self.serve(lane, worker, hooks); + } + *lock(&me.thread) = None; + } + + /// Sleep until an enqueue makes a worker ready (on this lane's list, or on a busy lane's list + /// with this lane chosen to steal it) or `timeout` passes. The wake is a `Thread::unpark` + /// from the producer, so an idle fleet costs no wake-ups; the timeout is the safety net. + pub fn park_lane(&self, lane: usize, timeout: Duration) { + let me = &self.lanes[lane]; + me.parked.store(true, Ordering::SeqCst); + // An enqueue that pushed before this check is seen here; one that pushes after it sees + // the flag (both go through the list's lock) and unparks, which a later park consumes. + if lock(&me.ready).is_empty() { + std::thread::park_timeout(timeout); + } + me.parked.store(false, Ordering::SeqCst); + } + + /// Whether some lane is serving with more workers on its list: the only reason for an idle + /// lane to wake before its own ingress does. One word per 64 lanes. + #[must_use] + pub fn has_stealable(&self) -> bool { + self.stealable + .iter() + .any(|word| word.load(Ordering::Relaxed) != 0) + } + + fn mark_stealable(&self, lane: usize, on: bool) { + let word = &self.stealable[lane / 64]; + let bit = 1u64 << (lane % 64); + if on { + word.fetch_or(bit, Ordering::Relaxed); + } else { + word.fetch_and(!bit, Ordering::Relaxed); + } + } + + /// Take the oldest ready worker of a lane whose bit says it is serving with more on its + /// list; one word read when there is none. + fn steal(&self, lane: usize, seed: &mut u64) -> Option { + if self.lanes.len() == 1 { + return None; + } + *seed ^= *seed << 13; + *seed ^= *seed >> 7; + *seed ^= *seed << 17; + let start = (*seed % self.stealable.len() as u64) as usize; + for step in 0..self.stealable.len() { + let index = (start + step) % self.stealable.len(); + let mut word = self.stealable[index].load(Ordering::Relaxed); + if index == lane / 64 { + word &= !(1u64 << (lane % 64)); + } + while word != 0 { + let bit = word.trailing_zeros() as usize; + word &= word - 1; + let victim = index * 64 + bit; + let Ok(mut ready) = self.lanes[victim].ready.try_lock() else { + continue; + }; + let taken = ready.pop_front(); + if ready.is_empty() { + self.mark_stealable(victim, false); + } + drop(ready); + if let Some(worker) = taken { + self.lanes[lane].steals.fetch_add(1, Ordering::Relaxed); + return Some(worker); + } + } + } + if self.config.steal_after == 0 { + return None; + } + // Nothing is stealable by the serving rule: look for a ready worker waiting on a lane that + // is not serving with `steal_after` events behind it, which is a lane that is not running. + // One try-lock per lane of the pool, only on an idle turn that found nothing above. + let count = self.lanes.len(); + let start = (*seed % count as u64) as usize; + for step in 0..count { + let victim = (start + step) % count; + if victim == lane || self.lanes[victim].serving.load(Ordering::Relaxed) { + continue; + } + let Ok(mut ready) = self.lanes[victim].ready.try_lock() else { + continue; + }; + let Some(&worker) = ready.front() else { + continue; + }; + if lock(&self.workers[worker as usize].queue).len() < self.config.steal_after { + continue; + } + ready.pop_front(); + drop(ready); + self.lanes[lane].steals.fetch_add(1, Ordering::Relaxed); + return Some(worker); + } + None + } + + fn serve(&self, lane: usize, worker: u32, hooks: &mut impl LaneHooks) { + let slot = &self.workers[worker as usize]; + let claimed = slot + .state + .compare_exchange(READY, RUNNING, Ordering::AcqRel, Ordering::Relaxed) + .is_ok(); + debug_assert!(claimed, "a ready worker is claimed by exactly one lane"); + let me = &self.lanes[lane]; + me.serving.store(true, Ordering::Relaxed); + if !lock(&me.ready).is_empty() { + self.mark_stealable(lane, true); + } + let started = Instant::now(); + let mut applied = 0; + { + let mut held = lock(&slot.held); + while applied < self.config.batch { + let event = { + let mut queue = lock(&slot.queue); + let event = queue.pop_front(); + if event.is_none() && queue.capacity() > KEEP_CAPACITY { + // The burst is over: give its buffer back. + queue.shrink_to(KEEP_CAPACITY); + } + event + }; + let Some(event) = event else { + break; + }; + hooks.apply( + Claimed { + worker, + state: &mut held, + }, + event, + ); + applied += 1; + } + } + me.busy_ns + .fetch_add(started.elapsed().as_nanos() as u64, Ordering::Relaxed); + let total = me.applied.fetch_add(applied as u64, Ordering::Relaxed) + applied as u64; + me.serving.store(false, Ordering::Relaxed); + self.mark_stealable(lane, false); + if total / SAMPLE_EVERY != (total - applied as u64) / SAMPLE_EVERY { + me.max_queued.fetch_max(self.queued(), Ordering::Relaxed); + } + slot.state.store(IDLE, Ordering::Release); + // Whoever observes the queue non-empty after this point schedules the worker once: an + // enqueue that saw RUNNING left scheduling to us, one that comes later wins the exchange. + let pending = !lock(&slot.queue).is_empty(); + if pending + && slot + .state + .compare_exchange(IDLE, READY, Ordering::AcqRel, Ordering::Relaxed) + .is_ok() + { + self.make_ready(lane, worker); + } + } +} + +#[cfg(test)] +mod tests { + use std::{ + sync::atomic::{AtomicBool, Ordering}, + time::Duration, + }; + + use super::{Claimed, Control, LanePool, LanePoolConfig}; + + struct Log { + seen: Vec, + } + + /// Every worker's events apply in order and under one lane at a time while three lanes + /// steal from each other and two producers wait for room instead of losing events. + #[test] + fn events_of_a_worker_apply_in_order_under_one_lane_at_a_time() { + const WORKERS: u32 = 16; + const EVENTS: u64 = 2_000; + const LANES: usize = 3; + let pool: LanePool = LanePool::new(LanePoolConfig { + lanes: LANES, + max_workers: WORKERS as usize, + depth_cap: 64, + batch: 8, + steal_after: 0, + }); + let produced = AtomicBool::new(false); + let claims: Vec = (0..WORKERS).map(|_| AtomicBool::new(false)).collect(); + std::thread::scope(|scope| { + for lane in 0..LANES { + let (pool, produced, claims) = (&pool, &produced, &claims); + scope.spawn(move || { + pool.run_lane( + lane, + &mut ( + |claimed: Claimed<'_, Log>, event: u64| { + let worker = claimed.worker as usize; + assert!( + !claims[worker].swap(true, Ordering::SeqCst), + "two lanes inside worker {worker}" + ); + claimed + .state + .get_or_insert_with(|| Log { seen: Vec::new() }) + .seen + .push(event); + std::thread::yield_now(); + claims[worker].store(false, Ordering::SeqCst); + }, + || Control::Continue, + || { + if produced.load(Ordering::SeqCst) && pool.queued() == 0 { + return Control::Stop; + } + // Long enough that progress depends on the producers' wake-ups. + pool.park_lane(lane, Duration::from_millis(50)); + Control::Continue + }, + ), + ); + }); + } + let producers: Vec<_> = (0..2u32) + .map(|half| { + let pool = &pool; + scope.spawn(move || { + for event in 0..EVENTS { + for worker in (half..WORKERS).step_by(2) { + let mut pending = event; + loop { + match pool.enqueue((worker as usize) % LANES, worker, pending) { + Ok(()) => break, + Err(full) => { + pending = full.0; + std::thread::yield_now(); + } + } + } + } + } + }) + }) + .collect(); + for producer in producers { + producer.join().expect("producer"); + } + produced.store(true, Ordering::SeqCst); + }); + let metrics = pool.metrics(); + assert_eq!(metrics.enqueued, u64::from(WORKERS) * EVENTS); + assert_eq!(metrics.applied, u64::from(WORKERS) * EVENTS); + assert_eq!(metrics.queued, 0); + assert!(metrics.max_depth <= 64); + for worker in 0..WORKERS { + let seen = super::lock(&pool.workers[worker as usize].held) + .as_ref() + .map(|log| log.seen.clone()) + .unwrap_or_default(); + assert_eq!(seen, (0..EVENTS).collect::>(), "worker {worker}"); + } + } + + /// The cap hands the event back instead of dropping it, and counts the refusal. + #[test] + fn the_depth_cap_refuses_without_dropping() { + let pool: LanePool<(), &'static str> = LanePool::new(LanePoolConfig { + lanes: 1, + max_workers: 2, + depth_cap: 2, + batch: 4, + steal_after: 0, + }); + assert!(pool.enqueue(0, 1, "a").is_ok()); + assert!(pool.enqueue(0, 1, "b").is_ok()); + let refused = pool + .enqueue(0, 1, "c") + .expect_err("third event is over the cap"); + assert_eq!(refused.0, "c"); + assert_eq!(pool.depth(1), 2); + assert_eq!(pool.depth(0), 0); + let metrics = pool.metrics(); + assert_eq!( + (metrics.enqueued, metrics.rejected, metrics.queued), + (2, 1, 2) + ); + assert_eq!(metrics.max_depth, 2); + } + + /// A lane serves its own ready workers a batch at a time; while a lane is stuck inside one + /// worker's event, another lane takes the rest of its list, whole workers at a time, and + /// every worker's events stay in order. Nothing is taken from a lane that is not serving. + #[test] + fn an_idle_lane_steals_whole_workers_from_a_serving_one() { + let pool: LanePool, u32> = LanePool::new(LanePoolConfig { + lanes: 2, + max_workers: 4, + depth_cap: 16, + batch: 2, + steal_after: 0, + }); + for round in 0..3 { + for worker in 0..4u32 { + // Workers 0 and 1 are lane 1's own, 2 and 3 are lane 0's. + let home = usize::from(worker >= 2); + pool.enqueue(1 - home, worker, round * 10 + worker).unwrap(); + } + } + // Lane 1 alone: lane 0 is not serving, so its workers stay where they are. + let mut order = Vec::new(); + pool.run_lane( + 1, + &mut ( + |claimed: Claimed<'_, Vec>, event: u32| { + claimed.state.get_or_insert_with(Vec::new).push(event); + order.push((claimed.worker, event)); + }, + || Control::Continue, + || Control::Stop, + ), + ); + assert_eq!(pool.metrics().steals, 0); + assert_eq!(pool.metrics().applied, 6); + // Batches of two: worker 0 gave the lane to worker 1 after two events and came back. + let lanes_own: Vec = order.iter().map(|(w, _)| *w).collect(); + assert_eq!(lanes_own, vec![0, 0, 1, 1, 0, 1]); + + // Lane 0 enters worker 2 and blocks there; lane 1 then takes worker 3 from its list. + let blocked = AtomicBool::new(true); + let stolen = std::sync::Mutex::new(Vec::new()); + std::thread::scope(|scope| { + let (pool, blocked) = (&pool, &blocked); + scope.spawn(move || { + pool.run_lane( + 0, + &mut ( + |claimed: Claimed<'_, Vec>, event: u32| { + while blocked.load(Ordering::SeqCst) { + std::thread::yield_now(); + } + claimed.state.get_or_insert_with(Vec::new).push(event); + }, + || Control::Continue, + || Control::Stop, + ), + ); + }); + while !pool.lanes[0].serving.load(Ordering::Relaxed) { + std::thread::yield_now(); + } + pool.run_lane( + 1, + &mut ( + |claimed: Claimed<'_, Vec>, event: u32| { + claimed.state.get_or_insert_with(Vec::new).push(event); + super::lock(&stolen).push((claimed.worker, event)); + }, + || Control::Continue, + || Control::Stop, + ), + ); + blocked.store(false, Ordering::SeqCst); + }); + let metrics = pool.metrics(); + assert_eq!(metrics.steals, 1, "{metrics:?}"); + assert_eq!(metrics.applied, 12); + assert_eq!(*super::lock(&stolen), vec![(3, 3), (3, 13), (3, 23)]); + let worker_two = super::lock(&pool.workers[2].held) + .clone() + .unwrap_or_default(); + assert_eq!(worker_two, vec![2, 12, 22]); + } + + /// A ready worker on a lane that is not running is taken by another lane once its queue is + /// `steal_after` deep, and never under the default: lane 0 is never run here, which is what a + /// preempted lane looks like to the pool. + #[test] + fn a_stalled_lane_is_drained_once_its_backlog_reaches_steal_after() { + for (steal_after, expect_applied, expect_steals) in [(0usize, 0u64, 0u64), (3, 4, 1)] { + let pool: LanePool<(), u32> = LanePool::new(LanePoolConfig { + lanes: 2, + max_workers: 1, + depth_cap: 16, + batch: 8, + steal_after, + }); + for event in 0..4u32 { + pool.enqueue(0, 0, event).unwrap(); + } + let mut applied = Vec::new(); + let mut turns = 0u32; + pool.run_lane( + 1, + &mut ( + |_claimed: Claimed<'_, ()>, event: u32| applied.push(event), + || Control::Continue, + || { + turns += 1; + if turns > 20 { + return Control::Stop; + } + pool.park_lane(1, Duration::from_millis(1)); + Control::Continue + }, + ), + ); + assert_eq!( + applied.len() as u64, + expect_applied, + "steal_after {steal_after}: lane 1 took the stalled lane's backlog" + ); + if expect_applied > 0 { + assert_eq!(applied, vec![0, 1, 2, 3], "in order"); + } + assert_eq!(pool.metrics().steals, expect_steals); + assert_eq!(pool.metrics().queued, 4 - expect_applied as usize); + } + } + + /// Under `steal_after`, a backlog shallower than the setting stays with its lane: a lane that + /// is merely slow to its next turn is not robbed of its own work. + #[test] + fn a_shallow_backlog_stays_with_its_lane() { + let pool: LanePool<(), u32> = LanePool::new(LanePoolConfig { + lanes: 2, + max_workers: 1, + depth_cap: 16, + batch: 8, + steal_after: 3, + }); + pool.enqueue(0, 0, 1).unwrap(); + pool.enqueue(0, 0, 2).unwrap(); + let mut turns = 0u32; + let mut applied = 0u64; + pool.run_lane( + 1, + &mut ( + |_claimed: Claimed<'_, ()>, _event: u32| applied += 1, + || Control::Continue, + || { + turns += 1; + if turns > 5 { + return Control::Stop; + } + Control::Continue + }, + ), + ); + assert_eq!(applied, 0); + assert_eq!(pool.metrics().steals, 0); + assert_eq!(pool.depth(0), 2); + } +} diff --git a/crates/kv_index/src/lib.rs b/crates/kv_index/src/lib.rs index fe62909c74..4f5707a35c 100644 --- a/crates/kv_index/src/lib.rs +++ b/crates/kv_index/src/lib.rs @@ -11,20 +11,37 @@ //! - Concurrent access via DashMap and RwLock //! - Efficient prefix matching with match counts +pub mod chain_index; +/// Churn generator shared by the churn bench and the gate test; not a public API. +#[doc(hidden)] +pub mod churn; mod common; mod event_tree; +mod lane_map; +pub mod lane_pool; mod path_hash; +mod prefetch; +pub mod reference; +pub mod salt; +pub mod sharded; pub mod snapshot; mod string_tree; mod token_tree; +pub use chain_index::{BlockRef, ChainIndex, ChainIndexStats}; pub use common::{MatchResult, TenantId}; pub use event_tree::{ chain_prefix_hash, compute_content_hash, compute_request_content_hashes, ApplyError, - ContentHash, OverlapScores, PositionalIndexer, PruneStats, SequenceHash, StoredBlock, - WorkerBlockMap, WorkerId, WorkerIdExhausted, XXH3_SEED, + ContentHash, ContentSeq, OverlapScores, PositionalIndexer, PruneStats, SequenceHash, + StoredBlock, WorkerBlockMap, WorkerId, WorkerIdExhausted, XXH3_SEED, +}; +pub use lane_map::ChainBlockMap; +pub use lane_pool::{ + Claimed, Control, LaneHooks, LanePool, LanePoolConfig, PoolMetrics, QueueFull, }; pub use path_hash::{hash_node_path, hash_token_path, GLOBAL_EVICTION_HASH}; +pub use reference::{request_prefix_hashes, ReferenceIndexer}; +pub use sharded::{ShardedChainIndex, SHARD_SHIFT}; // Re-export under names matching old tree.rs API for easier migration pub use string_tree::Tree; pub use string_tree::{ diff --git a/crates/kv_index/src/prefetch.rs b/crates/kv_index/src/prefetch.rs new file mode 100644 index 0000000000..d274fa6089 --- /dev/null +++ b/crates/kv_index/src/prefetch.rs @@ -0,0 +1,70 @@ +//! A read-prefetch hint, so a batch of hash-table probes can have its cache misses overlap: call +//! it on the home slots of the keys a few steps ahead, then probe. +//! +//! A prefetch is advice to the cache, never a load: it cannot fault, cannot change program state, +//! and may be ignored by the hardware. That is why [`prefetch_read`] is a safe function for any +//! pointer, aligned or not, mapped or not, dangling or null. The function is a no-op on targets +//! without a prefetch instruction. +//! +//! This is the crate's one unsafe line, behind a safe function; the workspace denies unsafe code +//! and the exception is scoped to that function, not to a crate of its own. + +#![deny(unsafe_op_in_unsafe_fn)] + +/// Hint that the cache line at `pointer` will be read soon (first-level cache, keep). Sound for +/// any pointer: a prefetch instruction never faults and has no effect on program state. +#[inline(always)] +#[cfg_attr( + any(target_arch = "aarch64", target_arch = "x86_64"), + expect( + unsafe_code, + reason = "a prefetch instruction is a hint: it never faults, reads nothing, writes nothing" + ) +)] +pub fn prefetch_read(pointer: *const T) { + #[cfg(target_arch = "aarch64")] + { + // SAFETY: `prfm pldl1keep` is a hint that never faults, whatever the address holds; it + // reads no memory, writes no memory and no register, and touches no flags. + unsafe { + core::arch::asm!( + "prfm pldl1keep, [{address}]", + address = in(reg) pointer, + options(nostack, preserves_flags, readonly) + ); + } + } + #[cfg(target_arch = "x86_64")] + { + // SAFETY: `prefetcht0` is a hint that never faults, whatever the address holds. + unsafe { + use core::arch::x86_64::{_mm_prefetch, _MM_HINT_T0}; + _mm_prefetch::<{ _MM_HINT_T0 }>(pointer.cast::()); + } + } + #[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))] + { + let _ = pointer; + } +} + +#[cfg(test)] +mod tests { + use super::prefetch_read; + + #[test] + fn prefetching_a_stack_value_is_harmless() { + let value = [7u64; 16]; + prefetch_read(value.as_ptr()); + prefetch_read(&value[15]); + assert_eq!(value[3], 7); + } + + #[test] + fn prefetching_dangling_and_null_pointers_does_not_fault() { + let dangling = 0x7f00_0000_0000usize as *const u64; + prefetch_read(dangling); + prefetch_read(core::ptr::null::()); + prefetch_read(usize::MAX as *const u8); + } +} diff --git a/crates/kv_index/src/reference.rs b/crates/kv_index/src/reference.rs new file mode 100644 index 0000000000..fffca733b3 --- /dev/null +++ b/crates/kv_index/src/reference.rs @@ -0,0 +1,276 @@ +//! A single-threaded reference model of what [`PositionalIndexer`](crate::PositionalIndexer) +//! promises, for the exactness guardrail (`docs/kv-router-leap.md`, section 3). +//! +//! Nothing here is meant to be fast. Every structure is the most literal one that states the +//! contract: +//! +//! - a worker holds a set of blocks; each block is identified by the engine's sequence hash and +//! sits at a position with a content hash and a prefix hash (the router's chain hash over the +//! content hashes up to and including that position); +//! - a store places its blocks after the parent block of the same worker (position 0 when there +//! is no parent) and fails exactly like the production indexer when the parent is unknown; +//! - a removal forgets the named blocks of that worker; a clear or a worker removal forgets all +//! of them; +//! - a lookup scores a worker with the length of the longest prefix of the request for which the +//! worker holds, at every position `i`, a block whose content hash and prefix hash equal the +//! request's at `i`; the first position without such a block ends the prefix; workers with an +//! empty prefix are not reported. +//! +//! The chain hash is the crate's own [`chain_prefix_hash`], not a copy of it, so both indexers +//! always agree on what a prefix hash is. + +use std::collections::{BTreeMap, BTreeSet}; + +use rustc_hash::{FxHashMap, FxHashSet}; + +use crate::event_tree::{chain_prefix_hash, ApplyError, ContentHash, SequenceHash, StoredBlock}; + +/// One worker's blocks, indexed two ways: by the engine hash (how removals name them) and by +/// `(position, content hash)` (how lookups find them), each position holding every prefix hash +/// that reaches it. +#[derive(Default, Clone)] +struct WorkerBlocks { + by_engine_hash: FxHashMap, + by_position: FxHashMap<(usize, ContentHash), FxHashSet>, +} + +impl WorkerBlocks { + fn insert( + &mut self, + engine_hash: SequenceHash, + position: usize, + content: ContentHash, + prefix: SequenceHash, + ) { + // A hash stored again at another place (a store without its parent followed by the + // whole chain) is held at the new place only, as the production indexers hold it. + if let Some(old) = self + .by_engine_hash + .insert(engine_hash, (position, content, prefix)) + { + if old != (position, content, prefix) { + let (old_position, old_content, old_prefix) = old; + if let Some(prefixes) = self.by_position.get_mut(&(old_position, old_content)) { + prefixes.remove(&old_prefix); + if prefixes.is_empty() { + self.by_position.remove(&(old_position, old_content)); + } + } + } + } + self.by_position + .entry((position, content)) + .or_default() + .insert(prefix); + } + + fn remove(&mut self, engine_hash: SequenceHash) { + let Some((position, content, prefix)) = self.by_engine_hash.remove(&engine_hash) else { + return; + }; + // Another block of this worker may reach the same (position, content) through the same + // chain only if it has the same engine hash, so the prefix is gone with this block. + if let Some(prefixes) = self.by_position.get_mut(&(position, content)) { + prefixes.remove(&prefix); + if prefixes.is_empty() { + self.by_position.remove(&(position, content)); + } + } + } + + fn holds(&self, position: usize, content: ContentHash, prefix: SequenceHash) -> bool { + self.by_position + .get(&(position, content)) + .is_some_and(|prefixes| prefixes.contains(&prefix)) + } +} + +/// The reference indexer. Workers are the `u32` ids the production indexer interns. +#[derive(Default, Clone)] +pub struct ReferenceIndexer { + workers: BTreeMap, +} + +/// The chain of prefix hashes of a request: position 0 is the bare content hash, every later +/// position chains the previous prefix with its content hash. +pub fn request_prefix_hashes(content_hashes: &[ContentHash]) -> Vec { + let mut chain = Vec::with_capacity(content_hashes.len()); + for (position, &content) in content_hashes.iter().enumerate() { + let prefix = if position == 0 { + SequenceHash(content.0) + } else { + chain_prefix_hash(chain[position - 1], content) + }; + chain.push(prefix); + } + chain +} + +impl ReferenceIndexer { + pub fn new() -> Self { + Self::default() + } + + /// Store `blocks` for `worker` after `parent`, with the production indexer's placement and + /// error rules. + pub fn apply_stored( + &mut self, + worker: u32, + blocks: &[StoredBlock], + parent: Option, + ) -> Result<(), ApplyError> { + if blocks.is_empty() { + return Ok(()); + } + let held = self.workers.entry(worker).or_default(); + let (start, mut previous_prefix) = match parent { + Some(parent_hash) => { + if held.by_engine_hash.is_empty() { + return Err(ApplyError::WorkerNotTracked); + } + let Some(&(parent_position, _, parent_prefix)) = + held.by_engine_hash.get(&parent_hash) + else { + return Err(ApplyError::ParentBlockNotFound); + }; + (parent_position + 1, Some(parent_prefix)) + } + None => (0, None), + }; + for (offset, block) in blocks.iter().enumerate() { + let position = start + offset; + let prefix = match previous_prefix { + Some(previous) => chain_prefix_hash(previous, block.content_hash), + None => SequenceHash(block.content_hash.0), + }; + held.insert(block.seq_hash, position, block.content_hash, prefix); + previous_prefix = Some(prefix); + } + Ok(()) + } + + /// Forget the named blocks of `worker`; unknown hashes are ignored, as in production. + pub fn apply_removed(&mut self, worker: u32, engine_hashes: &[SequenceHash]) { + if let Some(held) = self.workers.get_mut(&worker) { + for &engine_hash in engine_hashes { + held.remove(engine_hash); + } + } + } + + /// Forget every block of `worker` (the engine cleared its cache). + pub fn apply_cleared(&mut self, worker: u32) { + self.workers.remove(&worker); + } + + /// Forget every block of `worker` (the worker left). + pub fn remove_worker(&mut self, worker: u32) { + self.workers.remove(&worker); + } + + /// Score every worker by the length of the longest matching prefix of the request, omitting + /// workers that match nothing. + pub fn find_matches(&self, content_hashes: &[ContentHash]) -> BTreeMap { + let prefixes = request_prefix_hashes(content_hashes); + let mut scores = BTreeMap::new(); + for (&worker, held) in &self.workers { + let matched = content_hashes + .iter() + .zip(&prefixes) + .enumerate() + .take_while(|(position, (&content, &prefix))| { + held.holds(*position, content, prefix) + }) + .count(); + if matched > 0 { + scores.insert(worker, matched as u32); + } + } + scores + } + + /// Every block every worker holds, as `(worker, position, content hash, prefix hash)`. + pub fn blocks(&self) -> BTreeSet<(u32, usize, ContentHash, SequenceHash)> { + self.workers + .iter() + .flat_map(|(&worker, held)| { + held.by_engine_hash + .values() + .map(move |&(position, content, prefix)| (worker, position, content, prefix)) + }) + .collect() + } + + /// Number of blocks held by `worker`. + pub fn worker_block_count(&self, worker: u32) -> usize { + self.workers + .get(&worker) + .map_or(0, |held| held.by_engine_hash.len()) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn chain(contents: &[u64]) -> Vec { + let hashes: Vec = contents.iter().map(|&c| ContentHash(c)).collect(); + let prefixes = request_prefix_hashes(&hashes); + hashes + .iter() + .zip(prefixes) + .map(|(&content_hash, prefix)| StoredBlock { + seq_hash: prefix, + content_hash, + }) + .collect() + } + + #[test] + fn scores_the_longest_exact_prefix_only() { + let mut reference = ReferenceIndexer::new(); + let blocks = chain(&[1, 2, 3, 4]); + reference.apply_stored(7, &blocks, None).expect("store"); + let query: Vec = [1, 2, 3, 4].iter().map(|&c| ContentHash(c)).collect(); + assert_eq!(reference.find_matches(&query).get(&7), Some(&4)); + let diverged: Vec = [1, 9, 3, 4].iter().map(|&c| ContentHash(c)).collect(); + assert_eq!(reference.find_matches(&diverged).get(&7), Some(&1)); + let other: Vec = [5, 2, 3].iter().map(|&c| ContentHash(c)).collect(); + assert!(reference.find_matches(&other).is_empty()); + } + + #[test] + fn stores_after_the_parent_and_rejects_unknown_parents() { + let mut reference = ReferenceIndexer::new(); + let head = chain(&[1, 2]); + reference.apply_stored(1, &head, None).expect("head"); + let tail = chain(&[1, 2, 3, 4]); + reference + .apply_stored(1, &tail[2..], Some(head[1].seq_hash)) + .expect("tail after parent"); + let query: Vec = [1, 2, 3, 4].iter().map(|&c| ContentHash(c)).collect(); + assert_eq!(reference.find_matches(&query).get(&1), Some(&4)); + assert!(matches!( + reference.apply_stored(1, &tail[2..], Some(SequenceHash(0xdead))), + Err(ApplyError::ParentBlockNotFound) + )); + assert!(matches!( + reference.apply_stored(2, &tail[2..], Some(head[1].seq_hash)), + Err(ApplyError::WorkerNotTracked) + )); + } + + #[test] + fn removals_truncate_the_match() { + let mut reference = ReferenceIndexer::new(); + let blocks = chain(&[1, 2, 3, 4]); + reference.apply_stored(3, &blocks, None).expect("store"); + reference.apply_removed(3, &[blocks[3].seq_hash, blocks[2].seq_hash]); + let query: Vec = [1, 2, 3, 4].iter().map(|&c| ContentHash(c)).collect(); + assert_eq!(reference.find_matches(&query).get(&3), Some(&2)); + assert_eq!(reference.blocks().len(), 2); + reference.apply_cleared(3); + assert!(reference.find_matches(&query).is_empty()); + assert!(reference.blocks().is_empty()); + } +} diff --git a/crates/kv_index/src/salt.rs b/crates/kv_index/src/salt.rs new file mode 100644 index 0000000000..f0891de22f --- /dev/null +++ b/crates/kv_index/src/salt.rs @@ -0,0 +1,129 @@ +//! Namespaced content hashing. +//! +//! A block stored under a LoRA adapter or a cache salt is not reusable by a +//! plain prompt, and two salts are not reusable by each other. The engines +//! fold these into their own block hashes; this crate recomputes content +//! hashes from token ids, so it folds them into the XXH3 seed instead. The +//! mixing is fixed (seed = 1337 + xxh3(lora_name, 0) + xxh3(cache_salt, 1), +//! wrapping), so a corpus hashed elsewhere by the same rule stays comparable. Empty strings count as absent, as a client +//! sending `lora_name = ""` means the base model. + +use xxhash_rust::xxh3::{xxh3_64_with_seed, Xxh3}; + +use crate::{ContentHash, XXH3_SEED}; + +/// The XXH3 seed for a cache namespace; [`XXH3_SEED`] when there is none. +pub fn namespace_seed(lora_name: Option<&str>, cache_salt: Option<&str>) -> u64 { + let mut seed = XXH3_SEED; + if let Some(name) = lora_name.filter(|name| !name.is_empty()) { + seed = seed.wrapping_add(xxh3_64_with_seed(name.as_bytes(), 0)); + } + if let Some(salt) = cache_salt.filter(|salt| !salt.is_empty()) { + seed = seed.wrapping_add(xxh3_64_with_seed(salt.as_bytes(), 1)); + } + seed +} + +/// [`crate::compute_content_hash`] under an explicit seed. +pub fn content_hash_with_seed(token_ids: &[u32], seed: u64) -> ContentHash { + use std::hash::Hasher; + let mut hasher = Xxh3::with_seed(seed); + for &token in token_ids { + hasher.write(&token.to_le_bytes()); + } + ContentHash(hasher.finish()) +} + +/// [`crate::compute_request_content_hashes`] under an explicit seed (see +/// [`namespace_seed`]): one hash per full block of `block_size` tokens, the +/// partial tail ignored. +pub fn request_content_hashes_with_seed( + token_ids: &[u32], + block_size: usize, + seed: u64, +) -> Vec { + if block_size == 0 { + return Vec::new(); + } + token_ids + .chunks_exact(block_size) + .map(|block| content_hash_with_seed(block, seed)) + .collect() +} + +/// [`crate::compute_request_content_hashes`] under a cache namespace. +pub fn namespaced_request_content_hashes( + token_ids: &[u32], + block_size: usize, + lora_name: Option<&str>, + cache_salt: Option<&str>, +) -> Vec { + request_content_hashes_with_seed(token_ids, block_size, namespace_seed(lora_name, cache_salt)) +} + +#[cfg(test)] +mod tests { + use super::*; + + /// The content hash of one block stored under a cache namespace. Equal to + /// [`crate::compute_content_hash`] when the namespace is empty. + fn namespaced_content_hash( + token_ids: &[u32], + lora_name: Option<&str>, + cache_salt: Option<&str>, + ) -> ContentHash { + content_hash_with_seed(token_ids, namespace_seed(lora_name, cache_salt)) + } + + use crate::{compute_content_hash, compute_request_content_hashes}; + + const TOKENS: [u32; 6] = [11, 22, 33, 44, 55, 66]; + + #[test] + fn no_namespace_is_the_plain_content_hash() { + assert_eq!(namespace_seed(None, None), XXH3_SEED); + assert_eq!(namespace_seed(Some(""), Some("")), XXH3_SEED); + assert_eq!( + namespaced_content_hash(&TOKENS, None, None), + compute_content_hash(&TOKENS) + ); + assert_eq!( + namespaced_request_content_hashes(&TOKENS, 4, Some(""), None), + compute_request_content_hashes(&TOKENS, 4) + ); + } + + #[test] + fn lora_and_salt_each_change_the_hash_and_do_not_collide() { + let plain = namespaced_content_hash(&TOKENS, None, None); + let lora = namespaced_content_hash(&TOKENS, Some("adapter"), None); + let salt = namespaced_content_hash(&TOKENS, None, Some("adapter")); + let both = namespaced_content_hash(&TOKENS, Some("adapter"), Some("adapter")); + let other_salt = namespaced_content_hash(&TOKENS, None, Some("other")); + let distinct = [plain, lora, salt, both, other_salt]; + for (i, a) in distinct.iter().enumerate() { + for b in &distinct[i + 1..] { + assert_ne!(a, b); + } + } + } + + #[test] + fn seed_mixing_matches_the_documented_formula() { + let expected = XXH3_SEED + .wrapping_add(xxh3_64_with_seed(b"adapter", 0)) + .wrapping_add(xxh3_64_with_seed(b"salty", 1)); + assert_eq!(namespace_seed(Some("adapter"), Some("salty")), expected); + } + + #[test] + fn request_hashes_follow_the_block_grid() { + let hashes = namespaced_request_content_hashes(&TOKENS, 4, Some("adapter"), None); + assert_eq!(hashes.len(), 1); + assert_eq!( + hashes[0], + namespaced_content_hash(&TOKENS[..4], Some("adapter"), None) + ); + assert!(namespaced_request_content_hashes(&TOKENS, 0, None, None).is_empty()); + } +} diff --git a/crates/kv_index/src/sharded.rs b/crates/kv_index/src/sharded.rs new file mode 100644 index 0000000000..fe61ea535e --- /dev/null +++ b/crates/kv_index/src/sharded.rs @@ -0,0 +1,604 @@ +//! The chain index split into shards: each shard is a complete [`ChainIndex`] holding the workers +//! assigned to it, a lookup walks every shard and unions the scores. +//! +//! Why: with lanes on both sockets of a host, a single index has every run's version, length, +//! coverage words and child table, the arena's bump pointer and free lists, and the run slab +//! written from both sockets; the replay's duplicated corpus turns that into true sharing, and +//! the lane CPU per event doubles to triples against one socket. One shard per socket, each +//! written only by the lanes of that socket, removes every such line: no run state is shared +//! between shards, and the lookup, which writes nothing, reads them all. +//! +//! Exactness does not depend on the assignment: a shard is exact for its own workers whatever +//! the caller's affinity, the worker sets of the shards are disjoint, so the union of their +//! scores is the single index's answer. A caller with no control over which thread applies +//! which worker (the gateway's runtime workers on both sockets) is exact and merely pays the +//! sharing it did not avoid. `shards = 1` is the chain index itself behind one indirection: +//! the same ids, the same lookups, the same memory. +//! +//! Worker ids carry the shard: `shard << SHARD_SHIFT | local id`, so an id names its shard +//! without a table, and with one shard the ids are the chain index's own. Content held by workers +//! on two shards is stored once per shard (hash arrays and run headers; the lane maps are per +//! worker either way), which is the price of the split; `entry_count` counts it once per shard. + +use std::{ + collections::BTreeSet, + sync::atomic::{AtomicUsize, Ordering}, +}; + +use crate::{ + chain_index::{ChainIndex, ChainIndexStats}, + event_tree::{ + ApplyError, ContentHash, OverlapScores, SequenceHash, StoredBlock, WorkerIdExhausted, + }, + lane_map::ChainBlockMap, +}; + +/// Bits of a worker id below the shard number; a shard holds at most `1 << SHARD_SHIFT` worker +/// slots (the chain index allows 1,024). +pub const SHARD_SHIFT: u32 = 16; + +const LOCAL_MASK: u32 = (1 << SHARD_SHIFT) - 1; + +/// A chain index per shard; see the module documentation. +pub struct ShardedChainIndex { + shards: Box<[ChainIndex]>, + /// Where the next worker interned without a shard goes (round robin). + next_shard: AtomicUsize, +} + +impl ShardedChainIndex { + /// `shards` chain indexes (at least one) with `max_workers` worker slots each. + pub fn new(shards: usize, max_workers: usize) -> Self { + let shards = shards.max(1); + assert!( + shards <= 1 << (32 - SHARD_SHIFT), + "at most {} shards", + 1 << (32 - SHARD_SHIFT) + ); + Self { + shards: (0..shards) + .map(|_| ChainIndex::with_max_workers(max_workers)) + .collect(), + next_shard: AtomicUsize::new(0), + } + } + + /// `shards` chain indexes with the chain index's default worker slots. + pub fn with_shards(shards: usize) -> Self { + let shards = shards.max(1); + Self { + shards: (0..shards).map(|_| ChainIndex::new()).collect(), + next_shard: AtomicUsize::new(0), + } + } + + pub fn shards(&self) -> usize { + self.shards.len() + } + + /// One shard's index, for per-shard figures. + pub fn shard(&self, shard: usize) -> &ChainIndex { + &self.shards[shard] + } + + /// The shard a worker id names. + #[inline] + pub fn shard_of(worker: u32) -> usize { + (worker >> SHARD_SHIFT) as usize + } + + /// A worker's id within its shard. + #[inline] + pub fn local_id(worker: u32) -> u32 { + worker & LOCAL_MASK + } + + #[inline] + fn global(shard: usize, local: u32) -> u32 { + ((shard as u32) << SHARD_SHIFT) | local + } + + #[inline] + fn index_of(&self, worker: u32) -> &ChainIndex { + &self.shards[Self::shard_of(worker)] + } + + /// Intern `name` in `shard`; the id it already has if it is interned anywhere. + pub fn intern_worker_in(&self, shard: usize, name: &str) -> Result { + assert!( + shard < self.shards.len(), + "shard {shard} of {}", + self.shards.len() + ); + if let Some(id) = self.worker_id(name) { + return Ok(id); + } + self.shards[shard] + .intern_worker(name) + .map(|local| Self::global(shard, local)) + } + + /// Intern `name`, in the shard it already has or else the next one round robin: for a + /// caller without a placement of its own (exact whatever the shard, see the module notes). + pub fn intern_worker(&self, name: &str) -> Result { + if let Some(id) = self.worker_id(name) { + return Ok(id); + } + let shard = if self.shards.len() == 1 { + 0 + } else { + self.next_shard.fetch_add(1, Ordering::Relaxed) % self.shards.len() + }; + self.intern_worker_in(shard, name) + } + + /// The id of an interned worker. + pub fn worker_id(&self, name: &str) -> Option { + self.shards.iter().enumerate().find_map(|(shard, index)| { + index + .worker_id(name) + .map(|local| Self::global(shard, local)) + }) + } + + /// Store `blocks` for `worker` after `parent` (position 0 when `None`). + pub fn apply_stored( + &self, + worker: u32, + blocks: &[StoredBlock], + parent: Option, + map: &mut ChainBlockMap, + ) -> Result<(), ApplyError> { + self.index_of(worker) + .apply_stored(Self::local_id(worker), blocks, parent, map) + } + + /// Forget the named blocks of `worker`; unknown hashes are ignored. + pub fn apply_removed(&self, worker: u32, hashes: &[SequenceHash], map: &mut ChainBlockMap) { + self.index_of(worker) + .apply_removed(Self::local_id(worker), hashes, map); + } + + /// Forget every block of `worker` (the engine cleared its cache); the map is emptied. + pub fn apply_cleared(&self, worker: u32, map: &mut ChainBlockMap) { + self.index_of(worker) + .apply_cleared(Self::local_id(worker), map); + } + + /// Forget every block of `worker` (the worker left) and free its slot in its shard. + pub fn remove_worker(&self, worker: u32, map: ChainBlockMap) { + self.index_of(worker) + .remove_worker(Self::local_id(worker), map); + } + + /// Whether `worker`'s lane map holds the block with engine hash `key`. + pub fn is_held(&self, worker: u32, map: &ChainBlockMap, key: SequenceHash) -> bool { + self.index_of(worker).is_held(map, key) + } + + /// Blocks `worker` holds. + pub fn worker_block_count(&self, worker: u32) -> usize { + self.index_of(worker) + .worker_block_count(Self::local_id(worker)) + } + + /// Whether no worker of any shard holds a block (O(shards), each under its root's version). + pub fn is_empty(&self) -> bool { + self.shards.iter().all(ChainIndex::is_empty) + } + + /// Blocks held across all workers (a block two workers hold counts twice). + pub fn current_size(&self) -> usize { + self.shards.iter().map(ChainIndex::current_size).sum() + } + + /// Distinct blocks held, per shard and summed: content held on two shards counts twice, + /// which is exactly the memory the split costs. + pub fn entry_count(&self) -> usize { + self.shards.iter().map(ChainIndex::entry_count).sum() + } + + /// Score every worker by how many leading blocks of the request it holds; the union over + /// the shards. With `early_exit`, the workers holding the first block, each scored 1. + pub fn find_matches(&self, content_hashes: &[ContentHash], early_exit: bool) -> OverlapScores { + if let [only] = &*self.shards { + return only.find_matches(content_hashes, early_exit); + } + let mut out = OverlapScores::default(); + self.score_into( + content_hashes, + |content| content.0, + early_exit, + |worker, score| { + out.scores.insert(worker, score); + }, + ); + out + } + + /// The lookup behind [`find_matches`](Self::find_matches) for callers with their own hash + /// type and result shape, over every shard: `report` receives `(worker, score)` once per + /// worker with a non-empty prefix. Returns the runs walked, summed over the shards. + pub fn score_into( + &self, + content_hashes: &[T], + hash_of: impl Fn(&T) -> u64, + early_exit: bool, + mut report: impl FnMut(u32, u32), + ) -> usize { + if let [only] = &*self.shards { + return only.score_into(content_hashes, &hash_of, early_exit, report); + } + let Some(first) = content_hashes.first() else { + return 0; + }; + let first = hash_of(first); + let mut walked = 0; + for (shard, index) in self.shards.iter().enumerate() { + // A shard whose root has no run starting with the request's first block holds + // nothing of it: one probe of its root table and the shard is passed over. + let Some(entry) = index.head_entry(first) else { + continue; + }; + walked += index.score_from( + entry, + content_hashes, + &hash_of, + early_exit, + |worker, score| { + report(Self::global(shard, worker), score); + }, + ); + } + walked + } + + /// The lookup over one shard alone, for a caller that fans a lookup out to one thread per + /// shard (one per socket) and merges: the call walks only that shard's memory, and `report` + /// receives `(worker, score)` with the worker's global id, exactly the pairs + /// [`score_into`](Self::score_into) would report for the shard. Worker sets are disjoint + /// across shards, so the concatenation of every shard's reports is `score_into`'s answer. + /// Returns the runs walked on the shard. + pub fn score_shard_into( + &self, + shard: usize, + content_hashes: &[T], + hash_of: impl Fn(&T) -> u64, + early_exit: bool, + mut report: impl FnMut(u32, u32), + ) -> usize { + self.shards[shard].score_into(content_hashes, hash_of, early_exit, |worker, score| { + report(Self::global(shard, worker), score); + }) + } + + /// How many shards hold a chain starting with `first` (a diagnostic for the lookup's shard + /// filter: a shard without the head costs one root probe, a shard with it a walk). + pub fn shards_holding_head(&self, first: u64) -> usize { + self.shards + .iter() + .filter(|index| index.head_entry(first).is_some()) + .count() + } + + /// Mergeable adjacent pairs and the blocks their children hold, summed over the shards (see + /// [`ChainIndex::debug_mergeable`]). + #[doc(hidden)] + pub fn debug_mergeable(&self) -> (usize, usize) { + self.shards + .iter() + .map(ChainIndex::debug_mergeable) + .fold((0, 0), |(p, b), (q, c)| (p + q, b + c)) + } + + /// Shape and memory counters summed over the shards. + pub fn stats(&self) -> ChainIndexStats { + let mut total = ChainIndexStats::default(); + for stats in self.shards.iter().map(ChainIndex::stats) { + total.runs_allocated += stats.runs_allocated; + total.runs_free += stats.runs_free; + total.runs_live += stats.runs_live; + total.blocks_live += stats.blocks_live; + total.arena_bytes += stats.arena_bytes; + total.arena_free_bytes += stats.arena_free_bytes; + total.arena_chunk_bytes += stats.arena_chunk_bytes; + total.header_bytes += stats.header_bytes; + total.slab_bytes += stats.slab_bytes; + total.engine_conflicts += stats.engine_conflicts; + total.landing_mismatches += stats.landing_mismatches; + total.moved_hashes += stats.moved_hashes; + total.splits_by_branch += stats.splits_by_branch; + total.splits_by_hole += stats.splits_by_hole; + total.splits_by_mid_run_store += stats.splits_by_mid_run_store; + total.splits_by_prefix_holders += stats.splits_by_prefix_holders; + total.runs_died += stats.runs_died; + total.partial_entries += stats.partial_entries; + total.max_partials = total.max_partials.max(stats.max_partials); + total.child_entries += stats.child_entries; + total.child_tombstones += stats.child_tombstones; + } + total + } + + /// Shape and memory counters of each shard. + pub fn shard_stats(&self) -> Vec { + self.shards.iter().map(ChainIndex::stats).collect() + } + + /// Every block every worker holds, as `(worker, position, content hash, prefix hash)`, + /// with the workers' ids as the caller knows them. + pub fn debug_blocks(&self) -> BTreeSet<(u32, usize, ContentHash, SequenceHash)> { + self.shards + .iter() + .enumerate() + .flat_map(|(shard, index)| { + index + .debug_blocks() + .into_iter() + .map(move |(worker, position, content, prefix)| { + (Self::global(shard, worker), position, content, prefix) + }) + }) + .collect() + } +} + +impl Default for ShardedChainIndex { + fn default() -> Self { + Self::with_shards(1) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::reference::{request_prefix_hashes, ReferenceIndexer}; + + fn content(stream: u64, position: usize) -> ContentHash { + crate::compute_content_hash(&[stream as u32, (stream >> 32) as u32, position as u32]) + } + + fn blocks_of(contents: &[ContentHash]) -> Vec { + contents + .iter() + .zip(request_prefix_hashes(contents)) + .map(|(&content_hash, seq_hash)| StoredBlock { + seq_hash, + content_hash, + }) + .collect() + } + + fn sorted(scores: OverlapScores) -> Vec<(u32, u32)> { + let mut v: Vec<(u32, u32)> = scores.scores.into_iter().collect(); + v.sort_unstable(); + v + } + + /// One shard is the chain index: the same ids, scores, blocks and counters for the same + /// events. + #[test] + fn one_shard_is_the_chain_index() { + let single = ChainIndex::with_max_workers(8); + let sharded = ShardedChainIndex::new(1, 8); + let chain: Vec = (0..40).map(|p| content(1, p)).collect(); + let mut fork = chain[..10].to_vec(); + fork.extend((10..30).map(|p| content(2, p))); + for (name, contents) in [("a", &chain), ("b", &fork), ("c", &chain)] { + let (w1, w2) = ( + single.intern_worker(name).expect("id"), + sharded.intern_worker(name).expect("id"), + ); + assert_eq!(w1, w2); + let (mut m1, mut m2) = (ChainBlockMap::default(), ChainBlockMap::default()); + single + .apply_stored(w1, &blocks_of(contents), None, &mut m1) + .expect("store"); + sharded + .apply_stored(w2, &blocks_of(contents), None, &mut m2) + .expect("store"); + let tail: Vec = blocks_of(contents)[25..] + .iter() + .map(|block| block.seq_hash) + .collect(); + single.apply_removed(w1, &tail, &mut m1); + sharded.apply_removed(w2, &tail, &mut m2); + assert_eq!(m1.len(), m2.len()); + } + for query in [&chain, &fork] { + assert_eq!( + sorted(single.find_matches(query, false)), + sorted(sharded.find_matches(query, false)) + ); + assert_eq!( + sorted(single.find_matches(query, true)), + sorted(sharded.find_matches(query, true)) + ); + } + assert_eq!(single.debug_blocks(), sharded.debug_blocks()); + assert_eq!(single.entry_count(), sharded.entry_count()); + assert_eq!(single.current_size(), sharded.current_size()); + assert_eq!(single.stats().arena_bytes, sharded.stats().arena_bytes); + assert_eq!(sharded.shards(), 1); + assert_eq!( + ShardedChainIndex::shard_of(sharded.worker_id("b").expect("b")), + 0 + ); + assert_eq!(single.is_empty(), sharded.is_empty()); + assert!(ShardedChainIndex::new(2, 4).is_empty()); + } + + /// Two shards, workers placed round robin and by hand: every lookup, with and without the + /// early exit, and the final block set agree with the reference, through stores, + /// extensions, a divergence, a hole, a clear and a worker's removal. + #[test] + fn two_shards_are_exact_against_the_reference() { + let index = ShardedChainIndex::new(2, 8); + let mut reference = ReferenceIndexer::new(); + let chain: Vec = (0..60).map(|p| content(3, p)).collect(); + let mut fork = chain[..20].to_vec(); + fork.extend((20..50).map(|p| content(4, p))); + let names = ["a", "b", "c", "d"]; + let ids: Vec = names + .iter() + .map(|name| index.intern_worker(name).expect("id")) + .collect(); + let e = index.intern_worker_in(1, "e").expect("id"); + assert_eq!( + ids.iter() + .map(|&id| ShardedChainIndex::shard_of(id)) + .collect::>(), + vec![0, 1, 0, 1] + ); + assert_eq!(ShardedChainIndex::shard_of(e), 1); + assert_eq!(index.intern_worker("e").expect("again"), e); + let mut maps: Vec = (0..5).map(|_| ChainBlockMap::default()).collect(); + let workers = [ids[0], ids[1], ids[2], ids[3], e]; + let stores = [ + (0usize, blocks_of(&chain)), + (1, blocks_of(&chain)), + (2, blocks_of(&fork)), + (3, blocks_of(&fork)), + (4, blocks_of(&chain[..30])), + ]; + for (slot, blocks) in &stores { + index + .apply_stored(workers[*slot], blocks, None, &mut maps[*slot]) + .expect("store"); + reference + .apply_stored(workers[*slot], blocks, None) + .expect("ref store"); + } + // A hole in b's chain, e extends after its last block, c is cleared, d leaves. + let hole: Vec = blocks_of(&chain)[10..15] + .iter() + .map(|block| block.seq_hash) + .collect(); + index.apply_removed(workers[1], &hole, &mut maps[1]); + reference.apply_removed(workers[1], &hole); + let more = blocks_of(&chain); + index + .apply_stored( + workers[4], + &more[30..45], + Some(more[29].seq_hash), + &mut maps[4], + ) + .expect("extend"); + reference + .apply_stored(workers[4], &more[30..45], Some(more[29].seq_hash)) + .expect("ref extend"); + index.apply_cleared(workers[2], &mut maps[2]); + reference.apply_cleared(workers[2]); + index.remove_worker(workers[3], std::mem::take(&mut maps[3])); + reference.remove_worker(workers[3]); + let mut extended = chain[..45].to_vec(); + extended.extend((45..70).map(|p| content(5, p))); + for query in [&chain, &fork, &extended, &chain[..12].to_vec()] { + let expected: Vec<(u32, u32)> = reference.find_matches(query).into_iter().collect(); + assert_eq!( + sorted(index.find_matches(query, false)), + expected, + "{query:?}" + ); + let first: Vec<(u32, u32)> = expected.iter().map(|&(w, _)| (w, 1)).collect(); + assert_eq!(sorted(index.find_matches(query, true)), first); + } + assert_eq!(index.debug_blocks(), reference.blocks()); + assert_eq!( + index.current_size(), + reference.blocks().len(), + "memberships summed over the shards" + ); + assert!(index.is_held(workers[0], &maps[0], more[0].seq_hash)); + assert!(!index.is_held(workers[1], &maps[1], more[12].seq_hash)); + assert_eq!(index.worker_block_count(workers[1]), 55); + assert_eq!(index.worker_block_count(workers[2]), 0); + assert!(!index.is_empty()); + // The chain is held on both shards: once per shard in the distinct count and the arena. + let per_shard = index.shard_stats(); + assert_eq!(per_shard.len(), 2); + assert!(per_shard.iter().all(|stats| stats.blocks_live > 0)); + assert_eq!( + index.entry_count(), + index.shard(0).entry_count() + index.shard(1).entry_count() + ); + assert!(index.entry_count() > 60 + 30); + } + + /// A caller that interns into shards of its own choosing, including the same name from + /// two places, gets one id per name and the shard it asked for first. + /// A lookup fanned out one shard at a time reports, over the shards, exactly what the union + /// lookup reports, and a shard reports only the workers it holds. + #[test] + fn a_lookup_fanned_out_per_shard_is_the_union_lookup() { + let index = ShardedChainIndex::new(2, 8); + let contents: Vec = (0..40).map(|p| content(7, p)).collect(); + let mut workers = Vec::new(); + for (shard, name, len) in [(0, "a", 40), (1, "b", 25), (0, "c", 10), (1, "d", 40)] { + let worker = index.intern_worker_in(shard, name).expect("id"); + let mut map = ChainBlockMap::default(); + index + .apply_stored(worker, &blocks_of(&contents[..len]), None, &mut map) + .expect("store"); + workers.push(worker); + } + let query = &contents[..32]; + let mut fanned: Vec<(u32, u32)> = Vec::new(); + for shard in 0..2 { + index.score_shard_into( + shard, + query, + |c| c.0, + false, + |worker, score| { + fanned.push((worker, score)); + }, + ); + } + fanned.sort_unstable(); + assert_eq!(fanned, sorted(index.find_matches(query, false))); + let mut expected = vec![ + (workers[0], 32), + (workers[1], 25), + (workers[2], 10), + (workers[3], 32), + ]; + expected.sort_unstable(); + assert_eq!(fanned, expected); + let mut from_shard_1 = Vec::new(); + index.score_shard_into( + 1, + query, + |c| c.0, + false, + |worker, _| from_shard_1.push(worker), + ); + from_shard_1.sort_unstable(); + assert_eq!(from_shard_1, vec![workers[1], workers[3]]); + } + + #[test] + fn interning_is_idempotent_across_shards() { + let index = ShardedChainIndex::new(3, 4); + let a = index.intern_worker_in(2, "a").expect("a"); + assert_eq!(ShardedChainIndex::shard_of(a), 2); + assert_eq!(index.intern_worker_in(0, "a").expect("a again"), a); + assert_eq!(index.intern_worker("a").expect("a once more"), a); + assert_eq!(index.worker_id("a"), Some(a)); + assert_eq!(index.worker_id("zz"), None); + // Round robin fills the shards evenly for names without a placement. + let placed: Vec = ["p", "q", "r", "s", "t", "u"] + .iter() + .map(|name| ShardedChainIndex::shard_of(index.intern_worker(name).expect("id"))) + .collect(); + assert_eq!(placed.iter().filter(|&&s| s == 0).count(), 2); + assert_eq!(placed.iter().filter(|&&s| s == 1).count(), 2); + assert_eq!(placed.iter().filter(|&&s| s == 2).count(), 2); + // A shard's slots run out on their own. + for name in ["v", "w", "x", "y"] { + let _ = index.intern_worker_in(0, name); + } + assert!(index.intern_worker_in(0, "overflow").is_err()); + } +} diff --git a/crates/kv_index/tests/churn_gate.rs b/crates/kv_index/tests/churn_gate.rs new file mode 100644 index 0000000000..d412e5ddd1 --- /dev/null +++ b/crates/kv_index/tests/churn_gate.rs @@ -0,0 +1,94 @@ +//! Gate-sized churn: a few hundred thousand events of block-LRU churn with decode tails, at the +//! index's own speed, checked against the reference indexer and held to the content's shape. +//! +//! The churn bench (`benches/churn.rs`) runs the same generator for the series; this is the part +//! of it that fails a gate: exactness under churn always, and the bound on runs live that keeps +//! fragmentation from coming back silently where a divergence inside a run hangs a child off the +//! offset and leaves the run whole. + +use kv_index::{ + churn::{Churn, ChurnConfig, FreeOrder}, + ReferenceIndexer, ShardedChainIndex, +}; + +fn config(seed: u64) -> ChurnConfig { + ChurnConfig { + workers: 32, + chains: 600, + min_chain_len: 16, + max_chain_len: 128, + share: 0.7, + cache_blocks: 3000, + free_order: FreeOrder::TailFirst, + decode_blocks: 2, + decode_probability: 0.5, + restart_every: 7000, + refill_requests: 300, + seed, + } +} + +fn requests() -> u64 { + std::env::var("KV_INDEX_CHURN_REQUESTS") + .ok() + .and_then(|v| v.parse().ok()) + .unwrap_or(30_000) +} + +/// Every lookup during and after the churn agrees with the reference, under one shard and two. +#[test] +fn churn_with_decode_tails_and_restarts_is_exact() { + for shards in [1usize, 2] { + let index = ShardedChainIndex::new(shards, 64); + let mut reference = ReferenceIndexer::new(); + let mut churn = Churn::new(config(7 + shards as u64), &index); + let total = requests(); + for request in 1..=total { + let report = churn.step(&index, Some(&mut reference)); + assert!( + report.lookup_ns < 10_000_000_000, + "a lookup took {} ns", + report.lookup_ns + ); + if request % 5000 == 0 { + churn + .check_exact(&index, &reference, 100) + .unwrap_or_else(|why| panic!("shards {shards}, request {request}: {why}")); + } + } + churn + .check_exact(&index, &reference, 1000) + .unwrap_or_else(|why| panic!("shards {shards}, at the end: {why}")); + let stats = index.stats(); + assert_eq!( + stats.splits_by_branch, 0, + "a divergence inside a run hangs a child off it and splits nothing" + ); + assert!(churn.restarts > 0, "workers restarted"); + assert!(index.current_size() > 0); + } +} + +/// Runs live stay within a small multiple of the content's shape: one run per chain and per +/// divergence point, plus the live private decode tails, each a run of its own. Before a +/// divergence inside a run stopped splitting it, every prompt end left a boundary behind and the +/// count ran away with the requests; this keeps that from coming back. +#[test] +fn runs_live_stay_within_the_content_shape() { + let index = ShardedChainIndex::new(1, 64); + let mut churn = Churn::new(config(11), &index); + let total = requests() * 4; + let shape = churn.pool.chains.len() + churn.pool.divergence_points(); + for request in 1..=total { + churn.step(&index, None); + if request % 20_000 == 0 || request == total { + let runs_live = index.stats().runs_live; + let bound = 2 * shape + churn.live_decode_blocks(); + assert!( + runs_live <= bound, + "request {request}: runs live {runs_live} above the content's shape ({shape} chains and branch points, {} live decode blocks; bound {bound})", + churn.live_decode_blocks() + ); + } + } +} diff --git a/crates/kv_index/tests/common/mod.rs b/crates/kv_index/tests/common/mod.rs new file mode 100644 index 0000000000..228bcb67f2 --- /dev/null +++ b/crates/kv_index/tests/common/mod.rs @@ -0,0 +1,57 @@ +//! Helpers shared by the exactness and concurrency suites: a seeded xorshift generator and the +//! block builders that name a chain's blocks as an engine with chain hashes would. +#![allow(dead_code)] + +use kv_index::{request_prefix_hashes, ContentHash, StoredBlock}; + +/// xorshift64*, enough for a deterministic corpus without pulling a dependency into the test. +pub struct Rng(u64); + +impl Rng { + pub fn new(seed: u64) -> Self { + Self(seed.max(1)) + } + + pub fn next(&mut self) -> u64 { + let mut x = self.0; + x ^= x >> 12; + x ^= x << 25; + x ^= x >> 27; + self.0 = x; + x.wrapping_mul(0x2545_F491_4F6C_DD1D) + } + + pub fn below(&mut self, n: usize) -> usize { + (self.next() % n as u64) as usize + } + + pub fn range(&mut self, lo: usize, hi_inclusive: usize) -> usize { + lo + self.below(hi_inclusive - lo + 1) + } + + pub fn chance(&mut self, numerator: u64, denominator: u64) -> bool { + self.next() % denominator < numerator + } +} + +pub fn content(stream: u64, position: usize) -> ContentHash { + kv_index::compute_content_hash(&[ + (stream & 0xffff_ffff) as u32, + (stream >> 32) as u32, + position as u32, + ]) +} + +/// Blocks of a content sequence as an engine would hash them: the engine hash is the chain hash +/// of the contents so far, which is unique per distinct prefix and shared by every worker that +/// stores the same prefix. +pub fn blocks_of(contents: &[ContentHash]) -> Vec { + contents + .iter() + .zip(request_prefix_hashes(contents)) + .map(|(&content_hash, seq_hash)| StoredBlock { + seq_hash, + content_hash, + }) + .collect() +} diff --git a/crates/kv_index/tests/concurrency_chain.rs b/crates/kv_index/tests/concurrency_chain.rs new file mode 100644 index 0000000000..67793d640d --- /dev/null +++ b/crates/kv_index/tests/concurrency_chain.rs @@ -0,0 +1,474 @@ +//! Concurrency harness for `ChainIndex`: many event lanes writing shared chains at once while +//! readers look them up, then the end state against the `ReferenceIndexer`. +//! +//! Lanes own one worker each and replay random stores, extensions, tail and middle evictions, +//! clears and worker replacements over a shared pool of conversations, so runs are joined, split, +//! truncated and unlinked by different threads at the same time. Every lane logs what it applied; +//! after the lanes stop, the logs are replayed in application order into the reference, which +//! must hold exactly the same blocks and score every pool chain the same way. Readers run +//! throughout and check the invariants that hold under any interleaving: no score exceeds the +//! request, early-exit scores are 1, and nothing panics. +#![expect(clippy::expect_used)] + +use std::{ + collections::{BTreeMap, BTreeSet}, + io::Write, + sync::{ + atomic::{AtomicBool, AtomicU32, AtomicU64, AtomicU8, Ordering}, + Arc, Mutex, + }, + thread, + time::{Duration, Instant}, +}; + +use kv_index::{ + ChainBlockMap, ContentHash, ReferenceIndexer, SequenceHash, ShardedChainIndex, StoredBlock, +}; + +mod common; +use common::{blocks_of, content, Rng}; + +/// A shared pool of conversations: a few prompts, each continued by several turns, with sibling +/// turns branching off at various depths. +fn pool(rng: &mut Rng) -> Vec> { + let mut chains: Vec> = Vec::new(); + let mut stream = 1u64; + for _ in 0..6 { + stream += 1; + let prompt: Vec = (0..rng.range(8, 40)).map(|p| content(stream, p)).collect(); + for _ in 0..12 { + let base = if chains.is_empty() || rng.below(3) == 0 { + prompt.clone() + } else { + let parent = &chains[chains.len() - 1 - rng.below(chains.len().min(12))]; + parent[..rng.range(1, parent.len())].to_vec() + }; + stream += 1; + let mut chain = base; + chain.extend((0..rng.range(1, 24)).map(|p| content(stream, p))); + chains.push(chain); + } + } + chains +} + +enum Event { + Stored { + blocks: Vec, + parent: Option, + }, + Removed(Vec), + Cleared, + WorkerRemoved, +} + +struct Logged { + seq: u64, + worker: u32, + event: Event, +} + +struct Lane<'a> { + index: &'a ShardedChainIndex, + pool: &'a [Vec], + clock: &'a AtomicU64, + worker: u32, + name_counter: u64, + lane: usize, + map: ChainBlockMap, + rng: Rng, + log: Vec, + steps: u32, +} + +impl Lane<'_> { + /// The index credits this worker with exactly the blocks its lane map names. A holding the + /// index drops behind the map's back (or keeps after the map forgot it) shows here, at the + /// step that caused it, not at the end state or at a later store under the block. + fn check_count(&self, when: &str) { + let credited = self.index.worker_block_count(self.worker); + assert!( + credited == self.map.len(), + "lane {} worker {} step {}: the index credits {credited} blocks, the lane map names {} ({when})", + self.lane, + self.worker, + self.steps, + self.map.len() + ); + } + + fn record(&mut self, event: Event) { + let seq = self.clock.fetch_add(1, Ordering::Relaxed); + self.record_at(seq, event); + } + + fn record_at(&mut self, seq: u64, event: Event) { + self.log.push(Logged { + seq, + worker: self.worker, + event, + }); + } + + fn store(&mut self, contents: &[ContentHash]) { + let blocks = blocks_of(contents); + let mut known = 0; + while known < blocks.len() + && self + .index + .is_held(self.worker, &self.map, blocks[known].seq_hash) + { + known += 1; + } + let start = if known == blocks.len() { + blocks.len() - 1 + } else { + known + }; + let parent = (start > 0).then(|| blocks[start - 1].seq_hash); + let outcome = self + .index + .apply_stored(self.worker, &blocks[start..], parent, &mut self.map); + assert!( + outcome.is_ok(), + "lane {}: store after a held parent failed: {outcome:?}", + self.lane + ); + self.record(Event::Stored { + blocks: blocks[start..].to_vec(), + parent, + }); + self.check_count("after a store"); + } + + /// Remove a contiguous range of the blocks this lane holds on a pool chain. + fn remove_some(&mut self, contents: &[ContentHash]) { + let blocks = blocks_of(contents); + let held: Vec = (0..blocks.len()) + .filter(|&i| { + self.index + .is_held(self.worker, &self.map, blocks[i].seq_hash) + }) + .collect(); + if held.is_empty() { + return; + } + let from = self.rng.below(held.len()); + let to = match self.rng.below(4) { + 0 => held.len(), + _ => (from + self.rng.range(1, 3)).min(held.len()), + }; + let hashes: Vec = + held[from..to].iter().map(|&i| blocks[i].seq_hash).collect(); + self.index + .apply_removed(self.worker, &hashes, &mut self.map); + self.record(Event::Removed(hashes)); + self.check_count("after a removal"); + } + + fn step(&mut self) { + let chain = self.pool[self.rng.below(self.pool.len())].clone(); + // Worker replacement is part of every run, not a rare roll: each lane swaps its worker + // every 97 steps (so lanes do it at different times) besides the random 1%. + self.steps += 1; + if self.steps.is_multiple_of(97) { + self.replace_worker(); + return; + } + match self.rng.below(1000) { + 0..=549 => { + let len = if self.rng.below(2) == 0 { + chain.len() + } else { + self.rng.range(1, chain.len()) + }; + self.store(&chain[..len]); + } + 550..=899 => self.remove_some(&chain), + 900..=984 => { + // Walk a chain in two turns: a prefix now, the rest right after (decode extension). + let cut = self.rng.range(1, chain.len()); + self.store(&chain[..cut]); + self.store(&chain); + } + 985..=989 => { + self.index.apply_cleared(self.worker, &mut self.map); + self.record(Event::Cleared); + self.check_count("after a clear"); + } + _ => self.replace_worker(), + } + } + + /// Remove this lane's worker and intern a fresh one (its slot may come back reused). + fn replace_worker(&mut self) { + let map = std::mem::take(&mut self.map); + // Sequenced before the removal: another lane may intern the freed slot and store under + // the same id before this lane gets to record, and the replay must see the removal first. + let seq = self.clock.fetch_add(1, Ordering::Relaxed); + self.index.remove_worker(self.worker, map); + self.record_at(seq, Event::WorkerRemoved); + self.name_counter += 1; + let name = format!("lane-{}-{}", self.lane, self.name_counter); + self.worker = self.index.intern_worker(&name).expect("worker slot"); + self.check_count("after a worker replacement"); + } +} + +fn scores( + index: &ShardedChainIndex, + query: &[ContentHash], + early_exit: bool, +) -> BTreeMap { + index + .find_matches(query, early_exit) + .scores + .into_iter() + .collect() +} + +const READERS: usize = 4; +/// The run normally ends within seconds; past this the watchdog reports every thread's progress +/// to the process's stderr and aborts, so a wedged run fails instead of holding a gate. +const DEADLINE: Duration = Duration::from_secs(300); +const PHASE_LANES: u8 = 0; +const PHASE_READERS: u8 = 1; +const PHASE_DONE: u8 = 2; + +/// Write straight to the process's stderr: the test harness holds the test thread's printed +/// output back until the test ends (its capture sits under the print macros, not under a raw +/// write), and a report that precedes an abort must come out before it. +fn shout(message: &str) { + let line = format!("{message}\n"); + let _ = std::io::stderr().write_all(line.as_bytes()); +} + +fn panic_text(payload: &(dyn std::any::Any + Send)) -> String { + payload + .downcast_ref::() + .cloned() + .or_else(|| payload.downcast_ref::<&str>().map(|s| (*s).to_string())) + .unwrap_or_else(|| "a non-string panic payload".to_string()) +} + +#[test] +fn concurrent_lanes_and_readers_end_in_the_reference_state() { + let lanes = 16usize; + let steps = std::env::var("KV_INDEX_CONCURRENCY_STEPS") + .ok() + .and_then(|v| v.parse().ok()) + .unwrap_or(4000usize); + let mut rng = Rng::new(20261005); + let pool = pool(&mut rng); + let shards = std::env::var("KV_INDEX_CONCURRENCY_SHARDS") + .ok() + .and_then(|value| value.parse().ok()) + .unwrap_or(1usize); + // Lanes intern their workers round robin over the shards, so with two shards every + // other lane writes the other index and every lookup merges both. + let index = ShardedChainIndex::new(shards, 64); + let clock = AtomicU64::new(0); + let stop = Arc::new(AtomicBool::new(false)); + let logs: Mutex> = Mutex::new(Vec::new()); + let reader_lookups = Arc::new(AtomicU64::new(0)); + // What the watchdog reports if the run overstays its deadline: steps done per lane, rounds + // done per reader and the phase the run is in. + let lane_steps: Arc> = Arc::new((0..lanes).map(|_| AtomicU32::new(0)).collect()); + let reader_rounds: Arc> = + Arc::new((0..READERS).map(|_| AtomicU64::new(0)).collect()); + let phase = Arc::new(AtomicU8::new(PHASE_LANES)); + let watchdog = { + let (lane_steps, reader_rounds, phase) = ( + Arc::clone(&lane_steps), + Arc::clone(&reader_rounds), + Arc::clone(&phase), + ); + thread::spawn(move || { + let deadline = Instant::now() + DEADLINE; + while phase.load(Ordering::Relaxed) != PHASE_DONE { + if Instant::now() >= deadline { + let lane_steps: Vec = lane_steps + .iter() + .map(|s| s.load(Ordering::Relaxed)) + .collect(); + let reader_rounds: Vec = reader_rounds + .iter() + .map(|r| r.load(Ordering::Relaxed)) + .collect(); + shout(&format!( + "concurrency harness: {} s deadline exceeded in phase {} (0 lanes running, 1 lanes joined and readers stopping); steps per lane {lane_steps:?} of {steps}; rounds per reader {reader_rounds:?}; aborting", + DEADLINE.as_secs(), + phase.load(Ordering::Relaxed) + )); + std::process::abort(); + } + thread::sleep(Duration::from_millis(200)); + } + }) + }; + + thread::scope(|scope| { + let mut lane_threads = Vec::with_capacity(lanes); + for lane in 0..lanes { + let (index, pool, clock, logs, lane_steps) = + (&index, &pool, &clock, &logs, &lane_steps); + let seed = rng.next(); + lane_threads.push(scope.spawn(move || { + let worker = index + .intern_worker(&format!("lane-{lane}-0")) + .expect("worker slot"); + let mut state = Lane { + index, + pool, + clock, + worker, + name_counter: 0, + lane, + map: ChainBlockMap::default(), + rng: Rng::new(seed), + log: Vec::new(), + steps: 0, + }; + for _ in 0..steps { + state.step(); + lane_steps[lane].fetch_add(1, Ordering::Relaxed); + } + logs.lock().unwrap().extend(state.log); + })); + } + let mut reader_threads = Vec::with_capacity(READERS); + for reader in 0..READERS { + let (index, pool, stop, reader_lookups, reader_rounds) = + (&index, &pool, &stop, &reader_lookups, &reader_rounds); + let seed = rng.next() ^ reader as u64; + reader_threads.push(scope.spawn(move || { + let mut rng = Rng::new(seed); + while !stop.load(Ordering::Relaxed) { + let chain = &pool[rng.below(pool.len())]; + let mut query = chain.clone(); + if rng.below(3) == 0 { + let at = rng.below(query.len()); + query[at] = content(u64::MAX - reader as u64, at); + } + let full = scores(index, &query, false); + for (&worker, &score) in &full { + assert!( + score as usize <= query.len(), + "worker {worker} scored {score} on a {}-block request", + query.len() + ); + assert!(score > 0, "worker {worker} reported with score 0"); + } + let early = scores(index, &query, true); + assert!(early.values().all(|&score| score == 1)); + reader_lookups.fetch_add(1, Ordering::Relaxed); + reader_rounds[reader].fetch_add(1, Ordering::Relaxed); + } + })); + } + // Termination is unconditional. The lanes are finite and joined first; a panic among + // them is kept rather than raised here, because raising it would leave the readers + // running under a scope that waits for them (a lane's failed assertion then showed as + // a hang, its message held back by the test harness's output capture). The readers are + // told to stop and joined, and only then is the first panic re-raised. The log length + // is not a completion signal: a removal that finds nothing held logs nothing and the + // two-turn store logs twice. + let mut first_panic = None; + for (lane, handle) in lane_threads.into_iter().enumerate() { + if let Err(payload) = handle.join() { + shout(&format!( + "concurrency harness: lane {lane} panicked: {}", + panic_text(payload.as_ref()) + )); + first_panic.get_or_insert(payload); + } + } + phase.store(PHASE_READERS, Ordering::Relaxed); + stop.store(true, Ordering::Relaxed); + for (reader, handle) in reader_threads.into_iter().enumerate() { + if let Err(payload) = handle.join() { + shout(&format!( + "concurrency harness: reader {reader} panicked: {}", + panic_text(payload.as_ref()) + )); + first_panic.get_or_insert(payload); + } + } + phase.store(PHASE_DONE, Ordering::Relaxed); + if let Some(payload) = first_panic { + std::panic::resume_unwind(payload); + } + }); + watchdog.join().expect("the watchdog thread returned"); + + let mut logs = logs.into_inner().unwrap(); + logs.sort_by_key(|entry| entry.seq); + let mut reference = ReferenceIndexer::new(); + for entry in &logs { + match &entry.event { + Event::Stored { blocks, parent } => { + reference + .apply_stored(entry.worker, blocks, *parent) + .expect("the lane stored after a held parent"); + } + Event::Removed(hashes) => reference.apply_removed(entry.worker, hashes), + Event::Cleared => reference.apply_cleared(entry.worker), + Event::WorkerRemoved => reference.remove_worker(entry.worker), + } + } + assert!( + reader_lookups.load(Ordering::Relaxed) > 1000, + "readers barely ran: {} lookups", + reader_lookups.load(Ordering::Relaxed) + ); + + // Every engine hash here is a prefix hash of its chain, so a block has one place and no + // hash ever moves: a moved-hash count is a re-store taken for a move. + assert_eq!( + index.stats().moved_hashes, + 0, + "the index released holdings for hashes it took as moved" + ); + let produced = index.debug_blocks(); + let expected = reference.blocks(); + let missing: Vec<_> = expected.difference(&produced).take(5).collect(); + let phantom: Vec<_> = produced.difference(&expected).take(5).collect(); + assert!( + missing.is_empty() && phantom.is_empty(), + "end state differs after {} events: {} reference blocks, {} index blocks; missing e.g. \ + {missing:?}; phantom e.g. {phantom:?}", + logs.len(), + expected.len(), + produced.len() + ); + let mut checked = 0; + for chain in &pool { + for query in [chain.clone(), chain[..chain.len().div_ceil(2)].to_vec()] { + assert_eq!( + scores(&index, &query, false), + reference.find_matches(&query), + "lookup differs for a {}-block query", + query.len() + ); + checked += 1; + } + } + assert!(checked > 100); + // Distinct blocks are counted per shard: content held on two shards is stored twice. + assert_eq!( + index.entry_count(), + expected + .iter() + .map(|(worker, position, content, prefix)| { + ( + ShardedChainIndex::shard_of(*worker), + *position, + *content, + *prefix, + ) + }) + .collect::>() + .len(), + "distinct block counter" + ); +} diff --git a/crates/kv_index/tests/exactness_chain.rs b/crates/kv_index/tests/exactness_chain.rs new file mode 100644 index 0000000000..e83bd046d2 --- /dev/null +++ b/crates/kv_index/tests/exactness_chain.rs @@ -0,0 +1,937 @@ +//! Exactness harness: the run-compressed `ChainIndex` against the `ReferenceIndexer`. +//! +//! The same seeded corpus as `exactness_positional.rs` (stores, extensions, divergent siblings, +//! tail and whole-chain removals, clears, worker removals) plus evictions that leave holes: a +//! block in the middle of a chain, or its first block, goes away while the worker keeps the blocks +//! after it. +//! An engine with prefix caching produces exactly that when it evicts by block, and it keeps +//! reporting the later blocks as stored until it evicts them too. The reference says such a chain +//! matches up to the hole and no further; the chain index must say the same, and must heal when the +//! engine re-stores the missing block after its parent. +//! +//! After every round of events, lookups built from live chains and from mutated chains must score +//! identically in both indexers; at the end, the chain index must hold exactly the reference's +//! blocks. Scale with `KV_INDEX_EXACTNESS_EVENTS` (default 20000) and `KV_INDEX_EXACTNESS_SEED`. +#![expect(clippy::expect_used)] + +use std::collections::{BTreeMap, BTreeSet}; + +use kv_index::{ + ChainBlockMap, ChainIndex, ContentHash, ReferenceIndexer, SequenceHash, ShardedChainIndex, + StoredBlock, +}; +use rustc_hash::FxHashMap; + +mod common; +use common::{blocks_of, content, Rng}; + +struct Held { + worker: u32, + contents: Vec, +} + +struct Harness { + production: ShardedChainIndex, + reference: ReferenceIndexer, + maps: FxHashMap, + workers: Vec, + held: Vec, + prompts: Vec>, + next_stream: u64, + next_worker: u32, + rng: Rng, + events: usize, + stored_blocks: usize, + holes: usize, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)] +enum QueryKind { + Exact, + Prefix, + MiddleReplaced, + SuffixReplaced, + Extended, + Unknown, +} + +#[derive(Default)] +struct Mismatches { + by_kind: BTreeMap, + examples: Vec, + lookups: usize, +} + +impl Harness { + fn new(seed: u64, workers: usize) -> Self { + let mut rng = Rng::new(seed); + let mut harness = Self { + // `KV_INDEX_EXACTNESS_SHARDS` shards (default one, the chain index itself); workers are + // interned round robin, so two shards hold every other worker and share most content. + production: ShardedChainIndex::new( + env_or("KV_INDEX_EXACTNESS_SHARDS", 1) as usize, + 1024, + ), + reference: ReferenceIndexer::new(), + maps: FxHashMap::default(), + workers: Vec::new(), + held: Vec::new(), + prompts: Vec::new(), + next_stream: 1, + next_worker: 0, + rng: Rng::new(seed ^ 0x9e37_79b9_7f4a_7c15), + events: 0, + stored_blocks: 0, + holes: 0, + }; + for _ in 0..workers { + harness.add_worker(); + } + let prompt_count = rng.range(3, 8); + for _ in 0..prompt_count { + let len = rng.range(8, 48); + let stream = harness.fresh_stream(); + harness + .prompts + .push((0..len).map(|p| content(stream, p)).collect()); + } + harness + } + + fn fresh_stream(&mut self) -> u64 { + self.next_stream += 1; + self.next_stream + } + + fn add_worker(&mut self) -> u32 { + let url = format!("http://worker-{}:8000", self.next_worker); + self.next_worker += 1; + let id = self.production.intern_worker(&url).expect("worker id"); + self.maps.insert(id, ChainBlockMap::default()); + self.workers.push(id); + id + } + + fn random_worker(&mut self) -> u32 { + self.workers[self.rng.below(self.workers.len())] + } + + /// A new user turn: a prompt prefix plus fresh blocks. + fn new_chain(&mut self) -> Vec { + let prompt = &self.prompts[self.rng.below(self.prompts.len())]; + let mut contents = prompt.clone(); + let turn = self.rng.range(4, 32); + let stream = self.fresh_stream(); + contents.extend((0..turn).map(|p| content(stream, p))); + contents + } + + /// Store `contents` on `worker` as the engine would: only the suffix the worker does not hold, + /// after the last block it does hold. A fully held chain is re-stored for its last block, which + /// exercises the duplicate-store path. A chain with a hole is re-stored from the hole on, which + /// is how an engine heals one. + fn store(&mut self, worker: u32, contents: Vec) { + let blocks = blocks_of(&contents); + let held = self.maps.get(&worker).expect("worker map"); + let mut known = 0; + while known < blocks.len() + && self + .production + .is_held(worker, held, blocks[known].seq_hash) + { + known += 1; + } + let start = if known == blocks.len() { + blocks.len() - 1 + } else { + known + }; + let parent = if start == 0 { + None + } else { + Some(blocks[start - 1].seq_hash) + }; + let map = self.maps.get_mut(&worker).expect("worker map"); + let produced = self + .production + .apply_stored(worker, &blocks[start..], parent, map); + let referenced = self + .reference + .apply_stored(worker, &blocks[start..], parent); + assert_eq!( + produced.is_ok(), + referenced.is_ok(), + "store outcome differs: production {produced:?}, reference {referenced:?}" + ); + if produced.is_ok() { + self.stored_blocks += blocks.len() - start; + self.held.push(Held { worker, contents }); + } + self.events += 1; + } + + fn pick_held(&mut self) -> Option { + if self.held.is_empty() { + return None; + } + let index = self.rng.below(self.held.len()); + if self.maps.contains_key(&self.held[index].worker) { + Some(index) + } else { + self.held.swap_remove(index); + None + } + } + + fn remove(&mut self, worker: u32, hashes: &[SequenceHash]) { + let map = self.maps.get_mut(&worker).expect("worker map"); + self.production.apply_removed(worker, hashes, map); + self.reference.apply_removed(worker, hashes); + self.events += 1; + } + + /// Evict a tail: the chain's blocks from a random position on, together with the same + /// positions of every other held chain of that worker that runs through them. + fn remove_tail(&mut self) { + let Some(index) = self.pick_held() else { + return; + }; + let worker = self.held[index].worker; + let contents = self.held[index].contents.clone(); + let keep = self.rng.below(contents.len()); + let mut hashes: Vec = Vec::new(); + for held in self.held.iter_mut().filter(|h| h.worker == worker) { + let shared = held + .contents + .iter() + .zip(&contents) + .take_while(|(a, b)| a == b) + .count(); + if shared > keep { + let blocks = blocks_of(&held.contents); + hashes.extend(blocks[keep..].iter().map(|b| b.seq_hash)); + held.contents.truncate(keep); + } + } + hashes.sort_unstable_by_key(|h| h.0); + hashes.dedup(); + self.remove(worker, &hashes); + self.held.retain(|h| !h.contents.is_empty()); + } + + /// Evict one to three blocks strictly inside a chain, or its first block, and keep the rest: + /// a hole. The held chain is left as it is, so later lookups cross the hole and a later + /// extension re-stores from the hole on. + fn remove_middle(&mut self) { + let Some(index) = self.pick_held() else { + return; + }; + let worker = self.held[index].worker; + let contents = self.held[index].contents.clone(); + if contents.len() < 3 { + return; + } + let (from, to) = if self.rng.chance(1, 8) { + (0, 1) + } else { + let from = self.rng.range(1, contents.len() - 2); + let to = (from + self.rng.range(1, 3)).min(contents.len() - 1); + (from, to) + }; + let blocks = blocks_of(&contents); + let hashes: Vec = blocks[from..to].iter().map(|b| b.seq_hash).collect(); + self.remove(worker, &hashes); + self.holes += 1; + } + + /// Evict a whole conversation: the chain's blocks beyond the longest prefix it shares with + /// another held chain of the same worker. + fn remove_chain(&mut self) { + let Some(index) = self.pick_held() else { + return; + }; + let worker = self.held[index].worker; + let contents = self.held[index].contents.clone(); + let shared = self + .held + .iter() + .enumerate() + .filter(|(i, h)| *i != index && h.worker == worker) + .map(|(_, h)| { + h.contents + .iter() + .zip(&contents) + .take_while(|(a, b)| a == b) + .count() + }) + .max() + .unwrap_or(0); + let blocks = blocks_of(&contents); + let hashes: Vec = blocks[shared..].iter().map(|b| b.seq_hash).collect(); + self.remove(worker, &hashes); + self.held.swap_remove(index); + } + + fn clear_worker(&mut self) { + let worker = self.random_worker(); + let map = self.maps.get_mut(&worker).expect("worker map"); + self.production.apply_cleared(worker, map); + self.reference.apply_cleared(worker); + self.held.retain(|h| h.worker != worker); + self.events += 1; + } + + fn remove_worker(&mut self) { + if self.workers.len() < 2 { + return; + } + let position = self.rng.below(self.workers.len()); + let worker = self.workers.swap_remove(position); + let map = self.maps.remove(&worker).expect("worker map"); + self.production.remove_worker(worker, map); + self.reference.remove_worker(worker); + self.held.retain(|h| h.worker != worker); + self.add_worker(); + self.events += 1; + } + + /// Store `blocks` after `parent` exactly as given, on both indexers; the outcomes must agree. + /// Returns whether the store was accepted. + fn store_blocks( + &mut self, + worker: u32, + blocks: &[StoredBlock], + parent: Option, + ) -> bool { + let map = self.maps.get_mut(&worker).expect("worker map"); + let produced = self.production.apply_stored(worker, blocks, parent, map); + let referenced = self.reference.apply_stored(worker, blocks, parent); + assert_eq!( + produced.is_ok(), + referenced.is_ok(), + "store outcome differs after {} events: production {produced:?}, reference {referenced:?}", + self.events + ); + self.events += 1; + if produced.is_ok() { + self.stored_blocks += blocks.len(); + } + produced.is_ok() + } + + /// The engine stores a tail after a parent it names regardless of whether the index still + /// holds it (removed, or never stored on this worker): both indexers reject it the same way, + /// and the gateway's fallback then files the tail without a parent, at position 0 under the + /// tail's own hashes. + fn store_tail_under_named_parent(&mut self) { + let Some(index) = self.pick_held() else { + return; + }; + let worker = if self.rng.chance(1, 4) { + self.random_worker() + } else { + self.held[index].worker + }; + let contents = self.held[index].contents.clone(); + if contents.len() < 2 { + return; + } + let blocks = blocks_of(&contents); + let cut = self.rng.range(1, blocks.len() - 1); + let parent = Some(blocks[cut - 1].seq_hash); + if !self.store_blocks(worker, &blocks[cut..], parent) && self.rng.chance(3, 4) { + self.store_blocks(worker, &blocks[cut..], None); + } + } + + /// The engine announces a held chain whole, from the root: blocks it filed elsewhere (a + /// parentless tail) move back to their positions, blocks already in place stay. + fn restore_whole_chain(&mut self) { + let Some(index) = self.pick_held() else { + return; + }; + let worker = self.held[index].worker; + let contents = self.held[index].contents.clone(); + let blocks = blocks_of(&contents); + self.store_blocks(worker, &blocks, None); + } + + /// A second or third physical copy of a block the engine already reported: the same store + /// again, one block or a run of blocks, under its real parent. + fn store_again(&mut self) { + let Some(index) = self.pick_held() else { + return; + }; + let worker = self.held[index].worker; + let contents = self.held[index].contents.clone(); + let blocks = blocks_of(&contents); + let from = self.rng.below(blocks.len()); + let to = (from + self.rng.range(1, 4)).min(blocks.len()); + let parent = (from > 0).then(|| blocks[from - 1].seq_hash); + self.store_blocks(worker, &blocks[from..to], parent); + } + + /// A twin: the same content under the same parent, named by another engine hash (an engine + /// whose hash carries more than the chain, or a feed read against the wrong tokens). The + /// index files it onto the position it holds and counts the conflict; the reference keeps + /// one block per engine hash. `salt` keeps a twin's names distinct per round. + fn store_twins(&mut self, salt: u64) -> Vec { + let Some(index) = self.pick_held() else { + return Vec::new(); + }; + let worker = self.held[index].worker; + let contents = self.held[index].contents.clone(); + let blocks = blocks_of(&contents); + let from = self.rng.below(blocks.len()); + let to = (from + self.rng.range(1, 6)).min(blocks.len()); + let twins: Vec = blocks[from..to] + .iter() + .map(|b| StoredBlock { + seq_hash: SequenceHash(b.seq_hash.0 ^ salt.rotate_left(17) ^ 0x7777_0000_0000_0001), + content_hash: b.content_hash, + }) + .collect(); + let parent = (from > 0).then(|| blocks[from - 1].seq_hash); + if self.store_blocks(worker, &twins, parent) { + twins.iter().map(|b| b.seq_hash).collect() + } else { + Vec::new() + } + } + + /// One step of the wild corpus: the replayed corpus's events plus stores under parents the + /// index may not hold, whole-chain re-announcements, repeated stores and, with `twins`, + /// twin names; `twin_names` collects twin names for later removal by name. + fn step_wild(&mut self, twins: bool, twin_names: &mut Vec<(u32, SequenceHash)>) { + let roll = self.rng.below(1000); + match roll { + 0..=599 => self.step(), + 600..=729 => self.store_tail_under_named_parent(), + 730..=809 => self.restore_whole_chain(), + 810..=889 => self.store_again(), + 890..=949 => { + if twins { + let salt = self.rng.next(); + let worker = self.held.last().map_or(0, |h| h.worker); + let names = self.store_twins(salt); + twin_names.extend(names.into_iter().map(|n| (worker, n))); + } else { + self.remove_middle(); + } + } + 950..=979 => { + // Remove a twin by its own name, or the original name under a twin. + if !twin_names.is_empty() && self.rng.chance(1, 2) { + let at = self.rng.below(twin_names.len()); + let (worker, name) = twin_names.swap_remove(at); + if self.maps.contains_key(&worker) { + self.remove(worker, &[name]); + } + } else { + self.remove_tail(); + } + } + _ => self.clear_worker(), + } + } + + fn step(&mut self) { + let roll = self.rng.below(1000); + match roll { + 0..=369 => { + let chain = self.new_chain(); + let worker = self.random_worker(); + self.store(worker, chain); + } + 370..=599 => { + // Extend a held chain by another turn on the same worker. + let Some(index) = self.pick_held() else { + return; + }; + let worker = self.held[index].worker; + let mut contents = self.held[index].contents.clone(); + let turn = self.rng.range(4, 32); + let stream = self.fresh_stream(); + contents.extend((0..turn).map(|p| content(stream, p))); + self.store(worker, contents); + } + 600..=729 => { + // A sibling that diverges at a random position, including 1 and the last block, + // stored on a random worker (often another one, which shares the prefix blocks). + let Some(index) = self.pick_held() else { + return; + }; + let base = &self.held[index].contents; + if base.len() < 2 { + return; + } + let divergence = match self.rng.below(10) { + 0 => 1, + 1 => base.len() - 1, + 2 => base.len(), + _ => self.rng.range(1, base.len()), + }; + let mut contents: Vec = base[..divergence].to_vec(); + let turn = self.rng.range(1, 24); + let stream = self.fresh_stream(); + contents.extend((0..turn).map(|p| content(stream, p))); + let worker = self.random_worker(); + self.store(worker, contents); + } + 730..=819 => self.remove_tail(), + 820..=899 => self.remove_middle(), + 900..=949 => self.remove_chain(), + 950..=984 => { + if self.workers.len() < 64 && self.rng.chance(1, 4) { + self.add_worker(); + } else { + let Some(index) = self.pick_held() else { + return; + }; + let contents = self.held[index].contents.clone(); + let worker = self.random_worker(); + self.store(worker, contents); + } + } + 985..=994 => self.clear_worker(), + _ => self.remove_worker(), + } + } + + fn production_scores(&self, query: &[ContentHash]) -> BTreeMap { + self.production + .find_matches(query, false) + .scores + .into_iter() + .collect() + } + + fn check_lookup(&mut self, kind: QueryKind, query: Vec, out: &mut Mismatches) { + if query.is_empty() { + return; + } + out.lookups += 1; + let produced = self.production_scores(&query); + let expected = self.reference.find_matches(&query); + if produced != expected { + *out.by_kind.entry(kind).or_default() += 1; + if out.examples.len() < 12 { + let diff: Vec = expected + .keys() + .chain(produced.keys()) + .collect::>() + .into_iter() + .filter(|w| produced.get(w) != expected.get(w)) + .map(|w| { + format!( + "worker {w}: production {:?}, reference {:?}", + produced.get(w), + expected.get(w) + ) + }) + .collect(); + out.examples.push(format!( + "{kind:?} query of {} blocks after {} events: {}", + query.len(), + self.events, + diff.join("; ") + )); + } + } + // early_exit reports exactly the workers that match position 0, each with score 1. + let early: BTreeMap = self + .production + .find_matches(&query, true) + .scores + .into_iter() + .collect(); + let expected_early: BTreeMap = expected.keys().map(|&w| (w, 1)).collect(); + if early != expected_early { + *out.by_kind.entry(kind).or_default() += 1; + if out.examples.len() < 12 { + out.examples.push(format!( + "{kind:?} early-exit query of {} blocks: production {early:?}, reference {expected_early:?}", + query.len() + )); + } + } + } + + fn lookups(&mut self, out: &mut Mismatches) { + for _ in 0..8 { + let Some(index) = self.pick_held() else { + return; + }; + let base = self.held[index].contents.clone(); + if base.is_empty() { + continue; + } + self.check_lookup(QueryKind::Exact, base.clone(), out); + let prefix_len = self.rng.range(1, base.len()); + self.check_lookup(QueryKind::Prefix, base[..prefix_len].to_vec(), out); + if base.len() >= 2 { + let mut middle = base.clone(); + let at = self.rng.range(1, base.len() - 1); + let stream = self.fresh_stream(); + middle[at] = content(stream, 0); + self.check_lookup(QueryKind::MiddleReplaced, middle, out); + let mut suffix = base[..self.rng.range(1, base.len() - 1)].to_vec(); + let stream = self.fresh_stream(); + let extra = self.rng.range(1, 16); + suffix.extend((0..extra).map(|p| content(stream, p))); + self.check_lookup(QueryKind::SuffixReplaced, suffix, out); + } + let mut extended = base.clone(); + let stream = self.fresh_stream(); + let extra = self.rng.range(1, 40); + extended.extend((0..extra).map(|p| content(stream, p))); + self.check_lookup(QueryKind::Extended, extended, out); + } + let stream = self.fresh_stream(); + let len = self.rng.range(1, 32); + let unknown: Vec = (0..len).map(|p| content(stream, p)).collect(); + self.check_lookup(QueryKind::Unknown, unknown, out); + } + + /// The per-worker counters must agree with the maps the lanes keep. + fn check_counters(&self, label: &str) { + for (&worker, map) in &self.maps { + assert_eq!( + self.production.worker_block_count(worker), + map.len(), + "{label}: block count of worker {worker} after {} events", + self.events + ); + } + let total: usize = self.maps.values().map(|map| map.len()).sum(); + assert_eq!( + self.production.current_size(), + total, + "{label}: total blocks" + ); + // Distinct blocks are counted per shard: content held on two shards is stored twice. + assert_eq!( + self.production.entry_count(), + self.production + .debug_blocks() + .iter() + .map(|(worker, position, content, prefix)| { + ( + ShardedChainIndex::shard_of(*worker), + *position, + *content, + *prefix, + ) + }) + .collect::>() + .len(), + "{label}: distinct blocks" + ); + } +} + +fn env_or(name: &str, default: u64) -> u64 { + std::env::var(name) + .ok() + .and_then(|v| v.parse().ok()) + .unwrap_or(default) +} + +fn run_corpus(seed: u64, workers: usize, events: usize) -> (Harness, Mismatches) { + let mut harness = Harness::new(seed, workers); + let mut mismatches = Mismatches::default(); + while harness.events < events { + harness.step(); + if harness.events.is_multiple_of(256) { + harness.lookups(&mut mismatches); + } + } + harness.lookups(&mut mismatches); + (harness, mismatches) +} + +fn assert_exact(harness: &Harness, mismatches: &Mismatches, label: &str) { + let produced = harness.production.debug_blocks(); + let expected = harness.reference.blocks(); + let missing: Vec<_> = expected.difference(&produced).take(5).collect(); + let phantom: Vec<_> = produced.difference(&expected).take(5).collect(); + assert!( + missing.is_empty() && phantom.is_empty(), + "{label}: index content differs from the reference after {} events / {} stored blocks: \ + {} reference blocks, {} production blocks; missing e.g. {missing:?}; phantom e.g. {phantom:?}", + harness.events, + harness.stored_blocks, + expected.len(), + produced.len() + ); + assert!( + mismatches.by_kind.is_empty(), + "{label}: {} of {} lookups scored differently from the reference, by query kind {:?}; examples:\n{}", + mismatches.by_kind.values().sum::(), + mismatches.lookups, + mismatches.by_kind, + mismatches.examples.join("\n") + ); + harness.check_counters(label); +} + +#[test] +fn replayed_corpus_with_holes_matches_the_reference() { + let events = env_or("KV_INDEX_EXACTNESS_EVENTS", 20_000) as usize; + let seed = env_or("KV_INDEX_EXACTNESS_SEED", 20261005); + for (salt, workers) in [(8u64, 16usize), (64, 2), (3, 64)] { + let (harness, mismatches) = run_corpus(seed ^ salt, workers, events); + assert_exact( + &harness, + &mismatches, + &format!("{workers} workers, seed {seed}"), + ); + assert!( + harness.stored_blocks > events, + "corpus too small to mean anything: {} blocks for {} events", + harness.stored_blocks, + events + ); + assert!( + harness.holes * 20 > events, + "corpus made too few holes: {} for {} events", + harness.holes, + events + ); + } +} + +/// A worker holds [A, B, C]; a request [A, X, C] shares only A with it. +#[test] +fn a_divergence_at_the_tail_is_not_hidden() { + let production = ChainIndex::new(); + let mut reference = ReferenceIndexer::new(); + let worker = production.intern_worker("http://w:8000").unwrap(); + let mut map = ChainBlockMap::default(); + let held: Vec = (0..3).map(|p| content(1, p)).collect(); + let blocks = blocks_of(&held); + production + .apply_stored(worker, &blocks, None, &mut map) + .unwrap(); + reference.apply_stored(worker, &blocks, None).unwrap(); + let query = vec![held[0], content(2, 0), held[2]]; + let produced: BTreeMap = production + .find_matches(&query, false) + .scores + .into_iter() + .collect(); + assert_eq!(reference.find_matches(&query).get(&worker), Some(&1)); + assert_eq!(produced.get(&worker), Some(&1)); +} + +/// w1 holds the whole chain, w2 only its first 6 blocks, and w3 everything but block 0. +#[test] +fn partial_holders_score_their_own_prefix() { + let production = ChainIndex::new(); + let mut reference = ReferenceIndexer::new(); + let w1 = production.intern_worker("http://w1:8000").unwrap(); + let w2 = production.intern_worker("http://w2:8000").unwrap(); + let w3 = production.intern_worker("http://w3:8000").unwrap(); + let held: Vec = (0..20).map(|p| content(5, p)).collect(); + let blocks = blocks_of(&held); + let (mut m1, mut m2, mut m3) = ( + ChainBlockMap::default(), + ChainBlockMap::default(), + ChainBlockMap::default(), + ); + production.apply_stored(w1, &blocks, None, &mut m1).unwrap(); + reference.apply_stored(w1, &blocks, None).unwrap(); + production + .apply_stored(w2, &blocks[..6], None, &mut m2) + .unwrap(); + reference.apply_stored(w2, &blocks[..6], None).unwrap(); + production.apply_stored(w3, &blocks, None, &mut m3).unwrap(); + reference.apply_stored(w3, &blocks, None).unwrap(); + production.apply_removed(w3, &[blocks[0].seq_hash], &mut m3); + reference.apply_removed(w3, &[blocks[0].seq_hash]); + let expected = reference.find_matches(&held); + assert_eq!(expected.get(&w1), Some(&20)); + assert_eq!(expected.get(&w2), Some(&6)); + assert_eq!(expected.get(&w3), None); + let produced: BTreeMap = production + .find_matches(&held, false) + .scores + .into_iter() + .collect(); + assert_eq!(produced, expected); + assert_eq!(production.debug_blocks(), reference.blocks()); +} + +/// Two holes in one chain, healed one at a time, with a sibling branching off between them: the +/// score follows the first remaining hole until both are healed. +#[test] +fn holes_heal_independently() { + let production = ChainIndex::new(); + let mut reference = ReferenceIndexer::new(); + let worker = production.intern_worker("http://w:8000").unwrap(); + let mut map = ChainBlockMap::default(); + let held: Vec = (0..24).map(|p| content(7, p)).collect(); + let blocks = blocks_of(&held); + production + .apply_stored(worker, &blocks, None, &mut map) + .unwrap(); + reference.apply_stored(worker, &blocks, None).unwrap(); + let holes = [blocks[5].seq_hash, blocks[6].seq_hash, blocks[17].seq_hash]; + production.apply_removed(worker, &holes, &mut map); + reference.apply_removed(worker, &holes); + let mut sibling: Vec = held[..10].to_vec(); + sibling.extend((0..4).map(|p| content(8, p))); + let sibling_blocks = blocks_of(&sibling); + production + .apply_stored( + worker, + &sibling_blocks[10..], + Some(blocks[9].seq_hash), + &mut map, + ) + .unwrap(); + reference + .apply_stored(worker, &sibling_blocks[10..], Some(blocks[9].seq_hash)) + .unwrap(); + let agree = |production: &ChainIndex, reference: &ReferenceIndexer, query: &[ContentHash]| { + let produced: BTreeMap = production + .find_matches(query, false) + .scores + .into_iter() + .collect(); + assert_eq!(produced, reference.find_matches(query), "query {query:?}"); + }; + agree(&production, &reference, &held); + agree(&production, &reference, &sibling); + assert_eq!(reference.find_matches(&held).get(&worker), Some(&5)); + production + .apply_stored(worker, &blocks[5..7], Some(blocks[4].seq_hash), &mut map) + .unwrap(); + reference + .apply_stored(worker, &blocks[5..7], Some(blocks[4].seq_hash)) + .unwrap(); + agree(&production, &reference, &held); + agree(&production, &reference, &sibling); + assert_eq!(reference.find_matches(&held).get(&worker), Some(&17)); + assert_eq!(reference.find_matches(&sibling).get(&worker), Some(&14)); + production + .apply_stored(worker, &blocks[17..18], Some(blocks[16].seq_hash), &mut map) + .unwrap(); + reference + .apply_stored(worker, &blocks[17..18], Some(blocks[16].seq_hash)) + .unwrap(); + agree(&production, &reference, &held); + assert_eq!(reference.find_matches(&held).get(&worker), Some(&24)); + assert_eq!(production.debug_blocks(), reference.blocks()); +} + +/// The production index against the reference after every few events of a wild corpus: stores +/// under parents the worker no longer holds (and the gateway's parentless fallback after the +/// rejection), whole chains announced again from the root, repeated stores of held blocks, +/// tails, holes, clears and worker churn, from several seeds. The full state is compared every +/// 16 events, the lookups every 64, the counters at the end. Scale with +/// `KV_INDEX_WILD_EVENTS` (default 3000 per seed) and `KV_INDEX_WILD_SEEDS` (default 6). +#[test] +fn wild_event_sequences_match_the_reference_at_every_step() { + let events = env_or("KV_INDEX_WILD_EVENTS", 3000) as usize; + let seeds = env_or("KV_INDEX_WILD_SEEDS", 6); + let base = env_or("KV_INDEX_EXACTNESS_SEED", 20261006); + for seed in 0..seeds { + let mut harness = Harness::new(base.wrapping_add(seed.wrapping_mul(0x9e37)), 6); + let mut mismatches = Mismatches::default(); + let mut twin_names = Vec::new(); + let label = format!("wild seed {seed}"); + while harness.events < events { + let before = harness.events; + harness.step_wild(false, &mut twin_names); + if harness.events == before { + continue; + } + if harness.events.is_multiple_of(16) { + let produced = harness.production.debug_blocks(); + let expected = harness.reference.blocks(); + assert!( + produced == expected, + "{label}: state differs after {} events: {} reference blocks, {} production blocks; missing e.g. {:?}; phantom e.g. {:?}", + harness.events, + expected.len(), + produced.len(), + expected.difference(&produced).take(4).collect::>(), + produced.difference(&expected).take(4).collect::>() + ); + } + if harness.events.is_multiple_of(64) { + harness.lookups(&mut mismatches); + } + } + harness.lookups(&mut mismatches); + assert_exact(&harness, &mismatches, &label); + assert_eq!( + harness.production.stats().engine_conflicts, + 0, + "{label}: no twins in this corpus" + ); + } +} + +/// The same corpus with twins (one position named by two engine hashes): the index files a +/// twin onto the position it holds and counts the conflict, so it holds at most what the +/// reference holds and scores no worker higher than the reference does; removals by either +/// name, stores under either name and the cut-and-join they lead to must never panic. +#[test] +fn twins_never_panic_and_never_overstate_the_reference() { + let events = env_or("KV_INDEX_WILD_EVENTS", 3000) as usize; + let seeds = env_or("KV_INDEX_WILD_SEEDS", 6); + let base = env_or("KV_INDEX_EXACTNESS_SEED", 20261006); + for seed in 0..seeds { + let mut harness = Harness::new(base.wrapping_add(seed.wrapping_mul(0x51f1)), 6); + let mut twin_names = Vec::new(); + let label = format!("twins seed {seed}"); + while harness.events < events { + let before = harness.events; + harness.step_wild(true, &mut twin_names); + if harness.events == before { + continue; + } + if harness.events.is_multiple_of(16) { + let produced = harness.production.debug_blocks(); + let expected = harness.reference.blocks(); + let phantom: Vec<_> = produced.difference(&expected).take(4).collect(); + assert!( + phantom.is_empty(), + "{label}: production holds blocks the reference does not after {} events: {phantom:?}", + harness.events + ); + for (&worker, map) in &harness.maps { + assert!( + harness.production.worker_block_count(worker) <= map.len(), + "{label}: worker {worker} credited beyond its lane map after {} events", + harness.events + ); + } + } + if harness.events.is_multiple_of(64) { + for _ in 0..4 { + let Some(index) = harness.pick_held() else { + break; + }; + let query = harness.held[index].contents.clone(); + let produced = harness.production_scores(&query); + let expected = harness.reference.find_matches(&query); + for (worker, score) in &produced { + assert!( + expected.get(worker).is_some_and(|e| e >= score), + "{label}: worker {worker} scores {score} in production against {:?} in the reference after {} events", + expected.get(worker), + harness.events + ); + } + } + } + } + assert!( + harness.production.stats().engine_conflicts > 0, + "{label}: the corpus produced no twin" + ); + } +} diff --git a/crates/kv_index/tests/exactness_positional.rs b/crates/kv_index/tests/exactness_positional.rs new file mode 100644 index 0000000000..04667d5975 --- /dev/null +++ b/crates/kv_index/tests/exactness_positional.rs @@ -0,0 +1,731 @@ +//! Exactness harness: the production `PositionalIndexer` against the `ReferenceIndexer`. +//! +//! A seeded corpus of stores, extensions, divergent siblings, tail and whole-chain removals, +//! clears and worker removals is replayed into both indexers through the same per-worker +//! `WorkerBlockMap` handling the gateway's `KvEventMonitor` uses. After every round of events, +//! lookups built from live chains and from mutated chains must score identically in both; at the +//! end, the production index must hold exactly the reference's blocks. +//! +//! Two corpora run. One evicts tails and whole chains only, which is what a radix-tree cache such +//! as SGLang's produces (it evicts leaves). The other also evicts single blocks from the middle of +//! chains the worker keeps, which vLLM produces: a request that re-hits a shared prefix takes those +//! blocks out of the free queue and, when it finishes, re-queues them behind the unshared tail of +//! an earlier request, so that request's middle blocks fall out of the LRU before its later blocks +//! do. The engine's prefix match stops at the hole, so a worker's score must end there too. Scale +//! with `KV_INDEX_EXACTNESS_EVENTS` (default 20000) and `KV_INDEX_EXACTNESS_SEED`. +#![expect(clippy::expect_used)] + +use std::collections::{BTreeMap, BTreeSet}; + +use kv_index::{ContentHash, PositionalIndexer, ReferenceIndexer, SequenceHash, WorkerBlockMap}; +use rustc_hash::FxHashMap; + +mod common; +use common::{blocks_of, content, Rng}; + +struct Held { + worker: u32, + contents: Vec, +} + +/// A chain as it was when one of its middle blocks was evicted; its blocks after `hole` stay in +/// the index but are unreachable through the chain, so lookups along it must stop at `hole`. +struct Holed { + worker: u32, + contents: Vec, + hole: usize, +} + +struct Harness { + production: PositionalIndexer, + reference: ReferenceIndexer, + maps: FxHashMap, + workers: Vec, + held: Vec, + /// Whether the corpus evicts middle blocks (see the module doc). + holes: bool, + holed: Vec, + prompts: Vec>, + next_stream: u64, + next_worker: u32, + rng: Rng, + events: usize, + stored_blocks: usize, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)] +enum QueryKind { + Exact, + Prefix, + MiddleReplaced, + SuffixReplaced, + Extended, + Unknown, + /// A full chain through a block the worker evicted from its middle. + AcrossHole, + /// A prefix of such a chain that ends after the hole. + PastHole, +} + +#[derive(Default)] +struct Mismatches { + by_kind: BTreeMap, + issued: BTreeMap, + examples: Vec, + lookups: usize, +} + +impl Harness { + fn new(seed: u64, jump_size: usize, workers: usize, holes: bool) -> Self { + let mut rng = Rng::new(seed); + let production = PositionalIndexer::new(jump_size); + let mut harness = Self { + production, + reference: ReferenceIndexer::new(), + maps: FxHashMap::default(), + workers: Vec::new(), + held: Vec::new(), + holes, + holed: Vec::new(), + prompts: Vec::new(), + next_stream: 1, + next_worker: 0, + rng: Rng::new(seed ^ 0x9e37_79b9_7f4a_7c15), + events: 0, + stored_blocks: 0, + }; + for _ in 0..workers { + harness.add_worker(); + } + let prompt_count = rng.range(3, 8); + for _ in 0..prompt_count { + let len = rng.range(8, 48); + let stream = harness.fresh_stream(); + harness + .prompts + .push((0..len).map(|p| content(stream, p)).collect()); + } + harness + } + + fn fresh_stream(&mut self) -> u64 { + self.next_stream += 1; + self.next_stream + } + + fn add_worker(&mut self) -> u32 { + let url = format!("http://worker-{}:8000", self.next_worker); + self.next_worker += 1; + let id = self.production.intern_worker(&url).expect("worker id"); + self.maps.insert(id, WorkerBlockMap::default()); + self.workers.push(id); + id + } + + fn random_worker(&mut self) -> u32 { + self.workers[self.rng.below(self.workers.len())] + } + + /// A new user turn: a prompt prefix plus fresh blocks. + fn new_chain(&mut self) -> Vec { + let prompt = &self.prompts[self.rng.below(self.prompts.len())]; + let mut contents = prompt.clone(); + let turn = self.rng.range(4, 32); + let stream = self.fresh_stream(); + contents.extend((0..turn).map(|p| content(stream, p))); + contents + } + + /// Store `contents` on `worker` as the engine would: only the suffix the worker does not hold, + /// after the last block it does hold. A fully held chain is re-stored for its last block, which + /// exercises the duplicate-store path. + fn store(&mut self, worker: u32, contents: Vec) { + let blocks = blocks_of(&contents); + let held = self.maps.get(&worker).expect("worker map"); + let mut known = 0; + while known < blocks.len() && held.contains_key(&blocks[known].seq_hash) { + known += 1; + } + let start = if known == blocks.len() { + blocks.len() - 1 + } else { + known + }; + let parent = if start == 0 { + None + } else { + Some(blocks[start - 1].seq_hash) + }; + let map = self.maps.get_mut(&worker).expect("worker map"); + let produced = self + .production + .apply_stored(worker, &blocks[start..], parent, map); + let referenced = self + .reference + .apply_stored(worker, &blocks[start..], parent); + assert_eq!( + produced.is_ok(), + referenced.is_ok(), + "store outcome differs: production {produced:?}, reference {referenced:?}" + ); + if produced.is_ok() { + self.stored_blocks += blocks.len() - start; + self.held.push(Held { worker, contents }); + } + self.events += 1; + } + + fn pick_held(&mut self) -> Option { + if self.held.is_empty() { + return None; + } + let index = self.rng.below(self.held.len()); + if self.maps.contains_key(&self.held[index].worker) { + Some(index) + } else { + self.held.swap_remove(index); + None + } + } + + /// Evict a tail: the chain's blocks from a random position on, together with the same + /// positions of every other held chain of that worker that runs through them. That keeps every + /// chain the worker holds free of holes, which is what engines with prefix caching produce + /// (a block is only reusable through its predecessors) and what the production jump search + /// relies on when it skips from one landing to the next. + fn remove_tail(&mut self) { + let Some(index) = self.pick_held() else { + return; + }; + let worker = self.held[index].worker; + let contents = self.held[index].contents.clone(); + let keep = self.rng.below(contents.len()); + let mut hashes: Vec = Vec::new(); + for held in self.held.iter_mut().filter(|h| h.worker == worker) { + let shared = held + .contents + .iter() + .zip(&contents) + .take_while(|(a, b)| a == b) + .count(); + if shared > keep { + let blocks = blocks_of(&held.contents); + hashes.extend(blocks[keep..].iter().map(|b| b.seq_hash)); + held.contents.truncate(keep); + } + } + hashes.sort_unstable_by_key(|h| h.0); + hashes.dedup(); + let map = self.maps.get_mut(&worker).expect("worker map"); + self.production.apply_removed(worker, &hashes, map); + self.reference.apply_removed(worker, &hashes); + self.held.retain(|h| !h.contents.is_empty()); + self.events += 1; + } + + /// Evict a whole conversation: the chain's blocks beyond the longest prefix it shares with + /// another held chain of the same worker (the shared prefix stays, as it would in an engine + /// where the sibling still references it), so no chain of the worker is left with a hole. + fn remove_chain(&mut self) { + let Some(index) = self.pick_held() else { + return; + }; + let worker = self.held[index].worker; + let contents = self.held[index].contents.clone(); + let shared = self + .held + .iter() + .enumerate() + .filter(|(i, h)| *i != index && h.worker == worker) + .map(|(_, h)| { + h.contents + .iter() + .zip(&contents) + .take_while(|(a, b)| a == b) + .count() + }) + .max() + .unwrap_or(0); + let blocks = blocks_of(&contents); + let hashes: Vec = blocks[shared..].iter().map(|b| b.seq_hash).collect(); + let map = self.maps.get_mut(&worker).expect("worker map"); + self.production.apply_removed(worker, &hashes, map); + self.reference.apply_removed(worker, &hashes); + self.held.swap_remove(index); + self.events += 1; + } + + /// Evict one block from the middle of a chain while the blocks after it stay (the vLLM case in + /// the module doc). Every held chain of the worker that runs through the block loses it; the + /// full chains are kept aside so lookups can run across the hole. + fn remove_middle(&mut self) { + let Some(index) = self.pick_held() else { + return; + }; + if self.held[index].contents.len() < 3 { + return; + } + let worker = self.held[index].worker; + let contents = self.held[index].contents.clone(); + let hole = self.rng.range(1, contents.len() - 2); + let evicted = blocks_of(&contents)[hole].seq_hash; + for held in self.held.iter_mut().filter(|h| h.worker == worker) { + let shared = held + .contents + .iter() + .zip(&contents) + .take_while(|(a, b)| a == b) + .count(); + if shared > hole { + self.holed.push(Holed { + worker, + contents: held.contents.clone(), + hole, + }); + held.contents.truncate(hole); + } + } + let map = self.maps.get_mut(&worker).expect("worker map"); + self.production.apply_removed(worker, &[evicted], map); + self.reference.apply_removed(worker, &[evicted]); + self.events += 1; + } + + fn pick_holed(&mut self) -> Option { + if self.holed.is_empty() { + return None; + } + let index = self.rng.below(self.holed.len()); + if self.maps.contains_key(&self.holed[index].worker) { + Some(index) + } else { + self.holed.swap_remove(index); + None + } + } + + fn clear_worker(&mut self) { + let worker = self.random_worker(); + let map = self.maps.get_mut(&worker).expect("worker map"); + self.production.apply_cleared(worker, map); + self.reference.apply_cleared(worker); + self.held.retain(|h| h.worker != worker); + self.events += 1; + } + + fn remove_worker(&mut self) { + if self.workers.len() < 2 { + return; + } + let position = self.rng.below(self.workers.len()); + let worker = self.workers.swap_remove(position); + let map = self.maps.remove(&worker).expect("worker map"); + self.production.remove_worker(worker, map); + self.reference.remove_worker(worker); + self.held.retain(|h| h.worker != worker); + self.add_worker(); + self.events += 1; + } + + fn step(&mut self) { + let roll = self.rng.below(1000); + match roll { + 0..=399 => { + let chain = self.new_chain(); + let worker = self.random_worker(); + self.store(worker, chain); + } + 400..=649 => { + // Extend a held chain by another turn on the same worker. + let Some(index) = self.pick_held() else { + return; + }; + let worker = self.held[index].worker; + let mut contents = self.held[index].contents.clone(); + let turn = self.rng.range(4, 32); + let stream = self.fresh_stream(); + contents.extend((0..turn).map(|p| content(stream, p))); + self.store(worker, contents); + } + 650..=799 => { + // A sibling that diverges at a random position, including 1 and the last block, + // stored on a random worker (often another one, which shares the prefix blocks). + let Some(index) = self.pick_held() else { + return; + }; + let base = &self.held[index].contents; + if base.len() < 2 { + return; + } + let divergence = match self.rng.below(10) { + 0 => 1, + 1 => base.len() - 1, + 2 => base.len(), + _ => self.rng.range(1, base.len()), + }; + let mut contents: Vec = base[..divergence].to_vec(); + let turn = self.rng.range(1, 24); + let stream = self.fresh_stream(); + contents.extend((0..turn).map(|p| content(stream, p))); + let worker = self.random_worker(); + self.store(worker, contents); + } + 800..=899 => { + if self.holes && self.rng.chance(1, 2) { + self.remove_middle(); + } else { + self.remove_tail(); + } + } + 900..=949 => self.remove_chain(), + 950..=984 => { + if self.workers.len() < 64 && self.rng.chance(1, 4) { + self.add_worker(); + } else { + let Some(index) = self.pick_held() else { + return; + }; + let contents = self.held[index].contents.clone(); + let worker = self.random_worker(); + self.store(worker, contents); + } + } + 985..=994 => self.clear_worker(), + _ => self.remove_worker(), + } + } + + fn production_scores(&self, query: &[ContentHash]) -> BTreeMap { + self.production + .find_matches(query, false) + .scores + .into_iter() + .collect() + } + + fn check_lookup(&mut self, kind: QueryKind, query: Vec, out: &mut Mismatches) { + if query.is_empty() { + return; + } + out.lookups += 1; + *out.issued.entry(kind).or_default() += 1; + let produced = self.production_scores(&query); + let expected = self.reference.find_matches(&query); + if produced != expected { + *out.by_kind.entry(kind).or_default() += 1; + if out.examples.len() < 12 { + let diff: Vec = expected + .keys() + .chain(produced.keys()) + .collect::>() + .into_iter() + .filter(|w| produced.get(w) != expected.get(w)) + .map(|w| { + format!( + "worker {w}: production {:?}, reference {:?}", + produced.get(w), + expected.get(w) + ) + }) + .collect(); + out.examples.push(format!( + "{kind:?} query of {} blocks after {} events: {}", + query.len(), + self.events, + diff.join("; ") + )); + } + } + // early_exit reports exactly the workers that match position 0, each with score 1. + let early: BTreeMap = self + .production + .find_matches(&query, true) + .scores + .into_iter() + .collect(); + let expected_early: BTreeMap = expected.keys().map(|&w| (w, 1)).collect(); + if early != expected_early { + *out.by_kind.entry(kind).or_default() += 1; + if out.examples.len() < 12 { + out.examples.push(format!( + "{kind:?} early-exit query of {} blocks: production {early:?}, reference {expected_early:?}", + query.len() + )); + } + } + } + + fn lookups(&mut self, out: &mut Mismatches) { + for _ in 0..8 { + let Some(index) = self.pick_held() else { + return; + }; + let base = self.held[index].contents.clone(); + if base.is_empty() { + continue; + } + self.check_lookup(QueryKind::Exact, base.clone(), out); + let prefix_len = self.rng.range(1, base.len()); + self.check_lookup(QueryKind::Prefix, base[..prefix_len].to_vec(), out); + if base.len() >= 2 { + let mut middle = base.clone(); + let at = self.rng.range(1, base.len() - 1); + let stream = self.fresh_stream(); + middle[at] = content(stream, 0); + self.check_lookup(QueryKind::MiddleReplaced, middle, out); + let mut suffix = base[..self.rng.range(1, base.len() - 1)].to_vec(); + let stream = self.fresh_stream(); + let extra = self.rng.range(1, 16); + suffix.extend((0..extra).map(|p| content(stream, p))); + self.check_lookup(QueryKind::SuffixReplaced, suffix, out); + } + let mut extended = base.clone(); + let stream = self.fresh_stream(); + let extra = self.rng.range(1, 40); + extended.extend((0..extra).map(|p| content(stream, p))); + self.check_lookup(QueryKind::Extended, extended, out); + } + if self.holes { + for _ in 0..4 { + let Some(index) = self.pick_holed() else { + break; + }; + let (full, hole) = { + let holed = &self.holed[index]; + (holed.contents.clone(), holed.hole) + }; + let past = self.rng.range(hole + 1, full.len()); + self.check_lookup(QueryKind::PastHole, full[..past].to_vec(), out); + self.check_lookup(QueryKind::AcrossHole, full, out); + } + } + let stream = self.fresh_stream(); + let len = self.rng.range(1, 32); + let unknown: Vec = (0..len).map(|p| content(stream, p)).collect(); + self.check_lookup(QueryKind::Unknown, unknown, out); + } + + fn production_blocks(&self) -> BTreeSet<(u32, usize, ContentHash, SequenceHash)> { + self.production.debug_blocks().into_iter().collect() + } +} + +fn env_or(name: &str, default: u64) -> u64 { + std::env::var(name) + .ok() + .and_then(|v| v.parse().ok()) + .unwrap_or(default) +} + +fn run_corpus( + seed: u64, + jump_size: usize, + workers: usize, + events: usize, + holes: bool, +) -> (Harness, Mismatches) { + let mut harness = Harness::new(seed, jump_size, workers, holes); + let mut mismatches = Mismatches::default(); + while harness.events < events { + harness.step(); + if harness.events.is_multiple_of(256) { + harness.lookups(&mut mismatches); + } + } + harness.lookups(&mut mismatches); + (harness, mismatches) +} + +fn assert_exact(harness: &Harness, mismatches: &Mismatches, label: &str) { + let produced = harness.production_blocks(); + let expected = harness.reference.blocks(); + let missing: Vec<_> = expected.difference(&produced).take(5).collect(); + let phantom: Vec<_> = produced.difference(&expected).take(5).collect(); + assert!( + missing.is_empty() && phantom.is_empty(), + "{label}: index content differs from the reference after {} events / {} stored blocks: \ + {} reference blocks, {} production blocks; missing e.g. {missing:?}; phantom e.g. {phantom:?}", + harness.events, + harness.stored_blocks, + expected.len(), + produced.len() + ); + assert!( + mismatches.by_kind.is_empty(), + "{label}: {} of {} lookups scored differently from the reference, by query kind {:?} \ + (issued {:?}); examples:\n{}", + mismatches.by_kind.values().sum::(), + mismatches.lookups, + mismatches.by_kind, + mismatches.issued, + mismatches.examples.join("\n") + ); +} + +#[test] +fn replayed_corpus_matches_the_reference() { + let events = env_or("KV_INDEX_EXACTNESS_EVENTS", 20_000) as usize; + let seed = env_or("KV_INDEX_EXACTNESS_SEED", 20261005); + for (jump, workers) in [(8usize, 16usize), (64, 2), (3, 64)] { + let (harness, mismatches) = run_corpus(seed ^ jump as u64, jump, workers, events, false); + assert_exact( + &harness, + &mismatches, + &format!("jump {jump}, {workers} workers, seed {seed}"), + ); + assert!( + harness.stored_blocks > events, + "corpus too small to mean anything: {} blocks for {} events", + harness.stored_blocks, + events + ); + } +} + +/// The vLLM corpus: middle blocks get evicted while later blocks stay, and lookups run across and +/// past the holes. Before every position was verified, the jump search landed past a hole on an +/// entry that still named the worker and scored the whole chain. +#[test] +fn replayed_corpus_with_holes_matches_the_reference() { + let events = env_or("KV_INDEX_EXACTNESS_EVENTS", 20_000) as usize; + let seed = env_or("KV_INDEX_EXACTNESS_SEED", 20261005); + for (jump, workers) in [(8usize, 16usize), (64, 2), (3, 64)] { + let (harness, mismatches) = run_corpus(seed ^ jump as u64, jump, workers, events, true); + assert_exact( + &harness, + &mismatches, + &format!("holes, jump {jump}, {workers} workers, seed {seed}"), + ); + let across = mismatches + .issued + .get(&QueryKind::AcrossHole) + .copied() + .unwrap_or(0); + assert!( + across >= 64, + "hole corpus too small to mean anything: {across} lookups across holes" + ); + } +} + +/// A worker holds [A ..= J] (ten blocks) and evicts E while F ..= J stay, as vLLM's free queue +/// can order it. A request for the full chain hits A ..= D in the engine and recomputes the rest, +/// so the score is 4; the entries at F ..= J still name the worker and must not count. +#[test] +fn evicted_middle_block_ends_the_match() { + let index = PositionalIndexer::new(8); + let worker = index + .intern_worker("http://worker-0:8000") + .expect("worker id"); + let mut map = WorkerBlockMap::default(); + let contents: Vec = (0..10).map(|p| content(7, p)).collect(); + let blocks = blocks_of(&contents); + index + .apply_stored(worker, &blocks, None, &mut map) + .expect("store"); + index.apply_removed(worker, &[blocks[4].seq_hash], &mut map); + let scores = index.find_matches(&contents, false).scores; + assert_eq!(scores.get(&worker).copied(), Some(4), "scores {scores:?}"); + let past = index.find_matches(&contents[..7], false).scores; + assert_eq!(past.get(&worker).copied(), Some(4), "scores {past:?}"); + let before = index.find_matches(&contents[..4], false).scores; + assert_eq!(before.get(&worker).copied(), Some(4), "scores {before:?}"); +} + +/// A worker holds [A, B, C]; a request [A, X, C] shares only A with it. The jump search lands on +/// position 2, where the single stored entry for C must not be taken as a match without its prefix +/// hash: the request's chain differs from position 1 on. +#[test] +fn single_entry_shortcut_must_not_hide_a_divergence_at_the_tail() { + let production = PositionalIndexer::new(8); + let mut reference = ReferenceIndexer::new(); + let worker = production.intern_worker("http://w:8000").unwrap(); + let mut map = WorkerBlockMap::default(); + let held: Vec = (0..3).map(|p| content(1, p)).collect(); + let blocks = blocks_of(&held); + production + .apply_stored(worker, &blocks, None, &mut map) + .unwrap(); + reference.apply_stored(worker, &blocks, None).unwrap(); + let query = vec![held[0], content(2, 0), held[2]]; + let produced: BTreeMap = production + .find_matches(&query, false) + .scores + .into_iter() + .collect(); + assert_eq!(reference.find_matches(&query).get(&worker), Some(&1)); + assert_eq!( + produced.get(&worker), + Some(&1), + "scored {produced:?}, expected 1 (A only)" + ); +} + +/// A worker holds a 20-block chain; the request replaces block 10. The jump landings at 16 and 19 +/// find single entries whose content matches but whose prefix hash belongs to the stored chain, not +/// to the request's chain, which diverged at 10. +#[test] +fn single_entry_landing_must_not_hide_an_earlier_divergence() { + let production = PositionalIndexer::new(8); + let mut reference = ReferenceIndexer::new(); + let worker = production.intern_worker("http://w:8000").unwrap(); + let mut map = WorkerBlockMap::default(); + let held: Vec = (0..20).map(|p| content(3, p)).collect(); + let blocks = blocks_of(&held); + production + .apply_stored(worker, &blocks, None, &mut map) + .unwrap(); + reference.apply_stored(worker, &blocks, None).unwrap(); + let mut query = held.clone(); + query[10] = content(4, 0); + let produced: BTreeMap = production + .find_matches(&query, false) + .scores + .into_iter() + .collect(); + assert_eq!(reference.find_matches(&query).get(&worker), Some(&10)); + assert_eq!( + produced.get(&worker), + Some(&10), + "scored {produced:?}, expected 10" + ); +} + +/// Equal worker counts at a jump landing do not mean the same workers: w1 holds the whole chain, +/// w2 only its first 6 blocks, and w3 everything but block 0. At the landing (position 8) the +/// matching set is {w1, w3}, the same size as the active set {w1, w2}; w2 must still be scored 6. +#[test] +fn count_equality_at_a_landing_is_not_set_equality() { + let production = PositionalIndexer::new(8); + let mut reference = ReferenceIndexer::new(); + let w1 = production.intern_worker("http://w1:8000").unwrap(); + let w2 = production.intern_worker("http://w2:8000").unwrap(); + let w3 = production.intern_worker("http://w3:8000").unwrap(); + let held: Vec = (0..20).map(|p| content(5, p)).collect(); + let blocks = blocks_of(&held); + let (mut m1, mut m2, mut m3) = ( + WorkerBlockMap::default(), + WorkerBlockMap::default(), + WorkerBlockMap::default(), + ); + production.apply_stored(w1, &blocks, None, &mut m1).unwrap(); + reference.apply_stored(w1, &blocks, None).unwrap(); + production + .apply_stored(w2, &blocks[..6], None, &mut m2) + .unwrap(); + reference.apply_stored(w2, &blocks[..6], None).unwrap(); + production.apply_stored(w3, &blocks, None, &mut m3).unwrap(); + reference.apply_stored(w3, &blocks, None).unwrap(); + production.apply_removed(w3, &[blocks[0].seq_hash], &mut m3); + reference.apply_removed(w3, &[blocks[0].seq_hash]); + let expected = reference.find_matches(&held); + assert_eq!(expected.get(&w1), Some(&20)); + assert_eq!(expected.get(&w2), Some(&6)); + assert_eq!(expected.get(&w3), None); + let produced: BTreeMap = production + .find_matches(&held, false) + .scores + .into_iter() + .collect(); + assert_eq!(produced, expected); +} diff --git a/crates/kv_index/tests/split_counters.rs b/crates/kv_index/tests/split_counters.rs new file mode 100644 index 0000000000..a9764937cf --- /dev/null +++ b/crates/kv_index/tests/split_counters.rs @@ -0,0 +1,75 @@ +//! The split counters in the stats name the cause of every split: a removal leaving a hole, a +//! store entering a run under a parent the worker did not hold up to; a chain diverging inside a +//! run hangs the tail off the offset as a child and counts nothing. The churn harness reads them to attribute fragmentation; this keeps them honest. + +use kv_index::{ + compute_content_hash, request_prefix_hashes, ChainBlockMap, ContentHash, ShardedChainIndex, + StoredBlock, +}; + +fn chain(stream: u64, shared: &[ContentHash], len: usize) -> Vec { + let mut contents: Vec = shared.to_vec(); + for position in contents.len()..len { + contents.push(compute_content_hash(&[stream as u32, position as u32])); + } + contents + .iter() + .zip(request_prefix_hashes(&contents)) + .map(|(&content_hash, seq_hash)| StoredBlock { + seq_hash, + content_hash, + }) + .collect() +} + +#[test] +fn splits_are_counted_by_cause() { + let index = ShardedChainIndex::new(2, 8); + // Both on shard 0: a split needs the chains in one index (sharding is by worker). + let a = index.intern_worker_in(0, "a").unwrap(); + let b = index.intern_worker_in(0, "b").unwrap(); + let (mut ma, mut mb) = (ChainBlockMap::default(), ChainBlockMap::default()); + let chain_a = chain(1, &[], 100); + index.apply_stored(a, &chain_a, None, &mut ma).unwrap(); + let before = index.stats(); + assert_eq!( + ( + before.splits_by_branch, + before.splits_by_hole, + before.splits_by_mid_run_store + ), + (0, 0, 0) + ); + + // A chain sharing the first 50 blocks diverges inside a's run: the run stays whole and b's + // tail hangs off offset 50 as a child, so no split is counted and two runs are live. + let shared: Vec = chain_a[..50] + .iter() + .map(|block| block.content_hash) + .collect(); + let chain_b = chain(2, &shared, 100); + index.apply_stored(b, &chain_b, None, &mut mb).unwrap(); + let stats = index.stats(); + assert_eq!( + stats.splits_by_branch, 0, + "a divergence inside a run splits nothing" + ); + assert_eq!(stats.splits_by_hole, 0); + assert_eq!(stats.runs_live, 2); + + // a drops blocks 20..30 and keeps the rest: a hole, so the tail becomes its own run. + let hole: Vec<_> = chain_a[20..30].iter().map(|block| block.seq_hash).collect(); + index.apply_removed(a, &hole, &mut ma); + let stats = index.stats(); + assert_eq!(stats.splits_by_hole, 1, "a removal that leaves a hole"); + assert_eq!(stats.splits_by_branch, 0); + + // Deaths: when every holder of a run is gone, the run is unlinked and counted. + let died_before = index.stats().runs_died; + let tail: Vec<_> = chain_b[50..].iter().map(|block| block.seq_hash).collect(); + index.apply_removed(b, &tail, &mut mb); + assert!( + index.stats().runs_died > died_before, + "b's private tail run died with its only holder" + ); +} diff --git a/crates/mock_worker/Cargo.toml b/crates/mock_worker/Cargo.toml index 32027848a3..e1318cc412 100644 --- a/crates/mock_worker/Cargo.toml +++ b/crates/mock_worker/Cargo.toml @@ -3,7 +3,7 @@ name = "mock-worker" version = "0.1.0" edition = "2021" publish = false -description = "Multi-port mock HTTP/gRPC inference worker for SMG gateway scale testing" +description = "Multi-port mock HTTP/gRPC inference worker for SMG gateway scale testing, with the trace replayer that measures routing against it" [lib] name = "mock_worker" @@ -13,7 +13,16 @@ path = "src/lib.rs" name = "mock-worker" path = "src/main.rs" +# The trace replayer: a measurement tool belongs with the fleet it measures. +[[bin]] +name = "replay" +path = "src/bin/replay.rs" + [dependencies] +anyhow.workspace = true +clap = { version = "4", features = ["derive"] } +reqwest = { workspace = true, features = ["json", "stream", "rustls"] } +serde = { workspace = true, features = ["derive"] } smg-grpc-client.workspace = true engine-zmq-client = { workspace = true, features = ["mock-engine"] } tokio = { workspace = true, features = ["full"] } @@ -24,6 +33,10 @@ serde_json.workspace = true futures.workspace = true tracing.workspace = true tracing-subscriber.workspace = true +rmpv.workspace = true +zeromq.workspace = true [dev-dependencies] tempfile = "3" +engine-servicer = { path = "../engine_servicer" } +rmp-serde.workspace = true diff --git a/crates/mock_worker/README.md b/crates/mock_worker/README.md index def21937fc..87db6d95ce 100644 --- a/crates/mock_worker/README.md +++ b/crates/mock_worker/README.md @@ -28,6 +28,13 @@ ids): `HealthCheck`, `GetModelInfo`, `GetServerInfo`, `Generate` (streamed chunks + complete), `GetLoads`, `Abort`, and (realistic mode only) `SubscribeKvEvents`; other admin RPCs return `unimplemented`. +**ZMQ** (`--zmq-handshake --zmq-count N [--zmq-start-index I]`): +mock vLLM `EngineCore` ranks on the `engine-zmq-client` wire. Unlike the HTTP +and gRPC workers, which bind and are dialed, each rank dials the frontend's +handshake address, registers, takes `EngineCoreRequest`s and pushes +`EngineCoreOutputs` with the engine's load as `scheduler_stats`; the admin API +sees it as worker `zmq:`. + ## Run (canned) ```bash @@ -39,47 +46,187 @@ cargo run --release -p mock-worker -- \ Each worker is one port. Register them against an IGW gateway with `POST /workers` (`{"url":"http://127.0.0.1:9000"}`, or `grpc://…` with -`connection_mode`/`runtime`/`models` for gRPC). +`connection_mode`/`runtime`/`models` for gRPC). `mock-worker --help` lists +every flag with its default. ## Realistic engine `--engine realistic` backs each worker with a continuous-batching simulator -([`src/engine.rs`](src/engine.rs)) that reproduces the behaviors driving routing: - -- **prefill latency scales with input length** — TTFT grows with the *uncached* - prompt size, chunked across scheduler steps; -- **inter-token latency grows with batch size** — `ITL = base + slope · batch`, - so a busy replica is slower per token; -- **finite KV capacity + queueing** — when KV is full, requests wait, producing - the `num_waiting_uncached_tokens` signal `least_load` consumes; -- **prefix caching** — a request sharing a prefix with cached blocks pays less - prefill, reports `cached_tokens`, and the worker emits the KV-cache events - (`SubscribeKvEvents`) that drive event-driven `cache_aware` routing. +([`src/engine.rs`](src/engine.rs)) built the way vLLM schedules, so routing +experiments against it transfer to engines: -The cost model is parametric (defaults approximate one mid-size replica): +- **pass loop** — every pass has a token budget (`--max-batched-tokens`, 8192) + and a sequence cap (`--max-running`, 256). Running requests go first, each + taking one decode token or a prefill chunk; then the queue is admitted FCFS + while budget and KV room remain. With `--prefill-first true` (SGLang) a pass + that contains prefill runs prefill only. +- **block-level KV pool** — `--kv-blocks`/`--kv-tokens` physical blocks of + `--block-size` tokens, keyed by content hash with reference counts; idle + cached blocks sit in an LRU and are evicted head-first when an allocation + needs room; a waiting prompt is admitted only when the free and evictable + blocks could hold all of it beyond its cached blocks (vLLM's + `scheduler_reserve_full_isl`, a gate read at admission, head-of-line) while + only the pass's chunk is allocated, the next chunks allocating as they run, + and a running request that still cannot get a block for its + next tokens preempts the most recently admitted request (LIFO), which + recomputes later. A fully cached prompt recomputes its last block; a prefix + hit on an idle block references it again. +- **KV events** (`SubscribeKvEvents`) — per request within a pass, `Removed` + for the blocks evicted by its allocation, then one `Stored` per contiguous + run of blocks it completed (parent-chained, with token ids). `Removed` fires + only when the last copy of a hash leaves the pool, `Stored` only when a hash + first appears; completion and preemption emit nothing. Every event of a pass + becomes visible at the pass end, so a request arriving mid-pass cannot see + that pass's blocks. A reset (`POST /admin/reset`) publishes + `AllBlocksCleared`. +- **timing** — `--timing polynomial` (default) uses AISimulate's baseline: + prefill `16.50142 + 1.518344e-2·T + 4.209989e-7·T²` ms over the uncached + tokens `T` of the pass, decode `max(1, 5.74 + 54.01·u − 25.74·u²)` ms over + the KV utilisation `u` of the decoding requests; a pass lasts prefill plus + decode, and the first token of a prefill adds no decode time. + `--timing linear` is the simpler model: prefill at `--prefill-tps` tokens/s, + decode `base + per-request · batch` ms. +- **cached tokens** — a request sharing a prefix with cached blocks pays less + prefill and reports `cached_tokens` (gRPC chunks, HTTP + `usage.prompt_tokens_details.cached_tokens`). | Flag | Default | Meaning | |------|---------|---------| -| `--prefill-tps` | 8000 | prefill throughput (tokens/s) | -| `--decode-base-ms` | 6.0 | fixed decode-step latency (ms) | -| `--decode-per-req-ms` | 0.35 | added decode latency per running request | -| `--prefill-chunk` | 2048 | max prompt tokens prefilled per step | -| `--max-running` | 256 | continuous-batching width | -| `--kv-tokens` | 524288 | KV cache capacity (tokens) | -| `--block-size` | 16 | cache block/page size (tokens) | -| `--prefix-cache` | true | enable prefix caching + KV events | +| `--timing` | polynomial | `polynomial`, `linear`, or `fit:` (a hardware calibration JSON, below) | +| `--prefill-poly a,b,c` | AISimulate | prefill ms = a + b·T + c·T² | +| `--decode-poly a,b,c` | AISimulate | decode ms = max(1, a + b·u + c·u²) | +| `--request-overhead-ms` | 0 | fixed per-request latency added to every event of a stream (TTFT and e2e grow by it, ITL does not) | +| `--prefill-tps` | 8000 | linear model: prefill tokens/s (selects `linear`) | +| `--decode-base-ms` | 6.0 | linear model: fixed decode-pass ms | +| `--decode-per-req-ms` | 0.35 | linear model: decode ms per running request | +| `--max-batched-tokens` | 8192 | token budget per pass | +| `--max-running` | 256 | sequences per pass | +| `--kv-tokens` / `--kv-blocks` | 524288 tokens | KV pool capacity | +| `--block-size` | 16 | cache block/page size (tokens); must match the worker's `kv_block_size` | +| `--prefix-cache` | true | prefix caching + KV events | +| `--prefill-first` | false | SGLang-style prefill-only passes | +| `--reserve-full-isl` | true | admit a waiting prompt only when the free and evictable blocks could hold all of it beyond its cached blocks, stopping at the first that does not fit (vLLM's `scheduler_reserve_full_isl`: a gate, only the pass's chunk is allocated at admission); `false` skips the gate | +| `--context-length` | 32768 | advertised context length | +| `--loads-like` | mock | `vllm`: report only what the vLLM servicer reports (running, waiting, `token_usage`, maxima), so the gateway's expected-wait routes as it does on a vLLM fleet | +| `--admin-port` | off | process-wide admin API (below) | +| `--kv-events-zmq-base-port` | off | ZMQ KV-event publishers on vLLM's or SGLang's wire (below) | ```bash cargo run --release -p mock-worker -- \ - --engine realistic --grpc-base-port 19000 --grpc-count 8 --model mock-model + --engine realistic --grpc-base-port 19000 --grpc-count 8 --model mock-model --admin-port 19100 ``` +`--timing fit:` replaces the uncalibrated polynomials with a hardware +calibration the GPU harness writes. Prefill comes either from a measured +point table (`prefill_table_ms: [[tokens, ms], ...]`, or the harness's +`prefill_points_ms: {"": {"median_ms": ..}}`), interpolated linearly +and extrapolated with the last slope, together with a pass-total form for +batched prefills (`prefill_pass_ms: {"intercept_ms", "ms_per_token"}`, from a +concurrent sweep): a lone request costs its table value, a pass that prefills +several requests costs the pass-total form and never less than the table +value of its largest request; or, without a table, from the polynomial +`prefill_fit_ms` / `prefill_ms` (`{"a","b","c"}` or `[a,b,c]`). Decode is +`decode_fit_vs_utilisation_ms` / `decode_ms` (`{"d","e","f"}` or `[d,e,f]`, ms +over KV utilisation); capacity is `kv_capacity_tokens` or +`kv_capacity_blocks` with `block_size`; `request_overhead_ms` shifts every +event of a stream. Unknown keys are ignored, and `--block-size`, +`--kv-tokens`/`--kv-blocks` and `--request-overhead-ms` given explicitly win +over the file (a restricted pool is, for example, `--kv-blocks 12000 +--block-size 16`); the decode fit's utilisation is still read against the +file's `kv_capacity_tokens`, so a smaller pool changes how much fits, not +how long a decode step takes. The defaults stay AISimulate's uncalibrated +baseline. + +Agreement with hardware is the caller's problem: AISimulate's published +agreement for these polynomials (mean absolute percentage error 48.5% on TTFT, +28.9% on TPOT) was measured with prefix caching disabled, so cache-hit and +routing effects have no published validation. Treat the simulator as a relative +A/B harness for policies and validate absolute numbers on GPUs. + +### KV-event publisher on the engines' ZMQ wire + +`--kv-events-zmq-base-port ` gives every realistic engine the publisher +vLLM runs (`ZmqEventPublisher` in `vllm/distributed/kv_events.py`), so the +servicers' relays (`crates/engine_servicer`, the Python servicer) can be +exercised end to end without a GPU: worker `i` (gRPC workers by port offset, +then ZMQ ranks) publishes on `base + 2i` and answers replay on `base + 2i + 1`. + +- PUB frames: `[topic, sequence as u64 big-endian, msgpack]`, the sequence + counting from 0 per publisher (`--kv-events-topic`, empty by default as + vLLM's). +- Payload: vLLM's `EventBatch`, `[ts, events, data_parallel_rank]`; events are + tagged maps in vLLM's field order with its `omit_defaults`: `BlockStored` + (`type`, `block_hashes`, `parent_block_hash`, `token_ids`, `block_size`, + `lora_id`, `medium: "GPU"`, `lora_name`, `group_idx: 0`, + `kv_cache_spec_kind: "full_attention"`), `BlockRemoved` (`type`, + `block_hashes`, `medium`, `group_idx`), `AllBlocksCleared` (`type`). Block + hashes are unsigned 64-bit integers (vLLM's int form of the same bits the + gRPC stream carries signed). +- Replay (`--kv-events-replay`, on by default): a DEALER sends + `[b"", start as 8 bytes big-endian]`; the ROUTER answers every buffered + batch from `start` as `[b"", topic, seq, payload]` and then + `[b"", b"", END, b""]` with END = eight 0xff bytes; the last + `--kv-events-buffer-steps` (10000) batches are kept. + +`--kv-events-wire sglang` speaks SGLang's publisher instead: signed 64-bit +hashes (SGLang takes the first eight digest bytes signed), `BlockStored` with +only `block_hashes`, `parent_block_hash`, `token_ids`, `block_size`, `lora_id` +(one per radix node: here one per contiguous run a request completed), +`BlockRemoved` with one node's hashes (here one per evicted block), a nil +`attn_dp_rank`, an `AllBlocksCleared` batch at startup as the scheduler +publishes, and replay replies without the topic frame (`[b"", seq, payload]`, +then `[b"", END, b""]`). + +The unit tests decode both wires with the Rust relay's own normalizer +(`engine_servicer::kv_wire`), so a change on either side shows up here. + +### Admin API + +`--admin-port` serves the ground truth a routing benchmark needs and real +engines do not expose: + +- `GET /admin/health` — `ok`; +- `GET /admin/fleet` — every engine (`grpc:` / `http:` / + `zmq:`) with cache size, load, cached blocks and preemptions; +- `GET /admin/requests?since=&limit=` — admitted requests with the + serving worker, prompt/cached tokens, queue wait and the **arrival-time + oracle**: the most cached tokens any worker of the process held when the + request arrived (the best a router could have obtained). Join on the + gateway's response `id`; +- `GET /admin/cache/{worker}` — the worker's cached block keys; +- `POST /admin/reset[/{worker}]` — clear caches and publish `AllBlocksCleared` + (an engine restart, to the index). + **Tokenizer note (gRPC):** the gateway tokenizes prompts before routing, so it needs a real tokenizer for the model. Register each worker with a tokenizer label, e.g. `"labels":{"tokenizer_path":"gpt2"}`, and a `"kv_block_size":16`, and do **not** pass `--disable-tokenizer-autoload`. (HTTP workers need no tokenizer but cannot drive event-driven `cache_aware`, which requires token ids.) +### Fault hooks + +All under the admin API; `{worker}` is a worker name (`grpc:`, +`zmq:`) or `all`. A hook applies to every KV-event transport of the +worker (gRPC `SubscribeKvEvents` and the ZMQ publisher alike) and answers with +the worker's hook state (`drop_pending`, `dropped_total`, `delay_ms`, +`paused`, `generation`, `restarts`). + +| Hook | Effect | +|------|--------| +| `POST /admin/fault/{worker}/drop?batches=N` | the next N event batches are not published (lost on the wire; they stay in the replay buffer, so a gap replay recovers them) | +| `POST /admin/fault/{worker}/delay?ms=D` | every batch is published D ms after its pass ends (0 clears) | +| `POST /admin/fault/{worker}/restart-publisher` | the publisher restarts: gRPC sequence numbers start over at 1 and the ZMQ sequence at 0, the replay buffers are emptied, the cache is kept (no `AllBlocksCleared`); `generation` increments | +| `POST /admin/fault/{worker}/pause` | the engine freezes after its current pass: no passes, no tokens, no events; requests queue (and count as waiting); health and `GetLoads` keep answering | +| `POST /admin/fault/{worker}/resume` | the engine runs again | +| `GET /admin/fault/{worker}` | the hooks' current state (pending drops, delay, restarts, paused, generation) | +| `POST /admin/reset/{worker}` | (already there) clear the cache and publish `AllBlocksCleared` | + +### Truth endpoints + +| Endpoint | Answer | +|----------|--------| +| `POST /admin/truth/{worker}` with `{"token_ids": [...]}` | what the worker would serve from cache for that prompt right now: `cached_tokens`, `cached_blocks`, `block_size` (the engine's own prefix match, last-block rule included) | +| `GET /admin/truth` | per worker, over every admitted request: `requests`, `prompt_tokens`, `cached_tokens`, `oracle_tokens`, so a gateway's hit-rate claim can be checked against what the engines actually served | ## Capturing requests `--capture PATH` appends every gRPC `Generate` request a worker receives to @@ -142,3 +289,62 @@ scripts/sim_ab.sh --mode grpc --workers 8 --rps 120 --duration 30 # Offline-friendly: no tokenizer; drives least_load + latency + approximate cache_aware. scripts/sim_ab.sh --mode http --workers 16 --policies "random least_load" ``` + +## Replaying a trace +The `replay` binary of this crate (`cargo run --release -p mock-worker --bin replay`) +replays a Mooncake-format trace (`{timestamp, input_length, output_length, +hash_ids}` per line, one `hash_id` per 512-token block) through the gateway and +scores routing quality end to end. + +- Every `hash_id` becomes a deterministic text block (`--words-per-block`, 480 + words, about one token each), so rows that share ids share prompt prefixes + after the gateway tokenizes them. +- Requests are sent open-loop at `timestamp / --speedup` as streaming chat + completions with `stream_options.include_usage`, recording TTFT, inter-token + latencies, end-to-end latency, the serving worker (`system_fingerprint`, which + the gateway sets from the worker's `weight_version` label), and the + engine-reported `cached_tokens`. +- With `--admin ` each request is joined with the mock fleet's + record of it (`GET /admin/requests`), adding the arrival-time oracle (the most + cached tokens any worker held when it arrived) and the queue wait. + +Output: `summary.json` (mean/p50/p90/p99 TTFT, per-request mean ITL (TPOT) +distribution, e2e latency, goodput at the SLO `--slo-ttft-ms` / `--slo-itl-ms` +(default 500 ms TTFT and 50 ms per-request mean ITL; a strict variant uses the +per-request p99), prefix reuse = cached / prompt tokens, oracle prefix reuse, +hit-over-oracle, per-worker request and uncached-token counts, balance) and +`requests.csv` with one row per request. + +```bash +replay --trace mooncake_trace.jsonl --gateway http://127.0.0.1:31000 \ + --model mock-model --speedup 4 --limit 5000 --admin http://127.0.0.1:31002 --out out/ +``` + +With `--gateway-log ` (the gateway run at `--log-level debug`) the +gateway's routing decisions are joined to the requests by request id (the +`x-request-id` header, which is the response id without its uuid tail) and +`t4.md` lists them per request: `phase, idx, worker, branch, prompt_tokens, +engine cached_tokens, implied overlap, agree`. The implied overlap is the +gateway's stated credit when its log carries one (`overlap_tokens=`, +`overlap_blocks=` or the tree path's `matched_ratio=`); otherwise `agree` +compares the branch's claim of an overlap (`event_hit`, `event_spill`) with +whether the engine served cached tokens. With `--admin` the summary also +carries `engine_truth_per_worker`, each engine's own account of the prompt, +cached and oracle tokens it served, and `fleet.csv` samples every worker's +load, cached blocks, preemptions and KV batches once a second for the whole +run. + +The mock fleet it is meant for is `mock-worker --engine realistic --admin-port` +(the [Realistic engine](#realistic-engine) above; note the caveat that its timing +polynomials were validated against hardware only with prefix caching off). + +### Long runs + +For a run of hours, keep one gateway and one fleet up and start the replayer +window by window (`--skip` advancing by `--limit`, a `--label` per window, +`--admin` for the oracle join), driving the admin fault hooks on a cycle over +the workers in turn (lost batches, a publishing delay, a publisher restart, a +pause and resume, a cache reset) and sampling the gateway's `/metrics` once a +minute. The memory reading that matters is the gateway allocator's live bytes +at idle, hour over hour; its RSS shows how much the allocator keeps, not what +is live. diff --git a/crates/mock_worker/src/admin.rs b/crates/mock_worker/src/admin.rs new file mode 100644 index 0000000000..7dd36ac7f8 --- /dev/null +++ b/crates/mock_worker/src/admin.rs @@ -0,0 +1,336 @@ +//! Process-wide admin API for the simulated fleet: the ground truth a routing +//! benchmark needs that real engines do not expose, and the fault hooks a +//! recovery test switches on. +//! +//! - `GET /admin/fleet`: every registered engine with its cache size and load. +//! - `GET /admin/requests?since=&limit=`: admitted-request records +//! (`request_id`, serving worker, prompt/cached/oracle tokens, queue wait), +//! oldest first, `seq` strictly greater than `since`. +//! - `GET /admin/cache/{worker}`: the worker's cached block keys. +//! - `POST /admin/reset/{worker}` and `POST /admin/reset`: clear one or every +//! cache and publish `AllBlocksCleared` (an engine restart, to the index). +//! - `POST /admin/fault/{worker}/drop?batches=N`, `.../delay?ms=D`, +//! `.../restart-publisher`, `.../pause`, `.../resume` and +//! `GET /admin/fault/{worker}`: the fault hooks (see the README). +//! - `POST /admin/truth/{worker}` with `{"token_ids": [...]}`: what the +//! worker would serve from cache for that prompt right now; +//! `GET /admin/truth`: per worker, what it served over every admitted +//! request. +//! +//! Worker names are `grpc:` / `http:` / `zmq:`; `all` +//! addresses every worker. + +use std::{ + collections::{BTreeMap, HashMap}, + sync::Arc, +}; + +use axum::{ + extract::{Path, Query, State}, + http::StatusCode, + response::{IntoResponse, Response}, + routing::{get, post}, + Json, Router, +}; +use serde_json::{json, Value}; +use tokio::net::TcpListener; + +use crate::{ + config::Config, + engine::{self, Engine}, +}; + +struct AdminState { + cfg: Arc, +} + +fn router(state: Arc) -> Router { + Router::new() + .route("/admin/health", get(health)) + .route("/admin/fleet", get(fleet)) + .route("/admin/requests", get(requests)) + .route("/admin/cache/{worker}", get(cache)) + .route("/admin/reset", post(reset_all)) + .route("/admin/reset/{worker}", post(reset_one)) + .route("/admin/fault/{worker}", get(fault_status)) + .route("/admin/fault/{worker}/drop", post(fault_drop)) + .route("/admin/fault/{worker}/delay", post(fault_delay)) + .route( + "/admin/fault/{worker}/restart-publisher", + post(fault_restart_publisher), + ) + .route("/admin/fault/{worker}/pause", post(fault_pause)) + .route("/admin/fault/{worker}/resume", post(fault_resume)) + .route("/admin/truth", get(truth_served)) + .route("/admin/truth/{worker}", post(truth_prompt)) + .with_state(state) +} + +/// Serve the admin API on `port` until the process exits. +pub async fn serve(cfg: Arc, host: String, port: u16) { + let listener = match TcpListener::bind((host.as_str(), port)).await { + Ok(listener) => listener, + Err(e) => { + tracing::error!("admin bind {host}:{port} failed: {e}"); + return; + } + }; + let state = Arc::new(AdminState { cfg }); + if let Err(e) = axum::serve(listener, router(state)).await { + tracing::error!("admin server stopped: {e}"); + } +} + +async fn health() -> &'static str { + "ok" +} + +fn find(name: &str) -> Option { + engine::fleet_engines() + .into_iter() + .find(|e| e.name() == name) +} + +/// The engines a `{worker}` path segment addresses: one by name, or `all`. +fn select(worker: &str) -> Result, Response> { + if worker == "all" { + return Ok(engine::fleet_engines()); + } + find(worker) + .map(|e| vec![e]) + .ok_or_else(|| (StatusCode::NOT_FOUND, "unknown worker").into_response()) +} + +fn status_json(e: &Engine) -> Value { + let s = e.fault_status(); + json!({ + "worker": e.name(), + "drop_pending": s.drop_pending, + "dropped_total": s.dropped_total, + "delay_ms": s.delay_ms, + "paused": s.paused, + "generation": s.generation, + "restarts": s.restarts, + }) +} + +fn param(q: &HashMap, name: &str) -> Result { + q.get(name).and_then(|v| v.parse().ok()).ok_or_else(|| { + ( + StatusCode::BAD_REQUEST, + format!("missing or invalid query parameter {name}"), + ) + .into_response() + }) +} + +async fn fleet(State(state): State>) -> Json { + let workers: Vec = engine::fleet_engines() + .iter() + .map(|e| { + let load = e.load(); + json!({ + "worker": e.name(), + "cache_blocks": e.cache_keys().len(), + "block_size": state.cfg.engine.block_size, + "num_running_reqs": load.num_running_reqs, + "num_waiting_reqs": load.num_waiting_reqs, + "num_waiting_uncached_tokens": load.num_waiting_uncached_tokens, + "token_usage": load.token_usage, + "cache_hit_rate": load.cache_hit_rate, + "num_cached_blocks": load.num_cached_blocks, + "num_preemptions": load.num_preemptions, + "num_kv_batches": load.num_kv_batches, + }) + }) + .collect(); + Json(json!({ "workers": workers })) +} + +async fn requests(Query(q): Query>) -> Json { + let since = q.get("since").and_then(|v| v.parse().ok()).unwrap_or(0u64); + let limit = q + .get("limit") + .and_then(|v| v.parse().ok()) + .unwrap_or(100_000usize); + let records = engine::records_since(since, limit); + let next = records.last().map(|r| r.seq).unwrap_or(since); + let rows: Vec = records + .iter() + .map(|r| { + json!({ + "seq": r.seq, + "request_id": r.request_id, + "worker": r.worker, + "prompt_tokens": r.prompt_tokens, + "cached_tokens": r.cached_tokens, + "oracle_tokens": r.oracle_tokens, + "queued_ms": r.queued_ms, + "running_at_admit": r.running_at_admit, + "waiting_at_admit": r.waiting_at_admit, + "admitted_unix_ms": r.admitted_unix_ms, + }) + }) + .collect(); + Json(json!({ "records": rows, "next": next })) +} + +async fn cache(Path(worker): Path) -> Response { + match find(&worker) { + Some(e) => { + let mut keys = e.cache_keys(); + keys.sort_unstable(); + Json(json!({ "worker": worker, "blocks": keys })).into_response() + } + None => (StatusCode::NOT_FOUND, "unknown worker").into_response(), + } +} + +async fn reset_one(Path(worker): Path) -> Response { + match select(&worker) { + Ok(engines) => { + let names: Vec = engines + .iter() + .map(|e| { + e.reset(); + e.name().to_string() + }) + .collect(); + Json(json!({ "reset": names })).into_response() + } + Err(response) => response, + } +} + +async fn reset_all() -> Json { + let names: Vec = engine::fleet_engines() + .iter() + .map(|e| { + e.reset(); + e.name().to_string() + }) + .collect(); + Json(json!({ "reset": names })) +} + +async fn fault_status(Path(worker): Path) -> Response { + match select(&worker) { + Ok(engines) => { + let workers: Vec = engines.iter().map(status_json).collect(); + Json(json!({ "workers": workers })).into_response() + } + Err(response) => response, + } +} + +async fn fault_drop( + Path(worker): Path, + Query(q): Query>, +) -> Response { + let batches: u32 = match param(&q, "batches") { + Ok(v) => v, + Err(response) => return response, + }; + apply(&worker, |e| e.fault_drop(batches)) +} + +async fn fault_delay( + Path(worker): Path, + Query(q): Query>, +) -> Response { + let ms: u64 = match param(&q, "ms") { + Ok(v) => v, + Err(response) => return response, + }; + apply(&worker, |e| e.fault_delay_ms(ms)) +} + +async fn fault_restart_publisher(Path(worker): Path) -> Response { + match select(&worker) { + Ok(engines) => { + for e in &engines { + e.restart_publisher().await; + } + let workers: Vec = engines.iter().map(status_json).collect(); + Json(json!({ "workers": workers })).into_response() + } + Err(response) => response, + } +} + +async fn fault_pause(Path(worker): Path) -> Response { + apply(&worker, Engine::pause) +} + +async fn fault_resume(Path(worker): Path) -> Response { + apply(&worker, Engine::resume) +} + +/// Run `hook` on the addressed engines and answer with their fault state. +fn apply(worker: &str, hook: impl Fn(&Engine)) -> Response { + match select(worker) { + Ok(engines) => { + for e in &engines { + hook(e); + } + let workers: Vec = engines.iter().map(status_json).collect(); + Json(json!({ "workers": workers })).into_response() + } + Err(response) => response, + } +} + +async fn truth_prompt(Path(worker): Path, Json(body): Json) -> Response { + let Some(token_ids) = body.get("token_ids").and_then(Value::as_array).map(|ids| { + ids.iter() + .filter_map(Value::as_u64) + .map(|t| u32::try_from(t).unwrap_or(u32::MAX)) + .collect::>() + }) else { + return (StatusCode::BAD_REQUEST, "body needs token_ids: [u32]").into_response(); + }; + match select(&worker) { + Ok(engines) => { + let workers: Vec = engines + .iter() + .map(|e| { + let truth = e.cached_tokens_for(&token_ids); + json!({ + "worker": e.name(), + "cached_tokens": truth.cached_tokens, + "cached_blocks": truth.cached_blocks, + "block_size": truth.block_size, + "prompt_tokens": token_ids.len(), + }) + }) + .collect(); + Json(json!({ "workers": workers })).into_response() + } + Err(response) => response, + } +} + +/// Per worker, what it actually served over every admitted request. +async fn truth_served() -> Json { + let mut totals: BTreeMap = BTreeMap::new(); + for r in engine::records_since(0, usize::MAX) { + let t = totals.entry(r.worker).or_insert((0, 0, 0, 0)); + t.0 += 1; + t.1 += u64::from(r.prompt_tokens); + t.2 += u64::from(r.cached_tokens); + t.3 += u64::from(r.oracle_tokens); + } + let workers: Vec = totals + .iter() + .map(|(worker, (requests, prompt, cached, oracle))| { + json!({ + "worker": worker, + "requests": requests, + "prompt_tokens": prompt, + "cached_tokens": cached, + "oracle_tokens": oracle, + }) + }) + .collect(); + Json(json!({ "workers": workers })) +} diff --git a/crates/mock_worker/src/bin/replay.rs b/crates/mock_worker/src/bin/replay.rs new file mode 100644 index 0000000000..dd184859f5 --- /dev/null +++ b/crates/mock_worker/src/bin/replay.rs @@ -0,0 +1,1535 @@ +//! Replay a Mooncake-format trace through the gateway and score the routing. +//! +//! Each trace row is `{timestamp (ms), input_length, output_length, hash_ids}`, +//! where every `hash_id` stands for one 512-token block of prompt. The replayer +//! synthesizes a deterministic text block per `hash_id` (so rows sharing ids +//! share prompt prefixes exactly as the trace intends), sends each request as a +//! streaming chat completion at `timestamp / speedup` (open loop), and records +//! time to first token, inter-token latencies, the serving worker (the +//! gateway's `system_fingerprint`, set from the worker's `weight_version` +//! label) and the engine-reported `cached_tokens`. When the mock fleet's admin +//! API is given, every request is joined with the fleet's record of it, which +//! carries the arrival-time oracle: the most cached tokens any worker held. + +// A command-line tool: the summary goes to stdout, progress to stderr. +#![allow(clippy::print_stdout, clippy::print_stderr)] + +use std::{ + collections::{BTreeMap, HashMap}, + fs, + io::Write, + path::PathBuf, + sync::Arc, + time::{Duration, Instant}, +}; + +use anyhow::{anyhow, Context, Result}; +use clap::Parser; +use futures::StreamExt; +use serde::{Deserialize, Serialize}; +use serde_json::{json, Value}; +use tokio::{ + sync::{watch, Semaphore}, + task::JoinSet, +}; + +#[derive(Parser, Debug, Clone)] +#[command(about = "Replay a Mooncake trace through the gateway and score routing quality")] +struct Args { + /// Mooncake JSONL trace. + #[arg(long)] + trace: PathBuf, + /// Gateway base URL. + #[arg(long, default_value = "http://127.0.0.1:31000")] + gateway: String, + /// Model id to request. + #[arg(long, default_value = "mock-model")] + model: String, + /// Arrival speedup: trace time is divided by this (finite, greater than 0). + #[arg(long, default_value_t = 2.0, value_parser = parse_speedup)] + speedup: f64, + /// Rows to skip from the start of the trace. + #[arg(long, default_value_t = 0)] + skip: usize, + /// Rows to replay after `skip` (0 = all). + #[arg(long, default_value_t = 4000)] + limit: usize, + /// Words synthesized per 512-token trace block (about one token each). + #[arg(long, default_value_t = 480)] + words_per_block: usize, + /// Cap on `max_tokens` per request (the trace's output_length otherwise). + #[arg(long, default_value_t = 512)] + max_output: u32, + /// Safety cap on concurrently open requests. + #[arg(long, default_value_t = 4096)] + max_inflight: usize, + /// Mock fleet admin API base URL (enables the oracle join). + #[arg(long)] + admin: Option, + /// First rows whose results are excluded from the statistics (warm-up). + #[arg(long, default_value_t = 0)] + warmup: usize, + /// TTFT SLO in ms for goodput. + #[arg(long, default_value_t = 500.0)] + slo_ttft_ms: f64, + /// Per-request mean inter-token latency (TPOT) SLO in ms for goodput. + #[arg(long, default_value_t = 50.0)] + slo_itl_ms: f64, + /// Output directory for `summary.json` and `requests.csv`. + #[arg(long, default_value = "replay-out")] + out: PathBuf, + /// Label stored in the summary. + #[arg(long, default_value = "run")] + label: String, + /// Seed mixed into the synthesized text. + #[arg(long, default_value_t = 7)] + seed: u64, + /// The gateway's log at debug level: its routing decisions are joined to + /// the requests by request id, so the gateway's cache credit can be + /// compared with what the engine served. + #[arg(long)] + gateway_log: Option, + /// The fleet's block size, to turn a credit in blocks into tokens. + #[arg(long, default_value_t = 16)] + block_size: u32, + /// Rows of the per-request decision table written to `t4.md`. + #[arg(long, default_value_t = 40)] + t4_rows: usize, +} + +#[derive(Deserialize, Debug, Clone)] +struct TraceRow { + timestamp: u64, + input_length: u32, + output_length: u32, + hash_ids: Vec, +} + +#[derive(Serialize, Debug, Clone, Default)] +struct ReqResult { + row: usize, + trace_ts_ms: u64, + trace_input_length: u32, + sent_at_ms: f64, + status: String, + request_id: String, + worker: String, + prompt_tokens: u32, + completion_tokens: u32, + cached_tokens: u32, + oracle_tokens: Option, + queued_ms: Option, + ttft_ms: Option, + latency_ms: f64, + itl_mean_ms: Option, + itl_p99_ms: Option, + tokens_seen: u32, + /// Output tokens the engine counted but the stream never showed as text + /// (an incomplete UTF-8 piece the detokenizer holds back). + invisible_tokens: u32, + /// The gateway's `x-request-id` (the response id without its uuid tail). + gateway_request_id: String, + /// The gateway's routing branch for this request (from its debug log). + branch: String, + /// The gateway's cache credit in tokens, when its log states one. + credit_tokens: Option, + /// Whether the gateway's credit agrees with what the engine served. + agree: Option, +} + +const WORDS: &[&str] = &[ + "time", + "year", + "people", + "way", + "day", + "man", + "thing", + "woman", + "life", + "child", + "world", + "school", + "state", + "family", + "student", + "group", + "country", + "problem", + "hand", + "part", + "place", + "case", + "week", + "company", + "system", + "program", + "question", + "work", + "government", + "number", + "night", + "point", + "home", + "water", + "room", + "mother", + "area", + "money", + "story", + "fact", + "month", + "lot", + "right", + "study", + "book", + "eye", + "job", + "word", + "business", + "issue", + "side", + "kind", + "head", + "house", + "service", + "friend", + "father", + "power", + "hour", + "game", + "line", + "end", + "member", + "law", + "car", + "city", + "community", + "name", + "president", + "team", + "minute", + "idea", + "kid", + "body", + "information", + "back", + "parent", + "face", + "others", + "level", + "office", + "door", + "health", + "person", + "art", + "war", + "history", + "party", + "result", + "change", + "morning", + "reason", + "research", + "girl", + "guy", + "moment", + "air", + "teacher", + "force", + "education", + "foot", + "boy", + "age", + "policy", + "process", + "music", + "market", + "sense", + "nation", + "plan", + "college", + "interest", + "death", + "experience", + "effect", + "use", + "class", + "control", + "care", + "field", + "development", + "role", + "effort", + "rate", + "heart", + "drug", + "show", + "leader", + "light", + "voice", + "wife", + "police", + "mind", + "price", + "report", + "decision", + "son", + "view", + "relationship", + "town", + "road", + "arm", + "difference", + "value", + "building", + "action", + "model", + "season", + "society", + "tax", + "director", + "position", + "player", + "record", + "paper", + "space", + "ground", + "form", + "event", + "official", + "matter", + "center", + "couple", + "site", + "project", + "activity", + "star", + "table", + "need", + "court", + "oil", + "situation", + "cost", + "industry", + "figure", + "street", + "image", + "phone", + "data", + "picture", + "practice", + "piece", + "land", + "product", + "doctor", + "wall", + "patient", + "worker", + "news", + "test", + "movie", + "north", + "love", + "support", + "technology", + "step", + "baby", + "computer", + "type", + "attention", + "film", + "tree", + "source", + "nothing", + "network", + "trade", + "economy", + "author", + "window", + "energy", + "letter", + "church", + "cell", + "ship", + "island", + "plant", + "garden", + "river", + "bridge", + "engine", + "metal", + "glass", + "stone", + "storm", + "forest", + "valley", + "ocean", + "desert", + "signal", + "memory", + "logic", +]; + +/// Deterministic text for one trace block: `words_per_block` words drawn from +/// the vocabulary by a splitmix64 stream seeded with the block id. +fn block_text(hash_id: u64, seed: u64, words_per_block: usize) -> String { + let mut x = hash_id + .wrapping_mul(0x9E37_79B9_7F4A_7C15) + .wrapping_add(seed ^ 0xD1B5_4A32_D192_ED03); + let mut out = String::with_capacity(words_per_block * 7); + for i in 0..words_per_block { + x = x.wrapping_add(0x9E37_79B9_7F4A_7C15); + let mut z = x; + z = (z ^ (z >> 30)).wrapping_mul(0xBF58_476D_1CE4_E5B9); + z = (z ^ (z >> 27)).wrapping_mul(0x94D0_49BB_1331_11EB); + z ^= z >> 31; + if i > 0 { + out.push(' '); + } + out.push_str(WORDS[(z % WORDS.len() as u64) as usize]); + } + out +} + +fn prompt_for(row: &TraceRow, seed: u64, words_per_block: usize) -> String { + let mut prompt = String::new(); + for (i, id) in row.hash_ids.iter().enumerate() { + if i > 0 { + prompt.push('\n'); + } + prompt.push_str(&block_text(*id, seed, words_per_block)); + } + prompt +} + +/// `--speedup`: a divisor of trace time, so it must be a finite number +/// greater than zero; anything else would make a due time infinite, negative +/// or NaN. +fn parse_speedup(raw: &str) -> Result { + match raw.trim().parse::() { + Ok(v) if v.is_finite() && v > 0.0 => Ok(v), + _ => Err(format!( + "speedup must be a finite number greater than 0, got {raw}" + )), + } +} + +/// When a row stamped `trace_ts_ms` is due, relative to the window's first +/// row at `t0_ms`, at `speedup`. A row stamped before the first one (an +/// unsorted or hand-edited trace, or a window that starts on a reordered +/// row) is due at once rather than never: the difference saturates at zero +/// instead of wrapping. A due time too far ahead to represent is clamped. +fn due_after(trace_ts_ms: u64, t0_ms: u64, speedup: f64) -> Duration { + let secs = trace_ts_ms.saturating_sub(t0_ms) as f64 / 1000.0 / speedup; + Duration::try_from_secs_f64(secs).unwrap_or(Duration::MAX) +} + +fn percentile(sorted: &[f64], p: f64) -> f64 { + if sorted.is_empty() { + return f64::NAN; + } + let rank = ((sorted.len() - 1) as f64 * p).round() as usize; + sorted[rank.min(sorted.len() - 1)] +} + +fn mean(v: &[f64]) -> f64 { + if v.is_empty() { + f64::NAN + } else { + v.iter().sum::() / v.len() as f64 + } +} + +struct Job { + client: reqwest::Client, + url: String, + model: String, + row_index: usize, + row: TraceRow, + body_prompt: String, + max_output: u32, + sent_at_ms: f64, +} + +async fn run_one(job: Job) -> ReqResult { + let Job { + client, + url, + model, + row_index, + row, + body_prompt, + max_output, + sent_at_ms, + } = job; + let max_tokens = row.output_length.clamp(1, max_output); + let body = json!({ + "model": model, + "messages": [{"role": "user", "content": body_prompt}], + "max_tokens": max_tokens, + "temperature": 0, + "stream": true, + "stream_options": {"include_usage": true}, + }); + let mut result = ReqResult { + row: row_index, + trace_ts_ms: row.timestamp, + trace_input_length: row.input_length, + sent_at_ms, + status: "ok".to_string(), + ..Default::default() + }; + let started = Instant::now(); + let response = match client.post(&url).json(&body).send().await { + Ok(r) => r, + Err(e) => { + result.status = format!("send-error: {e}"); + result.latency_ms = started.elapsed().as_secs_f64() * 1000.0; + return result; + } + }; + if !response.status().is_success() { + result.status = format!("http-{}", response.status().as_u16()); + result.latency_ms = started.elapsed().as_secs_f64() * 1000.0; + return result; + } + result.gateway_request_id = response + .headers() + .get("x-request-id") + .and_then(|v| v.to_str().ok()) + .unwrap_or_default() + .to_string(); + let mut stream = response.bytes_stream(); + let mut buf: Vec = Vec::new(); + let mut observer = StreamObserver::default(); + while let Some(chunk) = stream.next().await { + let chunk = match chunk { + Ok(c) => c, + Err(e) => { + result.status = format!("stream-error: {e}"); + break; + } + }; + buf.extend_from_slice(&chunk); + // SSE events end with a blank line. + while let Some(pos) = find_double_newline(&buf) { + let event: Vec = buf.drain(..pos + 2).collect(); + let text = String::from_utf8_lossy(&event); + for line in text.lines() { + let Some(data) = line.strip_prefix("data:") else { + continue; + }; + let data = data.trim(); + if data == "[DONE]" { + continue; + } + let Ok(v) = serde_json::from_str::(data) else { + continue; + }; + observer.observe(&v, Instant::now()); + } + } + } + result.latency_ms = started.elapsed().as_secs_f64() * 1000.0; + observer.finish(started, &mut result); + result +} + +/// What one streamed response reveals, chunk by chunk. +#[derive(Default)] +struct StreamObserver { + request_id: String, + worker: String, + /// The first chunk that carried a token or the finish: a one-token + /// answer whose token has no visible text still arrives here. + first_signal: Option, + last_token: Option, + itls: Vec, + tokens_seen: u32, + prompt_tokens: u32, + completion_tokens: u32, + cached_tokens: u32, + saw_usage: bool, +} + +impl StreamObserver { + fn observe(&mut self, v: &Value, now: Instant) { + if self.request_id.is_empty() { + if let Some(id) = v.get("id").and_then(Value::as_str) { + self.request_id = id.to_string(); + } + } + if self.worker.is_empty() { + if let Some(fp) = v.get("system_fingerprint").and_then(Value::as_str) { + self.worker = fp.to_string(); + } + } + let choice = v + .get("choices") + .and_then(Value::as_array) + .and_then(|c| c.first()); + let has_content = choice + .and_then(|c| c.get("delta")) + .and_then(|d| d.get("content")) + .and_then(Value::as_str) + .is_some_and(|s| !s.is_empty()); + let finished = choice + .and_then(|c| c.get("finish_reason")) + .is_some_and(|f| !f.is_null()); + if has_content { + self.tokens_seen += 1; + if let Some(prev) = self.last_token { + self.itls.push((now - prev).as_secs_f64() * 1000.0); + } + self.first_signal.get_or_insert(now); + self.last_token = Some(now); + } else if finished { + self.first_signal.get_or_insert(now); + } + if let Some(usage) = v.get("usage").filter(|u| !u.is_null()) { + self.saw_usage = true; + self.prompt_tokens = usage + .get("prompt_tokens") + .and_then(Value::as_u64) + .unwrap_or(0) as u32; + self.completion_tokens = usage + .get("completion_tokens") + .and_then(Value::as_u64) + .unwrap_or(0) as u32; + self.cached_tokens = usage + .get("prompt_tokens_details") + .and_then(|d| d.get("cached_tokens")) + .or_else(|| usage.get("cached_tokens")) + .and_then(Value::as_u64) + .unwrap_or(0) as u32; + } + } + + /// Fold the observations into the result. A stream that showed no text + /// but finished with a counted output token is a served request whose + /// token was invisible; one with no token at all is `no-tokens`. + fn finish(self, started: Instant, result: &mut ReqResult) { + result.request_id = self.request_id; + result.worker = self.worker; + result.prompt_tokens = self.prompt_tokens; + result.completion_tokens = self.completion_tokens; + result.cached_tokens = self.cached_tokens; + result.tokens_seen = self.tokens_seen; + result.ttft_ms = self + .first_signal + .map(|t| (t - started).as_secs_f64() * 1000.0); + if !self.itls.is_empty() { + let mut sorted = self.itls.clone(); + sorted.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal)); + result.itl_mean_ms = Some(mean(&self.itls)); + result.itl_p99_ms = Some(percentile(&sorted, 0.99)); + } + if result.completion_tokens < self.tokens_seen { + // No usage arrived: count what was seen. + result.completion_tokens = self.tokens_seen; + } + if self.tokens_seen == 0 && self.completion_tokens > 0 && result.ttft_ms.is_some() { + result.invisible_tokens = self.completion_tokens; + } + if result.status == "ok" && (result.ttft_ms.is_none() || result.completion_tokens == 0) { + result.status = "no-tokens".to_string(); + } + } +} + +/// One routing decision from the gateway's debug log. +#[derive(Clone, Debug, Default, PartialEq)] +struct Decision { + worker: String, + branch: String, + /// `matched_ratio=` on the tree path. + matched_ratio: Option, + /// `overlap_tokens=` / `overlap_blocks=`, when the gateway logs them. + overlap_tokens: Option, + overlap_blocks: Option, +} + +impl Decision { + /// The credit in tokens, floored to the block, when the log states one. + fn credit_tokens(&self, prompt_tokens: u32, block_size: u32) -> Option { + let bs = block_size.max(1); + if let Some(t) = self.overlap_tokens { + return Some(t / bs * bs); + } + if let Some(b) = self.overlap_blocks { + return Some(b * bs); + } + self.matched_ratio + .map(|r| ((r * f64::from(prompt_tokens)) as u32) / bs * bs) + } + + fn claims_overlap(&self) -> bool { + matches!(self.branch.as_str(), "event_hit" | "event_spill") + || self.branch.starts_with("hit") + || self.branch.contains("cache_hit") + } +} + +/// The value of `name="..."` or `name=` in a log line. +fn log_field<'a>(line: &'a str, name: &str) -> Option<&'a str> { + let start = line.find(&format!("{name}="))? + name.len() + 1; + let rest = &line[start..]; + if let Some(quoted) = rest.strip_prefix('"') { + quoted.split('"').next() + } else { + rest.split(|c: char| c.is_whitespace() || c == ',' || c == '}') + .next() + } +} + +/// A log line without its ANSI colour sequences (`ESC [ ... m`), which the +/// gateway writes even into a file. +fn strip_ansi(line: &str) -> String { + let mut out = String::with_capacity(line.len()); + let mut chars = line.chars().peekable(); + while let Some(c) = chars.next() { + if c == '\u{1b}' && chars.peek() == Some(&'[') { + chars.next(); + for d in chars.by_ref() { + if ('@'..='~').contains(&d) { + break; + } + } + continue; + } + out.push(c); + } + out +} + +/// The gateway's routing decisions by request id: the last decision line of +/// each request (`Event-driven routing` and `Cache-aware selection` lines). +fn parse_decisions(text: &str) -> HashMap { + let mut out = HashMap::new(); + for raw in text.lines() { + if !(raw.contains("Event-driven routing") || raw.contains("Cache-aware selection")) { + continue; + } + let clean = strip_ansi(raw); + let line = clean.as_str(); + let Some(id) = log_field(line, "request_id") else { + continue; + }; + let branch = log_field(line, "branch").unwrap_or_else(|| { + if line.contains("no overlap") { + "expected_wait_fallback" + } else { + "?" + } + }); + out.insert( + id.to_string(), + Decision { + worker: log_field(line, "worker").unwrap_or_default().to_string(), + branch: branch.to_string(), + matched_ratio: log_field(line, "matched_ratio").and_then(|v| v.parse().ok()), + overlap_tokens: log_field(line, "overlap_tokens").and_then(|v| v.parse().ok()), + overlap_blocks: log_field(line, "overlap_blocks").and_then(|v| v.parse().ok()), + }, + ); + } + out +} + +/// The gateway's request id is the response id without its uuid tail +/// (`-xxxxxxxx-xxxx-xxxx-xxxx-xxxxxxxxxxxx`, 37 bytes). +fn request_id_prefix(id: &str) -> String { + let bytes = id.as_bytes(); + if bytes.len() > 37 { + let tail = &bytes[bytes.len() - 37..]; + let uuid_shaped = [0usize, 9, 14, 19, 24].iter().all(|&i| tail[i] == b'-') + && tail + .iter() + .enumerate() + .all(|(i, b)| matches!(i, 0 | 9 | 14 | 19 | 24) || b.is_ascii_hexdigit()); + if uuid_shaped { + return id[..bytes.len() - 37].to_string(); + } + } + id.to_string() +} + +fn decision_summary(ok: &[&ReqResult], decisions: usize) -> Value { + let joined: Vec<&&ReqResult> = ok.iter().filter(|r| !r.branch.is_empty()).collect(); + let agree = joined.iter().filter(|r| r.agree == Some(true)).count(); + let credit_known = joined.iter().filter(|r| r.credit_tokens.is_some()).count(); + let mut by_branch: BTreeMap = BTreeMap::new(); + for r in &joined { + let e = by_branch.entry(r.branch.clone()).or_insert((0, 0, 0)); + e.0 += 1; + e.1 += usize::from(r.agree == Some(true)); + e.2 += u64::from(r.cached_tokens); + } + json!({ + "decisions_in_log": decisions, + "joined": joined.len(), + "credit_known": credit_known, + "agree": agree, + "disagree": joined.len() - agree, + "by_branch": by_branch.iter().map(|(b, (n, a, cached))| json!({ + "branch": b, "requests": n, "agree": a, + "engine_cached_tokens_mean": if *n == 0 { 0.0 } else { *cached as f64 / *n as f64 }, + })).collect::>(), + }) +} + +/// The per-request decision table (`t4.md`): `implied overlap` is the +/// gateway's stated credit when its log carries one, `-` otherwise (then +/// `agree` compares the branch's claim of an overlap with the engine). +fn decision_table(results: &[ReqResult], rows: usize) -> String { + let mut out = String::from( + "| phase | idx | worker | branch | prompt_tokens | engine cached_tokens | implied overlap | agree | +|---|---|---|---|---|---|---|---| +", + ); + let joined: Vec<&ReqResult> = results + .iter() + .filter(|r| r.status == "ok" && !r.branch.is_empty()) + .collect(); + for r in joined.iter().take(rows) { + out.push_str(&format!( + "| replay | {} | {} | {} | {} | {} | {} | {} | +", + r.row, + r.worker.rsplit(':').next().unwrap_or(&r.worker), + r.branch, + r.prompt_tokens, + r.cached_tokens, + r.credit_tokens + .map(|v| v.to_string()) + .unwrap_or_else(|| "-".to_string()), + r.agree.map(|v| v.to_string()).unwrap_or_default() + )); + } + let agree = joined.iter().filter(|r| r.agree == Some(true)).count(); + out.push_str(&format!( + " +agreement: {agree}/{} (gateway credit vs engine truth; {} rows shown) +", + joined.len(), + joined.len().min(rows) + )); + out +} + +fn find_double_newline(buf: &[u8]) -> Option { + buf.windows(2).position(|w| w == b"\n\n") +} + +#[derive(Deserialize, Debug)] +struct AdminRecords { + records: Vec, + next: u64, +} + +#[derive(Deserialize, Debug, Clone)] +struct AdminRecord { + request_id: String, + worker: String, + cached_tokens: u32, + oracle_tokens: u32, + queued_ms: f64, +} + +async fn fetch_admin_records(client: &reqwest::Client, admin: &str) -> Result> { + let mut since = 0u64; + let mut all = Vec::new(); + loop { + let page: AdminRecords = client + .get(format!("{admin}/admin/requests?since={since}&limit=50000")) + .send() + .await? + .error_for_status()? + .json() + .await?; + if page.records.is_empty() { + break; + } + since = page.next; + all.extend(page.records); + } + Ok(all) +} + +/// One worker's line of the mock admin API's `GET /admin/fleet` snapshot. +#[derive(Deserialize, Debug, Clone)] +struct FleetWorker { + worker: String, + #[serde(default)] + num_running_reqs: u64, + #[serde(default)] + num_waiting_reqs: u64, + #[serde(default)] + num_waiting_uncached_tokens: u64, + #[serde(default)] + token_usage: f64, + #[serde(default)] + num_cached_blocks: u64, + #[serde(default)] + num_preemptions: u64, + #[serde(default)] + num_kv_batches: u64, +} + +#[derive(Deserialize, Debug)] +struct FleetSnapshot { + workers: Vec, +} + +const FLEET_CSV_HEADER: &str = + "t_s,worker,running,waiting,waiting_uncached_tokens,token_usage,cached_blocks,preemptions,kv_batches"; + +/// Samples the fleet's admin snapshot once a second into `path` (`fleet.csv`: +/// one row per worker per tick, seconds since `start`) until `stop` turns +/// true. The header is written before the first poll, so the file exists +/// with a header even when the admin API never answers; a failed poll is +/// skipped, not fatal. Returns the number of rows written. +async fn sample_fleet( + client: reqwest::Client, + admin: String, + path: PathBuf, + start: Instant, + mut stop: watch::Receiver, +) -> Result { + if let Some(dir) = path.parent() { + fs::create_dir_all(dir)?; + } + let mut out = String::from(FLEET_CSV_HEADER); + out.push('\n'); + fs::write(&path, &out)?; + let mut file = fs::OpenOptions::new().append(true).open(&path)?; + let url = format!("{}/admin/fleet", admin.trim_end_matches('/')); + let mut tick = tokio::time::interval(Duration::from_secs(1)); + tick.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay); + let mut rows = 0usize; + loop { + tokio::select! { + _ = tick.tick() => {} + changed = stop.changed() => { + if changed.is_err() || *stop.borrow() { + break; + } + continue; + } + } + let t_s = start.elapsed().as_secs_f64(); + let snapshot = match client.get(&url).send().await { + Ok(resp) => resp.json::().await, + Err(e) => Err(e), + }; + let Ok(snapshot) = snapshot else { + continue; + }; + let mut chunk = String::new(); + for w in &snapshot.workers { + chunk.push_str(&format!( + "{t_s:.1},{},{},{},{},{:.4},{},{},{}\n", + w.worker, + w.num_running_reqs, + w.num_waiting_reqs, + w.num_waiting_uncached_tokens, + w.token_usage, + w.num_cached_blocks, + w.num_preemptions, + w.num_kv_batches + )); + rows += 1; + } + file.write_all(chunk.as_bytes())?; + } + file.flush()?; + Ok(rows) +} + +#[tokio::main] +async fn main() -> Result<()> { + let args = Args::parse(); + let text = fs::read_to_string(&args.trace) + .with_context(|| format!("reading {}", args.trace.display()))?; + let rows: Vec = text + .lines() + .filter(|l| !l.trim().is_empty()) + .map(|l| serde_json::from_str::(l).map_err(|e| anyhow!("{e}: {l}"))) + .collect::>()?; + let end = if args.limit == 0 { + rows.len() + } else { + (args.skip + args.limit).min(rows.len()) + }; + let rows: Vec = rows[args.skip.min(rows.len())..end].to_vec(); + if rows.is_empty() { + return Err(anyhow!("no rows selected")); + } + let t0 = rows[0].timestamp; + let client = reqwest::Client::builder() + .pool_max_idle_per_host(args.max_inflight) + .timeout(Duration::from_secs(600)) + .build()?; + let url = format!("{}/v1/chat/completions", args.gateway.trim_end_matches('/')); + let inflight = Arc::new(Semaphore::new(args.max_inflight)); + let mut set: JoinSet = JoinSet::new(); + let start = Instant::now(); + // Per-second fleet series (admin snapshots) for the whole run, when the + // admin API is given; stopped once every request has finished. + let (stop_fleet, fleet_stop_rx) = watch::channel(false); + let mut fleet_set: JoinSet> = JoinSet::new(); + if let Some(admin) = &args.admin { + fleet_set.spawn(sample_fleet( + client.clone(), + admin.clone(), + args.out.join("fleet.csv"), + start, + fleet_stop_rx, + )); + } + eprintln!( + "replaying {} rows at {}x ({:.0} s of trace), {} words/block, max_output {}", + rows.len(), + args.speedup, + rows.last().map_or(0, |r| r.timestamp.saturating_sub(t0)) as f64 / 1000.0, + args.words_per_block, + args.max_output + ); + for (i, row) in rows.iter().enumerate() { + let due = due_after(row.timestamp, t0, args.speedup); + let elapsed = start.elapsed(); + if due > elapsed { + tokio::time::sleep(due - elapsed).await; + } + let permit = inflight.clone().acquire_owned().await?; + let prompt = prompt_for(row, args.seed, args.words_per_block); + let job = Job { + client: client.clone(), + url: url.clone(), + model: args.model.clone(), + row_index: i, + row: row.clone(), + body_prompt: prompt, + max_output: args.max_output, + sent_at_ms: start.elapsed().as_secs_f64() * 1000.0, + }; + set.spawn(async move { + let _permit = permit; + run_one(job).await + }); + if (i + 1) % 500 == 0 { + eprintln!( + " sent {} / {} ({:.0} s)", + i + 1, + rows.len(), + start.elapsed().as_secs_f64() + ); + } + } + let mut results: Vec = Vec::with_capacity(rows.len()); + while let Some(joined) = set.join_next().await { + match joined { + Ok(r) => results.push(r), + Err(e) => eprintln!("task failed: {e}"), + } + } + let wall_s = start.elapsed().as_secs_f64(); + results.sort_by_key(|r| r.row); + if !fleet_set.is_empty() { + let _ = stop_fleet.send(true); + match fleet_set.join_next().await { + Some(Ok(Ok(rows))) => eprintln!("fleet series: {rows} rows in fleet.csv"), + Some(Ok(Err(e))) => eprintln!("fleet series not written: {e}"), + Some(Err(e)) => eprintln!("fleet series task failed: {e}"), + None => {} + } + } + + // Oracle join (optional). + let mut admin_rows: HashMap = HashMap::new(); + if let Some(admin) = &args.admin { + match fetch_admin_records(&client, admin.trim_end_matches('/')).await { + Ok(recs) => { + for r in recs { + admin_rows.insert(r.request_id.clone(), r); + } + } + Err(e) => eprintln!("admin records unavailable: {e}"), + } + for r in &mut results { + if let Some(a) = admin_rows.get(&r.request_id) { + r.oracle_tokens = Some(a.oracle_tokens.max(a.cached_tokens)); + r.queued_ms = Some(a.queued_ms); + if r.worker.is_empty() { + r.worker = a.worker.clone(); + } + if r.cached_tokens == 0 && a.cached_tokens > 0 { + r.cached_tokens = a.cached_tokens; + } + } + } + } + + // Gateway routing decisions: join by request id. + let mut decisions: HashMap = HashMap::new(); + if let Some(log) = &args.gateway_log { + match fs::read_to_string(log) { + Ok(text) => decisions = parse_decisions(&text), + Err(e) => eprintln!("gateway log unreadable: {e}"), + } + for r in &mut results { + let key = if r.gateway_request_id.is_empty() { + request_id_prefix(&r.request_id) + } else { + r.gateway_request_id.clone() + }; + if let Some(d) = decisions.get(&key) { + r.branch.clone_from(&d.branch); + r.credit_tokens = d.credit_tokens(r.prompt_tokens, args.block_size); + r.agree = Some(match r.credit_tokens { + Some(credit) => credit == r.cached_tokens, + None => d.claims_overlap() == (r.cached_tokens > 0), + }); + } + } + } + let mut truth_per_worker: Value = Value::Null; + if let Some(admin) = &args.admin { + match client + .get(format!("{}/admin/truth", admin.trim_end_matches('/'))) + .send() + .await + { + Ok(resp) => truth_per_worker = resp.json().await.unwrap_or(Value::Null), + Err(e) => eprintln!("engine truth unavailable: {e}"), + } + } + + // Statistics over the non-warm-up rows. + let scored: Vec<&ReqResult> = results.iter().filter(|r| r.row >= args.warmup).collect(); + let ok: Vec<&ReqResult> = scored + .iter() + .copied() + .filter(|r| r.status == "ok") + .collect(); + let mut ttft: Vec = ok.iter().filter_map(|r| r.ttft_ms).collect(); + ttft.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal)); + let mut itl_means: Vec = ok.iter().filter_map(|r| r.itl_mean_ms).collect(); + itl_means.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal)); + let mut itl_p99s: Vec = ok.iter().filter_map(|r| r.itl_p99_ms).collect(); + itl_p99s.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal)); + let mut e2e: Vec = ok.iter().map(|r| r.latency_ms).collect(); + e2e.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal)); + let prompt_total: u64 = ok.iter().map(|r| u64::from(r.prompt_tokens)).sum(); + let cached_total: u64 = ok.iter().map(|r| u64::from(r.cached_tokens)).sum(); + let oracle_total: u64 = ok + .iter() + .filter_map(|r| r.oracle_tokens) + .map(u64::from) + .sum(); + let oracle_known = ok.iter().filter(|r| r.oracle_tokens.is_some()).count(); + let within_slo = ok + .iter() + .filter(|r| { + r.ttft_ms.is_some_and(|t| t <= args.slo_ttft_ms) + && r.itl_mean_ms.is_none_or(|m| m <= args.slo_itl_ms) + }) + .count(); + let within_slo_strict = ok + .iter() + .filter(|r| { + r.ttft_ms.is_some_and(|t| t <= args.slo_ttft_ms) + && r.itl_p99_ms.is_none_or(|p| p <= args.slo_itl_ms) + }) + .count(); + let mut per_worker: BTreeMap = BTreeMap::new(); + for r in &ok { + let e = per_worker.entry(r.worker.clone()).or_insert((0, 0)); + e.0 += 1; + e.1 += u64::from(r.prompt_tokens.saturating_sub(r.cached_tokens)); + } + let counts: Vec = per_worker.values().map(|v| v.0 as f64).collect(); + let balance_max_over_mean = if counts.is_empty() { + f64::NAN + } else { + counts.iter().copied().fold(0.0, f64::max) / mean(&counts) + }; + let summary = json!({ + "label": args.label, + "rows": rows.len(), + "scored": scored.len(), + "ok": ok.len(), + "errors": scored.len() - ok.len(), + "invisible_token_requests": ok.iter().filter(|r| r.invisible_tokens > 0).count(), + "speedup": args.speedup, + "wall_s": wall_s, + "req_per_s": ok.len() as f64 / wall_s, + "ttft_ms": {"mean": mean(&ttft), "p50": percentile(&ttft, 0.5), "p90": percentile(&ttft, 0.9), "p99": percentile(&ttft, 0.99)}, + "itl_mean_ms": {"mean": mean(&itl_means), "p90": percentile(&itl_means, 0.9), "p99": percentile(&itl_means, 0.99)}, + "itl_p99_ms_per_request": {"p50": percentile(&itl_p99s, 0.5), "p90": percentile(&itl_p99s, 0.9)}, + "e2e_ms": {"p50": percentile(&e2e, 0.5), "p99": percentile(&e2e, 0.99)}, + "goodput_req_per_s": within_slo as f64 / wall_s, + "within_slo_fraction": if ok.is_empty() { f64::NAN } else { within_slo as f64 / ok.len() as f64 }, + "within_slo_strict_fraction": if ok.is_empty() { f64::NAN } else { within_slo_strict as f64 / ok.len() as f64 }, + "prefix_reuse": if prompt_total == 0 { f64::NAN } else { cached_total as f64 / prompt_total as f64 }, + "oracle_prefix_reuse": if prompt_total == 0 || oracle_known == 0 { f64::NAN } else { oracle_total as f64 / prompt_total as f64 }, + "hit_over_oracle": if oracle_total == 0 { f64::NAN } else { cached_total as f64 / oracle_total as f64 }, + "oracle_known": oracle_known, + "per_worker_requests": per_worker.iter().map(|(w, v)| json!({"worker": w, "requests": v.0, "uncached_prompt_tokens": v.1})).collect::>(), + "t4": decision_summary(&ok, decisions.len()), + "engine_truth_per_worker": truth_per_worker.get("workers").cloned().unwrap_or(Value::Array(Vec::new())), + "balance_max_over_mean": balance_max_over_mean, + "slo": {"ttft_ms": args.slo_ttft_ms, "itl_ms": args.slo_itl_ms, "itl_metric": "per-request mean (strict variant: per-request p99)"}, + }); + fs::create_dir_all(&args.out)?; + fs::write( + args.out.join("summary.json"), + serde_json::to_string_pretty(&summary)?, + )?; + let mut csv = String::from("row,trace_ts_ms,trace_input_length,sent_at_ms,status,request_id,worker,prompt_tokens,completion_tokens,cached_tokens,oracle_tokens,queued_ms,ttft_ms,latency_ms,itl_mean_ms,itl_p99_ms,tokens_seen,invisible_tokens,gateway_request_id,branch,credit_tokens,agree\n"); + for r in &results { + csv.push_str(&format!( + "{},{},{},{:.1},{},{},{},{},{},{},{},{},{},{:.1},{},{},{},{},{},{},{},{}\n", + r.row, + r.trace_ts_ms, + r.trace_input_length, + r.sent_at_ms, + r.status.replace(',', ";"), + r.request_id, + r.worker, + r.prompt_tokens, + r.completion_tokens, + r.cached_tokens, + r.oracle_tokens.map(|v| v.to_string()).unwrap_or_default(), + r.queued_ms.map(|v| format!("{v:.1}")).unwrap_or_default(), + r.ttft_ms.map(|v| format!("{v:.1}")).unwrap_or_default(), + r.latency_ms, + r.itl_mean_ms.map(|v| format!("{v:.2}")).unwrap_or_default(), + r.itl_p99_ms.map(|v| format!("{v:.2}")).unwrap_or_default(), + r.tokens_seen, + r.invisible_tokens, + r.gateway_request_id, + r.branch, + r.credit_tokens.map(|v| v.to_string()).unwrap_or_default(), + r.agree.map(|v| v.to_string()).unwrap_or_default() + )); + } + fs::write(args.out.join("requests.csv"), csv)?; + fs::write( + args.out.join("t4.md"), + decision_table(&results, args.t4_rows), + )?; + println!("{}", serde_json::to_string_pretty(&summary)?); + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + + /// A one-endpoint HTTP/1.1 server answering every request with the given + /// body (what `GET /admin/fleet` returns), for the fleet sampler test. + async fn serve_json(body: &'static str) -> (String, JoinSet<()>) { + use tokio::io::{AsyncReadExt, AsyncWriteExt}; + let listener = tokio::net::TcpListener::bind("127.0.0.1:0") + .await + .expect("bind"); + let addr = listener.local_addr().expect("addr"); + let mut server: JoinSet<()> = JoinSet::new(); + server.spawn(async move { + let mut conns: JoinSet<()> = JoinSet::new(); + loop { + let Ok((mut sock, _)) = listener.accept().await else { + break; + }; + conns.spawn(async move { + let mut buf = vec![0u8; 4096]; + let mut n = 0; + while n < buf.len() { + let Ok(k) = sock.read(&mut buf[n..]).await else { + return; + }; + if k == 0 { + return; + } + n += k; + if buf[..n].windows(4).any(|w| w == b"\r\n\r\n") { + break; + } + } + let resp = format!( + "HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}", + body.len(), + body + ); + let _ = sock.write_all(resp.as_bytes()).await; + let _ = sock.shutdown().await; + }); + } + }); + (format!("http://{addr}"), server) + } + + #[tokio::test] + async fn fleet_series_has_a_header_and_one_row_per_worker_per_second() { + let (admin, _server) = serve_json( + r#"{"workers":[{"worker":"grpc:1","num_running_reqs":2,"num_waiting_reqs":1,"num_waiting_uncached_tokens":300,"token_usage":0.25,"num_cached_blocks":40,"num_preemptions":0,"num_kv_batches":9},{"worker":"grpc:2","num_running_reqs":0,"num_waiting_reqs":0,"token_usage":0.0,"num_cached_blocks":0,"num_preemptions":0,"num_kv_batches":0}]}"#, + ) + .await; + let dir = std::env::temp_dir().join(format!("replay-fleet-{}", std::process::id())); + let path = dir.join("fleet.csv"); + let (stop, rx) = watch::channel(false); + // The test server is local: no proxy from the environment. + let client = reqwest::Client::builder() + .no_proxy() + .build() + .expect("client"); + let mut sampler: JoinSet> = JoinSet::new(); + sampler.spawn(sample_fleet( + client, + admin, + path.clone(), + Instant::now(), + rx, + )); + tokio::time::sleep(Duration::from_millis(2600)).await; + let _ = stop.send(true); + let rows = sampler + .join_next() + .await + .expect("a task") + .expect("join") + .expect("sampler"); + let text = fs::read_to_string(&path).expect("fleet.csv"); + let mut lines = text.lines(); + assert_eq!(lines.next(), Some(FLEET_CSV_HEADER)); + let data: Vec> = lines.map(|l| l.split(',').collect()).collect(); + assert_eq!(data.len(), rows); + // ticks at 0, 1 and 2 s inside 2.6 s, two workers each + assert!((4..=8).contains(&rows) && rows % 2 == 0, "rows {rows}"); + assert_eq!(data[0][1], "grpc:1"); + assert_eq!(data[1][1], "grpc:2"); + assert_eq!(&data[0][2..], &["2", "1", "300", "0.2500", "40", "0", "9"]); + let t: Vec = data + .iter() + .step_by(2) + .map(|r| r[0].parse().expect("t_s")) + .collect(); + for w in t.windows(2) { + assert!( + (w[1] - w[0] - 1.0).abs() < 0.3, + "one tick per second: {t:?}" + ); + } + let _ = fs::remove_file(&path); + let _ = fs::remove_dir(&dir); + } + + #[test] + fn block_text_is_deterministic_per_hash_id() { + assert_eq!(block_text(42, 7, 32), block_text(42, 7, 32)); + assert_ne!(block_text(42, 7, 32), block_text(43, 7, 32)); + assert_ne!(block_text(42, 7, 32), block_text(42, 8, 32)); + assert_eq!(block_text(1, 7, 32).split(' ').count(), 32); + } + + #[test] + fn prompts_share_prefixes_when_rows_share_hash_ids() { + let a = TraceRow { + timestamp: 0, + input_length: 1024, + output_length: 1, + hash_ids: vec![0, 1], + }; + let b = TraceRow { + timestamp: 0, + input_length: 1024, + output_length: 1, + hash_ids: vec![0, 2], + }; + let pa = prompt_for(&a, 7, 16); + let pb = prompt_for(&b, 7, 16); + let shared = block_text(0, 7, 16); + assert!(pa.starts_with(&shared) && pb.starts_with(&shared)); + assert_ne!(pa, pb); + } + + #[test] + fn percentile_and_mean() { + let v = [1.0, 2.0, 3.0, 4.0, 5.0]; + assert_eq!(percentile(&v, 0.5), 3.0); + assert_eq!(percentile(&v, 0.0), 1.0); + assert_eq!(percentile(&v, 1.0), 5.0); + assert_eq!(mean(&v), 3.0); + assert!(percentile(&[], 0.5).is_nan()); + } + + #[test] + fn a_finish_only_stream_counts_as_served() { + // Captured from the gateway: a one-token answer whose token is an + // incomplete UTF-8 piece, so no chunk carries text; the finish chunk + // and the usage still arrive. + let finish: Value = serde_json::from_str( + r#"{"id":"chatcmpl-x","object":"chat.completion.chunk","created":1,"model":"m","system_fingerprint":"grpc:19600","choices":[{"index":0,"delta":{"reasoning_content":null},"logprobs":null,"finish_reason":"length"}]}"#, + ) + .unwrap(); + let usage: Value = serde_json::from_str( + r#"{"id":"chatcmpl-x","object":"chat.completion.chunk","created":1,"model":"m","system_fingerprint":"grpc:19600","choices":[],"usage":{"prompt_tokens":34,"completion_tokens":1,"total_tokens":35,"prompt_tokens_details":{"cached_tokens":16}}}"#, + ) + .unwrap(); + let started = Instant::now(); + let mut observer = StreamObserver::default(); + observer.observe(&finish, started + Duration::from_millis(40)); + observer.observe(&usage, started + Duration::from_millis(41)); + let mut result = ReqResult { + status: "ok".to_string(), + ..Default::default() + }; + observer.finish(started, &mut result); + assert_eq!(result.status, "ok"); + assert_eq!(result.request_id, "chatcmpl-x"); + assert_eq!(result.worker, "grpc:19600"); + assert_eq!(result.completion_tokens, 1); + assert_eq!(result.cached_tokens, 16); + assert_eq!(result.tokens_seen, 0); + assert_eq!(result.invisible_tokens, 1); + assert!((result.ttft_ms.unwrap() - 40.0).abs() < 1.0); + } + + #[test] + fn a_stream_with_no_output_at_all_is_no_tokens() { + let usage: Value = serde_json::from_str( + r#"{"id":"chatcmpl-y","choices":[],"usage":{"prompt_tokens":3,"completion_tokens":0,"total_tokens":3}}"#, + ) + .unwrap(); + let started = Instant::now(); + let mut observer = StreamObserver::default(); + observer.observe(&usage, started + Duration::from_millis(5)); + let mut result = ReqResult { + status: "ok".to_string(), + ..Default::default() + }; + observer.finish(started, &mut result); + assert_eq!(result.status, "no-tokens"); + assert!(result.ttft_ms.is_none()); + } + + #[test] + fn visible_tokens_give_ttft_and_itl() { + let tok = |s: &str| -> Value { + serde_json::from_str(&format!( + r#"{{"id":"chatcmpl-z","choices":[{{"index":0,"delta":{{"content":"{s}"}},"finish_reason":null}}]}}"# + )) + .unwrap() + }; + let started = Instant::now(); + let mut observer = StreamObserver::default(); + observer.observe(&tok("a"), started + Duration::from_millis(100)); + observer.observe(&tok("b"), started + Duration::from_millis(120)); + observer.observe(&tok("c"), started + Duration::from_millis(150)); + let mut result = ReqResult { + status: "ok".to_string(), + ..Default::default() + }; + observer.finish(started, &mut result); + assert_eq!(result.tokens_seen, 3); + assert!((result.ttft_ms.unwrap() - 100.0).abs() < 1.0); + assert!((result.itl_mean_ms.unwrap() - 25.0).abs() < 1.0); + assert_eq!(result.invisible_tokens, 0); + // No usage arrived: the stream is still a served one, counted as seen. + assert_eq!(result.status, "ok"); + assert_eq!(result.completion_tokens, 3); + } + + #[test] + fn gateway_decisions_are_parsed_and_joined_by_request_id() { + let log = concat!( + "2026-10-05 10:23:43 DEBUG http_request{method=POST uri=/v1/chat/completions version=HTTP/1.1 module=\"smg\" request_id=\"chatcmpl-nukU\"}: smg::policies::cache_aware: cache_aware.rs:1681: Cache-aware selection index=\"tree\" branch=\"expected_wait_fallback\" worker=\"grpc://127.0.0.1:19611\" model_id=\"mock-model\" matched_ratio=0.0 threshold=0.30000001192092896\n", + "2026-10-05 10:23:46 DEBUG http_request{method=POST uri=/v1/chat/completions version=HTTP/1.1 module=\"smg\" request_id=\"chatcmpl-UNiu\"}: smg::policies::cache_aware: cache_aware.rs:1493: Event-driven routing: overlap match worker=\"grpc://127.0.0.1:19611\" branch=\"event_hit\" model_id=\"mock-model\"\n", + "2026-10-05 10:23:48 DEBUG http_request{request_id=\"chatcmpl-fut\"}: smg::policies::cache_aware: Event-driven routing: overlap match worker=\"grpc://127.0.0.1:19610\" branch=\"event_hit\" overlap_blocks=12 model_id=\"mock-model\"\n", + "2026-10-05 10:23:49 DEBUG http_request{request_id=\"chatcmpl-none\"}: smg::policies::cache_aware: Event-driven routing: no overlap, expected-wait fallback worker=\"grpc://127.0.0.1:19610\" model_id=\"mock-model\"\n", + "2026-10-05 10:23:50 DEBUG some other line request_id=\"x\" branch=\"nope\"\n", + // As the gateway really writes it: ANSI colour sequences around every field. + "\u{1b}[2m2026-10-05 10:59:48\u{1b}[0m \u{1b}[34mDEBUG\u{1b}[0m \u{1b}[1mhttp_request\u{1b}[0m\u{1b}[1m{\u{1b}[0m\u{1b}[3mrequest_id\u{1b}[0m\u{1b}[2m=\u{1b}[0m\"chatcmpl-7lE1\"\u{1b}[1m}\u{1b}[0m\u{1b}[2m:\u{1b}[0m Event-driven routing: overlap match \u{1b}[3mworker\u{1b}[0m\u{1b}[2m=\u{1b}[0m\"grpc://127.0.0.1:19702\" \u{1b}[3mbranch\u{1b}[0m\u{1b}[2m=\u{1b}[0m\"event_spill\"\n", + ); + let d = parse_decisions(log); + assert_eq!(d.len(), 5); + assert_eq!(d["chatcmpl-7lE1"].branch, "event_spill"); + assert_eq!(d["chatcmpl-7lE1"].worker, "grpc://127.0.0.1:19702"); + assert_eq!(d["chatcmpl-nukU"].branch, "expected_wait_fallback"); + assert_eq!(d["chatcmpl-nukU"].matched_ratio, Some(0.0)); + assert_eq!(d["chatcmpl-nukU"].credit_tokens(1000, 16), Some(0)); + assert_eq!(d["chatcmpl-UNiu"].branch, "event_hit"); + assert!(d["chatcmpl-UNiu"].claims_overlap()); + assert_eq!(d["chatcmpl-UNiu"].credit_tokens(1000, 16), None); + assert_eq!(d["chatcmpl-fut"].credit_tokens(1000, 16), Some(192)); + assert_eq!(d["chatcmpl-none"].branch, "expected_wait_fallback"); + assert!(!d["chatcmpl-none"].claims_overlap()); + assert_eq!( + request_id_prefix( + "chatcmpl-nukUelurU5QeHF4GhOUknktv-01a10b97-428b-7f72-9728-a08fad5787ae" + ), + "chatcmpl-nukUelurU5QeHF4GhOUknktv" + ); + } + + #[test] + fn decision_table_has_the_expected_columns() { + let mut r = ReqResult { + row: 3, + status: "ok".to_string(), + worker: "grpc:19500".to_string(), + branch: "event_hit".to_string(), + prompt_tokens: 640, + cached_tokens: 512, + agree: Some(true), + ..Default::default() + }; + let table = decision_table(std::slice::from_ref(&r), 10); + assert!(table.starts_with("| phase | idx | worker | branch | prompt_tokens | engine cached_tokens | implied overlap | agree |")); + assert!(table.contains("| replay | 3 | 19500 | event_hit | 640 | 512 | - | true |")); + assert!(table.contains("agreement: 1/1")); + r.credit_tokens = Some(512); + let table = decision_table(std::slice::from_ref(&r), 10); + assert!(table.contains("| 640 | 512 | 512 | true |")); + } + + #[test] + fn a_row_stamped_before_the_first_is_due_at_once() { + // 3 s of trace at 2x is 1.5 s of wall time. + assert_eq!(due_after(3_010, 10, 2.0), Duration::from_millis(1_500)); + assert_eq!(due_after(10, 10, 2.0), Duration::ZERO); + // An earlier stamp (unsorted trace, or a window starting on a + // reordered row) saturates to "now" instead of wrapping around. + assert_eq!(due_after(5, 10, 2.0), Duration::ZERO); + assert_eq!(due_after(0, u64::MAX, 0.5), Duration::ZERO); + // A due time beyond what a Duration holds is clamped, not a panic. + assert_eq!(due_after(u64::MAX, 0, 1e-300), Duration::MAX); + } + + #[test] + fn speedup_must_be_finite_and_positive() { + assert_eq!(parse_speedup("2"), Ok(2.0)); + assert_eq!(parse_speedup(" 0.5 "), Ok(0.5)); + for bad in ["0", "-1", "inf", "-inf", "NaN", "fast", ""] { + assert!(parse_speedup(bad).is_err(), "{bad:?} must be rejected"); + } + // The flag rejects them at parsing, before any request is sent. + for bad in ["0", "-2", "inf", "nan"] { + assert!( + Args::try_parse_from(["replay", "--trace", "t.jsonl", "--speedup", bad]).is_err(), + "--speedup {bad} must be rejected" + ); + } + let ok = Args::try_parse_from(["replay", "--trace", "t.jsonl", "--speedup", "4"]) + .expect("a valid speedup parses"); + assert_eq!(ok.speedup, 4.0); + } + + #[test] + fn sse_event_boundary() { + assert_eq!(find_double_newline(b"data: x\n\ndata: y"), Some(7)); + assert_eq!(find_double_newline(b"data: x\n"), None); + } +} diff --git a/crates/mock_worker/src/config.rs b/crates/mock_worker/src/config.rs index 8140db7ba4..10149afb06 100644 --- a/crates/mock_worker/src/config.rs +++ b/crates/mock_worker/src/config.rs @@ -2,7 +2,7 @@ use std::{path::PathBuf, time::Duration}; -use crate::engine::EngineParams; +use crate::engine::{Calibration, EngineParams, LoadsLike, TimingModel}; /// Configuration shared by every mocked HTTP and gRPC worker in the process. #[derive(Debug, Clone)] @@ -38,6 +38,26 @@ pub struct Config { pub realistic: bool, /// Engine-simulator parameters (only used when `realistic`). pub engine: EngineParams, + /// Port of the process-wide admin API (fleet, request records, cache + /// dumps, resets); off when `None`. + pub admin_port: Option, + /// Context length advertised to the gateway (`max_context_length`, + /// `max_req_input_len`, `max_model_len`). + pub context_length: u32, + /// vLLM-wire ZMQ KV-event publishers (realistic engines only): the first + /// PUB port; worker `i` publishes on `base + 2i` and answers replay on + /// `base + 2i + 1`. Off when `None`. + pub kv_events_zmq_base_port: Option, + /// Whether the publishers serve replay requests (ROUTER at `port + 1`). + pub kv_events_replay: bool, + /// Topic frame of every published message (vLLM's default is empty). + pub kv_events_topic: String, + /// Batches each publisher keeps for replay. + pub kv_events_buffer_steps: usize, + /// Which engine's wire the publishers speak. + pub kv_events_wire: crate::kv_zmq::Wire, + /// Which backend's load report the workers imitate (`GetLoads`, `/v1/loads`). + pub loads_like: LoadsLike, /// Settings for replay testing (gRPC workers only). pub replay: ReplayConfig, } @@ -52,16 +72,71 @@ pub struct ReplayConfig { pub capture: Option, } -impl Config { - /// Parse the configuration from `std::env::args`, falling back to defaults. - pub fn from_args() -> Result { - Self::parse(std::env::args().skip(1)) +/// Timing flags collected while parsing; resolved into one [`TimingModel`] +/// at the end so flag order does not matter. +#[derive(Default)] +struct TimingFlags { + /// `polynomial`, `linear` or `fit:`. + kind: Option, + prefill_poly: Option<[f64; 3]>, + decode_poly: Option<[f64; 3]>, + prefill_tps: Option, + decode_base_ms: Option, + decode_per_req_ms: Option, +} + +impl TimingFlags { + /// The model, and the calibration it came from (for its capacity and overhead). + fn resolve(self) -> Result<(TimingModel, Option), String> { + if let Some(path) = self.kind.as_deref().and_then(|k| k.strip_prefix("fit:")) { + let calibration = Calibration::load(path)?; + return Ok((TimingModel::fitted(&calibration), Some(calibration))); + } + let linear_override = self.prefill_tps.is_some() + || self.decode_base_ms.is_some() + || self.decode_per_req_ms.is_some(); + let kind = match self.kind.as_deref() { + Some("polynomial") => "polynomial", + Some("linear") => "linear", + Some(other) => { + return Err(format!( + "--timing must be polynomial|linear|fit:, got {other}" + )) + } + None if linear_override => "linear", + None => "polynomial", + }; + let model = if kind == "linear" { + TimingModel::Linear { + prefill_tps: self.prefill_tps.unwrap_or(8000.0), + decode_base_ms: self.decode_base_ms.unwrap_or(6.0), + decode_per_req_ms: self.decode_per_req_ms.unwrap_or(0.35), + } + } else { + TimingModel::Polynomial { + prefill: self.prefill_poly.unwrap_or(TimingModel::POLY_PREFILL), + decode: self.decode_poly.unwrap_or(TimingModel::POLY_DECODE), + } + }; + Ok((model, None)) } +} - /// Parse the configuration from command-line flags, without the program - /// name. - fn parse(args: impl IntoIterator) -> Result { - let mut cfg = Self { +fn parse_poly(raw: String, flag: &str) -> Result<[f64; 3], String> { + let parts: Vec = raw + .split(',') + .map(|v| v.trim().parse::()) + .collect::>() + .map_err(|_| format!("invalid value for {flag}: {raw} (want a,b,c)"))?; + match parts.as_slice() { + [a, b, c] => Ok([*a, *b, *c]), + _ => Err(format!("invalid value for {flag}: {raw} (want a,b,c)")), + } +} + +impl Default for Config { + fn default() -> Self { + Self { host: "127.0.0.1".to_string(), http_base_port: 9000, http_count: 0, @@ -76,8 +151,32 @@ impl Config { output_tokens: 8, realistic: false, engine: EngineParams::default(), + admin_port: None, + context_length: 32768, + kv_events_zmq_base_port: None, + kv_events_replay: true, + kv_events_topic: String::new(), + kv_events_buffer_steps: 10_000, + kv_events_wire: crate::kv_zmq::Wire::Vllm, + loads_like: LoadsLike::Mock, replay: ReplayConfig::default(), - }; + } + } +} + +impl Config { + /// Parse the configuration from `std::env::args`, falling back to defaults. + pub fn from_args() -> Result { + Self::parse(std::env::args().skip(1)) + } + + /// Parse the configuration from command-line flags, without the program + /// name. + fn parse(args: impl IntoIterator) -> Result { + let mut cfg = Self::default(); + let mut timing = TimingFlags::default(); + let mut kv_blocks: Option = None; + let (mut kv_tokens_given, mut block_size_given, mut overhead_given) = (false, false, false); let mut args = args.into_iter(); while let Some(flag) = args.next() { @@ -108,21 +207,61 @@ impl Config { } } } - "--prefill-tps" => cfg.engine.prefill_tps = parse(value(&mut args, &flag)?, &flag)?, + "--timing" => timing.kind = Some(value(&mut args, &flag)?), + "--prefill-poly" => { + timing.prefill_poly = Some(parse_poly(value(&mut args, &flag)?, &flag)?); + } + "--decode-poly" => { + timing.decode_poly = Some(parse_poly(value(&mut args, &flag)?, &flag)?); + } + "--prefill-tps" => { + timing.prefill_tps = Some(parse(value(&mut args, &flag)?, &flag)?) + } "--decode-base-ms" => { - cfg.engine.decode_base_ms = parse(value(&mut args, &flag)?, &flag)?; + timing.decode_base_ms = Some(parse(value(&mut args, &flag)?, &flag)?); } "--decode-per-req-ms" => { - cfg.engine.decode_per_req_ms = parse(value(&mut args, &flag)?, &flag)?; + timing.decode_per_req_ms = Some(parse(value(&mut args, &flag)?, &flag)?); } - "--prefill-chunk" => { - cfg.engine.prefill_chunk_tokens = parse(value(&mut args, &flag)?, &flag)?; + "--max-batched-tokens" => { + cfg.engine.max_batched_tokens = parse(value(&mut args, &flag)?, &flag)?; } "--max-running" => cfg.engine.max_running = parse(value(&mut args, &flag)?, &flag)?, "--kv-tokens" => { cfg.engine.kv_capacity_tokens = parse(value(&mut args, &flag)?, &flag)?; + kv_tokens_given = true; + } + "--kv-blocks" => kv_blocks = Some(parse(value(&mut args, &flag)?, &flag)?), + "--request-overhead-ms" => { + cfg.engine.request_overhead_ms = parse(value(&mut args, &flag)?, &flag)?; + overhead_given = true; + } + "--prefill-first" => { + cfg.engine.prefill_first = parse(value(&mut args, &flag)?, &flag)?; } - "--block-size" => cfg.engine.block_size = parse(value(&mut args, &flag)?, &flag)?, + "--reserve-full-isl" => { + cfg.engine.reserve_full_isl = parse(value(&mut args, &flag)?, &flag)?; + } + "--context-length" => { + cfg.context_length = parse(value(&mut args, &flag)?, &flag)?; + } + "--kv-events-zmq-base-port" => { + cfg.kv_events_zmq_base_port = Some(parse(value(&mut args, &flag)?, &flag)?); + } + "--kv-events-replay" => { + cfg.kv_events_replay = parse(value(&mut args, &flag)?, &flag)?; + } + "--kv-events-topic" => cfg.kv_events_topic = value(&mut args, &flag)?, + "--kv-events-buffer-steps" => { + cfg.kv_events_buffer_steps = parse(value(&mut args, &flag)?, &flag)?; + } + "--kv-events-wire" => cfg.kv_events_wire = value(&mut args, &flag)?.parse()?, + "--loads-like" => cfg.loads_like = value(&mut args, &flag)?.parse()?, + "--block-size" => { + cfg.engine.block_size = parse(value(&mut args, &flag)?, &flag)?; + block_size_given = true; + } + "--admin-port" => cfg.admin_port = Some(parse(value(&mut args, &flag)?, &flag)?), "--prefix-cache" => { cfg.engine.prefix_cache = parse(value(&mut args, &flag)?, &flag)? } @@ -134,6 +273,30 @@ impl Config { if cfg.tokenizer_path.is_empty() { cfg.tokenizer_path = cfg.model_id.clone(); } + let (model, calibration) = timing.resolve()?; + cfg.engine.timing = model; + if let Some(c) = calibration { + // The calibration's block size, capacity and overhead apply unless + // the flags say otherwise. + if let (Some(bs), false) = (c.block_size, block_size_given) { + cfg.engine.block_size = bs; + } + if let (Some(tokens), false, None) = (c.kv_capacity_tokens, kv_tokens_given, kv_blocks) + { + cfg.engine.kv_capacity_tokens = tokens; + } + // The decode fit was measured against the calibration's pool: a + // smaller pool from the flags changes room, not step time. + if c.kv_capacity_tokens.is_some() { + cfg.engine.decode_reference_tokens = c.kv_capacity_tokens; + } + if !overhead_given { + cfg.engine.request_overhead_ms = c.request_overhead_ms; + } + } + if let Some(blocks) = kv_blocks { + cfg.engine.kv_capacity_tokens = blocks * u64::from(cfg.engine.block_size); + } // `--output-tokens` doubles as the realistic engine's default output // length when a request omits `max_tokens`. cfg.engine.max_new_default = cfg.output_tokens; @@ -153,6 +316,33 @@ impl Config { } } +impl Config { + /// The KV-event publisher of worker number `index` (gRPC workers first, + /// then ZMQ ranks), when publishing is on and the engine is realistic. + /// `dp_rank` is the rank the publisher stamps on the vLLM wire: 0 for a + /// gRPC worker, the engine index a ZMQ rank advertises. + pub(crate) fn kv_zmq_for( + &self, + index: u16, + dp_rank: i32, + ) -> Option { + if !self.realistic { + return None; + } + let base = self.kv_events_zmq_base_port?; + let port = base.checked_add(index.checked_mul(2)?)?; + Some(crate::kv_zmq::KvZmqConfig { + host: self.host.clone(), + port, + replay: self.kv_events_replay, + topic: self.kv_events_topic.clone(), + buffer_steps: self.kv_events_buffer_steps, + dp_rank, + wire: self.kv_events_wire, + }) + } +} + fn value(args: &mut impl Iterator, flag: &str) -> Result { args.next() .ok_or_else(|| format!("missing value for {flag}")) @@ -180,16 +370,39 @@ fn usage() -> String { --output-tokens output tokens per request when unspecified (default 8)\n\ --capture append each gRPC Generate request to as a JSON line\n\ \n\ - Realistic engine simulator (continuous batching; opt-in):\n\ + Realistic engine simulator (vLLM pass loop over a block-level KV pool; opt-in):\n\ --engine engine mode (default canned)\n\ - --prefill-tps prefill throughput, tokens/sec (default 8000)\n\ - --decode-base-ms fixed decode-step latency, ms (default 6.0)\n\ - --decode-per-req-ms added decode-step latency per running req (default 0.35)\n\ - --prefill-chunk max prompt tokens prefilled per step (default 2048)\n\ - --max-running max concurrent running requests (default 256)\n\ + --timing > pass duration model (default polynomial);\n\ + fit: reads a hardware calibration JSON (polynomials, KV capacity,\n\ + block size, per-request overhead; flags given explicitly win)\n\ + --prefill-poly prefill ms = a + b*T + c*T^2 over uncached tokens T in the pass\n\ + (default 16.50142,1.518344e-2,4.209989e-7)\n\ + --decode-poly decode ms = max(1, a + b*u + c*u^2) over KV utilisation u\n\ + (default 5.74,54.01,-25.74)\n\ + --prefill-tps linear model: prefill tokens/sec (default 8000; selects linear)\n\ + --decode-base-ms linear model: fixed decode-step ms (default 6.0)\n\ + --decode-per-req-ms linear model: decode ms per running request (default 0.35)\n\ + --max-batched-tokens token budget per pass (default 8192)\n\ + --max-running max sequences per pass (default 256)\n\ --kv-tokens KV cache capacity in tokens (default 524288)\n\ + --kv-blocks KV cache capacity in blocks (overrides --kv-tokens)\n\ + --request-overhead-ms fixed per-request latency added to every event of a stream (default 0)\n\ --block-size cache block/page size in tokens (default 16)\n\ - --prefix-cache enable prefix caching + KV events (default true)" + --prefix-cache enable prefix caching + KV events (default true)\n\ + --prefill-first SGLang-style: a pass with prefill runs prefill only (default false)\n\ + --reserve-full-isl admit a prompt only with KV room for all of it beyond its cached\n\ + blocks, head-of-line (vLLM's scheduler_reserve_full_isl; default true)\n\ + --context-length advertised context length (default 32768)\n\ + --kv-events-zmq-base-port vLLM-wire ZMQ KV-event publishers: worker i publishes\n\ + on base+2i (PUB) and replays on base+2i+1 (ROUTER) (default off)\n\ + --kv-events-replay serve replay requests on the ROUTER port (default true)\n\ + --kv-events-topic topic frame (default empty, as vLLM)\n\ + --kv-events-buffer-steps batches kept for replay (default 10000)\n\ + --kv-events-wire which engine's publisher to imitate (default vllm)\n\ + --loads-like load report: everything the simulator knows, or only what the\n\ + vLLM servicer reports (running, waiting, token_usage, maxima)\n\ + --admin-port process-wide admin API: fleet, request records with the\n\ + arrival-time oracle, cache dumps, resets (default off)" .to_string() } @@ -201,6 +414,27 @@ mod tests { Config::parse(flags.iter().map(|flag| (*flag).to_string())) } + #[test] + fn publishers_carry_their_worker_rank_and_port() { + let cfg = Config { + realistic: true, + kv_events_zmq_base_port: Some(19_000), + ..Config::default() + }; + // A gRPC worker publishes as rank 0 on base + 2 * index. + let grpc = cfg + .kv_zmq_for(1, 0) + .expect("publisher for a realistic engine"); + assert_eq!((grpc.port, grpc.dp_rank), (19_002, 0)); + // A ZMQ rank publishes under the rank it advertises, not rank 0. + let rank = cfg.kv_zmq_for(3, 2).expect("publisher for a ZMQ rank"); + assert_eq!((rank.port, rank.dp_rank), (19_006, 2)); + assert!( + Config::default().kv_zmq_for(0, 0).is_none(), + "no publisher for a canned worker" + ); + } + #[test] fn capture_flag_sets_the_capture_path() { let grpc = ["--grpc-base-port", "19000", "--grpc-count", "1"]; diff --git a/crates/mock_worker/src/engine.rs b/crates/mock_worker/src/engine.rs index 9da7cd2040..34240c8508 100644 --- a/crates/mock_worker/src/engine.rs +++ b/crates/mock_worker/src/engine.rs @@ -6,8 +6,9 @@ //! //! - **prefill latency that scales with input length** — time-to-first-token //! grows with the (uncached) prompt size, chunked across scheduler steps; -//! - **decode latency that grows with batch size** — inter-token latency is -//! `base + slope · batch`, so a busy replica is slower per token; +//! - **decode latency that grows with load** — a decode pass costs more as +//! the decoding requests' KV utilisation rises (with the linear model, as +//! the batch widens), so a busy replica is slower per token; //! - **finite KV capacity with admission/queueing** — when KV is full requests //! wait, producing the `num_waiting_uncached_tokens` signal `least_load` uses; //! - **prefix caching** — a request sharing a prefix with cached blocks pays @@ -17,53 +18,406 @@ //! The engine is an actor: one `tokio` task per virtual worker owns all mutable //! state and advances it in real wall-clock time, but the per-step work is plain //! arithmetic. Idle engines block on their request channel, so a fleet of mostly -//! idle workers stays cheap. The scheduling math lives in [`SchedulerState::step`], -//! a pure synchronous function returning the work done plus the time it took — +//! idle workers stays cheap. The scheduling math lives in `SchedulerState::step`, +//! a pure synchronous function returning the work done plus the time it took, //! which makes it deterministically unit-testable with no real timers. use std::{ collections::{BTreeSet, HashMap, HashSet, VecDeque}, pin::Pin, - sync::{Arc, Mutex, RwLock}, - time::Duration, + sync::{ + atomic::{AtomicBool, AtomicU32, AtomicU64, Ordering}, + Arc, Mutex, OnceLock, RwLock, + }, + time::{Duration, Instant, SystemTime, UNIX_EPOCH}, }; use futures::{stream, Stream, StreamExt}; use smg_grpc_client::common_proto as common; -use tokio::sync::{broadcast, mpsc}; +use tokio::sync::{broadcast, mpsc, Notify}; use tonic::Status; // --------------------------------------------------------------------------- // Configuration // --------------------------------------------------------------------------- -/// Tunable parameters of the simulated engine. Defaults approximate a single -/// mid-size model replica on one accelerator; override via mock-worker flags. +/// How a pass's duration follows from the work it contains. +#[derive(Clone, Debug, PartialEq)] +pub enum TimingModel { + /// Polynomials over the pass: prefill `a + b·T + c·T²` ms with `T` the + /// uncached tokens it processes, decode `max(1, a + b·u + c·u²)` ms with + /// `u` the KV utilisation of the decoding requests (their context tokens + /// over capacity). The defaults are AISimulate's uncalibrated baseline, so + /// results compare with other simulators built on the same polynomials. + Polynomial { prefill: [f64; 3], decode: [f64; 3] }, + /// Prefill at a fixed token rate; decode `base + per_req × batch` ms. + Linear { + prefill_tps: f64, + decode_base_ms: f64, + decode_per_req_ms: f64, + }, + /// A hardware calibration: a single request's prefill from a measured + /// point table (piecewise-linear over `(tokens, ms)`, extrapolated with + /// the last slope); a pass that prefills several requests costs the + /// pass-total form `intercept + slope × total uncached tokens in the + /// pass`, never less than the table value of its largest request (the + /// per-request plateau is the floor). Decode as the polynomial. + Calibrated { + table: Vec<(f64, f64)>, + pass_intercept_ms: f64, + pass_ms_per_token: f64, + decode: [f64; 3], + }, +} + +impl TimingModel { + /// AISimulate's prefill polynomial (ms over uncached tokens in the pass). + pub const POLY_PREFILL: [f64; 3] = [16.501_42, 1.518_344e-2, 4.209_989e-7]; + /// AISimulate's decode polynomial (ms over KV utilisation). + pub const POLY_DECODE: [f64; 3] = [5.74, 54.01, -25.74]; + + pub fn polynomial() -> Self { + Self::Polynomial { + prefill: Self::POLY_PREFILL, + decode: Self::POLY_DECODE, + } + } + + pub fn linear() -> Self { + Self::Linear { + prefill_tps: 8000.0, + decode_base_ms: 6.0, + decode_per_req_ms: 0.35, + } + } + + /// A calibration file's model (`Calibration::load`): the point table and + /// pass-total form when the file has them, else its polynomials. + pub(crate) fn fitted(c: &Calibration) -> Self { + match (&c.prefill_table, c.prefill_pass) { + (Some(table), pass) if table.len() >= 2 => { + let (pass_intercept_ms, pass_ms_per_token) = pass.unwrap_or_else(|| { + // No pass form: the table's last slope, from its first value. + let n = table.len(); + let slope = + (table[n - 1].1 - table[n - 2].1) / (table[n - 1].0 - table[n - 2].0); + (table[0].1, slope.max(0.0)) + }); + Self::Calibrated { + table: table.clone(), + pass_intercept_ms, + pass_ms_per_token, + decode: c.decode, + } + } + _ => Self::Polynomial { + prefill: c.prefill, + decode: c.decode, + }, + } + } + + /// Piecewise-linear value of a `(tokens, ms)` table at `tokens`: the + /// first point's value below the table, the last segment's slope (never + /// negative) above it. + fn table_ms(table: &[(f64, f64)], tokens: f64) -> f64 { + let Some(first) = table.first() else { + return 0.0; + }; + if tokens <= first.0 || table.len() == 1 { + return first.1; + } + for w in table.windows(2) { + let ((x0, y0), (x1, y1)) = (w[0], w[1]); + if tokens <= x1 { + return y0 + (y1 - y0) * (tokens - x0) / (x1 - x0).max(f64::EPSILON); + } + } + let n = table.len(); + let ((x0, y0), (x1, y1)) = (table[n - 2], table[n - 1]); + let slope = ((y1 - y0) / (x1 - x0).max(f64::EPSILON)).max(0.0); + y1 + slope * (tokens - x1) + } + + /// Prefill time of a pass that computes `total` uncached tokens, the + /// largest single request's share being `largest`: a lone request costs + /// its table value; a batched pass costs the pass-total form, never less + /// than the table value of its largest request. + fn prefill_pass_ms(&self, total: u32, largest: u32) -> f64 { + if total == 0 { + return 0.0; + } + match self { + Self::Calibrated { + table, + pass_intercept_ms, + pass_ms_per_token, + .. + } => { + let largest = largest.min(total).max(1); + let single = Self::table_ms(table, f64::from(largest)); + if largest >= total { + // A lone request costs its table value, whatever the pass form says. + return single; + } + let batched = pass_intercept_ms + pass_ms_per_token * f64::from(total); + single.max(batched) + } + other => other.prefill_ms(total), + } + } + + fn prefill_ms(&self, tokens: u32) -> f64 { + if tokens == 0 { + return 0.0; + } + match self { + Self::Polynomial { + prefill: [a, b, c], .. + } => { + let t = f64::from(tokens); + a + b * t + c * t * t + } + Self::Linear { prefill_tps, .. } => f64::from(tokens) / prefill_tps.max(1.0) * 1000.0, + Self::Calibrated { .. } => self.prefill_pass_ms(tokens, tokens), + } + } + + fn decode_ms(&self, batch: usize, active_tokens: u64, capacity_tokens: u64) -> f64 { + if batch == 0 { + return 0.0; + } + match self { + Self::Polynomial { + decode: [a, b, c], .. + } => { + let u = (active_tokens as f64 / capacity_tokens.max(1) as f64).min(1.0); + (a + b * u + c * u * u).max(1.0) + } + Self::Linear { + decode_base_ms, + decode_per_req_ms, + .. + } => decode_base_ms + decode_per_req_ms * batch as f64, + Self::Calibrated { + decode: [a, b, c], .. + } => { + let u = (active_tokens as f64 / capacity_tokens.max(1) as f64).min(1.0); + (a + b * u + c * u * u).max(1.0) + } + } + } +} + +/// A hardware calibration of the pass model, as the GPU harness writes it: +/// prefill `a + b·T + c·T²` ms over the uncached tokens of a pass, decode +/// `d + e·u + f·u²` ms over KV utilisation, the engine's KV capacity and a +/// fixed per-request overhead. The JSON accepts the harness's own layout +/// (`prefill_fit_ms: {a_ms, b_ms_per_token, c_ms_per_token2}`, +/// `decode_fit_vs_utilisation_ms: {d_ms, e_ms_per_u, f_ms_per_u2}`), or the +/// coefficients as objects (`{"a":..,"b":..,"c":..}` / `{"d":..,"e":..,"f":..}`) +/// or arrays under `prefill_ms`/`prefill` and `decode_ms`/`decode`; capacity +/// as `kv_capacity_tokens` or `kv_capacity_blocks` (with `block_size`); the +/// overhead only as an explicit `request_overhead_ms` (the harness's +/// `fixed_overhead_ms` is the measured one-token TTFT, which the prefill +/// intercept already covers, so it is not added again). Unknown keys are +/// ignored. +#[derive(Clone, Debug, PartialEq)] +pub(crate) struct Calibration { + pub prefill: [f64; 3], + pub decode: [f64; 3], + /// Measured single-request prefill `(tokens, ms)` points, sorted by tokens + /// (`prefill_table_ms: [[tokens, ms], ...]`, or the harness's + /// `prefill_points_ms: {"": {"median_ms": ..}}`). + pub prefill_table: Option>, + /// Pass-total form for batched prefills, `(intercept_ms, ms_per_token)` + /// (`prefill_pass_ms: {"intercept_ms": .., "ms_per_token": ..}`). + pub prefill_pass: Option<(f64, f64)>, + pub kv_capacity_tokens: Option, + pub block_size: Option, + pub request_overhead_ms: f64, +} + +impl Calibration { + pub(crate) fn load(path: &str) -> Result { + let text = std::fs::read_to_string(path) + .map_err(|e| format!("cannot read calibration {path}: {e}"))?; + let value: serde_json::Value = serde_json::from_str(&text) + .map_err(|e| format!("calibration {path} is not JSON: {e}"))?; + Self::from_value(&value).map_err(|e| format!("calibration {path}: {e}")) + } + + pub(crate) fn from_value(v: &serde_json::Value) -> Result { + let poly = |keys: [&str; 3], names: [&str; 3]| -> Result<[f64; 3], String> { + let node = keys + .iter() + .find_map(|k| v.get(*k)) + .ok_or_else(|| format!("missing {}", keys[0]))?; + if let Some(arr) = node.as_array() { + let got: Vec = arr.iter().filter_map(serde_json::Value::as_f64).collect(); + return match got.as_slice() { + [a, b, c] => Ok([*a, *b, *c]), + _ => Err(format!("{} needs three numbers", keys[0])), + }; + } + let mut out = [0.0; 3]; + for (slot, name) in out.iter_mut().zip(names) { + *slot = node + .get(name) + .and_then(serde_json::Value::as_f64) + .ok_or_else(|| format!("{} is missing {name}", keys[0]))?; + } + Ok(out) + }; + let prefill_table = Self::table_of(v); + let prefill_pass = v + .get("prefill_pass_ms") + .or_else(|| v.get("batched_prefill_ms")) + .and_then(|node| { + Some(( + node.get("intercept_ms")?.as_f64()?, + node.get("ms_per_token")?.as_f64()?, + )) + }); + let prefill = poly( + ["prefill_fit_ms", "prefill_ms", "prefill"], + ["a_ms", "b_ms_per_token", "c_ms_per_token2"], + ) + .or_else(|_| poly(["prefill_ms", "prefill", "prefill_fit_ms"], ["a", "b", "c"])) + .or_else(|e| { + // A table alone is a complete prefill model. + if prefill_table.is_some() { + Ok([0.0; 3]) + } else { + Err(e) + } + })?; + let decode_keys = ["decode_fit_vs_utilisation_ms", "decode_ms", "decode"]; + let decode = poly(decode_keys, ["d_ms", "e_ms_per_u", "f_ms_per_u2"]) + .or_else(|_| poly(decode_keys, ["d", "e", "f"])) + .or_else(|_| poly(decode_keys, ["a", "b", "c"]))?; + let block_size = v + .get("block_size") + .and_then(serde_json::Value::as_u64) + .map(|b| b as u32); + let kv_capacity_tokens = v + .get("kv_capacity_tokens") + .and_then(serde_json::Value::as_u64) + .or_else(|| { + let blocks = v.get("kv_capacity_blocks")?.as_u64()?; + Some(blocks * u64::from(block_size?)) + }); + let request_overhead_ms = v + .get("request_overhead_ms") + .or_else(|| v.get("per_request_overhead_ms")) + .and_then(serde_json::Value::as_f64) + .unwrap_or(0.0); + Ok(Self { + prefill, + decode, + prefill_table, + prefill_pass, + kv_capacity_tokens, + block_size, + request_overhead_ms, + }) + } + + /// The prefill point table, from `prefill_table_ms` (`[[tokens, ms], ..]`), + /// `prefill_points` (`[{tokens, ms | prefill_ms | median_ms}, ..]`) or the + /// harness's `prefill_points_ms` (`{"": {"median_ms": ..}}`). + fn table_of(v: &serde_json::Value) -> Option> { + let mut points: Vec<(f64, f64)> = Vec::new(); + if let Some(rows) = v + .get("prefill_table_ms") + .and_then(serde_json::Value::as_array) + { + for row in rows { + if let (Some(t), Some(ms)) = ( + row.get(0).and_then(serde_json::Value::as_f64), + row.get(1).and_then(serde_json::Value::as_f64), + ) { + points.push((t, ms)); + } + } + } else if let Some(rows) = v + .get("prefill_points") + .and_then(serde_json::Value::as_array) + { + for row in rows { + let ms = ["ms", "prefill_ms", "median_ms"] + .iter() + .find_map(|k| row.get(*k).and_then(serde_json::Value::as_f64)); + if let (Some(t), Some(ms)) = + (row.get("tokens").and_then(serde_json::Value::as_f64), ms) + { + points.push((t, ms)); + } + } + } else if let Some(map) = v + .get("prefill_points_ms") + .and_then(serde_json::Value::as_object) + { + for (tokens, entry) in map { + let ms = entry + .get("median_ms") + .and_then(serde_json::Value::as_f64) + .or_else(|| entry.as_f64()); + if let (Ok(t), Some(ms)) = (tokens.parse::(), ms) { + points.push((t, ms)); + } + } + } + if points.len() < 2 { + return None; + } + points.sort_by(|a, b| a.0.partial_cmp(&b.0).unwrap_or(std::cmp::Ordering::Equal)); + Some(points) + } +} + +/// Tunable parameters of the simulated engine. The scheduler is vLLM's pass +/// loop (a token budget per pass, running requests first, then FCFS waiting, +/// LIFO preemption when KV runs out) over a block-level KV pool with +/// reference counts and an LRU of idle cached blocks; override via +/// mock-worker flags. #[derive(Clone, Debug)] pub struct EngineParams { - /// Prefill throughput (prompt tokens processed per second). - pub prefill_tps: f64, - /// Fixed decode-step latency (ms) independent of batch size. - pub decode_base_ms: f64, - /// Added decode-step latency (ms) per running request — the batch slope. - pub decode_per_req_ms: f64, - /// Max prompt tokens prefilled per scheduler step (chunked prefill); keeps a - /// huge prompt from stalling the whole batch in a single step. - pub prefill_chunk_tokens: u32, - /// Max concurrent running requests (continuous-batching width). + /// Pass duration model. + pub timing: TimingModel, + /// Token budget per pass (`max_num_batched_tokens`): each decode token + /// costs one, prefill chunks take the rest. + pub max_batched_tokens: u32, + /// Max sequences in a pass (`max_num_seqs`). pub max_running: usize, - /// KV cache capacity in tokens. + /// KV cache capacity in tokens (`num_blocks × block_size`). pub kv_capacity_tokens: u64, + /// The KV capacity the decode cost was calibrated against. When set, a + /// decode step's utilisation is the decoding context over this reference, + /// not over the configured pool, so shrinking the pool (`--kv-blocks`) + /// changes how much fits, not how long a step takes: the engine's step + /// time depends on the tokens it attends to, not on the pool size. + pub decode_reference_tokens: Option, /// Cache block (page) size in tokens. pub block_size: u32, /// Whether prefix caching + KV-event emission are enabled. pub prefix_cache: bool, - /// Start evicting once KV usage exceeds this fraction of capacity. - pub kv_high_watermark: f64, - /// Evict down to this fraction of capacity once eviction starts. - pub kv_low_watermark: f64, + /// SGLang-style scheduling: a pass that contains prefill runs prefill only. + pub prefill_first: bool, + /// vLLM's `scheduler_reserve_full_isl`: a waiting request is admitted only + /// when the free and evictable blocks could hold its whole prompt beyond + /// its cached blocks, and admission stops at the first request that does + /// not fit; only the pass's chunk is allocated at admission, the next + /// chunks allocate as they run (the gate is read, not reserved, as in + /// vLLM's `full_sequence_must_fit`). Off, the gate is skipped (the older, + /// over-admitting behaviour). + pub reserve_full_isl: bool, /// Output tokens to generate when a request does not specify `max_new_tokens`. pub max_new_default: u32, + /// Fixed per-request overhead in ms (a calibration's intercept that the + /// pass polynomials do not explain): every event of a request's stream is + /// delivered that much later, so TTFT and e2e grow by it and ITL does not. + pub request_overhead_ms: f64, /// Capacity of the live KV-event broadcast channel. pub kv_broadcast_capacity: usize, /// How many recent KV-event batches to retain for subscriber replay. @@ -73,23 +427,30 @@ pub struct EngineParams { impl Default for EngineParams { fn default() -> Self { Self { - prefill_tps: 8000.0, - decode_base_ms: 6.0, - decode_per_req_ms: 0.35, - prefill_chunk_tokens: 2048, + timing: TimingModel::polynomial(), + max_batched_tokens: 8192, max_running: 256, kv_capacity_tokens: 524_288, + decode_reference_tokens: None, block_size: 16, prefix_cache: true, - kv_high_watermark: 0.92, - kv_low_watermark: 0.85, + prefill_first: false, + reserve_full_isl: true, max_new_default: 128, + request_overhead_ms: 0.0, kv_broadcast_capacity: 1024, kv_replay_capacity: 4096, } } } +impl EngineParams { + /// KV capacity in blocks. + fn capacity_blocks(&self) -> u64 { + (self.kv_capacity_tokens / u64::from(self.block_size.max(1))).max(1) + } +} + // --------------------------------------------------------------------------- // Public request / event types (transport-agnostic) // --------------------------------------------------------------------------- @@ -127,7 +488,7 @@ pub enum GenEvent { } /// A point-in-time view of engine load, served via `GetLoads` / `/v1/loads`. -#[derive(Clone, Debug)] +#[derive(Clone, Debug, PartialEq)] pub struct LoadSnapshot { pub num_running_reqs: i32, pub num_waiting_reqs: i32, @@ -138,6 +499,12 @@ pub struct LoadSnapshot { pub token_usage: f64, pub gen_throughput: f64, pub cache_hit_rate: f64, + /// Cached blocks (referenced or idle) in the KV pool. + pub num_cached_blocks: i32, + /// Requests preempted so far (LIFO, on KV exhaustion). + pub num_preemptions: i64, + /// KV-event batches produced so far (the current gRPC sequence number). + pub num_kv_batches: i64, } // --------------------------------------------------------------------------- @@ -147,15 +514,160 @@ pub struct LoadSnapshot { /// Shared state readable from the transport handlers while the actor runs. struct EngineShared { snapshot: RwLock, - kv_tx: broadcast::Sender, + /// Every published batch, as the publisher task releases it. + kv_tx: broadcast::Sender, + /// Hands batches from the actor to the publisher task with their release time. + publish_tx: mpsc::UnboundedSender<(Published, Instant)>, + /// Hands a request's events to the delivery task with their release time + /// (the per-request overhead); unused when the overhead is zero. + deliver_tx: mpsc::UnboundedSender<(Instant, mpsc::UnboundedSender, GenEvent)>, kv_replay: Mutex>, prefix_cache: bool, + block_size: u32, + /// Worker name for records and the admin API (`grpc:` / `http:`). + name: String, + /// Mirror of the actor's cache block keys, so sibling engines and the + /// admin API can read it without entering the actor: the fleet oracle + /// ("which worker holds the most of this prompt") is a read over these. + cache_mirror: RwLock>, + /// Fault hooks, switched on through the admin API. + faults: Faults, + /// Wakes a paused actor. + resume: Notify, +} + +/// Fault hooks the admin API switches on: batches lost on the wire, delayed +/// publishing, a frozen engine, a publisher restart. +#[derive(Default)] +struct Faults { + /// Event batches still to lose on the wire. + drop_batches: AtomicU32, + dropped_total: AtomicU64, + /// Publishing delay after a pass ends, in ms. + delay_ms: AtomicU64, + /// The engine is frozen: no passes until resumed. + paused: AtomicBool, + /// Publisher generation; a restart bumps it. + generation: AtomicU64, + restarts: AtomicU64, +} + +/// The hooks' current state. +#[derive(Clone, Debug, PartialEq, Eq)] +pub(crate) struct FaultStatus { + pub drop_pending: u32, + pub dropped_total: u64, + pub delay_ms: u64, + pub paused: bool, + pub generation: u64, + pub restarts: u64, +} + +/// What the engine itself would serve from cache for a prompt right now. +#[derive(Clone, Debug, PartialEq, Eq)] +pub(crate) struct CacheTruth { + pub cached_tokens: u32, + pub cached_blocks: u32, + pub block_size: u32, +} + +/// A batch as the publisher releases it to every transport. +#[derive(Clone, Debug)] +pub(crate) struct Published { + pub batch: common::KvEventBatch, + /// Lost on the wire by a drop hook: not delivered live, kept for replay. + pub dropped: bool, + /// The publisher generation the batch belongs to. + pub generation: u64, +} + +/// Messages into the engine actor. +enum EngineMsg { + Request(NewRequest), + /// Drop every cached block and announce `AllBlocksCleared` (an engine + /// restart, as far as the gateway's index is concerned). + Reset, + /// The publisher restarts: sequence numbers start over, the replay + /// buffer is emptied, the cache is kept. Acknowledged once applied. + RestartPublisher(tokio::sync::oneshot::Sender<()>), +} + +/// One admitted request as seen by the engine: the ground truth for routing +/// accuracy. `oracle_tokens` is the best cached prefix any worker of this +/// process held when the request arrived (the router's ideal choice). +#[derive(Clone, Debug)] +pub(crate) struct RequestRecord { + pub seq: u64, + pub request_id: String, + pub worker: String, + pub prompt_tokens: u32, + pub cached_tokens: u32, + pub oracle_tokens: u32, + pub queued_ms: f64, + pub running_at_admit: u32, + pub waiting_at_admit: u32, + pub admitted_unix_ms: u64, +} + +/// Every engine in this process, for the fleet oracle and the admin API. +fn fleet() -> &'static Mutex> { + static FLEET: OnceLock>> = OnceLock::new(); + FLEET.get_or_init(|| Mutex::new(Vec::new())) +} + +/// Recent admitted-request records across the fleet (a ring buffer). +fn records() -> &'static Mutex<(u64, VecDeque)> { + static RECORDS: OnceLock)>> = OnceLock::new(); + RECORDS.get_or_init(|| Mutex::new((0, VecDeque::new()))) +} + +const RECORD_CAPACITY: usize = 500_000; + +/// All engines registered in this process. +pub(crate) fn fleet_engines() -> Vec { + fleet().lock().unwrap_or_else(|p| p.into_inner()).clone() +} + +/// Records with `seq > since`, oldest first, at most `limit`. +pub(crate) fn records_since(since: u64, limit: usize) -> Vec { + let guard = records().lock().unwrap_or_else(|p| p.into_inner()); + guard + .1 + .iter() + .filter(|r| r.seq > since) + .take(limit) + .cloned() + .collect() +} + +fn push_record(mut record: RequestRecord) { + let mut guard = records().lock().unwrap_or_else(|p| p.into_inner()); + guard.0 += 1; + record.seq = guard.0; + guard.1.push_back(record); + while guard.1.len() > RECORD_CAPACITY { + guard.1.pop_front(); + } +} + +fn unix_ms(at: SystemTime) -> u64 { + at.duration_since(UNIX_EPOCH) + .map(|d| d.as_millis() as u64) + .unwrap_or(0) +} + +/// Seconds since the Unix epoch, as the engines stamp their event batches. +pub(crate) fn unix_seconds() -> f64 { + SystemTime::now() + .duration_since(UNIX_EPOCH) + .map(|d| d.as_secs_f64()) + .unwrap_or(0.0) } /// A cloneable handle to one simulated engine. #[derive(Clone)] pub struct Engine { - tx: mpsc::UnboundedSender, + tx: mpsc::UnboundedSender, shared: Arc, } @@ -168,26 +680,168 @@ impl Engine { /// The actor task is detached intentionally: when the last [`Engine`] handle /// drops, its request channel closes and `run` returns, so there is nothing /// to wait on at shutdown. + pub fn spawn(params: EngineParams) -> Engine { + Self::spawn_named(params, String::new(), false) + } + + /// Spawn a named engine; `register` adds it to the process fleet, which + /// the arrival-time oracle and the admin API read. #[expect( clippy::disallowed_methods, reason = "engine actor self-terminates when its request channel closes" )] - pub fn spawn(params: EngineParams) -> Engine { + pub fn spawn_named(params: EngineParams, name: String, register: bool) -> Engine { let (tx, rx) = mpsc::unbounded_channel(); let (kv_tx, _) = broadcast::channel(params.kv_broadcast_capacity.max(1)); + let (publish_tx, publish_rx) = mpsc::unbounded_channel(); + let (deliver_tx, deliver_rx) = mpsc::unbounded_channel(); let shared = Arc::new(EngineShared { snapshot: RwLock::new(LoadSnapshot::idle(¶ms)), - kv_tx, + kv_tx: kv_tx.clone(), + publish_tx, + deliver_tx, kv_replay: Mutex::new(VecDeque::new()), prefix_cache: params.prefix_cache, + block_size: params.block_size, + name, + cache_mirror: RwLock::new(HashSet::new()), + faults: Faults::default(), + resume: Notify::new(), }); + // The publisher task releases batches in order at their release time + // (the delay hook) and ends with the actor, which owns its sender. + tokio::spawn(publish(publish_rx, kv_tx)); + tokio::spawn(deliver(deliver_rx)); tokio::spawn(run(params, rx, shared.clone())); - Engine { tx, shared } + let engine = Engine { tx, shared }; + if register { + fleet() + .lock() + .unwrap_or_else(|p| p.into_inner()) + .push(engine.clone()); + } + engine } /// Submit a request. Dropped silently if the engine has shut down. pub fn submit(&self, req: NewRequest) { - let _ = self.tx.send(req); + let _ = self.tx.send(EngineMsg::Request(req)); + } + + /// Clear the prefix cache and announce it (`AllBlocksCleared`). + pub fn reset(&self) { + let _ = self.tx.send(EngineMsg::Reset); + } + + /// Lose the next `batches` event batches on the wire (they stay in the + /// replay buffer). + pub(crate) fn fault_drop(&self, batches: u32) { + self.shared + .faults + .drop_batches + .store(batches, Ordering::Relaxed); + } + + /// Publish every batch `ms` milliseconds after its pass ends (0 clears). + pub(crate) fn fault_delay_ms(&self, ms: u64) { + self.shared.faults.delay_ms.store(ms, Ordering::Relaxed); + } + + /// Restart the publisher: sequence numbers start over, the replay buffer + /// is emptied, the cache is kept. Returns once the actor has applied it + /// (after its current pass, at most), so a status read right after sees + /// the new generation. + pub(crate) async fn restart_publisher(&self) { + let (ack, applied) = tokio::sync::oneshot::channel(); + if self.tx.send(EngineMsg::RestartPublisher(ack)).is_ok() { + let _ = tokio::time::timeout(Duration::from_secs(2), applied).await; + } + } + + /// Freeze the engine: no passes, no tokens, no events; requests queue. + pub(crate) fn pause(&self) { + self.shared.faults.paused.store(true, Ordering::Relaxed); + } + + /// Run again after a pause. + pub(crate) fn resume(&self) { + self.shared.faults.paused.store(false, Ordering::Relaxed); + self.shared.resume.notify_one(); + } + + pub(crate) fn fault_status(&self) -> FaultStatus { + let f = &self.shared.faults; + FaultStatus { + drop_pending: f.drop_batches.load(Ordering::Relaxed), + dropped_total: f.dropped_total.load(Ordering::Relaxed), + delay_ms: f.delay_ms.load(Ordering::Relaxed), + paused: f.paused.load(Ordering::Relaxed), + generation: f.generation.load(Ordering::Relaxed), + restarts: f.restarts.load(Ordering::Relaxed), + } + } + + /// What this engine would serve from cache for `token_ids` right now: + /// its own prefix match over the mirror, last block recomputed, as + /// admission would compute it. + pub(crate) fn cached_tokens_for(&self, token_ids: &[u32]) -> CacheTruth { + let bs = self.shared.block_size.max(1); + let keys = prompt_blocks(token_ids, bs as usize).0; + let mut matched = self.match_prefix(&keys) as u32; + if matched > 0 && matched * bs >= token_ids.len() as u32 { + matched -= 1; + } + CacheTruth { + cached_tokens: matched * bs, + cached_blocks: matched, + block_size: bs, + } + } + + /// Every batch the publisher releases from now on, with its drop mark and + /// generation (for transports that keep their own replay buffer). + pub(crate) fn subscribe_published(&self) -> Pin + Send>> { + let live_rx = self.shared.kv_tx.subscribe(); + Box::pin(stream::unfold(live_rx, |mut rx| async move { + loop { + match rx.recv().await { + Ok(item) => return Some((item, rx)), + Err(broadcast::error::RecvError::Lagged(_)) => continue, + Err(broadcast::error::RecvError::Closed) => return None, + } + } + })) + } + + /// This engine's name (`grpc:` / `http:`; empty when unnamed). + pub(crate) fn name(&self) -> &str { + &self.shared.name + } + + /// Number of consecutive blocks from the start of `keys` this engine holds. + fn match_prefix(&self, keys: &[u64]) -> usize { + let mirror = self + .shared + .cache_mirror + .read() + .unwrap_or_else(|p| p.into_inner()); + keys.iter().take_while(|k| mirror.contains(k)).count() + } + + /// A copy of the cached block keys. + pub fn cache_keys(&self) -> Vec { + self.shared + .cache_mirror + .read() + .unwrap_or_else(|p| p.into_inner()) + .iter() + .copied() + .collect() + } + + /// Block keys of a prompt, exactly as the engine computes them. + pub fn block_keys(prompt_token_ids: &[u32], block_size: usize) -> Vec { + prompt_blocks(prompt_token_ids, block_size).0 } /// Current load snapshot. @@ -200,7 +854,7 @@ impl Engine { } /// Whether this engine emits KV-cache events. - pub fn kv_enabled(&self) -> bool { + pub(crate) fn kv_enabled(&self) -> bool { self.shared.prefix_cache } @@ -225,7 +879,10 @@ impl Engine { let live_stream = stream::unfold(live_rx, |mut rx| async move { loop { match rx.recv().await { - Ok(batch) => return Some((Ok(batch), rx)), + // A batch lost on the wire (drop hook) never reaches a live + // subscriber; it waits in the replay buffer. + Ok(item) if item.dropped => continue, + Ok(item) => return Some((Ok(item.batch), rx)), // A lagged slow consumer leaves a gap; the gateway detects it // and reconnects, replaying from the buffer. Skip and continue. Err(broadcast::error::RecvError::Lagged(_)) => continue, @@ -243,7 +900,7 @@ impl Engine { async fn run( params: EngineParams, - mut rx: mpsc::UnboundedReceiver, + mut rx: mpsc::UnboundedReceiver, shared: Arc, ) { let mut state = SchedulerState::new(); @@ -253,22 +910,71 @@ async fn run( if state.is_idle() { *shared.snapshot.write().unwrap_or_else(|p| p.into_inner()) = state.snapshot(¶ms); match rx.recv().await { - Some(req) => state.enqueue(req, ¶ms), + Some(msg) => handle_msg(&mut state, &shared, ¶ms, msg), None => return, // all handles dropped } } // Drain any other already-queued submissions without blocking. - while let Ok(req) = rx.try_recv() { - state.enqueue(req, ¶ms); + while let Ok(msg) = rx.try_recv() { + handle_msg(&mut state, &shared, ¶ms, msg); + } + // A paused engine takes messages (requests queue, scored at arrival) + // but runs no pass until resumed. + while shared.faults.paused.load(Ordering::Relaxed) { + *shared.snapshot.write().unwrap_or_else(|p| p.into_inner()) = state.snapshot(¶ms); + tokio::select! { + _ = shared.resume.notified() => {} + msg = rx.recv() => match msg { + Some(msg) => handle_msg(&mut state, &shared, ¶ms, msg), + None => return, + }, + } } let step = state.step(¶ms); if step.duration > Duration::ZERO { tokio::time::sleep(step.duration).await; } - // The step's outputs become observable only after its simulated time. - for (tx, ev) in step.sends { - let _ = tx.send(ev); + // The step's outputs become observable only after its simulated time, + // plus the fixed per-request overhead when one is calibrated (every + // event of a stream shifts by the same amount, so order is kept). + if params.request_overhead_ms > 0.0 { + let at = Instant::now() + Duration::from_secs_f64(params.request_overhead_ms / 1000.0); + for (tx, ev) in step.sends { + let _ = shared.deliver_tx.send((at, tx, ev)); + } + } else { + for (tx, ev) in step.sends { + let _ = tx.send(ev); + } + } + if step.cleared || !step.inserted.is_empty() || !step.evicted.is_empty() { + let mut mirror = shared + .cache_mirror + .write() + .unwrap_or_else(|p| p.into_inner()); + if step.cleared { + mirror.clear(); + } + for k in &step.evicted { + mirror.remove(k); + } + mirror.extend(step.inserted.iter().copied()); + } + let admitted_at = SystemTime::now(); + for admitted in step.admitted { + push_record(RequestRecord { + seq: 0, + request_id: admitted.request_id, + worker: shared.name.clone(), + prompt_tokens: admitted.prompt_tokens, + cached_tokens: admitted.cached_tokens, + oracle_tokens: admitted.oracle_tokens, + queued_ms: admitted.enqueued_at.elapsed().as_secs_f64() * 1000.0, + running_at_admit: admitted.running_at_admit, + waiting_at_admit: admitted.waiting_at_admit, + admitted_unix_ms: unix_ms(admitted_at), + }); } if let Some(batch) = step.batch { { @@ -278,275 +984,695 @@ async fn run( buf.pop_front(); } } - let _ = shared.kv_tx.send(batch); + let faults = &shared.faults; + let dropped = faults + .drop_batches + .fetch_update(Ordering::Relaxed, Ordering::Relaxed, |n| n.checked_sub(1)) + .is_ok(); + if dropped { + faults.dropped_total.fetch_add(1, Ordering::Relaxed); + } + let release_at = + Instant::now() + Duration::from_millis(faults.delay_ms.load(Ordering::Relaxed)); + let item = Published { + batch, + dropped, + generation: faults.generation.load(Ordering::Relaxed), + }; + let _ = shared.publish_tx.send((item, release_at)); } *shared.snapshot.write().unwrap_or_else(|p| p.into_inner()) = step.snapshot; } } +/// Release batches to every transport in order, each at its release time. +async fn publish( + mut rx: mpsc::UnboundedReceiver<(Published, Instant)>, + kv_tx: broadcast::Sender, +) { + while let Some((item, release_at)) = rx.recv().await { + tokio::time::sleep_until(release_at.into()).await; + let _ = kv_tx.send(item); + } +} + +/// Deliver request events in order, each at its release time. +async fn deliver( + mut rx: mpsc::UnboundedReceiver<(Instant, mpsc::UnboundedSender, GenEvent)>, +) { + while let Some((at, tx, ev)) = rx.recv().await { + tokio::time::sleep_until(at.into()).await; + let _ = tx.send(ev); + } +} + +/// Apply one actor message: enqueue a request (scoring the fleet oracle at +/// arrival, before this engine's own cache changes), reset the cache, or +/// restart the publisher. +fn handle_msg( + state: &mut SchedulerState, + shared: &Arc, + params: &EngineParams, + msg: EngineMsg, +) { + match msg { + EngineMsg::Request(req) => { + let oracle_tokens = if params.prefix_cache { + let keys = prompt_blocks(&req.prompt_token_ids, params.block_size as usize).0; + let fleet_best = fleet_engines() + .iter() + .map(|e| e.match_prefix(&keys)) + .max() + .unwrap_or(0); + let own_best = { + let own = shared + .cache_mirror + .read() + .unwrap_or_else(|p| p.into_inner()); + keys.iter().take_while(|k| own.contains(k)).count() + }; + (fleet_best.max(own_best) as u32 * params.block_size) + .min(req.prompt_token_ids.len() as u32) + } else { + 0 + }; + state.enqueue_with_oracle(req, params, oracle_tokens); + } + EngineMsg::Reset => state.reset(), + EngineMsg::RestartPublisher(ack) => { + state.restart_publisher(); + shared + .kv_replay + .lock() + .unwrap_or_else(|p| p.into_inner()) + .clear(); + shared.faults.generation.fetch_add(1, Ordering::Relaxed); + shared.faults.restarts.fetch_add(1, Ordering::Relaxed); + let _ = ack.send(()); + } + } +} + // --------------------------------------------------------------------------- -// Scheduler state + step (pure, deterministic, unit-testable) +// Scheduler state + pass (pure, deterministic, unit-testable) // --------------------------------------------------------------------------- -/// A request currently being prefilled or decoded. +/// A request in the running batch (being prefilled or decoded). struct RunningReq { + request_id: String, events: mpsc::UnboundedSender, + /// Tokens to (re)compute: the prompt, extended by the output generated + /// before a preemption. + seq_prompt: Vec, + /// Content keys of `seq_prompt`'s full blocks. + seq_keys: Vec, + /// Prompt / cached tokens as reported on the stream (constant per request). prompt_tokens: u32, cached_tokens: u32, + /// `seq_prompt` tokens computed so far (cached + prefilled); prefill ends + /// when it reaches the prompt length. + computed: u32, max_new: u32, + /// Output tokens produced in total, including those folded into + /// `seq_prompt` by a preemption. generated: u32, - /// Uncached prompt tokens still to be prefilled before the first token. - prefill_remaining: u32, - /// FNV hash of all tokens seen so far (prompt + generated) — the content key - /// of the next block once `pending_block` fills. - rolling_hash: u64, - /// Block key of the most recently completed block (for parent chaining). - prev_block_key: Option, - /// Tokens accumulated toward the next (not-yet-full) block. - pending_block: Vec, - /// Generated token ids (for the terminal `output_ids`). + /// Output tokens already folded into `seq_prompt`. + resumed_output: u32, output_ids: Vec, + /// Full blocks this request references, in sequence order. + held: Vec, + /// Whether an anonymous (not yet full) block is allocated for the tail. + partial: bool, + /// Tokens of the not-yet-full tail block (the next stored event's payload). + pending: Vec, + /// FNV over every token of the sequence so far (the next block's key). + rolling_hash: u64, /// Per-request seed so decode blocks never collide across requests. token_seed: u32, } +impl RunningReq { + fn prefilling(&self) -> bool { + (self.computed as usize) < self.seq_prompt.len() + } + + /// Tokens of the sequence resident in KV right now. + fn seq_len(&self) -> u32 { + if self.prefilling() { + self.computed + } else { + self.seq_prompt.len() as u32 + self.generated - self.resumed_output + } + } + + fn finished(&self) -> bool { + !self.prefilling() && self.generated >= self.max_new + } +} + +/// State a preempted request carries back to the queue. +struct Resume { + generated: u32, + output_ids: Vec, + cached_tokens: u32, + token_seed: u32, +} + /// A request admitted to the queue but not yet running. struct WaitingReq { req: NewRequest, + /// Content keys of the prompt's full blocks and the hash of every token. + keys: Vec, + rolling_hash: u64, + /// Prompt tokens as reported on the stream (the original prompt). prompt_tokens: u32, /// Uncached prompt tokens at enqueue time (the queued token-work it adds). uncached_tokens: u32, + /// Best cached prefix any fleet worker held at arrival (ground truth). + oracle_tokens: u32, + enqueued_at: Instant, + /// Set when this is a preempted request coming back for recompute. + resume: Option, +} + +/// What a pass reports about a request it admitted for the first time. +struct Admitted { + request_id: String, + prompt_tokens: u32, + cached_tokens: u32, + oracle_tokens: u32, + running_at_admit: u32, + waiting_at_admit: u32, + enqueued_at: Instant, } -/// Per-worker prefix cache: content-keyed block presence with LRU ordering. +/// Block-level KV pool: cached blocks keyed by content hash with reference +/// counts, idle (unreferenced) cached blocks in an LRU whose head is evicted +/// first, and anonymous blocks for tails that are not full yet. `allocated` +/// counts every physical block: referenced, idle-cached and anonymous. #[derive(Default)] -struct Cache { - present: HashSet, +struct BlockPool { + refs: HashMap, tick_of: HashMap, - lru: BTreeSet<(u64, u64)>, + free_lru: BTreeSet<(u64, u64)>, tick: u64, + allocated: u64, } -impl Cache { - fn len(&self) -> usize { - self.present.len() +impl BlockPool { + fn cached(&self) -> usize { + self.refs.len() } /// Number of consecutive cached blocks from the start of `keys`. fn match_prefix(&self, keys: &[u64]) -> usize { - let mut n = 0; - for &k in keys { - if self.present.contains(&k) { - n += 1; - } else { - break; + keys.iter() + .take_while(|k| self.refs.contains_key(k)) + .count() + } + + /// Reserve `n` anonymous blocks, evicting idle cached blocks (LRU first) + /// as needed. Reserves nothing and returns false when even that is not + /// enough. + fn reserve(&mut self, n: u64, capacity: u64, evicted: &mut Vec) -> bool { + if !self.fits(n, capacity) { + return false; + } + while capacity.saturating_sub(self.allocated) < n { + match self.evict_lru() { + Some(h) => evicted.push(h), + None => return false, } } - n + self.allocated += n; + true } - /// Mark `keys` as just-used (warms shared prefixes so eviction prefers - /// colder leaves, keeping the block set prefix-closed in the common case). - fn touch(&mut self, keys: &[u64]) { - self.tick += 1; - let t = self.tick; - for &k in keys { - if self.present.contains(&k) { - if let Some(old) = self.tick_of.insert(k, t) { - self.lru.remove(&(old, k)); + /// Whether `n` more blocks could be reserved now: never-used capacity + /// plus the idle cached blocks an eviction may take (what vLLM's + /// free-block count holds). + fn fits(&self, n: u64, capacity: u64) -> bool { + capacity.saturating_sub(self.allocated) + self.free_lru.len() as u64 >= n + } + + /// Reserve without a capacity check (a prompt larger than all of KV on + /// an otherwise empty engine must still run). + fn force_reserve(&mut self, n: u64) { + self.allocated += n; + } + + fn release_anonymous(&mut self, n: u64) { + self.allocated = self.allocated.saturating_sub(n); + } + + /// Take a reference to cached block `h` (leaving the idle LRU if it was there). + fn hit(&mut self, h: u64) { + if let Some(r) = self.refs.get_mut(&h) { + if *r == 0 { + if let Some(t) = self.tick_of.remove(&h) { + self.free_lru.remove(&(t, h)); } - self.lru.insert((t, k)); } + *r += 1; } } - /// Insert a block; returns true if it was newly added. - fn insert(&mut self, k: u64) -> bool { - if self.present.contains(&k) { - self.touch(&[k]); - return false; + /// Turn one of the caller's anonymous blocks into cached block `h`. + /// Returns true when the hash is new (a stored event is due); when it + /// already exists the anonymous block is given back and a reference taken. + fn register(&mut self, h: u64) -> bool { + match self.refs.entry(h) { + std::collections::hash_map::Entry::Occupied(_) => { + self.hit(h); + self.release_anonymous(1); + false + } + std::collections::hash_map::Entry::Vacant(slot) => { + slot.insert(1); + true + } + } + } + + /// Drop a reference; an unreferenced block becomes idle (evictable). + fn unref(&mut self, h: u64) { + if let Some(r) = self.refs.get_mut(&h) { + *r = r.saturating_sub(1); + if *r == 0 { + self.tick += 1; + self.tick_of.insert(h, self.tick); + self.free_lru.insert((self.tick, h)); + } + } + } + + /// Drop the references a finished (or preempted) request holds, tail + /// block first, so the head of its prefix carries the newest LRU stamp + /// and outlives its tail. vLLM's `free()` appends a request's blocks to + /// the free queue in reverse order for the same reason: a later repeat of + /// the prefix still finds its first blocks, and the chain behind a lost + /// head would be unhittable anyway. + fn unref_request(&mut self, held: &[u64]) { + for h in held.iter().rev() { + self.unref(*h); } - self.present.insert(k); - self.tick += 1; - let t = self.tick; - self.tick_of.insert(k, t); - self.lru.insert((t, k)); - true } - /// Evict the least-recently-used block; returns its key. - fn evict_one(&mut self) -> Option { - let &(t, k) = self.lru.iter().next()?; - self.lru.remove(&(t, k)); - self.tick_of.remove(&k); - self.present.remove(&k); - Some(k) + /// Evict the least recently idle cached block; returns its hash. + fn evict_lru(&mut self) -> Option { + let &(t, h) = self.free_lru.iter().next()?; + self.free_lru.remove(&(t, h)); + self.tick_of.remove(&h); + self.refs.remove(&h); + self.allocated = self.allocated.saturating_sub(1); + Some(h) } } -/// The actor-owned scheduler state. -pub(crate) struct SchedulerState { +/// The actor-owned scheduler state. `running` stays in admission order, which +/// is what LIFO preemption relies on. +struct SchedulerState { running: Vec, waiting: VecDeque, - cache: Cache, + pool: BlockPool, kv_seq: u64, kv_event_id: u64, gen_tp_ewma: f64, cache_hit_ewma: f64, + preemptions: u64, + /// A reset was requested; the next pass clears the cache and says so. + reset_pending: bool, } -/// The result of one scheduler step. -pub(crate) struct Step { +/// The result of one pass. +struct Step { duration: Duration, sends: Vec<(mpsc::UnboundedSender, GenEvent)>, batch: Option, snapshot: LoadSnapshot, + /// Cache deltas this pass, for the actor's mirror. + inserted: Vec, + evicted: Vec, + cleared: bool, + admitted: Vec, +} + +/// A block a request completed this pass: its position in the request's +/// block list, its content key and its tokens. +struct Completed { + index: usize, + key: u64, + tokens: Vec, +} + +/// Making room for a running request's next tokens preempted the request +/// itself (it is back in the queue; the pass moves on without it). +struct SelfPreempted; + +fn blocks_for(tokens: u32, block_size: u32) -> u64 { + u64::from(tokens.div_ceil(block_size.max(1))) } impl SchedulerState { - pub(crate) fn new() -> Self { + fn new() -> Self { Self { running: Vec::new(), waiting: VecDeque::new(), - cache: Cache::default(), + pool: BlockPool::default(), kv_seq: 0, kv_event_id: 0, gen_tp_ewma: 0.0, cache_hit_ewma: 0.0, + preemptions: 0, + reset_pending: false, } } - fn is_idle(&self) -> bool { - self.running.is_empty() && self.waiting.is_empty() + /// Ask the next pass to clear the cache and announce `AllBlocksCleared`. + fn reset(&mut self) { + self.reset_pending = true; } - /// Queue a request, recording the queued token-work it contributes. - pub(crate) fn enqueue(&mut self, req: NewRequest, p: &EngineParams) { - let prompt_tokens = req.prompt_token_ids.len() as u32; - let (block_keys, _, _) = prompt_blocks(&req.prompt_token_ids, p.block_size as usize); + /// The publisher restarted: batch sequence numbers start over. + fn restart_publisher(&mut self) { + self.kv_seq = 0; + } + + fn is_idle(&self) -> bool { + self.running.is_empty() && self.waiting.is_empty() && !self.reset_pending + } + + /// Queue a request with no oracle information (tests). + #[cfg(test)] + fn enqueue(&mut self, req: NewRequest, p: &EngineParams) { + self.enqueue_with_oracle(req, p, 0); + } + + /// Queue a request, recording the queued token-work it contributes, together + /// with the fleet oracle's cached-token count at arrival (what the + /// best-informed router could have obtained). + fn enqueue_with_oracle(&mut self, req: NewRequest, p: &EngineParams, oracle_tokens: u32) { + let prompt_tokens = req.prompt_token_ids.len() as u32; + let (keys, rolling_hash, _) = prompt_blocks(&req.prompt_token_ids, p.block_size as usize); let cached_blocks = if p.prefix_cache { - self.cache.match_prefix(&block_keys) + self.pool.match_prefix(&keys) } else { 0 }; let cached = (cached_blocks as u32 * p.block_size).min(prompt_tokens); - let uncached = prompt_tokens - cached; self.waiting.push_back(WaitingReq { req, + keys, + rolling_hash, prompt_tokens, - uncached_tokens: uncached, + uncached_tokens: prompt_tokens - cached, + oracle_tokens, + enqueued_at: Instant::now(), + resume: None, }); } - /// Tokens currently resident in KV. - fn used_tokens(&self, p: &EngineParams) -> u64 { - if p.prefix_cache { - // KV holds the shared radix cache (blocks persist across requests - // until evicted) plus each running request's not-yet-blocked tail. - let blocks = self.cache.len() as u64 * p.block_size as u64; - let partial: u64 = self - .running - .iter() - .map(|r| r.pending_block.len() as u64) - .sum(); - blocks + partial - } else { - // No sharing/persistence: each running request occupies its full - // current context; that KV frees when it leaves the batch. - self.running - .iter() - .map(|r| u64::from(r.prompt_tokens + r.generated)) - .sum() - } + /// Tokens of running sequences resident in KV (what `token_usage` reports; + /// idle cached blocks are evictable and not counted, as in SGLang). + fn active_tokens(&self) -> u64 { + self.running.iter().map(|r| u64::from(r.seq_len())).sum() } fn snapshot(&self, p: &EngineParams) -> LoadSnapshot { - let used = self.used_tokens(p); - let waiting_uncached: i64 = self.waiting.iter().map(|w| w.uncached_tokens as i64).sum(); + let used = self.active_tokens(); + let waiting_uncached: i64 = self + .waiting + .iter() + .map(|w| i64::from(w.uncached_tokens)) + .sum(); LoadSnapshot { num_running_reqs: self.running.len() as i32, num_waiting_reqs: self.waiting.len() as i32, - num_waiting_uncached_tokens: waiting_uncached.min(i32::MAX as i64) as i32, + num_waiting_uncached_tokens: waiting_uncached.min(i64::from(i32::MAX)) as i32, num_used_tokens: used.min(i32::MAX as u64) as i32, max_total_num_tokens: p.kv_capacity_tokens.min(i32::MAX as u64) as i32, max_running_requests: p.max_running.min(i32::MAX as usize) as i32, token_usage: (used as f64 / p.kv_capacity_tokens.max(1) as f64).clamp(0.0, 1.0), gen_throughput: self.gen_tp_ewma, cache_hit_rate: self.cache_hit_ewma, + num_cached_blocks: self.pool.cached().min(i32::MAX as usize) as i32, + num_preemptions: self.preemptions.min(i64::MAX as u64) as i64, + num_kv_batches: self.kv_seq.min(i64::MAX as u64) as i64, } } - /// Advance the engine by one scheduler iteration. Pure: mutates state and - /// returns the work produced plus how long it took, but performs no I/O. - pub(crate) fn step(&mut self, p: &EngineParams) -> Step { + /// Run one pass. Pure: mutates state and returns the work produced plus + /// how long it took, but performs no I/O. Order within the pass follows + /// vLLM: running requests first (a prefill chunk or one decode token + /// each), then FCFS admission from the queue while the token budget and + /// KV room last. Per request, KV events are the blocks evicted for its + /// allocation (`Removed`) followed by the blocks it completed (`Stored`, + /// contiguous, chained to a parent). The caller makes every event of the + /// pass visible at its end. + fn step(&mut self, p: &EngineParams) -> Step { let mut kv: Vec = Vec::new(); - - // ---- 1. Admission ---- - self.admit(p, &mut kv); - - // ---- 2. Chunked prefill ---- - let mut budget = p.prefill_chunk_tokens; - let mut prefill_tokens = 0u32; - for r in &mut self.running { - if r.prefill_remaining > 0 && budget > 0 { - let c = r.prefill_remaining.min(budget); - r.prefill_remaining -= c; - budget -= c; - prefill_tokens += c; + let mut inserted: Vec = Vec::new(); + let mut evicted: Vec = Vec::new(); + let mut sends: Vec<(mpsc::UnboundedSender, GenEvent)> = Vec::new(); + let mut admitted: Vec = Vec::new(); + let bs = p.block_size.max(1); + + // ---- 0. Reset (an engine restart, from the index's point of view) ---- + let cleared = std::mem::take(&mut self.reset_pending); + if cleared { + while !self.running.is_empty() { + self.preempt_last(p); + } + self.pool = BlockPool::default(); + if p.prefix_cache { + self.kv_event_id += 1; + kv.push(common::KvCacheEvent { + event_id: self.kv_event_id, + data: Some(common::kv_cache_event::Data::Cleared( + common::KvCacheCleared::default(), + )), + }); } } - // ---- 3. Decode (one token per ready request) ---- - let block_size = p.block_size as usize; - let mut sends: Vec<(mpsc::UnboundedSender, GenEvent)> = Vec::new(); - let mut new_blocks: Vec<(u64, Vec, Option)> = Vec::new(); + let mut budget = p.max_batched_tokens; + let mut prefill_tokens = 0u32; + let mut largest_chunk = 0u32; + let mut num_decode = 0usize; let mut decode_tokens = 0u32; - let num_decode = self - .running - .iter() - .filter(|r| { - r.prefill_remaining == 0 && r.generated < r.max_new && !r.events.is_closed() - }) - .count(); + let mut decode_ctx = 0u64; - for r in &mut self.running { - if r.prefill_remaining != 0 || r.generated >= r.max_new || r.events.is_closed() { + // ---- 1. Running requests: a prefill chunk or one decode token each ---- + let mut i = 0; + while i < self.running.len() { + if self.running[i].events.is_closed() || self.running[i].finished() { + i += 1; continue; } - let token_id = next_token(r); - r.generated += 1; - r.output_ids.push(token_id); - r.rolling_hash = fnv_step(r.rolling_hash, token_id); - r.pending_block.push(token_id); - decode_tokens += 1; - if r.pending_block.len() == block_size { - let key = r.rolling_hash; - new_blocks.push((key, std::mem::take(&mut r.pending_block), r.prev_block_key)); - r.prev_block_key = Some(key); + if budget == 0 { + break; } - sends.push(( - r.events.clone(), - GenEvent::Token { - token_id, - prompt_tokens: r.prompt_tokens, - cached_tokens: r.cached_tokens, - }, - )); + if self.running[i].prefilling() { + let remaining = self.running[i].seq_prompt.len() as u32 - self.running[i].computed; + let chunk = remaining.min(budget); + let completes = chunk == remaining; + // The chunk's blocks, plus the first output token's slot when + // this chunk finishes the prompt. + let add = chunk + u32::from(completes); + let mut freed = Vec::new(); + if self.ensure(i, add, p, &mut freed).is_err() { + continue; + } + evicted.extend(freed.iter().copied()); + self.push_removed(&mut kv, freed, p); + let stored = self.apply_prefill(i, chunk, p); + prefill_tokens += chunk; + largest_chunk = largest_chunk.max(chunk); + budget -= chunk; + let mut completed = stored; + if completes { + // The pass that finishes a prefill samples its first token; + // that token adds no decode time to the pass. + completed.extend(self.emit_token(i, &mut sends, p)); + } + inserted.extend(completed.iter().map(|c| c.key)); + self.push_stored(&mut kv, i, completed, p); + } else if !p.prefill_first { + let mut freed = Vec::new(); + if self.ensure(i, 1, p, &mut freed).is_err() { + continue; + } + evicted.extend(freed.iter().copied()); + self.push_removed(&mut kv, freed, p); + decode_ctx += u64::from(self.running[i].seq_len()); + let completed = self.emit_token(i, &mut sends, p); + num_decode += 1; + decode_tokens += 1; + budget -= 1; + inserted.extend(completed.iter().map(|c| c.key)); + self.push_stored(&mut kv, i, completed, p); + } + i += 1; } - // Commit newly completed decode blocks to the cache (+ stored events). - for (key, tokens, parent) in new_blocks { - if p.prefix_cache && self.cache.insert(key) { - kv.push(self.stored_event(key, tokens, parent, p.block_size)); + // ---- 2. FCFS admission while the budget and KV room last ---- + while budget > 0 && self.running.len() < p.max_running { + let Some(front) = self.waiting.front() else { + break; + }; + let prompt_len = front.req.prompt_token_ids.len() as u32; + let mut cached_blocks = if p.prefix_cache { + self.pool.match_prefix(&front.keys) + } else { + 0 + }; + // A fully cached prompt still recomputes its last block. + if cached_blocks > 0 && cached_blocks as u32 * bs >= prompt_len { + cached_blocks -= 1; + } + let cached = cached_blocks as u32 * bs; + let remaining = prompt_len - cached; + let chunk = remaining.min(budget); + let completes = chunk == remaining; + let add = chunk + u32::from(completes); + // This pass's chunk is what gets allocated. Under full-ISL + // admission the whole prompt (plus the first output token's slot) + // must also fit in the free and evictable blocks right now: vLLM's + // `full_sequence_must_fit` is a gate read at admission, not a + // reservation, so an admitted prompt holds no more than it has + // computed and the next chunks allocate (evict, preempt) as they run. + let need = blocks_for(cached + add, bs) - cached_blocks as u64; + let need_full = blocks_for(cached + remaining + 1, bs) - cached_blocks as u64; + // Reference the cached prefix before making room, so the eviction + // cannot take the very blocks this request is about to reuse. + for k in &front.keys[..cached_blocks] { + self.pool.hit(*k); + } + let mut freed = Vec::new(); + let room = (!p.reserve_full_isl || self.pool.fits(need_full, p.capacity_blocks())) + && self.pool.reserve(need, p.capacity_blocks(), &mut freed); + if !room { + if self.running.is_empty() { + self.pool.force_reserve(need); + } else { + let Some(front) = self.waiting.front() else { + break; + }; + for k in &front.keys[..cached_blocks] { + self.pool.unref(*k); + } + break; + } + } + let running_at_admit = self.running.len() as u32; + let waiting_at_admit = self.waiting.len() as u32; + let w = self.waiting.pop_front().expect("front exists"); + let NewRequest { + request_id, + prompt_token_ids, + max_new, + events, + } = w.req; + let (reported_cached, generated, resumed_output, output_ids, token_seed) = + match w.resume { + Some(r) => ( + r.cached_tokens, + r.generated, + r.generated, + r.output_ids, + r.token_seed, + ), + None => { + admitted.push(Admitted { + request_id: request_id.clone(), + prompt_tokens: w.prompt_tokens, + cached_tokens: cached, + oracle_tokens: w.oracle_tokens.max(cached), + running_at_admit, + waiting_at_admit, + enqueued_at: w.enqueued_at, + }); + let sample = if w.prompt_tokens > 0 { + f64::from(cached) / f64::from(w.prompt_tokens) + } else { + 0.0 + }; + self.cache_hit_ewma = ewma(self.cache_hit_ewma, sample, 0.2); + (cached, 0, 0, Vec::new(), fnv_hash_str(&request_id) as u32) + } + }; + let resolved_max_new = if max_new == 0 { + p.max_new_default + } else { + max_new + }; + self.running.push(RunningReq { + request_id, + events, + seq_keys: w.keys, + seq_prompt: prompt_token_ids, + prompt_tokens: w.prompt_tokens, + cached_tokens: reported_cached, + computed: cached, + max_new: resolved_max_new, + generated, + resumed_output, + output_ids, + held: Vec::new(), + partial: false, + pending: Vec::new(), + rolling_hash: w.rolling_hash, + token_seed, + }); + let idx = self.running.len() - 1; + self.running[idx].held = self.running[idx].seq_keys[..cached_blocks].to_vec(); + evicted.extend(freed.iter().copied()); + self.push_removed(&mut kv, freed, p); + let mut completed = self.apply_prefill(idx, chunk, p); + prefill_tokens += chunk; + largest_chunk = largest_chunk.max(chunk); + budget -= chunk; + if completes { + completed.extend(self.emit_token(idx, &mut sends, p)); + } + inserted.extend(completed.iter().map(|c| c.key)); + self.push_stored(&mut kv, idx, completed, p); + } + + // ---- 3. Prefill-first engines decode only in passes without prefill ---- + if p.prefill_first && prefill_tokens == 0 { + let mut i = 0; + while i < self.running.len() { + if self.running[i].events.is_closed() + || self.running[i].finished() + || self.running[i].prefilling() + || budget == 0 + { + i += 1; + continue; + } + let mut freed = Vec::new(); + if self.ensure(i, 1, p, &mut freed).is_err() { + continue; + } + evicted.extend(freed.iter().copied()); + self.push_removed(&mut kv, freed, p); + decode_ctx += u64::from(self.running[i].seq_len()); + let completed = self.emit_token(i, &mut sends, p); + num_decode += 1; + decode_tokens += 1; + budget -= 1; + inserted.extend(completed.iter().map(|c| c.key)); + self.push_stored(&mut kv, i, completed, p); + i += 1; } } - // ---- 4. Completion ---- + // ---- 4. Completion: release references, emit the terminal event ---- let mut still = Vec::with_capacity(self.running.len()); for r in std::mem::take(&mut self.running) { - let done = r.prefill_remaining == 0 && r.generated >= r.max_new; - if done || r.events.is_closed() { - if done { + if r.finished() || r.events.is_closed() { + if r.finished() { sends.push(( r.events.clone(), GenEvent::Done { @@ -557,28 +1683,26 @@ impl SchedulerState { }, )); } + self.release(&r, p); } else { still.push(r); } } self.running = still; - // ---- 5. Eviction under KV pressure ---- - self.evict(p, &mut kv); - - // ---- 6. Timing + bookkeeping ---- - let prefill_time = prefill_tokens as f64 / p.prefill_tps.max(1.0); - let decode_time = if num_decode > 0 { - (p.decode_base_ms + p.decode_per_req_ms * num_decode as f64) / 1000.0 - } else { - 0.0 - }; - let mut secs = prefill_time.max(decode_time); + // ---- 5. Timing + bookkeeping ---- + let prefill_ms = p.timing.prefill_pass_ms(prefill_tokens, largest_chunk); + let decode_ms = p.timing.decode_ms( + num_decode, + decode_ctx, + p.decode_reference_tokens.unwrap_or(p.kv_capacity_tokens), + ); + let mut secs = (prefill_ms + decode_ms) / 1000.0; if secs <= 0.0 && !self.is_idle() { secs = 0.001; // never busy-spin while work remains } let throughput_sample = if secs > 0.0 && decode_tokens > 0 { - decode_tokens as f64 / secs + f64::from(decode_tokens) / secs } else { 0.0 }; @@ -590,9 +1714,13 @@ impl SchedulerState { self.kv_seq += 1; Some(common::KvEventBatch { sequence_number: self.kv_seq, - timestamp: 0.0, + // Creation time, as the engines stamp their batches: a batch the + // delay hook holds back keeps it, so the delay shows up as lag. + timestamp: unix_seconds(), events: kv, dp_rank: Some(0), + snapshot: None, + load: None, }) }; @@ -601,151 +1729,282 @@ impl SchedulerState { sends, batch, snapshot: self.snapshot(p), + inserted, + evicted, + cleared, + admitted, } } - /// Admit waiting requests while batch width and KV capacity allow. - fn admit(&mut self, p: &EngineParams, kv: &mut Vec) { - let block_size = p.block_size as usize; - while self.running.len() < p.max_running { - let Some(front) = self.waiting.front() else { - break; - }; - let (block_keys, rolling, pending) = - prompt_blocks(&front.req.prompt_token_ids, block_size); - let cached_blocks = if p.prefix_cache { - self.cache.match_prefix(&block_keys) - } else { - 0 - }; - let cached_tokens = (cached_blocks as u32 * p.block_size).min(front.prompt_tokens); - let uncached = front.prompt_tokens - cached_tokens; - - // Admission control: require KV room for the uncached prompt unless - // the engine is empty (a prompt larger than all of KV must still run). - let free = p.kv_capacity_tokens.saturating_sub(self.used_tokens(p)); - if u64::from(uncached) > free && !self.running.is_empty() { - break; + /// Make room for `add` more tokens of `running[idx]`, evicting idle cached + /// blocks first and then preempting the most recently admitted request + /// (LIFO) until the allocation fits; `Err` when the victim was the request + /// itself. A lone request that cannot fit even then is over-allocated + /// rather than deadlocked. + fn ensure( + &mut self, + idx: usize, + add: u32, + p: &EngineParams, + freed: &mut Vec, + ) -> Result<(), SelfPreempted> { + let bs = p.block_size.max(1); + let before = self.running[idx].seq_len(); + let need = blocks_for(before + add, bs) - blocks_for(before, bs); + if need == 0 { + return Ok(()); + } + loop { + if self.pool.reserve(need, p.capacity_blocks(), freed) { + return Ok(()); + } + if self.running.len() == 1 { + self.pool.force_reserve(need); + return Ok(()); + } + // `running` is in admission order, so the LIFO victim is the last. + let victim = self.running.len() - 1; + self.preempt_last(p); + if victim == idx { + return Err(SelfPreempted); } + } + } - let w = self.waiting.pop_front().expect("front exists"); - let NewRequest { - request_id, - prompt_token_ids, - max_new, - events, - } = w.req; + /// Preempt the most recently admitted running request: free its KV and + /// put it back at the head of the queue to recompute (its output so far + /// becomes part of the prompt). + fn preempt_last(&mut self, p: &EngineParams) { + let Some(r) = self.running.pop() else { + return; + }; + self.release(&r, p); + self.preemptions += 1; + let mut seq = r.seq_prompt; + seq.extend_from_slice(&r.output_ids[r.resumed_output as usize..]); + let (keys, rolling_hash, _) = prompt_blocks(&seq, p.block_size.max(1) as usize); + let len = seq.len() as u32; + let cached = if p.prefix_cache { + (self.pool.match_prefix(&keys) as u32 * p.block_size).min(len) + } else { + 0 + }; + self.waiting.push_front(WaitingReq { + req: NewRequest { + request_id: r.request_id, + prompt_token_ids: seq, + max_new: r.max_new, + events: r.events, + }, + keys, + rolling_hash, + prompt_tokens: r.prompt_tokens, + uncached_tokens: len - cached, + oracle_tokens: 0, + enqueued_at: Instant::now(), + resume: Some(Resume { + generated: r.generated, + output_ids: r.output_ids, + cached_tokens: r.cached_tokens, + token_seed: r.token_seed, + }), + }); + } - if cached_blocks > 0 { - self.cache.touch(&block_keys[..cached_blocks]); - } - // Allocate + announce the uncached prompt blocks (resident during prefill). - let mut prev = cached_blocks.checked_sub(1).map(|i| block_keys[i]); - for j in cached_blocks..block_keys.len() { - let key = block_keys[j]; - if p.prefix_cache && self.cache.insert(key) { - let toks = prompt_token_ids[j * block_size..(j + 1) * block_size].to_vec(); - kv.push(self.stored_event(key, toks, prev, p.block_size)); + /// Drop every KV block the request holds. With prefix caching the full + /// blocks stay cached and become evictable (no events); without it they + /// were private and simply free. + fn release(&mut self, r: &RunningReq, p: &EngineParams) { + if p.prefix_cache { + self.pool.unref_request(&r.held); + } else { + self.pool.release_anonymous(r.held.len() as u64); + } + if r.partial { + self.pool.release_anonymous(1); + } + } + + /// Compute `chunk` more prompt tokens of `running[idx]`, registering the + /// blocks the chunk completes. Returns the newly cached blocks. + fn apply_prefill(&mut self, idx: usize, chunk: u32, p: &EngineParams) -> Vec { + let bs = p.block_size.max(1) as usize; + let mut stored = Vec::new(); + let r = &mut self.running[idx]; + let start = r.computed as usize; + let end = start + chunk as usize; + r.pending.reserve(bs); + for pos in start..end { + r.pending.push(r.seq_prompt[pos]); + if r.pending.len() == bs { + let index = r.held.len(); + let key = r.seq_keys[index]; + let tokens = std::mem::take(&mut r.pending); + r.held.push(key); + if p.prefix_cache && self.pool.register(key) { + stored.push(Completed { index, key, tokens }); } - prev = Some(key); } + } + let r = &mut self.running[idx]; + r.computed = end as u32; + r.partial = !r.pending.is_empty(); + stored + } - let resolved_max_new = if max_new == 0 { - p.max_new_default - } else { - max_new - }; - let sample = if w.prompt_tokens > 0 { - f64::from(cached_tokens) / f64::from(w.prompt_tokens) - } else { - 0.0 - }; - self.cache_hit_ewma = ewma(self.cache_hit_ewma, sample, 0.2); - - self.running.push(RunningReq { - events, - prompt_tokens: w.prompt_tokens, - cached_tokens, - max_new: resolved_max_new, - generated: 0, - prefill_remaining: uncached, - rolling_hash: rolling, - prev_block_key: prev, - pending_block: pending, - output_ids: Vec::new(), - token_seed: fnv_hash_str(&request_id) as u32, - }); + /// Generate one token for `running[idx]`, registering a block when the + /// tail fills. Returns the newly cached blocks. + fn emit_token( + &mut self, + idx: usize, + sends: &mut Vec<(mpsc::UnboundedSender, GenEvent)>, + p: &EngineParams, + ) -> Vec { + let bs = p.block_size.max(1) as usize; + let mut stored = Vec::new(); + let r = &mut self.running[idx]; + let token_id = next_token(r); + r.generated += 1; + r.output_ids.push(token_id); + r.rolling_hash = fnv_step(r.rolling_hash, token_id); + r.pending.reserve(bs); + r.pending.push(token_id); + if r.pending.len() == bs { + let index = r.held.len(); + let key = r.rolling_hash; + let tokens = std::mem::take(&mut r.pending); + r.held.push(key); + if p.prefix_cache && self.pool.register(key) { + stored.push(Completed { index, key, tokens }); + } } + let r = &mut self.running[idx]; + r.partial = !r.pending.is_empty(); + sends.push(( + r.events.clone(), + GenEvent::Token { + token_id, + prompt_tokens: r.prompt_tokens, + cached_tokens: r.cached_tokens, + }, + )); + stored } - /// Evict LRU blocks once KV usage crosses the high watermark. - fn evict(&mut self, p: &EngineParams, kv: &mut Vec) { - if !p.prefix_cache { + fn push_removed( + &mut self, + kv: &mut Vec, + freed: Vec, + p: &EngineParams, + ) { + if freed.is_empty() || !p.prefix_cache { return; } - let b = p.block_size as u64; - let partial: u64 = self - .running - .iter() - .map(|r| r.pending_block.len() as u64) - .sum(); - let mut blocks = self.cache.len() as u64; - let high = (p.kv_capacity_tokens as f64 * p.kv_high_watermark) as u64; - if blocks * b + partial <= high { + self.kv_event_id += 1; + kv.push(common::KvCacheEvent { + event_id: self.kv_event_id, + data: Some(common::kv_cache_event::Data::Removed( + common::KvBlocksRemoved { + block_hashes: freed.into_iter().map(|k| k as i64).collect(), + cache_level: None, + ..Default::default() + }, + )), + }); + } + + /// `Stored` events for the blocks `running[idx]` completed this pass: one + /// per contiguous run, chained to the block before the run's first. + fn push_stored( + &mut self, + kv: &mut Vec, + idx: usize, + completed: Vec, + p: &EngineParams, + ) { + if completed.is_empty() || !p.prefix_cache { return; } - let low = (p.kv_capacity_tokens as f64 * p.kv_low_watermark) as u64; - let mut removed: Vec = Vec::new(); - while blocks * b + partial > low { - match self.cache.evict_one() { - Some(key) => { - removed.push(key as i64); - blocks -= 1; - } - None => break, + let held = &self.running[idx].held; + let mut runs: Vec<(Option, Vec)> = Vec::new(); + let mut last_index: Option = None; + for c in completed { + let contiguous = last_index.is_some_and(|prev| prev + 1 == c.index); + if !contiguous || runs.is_empty() { + let parent = c.index.checked_sub(1).map(|pos| held[pos]); + runs.push((parent, Vec::new())); } + last_index = Some(c.index); + runs.last_mut() + .expect("run exists") + .1 + .push(common::KvBlock { + block_hash: c.key as i64, + token_ids: c.tokens, + block_size: p.block_size as i32, + lora_id: None, + cache_level: None, + ..Default::default() + }); } - if !removed.is_empty() { + for (parent, blocks) in runs { self.kv_event_id += 1; kv.push(common::KvCacheEvent { event_id: self.kv_event_id, - data: Some(common::kv_cache_event::Data::Removed( - common::KvBlocksRemoved { - block_hashes: removed, - cache_level: None, + data: Some(common::kv_cache_event::Data::Stored( + common::KvBlocksStored { + blocks, + parent_block_hash: parent.map(|k| k as i64), + ..Default::default() }, )), }); } } +} - fn stored_event( - &mut self, - key: u64, - token_ids: Vec, - parent: Option, - block_size: u32, - ) -> common::KvCacheEvent { - self.kv_event_id += 1; - common::KvCacheEvent { - event_id: self.kv_event_id, - data: Some(common::kv_cache_event::Data::Stored( - common::KvBlocksStored { - blocks: vec![common::KvBlock { - block_hash: key as i64, - token_ids, - block_size: block_size as i32, - lora_id: None, - cache_level: None, - }], - parent_block_hash: parent.map(|k| k as i64), - }, - )), +/// Which backend's load report the transports imitate. The vLLM servicer +/// fills only `num_running_reqs`, `num_waiting_reqs`, `token_usage` and the +/// maxima; the gateway's expected-wait then uses its default throughput and +/// `waiting_reqs × mean prefill` instead of the queued token-work and live +/// throughput the mock knows. Matching that makes routing on the mock agree +/// with routing on a vLLM fleet. +#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)] +pub enum LoadsLike { + /// Everything the simulator knows (SGLang-style report). + #[default] + Mock, + /// Only what the vLLM servicer reports. + Vllm, +} + +impl std::str::FromStr for LoadsLike { + type Err = String; + + fn from_str(s: &str) -> Result { + match s { + "mock" => Ok(Self::Mock), + "vllm" => Ok(Self::Vllm), + other => Err(format!("--loads-like must be mock|vllm, got {other}")), } } } impl LoadSnapshot { + /// The snapshot as a backend of kind `like` would report it. + pub(crate) fn as_reported_by(&self, like: LoadsLike) -> Self { + match like { + LoadsLike::Mock => self.clone(), + LoadsLike::Vllm => Self { + num_waiting_uncached_tokens: 0, + num_used_tokens: 0, + gen_throughput: 0.0, + cache_hit_rate: 0.0, + ..self.clone() + }, + } + } + fn idle(p: &EngineParams) -> Self { Self { num_running_reqs: 0, @@ -757,6 +2016,9 @@ impl LoadSnapshot { token_usage: 0.0, gen_throughput: 0.0, cache_hit_rate: 0.0, + num_cached_blocks: 0, + num_preemptions: 0, + num_kv_batches: 0, } } } @@ -845,27 +2107,127 @@ mod tests { ) } - /// Run steps until the given request emits its first Token, returning the - /// accumulated simulated time (TTFT) and step count. + /// Run passes until the given request emits its first Token, returning the + /// accumulated simulated time (TTFT) and pass count. fn run_to_first_token( st: &mut SchedulerState, p: &EngineParams, rx: &mut mpsc::UnboundedReceiver, ) -> (Duration, u32) { let mut total = Duration::ZERO; - for _ in 0..100_000 { + for n in 1..100_000 { let step = st.step(p); total += step.duration; for (tx, ev) in step.sends { let _ = tx.send(ev); } if let Ok(GenEvent::Token { .. }) = rx.try_recv() { - return (total, 1); + return (total, n); } } panic!("no token produced"); } + fn stored_events(batch: &common::KvEventBatch) -> Vec<&common::KvBlocksStored> { + batch + .events + .iter() + .filter_map(|e| match &e.data { + Some(common::kv_cache_event::Data::Stored(s)) => Some(s), + _ => None, + }) + .collect() + } + + fn has_removed(batch: &common::KvEventBatch) -> bool { + batch + .events + .iter() + .any(|e| matches!(e.data, Some(common::kv_cache_event::Data::Removed(_)))) + } + + #[test] + fn reset_clears_cache_and_publishes_cleared() { + let p = EngineParams::default(); + let mut st = SchedulerState::new(); + let (r, rx) = req("a", vec![3; 64], 2); + st.enqueue(r, &p); + let step = st.step(&p); + assert_eq!( + step.inserted.len(), + 4, + "64 tokens store four 16-token blocks" + ); + assert!(!step.cleared); + assert_ne!(st.pool.cached(), 0); + + st.reset(); + assert!(!st.is_idle(), "a pending reset keeps the actor stepping"); + let step = st.step(&p); + assert!(step.cleared, "the reset pass reports the clear"); + assert!( + step.batch.as_ref().is_some_and(|b| b + .events + .iter() + .any(|e| matches!(e.data, Some(common::kv_cache_event::Data::Cleared(_))))), + "the reset publishes AllBlocksCleared to subscribers" + ); + assert!( + step.batch.as_ref().is_some_and(|b| matches!( + b.events[0].data, + Some(common::kv_cache_event::Data::Cleared(_)) + )), + "the clear precedes the recompute's stored events" + ); + drop(rx); + } + + #[test] + fn admitted_records_carry_cached_and_oracle_tokens() { + let p = EngineParams::default(); + let mut st = SchedulerState::new(); + let (r1, rx1) = req("first", vec![9; 64], 1); + st.enqueue_with_oracle(r1, &p, 48); + let step = st.step(&p); + assert_eq!(step.admitted.len(), 1); + assert_eq!(step.admitted[0].prompt_tokens, 64); + assert_eq!( + step.admitted[0].cached_tokens, 0, + "a cold cache serves nothing" + ); + assert_eq!( + step.admitted[0].oracle_tokens, 48, + "the oracle is what the fleet could have served" + ); + + // The same prompt again: every block is cached, and the last one is + // recomputed anyway (vLLM's rule), so 48 of 64 tokens are served. + let (r2, rx2) = req("second", vec![9; 64], 1); + st.enqueue_with_oracle(r2, &p, 0); + let step = st.step(&p); + let admitted = step + .admitted + .iter() + .find(|a| a.request_id == "second") + .expect("second request admitted"); + assert_eq!(admitted.cached_tokens, 48); + assert_eq!( + admitted.oracle_tokens, 48, + "the oracle is never below what the worker actually served" + ); + drop((rx1, rx2)); + } + + #[test] + fn block_keys_are_deterministic_and_prefix_stable() { + let a = Engine::block_keys(&[1, 2, 3, 4, 5, 6, 7, 8], 4); + let b = Engine::block_keys(&[1, 2, 3, 4, 9, 9, 9, 9], 4); + assert_eq!(a.len(), 2); + assert_eq!(a[0], b[0], "a shared first block hashes the same"); + assert_ne!(a[1], b[1], "a different second block hashes differently"); + assert_eq!(a, Engine::block_keys(&[1, 2, 3, 4, 5, 6, 7, 8], 4)); + } + #[test] fn ttft_scales_with_uncached_prompt_length() { let p = EngineParams { @@ -888,11 +2250,26 @@ mod tests { ); } + #[test] + fn pass_time_follows_the_polynomials() { + let p = EngineParams::default(); + let mut st = SchedulerState::new(); + let (r, rx) = req("a", vec![7; 1024], 4); + st.enqueue(r, &p); + // One pass prefills 1024 uncached tokens: 16.50 + 15.55 + 0.44 ms. + let prefill = st.step(&p).duration.as_secs_f64() * 1000.0; + assert!((prefill - 32.49).abs() < 0.1, "prefill pass {prefill} ms"); + // A lone decoder at ~0 utilisation: 5.74 ms plus a sliver of 54u. + let decode = st.step(&p).duration.as_secs_f64() * 1000.0; + assert!((decode - 5.85).abs() < 0.1, "decode pass {decode} ms"); + drop(rx); + } + #[test] fn itl_grows_with_batch_size() { let p = EngineParams { + timing: TimingModel::linear(), prefix_cache: false, - prefill_chunk_tokens: 1_000_000, // finish prefill in one step ..Default::default() }; @@ -907,7 +2284,7 @@ mod tests { rxs.push(rx); } st.step(&p); // admit + prefill + first tokens - let duration = st.step(&p).duration; // a pure decode step + let duration = st.step(&p).duration; // a pure decode pass drop(rxs); duration }; @@ -916,8 +2293,66 @@ mod tests { let many = decode_step_duration(64); assert!( many > one, - "decode step should be slower with a bigger batch: one={one:?} many={many:?}" + "decode pass should be slower with a bigger batch: one={one:?} many={many:?}" + ); + } + + #[test] + fn a_calibrated_decode_reference_keeps_step_time_when_the_pool_shrinks() { + // The same 2048-token decoder: against a 4096-token pool the step + // reads u = 0.5; with the pool cut to 2560 tokens but the decode + // calibrated at 4096, the step must cost the same as before (the + // pool changed room, not the engine's speed); without the reference + // the smaller pool would read u = 0.8 and decode slower. + let full = EngineParams { + kv_capacity_tokens: 4096, + block_size: 16, + ..Default::default() + }; + let small_pinned = EngineParams { + kv_capacity_tokens: 2560, + decode_reference_tokens: Some(4096), + ..full.clone() + }; + let small_unpinned = EngineParams { + decode_reference_tokens: None, + ..small_pinned.clone() + }; + let step_time = |p: &EngineParams| { + let mut st = SchedulerState::new(); + let (r, _rx) = req("big", vec![1; 2048], 8); + st.enqueue(r, p); + st.step(p); + st.step(p).duration + }; + assert_eq!(step_time(&full), step_time(&small_pinned)); + assert!(step_time(&small_unpinned) > step_time(&full)); + } + + #[test] + fn decode_slows_with_kv_utilisation() { + // Polynomial decode depends on the decoding requests' context over + // capacity, not on the batch width. + let p = EngineParams { + kv_capacity_tokens: 4096, + block_size: 16, + ..Default::default() + }; + let mut st = SchedulerState::new(); + let (r, rx) = req("big", vec![1; 2048], 8); + st.enqueue(r, &p); + st.step(&p); + let busy = st.step(&p).duration; + let mut st2 = SchedulerState::new(); + let (r2, rx2) = req("small", vec![1; 16], 8); + st2.enqueue(r2, &p); + st2.step(&p); + let light = st2.step(&p).duration; + assert!( + busy > light, + "u=0.5 should decode slower than u~0: {busy:?} vs {light:?}" ); + drop((rx, rx2)); } #[test] @@ -939,7 +2374,8 @@ mod tests { } } - // Second request shares the whole prompt prefix. + // Second request shares the whole prompt: every block is cached, the + // last one is recomputed anyway, so 12 of 16 tokens come from cache. let (r2, mut rx2) = req("second", prompt, 2); st.enqueue(r2, &p); st.step(&p); // admission computes the cache hit @@ -953,7 +2389,10 @@ mod tests { } while let Ok(ev) = rx2.try_recv() { if let GenEvent::Token { cached_tokens, .. } = ev { - assert_eq!(cached_tokens, 16, "full prompt prefix should be cached"); + assert_eq!( + cached_tokens, 12, + "all but the last block served from cache" + ); saw_cached = true; } } @@ -983,7 +2422,295 @@ mod tests { } #[test] - fn prompt_blocks_emit_chained_kv_events() { + fn token_budget_chunks_prefill_across_passes() { + let p = EngineParams { + max_batched_tokens: 1000, + prefix_cache: false, + ..Default::default() + }; + let mut st = SchedulerState::new(); + let (r, mut rx) = req("long", vec![7; 2500], 1); + st.enqueue(r, &p); + let (_, passes) = run_to_first_token(&mut st, &p, &mut rx); + assert_eq!( + passes, 3, + "2500 tokens at a 1000-token budget take three passes" + ); + } + + #[test] + fn calibration_file_is_read_in_either_spelling() { + let v: serde_json::Value = serde_json::from_str( + r#"{"prefill_ms": {"a": 20.0, "b": 0.01, "c": 1e-7}, "decode_ms": {"d": 7.0, "e": 40.0, "f": -10.0}, + "kv_capacity_blocks": 6144, "block_size": 64, "request_overhead_ms": 12.5, "engine": "vllm"}"#, + ) + .unwrap(); + let c = Calibration::from_value(&v).unwrap(); + assert_eq!(c.prefill, [20.0, 0.01, 1e-7]); + assert_eq!(c.decode, [7.0, 40.0, -10.0]); + assert_eq!(c.kv_capacity_tokens, Some(6144 * 64)); + assert_eq!(c.block_size, Some(64)); + assert_eq!(c.request_overhead_ms, 12.5); + assert_eq!( + TimingModel::fitted(&c), + TimingModel::Polynomial { + prefill: [20.0, 0.01, 1e-7], + decode: [7.0, 40.0, -10.0] + } + ); + + let arrays: serde_json::Value = serde_json::from_str( + r#"{"prefill": [16.5, 0.015, 4e-7], "decode": [5.7, 54.0, -25.7], "kv_capacity_tokens": 393216}"#, + ) + .unwrap(); + let c = Calibration::from_value(&arrays).unwrap(); + assert_eq!(c.kv_capacity_tokens, Some(393_216)); + assert_eq!(c.block_size, None); + assert_eq!(c.request_overhead_ms, 0.0); + + let bad: serde_json::Value = serde_json::from_str(r#"{"decode": [1, 2, 3]}"#).unwrap(); + assert!(Calibration::from_value(&bad) + .unwrap_err() + .contains("prefill")); + } + + #[test] + fn calibrated_prefill_uses_the_table_for_one_request_and_the_pass_form_for_a_batch() { + let model = TimingModel::Calibrated { + table: vec![ + (1.0, 24.0), + (128.0, 37.0), + (512.0, 82.0), + (1024.0, 75.0), + (4096.0, 98.0), + (8192.0, 106.0), + (16384.0, 211.0), + ], + pass_intercept_ms: 35.0, + pass_ms_per_token: 0.0095, + decode: TimingModel::POLY_DECODE, + }; + let close = |a: f64, b: f64| (a - b).abs() < 0.5; + // A lone request: its interpolated table value, whatever the pass form says. + assert!(close(model.prefill_pass_ms(1024, 1024), 75.0)); + assert!( + close(model.prefill_pass_ms(768, 768), 78.5), + "{}", + model.prefill_pass_ms(768, 768) + ); + assert!( + close(model.prefill_pass_ms(64, 64), 30.45), + "{}", + model.prefill_pass_ms(64, 64) + ); + // Batched: four 1024-token prompts take the pass form once it exceeds the plateau. + assert!(close( + model.prefill_pass_ms(4096, 1024), + 75.0_f64.max(35.0 + 0.0095 * 4096.0) + )); + assert!(close( + model.prefill_pass_ms(16384, 1024), + 35.0 + 0.0095 * 16384.0 + )); + // One 16k request: its own table point wins over the pass form. + assert!(close(model.prefill_pass_ms(16384, 16384), 211.0)); + // Beyond the table: the last slope. + let slope = (211.0 - 106.0) / (16384.0 - 8192.0); + assert!(close( + model.prefill_pass_ms(20000, 20000), + 211.0 + slope * (20000.0 - 16384.0) + )); + assert!(close(model.prefill_ms(0), 0.0)); + } + + #[test] + fn calibration_table_and_pass_form_are_read() { + let v: serde_json::Value = serde_json::from_str( + r#"{"prefill_points_ms": {"1": {"median_ms": 23.84, "min_ms": 23.27}, "1024": {"median_ms": 74.94}, "128": {"median_ms": 37.33}}, + "prefill_pass_ms": {"intercept_ms": 38.3, "ms_per_token": 0.0102}, + "decode_fit_vs_utilisation_ms": {"d_ms": 3.76, "e_ms_per_u": 15.24, "f_ms_per_u2": -2.87}, + "kv_capacity_tokens": 676128}"#, + ) + .unwrap(); + let c = Calibration::from_value(&v).unwrap(); + assert_eq!( + c.prefill_table.as_deref(), + Some(&[(1.0, 23.84), (128.0, 37.33), (1024.0, 74.94)][..]), + "sorted by tokens" + ); + assert_eq!(c.prefill_pass, Some((38.3, 0.0102))); + assert!(matches!( + TimingModel::fitted(&c), + TimingModel::Calibrated { pass_intercept_ms, .. } if (pass_intercept_ms - 38.3).abs() < 1e-9 + )); + let canonical: serde_json::Value = serde_json::from_str( + r#"{"prefill_table_ms": [[1, 24], [4096, 98]], "decode_ms": [3.76, 15.24, -2.87]}"#, + ) + .unwrap(); + let c = Calibration::from_value(&canonical).unwrap(); + assert_eq!( + c.prefill_table.as_deref(), + Some(&[(1.0, 24.0), (4096.0, 98.0)][..]) + ); + assert_eq!( + c.prefill, [0.0; 3], + "a table alone is a complete prefill model" + ); + assert!(matches!( + TimingModel::fitted(&c), + TimingModel::Calibrated { .. } + )); + } + + #[test] + fn vllm_like_loads_drop_what_the_vllm_servicer_does_not_report() { + let full = LoadSnapshot { + num_running_reqs: 30, + num_waiting_reqs: 2, + num_waiting_uncached_tokens: 20_000, + num_used_tokens: 400_000, + max_total_num_tokens: 676_128, + max_running_requests: 128, + token_usage: 0.59, + gen_throughput: 3_000.0, + cache_hit_rate: 0.4, + num_cached_blocks: 42_000, + num_preemptions: 0, + num_kv_batches: 10, + }; + assert_eq!(full.as_reported_by(LoadsLike::Mock), full); + let v = full.as_reported_by(LoadsLike::Vllm); + assert_eq!((v.num_running_reqs, v.num_waiting_reqs), (30, 2)); + assert_eq!( + (v.max_total_num_tokens, v.max_running_requests), + (676_128, 128) + ); + assert_eq!(v.token_usage, 0.59); + assert_eq!( + (v.num_waiting_uncached_tokens, v.num_used_tokens), + (0, 0), + "no queued token-work or used-token count" + ); + assert_eq!((v.gen_throughput, v.cache_hit_rate), (0.0, 0.0)); + assert_eq!("vllm".parse::(), Ok(LoadsLike::Vllm)); + assert!("other".parse::().is_err()); + } + + #[test] + fn calibrated_model_reproduces_a_measured_batched_sweep() { + // A hardware calibration: the table of single-request medians and the + // pass form fitted to its batched sweep. + let model = TimingModel::Calibrated { + table: vec![ + (1.0, 23.84), + (128.0, 37.33), + (256.0, 48.16), + (512.0, 81.72), + (1024.0, 74.94), + (2048.0, 76.98), + (3072.0, 76.44), + (4096.0, 98.39), + (6144.0, 103.1), + (8192.0, 105.73), + (12288.0, 157.74), + (16384.0, 211.04), + ], + pass_intercept_ms: 36.12, + pass_ms_per_token: 0.01028, + decode: [3.7646, 15.2448, -2.8663], + }; + let within = |value: f64, lo: f64, hi: f64| (lo..=hi).contains(&value); + // Measured: 4x1024 in 79.6-80.8 ms, 8x1024 in 114.8-121.8 ms, 16x1024 in 200-210 ms. + assert!(within(model.prefill_pass_ms(4096, 1024), 74.0, 86.0)); + assert!(within(model.prefill_pass_ms(8192, 1024), 110.0, 126.0)); + assert!(within(model.prefill_pass_ms(16384, 1024), 195.0, 215.0)); + // The single-request plateau: 512-3072 tokens cost 75-82 ms alone. + assert!(within(model.prefill_pass_ms(2048, 2048), 74.0, 82.0)); + // 8x4096 is two passes of 16384 on a 16384-token budget: 409 ms by the + // pass form against ~335 ms measured; the second pass is cheaper on the + // hardware than the form says (recorded, not matched). + assert!(within( + 2.0 * model.prefill_pass_ms(16384, 4096), + 390.0, + 430.0 + )); + } + + #[test] + fn calibration_reads_the_gpu_harness_layout() { + // A calibration file as the GPU harness writes it, abridged. + let harness: serde_json::Value = serde_json::from_str( + r#"{"target": "127.0.0.1:20061", "kv_capacity_tokens": 676128, + "prefill_points_ms": {"1": {"median_ms": 23.84}}, + "prefill_fit_ms": {"a_ms": 35.24, "b_ms_per_token": 0.005654, "c_ms_per_token2": 2.076e-07, + "fixed_overhead_ms": 23.84, "residual_ms": [-22.48]}, + "decode_fit_vs_utilisation_ms": {"d_ms": 3.7645, "e_ms_per_u": 15.2448, "f_ms_per_u2": -2.8663, + "u": "KV tokens in use / capacity (computed)"}, + "decode_fit_vs_batch_ms": {"d_ms": 3.8, "e_ms_per_seq": 0.028, "f_ms_per_seq2": 1.3e-05}}"#, + ) + .unwrap(); + let c = Calibration::from_value(&harness).unwrap(); + assert_eq!(c.prefill, [35.24, 0.005654, 2.076e-07]); + assert_eq!(c.decode, [3.7645, 15.2448, -2.8663]); + assert_eq!(c.kv_capacity_tokens, Some(676_128)); + assert_eq!(c.block_size, None); + assert_eq!( + c.request_overhead_ms, 0.0, + "the one-token TTFT is in the prefill intercept already" + ); + } + + #[tokio::test] + async fn request_overhead_delays_the_stream_without_stretching_it() { + let quick = Engine::spawn(EngineParams::default()); + let slow = Engine::spawn(EngineParams { + request_overhead_ms: 300.0, + ..Default::default() + }); + async fn first_and_second(engine: &Engine) -> (Duration, Duration) { + let (r, mut rx) = req("a", vec![1; 32], 3); + let t0 = Instant::now(); + engine.submit(r); + let mut times = Vec::new(); + while times.len() < 2 { + let ev = tokio::time::timeout(Duration::from_secs(5), rx.recv()) + .await + .expect("events") + .expect("open"); + if matches!(ev, GenEvent::Token { .. }) { + times.push(t0.elapsed()); + } + } + (times[0], times[1] - times[0]) + } + let (ttft_quick, itl_quick) = first_and_second(&quick).await; + let (ttft_slow, itl_slow) = first_and_second(&slow).await; + assert!( + ttft_slow >= ttft_quick + Duration::from_millis(250), + "the overhead adds to TTFT: {ttft_quick:?} vs {ttft_slow:?}" + ); + assert!( + itl_slow < itl_quick + Duration::from_millis(50), + "and not to the inter-token gap: {itl_quick:?} vs {itl_slow:?}" + ); + } + + #[test] + fn batches_are_stamped_with_their_creation_time() { + let p = EngineParams::default(); + let mut st = SchedulerState::new(); + let (r, _rx) = req("x", vec![1; 32], 1); + st.enqueue(r, &p); + let batch = st.step(&p).batch.expect("a batch"); + let age = unix_seconds() - batch.timestamp; + assert!( + (0.0..5.0).contains(&age), + "fresh wall-clock stamp: {age} s old" + ); + } + + #[test] + fn prompt_blocks_emit_one_chained_stored_event() { let p = EngineParams { block_size: 4, ..Default::default() @@ -994,59 +2721,424 @@ mod tests { let step = st.step(&p); let batch = step.batch.expect("stored events expected"); assert_eq!(batch.sequence_number, 1); - let stored: Vec<_> = batch - .events - .iter() - .filter_map(|e| match &e.data { - Some(common::kv_cache_event::Data::Stored(s)) => Some(s), - _ => None, - }) - .collect(); - assert_eq!(stored.len(), 2, "two prompt blocks"); + let stored = stored_events(&batch); + assert_eq!(stored.len(), 1, "contiguous blocks share one event"); + assert_eq!(stored[0].blocks.len(), 2, "two prompt blocks"); assert!( stored[0].parent_block_hash.is_none(), "first block has no parent" ); - let first_hash = stored[0].blocks[0].block_hash; - assert_eq!( - stored[1].parent_block_hash, - Some(first_hash), - "second block chains to the first" - ); + assert_eq!(stored[0].blocks[0].token_ids, vec![0, 1, 2, 3]); + assert_eq!(stored[0].blocks[1].token_ids, vec![4, 5, 6, 7]); + + // A longer prompt sharing the prefix stores only its tail, chained to + // the last cached block. + let (r2, _rx2) = req("y", (0..12).collect(), 1); + st.enqueue(r2, &p); + let step = st.step(&p); + let batch = step.batch.expect("stored events expected"); + let stored = stored_events(&batch); + assert_eq!(stored.len(), 1); + assert_eq!(stored[0].blocks.len(), 1, "only the third block is new"); + let second_key = Engine::block_keys(&(0..8).collect::>(), 4)[1]; + assert_eq!(stored[0].parent_block_hash, Some(second_key as i64)); } #[test] - fn kv_pressure_evicts_and_emits_removed() { - // Tiny KV so a couple of prompts overflow it. + fn a_finished_request_frees_its_blocks_tail_first() { + // Request A held blocks 1-4; request B shared the prefix 1-2 and added 5. + let mut pool = BlockPool { + allocated: 5, + ..Default::default() + }; + for h in 1..=4u64 { + assert!(pool.register(h)); + } + pool.hit(1); + pool.hit(2); + assert!(pool.register(5)); + pool.unref_request(&[1, 2, 3, 4]); + pool.unref_request(&[1, 2, 5]); + // Tails go first, the shared prefix head last: 4, 3 (A's tail), then + // 5 (B's tail), then 2, then 1. + let order: Vec = std::iter::from_fn(|| pool.evict_lru()).collect(); + assert_eq!(order, vec![4, 3, 5, 2, 1]); + } + + #[test] + fn kv_pressure_evicts_lru_and_emits_removed_before_stored() { + // 16 blocks of 4 tokens; each prompt takes 4 blocks plus a tail slot. let p = EngineParams { block_size: 4, kv_capacity_tokens: 64, - kv_high_watermark: 0.5, - kv_low_watermark: 0.25, max_running: 64, - prefill_chunk_tokens: 1_000_000, ..Default::default() }; let mut st = SchedulerState::new(); + let mut rxs = Vec::new(); for i in 0..8 { - // Distinct prompts so each contributes its own blocks. let base = (i as u32) * 1000; - let (r, _rx) = req(&format!("r{i}"), (base..base + 16).collect(), 1); + let (r, rx) = req(&format!("r{i}"), (base..base + 16).collect(), 1); st.enqueue(r, &p); + rxs.push(rx); } let mut saw_removed = false; for _ in 0..50 { let step = st.step(&p); if let Some(batch) = step.batch { - if batch - .events - .iter() - .any(|e| matches!(e.data, Some(common::kv_cache_event::Data::Removed(_)))) - { + if has_removed(&batch) { saw_removed = true; + let first_removed = batch + .events + .iter() + .position(|e| { + matches!(e.data, Some(common::kv_cache_event::Data::Removed(_))) + }) + .expect("removed present"); + let first_stored = batch + .events + .iter() + .position(|e| { + matches!(e.data, Some(common::kv_cache_event::Data::Stored(_))) + }) + .expect("a stored event follows the eviction"); + assert!( + first_removed < first_stored, + "evictions precede the blocks they made room for" + ); } } } assert!(saw_removed, "KV pressure should emit a removed event"); + assert!( + st.pool.allocated <= 16, + "the pool never holds more than its capacity: {}", + st.pool.allocated + ); + drop(rxs); + } + + #[test] + fn full_isl_reservation_admits_only_prompts_that_fit_and_blocks_head_of_line() { + // 16 blocks of 4 tokens. Two 40-token prompts need 10 blocks (+1 for the + // first token) each: the second does not fit beside the first, and a + // small prompt behind it waits too (head-of-line), as vLLM does. + let p = EngineParams { + block_size: 4, + kv_capacity_tokens: 64, + max_batched_tokens: 60, + ..Default::default() + }; + let mut st = SchedulerState::new(); + let (a, _ra) = req("a", (0..40).collect(), 1); + let (b, _rb) = req("b", (100..140).collect(), 1); + let (c, _rc) = req("c", (200..208).collect(), 1); + st.enqueue(a, &p); + st.enqueue(b, &p); + st.enqueue(c, &p); + let step = st.step(&p); + assert_eq!(step.admitted.len(), 1, "only the first prompt fits"); + assert_eq!( + step.snapshot.num_waiting_reqs, 2, + "the small prompt waits behind the big one" + ); + assert!(st.pool.allocated <= 16); + + // Without the reservation the old behaviour admits chunks of everything. + let loose = EngineParams { + reserve_full_isl: false, + ..p.clone() + }; + let mut st = SchedulerState::new(); + let (a, _ra) = req("a", (0..40).collect(), 1); + let (b, _rb) = req("b", (100..140).collect(), 1); + st.enqueue(a, &loose); + st.enqueue(b, &loose); + let step = st.step(&loose); + assert_eq!( + step.admitted.len(), + 2, + "chunk-only reservation admits both long prompts" + ); + } + + #[test] + fn full_isl_admission_gates_on_the_whole_prompt_but_allocates_per_chunk() { + // A 2500-token prompt against a 1000-token pass budget: the gate reads + // the whole prompt, the first pass allocates its chunk only, the next + // chunks allocate as they run, and nothing is held beyond what is + // computed (plus the first token's slot once the prefill completes). + let p = EngineParams { + max_batched_tokens: 1000, + block_size: 16, + kv_capacity_tokens: 16 * 400, + ..Default::default() + }; + let mut st = SchedulerState::new(); + let (r, mut rx) = req("long", vec![7; 2500], 3); + st.enqueue(r, &p); + st.step(&p); + assert_eq!( + st.pool.allocated, + blocks_for(1000, 16), + "the first pass allocates its chunk, not the whole prompt" + ); + let (_, passes) = run_to_first_token(&mut st, &p, &mut rx); + assert_eq!(passes, 2, "two more passes finish the prefill"); + assert_eq!( + st.pool.allocated, + blocks_for(2501, 16), + "after the prefill the prompt and the first token's slot are held" + ); + for _ in 0..10 { + st.step(&p); + } + assert!(st.is_idle()); + assert_eq!( + st.pool.allocated as usize, + st.pool.cached(), + "once done, only cached blocks remain allocated (no leaked reservation)" + ); + } + + #[test] + fn lifo_preemption_recomputes_and_completes() { + // 16 blocks of 4 tokens. Two requests of 24 prompt tokens generating + // 40 tokens each need 32 blocks between them: the later one is + // preempted when the earlier one needs room, recomputes, and both + // still deliver every token exactly once. + let p = EngineParams { + block_size: 4, + kv_capacity_tokens: 64, + max_running: 64, + ..Default::default() + }; + let mut st = SchedulerState::new(); + let (r1, mut rx1) = req("first", (0..24).collect(), 40); + let (r2, mut rx2) = req("second", (100..124).collect(), 40); + st.enqueue(r1, &p); + st.enqueue(r2, &p); + for _ in 0..500 { + let step = st.step(&p); + for (tx, ev) in step.sends { + let _ = tx.send(ev); + } + if st.is_idle() { + break; + } + } + assert!(st.is_idle(), "both requests finish"); + assert!( + st.preemptions >= 1, + "KV exhaustion preempted the later request" + ); + for rx in [&mut rx1, &mut rx2] { + let mut tokens = 0; + let mut done = None; + while let Ok(ev) = rx.try_recv() { + match ev { + GenEvent::Token { .. } => tokens += 1, + GenEvent::Done { + completion_tokens, .. + } => done = Some(completion_tokens), + } + } + assert_eq!(tokens, 40, "every token delivered once"); + assert_eq!(done, Some(40)); + } + assert_eq!( + st.pool.allocated, + st.pool.free_lru.len() as u64, + "all blocks idle" + ); + } + + /// A live engine for the hook tests, with one request helper. + fn live() -> Engine { + Engine::spawn(EngineParams::default()) + } + + fn submit(engine: &Engine, id: &str, prompt: Vec) -> mpsc::UnboundedReceiver { + let (r, rx) = req(id, prompt, 1); + engine.submit(r); + rx + } + + async fn next_batch( + stream: &mut KvEventStream, + within: Duration, + ) -> Option { + tokio::time::timeout(within, stream.next()) + .await + .ok() + .flatten() + .and_then(Result::ok) + } + + #[tokio::test] + async fn drop_hook_loses_live_batches_but_keeps_them_for_replay() { + let engine = live(); + let mut live_stream = engine.subscribe_kv(0); + engine.fault_drop(1); + let mut rx1 = submit(&engine, "a", vec![1; 64]); + // Let the first request's pass (and its batch) complete before the + // second arrives, so the two batches are distinct. + while let Some(event) = tokio::time::timeout(Duration::from_secs(5), rx1.recv()) + .await + .expect("events") + { + if matches!(event, GenEvent::Done { .. }) { + break; + } + } + let _rx2 = submit(&engine, "b", vec![2; 64]); + let first_live = next_batch(&mut live_stream, Duration::from_secs(5)) + .await + .expect("a live batch"); + assert_eq!( + first_live.sequence_number, 2, + "the first batch was lost on the wire, the second arrives" + ); + let mut replay = engine.subscribe_kv(0); + let replayed = next_batch(&mut replay, Duration::from_secs(1)) + .await + .expect("replay"); + assert_eq!( + replayed.sequence_number, 1, + "the lost batch is still replayable" + ); + let status = engine.fault_status(); + assert_eq!((status.drop_pending, status.dropped_total), (0, 1)); + } + + #[tokio::test] + async fn pause_freezes_the_engine_until_resume() { + let engine = live(); + engine.pause(); + let mut rx = submit(&engine, "a", vec![3; 32]); + assert!( + tokio::time::timeout(Duration::from_millis(200), rx.recv()) + .await + .is_err(), + "no token while paused" + ); + assert!(engine.fault_status().paused); + assert_eq!(engine.load().num_waiting_reqs, 1, "the request queued"); + engine.resume(); + let event = tokio::time::timeout(Duration::from_secs(5), rx.recv()) + .await + .expect("a token after resume") + .expect("stream open"); + assert!(matches!(event, GenEvent::Token { .. })); + assert!(!engine.fault_status().paused); + } + + #[tokio::test] + async fn restart_publisher_starts_the_sequence_over_and_empties_replay() { + let engine = live(); + let mut live_stream = engine.subscribe_kv(0); + let _rx1 = submit(&engine, "a", vec![4; 64]); + let before = next_batch(&mut live_stream, Duration::from_secs(5)) + .await + .expect("a batch"); + assert_eq!(before.sequence_number, 1); + engine.restart_publisher().await; + let _rx2 = submit(&engine, "b", vec![5; 64]); + let after = next_batch(&mut live_stream, Duration::from_secs(5)) + .await + .expect("a batch after the restart"); + assert_eq!(after.sequence_number, 1, "sequence numbers start over"); + let status = engine.fault_status(); + assert_eq!((status.generation, status.restarts), (1, 1)); + let mut replay = engine.subscribe_kv(0); + let replayed = next_batch(&mut replay, Duration::from_secs(1)) + .await + .expect("replay"); + assert_eq!( + replayed.sequence_number, 1, + "only the new generation is replayable" + ); + } + + #[tokio::test] + async fn delay_hook_defers_publishing_but_not_tokens() { + let engine = live(); + let mut live_stream = engine.subscribe_kv(0); + engine.fault_delay_ms(300); + let mut rx = submit(&engine, "a", vec![6; 64]); + let token = tokio::time::timeout(Duration::from_secs(5), rx.recv()) + .await + .expect("a token") + .expect("stream open"); + assert!(matches!(token, GenEvent::Token { .. })); + let token_at = Instant::now(); + let batch = next_batch(&mut live_stream, Duration::from_secs(5)) + .await + .expect("the batch"); + let lag = token_at.elapsed(); + assert!( + lag >= Duration::from_millis(200), + "events trail the pass by the delay: {lag:?}" + ); + let age = unix_seconds() - batch.timestamp; + assert!( + (0.2..30.0).contains(&age), + "the batch keeps its creation time, so the delay is visible as lag: {age} s" + ); + } + + #[tokio::test] + async fn cached_tokens_for_reports_the_engine_truth() { + let engine = live(); + let prompt: Vec = (0..64).collect(); + assert_eq!(engine.cached_tokens_for(&prompt).cached_tokens, 0); + let mut rx = submit(&engine, "a", prompt.clone()); + while let Some(event) = tokio::time::timeout(Duration::from_secs(5), rx.recv()) + .await + .expect("events") + { + if matches!(event, GenEvent::Done { .. }) { + break; + } + } + // The mirror is updated at pass end, just before the token is sent. + let truth = engine.cached_tokens_for(&prompt); + assert_eq!(truth.block_size, 16); + assert_eq!( + (truth.cached_blocks, truth.cached_tokens), + (3, 48), + "all four blocks are cached; the last one is recomputed" + ); + let longer: Vec = (0..80).collect(); + assert_eq!(engine.cached_tokens_for(&longer).cached_tokens, 64); + } + + #[test] + fn prefill_first_pass_holds_decoders() { + let p = EngineParams { + prefill_first: true, + ..Default::default() + }; + let mut st = SchedulerState::new(); + let (r1, mut rx1) = req("decoder", vec![1; 32], 8); + st.enqueue(r1, &p); + st.step(&p); // prefill + first token + st.step(&p); // a decode pass + while rx1.try_recv().is_ok() {} + let (r2, _rx2) = req("arrival", vec![2; 32], 8); + st.enqueue(r2, &p); + let step = st.step(&p); // a pass with prefill: prefill only + for (tx, ev) in step.sends { + let _ = tx.send(ev); + } + assert!( + rx1.try_recv().is_err(), + "the decoder gets no token in a prefill pass" + ); + let step = st.step(&p); // no prefill left: everyone decodes + for (tx, ev) in step.sends { + let _ = tx.send(ev); + } + assert!(rx1.try_recv().is_ok(), "the decoder resumes afterwards"); } } diff --git a/crates/mock_worker/src/grpc.rs b/crates/mock_worker/src/grpc.rs index 187bd60fe7..0faab38b47 100644 --- a/crates/mock_worker/src/grpc.rs +++ b/crates/mock_worker/src/grpc.rs @@ -5,9 +5,10 @@ use std::{ net::{IpAddr, SocketAddr}, pin::Pin, sync::Arc, + time::{Duration, Instant}, }; -use futures::{stream, Stream}; +use futures::{stream, Stream, StreamExt}; use smg_grpc_client::{common_proto as common, tokenspeed_scheduler::tokenspeed_proto as ts}; use tokio::{net::TcpListener, sync::mpsc}; use tokio_stream::wrappers::TcpListenerStream; @@ -19,7 +20,7 @@ use ts::{ use crate::{ config::Config, - engine::{self, Engine, NewRequest}, + engine::{self, Engine, KvEventStream, LoadsLike, NewRequest}, replay::Capture, }; @@ -59,19 +60,39 @@ pub async fn serve_with_listener(cfg: Arc, listener: TcpListener) { }, None => None, }; - // One simulated engine per listener (i.e. per virtual worker). - let engine = cfg.realistic.then(|| Engine::spawn(cfg.engine.clone())); + // One simulated engine per listener (i.e. per virtual worker), registered + // in the process fleet under `grpc:` for the oracle and admin API. + let name = format!("grpc:{}", addr.map(|a| a.port()).unwrap_or(0)); + let engine = cfg + .realistic + .then(|| Engine::spawn_named(cfg.engine.clone(), name, true)); + // The worker's vLLM-wire KV-event publisher, numbered by its port offset. + let publisher = engine.as_ref().and_then(|engine| { + let port = addr.map(|a| a.port())?; + let index = port.checked_sub(cfg.grpc_base_port)?; + // A gRPC worker is a single-rank engine: it publishes as rank 0. + cfg.kv_zmq_for(index, 0) + .map(|kv| crate::kv_zmq::serve(engine.clone(), kv)) + }); let service = MockScheduler { cfg, engine, capture, }; - if let Err(e) = Server::builder() - .add_service(TokenSpeedSchedulerServer::new(service)) - .serve_with_incoming(TcpListenerStream::new(listener)) - .await - { - tracing::error!("grpc worker {addr:?} stopped: {e}"); + let server = async { + if let Err(e) = Server::builder() + .add_service(TokenSpeedSchedulerServer::new(service)) + .serve_with_incoming(TcpListenerStream::new(listener)) + .await + { + tracing::error!("grpc worker {addr:?} stopped: {e}"); + } + }; + match publisher { + Some(publisher) => { + tokio::join!(server, publisher); + } + None => server.await, } } @@ -85,7 +106,6 @@ struct MockScheduler { } type GenStream = Pin> + Send>>; -type KvEventStream = Pin> + Send>>; type TokenizerStream = Pin> + Send>>; @@ -204,8 +224,8 @@ impl TokenSpeedScheduler for MockScheduler { served_model_name: self.cfg.model_id.clone(), model_type: "mock".to_string(), architectures: vec!["MockForCausalLM".to_string()], - max_context_length: 32768, - max_req_input_len: 32768, + max_context_length: self.cfg.context_length.min(i32::MAX as u32) as i32, + max_req_input_len: self.cfg.context_length.min(i32::MAX as u32) as i32, vocab_size: 32000, eos_token_ids: vec![2], pad_token_id: 0, @@ -238,7 +258,9 @@ impl TokenSpeedScheduler for MockScheduler { _request: Request, ) -> Result, Status> { let load = match &self.engine { - Some(engine) => snapshot_to_scheduler_load(&engine.load()), + Some(engine) => { + snapshot_to_scheduler_load(&engine.load().as_reported_by(self.cfg.loads_like)) + } None => ts::SchedulerLoad { dp_rank: 0, num_running_reqs: 0, @@ -270,13 +292,18 @@ impl TokenSpeedScheduler for MockScheduler { request: Request, ) -> Result, Status> { match &self.engine { - // Realistic mode with prefix caching: stream the engine's KV events. + // Realistic mode with prefix caching: stream the engine's KV events, + // each carrying the load record a servicer attaches. Some(engine) if engine.kv_enabled() => { let start = request.into_inner().start_sequence_number; - Ok(Response::new(engine.subscribe_kv(start))) + Ok(Response::new(with_load_records( + engine.clone(), + self.cfg.loads_like, + engine.subscribe_kv(start), + ))) } - // Otherwise Unimplemented makes the gateway's KvEventMonitor give up - // cleanly (no idle per-worker task), exactly as before this RPC existed. + // Otherwise Unimplemented, on which the gateway's KvEventMonitor gives + // up cleanly instead of keeping an idle per-worker task. _ => Err(Status::unimplemented( "mock-worker (KV events require --engine realistic with --prefix-cache true)", )), @@ -382,6 +409,173 @@ fn generate_stream( } /// Map an engine load snapshot to the TokenSpeed `SchedulerLoad` wire type. +/// The load record the engine's `GetLoads` figures make, as reported for +/// `loads_like` (the vLLM shape has no queued token-work). +fn load_record( + snapshot: &engine::LoadSnapshot, + like: LoadsLike, + sample: u64, +) -> common::EngineLoad { + let unsigned = |value: i32| u32::try_from(value).unwrap_or(0); + common::EngineLoad { + running_requests: unsigned(snapshot.num_running_reqs), + waiting_requests: unsigned(snapshot.num_waiting_reqs), + waiting_uncached_tokens: matches!(like, LoadsLike::Mock) + .then(|| unsigned(snapshot.num_waiting_uncached_tokens)), + token_usage: snapshot.token_usage, + gen_throughput: snapshot.gen_throughput, + max_running_requests: unsigned(snapshot.max_running_requests), + age_ms: 0, + sample, + load_only: false, + // What both reports carry beyond the core (the mock has no + // memory, queue, speculative, LoRA or disaggregation sections). + cache_hit_rate: Some(snapshot.cache_hit_rate), + num_used_tokens: Some(snapshot.num_used_tokens), + max_total_num_tokens: Some(snapshot.max_total_num_tokens), + ..Default::default() + } +} + +/// Whether a record moved enough from `last` to be worth a `load_only` +/// batch: the Rust relay's rule (any queue, running or window change, KV +/// usage by half a percent, the rate by 5 % or 50 tokens per second). +fn load_changed(last: &common::EngineLoad, current: &common::EngineLoad) -> bool { + last.running_requests != current.running_requests + || last.waiting_requests != current.waiting_requests + || last.waiting_uncached_tokens != current.waiting_uncached_tokens + || last.max_running_requests != current.max_running_requests + || (last.token_usage - current.token_usage).abs() > 0.005 + || { + let delta = (last.gen_throughput - current.gen_throughput).abs(); + delta > 50.0 || delta > 0.05 * last.gen_throughput.max(current.gen_throughput) + } +} + +/// How often the stream checks the record while no batch flows, the silence +/// after which a heartbeat goes out, and the heartbeat interval once the +/// engine has been idle for two of them: the Rust relay's figures. +const LOAD_TICK: Duration = Duration::from_millis(100); +const HEARTBEAT_INTERVAL: Duration = Duration::from_secs(1); +const HEARTBEAT_BACKOFF: Duration = Duration::from_secs(5); + +struct LoadRecords { + stream: KvEventStream, + engine: engine::Engine, + like: LoadsLike, + last_seq: u64, + last_rank: Option, + last_record: Option, + last_sent_at: Instant, + idle_heartbeats: u32, + sample: u64, +} + +#[expect( + clippy::large_enum_variant, + reason = "the item is matched and moved out at once; boxing it would allocate per batch" +)] +enum Waited { + Item(Option>), + Tick, +} + +impl LoadRecords { + fn record(&mut self, load_only: bool) -> common::EngineLoad { + self.sample += 1; + let snapshot = self.engine.load().as_reported_by(self.like); + let mut record = load_record(&snapshot, self.like, self.sample); + record.load_only = load_only; + // As the Rust relay: telemetry on heartbeats and the first record. + if !load_only && self.sample > 1 { + smg_grpc_client::engine_load::core_only(&mut record); + } + self.last_record = Some(record.clone()); + record + } + + /// A `load_only` batch when the record moved since the last one sent, or + /// after the heartbeat interval of silence (the backoff interval after two + /// unchanged heartbeats); it repeats the last sequence sent, as the + /// gateway expects. + fn load_only_batch(&mut self) -> Option { + let current = load_record(&self.engine.load().as_reported_by(self.like), self.like, 0); + let changed = self + .last_record + .as_ref() + .is_none_or(|last| load_changed(last, ¤t)); + let due_after = if self.idle_heartbeats >= 2 { + HEARTBEAT_BACKOFF + } else { + HEARTBEAT_INTERVAL + }; + if !changed && self.last_sent_at.elapsed() < due_after { + return None; + } + let record = self.record(true); + self.last_sent_at = Instant::now(); + self.idle_heartbeats = if changed { + 0 + } else { + self.idle_heartbeats.saturating_add(1) + }; + Some(common::KvEventBatch { + sequence_number: self.last_seq, + timestamp: engine::unix_seconds(), + events: Vec::new(), + dp_rank: self.last_rank, + snapshot: None, + load: Some(record), + }) + } +} + +/// `stream` with the engine's load record on every batch and `load_only` +/// batches while the engine is quiet, as the Rust servicer's relay sends +/// them (`crates/engine_servicer/src/kv_events.rs`). +fn with_load_records( + engine: engine::Engine, + like: LoadsLike, + stream: KvEventStream, +) -> KvEventStream { + let records = LoadRecords { + stream, + engine, + like, + last_seq: 0, + last_rank: Some(0), + last_record: None, + last_sent_at: Instant::now(), + idle_heartbeats: 0, + sample: 0, + }; + Box::pin(stream::unfold(records, |mut records| async move { + loop { + let waited = tokio::select! { + item = records.stream.next() => Waited::Item(item), + () = tokio::time::sleep(LOAD_TICK) => Waited::Tick, + }; + match waited { + Waited::Item(None) => return None, + Waited::Item(Some(Err(status))) => return Some((Err(status), records)), + Waited::Item(Some(Ok(mut batch))) => { + records.last_seq = batch.sequence_number; + records.last_rank = batch.dp_rank; + records.last_sent_at = Instant::now(); + records.idle_heartbeats = 0; + batch.load = Some(records.record(false)); + return Some((Ok(batch), records)); + } + Waited::Tick => { + if let Some(batch) = records.load_only_batch() { + return Some((Ok(batch), records)); + } + } + } + } + })) +} + fn snapshot_to_scheduler_load(s: &engine::LoadSnapshot) -> ts::SchedulerLoad { ts::SchedulerLoad { dp_rank: 0, diff --git a/crates/mock_worker/src/http.rs b/crates/mock_worker/src/http.rs index 9c2a3df32d..f9e241109a 100644 --- a/crates/mock_worker/src/http.rs +++ b/crates/mock_worker/src/http.rs @@ -37,13 +37,13 @@ use crate::{ }; /// Per-listener HTTP state: shared config plus an optional engine simulator. -pub struct AppState { +struct AppState { cfg: Arc, engine: Option, } /// Build the router serving the mock HTTP worker contract. -pub fn router(state: Arc) -> Router { +fn router(state: Arc) -> Router { Router::new() .route("/health", get(health)) .route("/v1/models", get(models)) @@ -63,8 +63,11 @@ pub async fn serve(cfg: Arc, host: String, port: u16) { return; } }; - // One simulated engine per listener (i.e. per virtual worker). - let engine = cfg.realistic.then(|| Engine::spawn(cfg.engine.clone())); + // One simulated engine per listener (i.e. per virtual worker), registered + // in the process fleet under `http:` for the oracle and admin API. + let engine = cfg + .realistic + .then(|| Engine::spawn_named(cfg.engine.clone(), format!("http:{port}"), true)); let state = Arc::new(AppState { cfg, engine }); // TCP_NODELAY: without it Nagle holds each small SSE frame until the // gateway's delayed ACK (~40ms) arrives, which stalls every streamed @@ -92,14 +95,17 @@ async fn models(State(state): State>) -> Response { "created": 0, "owned_by": "sglang", "root": state.cfg.model_id, - "max_model_len": 32768, + "max_model_len": state.cfg.context_length, }], })) .into_response() } async fn loads(State(state): State>) -> Response { - let load = state.engine.as_ref().map(|e| e.load()); + let load = state + .engine + .as_ref() + .map(|e| e.load().as_reported_by(state.cfg.loads_like)); let value = match load { Some(s) => json!({ "dp_rank": 0, @@ -182,8 +188,8 @@ async fn handle(endpoint: Endpoint, state: Arc, body: Bytes) -> Respon }; } - // Canned mode: a single up-front delay, then a fixed response. Always - // chat-shaped (unchanged) so the existing scale rig is unaffected. + // Canned mode: a single up-front delay, then a fixed chat-shaped response, + // whichever endpoint was called. if !state.cfg.gen_delay.is_zero() { tokio::time::sleep(state.cfg.gen_delay).await; } @@ -264,6 +270,7 @@ async fn realistic_completion( "completion_tokens": completion_tokens, "total_tokens": prompt_tokens + completion_tokens, "cached_tokens": cached_tokens, + "prompt_tokens_details": { "cached_tokens": cached_tokens }, }); Json(match endpoint { Endpoint::Chat => json!({ diff --git a/crates/mock_worker/src/kv_zmq.rs b/crates/mock_worker/src/kv_zmq.rs new file mode 100644 index 0000000000..40b0933f28 --- /dev/null +++ b/crates/mock_worker/src/kv_zmq.rs @@ -0,0 +1,1006 @@ +//! A ZMQ KV-event publisher fed by a simulated engine, on either engine's +//! wire, so the servicers' relays (Rust and Python) can be exercised end to +//! end without a GPU. +//! +//! Common to both (`ZmqEventPublisher` in vLLM's `distributed/kv_events.py` +//! and SGLang's `disaggregation/kv_events.py`): +//! +//! - one PUB socket per worker; one multipart message per pass: +//! `[topic, sequence as u64 big-endian, msgpack payload]`, the sequence +//! counting from 0 per publisher; +//! - the payload is the array-like `EventBatch` `[ts, events, rank]` with +//! tagged-map events (`type`); +//! - optional replay on a ROUTER socket at the next port: a request's last +//! frame is the start sequence (8 bytes big-endian); the reply is every +//! buffered batch from it, then an END marker with the sequence slot set to +//! eight 0xff bytes ((-1) signed); the last `buffer_steps` batches are kept. +//! +//! vLLM wire (`Wire::Vllm`): events in vLLM's field order with its +//! `omit_defaults`: a `BlockStored` carries `block_hashes`, +//! `parent_block_hash`, `token_ids`, `block_size`, `lora_id`, `medium`, +//! `lora_name` (required, nil when unset) and then `group_idx` and +//! `kv_cache_spec_kind`; a `BlockRemoved` carries `block_hashes`, `medium`, +//! `group_idx`; hashes are unsigned 64-bit integers; replay replies are +//! `[routing…, topic, seq, payload]` and `[routing…, b"", END, b""]`. +//! +//! SGLang wire (`Wire::Sglang`): hashes are signed 64-bit integers (SGLang +//! takes the first eight digest bytes signed); a `BlockStored` carries +//! `block_hashes`, `parent_block_hash`, `token_ids`, `block_size`, `lora_id` +//! and nothing else (no medium, group or spec fields), one per radix node +//! (here: per contiguous run of blocks a request completed); a `BlockRemoved` +//! carries one node's hashes (here: one per evicted block); the batch's third +//! slot (`attn_dp_rank`) is nil; the publisher starts with an +//! `AllBlocksCleared` batch, as the scheduler does; replay replies are +//! `[routing…, seq, payload]` and `[routing…, END, b""]`. + +use std::collections::VecDeque; + +use futures::StreamExt; +use rmpv::Value; +use smg_grpc_client::common_proto as common; +use zeromq::{ + prelude::{Socket, SocketRecv, SocketSend}, + PubSocket, RouterSocket, ZmqError, ZmqMessage, +}; + +use crate::engine::{unix_seconds, Engine}; + +/// The engines' end-of-replay marker: `(-1).to_bytes(8, "big", signed=True)`. +const END_SEQ: [u8; 8] = [0xff; 8]; + +/// Which engine's publisher to imitate. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum Wire { + Vllm, + Sglang, +} + +impl std::str::FromStr for Wire { + type Err = String; + + fn from_str(s: &str) -> Result { + match s { + "vllm" => Ok(Self::Vllm), + "sglang" => Ok(Self::Sglang), + other => Err(format!("--kv-events-wire must be vllm|sglang, got {other}")), + } + } +} + +/// Where and how one worker publishes. +#[derive(Clone, Debug)] +pub(crate) struct KvZmqConfig { + pub(crate) host: String, + /// PUB port; the replay ROUTER, when enabled, binds `port + 1`. + pub(crate) port: u16, + pub(crate) replay: bool, + pub(crate) topic: String, + pub(crate) buffer_steps: usize, + /// The rank written into the vLLM wire's `data_parallel_rank` slot (a + /// gRPC worker is rank 0, a ZMQ rank is the engine index it advertises); + /// the SGLang wire writes a nil `attn_dp_rank` whatever the value. + pub(crate) dp_rank: i32, + pub(crate) wire: Wire, +} + +/// Bind the sockets and publish the engine's events until its event channel +/// closes. +pub(crate) async fn serve(engine: Engine, cfg: KvZmqConfig) { + let mut publisher = PubSocket::new(); + let endpoint = format!("tcp://{}:{}", cfg.host, cfg.port); + if let Err(e) = publisher.bind(&endpoint).await { + tracing::error!("kv-events publisher failed to bind {endpoint}: {e}"); + return; + } + let mut replay = None; + if cfg.replay { + let mut router = RouterSocket::new(); + let replay_endpoint = format!("tcp://{}:{}", cfg.host, cfg.port.saturating_add(1)); + if let Err(e) = router.bind(&replay_endpoint).await { + tracing::error!("kv-events replay failed to bind {replay_endpoint}: {e}"); + return; + } + replay = Some(router); + } + tracing::info!( + "kv-events publisher {} on {endpoint} ({:?} wire, replay {}, topic {:?})", + engine.name(), + cfg.wire, + if cfg.replay { "on" } else { "off" }, + cfg.topic + ); + run(engine, cfg, publisher, replay).await; +} + +/// Publish on already-bound sockets (tests bind port 0 and read it back). +async fn run( + engine: Engine, + cfg: KvZmqConfig, + mut publisher: PubSocket, + mut replay: Option, +) { + let mut events = engine.subscribe_published(); + let mut state = Publisher::new( + cfg.topic.into_bytes(), + cfg.buffer_steps, + cfg.dp_rank, + cfg.wire, + ); + let mut generation = engine.fault_status().generation; + if cfg.wire == Wire::Sglang { + // SGLang's scheduler clears its cache at startup and says so. + let message = state.publish(&startup_cleared()); + if let Err(e) = publisher.send(message).await { + tracing::warn!("kv-events publish failed: {e}"); + } + } + loop { + tokio::select! { + item = events.next() => { + let Some(item) = item else { break }; + if item.generation != generation { + // The engine's publisher restarted: so does this one. + generation = item.generation; + state.restart(); + } + let message = state.publish(&item.batch); + if item.dropped { + continue; // lost on the wire; replay still has it + } + if let Err(e) = publisher.send(message).await { + tracing::warn!("kv-events publish failed: {e}"); + } + } + request = recv_replay(&mut replay) => { + let Ok(request) = request else { continue }; + let Some(router) = replay.as_mut() else { continue }; + for reply in state.replay(&request) { + if let Err(e) = router.send(reply).await { + tracing::warn!("kv-events replay send failed: {e}"); + break; + } + } + } + } + } +} + +/// The next replay request, or never when replay is off. +async fn recv_replay(replay: &mut Option) -> Result { + match replay { + Some(router) => router.recv().await, + None => std::future::pending().await, + } +} + +/// The publisher's own state: sequence counter and replay buffer. +struct Publisher { + topic: Vec, + seq: u64, + buffer: VecDeque<(u64, Vec)>, + buffer_steps: usize, + dp_rank: i32, + wire: Wire, +} + +impl Publisher { + fn new(topic: Vec, buffer_steps: usize, dp_rank: i32, wire: Wire) -> Self { + Self { + topic, + seq: 0, + buffer: VecDeque::new(), + buffer_steps: buffer_steps.max(1), + dp_rank, + wire, + } + } + + /// The next sequence number to be published. + #[cfg(test)] + fn next_seq(&self) -> u64 { + self.seq + } + + /// Encode `batch`, assign it the next sequence, keep it for replay and + /// return the PUB message. + fn publish(&mut self, batch: &common::KvEventBatch) -> ZmqMessage { + let payload = encode_batch(batch, self.dp_rank, self.wire); + let seq = self.seq; + self.seq += 1; + self.buffer.push_back((seq, payload.clone())); + while self.buffer.len() > self.buffer_steps { + self.buffer.pop_front(); + } + frame(&self.topic, seq, payload) + } + + /// A publisher restart: the sequence starts over and the buffer is gone. + fn restart(&mut self) { + self.seq = 0; + self.buffer.clear(); + } + + /// Replies to a replay request as vLLM frames them (see the module doc); + /// a request without an 8-byte start sequence gets no reply. + fn replay(&self, request: &ZmqMessage) -> Vec { + let frames: Vec> = request.iter().map(|f| f.to_vec()).collect(); + let Some((start, routing)) = frames.split_last() else { + return Vec::new(); + }; + if start.len() != 8 || routing.is_empty() { + return Vec::new(); + } + let start = u64::from_be_bytes([ + start[0], start[1], start[2], start[3], start[4], start[5], start[6], start[7], + ]); + let mut replies = Vec::new(); + for (seq, payload) in self.buffer.iter().filter(|(seq, _)| *seq >= start) { + let mut message = routed(routing); + if self.wire == Wire::Vllm { + message.push_back(self.topic.clone().into()); + } + message.push_back(seq.to_be_bytes().to_vec().into()); + message.push_back(payload.clone().into()); + replies.push(message); + } + let mut end = routed(routing); + if self.wire == Wire::Vllm { + end.push_back(Vec::new().into()); + } + end.push_back(END_SEQ.to_vec().into()); + end.push_back(Vec::new().into()); + replies.push(end); + replies + } +} + +/// A reply's routing prefix (the request's frames before the start sequence: +/// the peer identity and the empty delimiter). +fn routed(routing: &[Vec]) -> ZmqMessage { + let mut message = ZmqMessage::from(routing[0].clone()); + for frame in &routing[1..] { + message.push_back(frame.clone().into()); + } + message +} + +/// A PUB message: `[topic, sequence (u64 big-endian), payload]`. +fn frame(topic: &[u8], seq: u64, payload: Vec) -> ZmqMessage { + let mut message = ZmqMessage::from(topic.to_vec()); + message.push_back(seq.to_be_bytes().to_vec().into()); + message.push_back(payload.into()); + message +} + +/// The batch SGLang's scheduler publishes first: a lone `AllBlocksCleared`. +fn startup_cleared() -> common::KvEventBatch { + common::KvEventBatch { + sequence_number: 0, + timestamp: 0.0, + events: vec![common::KvCacheEvent { + event_id: 0, + data: Some(common::kv_cache_event::Data::Cleared( + common::KvCacheCleared::default(), + )), + }], + dp_rank: None, + snapshot: None, + load: None, + } +} + +/// The engine's `EventBatch` for a proto batch: `[ts, events, rank]`, the +/// rank being `data_parallel_rank` on the vLLM wire and a nil `attn_dp_rank` +/// on SGLang's. +pub fn encode_batch(batch: &common::KvEventBatch, dp_rank: i32, wire: Wire) -> Vec { + let mut events: Vec = Vec::new(); + for event in &batch.events { + match (&event.data, wire) { + (Some(common::kv_cache_event::Data::Stored(stored)), Wire::Vllm) => { + events.push(stored_map(stored)); + } + (Some(common::kv_cache_event::Data::Stored(stored)), Wire::Sglang) => { + events.push(sglang_stored_map(stored)); + } + (Some(common::kv_cache_event::Data::Removed(removed)), Wire::Vllm) => { + events.push(removed_map(removed)); + } + (Some(common::kv_cache_event::Data::Removed(removed)), Wire::Sglang) => { + // One remove per node; a node here is one evicted block. + events.extend(removed.block_hashes.iter().map(|h| sglang_removed_map(*h))); + } + (Some(common::kv_cache_event::Data::Cleared(_)), _) => events.push(cleared_map()), + (None, _) => {} + } + } + let rank = match wire { + Wire::Vllm => Value::from(dp_rank), + Wire::Sglang => Value::Nil, + }; + // The batch's own creation time, as the engines' `ts`; a batch without one + // (the startup clear) is stamped now. + let ts = if batch.timestamp > 0.0 { + batch.timestamp + } else { + unix_seconds() + }; + let value = Value::Array(vec![Value::F64(ts), Value::Array(events), rank]); + let mut buf = Vec::new(); + // Writing into a Vec cannot fail. + let _ = rmpv::encode::write_value(&mut buf, &value); + buf +} + +fn key(name: &str) -> Value { + Value::String(name.into()) +} + +fn hash_value(hash: i64) -> Value { + // vLLM's int form is the unsigned low 64 bits; the proto carries the same + // bits as a signed value. + Value::from(hash as u64) +} + +/// `BlockStored` as vLLM encodes it: the block_size of the first block, the +/// token ids of every block in order. +fn stored_map(stored: &common::KvBlocksStored) -> Value { + let block_size = stored + .blocks + .first() + .map(|b| i64::from(b.block_size)) + .unwrap_or(0); + let token_ids: Vec = stored + .blocks + .iter() + .flat_map(|b| b.token_ids.iter().map(|t| Value::from(*t))) + .collect(); + Value::Map(vec![ + (key("type"), Value::String("BlockStored".into())), + ( + key("block_hashes"), + Value::Array( + stored + .blocks + .iter() + .map(|b| hash_value(b.block_hash)) + .collect(), + ), + ), + ( + key("parent_block_hash"), + stored + .parent_block_hash + .map(hash_value) + .unwrap_or(Value::Nil), + ), + (key("token_ids"), Value::Array(token_ids)), + (key("block_size"), Value::from(block_size)), + (key("lora_id"), Value::Nil), + (key("medium"), Value::String("GPU".into())), + (key("lora_name"), Value::Nil), + (key("group_idx"), Value::from(0)), + ( + key("kv_cache_spec_kind"), + Value::String("full_attention".into()), + ), + ]) +} + +fn removed_map(removed: &common::KvBlocksRemoved) -> Value { + Value::Map(vec![ + (key("type"), Value::String("BlockRemoved".into())), + ( + key("block_hashes"), + Value::Array( + removed + .block_hashes + .iter() + .map(|h| hash_value(*h)) + .collect(), + ), + ), + (key("medium"), Value::String("GPU".into())), + (key("group_idx"), Value::from(0)), + ]) +} + +/// `BlockStored` as SGLang encodes it: signed hashes, no medium, group or +/// spec fields (SGLang omits its optional fields when unset). +fn sglang_stored_map(stored: &common::KvBlocksStored) -> Value { + let block_size = stored + .blocks + .first() + .map(|b| i64::from(b.block_size)) + .unwrap_or(0); + let token_ids: Vec = stored + .blocks + .iter() + .flat_map(|b| b.token_ids.iter().map(|t| Value::from(*t))) + .collect(); + Value::Map(vec![ + (key("type"), Value::String("BlockStored".into())), + ( + key("block_hashes"), + Value::Array( + stored + .blocks + .iter() + .map(|b| Value::from(b.block_hash)) + .collect(), + ), + ), + ( + key("parent_block_hash"), + stored + .parent_block_hash + .map(Value::from) + .unwrap_or(Value::Nil), + ), + (key("token_ids"), Value::Array(token_ids)), + (key("block_size"), Value::from(block_size)), + (key("lora_id"), Value::Nil), + ]) +} + +fn sglang_removed_map(hash: i64) -> Value { + Value::Map(vec![ + (key("type"), Value::String("BlockRemoved".into())), + (key("block_hashes"), Value::Array(vec![Value::from(hash)])), + ]) +} + +fn cleared_map() -> Value { + Value::Map(vec![( + key("type"), + Value::String("AllBlocksCleared".into()), + )]) +} + +#[cfg(test)] +mod tests { + use std::time::Duration; + + use engine_servicer::kv_wire::{Normalizer, WireBatch}; + use tokio::time::timeout; + use zeromq::{DealerSocket, SubSocket}; + + use super::*; + use crate::engine::{EngineParams, GenEvent, NewRequest}; + + fn stored( + hashes: &[i64], + parent: Option, + tokens_per_block: &[&[u32]], + ) -> common::KvCacheEvent { + common::KvCacheEvent { + event_id: 1, + data: Some(common::kv_cache_event::Data::Stored( + common::KvBlocksStored { + blocks: hashes + .iter() + .zip(tokens_per_block) + .map(|(hash, tokens)| common::KvBlock { + block_hash: *hash, + token_ids: tokens.to_vec(), + block_size: tokens.len() as i32, + ..Default::default() + }) + .collect(), + parent_block_hash: parent, + ..Default::default() + }, + )), + } + } + + fn removed(hashes: &[i64]) -> common::KvCacheEvent { + common::KvCacheEvent { + event_id: 2, + data: Some(common::kv_cache_event::Data::Removed( + common::KvBlocksRemoved { + block_hashes: hashes.to_vec(), + ..Default::default() + }, + )), + } + } + + fn cleared() -> common::KvCacheEvent { + common::KvCacheEvent { + event_id: 3, + data: Some(common::kv_cache_event::Data::Cleared( + common::KvCacheCleared::default(), + )), + } + } + + fn sample_batch() -> common::KvEventBatch { + common::KvEventBatch { + sequence_number: 9, + timestamp: 0.0, + events: vec![ + stored(&[11, -2], Some(10), &[&[1, 2, 3, 4], &[5, 6, 7, 8]]), + removed(&[11]), + cleared(), + ], + dp_rank: Some(0), + snapshot: None, + load: None, + } + } + + fn keys(map: &Value) -> Vec { + match map { + Value::Map(entries) => entries + .iter() + .map(|(k, _)| k.as_str().expect("string key").to_string()) + .collect(), + other => panic!("not a map: {other:?}"), + } + } + + fn field<'a>(map: &'a Value, name: &str) -> &'a Value { + match map { + Value::Map(entries) => entries + .iter() + .find(|(k, _)| k.as_str() == Some(name)) + .map(|(_, v)| v) + .unwrap_or_else(|| panic!("no field {name}")), + other => panic!("not a map: {other:?}"), + } + } + + #[test] + fn encodes_vllm_event_batches_field_for_field() { + let payload = encode_batch(&sample_batch(), 3, Wire::Vllm); + let value = rmpv::decode::read_value(&mut payload.as_slice()).expect("msgpack"); + let Value::Array(batch) = value else { + panic!("batch is array-like"); + }; + assert_eq!(batch.len(), 3, "[ts, events, data_parallel_rank]"); + assert!(batch[0].as_f64().is_some_and(|ts| ts > 1.7e9)); + let mut stamped = sample_batch(); + stamped.timestamp = 1_700_000_000.25; + let again = encode_batch(&stamped, 3, Wire::Vllm); + let Value::Array(again) = rmpv::decode::read_value(&mut again.as_slice()).expect("msgpack") + else { + panic!("array"); + }; + assert_eq!( + again[0].as_f64(), + Some(1_700_000_000.25), + "the batch's creation time is the ts" + ); + assert_eq!(batch[2].as_i64(), Some(3)); + let events = batch[1].as_array().expect("events"); + assert_eq!(events.len(), 3); + + assert_eq!( + keys(&events[0]), + [ + "type", + "block_hashes", + "parent_block_hash", + "token_ids", + "block_size", + "lora_id", + "medium", + "lora_name", + "group_idx", + "kv_cache_spec_kind" + ] + ); + assert_eq!(field(&events[0], "type").as_str(), Some("BlockStored")); + let hashes = field(&events[0], "block_hashes") + .as_array() + .expect("hashes"); + assert_eq!(hashes[0].as_u64(), Some(11)); + assert_eq!( + hashes[1].as_u64(), + Some(u64::MAX - 1), + "a negative proto hash is the same bits as vLLM's unsigned int" + ); + assert_eq!(field(&events[0], "parent_block_hash").as_u64(), Some(10)); + assert_eq!( + field(&events[0], "token_ids").as_array().map(Vec::len), + Some(8), + "token ids of every block, in order" + ); + assert_eq!(field(&events[0], "block_size").as_i64(), Some(4)); + assert!(field(&events[0], "lora_id").is_nil()); + assert_eq!(field(&events[0], "medium").as_str(), Some("GPU")); + assert!(field(&events[0], "lora_name").is_nil()); + assert_eq!(field(&events[0], "group_idx").as_i64(), Some(0)); + assert_eq!( + field(&events[0], "kv_cache_spec_kind").as_str(), + Some("full_attention") + ); + + assert_eq!( + keys(&events[1]), + ["type", "block_hashes", "medium", "group_idx"] + ); + assert_eq!(field(&events[1], "type").as_str(), Some("BlockRemoved")); + assert_eq!(keys(&events[2]), ["type"]); + assert_eq!(field(&events[2], "type").as_str(), Some("AllBlocksCleared")); + } + + #[test] + fn the_relay_normalizes_the_payload_back_into_the_proto() { + let payload = encode_batch(&sample_batch(), 0, Wire::Vllm); + let wire: WireBatch = rmp_serde::from_slice(&payload).expect("the relay decodes it"); + let mut event_id = 0; + let relayed = Normalizer::new().normalize_batch(wire, 7, &mut event_id); + assert_eq!(relayed.sequence_number, 7); + let kinds: Vec<&str> = relayed + .events + .iter() + .map(|e| match &e.data { + Some(common::kv_cache_event::Data::Stored(_)) => "stored", + Some(common::kv_cache_event::Data::Removed(_)) => "removed", + Some(common::kv_cache_event::Data::Cleared(_)) => "cleared", + None => "none", + }) + .collect(); + assert_eq!(kinds, ["stored", "removed", "cleared"]); + let Some(common::kv_cache_event::Data::Stored(stored)) = &relayed.events[0].data else { + unreachable!() + }; + assert_eq!( + stored + .blocks + .iter() + .map(|b| b.block_hash) + .collect::>(), + [11, -2], + "hashes round-trip to the proto's signed identity" + ); + assert_eq!(stored.parent_block_hash, Some(10)); + assert_eq!(stored.blocks[1].token_ids, [5, 6, 7, 8]); + let Some(common::kv_cache_event::Data::Removed(removed)) = &relayed.events[1].data else { + unreachable!() + }; + assert_eq!(removed.block_hashes, [11]); + } + + #[test] + fn replay_answers_from_the_start_sequence_and_ends_with_the_marker() { + let mut publisher = Publisher::new(b"kv".to_vec(), 3, 0, Wire::Vllm); + for _ in 0..5 { + let message = publisher.publish(&sample_batch()); + assert_eq!(message.len(), 3); + assert_eq!(message.get(0).map(|t| t.as_ref()), Some(&b"kv"[..])); + } + assert_eq!(publisher.next_seq(), 5); + + let mut request = ZmqMessage::from(b"peer".to_vec()); + request.push_back(Vec::new().into()); + request.push_back(1u64.to_be_bytes().to_vec().into()); + let replies = publisher.replay(&request); + // The buffer keeps the last three (2, 3, 4); 1 is gone. + let seqs: Vec = replies[..replies.len() - 1] + .iter() + .map(|m| { + assert_eq!(m.len(), 5, "[peer, empty, topic, seq, payload]"); + assert_eq!(m.get(2).map(|t| t.as_ref()), Some(&b"kv"[..])); + let seq = m.get(3).expect("seq"); + u64::from_be_bytes(seq.as_ref().try_into().expect("8 bytes")) + }) + .collect(); + assert_eq!(seqs, [2, 3, 4]); + let end = replies.last().expect("end marker"); + assert_eq!(end.len(), 5, "[peer, empty, empty, END, empty]"); + assert_eq!(end.get(0).map(|t| t.as_ref()), Some(&b"peer"[..])); + assert!(end.get(2).is_some_and(|f| f.is_empty())); + assert_eq!(end.get(3).map(|t| t.as_ref()), Some(&END_SEQ[..])); + assert!(end.get(4).is_some_and(|f| f.is_empty())); + + let mut late = ZmqMessage::from(b"peer".to_vec()); + late.push_back(Vec::new().into()); + late.push_back(99u64.to_be_bytes().to_vec().into()); + assert_eq!(publisher.replay(&late).len(), 1, "only the end marker"); + + let mut bad = ZmqMessage::from(b"peer".to_vec()); + bad.push_back(vec![1, 2, 3, 4].into()); + assert!( + publisher.replay(&bad).is_empty(), + "a malformed request is ignored" + ); + + publisher.restart(); + assert_eq!(publisher.next_seq(), 0); + assert_eq!( + publisher.replay(&request).len(), + 1, + "the buffer is gone after a restart" + ); + } + + #[test] + fn sglang_wire_encodes_signed_hashes_without_group_or_spec_fields() { + let payload = encode_batch(&sample_batch(), 0, Wire::Sglang); + let value = rmpv::decode::read_value(&mut payload.as_slice()).expect("msgpack"); + let Value::Array(batch) = value else { + panic!("batch is array-like"); + }; + assert_eq!(batch.len(), 3, "[ts, events, attn_dp_rank]"); + assert!(batch[2].is_nil(), "attn_dp_rank is nil for a single rank"); + let events = batch[1].as_array().expect("events"); + assert_eq!(events.len(), 3, "stored, one remove per block, cleared"); + assert_eq!( + keys(&events[0]), + [ + "type", + "block_hashes", + "parent_block_hash", + "token_ids", + "block_size", + "lora_id" + ] + ); + let hashes = field(&events[0], "block_hashes") + .as_array() + .expect("hashes"); + assert_eq!(hashes[0].as_i64(), Some(11)); + assert_eq!(hashes[1].as_i64(), Some(-2), "signed, as SGLang's int64"); + assert_eq!(field(&events[0], "parent_block_hash").as_i64(), Some(10)); + assert_eq!(keys(&events[1]), ["type", "block_hashes"]); + assert_eq!(field(&events[1], "type").as_str(), Some("BlockRemoved")); + assert_eq!(keys(&events[2]), ["type"]); + } + + #[test] + fn the_relay_normalizes_the_sglang_payload_too() { + let payload = encode_batch(&sample_batch(), 0, Wire::Sglang); + let wire: WireBatch = rmp_serde::from_slice(&payload).expect("the relay decodes it"); + assert_eq!(wire.dp_rank, None); + let mut event_id = 0; + let relayed = Normalizer::new().normalize_batch(wire, 3, &mut event_id); + let Some(common::kv_cache_event::Data::Stored(stored)) = &relayed.events[0].data else { + panic!("first event is the store"); + }; + assert_eq!( + stored + .blocks + .iter() + .map(|b| b.block_hash) + .collect::>(), + [11, -2] + ); + assert_eq!(stored.parent_block_hash, Some(10)); + let Some(common::kv_cache_event::Data::Removed(removed)) = &relayed.events[1].data else { + panic!("second event is the remove"); + }; + assert_eq!(removed.block_hashes, [11]); + assert!(matches!( + relayed.events[2].data, + Some(common::kv_cache_event::Data::Cleared(_)) + )); + } + + #[test] + fn sglang_replay_frames_carry_no_topic() { + let mut publisher = Publisher::new(b"kv".to_vec(), 10, 0, Wire::Sglang); + publisher.publish(&startup_cleared()); + publisher.publish(&sample_batch()); + let mut request = ZmqMessage::from(b"peer".to_vec()); + request.push_back(Vec::new().into()); + request.push_back(0u64.to_be_bytes().to_vec().into()); + let replies = publisher.replay(&request); + assert_eq!(replies.len(), 3, "two batches and the end marker"); + assert_eq!(replies[0].len(), 4, "[peer, empty, seq, payload]"); + assert_eq!( + replies[0].get(2).map(|f| f.as_ref()), + Some(&0u64.to_be_bytes()[..]) + ); + let first: WireBatch = + rmp_serde::from_slice(replies[0].get(3).expect("payload")).expect("decodes"); + assert_eq!(first.events.len(), 1, "the startup clear comes first"); + let end = replies.last().expect("end"); + assert_eq!(end.len(), 4, "[peer, empty, END, empty]"); + assert_eq!(end.get(2).map(|f| f.as_ref()), Some(&END_SEQ[..])); + assert!(end.get(3).is_some_and(|f| f.is_empty())); + } + + #[tokio::test] + async fn sglang_publisher_starts_with_all_blocks_cleared() { + let engine = Engine::spawn(EngineParams::default()); + let mut pub_socket = PubSocket::new(); + pub_socket + .bind("tcp://127.0.0.1:0") + .await + .expect("pub binds"); + let mut router = RouterSocket::new(); + let replay_endpoint = router + .bind("tcp://127.0.0.1:0") + .await + .expect("router binds") + .to_string(); + let cfg = KvZmqConfig { + host: "127.0.0.1".to_string(), + port: 0, + replay: true, + topic: String::new(), + buffer_steps: 100, + dp_rank: 0, + wire: Wire::Sglang, + }; + #[expect( + clippy::disallowed_methods, + reason = "the publisher ends with the engine when the test drops it" + )] + let _publisher = tokio::spawn(run(engine.clone(), cfg, pub_socket, Some(router))); + let mut dealer = DealerSocket::new(); + dealer + .connect(&replay_endpoint) + .await + .expect("dealer connects"); + let mut request = ZmqMessage::from(Vec::new()); + request.push_back(0u64.to_be_bytes().to_vec().into()); + // The publisher binds before the startup batch is buffered; ask until + // the replay has it. + let mut first = None; + for _ in 0..50 { + dealer.send(request.clone()).await.expect("replay request"); + let reply = timeout(Duration::from_secs(5), dealer.recv()) + .await + .expect("replay answers") + .expect("reply"); + assert_eq!(reply.len(), 3, "[empty, seq, payload] on the SGLang wire"); + if reply.get(1).map(|f| f.as_ref()) == Some(&END_SEQ[..]) { + tokio::time::sleep(Duration::from_millis(20)).await; + continue; + } + first = Some(reply); + break; + } + let first = first.expect("the startup batch is replayable"); + assert_eq!( + first.get(1).map(|f| f.as_ref()), + Some(&0u64.to_be_bytes()[..]) + ); + let batch: WireBatch = + rmp_serde::from_slice(first.get(2).expect("payload")).expect("decodes"); + assert_eq!(batch.events.len(), 1); + let mut event_id = 0; + let relayed = Normalizer::new().normalize_batch(batch, 0, &mut event_id); + assert!(matches!( + relayed.events[0].data, + Some(common::kv_cache_event::Data::Cleared(_)) + )); + // Drain the end marker so the dealer is left clean. + let _ = timeout(Duration::from_secs(1), dealer.recv()).await; + } + + #[tokio::test] + async fn a_subscriber_gets_live_frames_and_can_replay_what_it_missed() { + let engine = Engine::spawn(EngineParams::default()); + let mut pub_socket = PubSocket::new(); + let endpoint = pub_socket + .bind("tcp://127.0.0.1:0") + .await + .expect("pub binds") + .to_string(); + let mut router = RouterSocket::new(); + let replay_endpoint = router + .bind("tcp://127.0.0.1:0") + .await + .expect("router binds") + .to_string(); + let cfg = KvZmqConfig { + host: "127.0.0.1".to_string(), + port: 0, + replay: true, + topic: "kv".to_string(), + buffer_steps: 100, + dp_rank: 0, + wire: Wire::Vllm, + }; + #[expect( + clippy::disallowed_methods, + reason = "the publisher ends with the engine when the test drops it" + )] + let _publisher = tokio::spawn(run(engine.clone(), cfg, pub_socket, Some(router))); + + let mut sub = SubSocket::new(); + sub.subscribe("kv").await.expect("subscribe"); + sub.connect(&endpoint).await.expect("connect"); + + // The subscription reaches the publisher a moment after the connect; + // keep the engine producing passes until a frame comes through. + let mut receivers = Vec::new(); + let mut first = None; + for i in 0..200 { + let (tx, rx) = tokio::sync::mpsc::unbounded_channel::(); + receivers.push(rx); + engine.submit(NewRequest { + request_id: format!("r{i}"), + prompt_token_ids: (0..64).map(|t| t + i * 64).collect(), + max_new: 1, + events: tx, + }); + if let Ok(Ok(message)) = timeout(Duration::from_millis(50), sub.recv()).await { + first = Some(message); + break; + } + } + let first = first.expect("a live frame arrived"); + assert_eq!(first.len(), 3); + assert_eq!(first.get(0).map(|t| t.as_ref()), Some(&b"kv"[..])); + let live_seq = u64::from_be_bytes( + first + .get(1) + .expect("seq") + .as_ref() + .try_into() + .expect("8 bytes"), + ); + let wire: WireBatch = + rmp_serde::from_slice(first.get(2).expect("payload")).expect("decodes"); + assert!(!wire.events.is_empty()); + + // Everything the subscriber missed is in the replay buffer, from 0. + let mut dealer = DealerSocket::new(); + dealer + .connect(&replay_endpoint) + .await + .expect("dealer connects"); + let mut request = ZmqMessage::from(Vec::new()); + request.push_back(0u64.to_be_bytes().to_vec().into()); + dealer.send(request).await.expect("replay request"); + let mut replayed = Vec::new(); + loop { + let reply = timeout(Duration::from_secs(5), dealer.recv()) + .await + .expect("replay answers") + .expect("reply"); + assert_eq!(reply.len(), 4, "[empty, topic, seq, payload]"); + let seq = reply.get(2).expect("seq"); + if seq.as_ref() == END_SEQ { + break; + } + replayed.push(u64::from_be_bytes( + seq.as_ref().try_into().expect("8 bytes"), + )); + let _: WireBatch = + rmp_serde::from_slice(reply.get(3).expect("payload")).expect("decodes"); + } + assert_eq!(replayed[0], 0, "replay starts at the requested sequence"); + assert!(replayed.windows(2).all(|w| w[1] == w[0] + 1), "contiguous"); + assert!( + replayed.contains(&live_seq), + "the live frame is in the buffer too" + ); + // A publisher restart starts the ZMQ sequence over as well. + engine.restart_publisher().await; + let mut restarted = None; + for i in 200..400 { + let (tx, rx) = tokio::sync::mpsc::unbounded_channel::(); + receivers.push(rx); + engine.submit(NewRequest { + request_id: format!("r{i}"), + prompt_token_ids: (0..64).map(|t| t + i * 64).collect(), + max_new: 1, + events: tx, + }); + if let Ok(Ok(message)) = timeout(Duration::from_millis(50), sub.recv()).await { + let seq = u64::from_be_bytes( + message + .get(1) + .expect("seq") + .as_ref() + .try_into() + .expect("8 bytes"), + ); + // Sequence 0 was consumed live (or missed) before the restart; + // seeing it again means the publisher started over. + if seq == 0 { + restarted = Some(seq); + break; + } + } + } + assert_eq!(restarted, Some(0), "the sequence restarted from 0"); + drop(receivers); + } +} diff --git a/crates/mock_worker/src/lib.rs b/crates/mock_worker/src/lib.rs index c76d153c1b..e6f8e0a491 100644 --- a/crates/mock_worker/src/lib.rs +++ b/crates/mock_worker/src/lib.rs @@ -1,9 +1,11 @@ //! Library surface for `mock-worker`'s HTTP/gRPC simulators, so both the //! standalone binary and in-process integration tests can drive them. +pub mod admin; pub mod config; pub mod engine; pub mod grpc; pub mod http; +pub mod kv_zmq; pub mod replay; pub mod zmq; diff --git a/crates/mock_worker/src/main.rs b/crates/mock_worker/src/main.rs index 1cacf5191e..7d74f71d04 100644 --- a/crates/mock_worker/src/main.rs +++ b/crates/mock_worker/src/main.rs @@ -8,7 +8,7 @@ use std::{process::ExitCode, sync::Arc}; -use mock_worker::{config::Config, grpc, http, replay::Capture, zmq}; +use mock_worker::{admin, config::Config, grpc, http, replay::Capture, zmq}; #[tokio::main] async fn main() -> ExitCode { @@ -64,6 +64,9 @@ async fn main() -> ExitCode { }; workers.spawn(grpc::serve(cfg.clone(), cfg.host.clone(), port)); } + if let Some(port) = cfg.admin_port { + workers.spawn(admin::serve(cfg.clone(), cfg.host.clone(), port)); + } if let Some(handshake) = cfg.zmq_handshake.clone() { for i in 0..cfg.zmq_count { let engine_index = cfg.zmq_start_index + i as u32; diff --git a/crates/mock_worker/src/replay.rs b/crates/mock_worker/src/replay.rs index 1e0c79b27f..55719fbe1e 100644 --- a/crates/mock_worker/src/replay.rs +++ b/crates/mock_worker/src/replay.rs @@ -1,5 +1,7 @@ -//! Replay testing: `--capture` records every gRPC `Generate` request the -//! worker receives, so a test can check what the gateway put on the wire. +//! Request capture: `--capture` records every gRPC `Generate` request the +//! worker receives, one JSON object per line, so a test can check what the +//! gateway put on the wire. (The trace replayer is this crate's `replay` +//! binary, `src/bin/replay.rs`.) #[cfg(unix)] use std::os::unix::fs::OpenOptionsExt; diff --git a/crates/mock_worker/src/zmq.rs b/crates/mock_worker/src/zmq.rs index c58d9079c7..d0d0bb6281 100644 --- a/crates/mock_worker/src/zmq.rs +++ b/crates/mock_worker/src/zmq.rs @@ -53,7 +53,24 @@ pub async fn serve(cfg: Arc, handshake_address: String, engine_index: u3 tracing::info!("zmq mock engine {engine_index} connected to {handshake_address}"); let (mut input, output) = mock.split(); - let engine = cfg.realistic.then(|| Engine::spawn(cfg.engine.clone())); + let engine = cfg + .realistic + .then(|| Engine::spawn_named(cfg.engine.clone(), format!("zmq:{engine_index}"), true)); + // The rank's KV-event publisher, numbered after the gRPC workers and + // stamping the batches with the rank this engine advertised. + let publisher = engine.as_ref().and_then(|engine| { + let index = cfg + .grpc_count + .checked_add(u16::try_from(engine_index).ok()?)?; + let dp_rank = i32::try_from(engine_index).ok()?; + cfg.kv_zmq_for(index, dp_rank) + .map(|kv| (engine.clone(), kv)) + }); + #[expect( + clippy::disallowed_methods, + reason = "publisher self-terminates when the engine's event channel closes" + )] + let _publisher = publisher.map(|(engine, kv)| tokio::spawn(crate::kv_zmq::serve(engine, kv))); // A single writer owns the output PUSH socket; per-request forwarders funnel // their outputs here so concurrent requests serialize onto the one socket. @@ -136,8 +153,8 @@ pub async fn serve(cfg: Arc, handshake_address: String, engine_index: u3 tracing::debug!("zmq engine {engine_index} ignoring start of wave {wave}"); } Ok(EngineInbound::Utility(call)) => { - // The mock holds no KV blocks, so a prefix-cache reset always - // succeeds; no other EngineCore method exists here. + // A prefix-cache reset is acknowledged and leaves the simulated + // engine's cache as it is; no other EngineCore method exists here. let outcome = if call.method == "reset_prefix_cache" { Ok(OpaqueValue::from(true)) } else { @@ -303,12 +320,12 @@ mod tests { zmq_count: 0, zmq_start_index: 0, model_id: "mock-model".to_string(), - tokenizer_path: "mock-model".to_string(), - gen_delay: Duration::ZERO, + tokenizer_path: String::new(), + gen_delay: Duration::from_millis(0), output_tokens: 4, - realistic: false, + realistic: true, engine: EngineParams::default(), - replay: Default::default(), + ..Config::default() } } @@ -358,7 +375,11 @@ mod tests { } } assert!(finished, "stream should reach a terminal output"); - assert_eq!(tokens.len(), 4, "canned mode emits output_tokens tokens"); + assert_eq!( + tokens.len(), + 4, + "a request without max_tokens gets output_tokens tokens" + ); } /// Two mock ranks dial one socket set — the grouped-worker topology the @@ -412,7 +433,11 @@ mod tests { } } assert!(finished, "rank {rank} should reach a terminal output"); - assert_eq!(tokens.len(), 4, "rank {rank} emits output_tokens tokens"); + assert_eq!( + tokens.len(), + 4, + "rank {rank} answers with output_tokens tokens" + ); } } } diff --git a/crates/mock_worker/tests/capture.rs b/crates/mock_worker/tests/capture.rs index 73dc652e7b..e075fc137e 100644 --- a/crates/mock_worker/tests/capture.rs +++ b/crates/mock_worker/tests/capture.rs @@ -12,7 +12,7 @@ use std::{ use futures::StreamExt; use mock_worker::{ config::{Config, ReplayConfig}, - engine::EngineParams, + engine::{EngineParams, TimingModel}, grpc::serve_with_listener, }; use serde_json::{json, Value}; @@ -38,6 +38,7 @@ fn config(capture: Option) -> Config { realistic: false, engine: EngineParams::default(), replay: ReplayConfig { capture }, + ..Config::default() } } @@ -145,7 +146,13 @@ async fn line_is_on_disk_before_the_first_frame() { let path = dir.path().join("generate.jsonl"); let mut cfg = config(Some(path.clone())); cfg.realistic = true; - cfg.engine.decode_base_ms = 200.0; + // The linear timing model with a 200 ms decode base holds the first + // frame back long enough for the check. + cfg.engine.timing = TimingModel::Linear { + prefill_tps: 8000.0, + decode_base_ms: 200.0, + decode_per_req_ms: 0.35, + }; let client = start(cfg).await; let mut req = request("early", &[1, 2, 3], "early"); diff --git a/crates/protocols/src/worker.rs b/crates/protocols/src/worker.rs index fcc327cd03..488cfc7bba 100644 --- a/crates/protocols/src/worker.rs +++ b/crates/protocols/src/worker.rs @@ -939,6 +939,12 @@ pub struct WorkerInfo { #[serde(default, skip_serializing_if = "Option::is_none")] pub pd_pairing: Option, + /// Why the gateway's liveness tracker currently keeps the worker out of + /// routing (`unreachable` or `wedged`) while its health status stands; + /// absent when it is routable. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub stalled: Option, + /// The worker's last polled engine load, as published by the load /// monitor. `None` when load monitoring has produced nothing for this /// worker yet. Unrelated to `load` above, which counts in-flight @@ -962,6 +968,7 @@ impl WorkerInfo { load: 0, http2: false, pd_pairing: None, + stalled: None, engine_load: None, job_status, } @@ -1455,6 +1462,14 @@ pub struct WorkerLoadResponse { pub loads: Vec, #[serde(skip_serializing_if = "Option::is_none")] pub aggregate: Option, + /// When the engine's state behind this report was sampled, on the + /// gateway's clock: the poll's receipt, or a pushed record's receipt + /// less its age and the one-way latency. A policy that books in-flight + /// work locally releases only what it dispatched before this instant. + /// Not serialized: it is meaningful only in the process that set it. + #[serde(skip)] + #[schemars(skip)] + pub sampled_at: Option, } impl WorkerLoadResponse { diff --git a/grpc_servicer/DEVELOPMENT.md b/grpc_servicer/DEVELOPMENT.md index f45fbbc73c..5553741354 100644 --- a/grpc_servicer/DEVELOPMENT.md +++ b/grpc_servicer/DEVELOPMENT.md @@ -11,6 +11,21 @@ pip install -e grpc_servicer/ No version concerns locally — editable installs always use the latest source. +## Proto stubs for local tests + +The tests import ``smg_grpc_proto``, whose stubs the released package builds at +install time from ``crates/grpc_client/proto``. While a proto change is in +flight, generate stubs from the checkout instead of installing the release: + +```bash +pip install 'grpcio-tools>=1.81.1,<1.82' +python3 grpc_servicer/scripts/gen_proto_stubs.py /tmp/smg-proto-gen +SMG_GRPC_PROTO_PATH=/tmp/smg-proto-gen pytest -q grpc_servicer/tests +``` + +Generated stubs are never committed; ``tests/conftest.py`` puts the directory +on ``sys.path`` when the variable is set. + ## CI — vLLM PR tests install both `smg-grpc-proto` and `smg-grpc-servicer` from source (not PyPI), diff --git a/grpc_servicer/README.md b/grpc_servicer/README.md index ef4115fe7a..277b00bdd4 100644 --- a/grpc_servicer/README.md +++ b/grpc_servicer/README.md @@ -269,14 +269,25 @@ To retain cache knowledge across a recoverable event gap, configure SGLang's The bridge subscribes to live events before requesting missed batches from the replay endpoint, preserves publisher sequence numbers, and removes overlap at handoff. Both subscriptions currently use DP rank 0; allocate non-overlapping -port ranges if multiple DP ranks publish events. +port ranges if multiple DP ranks publish events. The Rust servicer's relay +uses the same replay endpoint (the vLLM, SGLang and TokenSpeed launchers all +pass `replay_endpoint` from the engine's kv-events config) for gaps in flight +and for the batches published before its subscription joined the publisher; +without one those are counted as lost or unknown, never silently skipped. Without replay, or when history is expired, empty, malformed, or unavailable (timeout: five seconds), the bridge reports `OUT_OF_RANGE` before streaming or `DATA_LOSS` after streaming starts. SMG discards that worker's stale mappings and resubscribes with zero. A zero cursor rebuilds knowledge from subsequent live -events; it is not a complete cache snapshot. An empty replay is conservatively -reset because it cannot distinguish an idle publisher from a restarted one. +events; it is not a complete cache snapshot: this bridge relays per call and +keeps no history or block record between calls, so it has nothing older to +serve. The Rust servicer's relay (`crates/engine_servicer`, the vLLM servicer +and `SMG_SGLANG_SERVICER_IMPL=rust`) does: it keeps a bounded history and the +engine's live blocks for the servicer's lifetime, serves the history to a +subscription from zero while it is complete, and a state snapshot +(`KvSnapshotChunk`) before live events once the window has rolled. An empty +replay is conservatively reset because it cannot distinguish an idle publisher +from a restarted one. #### Rust request path (`SMG_SGLANG_SERVICER_IMPL=rust`) diff --git a/grpc_servicer/scripts/gen_proto_stubs.py b/grpc_servicer/scripts/gen_proto_stubs.py new file mode 100755 index 0000000000..79e0110176 --- /dev/null +++ b/grpc_servicer/scripts/gen_proto_stubs.py @@ -0,0 +1,80 @@ +#!/usr/bin/env python3 +"""Generate ``smg_grpc_proto`` stubs from this checkout's proto files, for tests. + +The released package builds its stubs at install time (``crates/grpc_client/ +python/setup.py``); while a proto change is in flight the tests need stubs of +the local ``.proto`` files instead. This writes an importable +``smg_grpc_proto`` package (the same layout as the release) into a directory, +without committing generated code: + + python3 grpc_servicer/scripts/gen_proto_stubs.py /tmp/smg-proto-gen + SMG_GRPC_PROTO_PATH=/tmp/smg-proto-gen pytest -q grpc_servicer/tests + +Requires ``grpcio-tools`` (the release caps it below 1.82 for protobuf 6). +""" + +from __future__ import annotations + +import pathlib +import shutil +import sys + + +def main(argv: list[str]) -> int: + if len(argv) != 2: + print(__doc__) + return 2 + target = pathlib.Path(argv[1]).resolve() + repo = pathlib.Path(__file__).resolve().parents[2] + proto_dir = repo / "crates" / "grpc_client" / "proto" + package_src = repo / "crates" / "grpc_client" / "python" / "smg_grpc_proto" + protos = sorted(proto_dir.glob("*.proto")) + if not protos: + print(f"no .proto files under {proto_dir}", file=sys.stderr) + return 1 + + import grpc_tools + from grpc_tools import protoc + + package = target / "smg_grpc_proto" + generated = package / "generated" + if package.exists(): + shutil.rmtree(package) + generated.mkdir(parents=True) + init = (package_src / "__init__.py").read_text() + init = init.replace( + '__version__ = version("smg-grpc-proto")', + 'try:\n __version__ = version("smg-grpc-proto")\n' + "except Exception: # locally generated stubs carry no distribution metadata\n" + ' __version__ = "0.0.0+local"', + ) + (package / "__init__.py").write_text(init) + (generated / "__init__.py").write_text('"""Auto-generated protobuf stubs. Do not edit."""\n') + + well_known = pathlib.Path(grpc_tools.__file__).parent / "_proto" + args = [ + "grpc_tools.protoc", + f"--proto_path={proto_dir}", + f"--proto_path={well_known}", + f"--python_out={generated}", + f"--grpc_python_out={generated}", + f"--pyi_out={generated}", + *map(str, protos), + ] + if protoc.main(args) != 0: + print("protoc failed", file=sys.stderr) + return 1 + # grpcio-tools emits absolute imports between the generated modules. + for module in generated.glob("*_pb2*.py"): + text = module.read_text() + for proto in protos: + name = proto.stem + "_pb2" + text = text.replace(f"import {name}", f"from . import {name}") + module.write_text("# mypy: ignore-errors\n" + text) + print(f"generated {len(protos)} protos into {package}") + print(f"export SMG_GRPC_PROTO_PATH={target}") + return 0 + + +if __name__ == "__main__": + sys.exit(main(sys.argv)) diff --git a/grpc_servicer/smg_grpc_servicer/kv_relay.py b/grpc_servicer/smg_grpc_servicer/kv_relay.py new file mode 100644 index 0000000000..2578e51622 --- /dev/null +++ b/grpc_servicer/smg_grpc_servicer/kv_relay.py @@ -0,0 +1,1169 @@ +"""The engines' KV-cache events relayed into the gateway's ``KvEventBatch`` stream. + +One relay per ``SubscribeKvEvents`` call: a SUB socket per data-parallel rank +(vLLM and SGLang publish one stream per rank on ``base_port + rank``, each +with its own sequence counter), each rank's sequence followed with gap replay +through the publisher's ROUTER, and every event normalized by the rules the +Rust relay applies (``crates/engine_servicer/src/kv_wire.rs``; the relay tests +on both sides build the same engine wire shapes in code and expect the same +output of them). + +What the gateway sees is one stream with its own contiguous sequence numbers +and the rank on every batch. Recovery is the relay's: a gap on a rank is +filled from that rank's replay endpoint before anything later is forwarded; +a replay that cannot be verified (history truncated, timeout, malformed) and +a publisher restart (its counter starts over) end the stream with +``DATA_LOSS``, which the gateway answers by clearing the worker and +resubscribing from zero. A non-zero ``start_sequence_number`` is refused with +``OUT_OF_RANGE`` for the same reason: the relay's numbering is per call. A +zero cursor is live only: nothing is kept between calls, so there is no +history and no state snapshot to serve (the Rust relay in +``crates/engine_servicer`` keeps both for the servicer's lifetime and serves a +``KvSnapshotChunk`` snapshot once its history window has rolled). + +Decoding is lenient, as the engines evolve their events by adding optional +keys: batches are ``[ts, events, rank]``, events are tagged maps (``type``) or +the legacy tag-first arrays, unknown keys are ignored, missing or unreadable +fields cost that event and not the batch, and ``token_ids`` cells may be an +int or a ``[token, next_token]`` bigram (Eagle-family speculative decoding), +folded to their tokens. + +Hash identity: an integer hash is used as is (vLLM sends the low 64 bits of +the digest unsigned, SGLang the high 64 bits signed; both are 64-bit patterns +the proto carries as ``int64``); a raw digest folds to its last eight bytes +big-endian, the integer vLLM would have sent for it. + +Stores and removals are forwarded one for one. vLLM keeps up to two physical +copies of one hash and removes them one at a time; every copy's store and +every removal go through, and the gateway counts copies per worker and tier. +The relay's own record of a hash (the namespace a child inherits, the digest +the hash check chains on) counts the copies too, capped as the gateway caps +them, and goes with the last removal. Of a hybrid model's KV-cache +groups only the main-attention ones are forwarded; this relay drops every +sliding-window and state-space group's events (the Rust relay additionally +forwards a rank that publishes sliding-window groups only, with the hashes +aligned to the tail of the tokens). + +``SMG_KV_EVENT_HASH_CHECK=sglang|vllm-sha256-cbor`` (or ``relay(..., +hash_check=...)``) turns on engine-hash verification: every store whose +parent is known is rehashed the worker's way (SGLang's per-page SHA-256 +chain; vLLM's ``sha256_cbor`` with the default seed) and mismatches are +counted, never dropped, so a worker with a different algorithm, seed or page +size shows up as a counter. The algorithms and their test vectors mirror +``crates/engine_servicer/src/engine_hash.rs``. +""" + +from __future__ import annotations + +import asyncio +import hashlib +import logging +import os +from collections.abc import AsyncIterator, Awaitable, Callable, Iterable +from dataclasses import dataclass, field +from enum import Enum +from typing import TYPE_CHECKING, Any + +import grpc +import msgspec +from smg_grpc_proto.generated import common_pb2 + +if TYPE_CHECKING: + import zmq + import zmq.asyncio + +# ``zmq`` is imported where the relay opens sockets (``replay_frames`` and +# ``relay``), not here: the rank discovery and normalization half of this +# module serves engines whose servicers carry no ``pyzmq`` (the vLLM +# servicer's unit tests import it on a bare interpreter). + +logger = logging.getLogger(__name__) + +__all__ = [ + "HASH_CHECKS", + "HASH_CHECK_ENV", + "Counts", + "Engine", + "Normalizer", + "RankSource", + "RelayFailure", + "WireBatch", + "WireEvent", + "decode_batch", + "endpoint_for_rank", + "fold_hash", + "rank_sources", + "relay", + "replay_frames", + "sglang_chain", + "sglang_event_int", + "sglang_page", + "sglang_salt_seed", + "vllm_block", + "vllm_chain", + "vllm_event_int", + "vllm_none_hash", +] + +_U64_MASK = 0xFFFF_FFFF_FFFF_FFFF +_I64_SIGN = 0x8000_0000_0000_0000 +_END_SEQ = _U64_MASK.to_bytes(8, "big") # (-1) signed and 2**64-1 unsigned are the same bytes + +MAIN_ATTENTION_KINDS = ("full_attention", "mla_attention", "sink_full_attention") +# The most physical copies of one block counted per record: the gateway's cap. +COPIES_CAP = 8 +_RESIDENCY_AGENT = "kvcr" + +_TIER_BY_MEDIUM = { + "GPU": common_pb2.KV_CACHE_TIER_DEVICE, + "DEVICE": common_pb2.KV_CACHE_TIER_DEVICE, + "CPU": common_pb2.KV_CACHE_TIER_HOST, + "CPU_PINNED": common_pb2.KV_CACHE_TIER_HOST, + "CPU_TIER1": common_pb2.KV_CACHE_TIER_HOST, + "CPU_TIER2": common_pb2.KV_CACHE_TIER_DISK, + "DISK": common_pb2.KV_CACHE_TIER_DISK, + "NVME": common_pb2.KV_CACHE_TIER_DISK, + "STORAGE": common_pb2.KV_CACHE_TIER_DISK, + "EXTERNAL": common_pb2.KV_CACHE_TIER_EXTERNAL, + "NETWORK": common_pb2.KV_CACHE_TIER_EXTERNAL, + "REMOTE": common_pb2.KV_CACHE_TIER_EXTERNAL, + "SHARED": common_pb2.KV_CACHE_TIER_EXTERNAL, +} +_CACHE_LEVEL = { + common_pb2.KV_CACHE_TIER_DEVICE: None, + common_pb2.KV_CACHE_TIER_HOST: 1, + common_pb2.KV_CACHE_TIER_DISK: 2, + common_pb2.KV_CACHE_TIER_EXTERNAL: 3, +} + + +class Engine(Enum): + """Which publisher is on the other end; only the replay reply framing differs.""" + + VLLM = "vllm" + SGLANG = "sglang" + + +# --------------------------------------------------------------------------- +# Wire decoding +# --------------------------------------------------------------------------- + + +def fold_hash(value: object) -> int | None: + """An engine block hash as the proto's signed 64-bit identity; ``None`` if unreadable.""" + if isinstance(value, bool): + return None + if isinstance(value, int): + masked = value & _U64_MASK + elif isinstance(value, (bytes, bytearray)): + masked = int.from_bytes(bytes(value)[-8:], "big") + else: + return None + return masked - (1 << 64) if masked >= _I64_SIGN else masked + + +def _text(value: object) -> str | None: + return value if isinstance(value, str) else None + + +def _unsigned32(value: object) -> int | None: + return ( + value + if isinstance(value, int) and not isinstance(value, bool) and 0 <= value < 2**32 + else None + ) + + +def _signed(value: object) -> int | None: + return value if isinstance(value, int) and not isinstance(value, bool) else None + + +def _hashes(value: object) -> list[int] | None: + if not isinstance(value, (list, tuple)): + return None + folded = [fold_hash(item) for item in value] + return None if any(item is None for item in folded) else folded # type: ignore[return-value] + + +def _tokens(value: object) -> tuple[list[int], list[int] | None] | None: + """Token ids as plain ints, or bigram ``[token, next]`` cells folded to their + tokens; the second item is both words of every bigram (for the engine-hash + check), ``None`` for plain ids.""" + if not isinstance(value, (list, tuple)): + return None + ids: list[int] = [] + words: list[int] = [] + pairs = 0 + for cell in value: + if isinstance(cell, int) and not isinstance(cell, bool): + if not 0 <= cell < 2**32: + return None + ids.append(cell) + elif ( + isinstance(cell, (list, tuple)) + and len(cell) == 2 + and all(isinstance(t, int) and not isinstance(t, bool) and 0 <= t < 2**32 for t in cell) + ): + ids.append(cell[0]) + words.extend(cell) + pairs += 1 + else: + return None + if pairs and pairs != len(value): + return None + return ids, (words if pairs else None) + + +def _extra_key(value: object) -> tuple | None: + """One item of vLLM's untagged per-block extra keys; ``None`` for shapes not modeled.""" + if isinstance(value, bool): + return None + if isinstance(value, str): + return ("text", value) + if isinstance(value, int): + return ("number", value) if -(2**63) <= value < 2**63 else None + if isinstance(value, (bytes, bytearray)): + return ("blob", bytes(value)) + if ( + isinstance(value, (list, tuple)) + and len(value) == 2 + and isinstance(value[0], str) + and isinstance(value[1], int) + and not isinstance(value[1], bool) + ): + return ("multimodal", value[0], value[1]) + return None + + +def _extra_keys(value: object) -> list[list[tuple] | None] | None: + if not isinstance(value, (list, tuple)): + return None + per_block: list[list[tuple] | None] = [] + for keys in value: + if isinstance(keys, (list, tuple)): + per_block.append( + [key for key in (_extra_key(item) for item in keys) if key is not None] + ) + else: + per_block.append(None) + return per_block + + +_TAIL_KEYS = ( + "medium", + "group_idx", + "kv_cache_spec_kind", + "kv_cache_spec_sliding_window", + "locality", + "ownership", + "session_id", +) + +# The legacy array layout, slots in the order vLLM's msgspec structs declare +# their fields; trailing defaults are omitted on the wire. +_STORED_SLOTS = ( + "block_hashes", + "parent_block_hash", + "token_ids", + "block_size", + "lora_id", + "medium", + "lora_name", + "extra_keys", + "group_idx", + "kv_cache_spec_kind", + "kv_cache_spec_sliding_window", + "locality", + "ownership", + "session_id", +) +_REMOVED_SLOTS = ("block_hashes", "medium", "group_idx", "locality", "ownership") + + +@dataclass +class WireEvent: + """One decoded event. ``kind`` is the engine's tag, ``"unknown"`` for a tag + the relay does not convert, or ``"malformed"`` (``missing`` names the field + that was absent or unreadable).""" + + kind: str + missing: str | None = None + block_hashes: list[int] = field(default_factory=list) + parent_block_hash: int | None = None + token_ids: list[int] = field(default_factory=list) + bigrams: bool = False + bigram_words: list[int] | None = None + block_size: int = 0 + lora_id: int | None = None + lora_name: str | None = None + cache_salt: str | None = None + extra_keys: list[list[tuple] | None] | None = None + medium: str | None = None + group_idx: int | None = None + kv_cache_spec_kind: str | None = None + kv_cache_spec_sliding_window: int | None = None + locality: str | None = None + ownership: str | None = None + session_id: str | None = None + + +@dataclass +class WireBatch: + ts: float + events: list[WireEvent] + dp_rank: int | None + + +def _event_from_fields(kind: str, raw: dict[str, Any]) -> WireEvent: + tail = {key: raw.get(key) for key in _TAIL_KEYS} + tail_parsed = { + "medium": _text(tail["medium"]), + "group_idx": _unsigned32(tail["group_idx"]), + "kv_cache_spec_kind": _text(tail["kv_cache_spec_kind"]), + "kv_cache_spec_sliding_window": _unsigned32(tail["kv_cache_spec_sliding_window"]), + "locality": _text(tail["locality"]), + "ownership": _text(tail["ownership"]), + "session_id": _text(tail["session_id"]), + } + if kind == "BlockStored": + hashes = _hashes(raw.get("block_hashes")) + if hashes is None: + return WireEvent("malformed", "block_hashes") + tokens = _tokens(raw.get("token_ids")) + if tokens is None: + return WireEvent("malformed", "token_ids") + block_size = _signed(raw.get("block_size")) + if block_size is None: + return WireEvent("malformed", "block_size") + return WireEvent( + kind, + block_hashes=hashes, + parent_block_hash=fold_hash(raw.get("parent_block_hash")), + token_ids=tokens[0], + bigrams=tokens[1] is not None, + bigram_words=tokens[1], + block_size=block_size, + lora_id=_signed(raw.get("lora_id")), + lora_name=_text(raw.get("lora_name")), + cache_salt=_text(raw.get("cache_salt")), + extra_keys=_extra_keys(raw.get("extra_keys")), + **tail_parsed, + ) + if kind == "BlockRemoved": + hashes = _hashes(raw.get("block_hashes")) + if hashes is None: + return WireEvent("malformed", "block_hashes") + return WireEvent(kind, block_hashes=hashes, **tail_parsed) + if kind == "AllBlocksCleared": + return WireEvent(kind, ownership=tail_parsed["ownership"]) + return WireEvent("unknown") + + +def parse_event(raw: object) -> WireEvent: + """A tagged map or a tag-first array as a :class:`WireEvent`.""" + if isinstance(raw, dict): + kind = raw.get("type") + if not isinstance(kind, str): + return WireEvent("malformed", "type") + return _event_from_fields(kind, raw) + if isinstance(raw, (list, tuple)): + if not raw or not isinstance(raw[0], str): + return WireEvent("malformed", "type") + kind, slots = raw[0], raw[1:] + names = { + "BlockStored": _STORED_SLOTS, + "BlockRemoved": _REMOVED_SLOTS, + "AllBlocksCleared": (), + }.get(kind) + if names is None: + return WireEvent("unknown") + return _event_from_fields(kind, dict(zip(names, slots))) + return WireEvent("malformed", "event") + + +_decoder = msgspec.msgpack.Decoder() + + +def decode_batch(payload: bytes) -> WireBatch: + """A publisher payload, ``[ts, events, rank]``, leniently decoded. + + Raises ``ValueError`` when the envelope itself is not a batch; a bad event + inside a good envelope becomes a ``malformed`` event instead. + """ + raw = _decoder.decode(payload) + if not isinstance(raw, (list, tuple)) or len(raw) < 2 or not isinstance(raw[1], (list, tuple)): + raise ValueError("KV event batch is not [ts, events, rank]") + ts = raw[0] if isinstance(raw[0], (int, float)) and not isinstance(raw[0], bool) else 0.0 + rank = raw[2] if len(raw) > 2 else None + rank = rank if isinstance(rank, int) and not isinstance(rank, bool) else None + return WireBatch(ts=float(ts), events=[parse_event(item) for item in raw[1]], dp_rank=rank) + + +# --------------------------------------------------------------------------- +# Normalization +# --------------------------------------------------------------------------- + + +@dataclass +class Counts: + """What one stream has forwarded and dropped, by ``DropReason::as_str`` name.""" + + forwarded_stored: int = 0 + forwarded_removed: int = 0 + forwarded_cleared: int = 0 + duplicate_stores: int = 0 + bigram_stores: int = 0 + hash_checked: int = 0 + hash_mismatch: int = 0 + hash_unverifiable: int = 0 + dropped: dict[str, int] = field(default_factory=dict) + + +def tier_of(medium: str | None) -> int | None: + """The tier a medium names: the device when absent, ``None`` when unknown.""" + if medium is None: + return common_pb2.KV_CACHE_TIER_DEVICE + return _TIER_BY_MEDIUM.get(medium.upper()) + + +def _locality_of(locality: str | None) -> int | None: + if locality is None or locality.upper() == "LOCAL": + return common_pb2.KV_CACHE_LOCALITY_LOCAL + return None + + +def _is_residency_agent(ownership: str | None) -> bool: + return ownership is not None and ownership.lower() == _RESIDENCY_AGENT + + +def _salt_from_extra_keys(extra_keys, lora_name: str | None) -> str | None: + """vLLM's cache salt: the first text key of block 0 that is not the LoRA name.""" + if not extra_keys or not extra_keys[0]: + return None + for key in extra_keys[0]: + if key[0] == "text" and key[1] and key[1] != lora_name: + return key[1] + return None + + +def _extra_key_proto(key: tuple) -> common_pb2.KvBlockExtraKey: + if key[0] == "text": + return common_pb2.KvBlockExtraKey(text=key[1]) + if key[0] == "number": + return common_pb2.KvBlockExtraKey(number=key[1]) + if key[0] == "blob": + return common_pb2.KvBlockExtraKey(blob=key[1]) + return common_pb2.KvBlockExtraKey( + multimodal=common_pb2.KvMultimodalKey(identifier=key[1], offset=key[2]) + ) + + +class _RankState: + __slots__ = ("tiers", "groups") + + def __init__(self) -> None: + # tier -> engine hash -> ((lora_name, cache_salt), recomputed digest or + # None, physical copies stored and not yet removed). A record lives + # until its last copy is removed, so a child stored after one copy + # went still finds its parent's namespace and digest. + self.tiers: dict[ + int, dict[int, tuple[tuple[str | None, str | None], bytes | None, int]] + ] = {} + # cache group -> whether it is a main-attention group + self.groups: dict[int, bool] = {} + + +class Normalizer: + """Per-stream normalization: the drop rules, the namespace inheritance and + the counters, as ``kv_wire::Normalizer`` keeps them. ``hash_check`` names + the engine hash to recompute per store (one of :data:`HASH_CHECKS`), or + ``None`` for no verification; an unknown name is logged and ignored.""" + + def __init__(self, hash_check: str | None = None) -> None: + self.ranks: dict[int, _RankState] = {} + self.counts = Counts() + self._event_id = 0 + self.hash_check = _hash_check_name(hash_check) + + def _drop(self, reason: str, event_id: int) -> None: + count = self.counts.dropped.get(reason, 0) + 1 + self.counts.dropped[reason] = count + if count <= 3: + logger.debug("KV event %d not forwarded: %s", event_id, reason) + + def normalize_batch( + self, batch: WireBatch, sequence_number: int, dp_rank: int | None = None + ) -> common_pb2.KvEventBatch: + """A whole batch as its proto. ``dp_rank`` (the socket's rank) wins over + the payload's; event ids advance once per event, forwarded or not.""" + rank = batch.dp_rank if dp_rank is None else dp_rank + proto = common_pb2.KvEventBatch(sequence_number=sequence_number, timestamp=batch.ts) + if rank is not None: + proto.dp_rank = rank + for event in batch.events: + self._event_id += 1 + converted = self.normalize(event, rank, self._event_id) + if converted is not None: + proto.events.append(converted) + return proto + + def normalize( + self, event: WireEvent, dp_rank: int | None, event_id: int + ) -> common_pb2.KvCacheEvent | None: + rank = -1 if dp_rank is None else dp_rank + if event.kind == "unknown": + self._drop("unknown_type", event_id) + return None + if event.kind == "malformed": + logger.debug("KV event %d field unreadable: %s", event_id, event.missing) + self._drop("malformed", event_id) + return None + if event.kind == "AllBlocksCleared": + if _is_residency_agent(event.ownership): + self._drop("unsupported_ownership", event_id) + return None + self.ranks.pop(rank, None) + self.counts.forwarded_cleared += 1 + cleared = common_pb2.KvCacheCleared() + if event.ownership is not None: + cleared.ownership = event.ownership + return common_pb2.KvCacheEvent(event_id=event_id, cleared=cleared) + if event.kind == "BlockStored": + data = self._stored(event, rank, event_id) + return None if data is None else common_pb2.KvCacheEvent(event_id=event_id, stored=data) + if event.kind == "BlockRemoved": + data = self._removed(event, rank, event_id) + return ( + None if data is None else common_pb2.KvCacheEvent(event_id=event_id, removed=data) + ) + self._drop("unknown_type", event_id) + return None + + def _admit(self, event: WireEvent, rank: int, event_id: int, learn_group: bool): + """The shared gates: ownership, locality, medium, cache group.""" + if _is_residency_agent(event.ownership): + self._drop("unsupported_ownership", event_id) + return None + locality = _locality_of(event.locality) + if locality is None: + self._drop("non_local_locality", event_id) + return None + tier = tier_of(event.medium) + if tier is None: + self._drop("unknown_medium", event_id) + return None + if event.group_idx is not None: + state = self.ranks.setdefault(rank, _RankState()) + if event.kv_cache_spec_kind is not None: + main = event.kv_cache_spec_kind in MAIN_ATTENTION_KINDS + if learn_group: + state.groups[event.group_idx] = main + else: + # A kind-less event follows its learned group; an unknown + # group counts as main, as a single-group publisher's does. + main = state.groups.get(event.group_idx, True) + if not main: + self._drop("non_main_attention_group", event_id) + return None + return tier, locality + + def _stored(self, event: WireEvent, rank: int, event_id: int): + admitted = self._admit(event, rank, event_id, learn_group=True) + if admitted is None: + return None + tier, locality = admitted + if not event.block_hashes or not event.token_ids: + self._drop("placeholder", event_id) + return None + width = event.block_size + if width <= 0 or width >= 2**31 or len(event.block_hashes) * width != len(event.token_ids): + self._drop("unaligned_blocks", event_id) + return None + seen = set() + if event.parent_block_hash is not None: + seen.add(event.parent_block_hash) + for block_hash in event.block_hashes: + if block_hash in seen: + self._drop("self_referencing_hashes", event_id) + return None + seen.add(block_hash) + + lora_name = event.lora_name or None + cache_salt = (event.cache_salt or None) or _salt_from_extra_keys( + event.extra_keys, lora_name + ) + blocks_state = self.ranks.setdefault(rank, _RankState()).tiers.setdefault(tier, {}) + if (lora_name is None or cache_salt is None) and event.parent_block_hash is not None: + parent = blocks_state.get(event.parent_block_hash) + if parent is not None: + lora_name = lora_name if lora_name is not None else parent[0][0] + cache_salt = cache_salt if cache_salt is not None else parent[0][1] + namespace = (lora_name, cache_salt) + all_seen = True + for block_hash in event.block_hashes: + record = blocks_state.get(block_hash) + if record is None: + all_seen = False + blocks_state[block_hash] = (namespace, None, 1) + else: + # A second physical copy: the digest stays, one more copy counted. + blocks_state[block_hash] = (namespace, record[1], min(record[2] + 1, COPIES_CAP)) + if all_seen: + self.counts.duplicate_stores += 1 + if event.bigrams: + self.counts.bigram_stores += 1 + if self.hash_check is not None: + self._verify_hashes(blocks_state, event, width, lora_name, cache_salt) + + cache_level = _CACHE_LEVEL[tier] + extra_keys = event.extra_keys or [] + blocks = [] + for index, block_hash in enumerate(event.block_hashes): + block = common_pb2.KvBlock( + block_hash=block_hash, + token_ids=event.token_ids[index * width : (index + 1) * width], + block_size=width, + ) + if event.lora_id is not None: + block.lora_id = event.lora_id + if cache_level is not None: + block.cache_level = cache_level + keys = extra_keys[index] if index < len(extra_keys) else None + if keys: + block.extra_keys.extend(_extra_key_proto(key) for key in keys) + blocks.append(block) + stored = common_pb2.KvBlocksStored(blocks=blocks, tier=tier, locality=locality) + if event.parent_block_hash is not None: + stored.parent_block_hash = event.parent_block_hash + for name, value in ( + ("medium", event.medium), + ("group_idx", event.group_idx), + ("kv_cache_spec_kind", event.kv_cache_spec_kind), + ("kv_cache_spec_sliding_window", event.kv_cache_spec_sliding_window), + ("ownership", event.ownership), + ("session_id", event.session_id), + ("lora_name", lora_name), + ("cache_salt", cache_salt), + ): + if value is not None: + setattr(stored, name, value) + self.counts.forwarded_stored += 1 + return stored + + def _removed(self, event: WireEvent, rank: int, event_id: int): + admitted = self._admit(event, rank, event_id, learn_group=False) + if admitted is None: + return None + tier, locality = admitted + state = self.ranks.get(rank) + if state is not None: + blocks_state = state.tiers.get(tier) + if blocks_state: + # One physical copy goes; the record goes with the last. + for block_hash in event.block_hashes: + record = blocks_state.get(block_hash) + if record is None: + continue + if record[2] <= 1: + del blocks_state[block_hash] + else: + blocks_state[block_hash] = (record[0], record[1], record[2] - 1) + removed = common_pb2.KvBlocksRemoved( + block_hashes=event.block_hashes, tier=tier, locality=locality + ) + cache_level = _CACHE_LEVEL[tier] + if cache_level is not None: + removed.cache_level = cache_level + for name, value in ( + ("medium", event.medium), + ("group_idx", event.group_idx), + ("ownership", event.ownership), + ): + if value is not None: + setattr(removed, name, value) + self.counts.forwarded_removed += 1 + return removed + + def _verify_hashes(self, blocks_state, event: WireEvent, width: int, lora_name, cache_salt): + """Rehash the store's blocks the worker's way and count the outcome; + what is forwarded never changes. Digests stay on the records so + children can chain on them.""" + check = self.hash_check + blocks = len(event.block_hashes) + if event.parent_block_hash is not None: + record = blocks_state.get(event.parent_block_hash) + prior = record[1] if record is not None else None + if prior is None: + self.counts.hash_unverifiable += blocks + return + elif check == "sglang" and cache_salt: + prior = sglang_salt_seed(cache_salt) + else: + prior = None # vllm_block applies NONE_HASH itself + if check == "vllm-sha256-cbor" and ( + lora_name is not None or event.lora_id is not None or event.bigram_words is not None + ): + self.counts.hash_unverifiable += blocks + return + extra_keys = event.extra_keys or [] + for index, block_hash in enumerate(event.block_hashes): + tokens = event.token_ids[index * width : (index + 1) * width] + if check == "sglang": + words = ( + event.bigram_words[index * 2 * width : (index + 1) * 2 * width] + if event.bigram_words is not None + else tokens + ) + digest = sglang_page(prior, words) + expected = sglang_event_int(digest) + else: + keys = _vllm_hash_keys( + extra_keys[index] if index < len(extra_keys) else None, index + ) + if keys is _UNVERIFIABLE: + self.counts.hash_unverifiable += blocks - index + return + digest = vllm_block(prior, tokens, keys) + expected = vllm_event_int(digest) + self.counts.hash_checked += 1 + if expected != block_hash: + self.counts.hash_mismatch += 1 + record = blocks_state.get(block_hash) + if record is not None: + blocks_state[block_hash] = (record[0], digest, record[2]) + prior = digest + + +# --------------------------------------------------------------------------- +# Engine-exact hashes (verification only; the index uses SMG content hashes) +# --------------------------------------------------------------------------- + +HASH_CHECK_ENV = "SMG_KV_EVENT_HASH_CHECK" +HASH_CHECKS = ("sglang", "vllm-sha256-cbor") +_UNVERIFIABLE = object() + + +def _hash_check_name(value: str | None) -> str | None: + if value is None: + return None + name = value.strip().lower().replace("_", "-") + if not name: + return None + if name not in HASH_CHECKS: + logger.warning("%s names no known engine hash (%r); check off", HASH_CHECK_ENV, value) + return None + return name + + +def sglang_salt_seed(cache_salt: str) -> bytes: + """The chain seed of a salted SGLang request.""" + return hashlib.sha256(b"sglang-cache-salt-v1\0" + cache_salt.encode("utf-8")).digest() + + +def sglang_page(prior: bytes | None, words) -> bytes: + """One SGLang page: the prior digest (if any), then each word as four + little-endian bytes (both words of every bigram under Eagle hashing).""" + hasher = hashlib.sha256() + if prior is not None: + hasher.update(prior) + for word in words: + hasher.update(int(word).to_bytes(4, "little")) + return hasher.digest() + + +def sglang_event_int(digest: bytes) -> int: + """The integer SGLang publishes: the first eight digest bytes, big-endian, signed.""" + return int.from_bytes(digest[:8], "big", signed=True) + + +def sglang_chain(tokens, page_size: int, prior: bytes | None = None) -> list[tuple[bytes, int]]: + """Every full page of ``tokens`` chained from ``prior``: (digest, published int).""" + out = [] + if page_size <= 0: + return out + for start in range(0, len(tokens) - page_size + 1, page_size): + prior = sglang_page(prior, tokens[start : start + page_size]) + out.append((prior, sglang_event_int(prior))) + return out + + +def _cbor_head(out: bytearray, major: int, value: int) -> None: + major <<= 5 + if value < 24: + out.append(major | value) + elif value <= 0xFF: + out += bytes((major | 24, value)) + elif value <= 0xFFFF: + out += bytes((major | 25,)) + value.to_bytes(2, "big") + elif value <= 0xFFFF_FFFF: + out += bytes((major | 26,)) + value.to_bytes(4, "big") + else: + out += bytes((major | 27,)) + value.to_bytes(8, "big") + + +def _cbor(out: bytearray, value) -> None: + """Canonical CBOR for the shapes vLLM's hash input uses (what ``cbor2.dumps( + value, canonical=True)`` emits for them): ints, bytes, text, tuples/lists, None.""" + if value is None: + out.append(0xF6) + elif isinstance(value, bool): + raise TypeError("booleans are not part of a block hash input") + elif isinstance(value, int): + if value >= 0: + _cbor_head(out, 0, value) + else: + _cbor_head(out, 1, -1 - value) + elif isinstance(value, (bytes, bytearray)): + _cbor_head(out, 2, len(value)) + out += value + elif isinstance(value, str): + encoded = value.encode("utf-8") + _cbor_head(out, 3, len(encoded)) + out += encoded + elif isinstance(value, (list, tuple)): + _cbor_head(out, 4, len(value)) + for item in value: + _cbor(out, item) + else: + raise TypeError(f"unsupported hash input {type(value).__name__}") + + +_VLLM_NONE_HASH_SEED = "vllm-none-hash" + + +def vllm_none_hash() -> bytes: + """``NONE_HASH``: sha256 of the CBOR text ``vllm-none-hash`` (PYTHONHASHSEED unset).""" + out = bytearray() + _cbor(out, _VLLM_NONE_HASH_SEED) + return hashlib.sha256(out).digest() + + +def vllm_block(parent: bytes | None, tokens, extra_keys=None) -> bytes: + """One vLLM ``sha256_cbor`` block hash: ``sha256(cbor([parent or NONE, + [tokens...], extra_keys or None]))`` with the keys as the tagged tuples the + engine folds in (``("lora", name, path)``, ``("mm", identifier, offset)``, + ``("cache_salt", salt)``, ``("prompt_embeds", digest)``).""" + out = bytearray() + _cbor( + out, + ( + parent if parent is not None else vllm_none_hash(), + tuple(int(t) for t in tokens), + tuple(extra_keys) if extra_keys else None, + ), + ) + return hashlib.sha256(out).digest() + + +def vllm_event_int(digest: bytes) -> int: + """The integer vLLM publishes: the low 64 bits, carried as the same bits in an int64.""" + return fold_hash(digest) + + +def vllm_chain(tokens, block_size: int, parent: bytes | None = None) -> list[tuple[bytes, int]]: + """Every full block of ``tokens`` chained from ``parent`` with no extra keys.""" + out = [] + if block_size <= 0: + return out + for start in range(0, len(tokens) - block_size + 1, block_size): + parent = vllm_block(parent, tokens[start : start + block_size]) + out.append((parent, vllm_event_int(parent))) + return out + + +def _vllm_hash_keys(keys, index: int): + """vLLM's untagged event keys as the tagged keys inside the hash: block 0's + text is the cache salt (a LoRA request was excluded before), a pair is a + multimodal item, bytes are a prompt-embeddings digest.""" + if not keys: + return None + tagged = [] + for key in keys: + if key[0] == "text" and index == 0: + tagged.append(("cache_salt", key[1])) + elif key[0] == "multimodal": + tagged.append(("mm", key[1], key[2])) + elif key[0] == "blob": + tagged.append(("prompt_embeds", key[1])) + else: + return _UNVERIFIABLE + return tagged + + +# --------------------------------------------------------------------------- +# Endpoints and ranks +# --------------------------------------------------------------------------- + + +def endpoint_for_rank(endpoint: str, dp_rank: int) -> str: + """A connectable SUB address for rank ``dp_rank`` of a publisher at ``endpoint``. + + Bind wildcards (``*``, ``0.0.0.0``) become loopback; tcp ports are offset + by the rank, as both engines' publishers offset theirs; other transports + get no port arithmetic. + """ + resolved = endpoint.replace("*", "127.0.0.1").replace("0.0.0.0", "127.0.0.1") + if resolved.startswith("tcp://") and dp_rank: + host, sep, port = resolved.rpartition(":") + if sep and port.isdigit() and int(port) != 0: + return f"{host}:{int(port) + dp_rank}" + return resolved + + +@dataclass(frozen=True) +class RankSource: + """One rank's publisher: where to subscribe and where to ask for replay.""" + + rank: int + endpoint: str + replay_endpoint: str | None = None + + +def rank_sources(config: object, ranks: Iterable[int]) -> list[RankSource]: + """Sources for ``ranks`` of a publisher configured by ``config`` (an + object with ``endpoint`` and optional ``replay_endpoint``).""" + endpoint = str(getattr(config, "endpoint", "") or "") + replay = getattr(config, "replay_endpoint", None) + return [ + RankSource( + rank=rank, + endpoint=endpoint_for_rank(endpoint, rank), + replay_endpoint=endpoint_for_rank(str(replay), rank) if replay else None, + ) + for rank in ranks + ] + + +# --------------------------------------------------------------------------- +# Streaming +# --------------------------------------------------------------------------- + + +class RelayFailure(Exception): + """The stream can no longer be trusted; the gateway must clear and resubscribe.""" + + +async def replay_frames( + endpoint: str, + start: int, + *, + timeout: float, + context: zmq.asyncio.Context | None = None, +) -> AsyncIterator[tuple[int, bytes]]: + """Ask a publisher's replay ROUTER for every buffered batch from ``start``. + + The request is ``[b"", start as 8-byte big-endian]`` on a DEALER. Replies + are ``[b"", seq, payload]`` (SGLang) or ``[b"", topic, seq, payload]`` + (vLLM) and end with the all-ones sequence. The first reply must be + ``start`` itself and the rest contiguous; anything else, or no reply in + ``timeout`` seconds, raises :class:`RelayFailure`. + """ + import zmq + import zmq.asyncio + + ctx = context or zmq.asyncio.Context.instance() + dealer = ctx.socket(zmq.DEALER) + dealer.setsockopt(zmq.LINGER, 0) + try: + dealer.connect(endpoint) + await dealer.send_multipart([b"", start.to_bytes(8, "big")]) + expected = start + while True: + if not await dealer.poll(timeout=int(timeout * 1000)): + raise RelayFailure(f"KV event replay from {endpoint} timed out at {expected}") + frames = await dealer.recv_multipart() + if len(frames) == 3: + seq_bytes, payload = frames[1], frames[2] + elif len(frames) == 4: + seq_bytes, payload = frames[2], frames[3] + else: + raise RelayFailure(f"malformed KV event replay reply ({len(frames)} frames)") + if len(seq_bytes) != 8: + raise RelayFailure("malformed KV event replay sequence") + if seq_bytes == _END_SEQ: + if expected == start: + raise RelayFailure( + f"KV event replay history at {endpoint} no longer holds {start}" + ) + return + seq = int.from_bytes(seq_bytes, "big") + if seq != expected: + raise RelayFailure( + f"KV event replay from {endpoint} is not contiguous: expected {expected}, got {seq}" + ) + yield seq, payload + expected += 1 + except zmq.ZMQError as error: + raise RelayFailure(f"KV event replay transport failed: {error}") from error + finally: + dealer.close(linger=0) + + +class _RankStream: + __slots__ = ("source", "socket", "cursor", "replayed") + + def __init__(self, source: RankSource, socket: zmq.asyncio.Socket) -> None: + self.source = source + self.socket = socket + self.cursor: int | None = None + # The sequence range the last gap replay forwarded, while it is still + # open. A replay usually runs past the live batch that exposed the gap, + # and the live batches the publisher sent meanwhile are still queued + # on the SUB socket: they arrive below the cursor the replay left and + # are duplicates of what it forwarded, not a restarted publisher. The + # queue drains in order, so the first live batch past the range closes + # it; a sequence inside it after that can only be a restart. + self.replayed: tuple[int, int] | None = None + + def covered_by_replay(self, seq: int) -> bool: + return self.replayed is not None and self.replayed[0] <= seq <= self.replayed[1] + + +async def relay( + sources: list[RankSource], + engine: Engine, + start_sequence_number: int, + context: grpc.aio.ServicerContext, + *, + topic: str = "", + hwm: int | None = None, + recv_timeout: float = 1.0, + replay_timeout: float = 5.0, + zmq_context: zmq.asyncio.Context | None = None, + normalizer: Normalizer | None = None, + hash_check: str | None = None, + on_counts: Callable[[Counts], Awaitable[None] | None] | None = None, +) -> AsyncIterator[common_pb2.KvEventBatch]: + """Relay every rank in ``sources`` into one ``KvEventBatch`` stream. + + See the module docstring for the contract. ``engine`` is informational + (both replay reply framings are accepted); ``hash_check`` (default: the + ``SMG_KV_EVENT_HASH_CHECK`` environment variable) turns engine-hash + verification on; ``on_counts`` is called with the stream's counters when + it ends. + """ + import zmq + import zmq.asyncio + + if start_sequence_number: + await context.abort( + grpc.StatusCode.OUT_OF_RANGE, + "KV event replay cursors are per subscription; resubscribe with zero and rebuild", + ) + return + if not sources: + await context.abort( + grpc.StatusCode.UNIMPLEMENTED, "KV cache events have no publisher to relay" + ) + return + + ctx = zmq_context or zmq.asyncio.Context.instance() + if normalizer is None: + normalizer = Normalizer( + hash_check if hash_check is not None else os.environ.get(HASH_CHECK_ENV) + ) + streams: dict[Any, _RankStream] = {} + poller = zmq.asyncio.Poller() + relay_seq = 0 + sent_headers = False + + def forward(payload: bytes, stream: _RankStream, seq: int, kind: str): + nonlocal relay_seq + try: + batch = decode_batch(payload) + except Exception as error: # noqa: BLE001 - one bad payload must not end the stream + logger.warning( + "Failed to decode %s KV event batch %d from rank %d: %s", + kind, + seq, + stream.source.rank, + error, + ) + return None + relay_seq += 1 + return normalizer.normalize_batch(batch, relay_seq, stream.source.rank) + + try: + for source in sources: + socket = ctx.socket(zmq.SUB) + if hwm is not None: + socket.setsockopt(zmq.RCVHWM, int(hwm)) + socket.subscribe(topic.encode("utf-8")) + socket.connect(source.endpoint) + streams[socket] = _RankStream(source, socket) + poller.register(socket, zmq.POLLIN) + logger.info( + "SubscribeKvEvents: subscribed to %d %s rank(s): %s", + len(sources), + engine.value, + ", ".join(source.endpoint for source in sources), + ) + await context.send_initial_metadata(()) + sent_headers = True + + while not context.cancelled(): + ready = await poller.poll(timeout=int(recv_timeout * 1000)) + for socket, _ in ready: + stream = streams[socket] + frames = await socket.recv_multipart() + if len(frames) < 3 or len(frames[1]) != 8: + continue + seq = int.from_bytes(frames[1], "big") + cursor = stream.cursor + if cursor is not None: + if seq == cursor or stream.covered_by_replay(seq): + continue # a duplicate of what a replay already forwarded + if seq < cursor: + raise RelayFailure( + f"publisher of rank {stream.source.rank} restarted its sequence at " + f"{seq} after {cursor}" + ) + if seq > cursor + 1: + replay_endpoint = stream.source.replay_endpoint + if replay_endpoint is None: + raise RelayFailure( + f"rank {stream.source.rank} skipped from {cursor} to {seq} and has " + "no replay endpoint" + ) + logger.warning( + "KV events of rank %d skipped from %d to %d; replaying", + stream.source.rank, + cursor, + seq, + ) + async for replayed_seq, payload in replay_frames( + replay_endpoint, cursor + 1, timeout=replay_timeout, context=ctx + ): + proto = forward(payload, stream, replayed_seq, "replayed") + stream.cursor = replayed_seq + stream.replayed = (cursor + 1, replayed_seq) + if proto is not None: + yield proto + if stream.cursor < seq - 1: + raise RelayFailure( + f"KV event replay of rank {stream.source.rank} ended at " + f"{stream.cursor}, before {seq - 1}" + ) + if seq <= stream.cursor: + continue # the replay ran past the live message + stream.cursor = seq + if stream.replayed is not None and seq > stream.replayed[1]: + stream.replayed = None # the live stream is past the replay's range + proto = forward(frames[2], stream, seq, "live") + if proto is not None: + yield proto + except asyncio.CancelledError: + pass + except RelayFailure as error: + logger.warning("SubscribeKvEvents: %s; the gateway will clear and resubscribe", error) + await context.abort( + grpc.StatusCode.DATA_LOSS if sent_headers else grpc.StatusCode.OUT_OF_RANGE, str(error) + ) + finally: + for socket in streams: + socket.close(linger=0) + if on_counts is not None: + result = on_counts(normalizer.counts) + if asyncio.iscoroutine(result): + await result + logger.info("SubscribeKvEvents: stream closed; %s", normalizer.counts) diff --git a/grpc_servicer/smg_grpc_servicer/sglang/kv_events.py b/grpc_servicer/smg_grpc_servicer/sglang/kv_events.py index a9f1f54018..9d70ccfd22 100644 --- a/grpc_servicer/smg_grpc_servicer/sglang/kv_events.py +++ b/grpc_servicer/smg_grpc_servicer/sglang/kv_events.py @@ -1,133 +1,59 @@ -"""SGLang KV-event transport, independent of engine imports.""" +"""SGLang KV-event transport: every DP rank's publisher relayed through +:mod:`smg_grpc_servicer.kv_relay`, independent of engine imports. + +SGLang runs one KV-event publisher per independent KV cache, ``dp_size`` of +them (under DP attention the attention-DP ranks are those replicas), each on +``endpoint_port_base + rank`` with its own sequence counter and, when +configured, its own replay ROUTER on ``replay_endpoint_port_base + rank``. +This is what ``/server_info.kv_events`` advertises; in-process the same +figures come from the server args and the KV-events config. +""" + +from __future__ import annotations import logging -from collections.abc import AsyncIterator, Callable -from contextlib import aclosing +from collections.abc import AsyncIterator import grpc -import zmq -import zmq.asyncio from smg_grpc_proto.generated import common_pb2 -from smg_grpc_servicer.kv_events import endpoint_for_rank +from smg_grpc_servicer.kv_relay import ( + Engine, + RankSource, + endpoint_for_rank, + rank_sources, + relay, +) -logger = logging.getLogger(__name__) -_REPLAY_TIMEOUT_MS = 5000 -_END_SEQ = (-1).to_bytes(8, "big", signed=True) +__all__ = ["dp_rank_count", "endpoint_for_rank", "sources", "subscribe_kv_events"] +logger = logging.getLogger(__name__) -class ReplayUnavailable(Exception): - """The publisher cannot provide a contiguous, decodable replay.""" +def dp_rank_count(server_args: object) -> int: + """How many KV-event publishers this SGLang instance runs: ``dp_size``.""" + dp_size = getattr(server_args, "dp_size", None) + return int(dp_size) if isinstance(dp_size, int) and dp_size > 0 else 1 -async def replay_frames(endpoint: str, cursor: int) -> AsyncIterator[tuple[int, bytes]]: - """Read SGLang's ROUTER replay protocol through a DEALER socket. - The publisher accepts an inclusive start sequence and ends with -1. - Request the first missing batch, verifying continuity before forwarding. - """ - replay = zmq.asyncio.Context.instance().socket(zmq.DEALER) - try: - replay.connect(endpoint_for_rank(endpoint, 0)) - await replay.send_multipart([b"", (cursor + 1).to_bytes(8, "big")]) - expected = cursor + 1 - received_any = False - while True: - if not await replay.poll(timeout=_REPLAY_TIMEOUT_MS): - raise ReplayUnavailable("KV event replay timed out") - frames = await replay.recv_multipart() - if len(frames) != 3 or frames[0] != b"" or len(frames[1]) != 8: - raise ReplayUnavailable("Malformed KV event replay response") - if frames[1] == _END_SEQ: - # Empty replay cannot distinguish an idle publisher from a - # restarted publisher whose counter is behind our cursor. - if not received_any: - raise ReplayUnavailable("KV event replay history is unavailable") - return - seq = int.from_bytes(frames[1], "big") - if seq != expected: - raise ReplayUnavailable( - f"KV event replay is incomplete: expected {expected}, received {seq}" - ) - yield seq, frames[2] - received_any = True - expected += 1 - except zmq.ZMQError as exc: - raise ReplayUnavailable("KV event replay transport failed") from exc - finally: - replay.close(linger=0) +def sources(config: object, server_args: object) -> list[RankSource]: + """One source per rank of the publisher ``config`` describes.""" + return rank_sources(config, range(dp_rank_count(server_args))) async def subscribe_kv_events( config: object, + server_args: object, start_sequence_number: int, context: grpc.aio.ServicerContext, - decode: Callable[[bytes], object], - convert: Callable[[object, int], common_pb2.KvEventBatch], ) -> AsyncIterator[common_pb2.KvEventBatch]: - """Replay a resume cursor when configured, then forward live events. - - Unrecoverable replay clears stale gateway mappings via OUT_OF_RANGE - before headers, or DATA_LOSS once streaming. Zero starts live rebuilding, - not a full cache snapshot. - """ - replay_endpoint = getattr(config, "replay_endpoint", None) - if start_sequence_number and (not replay_endpoint or start_sequence_number == 2**64 - 1): - await context.abort( - grpc.StatusCode.OUT_OF_RANGE, - "KV event replay is unavailable; resubscribe with zero for live events", - ) - return - - # Each DP rank has independent sequence numbers. Keep the existing rank-0 - # subscription until the gateway supports per-rank indexes. - endpoint = endpoint_for_rank(config.endpoint, 0) - sub = zmq.asyncio.Context.instance().socket(zmq.SUB) - sent_headers = False - replayed_seq = None - try: - sub.setsockopt(zmq.RCVHWM, getattr(config, "hwm", 100_000)) - sub.subscribe(config.topic.encode("utf-8")) - sub.connect(endpoint) - logger.info("SubscribeKvEvents: connected to ZMQ endpoint %s", endpoint) - # Subscribe before replay so live events can queue during recovery. - # A PUB/SUB join race still surfaces as a native gap for another replay. - if start_sequence_number: - async with aclosing(replay_frames(replay_endpoint, start_sequence_number)) as frames: - async for seq, payload in frames: - try: - batch = convert(decode(payload), seq) - except Exception as exc: - raise ReplayUnavailable("Failed to decode KV event replay") from exc - if not sent_headers: - await context.send_initial_metadata(()) - sent_headers = True - yield batch - replayed_seq = seq - if not sent_headers: - await context.send_initial_metadata(()) - sent_headers = True - while not context.cancelled(): - # Cancelling recv_multipart on an idle timeout can lose a message. - if not await sub.poll(timeout=1000): - continue - frames = await sub.recv_multipart() - if len(frames) < 3: - continue - seq = int.from_bytes(frames[1], "big") - if replayed_seq is not None and seq <= replayed_seq: - continue - try: - batch = decode(frames[2]) - except Exception as exc: # noqa: BLE001 - logger.warning("Failed to decode KV event batch: %s", exc) - continue - yield convert(batch, seq) - except ReplayUnavailable as exc: - await context.abort( - grpc.StatusCode.DATA_LOSS if sent_headers else grpc.StatusCode.OUT_OF_RANGE, - str(exc), - ) - finally: - sub.close(linger=0) - logger.info("SubscribeKvEvents: stream closed") + """Relay every rank's events; see :func:`smg_grpc_servicer.kv_relay.relay`.""" + async for batch in relay( + sources(config, server_args), + Engine.SGLANG, + start_sequence_number, + context, + topic=str(getattr(config, "topic", "") or ""), + hwm=getattr(config, "hwm", None), + ): + yield batch diff --git a/grpc_servicer/smg_grpc_servicer/sglang/rust.py b/grpc_servicer/smg_grpc_servicer/sglang/rust.py index 73b8d123b7..affb8e456e 100644 --- a/grpc_servicer/smg_grpc_servicer/sglang/rust.py +++ b/grpc_servicer/smg_grpc_servicer/sglang/rust.py @@ -173,6 +173,9 @@ def server_facts(server_args: Any) -> dict[str, Any]: except ImportError: # the launcher's unit tests run without SGLang sglang_version = "" dp_size = getattr(server_args, "dp_size", None) + kv_events_endpoint, kv_events_replay_endpoint, kv_events_topic = kv_events_publisher( + server_args + ) return { # allow_nan=False: a non-finite float would otherwise become `NaN` text, # which the Rust side rejects as a whole object. @@ -181,9 +184,35 @@ def server_facts(server_args: Any) -> dict[str, Any]: "sglang_version": str(sglang_version), "max_running_requests": int(getattr(server_args, "max_running_requests", None) or 0), "data_parallel_size": int(dp_size) if isinstance(dp_size, int) and dp_size > 0 else 1, + "kv_events_endpoint": kv_events_endpoint, + "kv_events_replay_endpoint": kv_events_replay_endpoint, + "kv_events_topic": kv_events_topic, } +def kv_events_publisher(server_args: Any) -> tuple[str, str, str]: + """The ZMQ KV-event publisher SGLang was told to run (``--kv-events-config``) + as ``(endpoint, replay_endpoint, topic)``, with SGLang's defaults filled in; + empty strings when events are off or the publisher is not ZMQ, and an empty + replay endpoint when SGLang runs no replay socket. The Rust servicer relays + this publisher on ``SubscribeKvEvents`` and asks the replay socket for gaps + and for the batches published before its subscription joined.""" + raw = getattr(server_args, "kv_events_config", None) + if not raw: + return "", "", "" + try: + config = json.loads(raw) if isinstance(raw, str) else dict(raw) + except (TypeError, ValueError): + return "", "", "" + if not isinstance(config, dict) or config.get("publisher", "null") != "zmq": + return "", "", "" + return ( + str(config.get("endpoint") or "tcp://*:5557"), + str(config.get("replay_endpoint") or ""), + str(config.get("topic") or ""), + ) + + # --------------------------------------------------------------------------- # Headless scheduler: this package's launcher, dialing this servicer # --------------------------------------------------------------------------- diff --git a/grpc_servicer/smg_grpc_servicer/sglang/servicer.py b/grpc_servicer/smg_grpc_servicer/sglang/servicer.py index d8c72fb603..e4859116c4 100644 --- a/grpc_servicer/smg_grpc_servicer/sglang/servicer.py +++ b/grpc_servicer/smg_grpc_servicer/sglang/servicer.py @@ -23,10 +23,6 @@ from google.protobuf.timestamp_pb2 import Timestamp from sglang.srt.configs.model_config import ModelConfig from sglang.srt.disaggregation.kv_events import ( - AllBlocksCleared, - BlockRemoved, - BlockStored, - KVEventBatch, KVEventsConfig, ) from sglang.srt.disaggregation.utils import FAKE_BOOTSTRAP_HOST, DisaggregationMode @@ -162,7 +158,6 @@ def __init__( # Parse KV events config for SubscribeKvEvents support self._kv_events_config: KVEventsConfig | None = None - self._kv_event_id_counter = 0 if server_args.kv_events_config: try: self._kv_events_config = KVEventsConfig.from_cli(server_args.kv_events_config) @@ -734,11 +729,8 @@ async def SubscribeKvEvents( request: common_pb2.SubscribeKvEventsRequest, context: grpc.aio.ServicerContext, ) -> AsyncIterator[common_pb2.KvEventBatch]: - """Bridge internal ZMQ KV cache events to gRPC server-streaming. - - Uses the ZMQ publisher's native sequence numbers as gRPC sequence - numbers directly. - """ + """Relay the scheduler's ZMQ KV cache events, every DP rank's publisher, + as one gRPC stream (see ``smg_grpc_servicer.kv_relay``).""" if self._kv_events_config is None: await context.abort( grpc.StatusCode.UNIMPLEMENTED, @@ -747,72 +739,14 @@ async def SubscribeKvEvents( ) return - decoder = msgspec.msgpack.Decoder(KVEventBatch) async for batch in subscribe_kv_events( self._kv_events_config, + self.server_args, request.start_sequence_number, context, - decoder.decode, - self._convert_kv_event_batch, ): yield batch - def _convert_kv_event_batch( - self, raw_batch: KVEventBatch, seq_num: int - ) -> common_pb2.KvEventBatch: - """Convert a ZMQ KVEventBatch to proto KvEventBatch.""" - proto_batch = common_pb2.KvEventBatch( - sequence_number=seq_num, - timestamp=raw_batch.ts, - ) - if raw_batch.attn_dp_rank is not None: - proto_batch.dp_rank = raw_batch.attn_dp_rank - - for event in raw_batch.events: - proto_event = self._convert_kv_event(event) - if proto_event is not None: - proto_batch.events.append(proto_event) - - return proto_batch - - def _convert_kv_event(self, event) -> common_pb2.KvCacheEvent | None: - """Convert a single raw KV event to proto KvCacheEvent.""" - self._kv_event_id_counter += 1 - event_id = self._kv_event_id_counter - - if isinstance(event, BlockStored): - # SGLang emits one BlockStored per page with block_hashes=[single_hash] - # and token_ids containing only that page's tokens. - blocks = [] - for i, bh in enumerate(event.block_hashes): - start = i * event.block_size - end = start + event.block_size - block = common_pb2.KvBlock( - block_hash=bh, - token_ids=event.token_ids[start:end], - block_size=event.block_size, - ) - if event.lora_id is not None: - block.lora_id = event.lora_id - blocks.append(block) - - stored = common_pb2.KvBlocksStored(blocks=blocks) - if event.parent_block_hash is not None: - stored.parent_block_hash = event.parent_block_hash - - return common_pb2.KvCacheEvent(event_id=event_id, stored=stored) - - elif isinstance(event, BlockRemoved): - return common_pb2.KvCacheEvent( - event_id=event_id, - removed=common_pb2.KvBlocksRemoved(block_hashes=event.block_hashes), - ) - - elif isinstance(event, AllBlocksCleared): - return common_pb2.KvCacheEvent(event_id=event_id, cleared=common_pb2.KvCacheCleared()) - - return None - def _handle_epd_disaggregation_encode_request( self, grpc_req: sglang_scheduler_pb2.GenerateRequest, diff --git a/grpc_servicer/smg_grpc_servicer/tokenspeed/kv_events.py b/grpc_servicer/smg_grpc_servicer/tokenspeed/kv_events.py index 0d0dbaca94..3ee0e5b0a2 100644 --- a/grpc_servicer/smg_grpc_servicer/tokenspeed/kv_events.py +++ b/grpc_servicer/smg_grpc_servicer/tokenspeed/kv_events.py @@ -24,10 +24,13 @@ @dataclass(frozen=True) class ResolvedKvEventsConfig: - """The subset of KVEventsConfig the bridge needs to open a SUB socket.""" + """The subset of KVEventsConfig the bridge needs to open a SUB socket, and + the replay ROUTER endpoint (empty when the publisher runs none) the relay + asks for gaps and for the batches published before its subscription joined.""" endpoint: str topic: str + replay_endpoint: str = "" def resolve_kv_events_config(server_args: object) -> ResolvedKvEventsConfig | None: @@ -79,5 +82,16 @@ def resolve_kv_events_config(server_args: object) -> ResolvedKvEventsConfig | No topic, ) return None - logger.info("TokenSpeed KV events enabled: endpoint=%s", endpoint) - return ResolvedKvEventsConfig(endpoint=endpoint, topic=topic) + replay_endpoint = cfg.get("replay_endpoint") or "" + if not isinstance(replay_endpoint, str): + logger.warning( + "TokenSpeed kv-events replay_endpoint must be a string (got %r); replay disabled", + replay_endpoint, + ) + replay_endpoint = "" + logger.info( + "TokenSpeed KV events enabled: endpoint=%s replay_endpoint=%s", + endpoint, + replay_endpoint or "(none)", + ) + return ResolvedKvEventsConfig(endpoint=endpoint, topic=topic, replay_endpoint=replay_endpoint) diff --git a/grpc_servicer/smg_grpc_servicer/tokenspeed/rust.py b/grpc_servicer/smg_grpc_servicer/tokenspeed/rust.py index fc0a9d260f..2bc360cfc2 100644 --- a/grpc_servicer/smg_grpc_servicer/tokenspeed/rust.py +++ b/grpc_servicer/smg_grpc_servicer/tokenspeed/rust.py @@ -163,6 +163,7 @@ def server_facts(server_args: Any) -> dict[str, Any]: "max_running_requests": running_window(server_args), "data_parallel_size": int(dp_size) if isinstance(dp_size, int) and dp_size > 0 else 1, "kv_events_endpoint": kv_events.endpoint if kv_events else "", + "kv_events_replay_endpoint": kv_events.replay_endpoint if kv_events else "", "kv_events_topic": kv_events.topic if kv_events else "", } diff --git a/grpc_servicer/smg_grpc_servicer/vllm/kv_events.py b/grpc_servicer/smg_grpc_servicer/vllm/kv_events.py index 619c65b398..8c62ee6f75 100644 --- a/grpc_servicer/smg_grpc_servicer/vllm/kv_events.py +++ b/grpc_servicer/smg_grpc_servicer/vllm/kv_events.py @@ -1,8 +1,10 @@ -"""vLLM-specific KV-events config resolution. +"""vLLM-specific KV-events config resolution and per-rank publisher discovery. -The wire-format conversion + ZMQ streaming helpers live in the engine-neutral -``smg_grpc_servicer.kv_events`` module and are re-exported here for backwards -compatibility (existing imports and tests reference them via this module). +The relay itself (lenient decoding, normalization, per-rank cursors and +replay) is the engine-neutral ``smg_grpc_servicer.kv_relay``. The older +single-rank helpers in ``smg_grpc_servicer.kv_events`` are re-exported here +for backwards compatibility (existing imports and tests reference them via +this module). """ import logging @@ -14,11 +16,13 @@ stream_kv_events, to_int64, ) +from smg_grpc_servicer.kv_relay import RankSource, rank_sources __all__ = [ "convert_batch", "convert_event", "endpoint_for_rank", + "rank_sources_for", "resolve_kv_events_config", "stream_kv_events", "to_int64", @@ -48,3 +52,44 @@ def resolve_kv_events_config(engine: object): return None logger.info("vLLM KV events enabled: endpoint=%s", getattr(cfg, "endpoint", "?")) return cfg + + +def rank_sources_for(config: object, engine: object) -> list[RankSource]: + """One KV-event source per data-parallel rank. + + vLLM's engine client reports each ready engine core's publisher config + (``get_kv_event_sources``, keyed by DP rank; the endpoints in it are the + rank's own, already offset). Older engines without it get + ``data_parallel_size`` ranks offset from the base endpoint, as the + publishers offset theirs. + """ + reported: dict = {} + getter = getattr(engine, "get_kv_event_sources", None) + if callable(getter): + try: + reported = dict(getter() or {}) + except Exception as error: # noqa: BLE001 - discovery is best effort + logger.warning("get_kv_event_sources failed; deriving ranks from the config: %s", error) + reported = {} + if reported: + sources = [] + for rank, rank_config in sorted(reported.items()): + if not isinstance(rank, int): + continue + endpoint = str(getattr(rank_config, "endpoint", "") or "") + replay = getattr(rank_config, "replay_endpoint", None) + if not endpoint: + continue + sources.append( + RankSource( + rank=rank, + endpoint=endpoint_for_rank(endpoint, 0), + replay_endpoint=endpoint_for_rank(str(replay), 0) if replay else None, + ) + ) + if sources: + return sources + parallel = getattr(getattr(engine, "vllm_config", None), "parallel_config", None) + dp_size = getattr(parallel, "data_parallel_size", None) + dp_size = int(dp_size) if isinstance(dp_size, int) and dp_size > 0 else 1 + return rank_sources(config, range(dp_size)) diff --git a/grpc_servicer/smg_grpc_servicer/vllm/loads.py b/grpc_servicer/smg_grpc_servicer/vllm/loads.py new file mode 100644 index 0000000000..f8f1c03cfc --- /dev/null +++ b/grpc_servicer/smg_grpc_servicer/vllm/loads.py @@ -0,0 +1,167 @@ +"""What the vLLM servicer can tell about its engine's load from the requests it +forwards: queued token-work, generation throughput and the prefix-cache hit +rate, the ``GetLoads`` fields the gateway's expected-wait score reads and +vLLM's own ``SchedulerStats`` do not carry (they report running, waiting and +KV usage). Without them the gateway scores a vLLM worker on defaults and +routes it unlike its peers. + +The servicer sees every request it submits, the engine's first output for it +(with the prompt and cached token counts) and every token streamed, so: + +- queued token-work is the uncached prompt tokens of the requests the engine + has not started: of the submitted requests without a first output, the + youngest ``num_waiting_reqs`` (the engine admits FCFS, so the older ones are + in prefill), each prompt discounted by the recent hit rate since its own + cached count is unknown until it starts; +- generation throughput is the tokens streamed over the last + :data:`THROUGHPUT_WINDOW` seconds; +- the hit rate is cached over prompt tokens across the last + :data:`HIT_RATE_SAMPLES` first outputs. + +These are the field semantics the SGLang servicer reports from its scheduler +and the mock engine from its queue; the Rust servicer keeps the same +bookkeeping (``crates/engine_servicer/src/load_tracker.rs``). Engine-free, so +it is unit-tested without vLLM installed. +""" + +from __future__ import annotations + +import threading +import time +from collections import deque +from typing import NamedTuple + +THROUGHPUT_WINDOW = 2.0 +HIT_RATE_SAMPLES = 64 + + +class LoadEstimate(NamedTuple): + """The three fields, as ``GetLoads`` reports them.""" + + queued_token_work: int = 0 + gen_throughput: float = 0.0 + cache_hit_rate: float = 0.0 + + +class LoadTracker: + """Per-servicer load bookkeeping; thread-safe, every call is cheap.""" + + def __init__(self, clock=time.monotonic) -> None: + self._clock = clock + self._lock = threading.Lock() + # (request_id, prompt_tokens), oldest first: submitted, no first output yet. + self._pending: deque[tuple[str, int]] = deque() + # (prompt tokens, cached tokens) of recent first outputs. + self._hits: deque[tuple[int, int]] = deque() + self._prompt_sum = 0 + self._cached_sum = 0 + # (time, tokens) streamed within the window. + self._generated: deque[tuple[float, int]] = deque() + self._generated_sum = 0 + + def submitted(self, request_id: str, prompt_tokens: int) -> None: + """A request with ``prompt_tokens`` was handed to the engine.""" + with self._lock: + self._pending.append((request_id, max(0, int(prompt_tokens)))) + + def first_output(self, request_id: str, prompt_tokens: int, cached_tokens: int) -> None: + """The engine's first output for a request: it has started, and its + prompt had ``cached_tokens`` of ``prompt_tokens`` in the prefix cache.""" + with self._lock: + self._drop_pending(request_id) + prompt_tokens = max(0, int(prompt_tokens or 0)) + if prompt_tokens == 0: + return + cached_tokens = min(max(0, int(cached_tokens or 0)), prompt_tokens) + self._hits.append((prompt_tokens, cached_tokens)) + self._prompt_sum += prompt_tokens + self._cached_sum += cached_tokens + if len(self._hits) > HIT_RATE_SAMPLES: + prompt, cached = self._hits.popleft() + self._prompt_sum -= prompt + self._cached_sum -= cached + + def generated(self, tokens: int, at: float | None = None) -> None: + """``tokens`` were streamed to the client just now (or ``at``).""" + tokens = int(tokens) + if tokens <= 0: + return + now = self._clock() if at is None else at + with self._lock: + self._generated.append((now, tokens)) + self._generated_sum += tokens + self._trim(now) + + def finished(self, request_id: str) -> None: + """A request ended (or was aborted) without the engine starting it.""" + with self._lock: + self._drop_pending(request_id) + + def pending(self) -> int: + """Submitted requests the engine has not started.""" + with self._lock: + return len(self._pending) + + def estimate(self, num_waiting_reqs: int, now: float | None = None) -> LoadEstimate: + """The estimate now, given the engine's own count of waiting requests.""" + now = self._clock() if now is None else now + with self._lock: + self._trim(now) + hit_rate = self._hit_rate() + waiting = min(max(0, int(num_waiting_reqs or 0)), len(self._pending)) + queued = 0.0 + for _, prompt in list(self._pending)[len(self._pending) - waiting :]: + queued += prompt * (1.0 - hit_rate) + return LoadEstimate( + queued_token_work=int(round(queued)), + gen_throughput=self._generated_sum / THROUGHPUT_WINDOW, + cache_hit_rate=hit_rate, + ) + + def _drop_pending(self, request_id: str) -> None: + for index, (pending_id, _) in enumerate(self._pending): + if pending_id == request_id: + del self._pending[index] + return + + def _trim(self, now: float) -> None: + while self._generated and now - self._generated[0][0] > THROUGHPUT_WINDOW: + _, tokens = self._generated.popleft() + self._generated_sum -= tokens + + def _hit_rate(self) -> float: + if self._prompt_sum <= 0: + return 0.0 + return min(1.0, max(0.0, self._cached_sum / self._prompt_sum)) + + +def scheduler_load_fields( + num_running: int, + num_waiting: int, + kv_usage: float, + estimate: LoadEstimate, + *, + max_total_num_tokens: int = 0, + max_running_requests: int = 0, +) -> dict: + """The ``SchedulerLoad`` fields for one rank: vLLM's own counts plus the + tracker's estimate, with ``num_used_tokens`` and ``utilization`` derived + from the KV usage as the SGLang servicer derives them.""" + kv_usage = max(0.0, float(kv_usage or 0.0)) + fields = { + "dp_rank": 0, + "num_running_reqs": int(num_running), + "num_waiting_reqs": int(num_waiting), + "num_waiting_uncached_tokens": int(estimate.queued_token_work), + "num_total_reqs": int(num_running) + int(num_waiting), + "token_usage": kv_usage, + "utilization": kv_usage, + "gen_throughput": float(estimate.gen_throughput), + "cache_hit_rate": float(estimate.cache_hit_rate), + } + if max_total_num_tokens > 0: + fields["max_total_num_tokens"] = int(max_total_num_tokens) + fields["num_used_tokens"] = int(round(kv_usage * max_total_num_tokens)) + if max_running_requests > 0: + fields["max_running_requests"] = int(max_running_requests) + return fields diff --git a/grpc_servicer/smg_grpc_servicer/vllm/servicer.py b/grpc_servicer/smg_grpc_servicer/vllm/servicer.py index 98237e7183..38e46cf74e 100755 --- a/grpc_servicer/smg_grpc_servicer/vllm/servicer.py +++ b/grpc_servicer/smg_grpc_servicer/vllm/servicer.py @@ -14,15 +14,11 @@ from pathlib import Path import grpc -import msgspec import torch -import zmq -import zmq.asyncio from smg_grpc_proto import vllm_engine_pb2, vllm_engine_pb2_grpc from smg_grpc_proto.generated import common_pb2 from transformers import BatchFeature from vllm import PoolingParams, SamplingParams, TokensPrompt -from vllm.distributed.kv_events import KVEventBatch from vllm.engine.protocol import EngineClient from vllm.inputs.engine import MultiModalInput as VllmMultiModalInput from vllm.inputs.engine import mm_input, tokens_input @@ -36,14 +32,14 @@ from vllm.outputs import CompletionOutput, RequestOutput from vllm.sampling_params import RequestOutputKind, StructuredOutputsParams +from smg_grpc_servicer.kv_relay import Engine, relay from smg_grpc_servicer.tokenizer_bundle import CHUNK_SIZE, build_tokenizer_zip from smg_grpc_servicer.vllm import attach_vllm_logging from smg_grpc_servicer.vllm.admin import flush_cache from smg_grpc_servicer.vllm.errors import grpc_code_for from smg_grpc_servicer.vllm.kv_events import ( - endpoint_for_rank, + rank_sources_for, resolve_kv_events_config, - stream_kv_events, ) from smg_grpc_servicer.vllm.kv_transfer import ( params_from_request, @@ -53,6 +49,7 @@ # The launcher imports this module before it defines serve_grpc: the moment # the servicer switch has to be in place (see launcher_switch). from smg_grpc_servicer.vllm.launcher_switch import install_launcher_switch +from smg_grpc_servicer.vllm.loads import LoadTracker, scheduler_load_fields from smg_grpc_servicer.vllm.media_identity import build_media_identity, media_identity_supported from smg_grpc_servicer.vllm.media_refs import parse_media_refs, validate_schemes from smg_grpc_servicer.vllm.mm_processor import ( @@ -94,6 +91,31 @@ VLLM_VERSION = "" +def _prompt_length(prompt) -> int: + """The token count of an engine prompt (0 when it is text the engine tokenizes).""" + if isinstance(prompt, dict): + ids = prompt.get("prompt_token_ids") + return len(ids) if ids is not None else 0 + return 0 + + +def _kv_capacity_tokens(engine) -> int: + """The KV cache's token capacity, when the engine config exposes it.""" + cache = getattr(getattr(engine, "vllm_config", None), "cache_config", None) + blocks = getattr(cache, "num_gpu_blocks", None) + block_size = getattr(cache, "block_size", None) + if isinstance(blocks, int) and isinstance(block_size, int) and blocks > 0 and block_size > 0: + return blocks * block_size + return 0 + + +def _max_running_requests(engine) -> int: + """The scheduler's running window (``max_num_seqs``), when exposed.""" + scheduler = getattr(getattr(engine, "vllm_config", None), "scheduler_config", None) + window = getattr(scheduler, "max_num_seqs", None) + return window if isinstance(window, int) and window > 0 else 0 + + def _latest_scheduler_stats(engine, engine_idx: int = 0): """Best-effort read of the most recent ``SchedulerStats`` snapshot. @@ -168,6 +190,9 @@ def __init__( # Resolve KV-event publishing config from the engine. Non-None only when # vLLM was started with --kv-events-config enabling the ZMQ publisher. self._kv_events_config = resolve_kv_events_config(async_llm) + # Queued token-work, generation throughput and hit rate for GetLoads, + # from the requests this servicer forwards (vLLM's stats carry none). + self._loads = LoadTracker() # Flag > env > default, resolved once so each value names its source. self._mm_settings = (mm_settings or MmSettings()).resolve() # Worker-side media processing (media_refs); None keeps refs rejected. @@ -403,6 +428,7 @@ async def Generate( # Track which indices have sent their first chunk seen_indices: set[int] = set() + self._loads.submitted(request_id, _prompt_length(prompt)) async for output in self.engine.generate( prompt=prompt, sampling_params=sampling_params, @@ -412,7 +438,14 @@ async def Generate( request.data_parallel_rank if request.HasField("data_parallel_rank") else None ), ): + if not engine_started: + self._loads.first_output( + request_id, + len(output.prompt_token_ids or ()), + getattr(output, "num_cached_tokens", 0) or 0, + ) engine_started = True + self._loads.generated(sum(len(c.token_ids) for c in output.outputs)) # For streaming, send chunks for EACH completion output (n outputs) if request.stream: for completion in output.outputs: @@ -471,6 +504,8 @@ async def Generate( logger.warning("Generate request %s rejected (%s): %s", request_id, code.name, e) await self._notify_kv_transfer_rejected(request_id, kv_transfer_params, engine_started) await context.abort(code, str(e)) + finally: + self._loads.finished(request_id) async def _notify_kv_transfer_rejected( self, @@ -701,7 +736,10 @@ async def GetLoads( Reads the latest SchedulerStats snapshot cached on the engine's stat loggers and maps it onto a single-DP-rank SchedulerLoad: ``token_usage`` carries KV-cache utilization ([0,1)) and ``num_running_reqs`` / - ``num_waiting_reqs`` report queue depth. + ``num_waiting_reqs`` report queue depth. ``num_waiting_uncached_tokens``, + ``gen_throughput`` and ``cache_hit_rate`` come from this servicer's own + bookkeeping of the requests it forwards (``smg_grpc_servicer.vllm.loads``), + since vLLM's stats do not carry them. Always returns exactly one SchedulerLoad entry (zero-filled when no snapshot is available yet, e.g. with --disable-log-stats or before the @@ -731,11 +769,14 @@ async def GetLoads( kv_usage = 0.0 load = vllm_engine_pb2.SchedulerLoad( - dp_rank=0, - num_running_reqs=num_running, - num_waiting_reqs=num_waiting, - num_total_reqs=num_running + num_waiting, - token_usage=max(0.0, kv_usage), + **scheduler_load_fields( + num_running, + num_waiting, + kv_usage, + self._loads.estimate(num_waiting), + max_total_num_tokens=_kv_capacity_tokens(self.engine), + max_running_requests=_max_running_requests(self.engine), + ) ) return vllm_engine_pb2.GetLoadsResponse( @@ -1274,11 +1315,8 @@ async def SubscribeKvEvents( request: common_pb2.SubscribeKvEventsRequest, context: grpc.aio.ServicerContext, ) -> AsyncIterator[common_pb2.KvEventBatch]: - """Bridge vLLM's in-process ZMQ KV cache events to a gRPC stream. - - The ZMQ publisher's sequence numbers are used directly as the gRPC - batch sequence numbers. - """ + """Relay vLLM's ZMQ KV cache events, every DP rank's publisher, as one + gRPC stream (see ``smg_grpc_servicer.kv_relay``).""" if self._kv_events_config is None: await context.abort( grpc.StatusCode.UNIMPLEMENTED, @@ -1288,34 +1326,12 @@ async def SubscribeKvEvents( ) config = self._kv_events_config - - # For DP attention each rank publishes on port + rank with independent - # sequence counters; subscribing to several on one socket interleaves - # them and breaks gap detection. Subscribe to rank 0 only for now. - # TODO(phase3): per-rank virtual workers or merged renumbering. - pub_endpoint = endpoint_for_rank(config.endpoint, 0) - - zmq_ctx = zmq.asyncio.Context.instance() - sub_socket = zmq_ctx.socket(zmq.SUB) - sub_socket.subscribe(config.topic.encode("utf-8")) - sub_socket.connect(pub_endpoint) - logger.info("SubscribeKvEvents: connected to ZMQ endpoint %s", pub_endpoint) - - decoder = msgspec.msgpack.Decoder(KVEventBatch) - - try: - async for proto_batch in stream_kv_events( - sub_socket, - decoder.decode, - lambda: context.send_initial_metadata(()), - context.cancelled, - ): - yield proto_batch - except asyncio.CancelledError: - pass - except Exception as e: - logger.exception("SubscribeKvEvents failed") - await context.abort(grpc.StatusCode.INTERNAL, str(e)) - finally: - sub_socket.close(linger=0) - logger.info("SubscribeKvEvents: stream closed") + async for proto_batch in relay( + rank_sources_for(config, self.engine), + Engine.VLLM, + request.start_sequence_number, + context, + topic=str(getattr(config, "topic", "") or ""), + hwm=getattr(config, "hwm", None), + ): + yield proto_batch diff --git a/grpc_servicer/tests/conftest.py b/grpc_servicer/tests/conftest.py index d11ac1c6b4..6df77bbe50 100644 --- a/grpc_servicer/tests/conftest.py +++ b/grpc_servicer/tests/conftest.py @@ -1,11 +1,20 @@ """Put this repo's ``grpc_servicer/`` at the front of ``sys.path`` so ``smg_grpc_servicer`` imports resolve in-repo when tests run from the repo root (as CI does), without an editable install. + +``SMG_GRPC_PROTO_PATH`` names a directory holding locally generated +``smg_grpc_proto`` stubs (``scripts/gen_proto_stubs.py``), so the tests see +this checkout's proto instead of the released package. """ +import os import sys from pathlib import Path _GRPC_SERVICER_ROOT = Path(__file__).resolve().parent.parent if sys.path[:1] != [str(_GRPC_SERVICER_ROOT)]: sys.path.insert(0, str(_GRPC_SERVICER_ROOT)) + +_PROTO_PATH = os.environ.get("SMG_GRPC_PROTO_PATH") +if _PROTO_PATH and _PROTO_PATH not in sys.path: + sys.path.insert(0, _PROTO_PATH) diff --git a/grpc_servicer/tests/test_kv_relay.py b/grpc_servicer/tests/test_kv_relay.py new file mode 100644 index 0000000000..36535c2e62 --- /dev/null +++ b/grpc_servicer/tests/test_kv_relay.py @@ -0,0 +1,1265 @@ +"""The KV-event relay: parity with the Rust normalizer on the engines' wire +shapes, lenient decoding, and multi-rank streaming with per-rank replay over +real ZMQ and gRPC (no engine needed).""" + +from __future__ import annotations + +import asyncio +import hashlib +from types import SimpleNamespace + +import pytest +import pytest_asyncio + +pytest.importorskip("smg_grpc_proto") +grpc = pytest.importorskip("grpc") +zmq = pytest.importorskip("zmq") +msgspec = pytest.importorskip("msgspec") +import zmq.asyncio # noqa: E402, F811 +from smg_grpc_proto.generated import common_pb2 # noqa: E402 +from smg_grpc_servicer import kv_relay # noqa: E402 + +TIERS = { + "device": common_pb2.KV_CACHE_TIER_DEVICE, + "host": common_pb2.KV_CACHE_TIER_HOST, + "disk": common_pb2.KV_CACHE_TIER_DISK, + "external": common_pb2.KV_CACHE_TIER_EXTERNAL, +} +U64 = (1 << 64) - 1 + + +def _optional(message, name): + return getattr(message, name) if message.HasField(name) else None + + +def _key_shape(key): + which = key.WhichOneof("key") + if which == "blob": + return {"blob_len": len(key.blob)} + if which == "multimodal": + return {"multimodal": [key.multimodal.identifier, key.multimodal.offset]} + return {which: getattr(key, which)} + + +# --------------------------------------------------------------------------- +# Parity with the Rust relay +# +# The scenarios of ``crates/engine_servicer/src/kv_wire_shapes.rs``, encoded +# here with msgspec through mirrors of the engines' own structs so the bytes +# are what the publishers put on the wire: +# +# - vLLM ``vllm/distributed/kv_events.py``: ``EventBatch`` is ``array_like`` +# (``[ts, events, data_parallel_rank]``), events are tagged maps (``type``) +# with ``omit_defaults``; required fields are present even when nil. +# - SGLang ``python/sglang/srt/disaggregation/kv_events.py``: the same shape, +# ``attn_dp_rank`` in the batch's third slot (nil when unset, SGLang's batch +# does not omit defaults), ``medium``, ``cache_salt`` and ``session_id`` +# omitted when None, bigram pages as ``[t, t+1]`` cells. +# - The legacy layout both engines used before the map encoding: the same +# structs with ``array_like=True``, tag first, every field in declaration +# order and nil when unset (``omit_defaults`` does not thin an array). +# +# Each scenario is a sequence of batches (one ZMQ message each) with the +# normalizer's expected output: the forwarded events in order and the +# counters, keyed by the relay's ``DropReason`` names. The expectations are +# written by hand next to the events, as the rules in +# ``crates/engine_servicer/src/kv_wire.rs`` say; the Rust test asserts the +# same ones, which keeps the two relays in step. +# --------------------------------------------------------------------------- + + +def _vllm_structs(array_like): + """vLLM's event structs; ``array_like`` selects the legacy layout.""" + + class EventBatch(msgspec.Struct, array_like=True, omit_defaults=True, gc=False): + ts: float + events: list + data_parallel_rank: int | None = None + + class KVCacheEvent( + msgspec.Struct, array_like=array_like, omit_defaults=True, gc=False, tag=True + ): + pass + + class BlockStored(KVCacheEvent): + block_hashes: list + parent_block_hash: int | bytes | None + token_ids: list + block_size: int + lora_id: int | None + medium: str | None + lora_name: str | None + extra_keys: list | None = None + group_idx: int | None = None + kv_cache_spec_kind: str | None = None + kv_cache_spec_sliding_window: int | None = None + locality: str | None = None + ownership: str | None = None + session_id: str | None = None + + class BlockRemoved(KVCacheEvent): + block_hashes: list + medium: str | None + group_idx: int | None = None + locality: str | None = None + ownership: str | None = None + + class AllBlocksCleared(KVCacheEvent): + pass + + class BlockMigrated(KVCacheEvent): + """An event type the relay does not know (stands in for a future one).""" + + block_hashes: list + destination: str + + return EventBatch, BlockStored, BlockRemoved, AllBlocksCleared, BlockMigrated + + +def _sglang_structs(array_like): + """SGLang's event structs; ``array_like`` selects the legacy layout.""" + + class EventBatch(msgspec.Struct, array_like=True, gc=False): + ts: float + events: list + attn_dp_rank: int | None = None + + class KVCacheEvent( + msgspec.Struct, array_like=array_like, omit_defaults=True, gc=False, tag=True + ): + pass + + class BlockStored(KVCacheEvent): + block_hashes: list + parent_block_hash: int | None + token_ids: list # ints, or [t, t+1] pairs under bigram hashing + block_size: int + lora_id: int | None + medium: str | None = None + cache_salt: str | None = None + session_id: str | None = None + + class BlockRemoved(KVCacheEvent): + block_hashes: list + medium: str | None = None + + class AllBlocksCleared(KVCacheEvent): + pass + + class BlockMigrated(KVCacheEvent): + block_hashes: list + destination: str + + return EventBatch, BlockStored, BlockRemoved, AllBlocksCleared, BlockMigrated + + +def _digest(label): + return hashlib.sha256(label.encode()).digest() + + +def _as_i64(value): + value &= U64 + return value - (1 << 64) if value >= 1 << 63 else value + + +def _vllm_int(label): + """vLLM's integer form: the low 64 bits of the digest, unsigned.""" + return int.from_bytes(_digest(label), "big") & U64 + + +def _vllm_expected(label): + """What the relay forwards for either form of a vLLM hash.""" + return _as_i64(_vllm_int(label)) + + +def _sglang_int(label): + """SGLang's integer form: the high 64 bits of the digest, signed.""" + return int.from_bytes(_digest(label)[:8], "big", signed=True) + + +def _stored_expect( + hashes, + tokens, + *, + dp_rank=0, + parent=None, + tier="device", + cache_level=None, + lora_name=None, + cache_salt=None, + group_idx=None, + session_id=None, + extra_keys=None, +): + return { + "kind": "stored", + "dp_rank": dp_rank, + "hashes": hashes, + "parent": parent, + "tier": tier, + "cache_level": cache_level, + "tokens": tokens, + "lora_name": lora_name, + "cache_salt": cache_salt, + "group_idx": group_idx, + "session_id": session_id, + "extra_keys": extra_keys, + } + + +def _removed_expect(hashes, *, dp_rank=0, tier="device", cache_level=None): + return { + "kind": "removed", + "dp_rank": dp_rank, + "hashes": hashes, + "tier": tier, + "cache_level": cache_level, + } + + +def _cleared_expect(*, dp_rank=0): + return {"kind": "cleared", "dp_rank": dp_rank} + + +def _vllm_scenario(array_like): + """Both hash forms, a sliding-window group in both of its shapes, a second + physical copy with per-copy removals, the offload tiers and every drop + rule, a pool reset, a LoRA request with a multimodal item, a cache salt + and a prompt embeddings digest whose child inherits the salt, and a + second DP rank.""" + EventBatch, Stored, Removed, Cleared, Migrated = _vllm_structs(array_like) + bs = 4 + h = {name: _vllm_int(f"vllm-{name}") for name in "abcdefghijk"} + e = {name: _vllm_expected(f"vllm-{name}") for name in "abcdefghijk"} + d1, d2 = _digest("vllm-digest-1"), _digest("vllm-digest-2") + embeds = _digest("vllm-prompt-embeds") + # The low 64 bits of a digest can exceed i64::MAX; make sure one does. + assert any(v >= 1 << 63 for v in h.values()), "pick labels with a high bit set" + + def gpu_stored(hashes, parent, tokens, **kw): + kw.setdefault("medium", "GPU") + kw.setdefault("group_idx", 0) + kw.setdefault("kv_cache_spec_kind", "full_attention") + return Stored( + block_hashes=hashes, + parent_block_hash=parent, + token_ids=tokens, + block_size=kw.pop("block_size", bs), + lora_id=kw.pop("lora_id", None), + lora_name=kw.pop("lora_name", None), + **kw, + ) + + batches = [] + forwarded = [] + + # Batch 0: a plain chain in both hash forms; the sliding-window group in + # both of its shapes. + batches.append( + EventBatch( + ts=1700000000.0, + data_parallel_rank=0, + events=[ + gpu_stored([h["a"], h["b"]], None, list(range(1, 9)), session_id="req-1"), + # Sliding-window group: more tokens than hashes x block size, + # no hashes at all. Dropped by the group gate. + gpu_stored( + [], + None, + list(range(1, 9)), + group_idx=1, + kv_cache_spec_kind="sliding_window", + kv_cache_spec_sliding_window=128, + ), + # The same group's usual shape: token_ids span the whole + # computed range and block_hashes name only the window's last + # blocks. Dropped whole by the group gate, never sliced from + # the head. + gpu_stored( + [h["k"]], + None, + list(range(1, 25)), + group_idx=1, + kv_cache_spec_kind="sliding_window", + kv_cache_spec_sliding_window=128, + ), + # Raw digests (VLLM_KV_EVENTS_USE_INT_BLOCK_HASHES=0), the + # parent given as an int, no extra keys on either block. + gpu_stored([d1, d2], h["b"], list(range(9, 17)), extra_keys=[None, None]), + ], + ) + ) + forwarded += [ + _stored_expect( + [e["a"], e["b"]], + [[1, 2, 3, 4], [5, 6, 7, 8]], + group_idx=0, + session_id="req-1", + ), + _stored_expect( + [_as_i64(int.from_bytes(d1[-8:], "big")), _as_i64(int.from_bytes(d2[-8:], "big"))], + [[9, 10, 11, 12], [13, 14, 15, 16]], + parent=e["b"], + group_idx=0, + ), + ] + + # Batch 1: a second physical copy, per-copy removals, offload tiers and + # every other drop rule. + batches.append( + EventBatch( + ts=1700000001.0, + data_parallel_rank=0, + events=[ + gpu_stored([h["a"], h["b"]], None, list(range(1, 9))), # duplicate copy + Removed(block_hashes=[h["a"]], medium="GPU", group_idx=0), + Removed(block_hashes=[h["a"]], medium="GPU", group_idx=0), # other copy + # CPU offload placeholder: a chunk key, no tokens, block_size 0. + gpu_stored([h["c"]], None, [], medium="CPU", block_size=0, kv_cache_spec_kind=None), + Removed(block_hashes=[h["c"]], medium="CPU", group_idx=0), + gpu_stored([h["d"]], None, [1, 2, 3, 4], medium="STORAGE", locality="REMOTE"), + gpu_stored( + [h["d"]], + None, + [1, 2, 3, 4], + medium="STORAGE", + locality="LOCAL", + ownership="kvcr", + ), + gpu_stored([h["d"]], None, [1, 2, 3, 4], medium="STORAGE", locality="LOCAL"), + gpu_stored([h["f"]], None, [1, 2, 3, 4], medium="MARS"), + gpu_stored([h["g"]], None, [1, 2, 3, 4, 5, 6]), # unaligned + gpu_stored([h["i"]], h["i"], [1, 2, 3, 4]), # parent is itself + Migrated(block_hashes=[h["j"]], destination="peer"), + # Malformed: hashes are not a list. Written out by hand + # because no struct produces it. + ( + ["BlockStored", "nope", None, [1, 2, 3, 4], bs] + if array_like + else { + "type": "BlockStored", + "block_hashes": "nope", + "parent_block_hash": None, + "token_ids": [1, 2, 3, 4], + "block_size": bs, + } + ), + ], + ) + ) + forwarded += [ + _stored_expect([e["a"], e["b"]], [[1, 2, 3, 4], [5, 6, 7, 8]], group_idx=0), + _removed_expect([e["a"]]), + _removed_expect([e["a"]]), + _removed_expect([e["c"]], tier="host", cache_level=1), + _stored_expect([e["d"]], [[1, 2, 3, 4]], tier="disk", cache_level=2, group_idx=0), + ] + + # Batch 2: the pool reset, then the chain again (not a duplicate any more). + batches.append( + EventBatch( + ts=1700000002.0, + data_parallel_rank=0, + events=[ + Cleared(), + gpu_stored([h["a"], h["b"]], None, list(range(1, 9))), + ], + ) + ) + forwarded += [ + _cleared_expect(), + _stored_expect([e["a"], e["b"]], [[1, 2, 3, 4], [5, 6, 7, 8]], group_idx=0), + ] + + # Batch 3: a LoRA request with a multimodal item, a cache salt and prompt + # embeddings; the salt rides in block 0's extra keys only and the child + # inherits it. + batches.append( + EventBatch( + ts=1700000003.0, + data_parallel_rank=0, + events=[ + gpu_stored( + [h["e"]], + None, + [1, 2, 3, 4], + lora_id=7, + lora_name="adapter", + extra_keys=[("adapter", ("mm-abc", 0), "salt-1", embeds)], + session_id="req-2", + ), + gpu_stored( + [h["h"]], + h["e"], + [5, 6, 7, 8], + lora_id=7, + lora_name="adapter", + extra_keys=[("adapter",)], + session_id="req-2", + ), + ], + ) + ) + forwarded += [ + _stored_expect( + [e["e"]], + [[1, 2, 3, 4]], + lora_name="adapter", + cache_salt="salt-1", + group_idx=0, + session_id="req-2", + extra_keys=[ + [ + {"text": "adapter"}, + {"multimodal": ["mm-abc", 0]}, + {"text": "salt-1"}, + {"blob_len": 32}, + ] + ], + ), + _stored_expect( + [e["h"]], + [[5, 6, 7, 8]], + parent=e["e"], + lora_name="adapter", + cache_salt="salt-1", + group_idx=0, + session_id="req-2", + extra_keys=[[{"text": "adapter"}]], + ), + ] + + # Batch 4: another DP rank stores the same hashes; seen-sets are per rank. + batches.append( + EventBatch( + ts=1700000004.0, + data_parallel_rank=1, + events=[gpu_stored([h["a"], h["b"]], None, list(range(1, 9)))], + ) + ) + forwarded += [ + _stored_expect([e["a"], e["b"]], [[1, 2, 3, 4], [5, 6, 7, 8]], dp_rank=1, group_idx=0), + ] + + counts = { + "forwarded_stored": 8, + "forwarded_removed": 3, + "forwarded_cleared": 1, + "duplicate_stores": 1, + "bigram_stores": 0, + "dropped": { + "non_main_attention_group": 2, + "placeholder": 1, + "non_local_locality": 1, + "unsupported_ownership": 1, + "unknown_medium": 1, + "unaligned_blocks": 1, + "self_referencing_hashes": 1, + "unknown_type": 1, + "malformed": 1, + }, + } + return batches, {"forwarded": forwarded, "counts": counts} + + +def _sglang_scenario(array_like): + """The startup clear, a chain with a coalesced two-page store, HiCache + write-through, a salted chain, an Eagle bigram page, the DISK and + EXTERNAL media, an unknown event type and a second attention DP rank.""" + EventBatch, Stored, Removed, Cleared, Migrated = _sglang_structs(array_like) + bs = 4 + s = {name: _sglang_int(f"sglang-{name}") for name in "abcdefgh"} + assert any(v < 0 for v in s.values()), "pick labels with a negative i64" + + def stored(hashes, parent, tokens, **kw): + kw.setdefault("medium", "GPU") + return Stored( + block_hashes=hashes, + parent_block_hash=parent, + token_ids=tokens, + block_size=bs, + lora_id=None, + **kw, + ) + + batches = [] + forwarded = [] + + # Batch 0: the first batch after startup clears. + batches.append(EventBatch(ts=1700000000.0, attn_dp_rank=0, events=[Cleared()])) + forwarded += [_cleared_expect()] + + # Batch 1: a chain; the second store is coalesced over two pages. + batches.append( + EventBatch( + ts=1700000001.0, + attn_dp_rank=0, + events=[ + stored([s["a"]], None, [1, 2, 3, 4], session_id="req-1"), + stored([s["b"], s["c"]], s["a"], list(range(5, 13)), session_id="req-1"), + ], + ) + ) + forwarded += [ + _stored_expect([s["a"]], [[1, 2, 3, 4]], session_id="req-1"), + _stored_expect( + [s["b"], s["c"]], [[5, 6, 7, 8], [9, 10, 11, 12]], parent=s["a"], session_id="req-1" + ), + ] + + # Batch 2: HiCache write-through: back up to host, demote (device copy + # goes, host stays), load back, evict the host copy. + batches.append( + EventBatch( + ts=1700000002.0, + attn_dp_rank=0, + events=[ + stored([s["a"]], None, [1, 2, 3, 4], medium="CPU_PINNED"), + Removed(block_hashes=[s["a"]], medium="GPU"), + stored([s["a"]], None, [1, 2, 3, 4]), + Removed(block_hashes=[s["a"]], medium="CPU_PINNED"), + ], + ) + ) + forwarded += [ + _stored_expect([s["a"]], [[1, 2, 3, 4]], tier="host", cache_level=1), + _removed_expect([s["a"]]), + _stored_expect([s["a"]], [[1, 2, 3, 4]]), + _removed_expect([s["a"]], tier="host", cache_level=1), + ] + + # Batch 3: a salted request's chain (the legacy array layout has no + # readable salt slot, so that variant stores the chain unsalted). + salt = None if array_like else "tenant-a" + batches.append( + EventBatch( + ts=1700000003.0, + attn_dp_rank=0, + events=[ + stored([s["d"]], None, [1, 2, 3, 4], cache_salt=salt), + stored([s["e"]], s["d"], [5, 6, 7, 8], cache_salt=salt), + ], + ) + ) + forwarded += [ + _stored_expect([s["d"]], [[1, 2, 3, 4]], cache_salt=salt), + _stored_expect([s["e"]], [[5, 6, 7, 8]], parent=s["d"], cache_salt=salt), + ] + + # Batch 4: an Eagle bigram page, removed again in the same batch. (Its + # tokens differ from the plain chain's: two engine hashes with the same + # tokens at the same position share one index membership per worker.) + batches.append( + EventBatch( + ts=1700000004.0, + attn_dp_rank=0, + events=[ + stored([s["f"]], None, [[21, 22], [22, 23], [23, 24], [24, 25]]), + Removed(block_hashes=[s["f"]], medium="GPU"), + ], + ) + ) + forwarded += [ + _stored_expect([s["f"]], [[21, 22, 23, 24]]), + _removed_expect([s["f"]]), + ] + + # Batch 5: the tiers the default core never emits but defines; an event + # type the relay does not know. + batches.append( + EventBatch( + ts=1700000005.0, + attn_dp_rank=0, + events=[ + stored([s["g"]], None, [1, 2, 3, 4], medium="DISK"), + stored([s["h"]], None, [1, 2, 3, 4], medium="EXTERNAL"), + Migrated(block_hashes=[s["h"]], destination="peer"), + ], + ) + ) + forwarded += [ + _stored_expect([s["g"]], [[1, 2, 3, 4]], tier="disk", cache_level=2), + _stored_expect([s["h"]], [[1, 2, 3, 4]], tier="external", cache_level=3), + ] + + # Batch 6: another attention DP rank; its batch carries its rank. + batches.append( + EventBatch( + ts=1700000006.0, + attn_dp_rank=1, + events=[stored([s["a"]], None, [1, 2, 3, 4])], + ) + ) + forwarded += [_stored_expect([s["a"]], [[1, 2, 3, 4]], dp_rank=1)] + + if array_like: + # SGLang's legacy arrays put session_id where vLLM has extra_keys; the + # relay cannot read it there. + for item in forwarded: + if "session_id" in item: + item["session_id"] = None + + counts = { + "forwarded_stored": 10, + "forwarded_removed": 3, + "forwarded_cleared": 1, + "duplicate_stores": 0, + "bigram_stores": 1, + "dropped": {"unknown_type": 1}, + } + return batches, {"forwarded": forwarded, "counts": counts} + + +SCENARIOS = {"vllm": _vllm_scenario, "sglang": _sglang_scenario} + + +@pytest.mark.parametrize( + "engine,layout", + [("vllm", "map"), ("vllm", "array"), ("sglang", "map"), ("sglang", "array")], +) +def test_the_engines_wire_shapes_normalize_like_the_rust_relay(engine, layout): + name = f"{engine}-{layout}" + batches, expect = SCENARIOS[engine](layout == "array") + encoder = msgspec.msgpack.Encoder() + normalizer = kv_relay.Normalizer() + normalized = [ + normalizer.normalize_batch(kv_relay.decode_batch(encoder.encode(batch)), seq) + for seq, batch in enumerate(batches) + ] + forwarded = [ + (_optional(batch, "dp_rank"), event) for batch in normalized for event in batch.events + ] + assert len(forwarded) == len(expect["forwarded"]), name + for index, ((rank, event), want) in enumerate(zip(forwarded, expect["forwarded"])): + at = f"{name} forwarded event {index}" + assert rank == want["dp_rank"], at + kind = event.WhichOneof("data") + assert kind == want["kind"], at + if kind == "stored": + stored = event.stored + assert [b.block_hash for b in stored.blocks] == want["hashes"], at + assert _optional(stored, "parent_block_hash") == want["parent"], at + assert stored.tier == TIERS[want["tier"]], at + assert [list(b.token_ids) for b in stored.blocks] == want["tokens"], at + for block in stored.blocks: + assert _optional(block, "cache_level") == want["cache_level"], at + assert block.block_size == len(block.token_ids), at + assert _optional(stored, "lora_name") == want["lora_name"], at + assert _optional(stored, "cache_salt") == want["cache_salt"], at + assert _optional(stored, "group_idx") == want["group_idx"], at + assert _optional(stored, "session_id") == want.get("session_id"), at + if want.get("extra_keys") is not None: + got = [[_key_shape(key) for key in block.extra_keys] for block in stored.blocks] + assert got == want["extra_keys"], at + elif kind == "removed": + removed = event.removed + assert list(removed.block_hashes) == want["hashes"], at + assert removed.tier == TIERS[want["tier"]], at + assert _optional(removed, "cache_level") == want["cache_level"], at + counts = normalizer.counts + want = expect["counts"] + assert counts.forwarded_stored == want["forwarded_stored"], name + assert counts.forwarded_removed == want["forwarded_removed"], name + assert counts.forwarded_cleared == want["forwarded_cleared"], name + assert counts.duplicate_stores == want["duplicate_stores"], name + assert counts.bigram_stores == want["bigram_stores"], name + assert counts.dropped == want["dropped"], name + + +# --------------------------------------------------------------------------- +# Lenient decoding +# --------------------------------------------------------------------------- + + +def _store(hashes, tokens, **extra): + event = { + "type": "BlockStored", + "block_hashes": hashes, + "parent_block_hash": None, + "token_ids": tokens, + "block_size": 4, + "lora_id": None, + "medium": "GPU", + "lora_name": None, + } + event.update(extra) + return event + + +def _batch(events, rank=0, ts=1.5): + return msgspec.msgpack.encode([ts, events, rank]) + + +def _normalize(payloads, rank=None): + normalizer = kv_relay.Normalizer() + batches = [ + normalizer.normalize_batch(kv_relay.decode_batch(payload), seq + 1, rank) + for seq, payload in enumerate(payloads) + ] + return batches, normalizer.counts + + +def test_hashes_fold_like_the_engines_send_them(): + assert kv_relay.fold_hash(7) == 7 + assert kv_relay.fold_hash(2**63) == -(2**63) + assert kv_relay.fold_hash(-3) == -3 + digest = bytes(range(32)) + assert kv_relay.fold_hash(digest) == int.from_bytes(digest[-8:], "big", signed=True) + assert kv_relay.fold_hash("nope") is None + + +def test_unknown_keys_and_bigram_cells_decode_and_a_bad_event_costs_itself(): + payload = _batch( + [ + _store([1], [[1, 2], [2, 3], [3, 4], [4, 5]], future_field={"nested": True}), + {"type": "BlockStored", "block_hashes": "nope", "token_ids": [1], "block_size": 1}, + {"type": "BlockMigrated", "block_hashes": [1]}, + ["BlockRemoved", [1], "GPU"], + ] + ) + batches, counts = _normalize([payload]) + events = batches[0].events + assert [event.WhichOneof("data") for event in events] == ["stored", "removed"] + assert list(events[0].stored.blocks[0].token_ids) == [1, 2, 3, 4] + assert events[1].event_id == 4, "ids advance for dropped events too" + assert counts.bigram_stores == 1 + assert counts.dropped == {"malformed": 1, "unknown_type": 1} + + +def test_a_parent_record_outlives_all_but_its_last_copy(): + """vLLM keeps up to two physical copies of a hash and removes them one at + a time. The record a child inherits its namespace from, and the digest the + hash check chains on, stay until the last copy is removed and survive a + repeat store, as the Rust normalizer keeps them; after the last removal a + child is an orphan: no inherited salt, unverifiable.""" + seed = kv_relay.sglang_salt_seed("tenant-a") + root = kv_relay.sglang_chain([1, 2, 3, 4], 4, seed)[0][1] + child = kv_relay.sglang_chain([1, 2, 3, 4, 5, 6, 7, 8], 4, seed)[1][1] + orphan = kv_relay.sglang_chain([1, 2, 3, 4, 13, 14, 15, 16], 4, seed)[1][1] + normalizer = kv_relay.Normalizer(hash_check="sglang") + + def normalize(events, seq): + return normalizer.normalize_batch(kv_relay.decode_batch(_batch(events)), seq) + + first = normalize( + [ + _store([root], [1, 2, 3, 4], cache_salt="tenant-a"), + _store([root], [1, 2, 3, 4], cache_salt="tenant-a"), # the second physical copy + {"type": "BlockRemoved", "block_hashes": [root], "medium": "GPU"}, # one copy goes + _store([child], [5, 6, 7, 8], parent_block_hash=root), # no salt of its own + ], + 1, + ) + stored = [event.stored for event in first.events if event.WhichOneof("data") == "stored"] + assert stored[2].cache_salt == "tenant-a", "inherited from the copy still cached" + counts = normalizer.counts + assert ( + counts.hash_checked, + counts.hash_mismatch, + counts.hash_unverifiable, + counts.duplicate_stores, + ) == (3, 0, 0, 1) + second = normalize( + [ + {"type": "BlockRemoved", "block_hashes": [root], "medium": "GPU"}, # the last copy + _store([orphan], [13, 14, 15, 16], parent_block_hash=root), + ], + 2, + ) + assert not second.events[1].stored.HasField("cache_salt") + assert (counts.hash_checked, counts.hash_unverifiable) == (3, 1) + assert counts.forwarded_removed == 2 + + +def test_socket_rank_wins_over_the_payload_rank(): + batches, _ = _normalize([_batch([_store([1], [1, 2, 3, 4])], rank=0)], rank=3) + assert batches[0].dp_rank == 3 + batches, _ = _normalize([msgspec.msgpack.encode([1.5, [_store([1], [1, 2, 3, 4])]])]) + assert not batches[0].HasField("dp_rank") + + +def test_a_non_batch_payload_is_an_error(): + with pytest.raises(ValueError): + kv_relay.decode_batch(msgspec.msgpack.encode({"not": "a batch"})) + + +# --------------------------------------------------------------------------- +# Streaming: ranks, cursors, replay +# --------------------------------------------------------------------------- + + +def _bind_consecutive(ctx, kind, count, attempts=20): + """``count`` sockets of ``kind`` on consecutive ports, as the engines lay out ranks.""" + for _ in range(attempts): + sockets = [] + try: + first = ctx.socket(kind) + if kind == zmq.XPUB: + first.setsockopt(zmq.XPUB_VERBOSE, 1) + base = first.bind_to_random_port("tcp://127.0.0.1") + sockets.append(first) + for offset in range(1, count): + sock = ctx.socket(kind) + if kind == zmq.XPUB: + sock.setsockopt(zmq.XPUB_VERBOSE, 1) + sock.bind(f"tcp://127.0.0.1:{base + offset}") + sockets.append(sock) + return base, sockets + except zmq.ZMQError: + for sock in sockets: + sock.close(linger=0) + raise RuntimeError("no consecutive ports available") + + +@pytest_asyncio.fixture +async def bridge(): + ctx = zmq.asyncio.Context() + base, pubs = _bind_consecutive(ctx, zmq.XPUB, 2) + replay_base, routers = _bind_consecutive(ctx, zmq.ROUTER, 2) + config = SimpleNamespace( + endpoint=f"tcp://127.0.0.1:{base}", + replay_endpoint=f"tcp://127.0.0.1:{replay_base}", + topic="kv", + ) + options = {"replay_timeout": 0.3, "recv_timeout": 0.1} + + async def handler(request, context): + sources = kv_relay.rank_sources(config, range(len(pubs))) + async for batch in kv_relay.relay( + sources, + kv_relay.Engine.SGLANG, + request.start_sequence_number, + context, + topic=config.topic, + zmq_context=ctx, + **options, + ): + yield batch + + server = grpc.aio.server() + server.add_generic_rpc_handlers( + ( + grpc.method_handlers_generic_handler( + "test.KvEvents", + { + "Subscribe": grpc.unary_stream_rpc_method_handler( + handler, + request_deserializer=common_pb2.SubscribeKvEventsRequest.FromString, + response_serializer=common_pb2.KvEventBatch.SerializeToString, + ) + }, + ), + ) + ) + grpc_port = server.add_insecure_port("127.0.0.1:0") + await server.start() + channel = grpc.aio.insecure_channel(f"127.0.0.1:{grpc_port}") + rpc = channel.unary_stream( + "/test.KvEvents/Subscribe", + request_serializer=common_pb2.SubscribeKvEventsRequest.SerializeToString, + response_deserializer=common_pb2.KvEventBatch.FromString, + ) + + def subscribe(cursor=0): + return rpc(common_pb2.SubscribeKvEventsRequest(start_sequence_number=cursor)) + + async def subscribed(): + # XPUB acknowledges each actual subscription; no timing sleeps needed. + for pub in pubs: + while await asyncio.wait_for(pub.recv(), 3) != b"\x01kv": + pass + + async def publish(rank, seq, payload=None): + payload = ( + _batch([_store([seq + 1], [1, 2, 3, 4])], rank=None) if payload is None else payload + ) + await pubs[rank].send_multipart([b"kv", seq.to_bytes(8, "big"), payload]) + + async def replay_request(rank): + frames = await asyncio.wait_for(routers[rank].recv_multipart(), 3) + assert frames[1] == b"" + return frames[0], int.from_bytes(frames[2], "big") + + async def replay_send(rank, identity, seq, payload=None, framing="sglang"): + payload = _batch([_store([seq + 1], [1, 2, 3, 4])]) if payload is None else payload + seq_bytes = kv_relay._END_SEQ if seq == -1 else seq.to_bytes(8, "big") + frames = [identity, b"", seq_bytes, payload if seq != -1 else b""] + if framing == "vllm": + frames.insert(2, b"kv" if seq != -1 else b"") + await routers[rank].send_multipart(frames) + + try: + yield SimpleNamespace( + subscribe=subscribe, + subscribed=subscribed, + publish=publish, + replay_request=replay_request, + replay_send=replay_send, + pubs=pubs, + routers=routers, + config=config, + options=options, + ) + finally: + await channel.close() + await server.stop(None) + for sock in pubs + routers: + sock.close(linger=0) + ctx.term() + + +async def read(call): + return await asyncio.wait_for(call.read(), 3) + + +async def read_error(call, after_at_most=3): + """The status the stream ends with, allowing a few batches before it.""" + with pytest.raises(grpc.aio.AioRpcError) as error: + for _ in range(after_at_most + 1): + await read(call) + return error.value.code() + + +@pytest.mark.asyncio +async def test_ranks_are_tagged_from_their_socket_and_numbered_contiguously(bridge): + call = bridge.subscribe() + await bridge.subscribed() + await bridge.publish(0, 0) + await bridge.publish(0, 1) + await bridge.publish(1, 0) # rank 1 starts its own count; not stale + await bridge.publish(1, 1) + seen = [await read(call) for _ in range(4)] + assert [batch.sequence_number for batch in seen] == [1, 2, 3, 4] + # The two sockets are polled together, so ranks interleave; each rank's + # own order holds and every batch carries its socket's rank. + per_rank = {0: [], 1: []} + for batch in seen: + per_rank[batch.dp_rank].append(batch.events[0].stored.blocks[0].block_hash) + assert per_rank == {0: [1, 2], 1: [1, 2]} + call.cancel() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("framing", ["sglang", "vllm"]) +async def test_a_gap_is_filled_from_that_ranks_replay_before_later_batches(bridge, framing): + call = bridge.subscribe() + await bridge.subscribed() + await bridge.publish(1, 0) + assert (await read(call)).dp_rank == 1 + # Rank 1 skips 1 and 2; rank 0 keeps publishing meanwhile. + await bridge.publish(1, 3) + identity, start = await bridge.replay_request(1) + assert start == 1 + await bridge.publish(0, 0) + for seq in (1, 2, 3): # the replay overlaps the live batch 3 + await bridge.replay_send(1, identity, seq, framing=framing) + await bridge.replay_send(1, identity, -1, framing=framing) + hashes = [] + for _ in range(4): + batch = await read(call) + hashes.append((batch.dp_rank, batch.events[0].stored.blocks[0].block_hash)) + # Replayed 1, 2, 3 for rank 1 in order, the live 3 deduplicated, then rank 0. + assert hashes[:3] == [(1, 2), (1, 3), (1, 4)] + assert hashes[3] == (0, 1) + await bridge.publish(1, 4) + assert (await read(call)).events[0].stored.blocks[0].block_hash == 5 + call.cancel() + + +@pytest.mark.asyncio +async def test_live_batches_the_replay_already_covered_are_duplicates_not_a_restart(bridge): + """A gap's replay usually runs past the live batch that exposed it, and the + live batches the publisher sent meanwhile are queued on the SUB socket: + they arrive below the cursor the replay left. They are duplicates of what + the replay forwarded, not a publisher restart, and the stream goes on; a + sequence below what the replay covered is still a restart.""" + call = bridge.subscribe() + await bridge.subscribed() + await bridge.publish(0, 0) + assert (await read(call)).sequence_number == 1 + await bridge.publish(0, 3) # 1 and 2 skipped + identity, start = await bridge.replay_request(0) + assert start == 1 + await bridge.publish(0, 4) # published while the replay runs: queued behind it + await bridge.publish(0, 5) + for seq in (1, 2, 3, 4, 5): # the replay fills the gap and runs past the live 3, 4 and 5 + await bridge.replay_send(0, identity, seq) + await bridge.replay_send(0, identity, -1) + hashes = [(await read(call)).events[0].stored.blocks[0].block_hash for _ in range(5)] + assert hashes == [2, 3, 4, 5, 6], "replayed 1..5, each once" + await bridge.publish(0, 6) # the queued 4 and 5 were duplicates; live continues + assert (await read(call)).events[0].stored.blocks[0].block_hash == 7 + await bridge.publish(0, 0) # below everything the replay covered: a restart + assert await read_error(call) == grpc.StatusCode.DATA_LOSS + + +@pytest.mark.asyncio +async def test_the_replay_window_closes_once_the_live_stream_is_past_it(bridge): + """The duplicates a replay leaves on the SUB queue can only arrive until + the first live batch past the replay; after that a sequence inside the + old window can only come from a restarted publisher, and it must end + the stream like any other restart instead of being swallowed.""" + call = bridge.subscribe() + await bridge.subscribed() + await bridge.publish(0, 0) + assert (await read(call)).sequence_number == 1 + await bridge.publish(0, 3) # 1 and 2 skipped + identity, start = await bridge.replay_request(0) + assert start == 1 + await bridge.publish(0, 4) # queued behind the replay + for seq in (1, 2, 3, 4): + await bridge.replay_send(0, identity, seq) + await bridge.replay_send(0, identity, -1) + for _ in range(4): + await read(call) + await bridge.publish(0, 5) # live catches up past the replay: the window closes + assert (await read(call)).events[0].stored.blocks[0].block_hash == 6 + await bridge.publish(0, 2) # inside the old window now: a restarted publisher + assert await read_error(call) == grpc.StatusCode.DATA_LOSS + + +@pytest.mark.asyncio +@pytest.mark.parametrize("fault", ["truncated", "empty", "timeout", "short", "malformed"]) +async def test_an_unverifiable_replay_ends_the_stream_with_data_loss(bridge, fault): + call = bridge.subscribe() + await bridge.subscribed() + await bridge.publish(0, 0) + assert (await read(call)).sequence_number == 1 + await bridge.publish(0, 4) + identity, start = await bridge.replay_request(0) + assert start == 1 + if fault == "truncated": + await bridge.replay_send(0, identity, 2) # history no longer holds 1 + elif fault == "empty": + await bridge.replay_send(0, identity, -1) + elif fault == "short": + await bridge.replay_send(0, identity, 1) + await bridge.replay_send(0, identity, -1) # ends before 3 + elif fault == "malformed": + await bridge.routers[0].send_multipart([identity, b"", b"bad"]) + assert await read_error(call) == grpc.StatusCode.DATA_LOSS + + +@pytest.mark.asyncio +async def test_a_gap_without_a_replay_endpoint_ends_the_stream(bridge): + bridge.config.replay_endpoint = None + call = bridge.subscribe() + await bridge.subscribed() + await bridge.publish(0, 0) + assert (await read(call)).sequence_number == 1 + await bridge.publish(0, 2) + assert await read_error(call) == grpc.StatusCode.DATA_LOSS + + +@pytest.mark.asyncio +async def test_a_publisher_restart_ends_the_stream_with_data_loss(bridge): + call = bridge.subscribe() + await bridge.subscribed() + await bridge.publish(0, 5) + assert (await read(call)).sequence_number == 1 + await bridge.publish(0, 0) + assert await read_error(call) == grpc.StatusCode.DATA_LOSS + + +@pytest.mark.asyncio +async def test_a_nonzero_cursor_is_refused_before_subscribing(bridge): + call = bridge.subscribe(100) + assert await read_error(call) == grpc.StatusCode.OUT_OF_RANGE + assert not await bridge.pubs[0].poll(timeout=50) + + +@pytest.mark.asyncio +async def test_bad_payloads_and_duplicates_are_skipped_without_replay(bridge): + call = bridge.subscribe() + await bridge.subscribed() + await bridge.publish(0, 0) + await bridge.publish(0, 1, b"not msgpack") + await bridge.publish(0, 1, b"not msgpack") + await bridge.pubs[0].send_multipart([b"kv", b"short"]) + await bridge.publish(0, 2) + await bridge.publish(0, 2) + first = await read(call) + second = await read(call) + assert (first.sequence_number, second.sequence_number) == (1, 2) + assert second.events[0].stored.blocks[0].block_hash == 3 + assert not await bridge.routers[0].poll(timeout=50), "no replay for a consumed sequence" + call.cancel() + + +@pytest.mark.asyncio +async def test_cancellation_releases_every_subscription(bridge): + call = bridge.subscribe() + await bridge.subscribed() + call.cancel() + for pub in bridge.pubs: + assert await asyncio.wait_for(pub.recv(), 3) == b"\x00kv" + + +# --------------------------------------------------------------------------- +# Engine-exact hashes and the opt-in check +# --------------------------------------------------------------------------- + + +def test_sglang_hashes_reproduce_the_published_vectors(): + def ints(tokens, page, prior=None): + return [i for _, i in kv_relay.sglang_chain(tokens, page, prior)] + + assert ints([1, 2, 3, 4], 4) == [-3488128144981237669] + assert ints([10, 20, 30, 40, 50, 60, 70, 80], 2) == [ + 978178666101069530, + -895308556211281782, + -8033692805846017938, + 835415944263129316, + ] + assert kv_relay.sglang_event_int(hashlib.sha256(b"").digest()) == -2039914840885289964 + seed = kv_relay.sglang_salt_seed("tenant-a") + assert seed.hex() == "f5d0f785efe6042a4e4b0297a4d712917e1763c850835e93df58b373fd47fd2c" + assert ints([1, 2, 3, 4, 5, 6, 7, 8], 4, seed) == [3718046898735569995, 3308615664605479373] + assert ( + kv_relay.sglang_event_int(kv_relay.sglang_page(None, [1, 2, 2, 3, 3, 4, 4, 5])) + == -638950109823820341 + ) + assert ints([1, 2, 3], 4) == [], "only full pages" + + +def test_vllm_sha256_cbor_hashes_reproduce_the_reference_run(): + # hash_block_tokens with sha256_cbor from vLLM at 0c16eee3f1, PYTHONHASHSEED unset. + assert ( + kv_relay.vllm_none_hash().hex() + == "9bd96a485ad84efdafb72ee48a1d7a69bcead0f8f0433173941b276b9581eef0" + ) + a = kv_relay.vllm_block(None, [1, 2, 3, 4]) + assert a.hex() == "58d0879dff3800f65f8c5fd449d73048c0f7699dd238a03b84b146b151d55111" + assert kv_relay.vllm_event_int(a) == -8885242862429187823 + b = kv_relay.vllm_block(a, [5, 6, 7, 8]) + assert b.hex() == "e1bc29393a2da9fd798f35ff04d215f846aee6df4b803ce3d43b57df56b58fb1" + assert [i for _, i in kv_relay.vllm_chain([1, 2, 3, 4, 5, 6, 7, 8], 4)] == [ + -8885242862429187823, + -3153830497298837583, + ] + lora = ("lora", "adapter", "/adapters/adapter") + c = kv_relay.vllm_block( + None, + [1, 2, 3, 4], + [lora, ("mm", "mm-abc", 0), ("cache_salt", "salt-1"), ("prompt_embeds", bytes(range(32)))], + ) + assert c.hex() == "280bce663478b44320e68be593605214c94f7648ac34b4b94274fdcd0f2a8a04" + d = kv_relay.vllm_block(c, [5, 6, 7, 8], [lora, ("mm", "mm-abc", -4)]) + assert d.hex() == "33e7394e25b7be46c7bf0c090da02e903b4e9d9536b1c18eed4d6e57e7a3658a" + e = kv_relay.vllm_block( + None, + [1, 2, 3, 4], + [("mm", "mm-abc", 0), ("cache_salt", "salt-1"), ("prompt_embeds", bytes(range(32)))], + ) + assert e.hex() == "0d7b0d8344b79ad5e183b117cacc04aeb415bfdfa968fd28f611bc6971f4abba" + f = kv_relay.vllm_block(e, [5, 6, 7, 8], [("mm", "mm-abc", -4)]) + assert f.hex() == "b81a4631bec909b72bcd0081cc2d7c87bcce5652c48318f9615e4aa3370c1d69" + g = kv_relay.vllm_block(f, [9, 10, 11, 12]) + assert g.hex() == "024cf74ecdb93c6049fe59a66321192aab04701a75ac3e86619395104df937c4" + + +def _checked(normalizer): + c = normalizer.counts + return c.hash_checked, c.hash_mismatch, c.hash_unverifiable + + +def test_hash_check_verifies_sglang_chains_and_counts_mismatches(): + chain = kv_relay.sglang_chain([1, 2, 3, 4, 5, 6, 7, 8], 4) + first, second = chain[0][1], chain[1][1] + normalizer = kv_relay.Normalizer(hash_check="sglang") + batch = normalizer.normalize_batch( + kv_relay.decode_batch( + _batch( + [ + _store([first], [1, 2, 3, 4]), + _store([second], [5, 6, 7, 8], parent_block_hash=first), + _store([second + 1], [5, 6, 7, 8], parent_block_hash=first), # tampered + _store([99], [9, 10, 11, 12], parent_block_hash=12345), # unknown parent + ] + ) + ), + 1, + ) + assert len(batch.events) == 4, "a mismatch never drops" + assert _checked(normalizer) == (3, 1, 1) + + seed = kv_relay.sglang_salt_seed("tenant-a") + salted = kv_relay.sglang_chain([1, 2, 3, 4, 5, 6, 7, 8], 4, seed) + normalizer = kv_relay.Normalizer(hash_check="SGLANG") + normalizer.normalize_batch( + kv_relay.decode_batch( + _batch([_store([salted[0][1], salted[1][1]], list(range(1, 9)), cache_salt="tenant-a")]) + ), + 1, + ) + assert _checked(normalizer) == (2, 0, 0) + + normalizer = kv_relay.Normalizer(hash_check="sglang") + normalizer.normalize_batch( + kv_relay.decode_batch( + _batch([_store([-638950109823820341], [[1, 2], [2, 3], [3, 4], [4, 5]])]) + ), + 1, + ) + assert _checked(normalizer) == (1, 0, 0) + + +def test_hash_check_verifies_vllm_sha256_cbor_chains(): + a, b = -8885242862429187823, -3153830497298837583 + normalizer = kv_relay.Normalizer(hash_check="vllm_sha256_cbor") + batch = normalizer.normalize_batch( + kv_relay.decode_batch( + _batch( + [ + _store([a, b], list(range(1, 9))), + _store([7], [9, 10, 11, 12], parent_block_hash=b, lora_name="adapter"), + _store( + [8], [9, 10, 11, 12, 13], parent_block_hash=b + ), # unaligned, dropped first + ] + ) + ), + 1, + ) + assert len(batch.events) == 2 + assert _checked(normalizer) == (2, 0, 1) + + embeds = bytes(range(32)) + e = kv_relay.vllm_block( + None, + [1, 2, 3, 4], + [("mm", "mm-abc", 0), ("cache_salt", "salt-1"), ("prompt_embeds", embeds)], + ) + f = kv_relay.vllm_block(e, [5, 6, 7, 8], [("mm", "mm-abc", -4)]) + g = kv_relay.vllm_block(f, [9, 10, 11, 12]) + ei, fi, gi = (kv_relay.vllm_event_int(x) for x in (e, f, g)) + normalizer = kv_relay.Normalizer(hash_check="vllm-sha256-cbor") + batch = normalizer.normalize_batch( + kv_relay.decode_batch( + _batch( + [ + _store([ei], [1, 2, 3, 4], extra_keys=[[["mm-abc", 0], "salt-1", embeds]]), + _store([fi], [5, 6, 7, 8], parent_block_hash=ei, extra_keys=[[["mm-abc", -4]]]), + _store([gi], [9, 10, 11, 12], parent_block_hash=fi), + _store( + [5], [13, 14, 15, 16], parent_block_hash=gi, extra_keys=[[3]] + ), # unknown key shape + ] + ) + ), + 1, + ) + assert len(batch.events) == 4 + assert _checked(normalizer) == (3, 0, 1) + assert batch.events[0].stored.cache_salt == "salt-1" + + +def test_hash_check_is_off_unless_asked(caplog): + assert kv_relay.Normalizer().hash_check is None + assert kv_relay.Normalizer(hash_check="").hash_check is None + with caplog.at_level("WARNING"): + assert kv_relay.Normalizer(hash_check="xxhash").hash_check is None + assert "names no known engine hash" in caplog.text + normalizer = kv_relay.Normalizer() + normalizer.normalize_batch(kv_relay.decode_batch(_batch([_store([1], [1, 2, 3, 4])])), 1) + assert _checked(normalizer) == (0, 0, 0) diff --git a/grpc_servicer/tests/test_sglang_kv_events.py b/grpc_servicer/tests/test_sglang_kv_events.py index 4f0bb20dea..710daa20aa 100644 --- a/grpc_servicer/tests/test_sglang_kv_events.py +++ b/grpc_servicer/tests/test_sglang_kv_events.py @@ -1,379 +1,85 @@ -"""Exercise the SGLang bridge over real ZMQ and gRPC without loading an engine.""" +"""The SGLang bridge's KV-event wiring: every DP rank's publisher goes through +the shared relay (its protocol has its own tests in ``test_kv_relay.py``).""" import ast import asyncio -import importlib.util -import logging from collections.abc import AsyncIterator from pathlib import Path from types import SimpleNamespace import pytest -import pytest_asyncio pytest.importorskip("smg_grpc_proto") grpc = pytest.importorskip("grpc") -zmq = pytest.importorskip("zmq") -import zmq.asyncio # noqa: E402, F811 +pytest.importorskip("zmq") +pytest.importorskip("msgspec") from smg_grpc_proto.generated import common_pb2 # noqa: E402 - -_PATH = Path(__file__).parents[1] / "smg_grpc_servicer" / "sglang" / "kv_events.py" -_spec = importlib.util.spec_from_file_location("sglang_kv_transport", _PATH) -transport = importlib.util.module_from_spec(_spec) -_spec.loader.exec_module(transport) - - -def decode(payload): - return int(payload) - - -def convert(value, seq): - return common_pb2.KvEventBatch(sequence_number=seq, timestamp=value) +from smg_grpc_servicer.sglang import kv_events as transport # noqa: E402 + +_SERVICER = Path(__file__).parents[1] / "smg_grpc_servicer" / "sglang" / "servicer.py" + + +def test_one_source_per_data_parallel_rank(): + config = SimpleNamespace(endpoint="tcp://*:5557", replay_endpoint="tcp://*:5558", topic="") + assert transport.dp_rank_count(SimpleNamespace(dp_size=3)) == 3 + assert transport.dp_rank_count(SimpleNamespace(dp_size=None)) == 1 + assert transport.dp_rank_count(SimpleNamespace()) == 1 + sources = transport.sources(config, SimpleNamespace(dp_size=2)) + assert [(s.rank, s.endpoint, s.replay_endpoint) for s in sources] == [ + (0, "tcp://127.0.0.1:5557", "tcp://127.0.0.1:5558"), + (1, "tcp://127.0.0.1:5558", "tcp://127.0.0.1:5559"), + ] + without_replay = transport.sources( + SimpleNamespace(endpoint="tcp://*:6000"), SimpleNamespace(dp_size=1) + ) + assert without_replay[0].replay_endpoint is None -def subscribe_method(): - # Load the real RPC method without importing SGLang/torch. Its transport - # and gRPC context are real; only engine batch decoding is substituted. - path = _PATH.with_name("servicer.py") - tree = ast.parse(path.read_text()) +def _subscribe_method(namespace): + # The real RPC method without importing SGLang/torch. + tree = ast.parse(_SERVICER.read_text()) method = next( node for node in ast.walk(tree) if isinstance(node, ast.AsyncFunctionDef) and node.name == "SubscribeKvEvents" ) - namespace = { - "AsyncIterator": AsyncIterator, - "common_pb2": common_pb2, - "grpc": grpc, - "asyncio": asyncio, - "zmq": zmq, - "logger": logging.getLogger(__name__), - "KVEventBatch": object, - "msgspec": SimpleNamespace( - msgpack=SimpleNamespace(Decoder=lambda _: SimpleNamespace(decode=decode)) - ), - "ZmqEventPublisher": SimpleNamespace(offset_endpoint_port=transport.endpoint_for_rank), - "subscribe_kv_events": transport.subscribe_kv_events, - } - exec(compile(ast.Module(body=[method], type_ignores=[]), str(path), "exec"), namespace) + namespace.update({"AsyncIterator": AsyncIterator, "common_pb2": common_pb2, "grpc": grpc}) + exec(compile(ast.Module(body=[method], type_ignores=[]), str(_SERVICER), "exec"), namespace) return namespace["SubscribeKvEvents"] -@pytest_asyncio.fixture -async def bridge(): - ctx = zmq.asyncio.Context() - pub = ctx.socket(zmq.XPUB) - pub.setsockopt(zmq.XPUB_VERBOSE, 1) - port = pub.bind_to_random_port("tcp://127.0.0.1") - config = SimpleNamespace(endpoint=f"tcp://127.0.0.1:{port}", topic="kv") - cursors = [] - method = subscribe_method() - servicer = SimpleNamespace(_kv_events_config=config, _convert_kv_event_batch=convert) - - async def handler(request, context): - cursors.append(request.start_sequence_number) - async for batch in method(servicer, request, context): - yield batch - - server = grpc.aio.server() - server.add_generic_rpc_handlers( - ( - grpc.method_handlers_generic_handler( - "test.KvEvents", - { - "Subscribe": grpc.unary_stream_rpc_method_handler( - handler, - request_deserializer=common_pb2.SubscribeKvEventsRequest.FromString, - response_serializer=common_pb2.KvEventBatch.SerializeToString, - ) - }, - ), - ) - ) - grpc_port = server.add_insecure_port("127.0.0.1:0") - await server.start() - channel = grpc.aio.insecure_channel(f"127.0.0.1:{grpc_port}") - rpc = channel.unary_stream( - "/test.KvEvents/Subscribe", - request_serializer=common_pb2.SubscribeKvEventsRequest.SerializeToString, - response_deserializer=common_pb2.KvEventBatch.FromString, - ) - - def subscribe(cursor=0): - return rpc(common_pb2.SubscribeKvEventsRequest(start_sequence_number=cursor)) - - async def subscribed(): - # XPUB acknowledges the actual subscription; no timing sleeps needed. - while await asyncio.wait_for(pub.recv(), 3) != b"\x01kv": - pass - - async def publish(seq, payload=None): - await pub.send_multipart([b"kv", seq.to_bytes(8, "big"), payload or str(seq).encode()]) - - try: - yield SimpleNamespace( - subscribe=subscribe, - subscribed=subscribed, - publish=publish, - config=config, - cursors=cursors, - pub=pub, - ctx=ctx, - ) - finally: - await channel.close() - await server.stop(None) - pub.close(linger=0) - ctx.term() - - -async def read(call): - return await asyncio.wait_for(call.read(), 3) - - -@pytest.mark.asyncio -async def test_gap_resume_is_rejected_then_live_cursor_advances(bridge): - call = bridge.subscribe() - await bridge.subscribed() - await bridge.publish(100) - assert (await read(call)).sequence_number == 100 - await bridge.publish(103) - assert (await read(call)).sequence_number == 103 - call.cancel() - - # The router retries from its last contiguous batch. Do not silently - # reconnect to live seq=110 while it still expects seq=101. - retry = bridge.subscribe(100) - with pytest.raises(grpc.aio.AioRpcError) as error: - await read(retry) - assert error.value.code() == grpc.StatusCode.OUT_OF_RANGE - - # OUT_OF_RANGE triggers the gateway's existing per-worker clear/reset. - fresh = bridge.subscribe() - await bridge.subscribed() - for seq in (110, 111): - await bridge.publish(seq) - assert (await read(fresh)).sequence_number == seq - assert bridge.cursors == [0, 100, 0] - fresh.cancel() - - -@pytest.mark.asyncio -@pytest.mark.parametrize("cursor", [1, 2**64 - 1]) -async def test_replay_rejected_before_opening_live_subscription(bridge, cursor): - call = bridge.subscribe(cursor) - with pytest.raises(grpc.aio.AioRpcError) as error: - await read(call) - assert error.value.code() == grpc.StatusCode.OUT_OF_RANGE - assert not await bridge.pub.poll(timeout=50) - - -@pytest.mark.asyncio -async def test_idle_poll_preserves_later_events_and_cancellation(bridge): - call = bridge.subscribe() - await bridge.subscribed() - await asyncio.wait_for(call.initial_metadata(), 3) - await asyncio.sleep(1.1) # exercise the idle poll timeout - await bridge.publish(0) - assert (await read(call)).sequence_number == 0 - call.cancel() - assert await asyncio.wait_for(bridge.pub.recv(), 3) == b"\x00kv" - - -@pytest.mark.asyncio -async def test_bad_payload_does_not_hide_native_sequence_gap(bridge): - call = bridge.subscribe() - await bridge.subscribed() - await bridge.publish(10) - assert (await read(call)).sequence_number == 10 - await bridge.pub.send_multipart([b"kv", b"short"]) - await bridge.publish(11, b"undecodable") - await bridge.publish(12) - assert (await read(call)).sequence_number == 12 - call.cancel() - - -@pytest_asyncio.fixture -async def replay(bridge): - socket = bridge.ctx.socket(zmq.ROUTER) - port = socket.bind_to_random_port("tcp://127.0.0.1") - bridge.config.replay_endpoint = f"tcp://127.0.0.1:{port}" - - async def request(cursor): - frames = await asyncio.wait_for(socket.recv_multipart(), 3) - assert frames[1:] == [b"", (cursor + 1).to_bytes(8, "big")] - return frames[0] - - async def send(identity, seq, payload=None): - wire_seq = transport._END_SEQ if seq == -1 else seq.to_bytes(8, "big") - await socket.send_multipart( - [identity, b"", wire_seq, str(seq).encode() if payload is None else payload] - ) - - try: - yield SimpleNamespace(request=request, send=send, socket=socket) - finally: - socket.close(linger=0) - - -@pytest.mark.asyncio -async def test_replay_handoff_keeps_order_and_deduplicates_live_overlap(bridge, replay): - call = bridge.subscribe(100) - identity = await replay.request(100) - await bridge.subscribed() - # Queue both overlap and new live traffic while historical replay runs. - for seq in (101, 102, 103): - await bridge.publish(seq) - for seq in (101, 102): - await replay.send(identity, seq) - batch = await read(call) - assert (batch.sequence_number, batch.timestamp) == (seq, seq) - await replay.send(identity, -1) - assert (await read(call)).sequence_number == 103 - await bridge.publish(104) - assert (await read(call)).sequence_number == 104 - assert bridge.cursors == [100] - call.cancel() - - -@pytest.mark.asyncio -@pytest.mark.parametrize("first", [-1, 90, 103]) -async def test_empty_or_expired_history_requires_fresh_subscription(bridge, replay, first): - call = bridge.subscribe(100) - identity = await replay.request(100) - await replay.send(identity, first) - with pytest.raises(grpc.aio.AioRpcError) as error: - await read(call) - assert error.value.code() == grpc.StatusCode.OUT_OF_RANGE - fresh = bridge.subscribe() - await bridge.subscribed() - # The failed call also subscribed; drain its subscribe/unsubscribe so the - # new live subscriber is established before sending. - await bridge.subscribed() - await bridge.publish(120) - assert (await read(fresh)).sequence_number == 120 - fresh.cancel() - - -@pytest.mark.asyncio -@pytest.mark.parametrize("fault", ["gap", "decode", "timeout", "malformed"]) -async def test_partial_replay_failure_signals_data_loss(bridge, replay, monkeypatch, fault): - monkeypatch.setattr(transport, "_REPLAY_TIMEOUT_MS", 100) - call = bridge.subscribe(100) - identity = await replay.request(100) - await replay.send(identity, 101) - assert (await read(call)).sequence_number == 101 - if fault == "gap": - await replay.send(identity, 103) - elif fault == "decode": - await replay.send(identity, 102, b"bad") - elif fault == "malformed": - await replay.socket.send_multipart([identity, b"", b"bad", b""]) - with pytest.raises(grpc.aio.AioRpcError) as error: - await read(call) - assert error.value.code() == grpc.StatusCode.DATA_LOSS - - -@pytest.mark.asyncio -@pytest.mark.parametrize("fault", ["timeout", "malformed", "decode"]) -async def test_replay_failure_before_headers_signals_out_of_range( - bridge, replay, monkeypatch, fault -): - monkeypatch.setattr(transport, "_REPLAY_TIMEOUT_MS", 100) - call = bridge.subscribe(100) - identity = await replay.request(100) - if fault == "malformed": - await replay.socket.send_multipart([identity, b"", b"bad", b""]) - elif fault == "decode": - await replay.send(identity, 101, b"bad") - with pytest.raises(grpc.aio.AioRpcError) as error: - await read(call) - assert error.value.code() == grpc.StatusCode.OUT_OF_RANGE - - -@pytest.mark.asyncio -async def test_zero_cursor_does_not_replay_old_history(bridge, replay): - call = bridge.subscribe() - await bridge.subscribed() - await bridge.publish(7) - assert (await read(call)).sequence_number == 7 - assert not await replay.socket.poll(timeout=50) - call.cancel() - - -@pytest.mark.asyncio -async def test_cancellation_during_replay_releases_live_subscription(bridge, replay): - call = bridge.subscribe(100) - await replay.request(100) - await bridge.subscribed() - call.cancel() - assert await asyncio.wait_for(bridge.pub.recv(), 3) == b"\x00kv" - - @pytest.mark.asyncio -async def test_live_gap_after_replay_remains_visible_for_next_recovery(bridge, replay): - call = bridge.subscribe(100) - identity = await replay.request(100) - await bridge.subscribed() - await replay.send(identity, 101) - await replay.send(identity, -1) - assert (await read(call)).sequence_number == 101 - await bridge.publish(103) - assert (await read(call)).sequence_number == 103 - call.cancel() +async def test_servicer_hands_config_server_args_and_cursor_to_the_relay(): + calls = [] + async def fake_subscribe(config, server_args, start, context): + calls.append((config, server_args, start, context)) + yield common_pb2.KvEventBatch(sequence_number=1) -@pytest.mark.asyncio -async def test_invalid_replay_endpoint_falls_back_without_retaining_cursor(bridge): - bridge.config.replay_endpoint = "invalid://endpoint" - call = bridge.subscribe(100) - with pytest.raises(grpc.aio.AioRpcError) as error: - await read(call) - assert error.value.code() == grpc.StatusCode.OUT_OF_RANGE - - -@pytest.mark.asyncio -async def test_decode_failure_closes_both_sockets(bridge, replay, monkeypatch): - ctx = zmq.asyncio.Context.instance() - sockets = [] - - def socket(kind): - result = ctx.socket(kind) - sockets.append(result) - return result - - monkeypatch.setattr(zmq.asyncio.Context, "instance", lambda: SimpleNamespace(socket=socket)) - call = bridge.subscribe(100) - identity = await replay.request(100) - await replay.send(identity, 101, b"bad") - with pytest.raises(grpc.aio.AioRpcError) as error: - await read(call) - assert error.value.code() == grpc.StatusCode.OUT_OF_RANGE - assert len(sockets) == 2 - assert all(socket.closed for socket in sockets) + method = _subscribe_method({"subscribe_kv_events": fake_subscribe}) + config = SimpleNamespace(endpoint="tcp://*:5557", topic="kv") + server_args = SimpleNamespace(dp_size=2) + servicer = SimpleNamespace(_kv_events_config=config, server_args=server_args) + context = SimpleNamespace() + request = common_pb2.SubscribeKvEventsRequest(start_sequence_number=7) + batches = [batch async for batch in method(servicer, request, context)] + assert [b.sequence_number for b in batches] == [1] + assert calls == [(config, server_args, 7, context)] @pytest.mark.asyncio -@pytest.mark.parametrize("hwm", [None, 4096, 0]) -async def test_live_backlog_survives_replay_with_publisher_hwm(bridge, replay, monkeypatch, hwm): - if hwm is not None: - bridge.config.hwm = hwm - # In-process transport makes queue capacity deterministic, without TCP - # kernel buffers masking a too-small SUB HWM. Limit the sender's share. - bridge.pub.setsockopt(zmq.SNDHWM, 1) - bridge.config.endpoint = "inproc://kv-replay-backlog" - bridge.pub.bind(bridge.config.endpoint) - monkeypatch.setattr(zmq.asyncio.Context, "instance", lambda: bridge.ctx) - - call = bridge.subscribe(100) - identity = await replay.request(100) - await bridge.subscribed() - await replay.send(identity, 101) - assert (await read(call)).sequence_number == 101 - # Keep replay open while more than the default 1000 live batches queue. - for seq in range(102, 2150): - await bridge.publish(seq) - await replay.send(identity, -1) - for seq in range(102, 2150): - assert (await read(call)).sequence_number == seq - call.cancel() +async def test_servicer_without_a_publisher_is_unimplemented(): + aborted = [] + + class Context: + async def abort(self, code, message): + aborted.append((code, message)) + raise asyncio.CancelledError # grpc's abort never returns + + method = _subscribe_method({"subscribe_kv_events": None}) + servicer = SimpleNamespace(_kv_events_config=None, server_args=SimpleNamespace(dp_size=1)) + with pytest.raises(asyncio.CancelledError): + async for _ in method(servicer, common_pb2.SubscribeKvEventsRequest(), Context()): + pass + assert aborted[0][0] == grpc.StatusCode.UNIMPLEMENTED + assert "--kv-events-config" in aborted[0][1] diff --git a/grpc_servicer/tests/test_sglang_rust_servicer.py b/grpc_servicer/tests/test_sglang_rust_servicer.py index c8b1d12ff6..d1438b5d9a 100644 --- a/grpc_servicer/tests/test_sglang_rust_servicer.py +++ b/grpc_servicer/tests/test_sglang_rust_servicer.py @@ -121,6 +121,35 @@ def test_server_facts_carry_the_router_labels_and_the_window(monkeypatch): assert facts["scheduler_info_json"] == "{}" +def test_server_facts_carry_the_kv_events_publisher(monkeypatch): + monkeypatch.setattr(rust, "pairing_protocol_from_env", lambda: "") + off = rust.server_facts(_server_args()) + assert (off["kv_events_endpoint"], off["kv_events_topic"]) == ("", "") + assert off["kv_events_replay_endpoint"] == "" + null_publisher = rust.server_facts(_server_args(kv_events_config='{"publisher": "null"}')) + assert null_publisher["kv_events_endpoint"] == "" + zmq_defaults = rust.server_facts(_server_args(kv_events_config='{"publisher": "zmq"}')) + assert (zmq_defaults["kv_events_endpoint"], zmq_defaults["kv_events_topic"]) == ( + "tcp://*:5557", + "", + ) + explicit = rust.server_facts( + _server_args( + kv_events_config='{"publisher": "zmq", "endpoint": "tcp://*:6100", "topic": "kv"}' + ) + ) + assert (explicit["kv_events_endpoint"], explicit["kv_events_topic"]) == ("tcp://*:6100", "kv") + assert explicit["kv_events_replay_endpoint"] == "", "no replay socket configured" + with_replay = rust.server_facts( + _server_args( + kv_events_config='{"publisher": "zmq", "endpoint": "tcp://*:6100", ' + '"replay_endpoint": "tcp://*:6101", "topic": "kv"}' + ) + ) + assert with_replay["kv_events_replay_endpoint"] == "tcp://*:6101" + assert rust.kv_events_publisher(_server_args(kv_events_config="not json")) == ("", "", "") + + def test_disaggregated_workers_are_refused_up_front(): rust.refuse_disaggregation(_server_args()) rust.refuse_disaggregation(_server_args(disaggregation_mode="null")) diff --git a/grpc_servicer/tests/test_tokenspeed_kv_events.py b/grpc_servicer/tests/test_tokenspeed_kv_events.py index b72650bfa2..4687050c1a 100644 --- a/grpc_servicer/tests/test_tokenspeed_kv_events.py +++ b/grpc_servicer/tests/test_tokenspeed_kv_events.py @@ -192,3 +192,28 @@ def test_attn_dp_rank_zero_is_set(self): proto, _ = shared.convert_batch(batch, seq_num=1, event_id_start=0) assert proto.HasField("dp_rank") assert proto.dp_rank == 0 + + +class TestResolvedReplayEndpoint: + def test_replay_endpoint_is_carried_when_configured(self): + cfg = ts_kv_events.resolve_kv_events_config( + _Args( + _cfg( + enable_kv_cache_events=True, + publisher="zmq", + endpoint="tcp://*:5600", + replay_endpoint="tcp://*:5601", + ) + ) + ) + assert (cfg.endpoint, cfg.replay_endpoint) == ("tcp://*:5600", "tcp://*:5601") + + def test_replay_endpoint_is_empty_when_absent_or_not_a_string(self): + absent = ts_kv_events.resolve_kv_events_config( + _Args(_cfg(enable_kv_cache_events=True, publisher="zmq")) + ) + assert absent.replay_endpoint == "" + odd = ts_kv_events.resolve_kv_events_config( + _Args(_cfg(enable_kv_cache_events=True, publisher="zmq", replay_endpoint=5601)) + ) + assert odd is not None and odd.replay_endpoint == "" diff --git a/grpc_servicer/tests/test_tokenspeed_rust_servicer.py b/grpc_servicer/tests/test_tokenspeed_rust_servicer.py index 14c05e5f2f..c77fd2c019 100644 --- a/grpc_servicer/tests/test_tokenspeed_rust_servicer.py +++ b/grpc_servicer/tests/test_tokenspeed_rust_servicer.py @@ -147,7 +147,7 @@ def test_server_facts_carry_labels_window_and_kv_events(monkeypatch): monkeypatch.setenv("SMG_PAIRING_PROTOCOL", " nixl ") args = FakeServerArgs( kv_events_config='{"enable_kv_cache_events": true, "publisher": "zmq", ' - '"endpoint": "tcp://*:5600", "topic": "kv"}', + '"endpoint": "tcp://*:5600", "replay_endpoint": "tcp://*:5601", "topic": "kv"}', mapping=SimpleNamespace(attn=SimpleNamespace(dp_size=2)), ) facts = rust.server_facts(args) @@ -162,6 +162,7 @@ def test_server_facts_carry_labels_window_and_kv_events(monkeypatch): assert facts["max_running_requests"] == 64 assert facts["data_parallel_size"] == 2 assert (facts["kv_events_endpoint"], facts["kv_events_topic"]) == ("tcp://*:5600", "kv") + assert facts["kv_events_replay_endpoint"] == "tcp://*:5601" monkeypatch.delenv("SMG_PAIRING_PROTOCOL") plain = rust.server_facts(FakeServerArgs()) @@ -169,6 +170,7 @@ def test_server_facts_carry_labels_window_and_kv_events(monkeypatch): assert "dp_size" not in server_args and "pairing_protocol" not in server_args assert plain["data_parallel_size"] == 1 assert plain["kv_events_endpoint"] == "" + assert plain["kv_events_replay_endpoint"] == "" def test_headless_server_args_dial_the_servicer(): diff --git a/grpc_servicer/tests/test_vllm_loads.py b/grpc_servicer/tests/test_vllm_loads.py new file mode 100644 index 0000000000..abbde9f5c2 --- /dev/null +++ b/grpc_servicer/tests/test_vllm_loads.py @@ -0,0 +1,121 @@ +"""The vLLM servicer's load bookkeeping: queued token-work, generation throughput +and the hit rate from what it forwards and streams (no engine needed).""" + +from __future__ import annotations + +import importlib.util +from pathlib import Path + +import pytest + +pytest.importorskip("smg_grpc_proto") + + +@pytest.fixture(scope="module") +def loads_mod(): + """Load loads.py by path: the vllm package __init__ imports the engine.""" + path = Path(__file__).resolve().parent.parent / "smg_grpc_servicer" / "vllm" / "loads.py" + spec = importlib.util.spec_from_file_location("vllm_loads_under_test", path) + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return module + + +class Clock: + def __init__(self): + self.now = 1000.0 + + def __call__(self): + return self.now + + +def test_queued_token_work_is_the_youngest_waiting_prompts_discounted_by_the_hit_rate(loads_mod): + tracker = loads_mod.LoadTracker(clock=Clock()) + assert tracker.estimate(5) == loads_mod.LoadEstimate() + + # Three in flight, none started: the engine says two are waiting, so the + # oldest is in prefill and the two youngest are queued work. + tracker.submitted("a", 1000) + tracker.submitted("b", 300) + tracker.submitted("c", 200) + assert tracker.estimate(2).queued_token_work == 500 + assert tracker.estimate(0).queued_token_work == 0 + assert tracker.estimate(10).queued_token_work == 1500, "never more than we hold" + + # The first output of `a`: half its prompt was cached; queued prompts are + # discounted by that rate. + tracker.first_output("a", 1000, 500) + assert tracker.pending() == 2 + estimate = tracker.estimate(2) + assert estimate.cache_hit_rate == 0.5 + assert estimate.queued_token_work == 250 + + tracker.finished("b") # ended before it started + assert tracker.estimate(2).queued_token_work == 100 + tracker.first_output("c", 200, 200) + assert tracker.pending() == 0 + assert tracker.estimate(3).queued_token_work == 0 + assert tracker.estimate(0).cache_hit_rate == pytest.approx(700 / 1200) + + +def test_throughput_counts_tokens_inside_the_window_only(loads_mod): + clock = Clock() + tracker = loads_mod.LoadTracker(clock=clock) + tracker.generated(1000) + clock.now += 1.0 + tracker.generated(3000) + assert tracker.estimate(0).gen_throughput == 2000.0 + clock.now += 1.5 # the first sample falls out of the two-second window + assert tracker.estimate(0).gen_throughput == 1500.0 + clock.now += 2.0 + assert tracker.estimate(0).gen_throughput == 0.0 + tracker.generated(0) + assert tracker.estimate(0).gen_throughput == 0.0 + + +def test_the_hit_rate_averages_recent_first_outputs_and_ignores_empty_prompts(loads_mod): + tracker = loads_mod.LoadTracker(clock=Clock()) + tracker.submitted("x", 0) + tracker.first_output("x", 0, 0) + assert tracker.estimate(0).cache_hit_rate == 0.0 + for index in range(loads_mod.HIT_RATE_SAMPLES): + tracker.first_output(f"old-{index}", 100, 0) + assert tracker.estimate(0).cache_hit_rate == 0.0 + for index in range(loads_mod.HIT_RATE_SAMPLES): + tracker.first_output(f"new-{index}", 100, 100) + assert tracker.estimate(0).cache_hit_rate == 1.0, "the old samples aged out" + tracker.first_output("odd", 10, 50) + assert tracker.estimate(0).cache_hit_rate <= 1.0 + + +def test_scheduler_load_fields_carry_vllms_counts_and_the_estimate(loads_mod): + estimate = loads_mod.LoadEstimate( + queued_token_work=640, gen_throughput=1234.5, cache_hit_rate=0.25 + ) + fields = loads_mod.scheduler_load_fields( + 3, 2, 0.4, estimate, max_total_num_tokens=10_000, max_running_requests=64 + ) + assert fields == { + "dp_rank": 0, + "num_running_reqs": 3, + "num_waiting_reqs": 2, + "num_waiting_uncached_tokens": 640, + "num_total_reqs": 5, + "token_usage": 0.4, + "utilization": 0.4, + "gen_throughput": 1234.5, + "cache_hit_rate": 0.25, + "max_total_num_tokens": 10_000, + "num_used_tokens": 4_000, + "max_running_requests": 64, + } + # Unknown capacity figures stay unset rather than zero-filled; a negative + # KV usage (a stats race) reads as empty. + bare = loads_mod.scheduler_load_fields(0, 0, -0.1, loads_mod.LoadEstimate()) + assert "max_total_num_tokens" not in bare and "num_used_tokens" not in bare + assert bare["token_usage"] == 0.0 + + from smg_grpc_proto.generated import vllm_engine_pb2 + + load = vllm_engine_pb2.SchedulerLoad(**fields) + assert (load.num_waiting_uncached_tokens, load.gen_throughput) == (640, 1234.5) diff --git a/model_gateway/Cargo.toml b/model_gateway/Cargo.toml index bdb6622614..7c835dad0c 100644 --- a/model_gateway/Cargo.toml +++ b/model_gateway/Cargo.toml @@ -187,6 +187,10 @@ wasmtime-wasi = { workspace = true } lru = { workspace = true } wat = "1.252" zip = "8.6" +# The KV-index exactness tests and the decision bench run the mock engine's +# streams through the relay's decoder and normalizer into the gateway's index. +engine-servicer.workspace = true +rmp-serde.workspace = true [[bench]] name = "wasm_middleware_latency" @@ -252,6 +256,16 @@ name = "scheduler_load" harness = false path = "benches/scheduler_load.rs" +[[bench]] +name = "policy_selection" +harness = false +path = "benches/policy_selection.rs" + +[[bench]] +name = "kv_index_decision" +harness = false +path = "benches/kv_index_decision.rs" + [lints] workspace = true diff --git a/model_gateway/benches/kv_index_decision.rs b/model_gateway/benches/kv_index_decision.rs new file mode 100644 index 0000000000..eab8865ba1 --- /dev/null +++ b/model_gateway/benches/kv_index_decision.rs @@ -0,0 +1,513 @@ +//! T6, the routing decision's cost: what one request pays in `cache_aware`'s +//! `select_worker` (block hashing, the KV index lookup, the selection) against +//! an index populated to 128 workers, for the positional indexer and the chain +//! index (`--kv-index`). +//! +//! The index is fed through the KV event monitor's own apply path: the mock +//! engine's vLLM-shaped stream (`mock_streams`, the exactness tests' generator) +//! into every worker, then a synthetic fleet state shaped like Mooncake +//! conversation traffic (sessions of 8 to 512 blocks, most starting with one of +//! a few shared chat-template prefixes, each session resident on one to four +//! workers, every worker filled to its block budget). Requests are new turns on +//! resident sessions, prefixes of them, and novel prompts on a shared prefix. +//! +//! Two measurements per backend: the criterion sample of one decision, and the +//! contract's condition, a sustained loop at `T6_RATE` decisions per second +//! (default 10,000) for `T6_SECS` seconds (default 12, the first 2 discarded) +//! on whatever core the process is pinned to, reporting p50/p99/p999 over the +//! decisions with the lookup's and the hashing's own distributions beside +//! them. `T6_WORKERS` (128) and `T6_BLOCKS_PER_WORKER` (8192) size the fleet; +//! `T6_INDEX=positional|chain` runs one backend alone, which is how the index's +//! RSS delta is measured (the second backend in a process reuses the first's +//! freed pages and reads zero). +//! +//! Run with (one pinned core, outside the measurement set): +//! taskset -c 100 cargo bench -p smg --bench kv_index_decision +#![expect( + clippy::unwrap_used, + clippy::expect_used, + clippy::print_stderr, + reason = "benchmark code: panicking on setup failure is expected, eprintln is the report" +)] + +use std::{ + fs, + hint::black_box, + sync::Arc, + time::{Duration, Instant}, +}; + +use criterion::Criterion; +use kv_index::{compute_content_hash, compute_request_content_hashes, request_prefix_hashes}; +use openai_protocol::worker::HealthCheckConfig; +use smg::{ + config::KvIndexKind, + policies::{CacheAwareConfig, CacheAwarePolicy, LoadBalancingPolicy, SelectWorkerInfo}, + worker::{ + kv_event_monitor::bench_support::IndexFeed, BasicWorkerBuilder, KvEventMonitor, KvIndex, + Worker, WorkerType, + }, +}; +use smg_grpc_client::common_proto::{ + kv_cache_event, KvBlock, KvBlocksStored, KvCacheEvent, KvEventBatch, +}; + +/// The exactness tests' stream generator, by path: the mock engine's streams +/// through the relay's decoder and normalizer. +#[path = "../src/worker/kv_index_backend/mock_streams.rs"] +#[expect( + dead_code, + reason = "the bench takes the payloads; the exactness tests the checkpoints and prompts too" +)] +mod mock_streams; + +const BLOCK: usize = 16; +/// Workers whose model id is empty route under this key. +const MODEL: &str = "unknown"; + +fn env_or(name: &str, default: u64) -> u64 { + std::env::var(name) + .ok() + .and_then(|value| value.parse().ok()) + .unwrap_or(default) +} + +/// SplitMix64: a deterministic fleet and request stream. +struct Rng(u64); + +impl Rng { + fn next(&mut self) -> u64 { + self.0 = self.0.wrapping_add(0x9e37_79b9_7f4a_7c15); + let mut z = self.0; + z = (z ^ (z >> 30)).wrapping_mul(0xbf58_476d_1ce4_e5b9); + z = (z ^ (z >> 27)).wrapping_mul(0x94d0_49bb_1331_11eb); + z ^ (z >> 31) + } + + fn below(&mut self, n: usize) -> usize { + (self.next() % n as u64) as usize + } + + /// Log-uniform in `[lo, hi]`. + fn log_uniform(&mut self, lo: usize, hi: usize) -> usize { + let (lo_f, hi_f) = ((lo as f64).ln(), (hi as f64).ln()); + let unit = (self.next() >> 11) as f64 / (1u64 << 53) as f64; + (lo_f + (hi_f - lo_f) * unit).exp().round() as usize + } +} + +/// Token ids of block `position` of stream `stream`: distinct per pair. +fn tokens(stream: u64, position: usize) -> Vec { + (0..BLOCK as u32) + .map(|i| { + let word = stream + .wrapping_mul(0x2545_f491_4f6c_dd1d) + .wrapping_add(position as u64 * 0x9e37_79b9) + .wrapping_add(u64::from(i)); + (word >> 7) as u32 & 0x3_ffff + }) + .collect() +} + +/// A resident chain: its token blocks and the engine hashes an engine would +/// publish (the chain hash of the contents so far, shared by every holder). +struct Chain { + token_blocks: Vec>, + hashes: Vec, +} + +impl Chain { + fn new(token_blocks: Vec>) -> Self { + let contents: Vec<_> = token_blocks + .iter() + .map(|tokens| compute_content_hash(tokens)) + .collect(); + let hashes = request_prefix_hashes(&contents) + .into_iter() + .map(|prefix| prefix.0 as i64) + .collect(); + Self { + token_blocks, + hashes, + } + } + + fn tokens(&self, blocks: usize) -> Vec { + self.token_blocks[..blocks].concat() + } + + /// Stores of at most 16 blocks each, chained by parent, as vLLM publishes. + fn stores(&self) -> Vec { + (0..self.token_blocks.len()) + .step_by(16) + .map(|start| { + let end = (start + 16).min(self.token_blocks.len()); + KvCacheEvent { + event_id: 0, + data: Some(kv_cache_event::Data::Stored(KvBlocksStored { + blocks: (start..end) + .map(|i| KvBlock { + block_hash: self.hashes[i], + token_ids: self.token_blocks[i].clone(), + block_size: BLOCK as i32, + ..Default::default() + }) + .collect(), + parent_block_hash: (start > 0).then(|| self.hashes[start - 1]), + ..Default::default() + })), + } + }) + .collect() + } +} + +/// The fleet's sessions: shared chat-template prefixes, then conversations. +struct Sessions { + prefixes: Vec>>, + chains: Vec, +} + +impl Sessions { + fn generate(rng: &mut Rng, count: usize) -> Self { + let prefixes: Vec>> = (0..8) + .map(|p| { + let len = rng.log_uniform(16, 64); + (0..len).map(|i| tokens(1_000 + p, i)).collect() + }) + .collect(); + let chains = (0..count) + .map(|s| { + let mut token_blocks = if rng.below(10) < 6 { + prefixes[rng.below(prefixes.len())].clone() + } else { + Vec::new() + }; + let body = rng.log_uniform(8, 512); + token_blocks.extend((0..body).map(|i| tokens(10_000 + s as u64, i))); + Chain::new(token_blocks) + }) + .collect(); + Self { prefixes, chains } + } +} + +/// The engine stream every worker starts from: the mock engine's vLLM-shaped +/// run, as the relay forwards it. +fn engine_batches() -> Vec { + mock_streams::Stream::generate(mock_streams::Shape::Vllm, 1, 320).normalized() +} + +fn rss_mb() -> f64 { + fs::read_to_string("/proc/self/status") + .ok() + .and_then(|status| { + status + .lines() + .find(|line| line.starts_with("VmRSS:")) + .and_then(|line| line.split_whitespace().nth(1)) + .and_then(|kb| kb.parse::().ok()) + }) + .map_or(f64::NAN, |kb| kb / 1024.0) +} + +/// A populated gateway: the policy with its monitor and index, the workers, +/// and the request stream. +struct Setup { + policy: CacheAwarePolicy, + index: Arc, + workers: Vec>, + requests: Vec>, + memberships: usize, + index_mb: f64, + fill_secs: f64, +} + +fn setup(kind: KvIndexKind, worker_count: usize, blocks_per_worker: usize) -> Setup { + let rss_before = rss_mb(); + let started = Instant::now(); + // One seed for both backends: the same fleet, the same requests. + let mut rng = Rng(1); + let sessions = Sessions::generate(&mut rng, worker_count * 48); + let engine_stream = engine_batches(); + + let index = Arc::new(KvIndex::new(kind, 64)); + let workers: Vec> = (0..worker_count) + .map(|i| { + Arc::new( + BasicWorkerBuilder::new(format!("grpc://worker-{i:03}:9000")) + .worker_type(WorkerType::Regular) + .health_config(HealthCheckConfig { + disable_health_check: true, + ..Default::default() + }) + .build(), + ) as Arc + }) + .collect(); + // Every session lives on one to four workers; each worker is filled to + // its budget from the sessions assigned to it, the engine stream first. + let mut assigned: Vec> = vec![Vec::new(); worker_count]; + for (s, _) in sessions.chains.iter().enumerate() { + let holders = 1 + rng.below(4); + for _ in 0..holders { + assigned[rng.below(worker_count)].push(s); + } + } + let mut memberships = 0usize; + for (w, worker) in workers.iter().enumerate() { + let mut feed = IndexFeed::new(&index, worker.url()).unwrap(); + for batch in &engine_stream { + feed.apply(&index, batch); + } + let mut budget = blocks_per_worker; + let mut sessions_here = assigned[w].clone(); + let mut next_unique = 0u64; + while budget > 0 { + let chain = match sessions_here.pop() { + Some(s) => &sessions.chains[s], + None => { + // A chain only this worker holds. + next_unique += 1; + let len = rng.log_uniform(8, 256).min(budget.max(8)); + let stream = 1_000_000 + (w as u64) * 1_000_000 + next_unique; + let token_blocks = (0..len).map(|i| tokens(stream, i)).collect(); + let chain = Chain::new(token_blocks); + feed.apply( + &index, + &KvEventBatch { + events: chain.stores(), + ..Default::default() + }, + ); + budget = budget.saturating_sub(len); + continue; + } + }; + feed.apply( + &index, + &KvEventBatch { + events: chain.stores(), + ..Default::default() + }, + ); + budget = budget.saturating_sub(chain.token_blocks.len()); + } + memberships += index.worker_block_count(feed.worker_id()); + } + let index_mb = rss_mb() - rss_before; + + let monitor = Arc::new(KvEventMonitor::with_kind(kind, None)); + monitor.set_index(MODEL, Arc::clone(&index)); + monitor.set_block_size(MODEL, BLOCK); + let policy = CacheAwarePolicy::with_config(CacheAwareConfig { + eviction_interval_secs: 0, + block_size: BLOCK, + ..Default::default() + }); + policy.init_workers(&workers); + policy.set_kv_event_monitor(Some(monitor)); + + // Requests: a new turn on a resident session (its chain plus 1 to 16 new + // blocks), a prefix of a resident session, or a novel prompt on a shared + // chat-template prefix. + let requests: Vec> = (0..8192) + .map(|r| { + let draw = rng.below(10); + if draw < 7 { + let chain = &sessions.chains[rng.below(sessions.chains.len())]; + let mut turn = chain.tokens(chain.token_blocks.len()); + let more = 1 + rng.below(16); + for i in 0..more { + turn.extend(tokens(5_000_000 + r as u64, i)); + } + turn + } else if draw < 9 { + let chain = &sessions.chains[rng.below(sessions.chains.len())]; + let cut = 1 + rng.below(chain.token_blocks.len()); + chain.tokens(cut) + } else { + let mut novel = sessions.prefixes[rng.below(sessions.prefixes.len())].concat(); + for i in 0..rng.log_uniform(4, 256) { + novel.extend(tokens(6_000_000 + r as u64, i)); + } + novel + } + }) + .collect(); + Setup { + policy, + index, + workers, + requests, + memberships, + index_mb, + fill_secs: started.elapsed().as_secs_f64(), + } +} + +fn percentile(sorted: &[u64], q: f64) -> f64 { + if sorted.is_empty() { + return f64::NAN; + } + let rank = ((sorted.len() - 1) as f64 * q).round() as usize; + sorted[rank] as f64 / 1000.0 +} + +struct Dist { + p50: f64, + p90: f64, + p99: f64, + p999: f64, + max: f64, + mean: f64, + count: usize, +} + +fn dist(mut samples: Vec) -> Dist { + samples.sort_unstable(); + let count = samples.len(); + let mean = samples.iter().sum::() as f64 / count.max(1) as f64 / 1000.0; + Dist { + p50: percentile(&samples, 0.5), + p90: percentile(&samples, 0.9), + p99: percentile(&samples, 0.99), + p999: percentile(&samples, 0.999), + max: percentile(&samples, 1.0), + mean, + count, + } +} + +/// The contract's condition: one decision every `1 / rate` seconds on this +/// core, for `secs` seconds, the first two discarded as warm-up. Returns the +/// decision latencies and the achieved rate. +fn sustained( + setup: &Setup, + rate: u64, + secs: u64, + mut work: impl FnMut(&Setup, &[u32]), +) -> (Vec, f64) { + let interval = Duration::from_nanos(1_000_000_000 / rate.max(1)); + let warmup = Duration::from_secs(2.min(secs)); + let total = Duration::from_secs(secs); + let start = Instant::now(); + let mut next = start; + let mut samples = Vec::with_capacity((rate * secs) as usize); + let mut decisions = 0u64; + let mut r = 0usize; + loop { + while Instant::now() < next { + std::hint::spin_loop(); + } + let now = Instant::now(); + if now - start >= total { + break; + } + let request = &setup.requests[r % setup.requests.len()]; + r += 1; + let t0 = Instant::now(); + work(setup, request); + let elapsed = t0.elapsed(); + decisions += 1; + if now - start >= warmup { + samples.push(elapsed.as_nanos() as u64); + } + next += interval; + if next < Instant::now() - interval { + // Fell behind by more than a tick: resume the schedule from now. + next = Instant::now(); + } + } + (samples, decisions as f64 / start.elapsed().as_secs_f64()) +} + +fn decide(setup: &Setup, tokens: &[u32]) { + let info = SelectWorkerInfo { + tokens: Some(tokens), + ..Default::default() + }; + black_box(setup.policy.select_worker(&setup.workers, &info)); +} + +fn lookup(setup: &Setup, tokens: &[u32]) { + let hashes = compute_request_content_hashes(tokens, BLOCK); + black_box(setup.index.find_matches(&hashes, false)); +} + +fn hashing(_: &Setup, tokens: &[u32]) { + black_box(compute_request_content_hashes(tokens, BLOCK)); +} + +fn report(kind: KvIndexKind, setup: &Setup, rate: u64, secs: u64) { + let (decision, achieved) = sustained(setup, rate, secs, decide); + let (lookup_samples, _) = sustained(setup, rate, secs.min(4), lookup); + let (hash_samples, _) = sustained(setup, rate, secs.min(4), hashing); + let decision = dist(decision); + let lookup = dist(lookup_samples); + let hash = dist(hash_samples); + let mut lens: Vec = setup.requests.iter().map(Vec::len).collect(); + lens.sort_unstable(); + eprintln!( + "| {kind:?} | {workers} | {memberships} | {index_mb:.0} | {fill:.1} | {achieved:.0} | {n} | \ + {d50:.2} | {d90:.2} | {d99:.2} | {d999:.2} | {dmax:.1} | {l50:.2} | {l99:.2} | {lshare:.0}% | \ + {h50:.2} | {h99:.2} | {tok50} | {tok99} |", + workers = setup.workers.len(), + memberships = setup.memberships, + index_mb = setup.index_mb, + fill = setup.fill_secs, + n = decision.count, + d50 = decision.p50, + d90 = decision.p90, + d99 = decision.p99, + d999 = decision.p999, + dmax = decision.max, + l50 = lookup.p50, + l99 = lookup.p99, + lshare = 100.0 * lookup.mean / decision.mean, + h50 = hash.p50, + h99 = hash.p99, + tok50 = lens[lens.len() / 2], + tok99 = lens[lens.len() * 99 / 100], + ); +} + +fn main() { + let workers = env_or("T6_WORKERS", 128) as usize; + let blocks_per_worker = env_or("T6_BLOCKS_PER_WORKER", 8192) as usize; + let rate = env_or("T6_RATE", 10_000); + let secs = env_or("T6_SECS", 12); + // A Prometheus recorder, so the decision pays for its metrics as in the gateway. + let _recorder = metrics_exporter_prometheus::PrometheusBuilder::new() + .install_recorder() + .expect("recorder"); + let mut criterion = Criterion::default().configure_from_args(); + eprintln!( + "| index | workers | memberships | index MB | fill s | achieved/s | decisions | dec p50 us | p90 | p99 | p999 | max | \ + lookup p50 | lookup p99 | lookup share | hash p50 | hash p99 | req tokens p50 | p99 |" + ); + eprintln!("|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|"); + // `T6_INDEX=positional|chain` runs one backend in its own process: the second + // backend in one process reuses the pages the first freed, so its RSS + // delta reads zero. + let only = std::env::var("T6_INDEX") + .ok() + .and_then(|value| KvIndexKind::parse(&value)); + for kind in [KvIndexKind::Positional, KvIndexKind::Chain] { + if only.is_some_and(|only| only != kind) { + continue; + } + let setup = setup(kind, workers, blocks_per_worker); + let mut r = 0usize; + criterion.bench_function(&format!("decision/{kind:?}"), |b| { + b.iter(|| { + let request = &setup.requests[r % setup.requests.len()]; + r += 1; + decide(&setup, request); + }); + }); + report(kind, &setup, rate, secs); + drop(setup); + } + criterion.final_summary(); +} diff --git a/model_gateway/benches/policy_selection.rs b/model_gateway/benches/policy_selection.rs new file mode 100644 index 0000000000..5de831f336 --- /dev/null +++ b/model_gateway/benches/policy_selection.rs @@ -0,0 +1,120 @@ +//! Selection-policy stage cost: what one request pays inside `WorkerSelectionPolicy::select` +//! for each policy, and what the host pays to gather the inputs from a 128-worker overlap map +//! before calling it. +//! +//! Run with: cargo bench --bench policy_selection +#![expect( + clippy::unwrap_used, + reason = "benchmark code: panicking on setup failure is expected" +)] + +use std::{collections::HashMap, hint::black_box}; + +use criterion::{criterion_group, criterion_main, BenchmarkId, Criterion, Throughput}; +use smg::policies::cost::{build, CandidateInputs, RequestInputs, POLICY_NAMES}; + +const BLOCK_SIZE: usize = 16; +const PROMPT_TOKENS: usize = 4_096; + +/// Deterministic pseudo-random fleet state (SplitMix64). +fn mix(mut value: u64) -> u64 { + value = (value ^ (value >> 30)).wrapping_mul(0xbf58_476d_1ce4_e5b9); + value = (value ^ (value >> 27)).wrapping_mul(0x94d0_49bb_1331_11eb); + value ^ (value >> 31) +} + +struct Fleet { + urls: Vec, + /// Overlap in blocks for the workers that hold part of the prompt (the indexer's map). + overlap: HashMap, + prefix_hashes: Vec, +} + +fn fleet(workers: usize) -> Fleet { + let urls: Vec = (0..workers) + .map(|i| format!("http://worker-{i:03}:8000")) + .collect(); + let overlap = (0..workers as u32) + .filter(|&i| mix(u64::from(i)).is_multiple_of(3)) + .map(|i| (i, (mix(u64::from(i) + 7) % 256) as u32 + 1)) + .collect(); + let prefix_hashes = (0..PROMPT_TOKENS / BLOCK_SIZE) + .map(|i| mix(i as u64 + 1_000)) + .collect(); + Fleet { + urls, + overlap, + prefix_hashes, + } +} + +fn gather<'a>(fleet: &'a Fleet, all_workers: bool) -> Vec> { + fleet + .urls + .iter() + .enumerate() + .filter_map(|(idx, url)| { + let overlap = fleet.overlap.get(&(idx as u32)).copied().unwrap_or(0); + if overlap == 0 && !all_workers { + return None; + } + Some(CandidateInputs { + idx, + url, + device_blocks: f64::from(overlap), + effective_score: f64::from(overlap), + }) + }) + .collect() +} + +fn request(fleet: &Fleet) -> RequestInputs<'_> { + RequestInputs { + prompt_tokens: PROMPT_TOKENS, + block_size: BLOCK_SIZE, + request_blocks: PROMPT_TOKENS / BLOCK_SIZE, + avg_load: 8.0, + prefix_hashes: Some(&fleet.prefix_hashes), + } +} + +/// The policy stage alone: inputs are prepared once, `select` runs per iteration. +fn bench_select(c: &mut Criterion) { + let mut group = c.benchmark_group("policy_selection/select"); + for workers in [8, 32, 128] { + let fleet = fleet(workers); + for name in POLICY_NAMES { + let policy = build(name, 0.0).unwrap(); + let inputs = gather(&fleet, policy.needs().all_workers); + let req = request(&fleet); + group.throughput(Throughput::Elements(1)); + group.bench_with_input(BenchmarkId::new(*name, workers), &inputs, |b, inputs| { + b.iter(|| black_box(policy.select(black_box(&req), black_box(inputs)))); + }); + } + } + group.finish(); +} + +/// Gather plus select: per iteration, build the candidate inputs from the 128-worker overlap +/// map (what the cache-aware host does), then select. +fn bench_gather_and_select(c: &mut Criterion) { + let mut group = c.benchmark_group("policy_selection/gather_and_select"); + let fleet = fleet(128); + for name in POLICY_NAMES { + let policy = build(name, 0.0).unwrap(); + let all_workers = policy.needs().all_workers; + let req = request(&fleet); + group.throughput(Throughput::Elements(1)); + group.bench_function(BenchmarkId::new(*name, 128), |b| { + b.iter(|| { + let inputs = gather(black_box(&fleet), all_workers); + black_box(policy.select(black_box(&req), &inputs)) + }); + }); + } + group.finish(); +} + +criterion_group!(benches, bench_select, bench_gather_and_select); +criterion_main!(benches); diff --git a/model_gateway/benches/workers_endpoint.rs b/model_gateway/benches/workers_endpoint.rs index 4198d1c6e4..93e41161c7 100644 --- a/model_gateway/benches/workers_endpoint.rs +++ b/model_gateway/benches/workers_endpoint.rs @@ -98,6 +98,7 @@ fn deep_clone_info(w: &Arc) -> WorkerInfo { load: w.load(), http2: meta.http2, pd_pairing: None, + stalled: None, engine_load: None, job_status: None, } diff --git a/model_gateway/src/app_context.rs b/model_gateway/src/app_context.rs index bb9445b4eb..cc1d0db332 100644 --- a/model_gateway/src/app_context.rs +++ b/model_gateway/src/app_context.rs @@ -13,10 +13,10 @@ use smg_data_connector::{ use smg_mcp::McpOrchestrator; use tokio::sync::broadcast::error::RecvError; use tool_parser::ParserFactory as ToolParserFactory; -use tracing::debug; +use tracing::{debug, warn}; use crate::{ - config::RouterConfig, + config::{KvIndexKind, RouterConfig}, middleware::{AuthConfig, TokenBucket}, observability::inflight_tracker::InFlightRequestTracker, policies::PolicyRegistry, @@ -30,8 +30,8 @@ use crate::{ }, wasm::{config::WasmRuntimeConfig, module_manager::WasmModuleManager}, worker::{ - KvEventMonitor, PrefillAdmission, WorkerHttpClientCache, WorkerMonitor, WorkerRegistry, - WorkerService, + liveness, KvEventMonitor, PrefillAdmission, WorkerHttpClientCache, WorkerMonitor, + WorkerRegistry, WorkerService, }, workflow::{JobQueue, WorkflowEngines}, }; @@ -687,6 +687,21 @@ impl AppContextBuilder { // The overload shed advertises the poll interval as Retry-After — the // veto cannot clear between polls. overload::set_shed_retry_after_secs(config.load_monitor_interval_secs); + if let Some(registry) = self.worker_registry.as_ref() { + registry.set_overload_shed(config.worker_overload_shed); + } + // Progress-based liveness thresholds (see `worker::liveness`). + liveness::configure( + Duration::from_secs(config.worker_stall_secs), + Duration::from_secs(config.worker_wedge_secs), + ); + liveness::configure_warmup(liveness::Warmup { + secs: Duration::from_secs(config.worker_warmup_secs), + share: config.worker_warmup_share, + blocks: config.worker_warmup_blocks, + thin_ratio: config.worker_warmup_thin_ratio, + divert_every: config.worker_warmup_divert_every, + }); // PD dispatch waits here, not in the decode engine's queue, when the // pair's running window is full. pd_admission::set_pd_admission_wait_secs(config.pd_admission_wait_secs); @@ -778,8 +793,18 @@ impl AppContextBuilder { }; if is_cache_aware { - let monitor = Arc::new(KvEventMonitor::new(None)); - debug!("Created KV event monitor for event-driven cache-aware routing"); + let monitor = Arc::new(KvEventMonitor::with_kind(config.kv_index, None)); + debug!( + kv_index = config.kv_index.as_str(), + "Created KV event monitor for event-driven cache-aware routing" + ); + // The load records on the event streams are polls of the worker. + if let Some(worker_monitor) = &self.worker_monitor { + monitor.set_load_sink(worker_monitor); + } + if KvIndexKind::deprecated_alias_used() { + warn!("--kv-index run is the deprecated spelling of --kv-index chain"); + } // Optional indexer bounding: prune entries by last-touch TTL and/or // capacity ceiling. Both default off (unbounded, prior behavior). @@ -787,6 +812,7 @@ impl AppContextBuilder { config.kv_indexer_ttl_secs.unwrap_or(0), config.kv_indexer_max_entries.unwrap_or(0), ); + monitor.start_stats_task(); // Inject monitor into PolicyRegistry — propagates to default_policy // and any other existing cache-aware policies. @@ -982,6 +1008,8 @@ mod tests { cache_index: Default::default(), cache_ttl_secs: 180, cache_boundaries: Vec::new(), + selection_policy: None, + selection_accounting_ttl_ms: 0, }); let builder = AppContextBuilder::new() .with_client(&config, 5) @@ -1025,6 +1053,8 @@ mod tests { cache_index: Default::default(), cache_ttl_secs: 180, cache_boundaries: Vec::new(), + selection_policy: None, + selection_accounting_ttl_ms: 0, })); } @@ -1048,6 +1078,8 @@ mod tests { cache_index: Default::default(), cache_ttl_secs: 180, cache_boundaries: Vec::new(), + selection_policy: None, + selection_accounting_ttl_ms: 0, }; let mut config = config_with_policy(PolicyConfig::Random); diff --git a/model_gateway/src/config/builder.rs b/model_gateway/src/config/builder.rs index 81e7acd590..3979627347 100644 --- a/model_gateway/src/config/builder.rs +++ b/model_gateway/src/config/builder.rs @@ -5,9 +5,10 @@ use smg_mcp::McpConfig; use super::{ CacheIndexKind, CircuitBreakerConfig, ConfigError, ConfigResult, DiscoveryConfig, - HealthCheckConfig, HistoryBackend, KubernetesDiscoveryConfig, MetricsConfig, OracleConfig, - PdPairingMode, PolicyConfig, PostgresConfig, RedisConfig, RetryConfig, RouterConfig, - RoutingKeyOverrideConfig, RoutingMode, TenantApiKeyEntry, TokenizerCacheConfig, TraceConfig, + HealthCheckConfig, HistoryBackend, KubernetesDiscoveryConfig, KvIndexKind, MetricsConfig, + OracleConfig, PdPairingMode, PolicyConfig, PostgresConfig, RedisConfig, RetryConfig, + RouterConfig, RoutingKeyOverrideConfig, RoutingMode, TenantApiKeyEntry, TokenizerCacheConfig, + TraceConfig, }; use crate::worker::{ConnectionMode, RuntimeType}; @@ -144,6 +145,8 @@ impl RouterConfigBuilder { cache_index: CacheIndexKind::Tree, cache_ttl_secs: 180, cache_boundaries: Vec::new(), + selection_policy: None, + selection_accounting_ttl_ms: 0, }; self } @@ -263,6 +266,32 @@ impl RouterConfigBuilder { self } + pub fn worker_stall_secs(mut self, secs: u64) -> Self { + self.config.worker_stall_secs = secs; + self + } + + pub fn worker_wedge_secs(mut self, secs: u64) -> Self { + self.config.worker_wedge_secs = secs; + self + } + + pub fn worker_warmup( + mut self, + secs: u64, + share: f32, + blocks: usize, + thin_ratio: f32, + divert_every: u64, + ) -> Self { + self.config.worker_warmup_secs = secs; + self.config.worker_warmup_share = share; + self.config.worker_warmup_blocks = blocks; + self.config.worker_warmup_thin_ratio = thin_ratio; + self.config.worker_warmup_divert_every = divert_every; + self + } + pub fn pd_admission_wait_secs(mut self, secs: u64) -> Self { self.config.pd_admission_wait_secs = secs; self @@ -283,6 +312,11 @@ impl RouterConfigBuilder { self } + pub fn worker_overload_shed(mut self, shed: bool) -> Self { + self.config.worker_overload_shed = shed; + self + } + pub fn worker_overload_token_usage(mut self, threshold: Option) -> Self { self.config.worker_overload_token_usage = threshold; self @@ -298,6 +332,11 @@ impl RouterConfigBuilder { self } + pub fn kv_index(mut self, kind: KvIndexKind) -> Self { + self.config.kv_index = kind; + self + } + pub fn engine_metrics(mut self, enabled: bool) -> Self { self.config.engine_metrics = enabled; self diff --git a/model_gateway/src/config/types.rs b/model_gateway/src/config/types.rs index 3883d1b71d..38991d55ed 100755 --- a/model_gateway/src/config/types.rs +++ b/model_gateway/src/config/types.rs @@ -1,4 +1,7 @@ -use std::collections::HashMap; +use std::{ + collections::HashMap, + sync::atomic::{AtomicBool, Ordering}, +}; use openai_protocol::worker::HealthCheckConfig as ProtocolHealthCheckConfig; pub use openai_protocol::worker::{MmProcessingMode, TransportMode}; @@ -12,7 +15,10 @@ use super::{validation::ConfigValidator, ConfigResult}; use crate::{ routers::common::pd_admission::DEFAULT_PD_ADMISSION_WAIT_SECS, tenant::DEFAULT_TENANT_HEADER_NAME, - worker::{ConnectionMode, RuntimeType}, + worker::{ + overload::{DEFAULT_TOKEN_USAGE_CEILING, DEFAULT_WAITING_REQUESTS}, + ConnectionMode, RuntimeType, + }, }; /// Main router configuration @@ -96,6 +102,42 @@ pub struct RouterConfig { pub job_queue_concurrency: usize, #[serde(default = "default_load_monitor_interval_secs")] pub load_monitor_interval_secs: u64, + /// Seconds without any contact from a worker (a load poll, a health probe, + /// a KV event, a response) after which a transport failure excludes it + /// from routing; the first successful contact re-admits it. + #[serde(default = "default_worker_stall_secs")] + pub worker_stall_secs: u64, + /// Seconds without a token or a completion from a worker that still + /// answers polls, with requests in flight and a growing queue, after which + /// new requests stop being routed to it until it makes progress. + #[serde(default = "default_worker_wedge_secs")] + pub worker_wedge_secs: u64, + /// Warm-up slice for cache-aware routing: for this many seconds after a + /// worker becomes routable, until its index has grown by + /// `worker_warmup_blocks` blocks, one cache miss in `1 / share` goes to it. + #[serde(default = "default_worker_warmup_secs")] + pub worker_warmup_secs: u64, + #[serde(default = "default_worker_warmup_share")] + pub worker_warmup_share: f32, + #[serde(default = "default_worker_warmup_blocks")] + pub worker_warmup_blocks: usize, + /// A worker whose index holds less than this share of the fleet's level + /// (the median over healthy workers), or nothing, is thin and receives + /// the warm-up slice until it has grown by `worker_warmup_blocks`, + /// whatever emptied it (a resync after a publisher restart, an + /// out-of-range or data-loss resubscription, an engine that came back + /// empty). 0 keeps the age rule alone. + #[serde(default = "default_worker_warmup_thin_ratio")] + pub worker_warmup_thin_ratio: f32, + /// One cache hit in this many is diverted to a thin worker although + /// another worker holds its prefix, so an index emptied by a resync + /// refills on a workload where every request has a holder; shallow + /// overlaps first, never while the thin worker has a request in flight + /// (unless its last diversion is older than two seconds), until its + /// index crosses `worker_warmup_thin_ratio` of the fleet's level. 0 + /// disables. + #[serde(default = "default_worker_warmup_divert_every")] + pub worker_warmup_divert_every: u64, /// How long a disaggregated (PD) dispatch waits for a slot in the decode /// engine's running window before shedding. Must stay well under the /// engine's bootstrap deadline (120s on TokenSpeed): a request that waits @@ -111,28 +153,34 @@ pub struct RouterConfig { /// always fed regardless of this flag. #[serde(default)] pub disable_load_monitoring: bool, - /// Enable absolute worker overload protection with the gateway default of - /// `worker_overload_token_usage = 0.9` (KV token usage is engine-universal; - /// a waiting-requests default would be workload-dependent, so that signal - /// stays unset). Redundant when either explicit threshold below is set — - /// those enable protection on their own, exactly as before this flag. - #[serde(default)] + /// Absolute worker overload protection, on by default. A worker whose load + /// report is at or above either threshold below is left out of routing + /// while another worker is under them; when every worker is over them the + /// request goes to the least-loaded one (see `worker_overload_shed` for + /// refusing instead). Evaluated once per ingested load report, never per + /// request. `false` switches both gateway thresholds off; per-worker + /// `overload` blocks on a WorkerSpec still apply. + #[serde(default = "default_worker_overload_protection")] pub worker_overload_protection: bool, - /// Queued-request count at or above which a worker is considered - /// overloaded and excluded from routing until the signal recovers; when all - /// workers are overloaded, requests are shed immediately rather than - /// queued. Evaluated once per ingested load report, never per request. - /// `None` (default) disables this signal. - #[serde(default, skip_serializing_if = "Option::is_none")] + /// Queued (waiting) requests, summed across DP ranks, at or above which a + /// worker counts as overloaded. Default 8; `null` switches this signal off. + #[serde(default = "default_worker_overload_waiting_requests")] pub worker_overload_waiting_requests: Option, /// KV-cache token usage (0.0-1.0, averaged across DP ranks) at or above - /// which a worker is considered overloaded — the same signal + /// which a worker counts as overloaded — the same signal /// `balance_token_usage_threshold` reads, applied as an absolute per-worker - /// ceiling instead of a fleet-relative spread. `None` (default) disables - /// this signal; with both signals unset, overload protection is off and - /// routing behaves exactly as before. - #[serde(default, skip_serializing_if = "Option::is_none")] + /// ceiling instead of a fleet-relative spread. Default 0.8; `null` + /// switches this signal off. + #[serde(default = "default_worker_overload_token_usage")] pub worker_overload_token_usage: Option, + /// Refuse a request with a 503 (`worker_overload_protection_shed`, + /// Retry-After the poll interval) when every worker it could use is + /// overloaded, instead of steering it to the least-loaded one. Off by + /// default: a fleet that is uniformly over the thresholds is still a + /// fleet, and the dispatch-time re-check that sheds a worker flagged + /// between selection and dispatch is part of the same opt-in. + #[serde(default)] + pub worker_overload_shed: bool, /// TTL in seconds for entries in the event-driven cache-aware positional /// indexer: entries neither stored to nor read by a query within this /// window are evicted by a periodic background prune. Bounds index growth @@ -145,9 +193,16 @@ pub struct RouterConfig { /// to 90% of the ceiling. `None`/`0` disables the ceiling (default). #[serde(default, skip_serializing_if = "Option::is_none")] pub kv_indexer_max_entries: Option, + /// Which event-driven KV index cache-aware routing reads: the positional + /// indexer (default) or the chain index. The prune bounds above + /// apply to the positional indexer only. + #[serde(default)] + pub kv_index: KvIndexKind, /// Force `GetLoads` polling for `smg_engine_*` gauges even when no /// load-aware routing policy is active. Successful routing-owned polls are - /// always re-exported without an additional Engine RPC. + /// always re-exported without an additional Engine RPC. A worker whose + /// KV-event stream pushes its load feeds the gauges from those records + /// and is not polled while they flow; the poll is its fallback. #[serde(default)] pub engine_metrics: bool, /// Global multimodal tensor transport mode (`inline` | `shm` | `auto` | `rdma`). @@ -395,6 +450,55 @@ fn default_load_monitor_interval_secs() -> u64 { 10 } +fn default_worker_overload_protection() -> bool { + true +} + +// A serde default returns the field's type, `Option` included. +#[expect( + clippy::unnecessary_wraps, + reason = "serde default for an optional field" +)] +fn default_worker_overload_waiting_requests() -> Option { + Some(DEFAULT_WAITING_REQUESTS) +} + +#[expect( + clippy::unnecessary_wraps, + reason = "serde default for an optional field" +)] +fn default_worker_overload_token_usage() -> Option { + Some(DEFAULT_TOKEN_USAGE_CEILING) +} + +fn default_worker_stall_secs() -> u64 { + 2 +} + +fn default_worker_wedge_secs() -> u64 { + 3 +} + +fn default_worker_warmup_secs() -> u64 { + 60 +} + +fn default_worker_warmup_share() -> f32 { + 0.25 +} + +fn default_worker_warmup_blocks() -> usize { + 1024 +} + +fn default_worker_warmup_thin_ratio() -> f32 { + 0.5 +} + +fn default_worker_warmup_divert_every() -> u64 { + 8 +} + fn default_pd_admission_wait_secs() -> u64 { DEFAULT_PD_ADMISSION_WAIT_SECS } @@ -690,6 +794,55 @@ impl Default for RoutingKeyOverrideConfig { } } +/// The event-driven KV index behind cache-aware routing (`--kv-index`). +#[derive(Debug, Clone, Copy, Serialize, Deserialize, Default, PartialEq, Eq)] +#[serde(rename_all = "snake_case")] +pub enum KvIndexKind { + /// The positional indexer: one entry per `(position, content hash)`, + /// probed per block of a request (default). + #[default] + Positional, + /// The chain index: chains stored as runs with per-run worker coverage; + /// lock-free, store-free lookups, memory proportional to the blocks the + /// engines report. `run`, its name before 2026-10-06, is accepted as a + /// deprecated alias until the positional indexer is removed. + #[serde(alias = "run")] + Chain, +} + +/// Whether the deprecated `run` spelling of the chain index was given on the +/// command line; read once at startup to log the deprecation. +static DEPRECATED_KV_INDEX_ALIAS: AtomicBool = AtomicBool::new(false); + +impl KvIndexKind { + /// Parse from a case-insensitive string (`positional` | `chain`, with + /// `run` as the deprecated spelling of `chain`). + pub fn parse(value: &str) -> Option { + match value.trim().to_ascii_lowercase().as_str() { + "positional" => Some(Self::Positional), + "chain" => Some(Self::Chain), + "run" => { + DEPRECATED_KV_INDEX_ALIAS.store(true, Ordering::Relaxed); + Some(Self::Chain) + } + _ => None, + } + } + + /// Whether `parse` was given the deprecated `run` spelling. + pub fn deprecated_alias_used() -> bool { + DEPRECATED_KV_INDEX_ALIAS.load(Ordering::Relaxed) + } + + /// Canonical lowercase name. + pub fn as_str(self) -> &'static str { + match self { + Self::Positional => "positional", + Self::Chain => "chain", + } + } +} + /// Under-layer index the cache_aware policy keeps per model. #[derive(Debug, Clone, Copy, Serialize, Deserialize, Default, PartialEq, Eq)] #[serde(rename_all = "snake_case")] @@ -767,6 +920,15 @@ pub enum PolicyConfig { /// shared `cache_boundaries` config). #[serde(default, skip_serializing_if = "Vec::is_empty")] cache_boundaries: Vec, + /// Worker selection policy run over the gathered per-worker inputs + /// (`cache-aware-default`). Unset is the cache-aware default. + #[serde(default, skip_serializing_if = "Option::is_none")] + selection_policy: Option, + /// Lifetime in milliseconds of optimistic dispatch bookings + /// (predicted prefill and prefix placement charged to the chosen + /// worker before the engine reports it). `0` disables. + #[serde(default)] + selection_accounting_ttl_ms: u64, }, /// Power-of-two choices policy: samples two workers and routes to the one @@ -1255,13 +1417,22 @@ impl Default for RouterConfig { job_queue_capacity: default_job_queue_capacity(), job_queue_concurrency: default_job_queue_concurrency(), load_monitor_interval_secs: 10, + worker_stall_secs: default_worker_stall_secs(), + worker_wedge_secs: default_worker_wedge_secs(), + worker_warmup_secs: default_worker_warmup_secs(), + worker_warmup_share: default_worker_warmup_share(), + worker_warmup_blocks: default_worker_warmup_blocks(), + worker_warmup_thin_ratio: default_worker_warmup_thin_ratio(), + worker_warmup_divert_every: default_worker_warmup_divert_every(), pd_admission_wait_secs: default_pd_admission_wait_secs(), disable_load_monitoring: false, - worker_overload_protection: false, - worker_overload_waiting_requests: None, - worker_overload_token_usage: None, + worker_overload_protection: default_worker_overload_protection(), + worker_overload_waiting_requests: default_worker_overload_waiting_requests(), + worker_overload_token_usage: default_worker_overload_token_usage(), + worker_overload_shed: false, kv_indexer_ttl_secs: None, kv_indexer_max_entries: None, + kv_index: KvIndexKind::default(), engine_metrics: false, multimodal_tensor_transport: None, multimodal_shm_min_bytes: None, @@ -1493,6 +1664,44 @@ mod tests { assert!(deserialized.trace_config.is_none()); } + /// A switched-off overload signal (`null`) must come back off from a + /// serialize/deserialize round-trip. The field's default is `Some`, so + /// skipping `None` on serialize would hand the default back on + /// deserialize and turn the signal on again. + #[test] + fn disabled_overload_signals_survive_a_serde_round_trip() { + let config = RouterConfig { + worker_overload_waiting_requests: None, + worker_overload_token_usage: None, + ..Default::default() + }; + let json = serde_json::to_string(&config).unwrap(); + let value: serde_json::Value = serde_json::from_str(&json).unwrap(); + assert!( + value["worker_overload_waiting_requests"].is_null(), + "a disabled signal is written as null, not left out" + ); + assert!(value["worker_overload_token_usage"].is_null()); + + let back: RouterConfig = serde_json::from_str(&json).unwrap(); + assert_eq!(back.worker_overload_waiting_requests, None); + assert_eq!(back.worker_overload_token_usage, None); + + // A document that does not mention the signals still gets the defaults. + let mut without = value; + without + .as_object_mut() + .unwrap() + .remove("worker_overload_waiting_requests"); + without + .as_object_mut() + .unwrap() + .remove("worker_overload_token_usage"); + let defaulted: RouterConfig = serde_json::from_value(without).unwrap(); + assert_eq!(defaulted.worker_overload_waiting_requests, Some(8)); + assert_eq!(defaulted.worker_overload_token_usage, Some(0.8)); + } + #[test] fn test_health_check_port_serde_roundtrip_and_backward_compat() { // Default: dedicated probe listener off, and `skip_serializing_if` @@ -1883,6 +2092,8 @@ mod tests { cache_index: Default::default(), cache_ttl_secs: 180, cache_boundaries: Vec::new(), + selection_policy: None, + selection_accounting_ttl_ms: 0, }; assert_eq!(cache_aware.name(), "cache_aware"); @@ -1912,6 +2123,8 @@ mod tests { cache_index: Default::default(), cache_ttl_secs: 180, cache_boundaries: Vec::new(), + selection_policy: None, + selection_accounting_ttl_ms: 0, }; let json = serde_json::to_string(&cache_aware).unwrap(); assert!(json.contains("\"type\":\"cache_aware\"")); @@ -1942,6 +2155,8 @@ mod tests { cache_index: Default::default(), cache_ttl_secs: 180, cache_boundaries: Vec::new(), + selection_policy: None, + selection_accounting_ttl_ms: 0, }; match cache_aware { @@ -2585,6 +2800,8 @@ discovery: cache_index: Default::default(), cache_ttl_secs: 180, cache_boundaries: Vec::new(), + selection_policy: None, + selection_accounting_ttl_ms: 0, }), decode_policy: Some(PolicyConfig::PowerOfTwo { load_check_interval_secs: 60, @@ -2623,6 +2840,8 @@ discovery: cache_index: Default::default(), cache_ttl_secs: 180, cache_boundaries: Vec::new(), + selection_policy: None, + selection_accounting_ttl_ms: 0, }), decode_policy: None, }; @@ -2687,6 +2906,8 @@ discovery: cache_index: Default::default(), cache_ttl_secs: 180, cache_boundaries: Vec::new(), + selection_policy: None, + selection_accounting_ttl_ms: 0, }; match pd.get_prefill_policy(&main_policy) { diff --git a/model_gateway/src/config/validation.rs b/model_gateway/src/config/validation.rs index 554c28fc80..0fc1760fe2 100644 --- a/model_gateway/src/config/validation.rs +++ b/model_gateway/src/config/validation.rs @@ -2,6 +2,7 @@ use axum::http::HeaderName; use sha2::{Digest, Sha256}; use super::*; +use crate::policies::cost as selection_cost; /// Validate a user-supplied mesh server name. The name keys rate-limit /// shards as `rl:{counter}:{name}`, so an empty name or one containing the @@ -515,9 +516,26 @@ impl ConfigValidator { cache_index, cache_ttl_secs, cache_boundaries, + selection_policy, + selection_accounting_ttl_ms: _, } => { Self::validate_cache_boundaries(cache_boundaries)?; + // Build the selection policy once here so a bad name or + // parameter fails configuration instead of routing. + let selection_policy_name = selection_policy + .as_deref() + .unwrap_or(selection_cost::DEFAULT_POLICY); + if let Err(err) = + selection_cost::build(selection_policy_name, *selection_temperature) + { + return Err(ConfigError::InvalidValue { + field: "selection_policy".to_string(), + value: selection_policy_name.to_string(), + reason: err.to_string(), + }); + } + if *cache_ttl_secs == 0 { return Err(ConfigError::InvalidValue { field: "cache_ttl_secs".to_string(), @@ -1734,6 +1752,8 @@ mod tests { cache_index: Default::default(), cache_ttl_secs: 180, cache_boundaries: Vec::new(), + selection_policy: None, + selection_accounting_ttl_ms: 0, }, ); @@ -1764,6 +1784,8 @@ mod tests { cache_index: Default::default(), cache_ttl_secs: 180, cache_boundaries: Vec::new(), + selection_policy: None, + selection_accounting_ttl_ms: 0, }, ) }; @@ -1799,6 +1821,8 @@ mod tests { cache_index, cache_ttl_secs, cache_boundaries: boundaries, + selection_policy: None, + selection_accounting_ttl_ms: 0, }, ) }; @@ -1851,6 +1875,8 @@ mod tests { cache_index: Default::default(), cache_ttl_secs: 180, cache_boundaries: Vec::new(), + selection_policy: None, + selection_accounting_ttl_ms: 0, }, ); @@ -1972,6 +1998,8 @@ mod tests { cache_index: Default::default(), cache_ttl_secs: 180, cache_boundaries: Vec::new(), + selection_policy: None, + selection_accounting_ttl_ms: 0, }, ); @@ -2024,6 +2052,8 @@ mod tests { cache_index: Default::default(), cache_ttl_secs: 180, cache_boundaries: Vec::new(), + selection_policy: None, + selection_accounting_ttl_ms: 0, }), decode_policy: Some(PolicyConfig::PowerOfTwo { load_check_interval_secs: 60, @@ -2155,6 +2185,8 @@ mod tests { cache_index: Default::default(), cache_ttl_secs: 180, cache_boundaries: Vec::new(), + selection_policy: None, + selection_accounting_ttl_ms: 0, }), prefill_policy: None, decode_policy: None, diff --git a/model_gateway/src/main.rs b/model_gateway/src/main.rs index 1c398ba9c8..45eb644b80 100644 --- a/model_gateway/src/main.rs +++ b/model_gateway/src/main.rs @@ -8,22 +8,42 @@ use clap::{ArgAction, Parser, Subcommand, ValueEnum}; #[cfg(all(not(target_env = "msvc"), not(target_env = "musl")))] #[global_allocator] static GLOBAL_ALLOCATOR: tikv_jemallocator::Jemalloc = tikv_jemallocator::Jemalloc; + +// Jemalloc's run-time options for a long-running server. With the stock +// settings a thread's freed pages decay back to the OS only when that thread +// allocates again, so after a traffic burst an idle gateway kept ~1.5 GB +// resident over ~270 MB of live objects (soak s2, ten hours). A background +// thread purges on schedule instead; dirty pages are returned after 10 s and +// muzzy pages at once. This is jemalloc's application-provided `malloc_conf` +// string under the vendored build's `_rjem_` prefix; the `_RJEM_MALLOC_CONF` +// environment variable is read after it and overrides it entry by entry. +#[cfg(all(not(target_env = "msvc"), not(target_env = "musl")))] +#[expect( + unsafe_code, + reason = "jemalloc reads its options from this exported symbol; a NUL-terminated byte string nothing in Rust dereferences" +)] +#[export_name = "_rjem_malloc_conf"] +pub static MALLOC_CONF: &[u8; 61] = + b"background_thread:true,dirty_decay_ms:10000,muzzy_decay_ms:0\0"; + use openai_protocol::worker::{MmProcessingMode, TransportMode}; use rand::{distr::Alphanumeric, RngExt}; use smg::{ config::{ resolve_worker_auto_recovery, validate_mesh_server_name, CacheIndexKind, CircuitBreakerConfig, ConfigError, ConfigResult, DiscoveryConfig, HealthCheckConfig, - HistoryBackend, KubernetesDiscoveryConfig, ManualAssignmentMode, MetricsConfig, - OracleConfig, PdPairingMode, PolicyConfig, PostgresConfig, RedisConfig, RetryConfig, - RouterConfig, RoutingKeyOverrideConfig, RoutingMode, SchemaConfig, TenantApiKeyEntry, - TokenizerCacheConfig, TraceConfig, + HistoryBackend, KubernetesDiscoveryConfig, KvIndexKind, ManualAssignmentMode, + MetricsConfig, OracleConfig, PdPairingMode, PolicyConfig, PostgresConfig, RedisConfig, + RetryConfig, RouterConfig, RoutingKeyOverrideConfig, RoutingMode, SchemaConfig, + TenantApiKeyEntry, TokenizerCacheConfig, TraceConfig, }, mesh_discovery::MeshDiscoveryConfig, observability::{ + logging::close_logging, metrics::{register_jemalloc_as_global_allocator, PrometheusConfig}, otel_trace::{is_otel_enabled, shutdown_otel}, }, + policies::cost::DEFAULT_POLICY as DEFAULT_SELECTION_POLICY, server::{self, ServerConfig}, service_discovery::{ModelIdSource, RuntimeDiscoveryConfig}, version, @@ -173,6 +193,12 @@ fn parse_transport_mode(value: &str) -> Result { .ok_or_else(|| format!("invalid value '{value}'; expected inline, shm, auto, or rdma")) } +/// Parse the `--kv-index` value into a `KvIndexKind`. +fn parse_kv_index_kind(value: &str) -> Result { + KvIndexKind::parse(value) + .ok_or_else(|| format!("invalid value '{value}'; expected positional or chain")) +} + /// Parse the `--mm-processing` value into an `MmProcessingMode`. fn parse_mm_processing(value: &str) -> Result { MmProcessingMode::parse(value) @@ -293,50 +319,46 @@ struct CliArgs { #[arg(long, default_value_t = 1.0, help_heading = "Routing Policy")] overload_token_usage_threshold: f32, - /// Enable worker overload protection with the gateway default thresholds. - /// - /// A worker whose load signal crosses a threshold is considered overloaded - /// and excluded from routing until the signal recovers; when every worker - /// is overloaded, requests are shed immediately rather than queued. - /// - /// This flag alone applies --worker-overload-token-usage 0.9 and leaves - /// --worker-overload-waiting-requests unset: KV token usage means the same - /// thing on every engine, while a sensible waiting-requests ceiling is - /// workload-dependent, so it has no universal default. Explicit thresholds - /// override the default, and either threshold set on its own enables - /// protection without this flag — exactly as before it existed. Per-worker - /// `overload` blocks on a WorkerSpec override the gateway values per - /// signal, and enable protection for that worker even with everything here - /// unset. - #[arg(long, default_value_t = false, help_heading = "Routing Policy")] + /// Worker overload protection (the default; kept so existing command + /// lines still parse). A worker whose load report is at or above + /// --worker-overload-waiting-requests or --worker-overload-token-usage is + /// left out of routing while another worker is under them; when every + /// worker is over them the request goes to the least-loaded one, unless + /// --worker-overload-shed asks for a 503 instead. Evaluated once per load + /// report, never per request. Per-worker `overload` blocks on a WorkerSpec + /// override the gateway values per signal + #[arg(long, default_value_t = true, help_heading = "Routing Policy")] worker_overload_protection: bool, - /// Queued-request count at or above which a worker is considered - /// overloaded and excluded from routing until the signal recovers; when - /// every worker is overloaded, requests are shed immediately rather than - /// queued. Unset disables overload protection. - /// - /// Queued (waiting) requests, summed across DP ranks. Must be >= 1: the - /// comparison is inclusive, so 0 would veto every worker unconditionally. - #[arg(long, value_parser = parse_positive_usize, help_heading = "Routing Policy")] - worker_overload_waiting_requests: Option, - - /// KV-cache token usage at or above which a worker is considered - /// overloaded and excluded from routing until the signal recovers; when - /// every worker is overloaded, requests are shed immediately rather than - /// queued. Unset disables overload protection. - /// - /// Mean KV-cache token usage across DP ranks, the same signal - /// `--balance-token-usage-threshold` reads, applied as an absolute - /// per-worker ceiling rather than a fleet-relative spread. Backend must - /// report token_usage. Must be in (0.0, 1.0]: the comparison is inclusive, - /// so 0.0 would veto every worker unconditionally. - /// - /// Distinct from `--overload-token-usage-threshold`, which only de-ranks - /// the hottest backend within cache-aware affinity; this flag removes the - /// worker from routing entirely and sheds when every worker crosses it. - #[arg(long, value_parser = parse_unit_fraction, help_heading = "Routing Policy")] - worker_overload_token_usage: Option, + /// Switch worker overload protection off: no worker is ever left out of + /// routing for its waiting queue or KV usage (per-worker `overload` blocks + /// on a WorkerSpec still apply) + #[arg(long, default_value_t = false, help_heading = "Routing Policy")] + disable_worker_overload_protection: bool, + + /// Queued (waiting) requests, summed across DP ranks, at or above which a + /// worker counts as overloaded. Must be >= 1: the comparison is inclusive, + /// so 0 would veto every worker unconditionally + #[arg(long, default_value_t = 8, value_parser = parse_positive_usize, help_heading = "Routing Policy")] + worker_overload_waiting_requests: usize, + + /// Mean KV-cache token usage across DP ranks at or above which a worker + /// counts as overloaded: the same signal --balance-token-usage-threshold + /// reads, applied as an absolute per-worker ceiling rather than a + /// fleet-relative spread. Backend must report token_usage. Must be in + /// (0.0, 1.0]: the comparison is inclusive, so 0.0 would veto every worker + /// unconditionally. Distinct from --overload-token-usage-threshold, which + /// only de-ranks the hottest backend within cache-aware affinity + #[arg(long, default_value_t = 0.8, value_parser = parse_unit_fraction, help_heading = "Routing Policy")] + worker_overload_token_usage: f64, + + /// Refuse a request with a 503 (worker_overload_protection_shed, + /// Retry-After the load poll interval) when every worker it could use is + /// overloaded, instead of routing it to the least-loaded one; also sheds a + /// worker that crossed a threshold between selection and dispatch. Off by + /// default: steering never turns a load signal into an outage + #[arg(long, default_value_t = false, help_heading = "Routing Policy")] + worker_overload_shed: bool, /// Anti-hotspot decay: de-rank cache-affine candidates by their /// waiting-prefill backlog (overlap score divided by 1 + overlap_decay @@ -351,6 +373,24 @@ struct CliArgs { #[arg(long, default_value_t = 0.0, help_heading = "Routing Policy")] selection_temperature: f32, + /// Worker selection policy for cache_aware, run over the per-worker + /// inputs the router gathers (prefix overlap, in-flight requests, + /// backend load reports); cache-aware-default, the affinity-group + /// decision, is the one policy + #[arg( + long, + default_value = "cache-aware-default", + help_heading = "Routing Policy" + )] + selection_policy: String, + + /// Lifetime in milliseconds of optimistic dispatch bookings for + /// cache_aware: predicted prefill and prefix placement are charged to the + /// chosen worker until the engine reports them or the booking expires. + /// 0 disables; set a little above the engine's KV-event lag + #[arg(long, default_value_t = 0, help_heading = "Routing Policy")] + selection_accounting_ttl_ms: u64, + /// Interval in seconds between cache-tree eviction cycles #[arg(long, default_value_t = 120, help_heading = "Routing Policy")] eviction_interval: u64, @@ -577,6 +617,54 @@ struct CliArgs { #[arg(long, default_value_t = 10, help_heading = "Load Monitoring")] load_monitor_interval: u64, + /// Seconds without any contact from a worker (a load poll, a health probe, + /// a KV event, a response) after which a transport failure excludes it + /// from routing; the first successful contact re-admits it. + #[arg(long, default_value_t = 2, help_heading = "Load Monitoring")] + worker_stall_secs: u64, + + /// Seconds without a token or a completion from a worker with requests + /// in flight whose waiting queue grows, or whose in-flight pile grows or + /// is four deep, after which new requests stop being routed to it until + /// it makes progress. The bound stretches to the time its in-flight + /// prompts may still need in prefill, up to 120 seconds. The pile is the + /// streaming gRPC generations in flight; HTTP workers, PD legs and + /// non-streaming generations give no signal and never form one. 0 + /// disables the rule. + #[arg(long, default_value_t = 3, help_heading = "Load Monitoring")] + worker_wedge_secs: u64, + + /// Warm-up slice for cache-aware routing: for this many seconds after a + /// worker becomes routable, until its index has grown by + /// --worker-warmup-blocks blocks, one cache miss in 1/--worker-warmup-share + /// is routed to it so it builds a cache instead of idling. 0 disables. + #[arg(long, default_value_t = 60, help_heading = "Routing Policy")] + worker_warmup_secs: u64, + + /// Share of cache misses offered to warming workers (0.0 to 1.0). + #[arg(long, default_value_t = 0.25, help_heading = "Routing Policy")] + worker_warmup_share: f32, + + /// A worker whose index has grown by this many blocks since it became + /// thin is warm. + #[arg(long, default_value_t = 1024, help_heading = "Routing Policy")] + worker_warmup_blocks: usize, + + /// A worker whose index holds less than this share of the fleet's median + /// (or nothing) is thin and receives the warm-up slice until it has grown + /// by --worker-warmup-blocks, whatever emptied it (a resync after a + /// publisher restart, an out-of-range or data-loss resubscription, an + /// engine that came back empty). 0 keeps the age rule alone. + #[arg(long, default_value_t = 0.5, help_heading = "Routing Policy")] + worker_warmup_thin_ratio: f32, + + /// One cache hit in this many is diverted to a thin worker although + /// another worker holds its prefix (shallow overlaps first, one in flight + /// per thin worker), so an index emptied by a resync refills on a + /// workload where every request has a holder. 0 disables. + #[arg(long, default_value_t = 8, help_heading = "Routing Policy")] + worker_warmup_divert_every: u64, + /// Only poll worker loads when a load-aware routing policy, /// --engine-metrics, or worker overload protection needs the data. By /// default every worker group is polled from registration onward; this @@ -587,6 +675,9 @@ struct CliArgs { /// Force GetLoads polling for smg_engine_* Prometheus gauges even without /// a load-aware routing policy. Routing-owned polls are always re-exported. + /// A worker whose KV-event stream pushes its load feeds the gauges from + /// those records and is not polled while they flow (GetLoads is the + /// fallback). #[arg(long, default_value_t = false, help_heading = "Load Monitoring")] engine_metrics: bool, @@ -613,6 +704,15 @@ struct CliArgs { #[arg(long, help_heading = "Routing Policy")] kv_indexer_max_entries: Option, + /// The event-driven KV index behind cache-aware routing: `positional` + /// (one entry per block position, the default) or `chain` (chains as runs + /// with per-run worker coverage; lock-free, store-free lookups; `run` is + /// its deprecated spelling). The --kv-indexer-* prune bounds apply to the + /// positional index only; the chain index holds what the engines report + /// and shrinks with their removals. + #[arg(long, default_value = "positional", value_parser = parse_kv_index_kind, help_heading = "Routing Policy")] + kv_index: KvIndexKind, + /// Multimodal tensor transport mode: `inline` (default), `shm` (same-host /// /dev/shm), or `auto` (shm only when the worker shares /dev/shm). A /// per-worker `WorkerSpec.multimodal_tensor_transport` overrides this. @@ -1560,6 +1660,9 @@ impl CliArgs { cache_index: Self::parse_cache_index(&self.cache_index), cache_ttl_secs: self.cache_ttl_secs, cache_boundaries: self.cache_boundaries.clone(), + selection_policy: (self.selection_policy != DEFAULT_SELECTION_POLICY) + .then(|| self.selection_policy.clone()), + selection_accounting_ttl_ms: self.selection_accounting_ttl_ms, }, "power_of_two" => PolicyConfig::PowerOfTwo { load_check_interval_secs: 5, @@ -1967,13 +2070,26 @@ impl CliArgs { .job_queue_capacity(self.job_queue_capacity) .job_queue_concurrency(self.job_queue_concurrency) .load_monitor_interval_secs(self.load_monitor_interval) + .worker_stall_secs(self.worker_stall_secs) + .worker_wedge_secs(self.worker_wedge_secs) + .worker_warmup( + self.worker_warmup_secs, + self.worker_warmup_share, + self.worker_warmup_blocks, + self.worker_warmup_thin_ratio, + self.worker_warmup_divert_every, + ) .pd_admission_wait_secs(self.pd_admission_wait_secs) .disable_load_monitoring(self.disable_load_monitoring) - .worker_overload_protection(self.worker_overload_protection) - .worker_overload_waiting_requests(self.worker_overload_waiting_requests) - .worker_overload_token_usage(self.worker_overload_token_usage) + .worker_overload_protection( + self.worker_overload_protection && !self.disable_worker_overload_protection, + ) + .worker_overload_waiting_requests(Some(self.worker_overload_waiting_requests)) + .worker_overload_token_usage(Some(self.worker_overload_token_usage)) + .worker_overload_shed(self.worker_overload_shed) .kv_indexer_ttl_secs(self.kv_indexer_ttl_secs) .kv_indexer_max_entries(self.kv_indexer_max_entries) + .kv_index(self.kv_index) .engine_metrics(self.engine_metrics) .multimodal_tensor_transport(self.multimodal_tensor_transport) .multimodal_shm_min_bytes(self.multimodal_shm_min_bytes) @@ -2261,7 +2377,11 @@ fn main() -> Result<(), Box> { tokio::runtime::Runtime::new()? } }; - runtime.block_on(Box::pin(server::startup(server_config)))?; + let served = runtime.block_on(Box::pin(server::startup(server_config))); + // The writer threads outlive `startup` (a process may start more than one + // router); flush and stop them now that this process is done logging. + close_logging(); + served?; if is_otel_enabled() { shutdown_otel(); } @@ -2319,6 +2439,33 @@ mod tests { assert_eq!(defaults.kv_indexer_max_entries, None); } + /// `--kv-index` selects the event-driven index for cache-aware routing; + /// the positional indexer stays the default until the chain index has + /// passed a soak in the gateway. + #[test] + fn kv_index_flag_flows_into_router_config() { + let defaults = cli_args_from(&[]).to_router_config(vec![], vec![]).unwrap(); + assert_eq!(defaults.kv_index, KvIndexKind::Positional); + + let cli = cli_args_from(&["--kv-index", "chain"]); + let router_config = cli.to_router_config(vec![], vec![]).unwrap(); + assert_eq!(router_config.kv_index, KvIndexKind::Chain); + let server_config = cli.to_server_config(router_config).unwrap(); + assert_eq!(server_config.router_config.kv_index, KvIndexKind::Chain); + // The spelling before the rename still selects the chain index. + assert_eq!( + cli_args_from(&["--kv-index", "run"]).kv_index, + KvIndexKind::Chain + ); + assert!(KvIndexKind::deprecated_alias_used()); + + assert_eq!( + cli_args_from(&["--kv-index", "Positional"]).kv_index, + KvIndexKind::Positional + ); + assert!(parse_kv_index_kind("tree").is_err()); + } + /// The retry-buffer cap must flow through both conversion paths. #[test] fn max_buffered_request_bytes_flows_into_both_config_paths() { @@ -2835,25 +2982,48 @@ mod tests { ); } - /// Unset means off on both paths: the feature must be byte-identical to - /// pre-feature behavior until an operator opts in. + /// Protection is on by default on both signals, steering only, on both + /// config paths. #[test] - fn worker_overload_thresholds_default_to_unset_in_both_configs() { + fn worker_overload_protection_defaults_to_on_and_steering_in_both_configs() { let cli = cli_args_from(&[]); let router_config = cli.to_router_config(vec![], vec![]).unwrap(); - assert_eq!(router_config.worker_overload_waiting_requests, None); - assert_eq!(router_config.worker_overload_token_usage, None); + assert!(router_config.worker_overload_protection); + assert_eq!(router_config.worker_overload_waiting_requests, Some(8)); + assert_eq!(router_config.worker_overload_token_usage, Some(0.8)); + assert!(!router_config.worker_overload_shed); let server_config = cli.to_server_config(router_config).unwrap(); + assert!(server_config.router_config.worker_overload_protection); assert_eq!( server_config.router_config.worker_overload_waiting_requests, - None + Some(8) ); assert_eq!( server_config.router_config.worker_overload_token_usage, - None + Some(0.8) ); + assert!(!server_config.router_config.worker_overload_shed); + } + + /// The two opt-outs: `--disable-worker-overload-protection` switches the + /// gateway thresholds off, `--worker-overload-shed` turns steering into + /// refusal; both must survive into `ServerConfig.router_config`. + #[test] + fn worker_overload_opt_outs_flow_into_both_configs() { + let cli = cli_args_from(&["--disable-worker-overload-protection"]); + let router_config = cli.to_router_config(vec![], vec![]).unwrap(); + assert!(!router_config.worker_overload_protection); + let server_config = cli.to_server_config(router_config).unwrap(); + assert!(!server_config.router_config.worker_overload_protection); + + let cli = cli_args_from(&["--worker-overload-shed"]); + let router_config = cli.to_router_config(vec![], vec![]).unwrap(); + assert!(router_config.worker_overload_shed); + assert!(router_config.worker_overload_protection); + let server_config = cli.to_server_config(router_config).unwrap(); + assert!(server_config.router_config.worker_overload_shed); } /// Both thresholds are `>=` comparisons, so the excluded ends of their @@ -2894,10 +3064,9 @@ mod tests { router_config.disable_load_monitoring, "disable_load_monitoring must reach RouterConfig via to_router_config" ); - // The flag alone carries no thresholds; the token default is applied - // at resolution, not stored in config. - assert_eq!(router_config.worker_overload_waiting_requests, None); - assert_eq!(router_config.worker_overload_token_usage, None); + // The flag changes nothing: the defaults are already on. + assert_eq!(router_config.worker_overload_waiting_requests, Some(8)); + assert_eq!(router_config.worker_overload_token_usage, Some(0.8)); let server_config = cli.to_server_config(router_config).unwrap(); assert!( @@ -2910,18 +3079,18 @@ mod tests { ); } - /// Defaults: protection off, monitoring default-on (opt-out false) — the - /// behavior change is monitoring, and it is carried by the default here. + /// Defaults: protection on, monitoring on (opt-out false) — both carried + /// by the defaults here. #[test] - fn overload_protection_and_monitoring_flags_default_off_in_both_configs() { + fn overload_protection_and_monitoring_default_on_in_both_configs() { let cli = cli_args_from(&[]); let router_config = cli.to_router_config(vec![], vec![]).unwrap(); - assert!(!router_config.worker_overload_protection); + assert!(router_config.worker_overload_protection); assert!(!router_config.disable_load_monitoring); let server_config = cli.to_server_config(router_config).unwrap(); - assert!(!server_config.router_config.worker_overload_protection); + assert!(server_config.router_config.worker_overload_protection); assert!(!server_config.router_config.disable_load_monitoring); } diff --git a/model_gateway/src/mesh/wiring.rs b/model_gateway/src/mesh/wiring.rs index 404ff439d2..6bf65b3d75 100644 --- a/model_gateway/src/mesh/wiring.rs +++ b/model_gateway/src/mesh/wiring.rs @@ -475,6 +475,8 @@ mod tests { cache_index: Default::default(), cache_ttl_secs: 180, cache_boundaries: Vec::new(), + selection_policy: None, + selection_accounting_ttl_ms: 0, } } diff --git a/model_gateway/src/observability/logging.rs b/model_gateway/src/observability/logging.rs index 74a5463393..e280f2ce47 100644 --- a/model_gateway/src/observability/logging.rs +++ b/model_gateway/src/observability/logging.rs @@ -1,10 +1,28 @@ -//! Logging infrastructure with non-blocking file I/O. - -use std::path::PathBuf; +//! Logging infrastructure with non-blocking I/O for every sink. +//! +//! Each formatting layer hands its lines to a dedicated writer thread through +//! a bounded, lossy channel ([`tracing_appender::non_blocking`]), so the tokio +//! runtime threads that emit the events never wait for the sink. That matters +//! for stdout as much as for the log file: when stdout is redirected to a +//! regular file, every `write(2)` is subject to the kernel's dirty-page +//! throttling and can sleep for hundreds of milliseconds under writeback +//! pressure, which with a synchronous `std::io::stdout()` writer parked a +//! runtime worker, and every connection it was driving, once per access-log +//! line. + +use std::{ + io::Write, + path::PathBuf, + sync::{ + atomic::{AtomicBool, Ordering}, + Mutex, + }, + time::Duration, +}; -use tracing::Level; +use tracing::{warn, Level}; use tracing_appender::{ - non_blocking::WorkerGuard, + non_blocking::{ErrorCounter, NonBlocking, NonBlockingBuilder, WorkerGuard}, rolling::{RollingFileAppender, Rotation}, }; use tracing_log::{AsLog, LogTracer}; @@ -17,7 +35,7 @@ use tracing_subscriber::{ }; use super::otel_trace::get_otel_layer; -use crate::config::TraceConfig; +use crate::{config::TraceConfig, observability::metrics::Metrics}; const TIME_FORMAT: &str = "%Y-%m-%d %H:%M:%S"; @@ -67,9 +85,145 @@ impl Default for LoggingConfig { } } -/// Guard that keeps the file appender thread alive. -pub struct LogGuard { - _file_guard: Option, +/// What `init_logging` hands back. The writer threads live in [`SINKS`] for +/// the process, not in this value: the subscriber is installed once per +/// process and every router started after the first (the Python binding +/// starts them in one process) logs through the same writers, so a guard +/// that died with the first `startup` call would leave every later line +/// queued behind a thread that had exited. Dropping this stops nothing; +/// [`close_logging`] flushes and stops the writers at process exit. +pub struct LogGuard(()); + +/// One writer thread: its handle, and the count of lines its lossy queue +/// dropped. +struct Sink { + name: &'static str, + counter: ErrorCounter, + /// The count at the last report, for the delta. + reported: usize, + _guard: WorkerGuard, +} + +/// The writer threads of every sink installed in this process. +struct Sinks { + sinks: Vec, + reporter_started: bool, +} + +static SINKS: Mutex = Mutex::new(Sinks { + sinks: Vec::new(), + reporter_started: false, +}); + +/// Set by [`close_logging`]; the reporter thread ends at its next tick. +static CLOSED: AtomicBool = AtomicBool::new(false); + +/// How often the dropped-line counts are published. +const DROPPED_LINES_REPORT_SECS: u64 = 30; + +/// Move `sink` to its own writer thread held for the process, and return +/// the writer for a subscriber layer. The thread is counted among the +/// sinks whose dropped lines the reporter publishes. +fn install(sink: W, name: &'static str) -> NonBlocking { + let (writer, guard) = non_blocking(sink, name); + let mut sinks = SINKS + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()); + sinks.sinks.push(Sink { + name, + counter: writer.error_counter(), + reported: 0, + _guard: guard, + }); + if !sinks.reporter_started { + sinks.reporter_started = std::thread::Builder::new() + .name("smg-log-dropped".to_string()) + .spawn(|| { + let tick = Duration::from_millis(100); + let ticks_per_report = DROPPED_LINES_REPORT_SECS * 10; + loop { + for _ in 0..ticks_per_report { + if CLOSED.load(Ordering::Relaxed) { + return; + } + std::thread::sleep(tick); + } + report_dropped_lines(); + } + }) + .is_ok(); + } + writer +} + +/// Lines each writer's queue has dropped since the process started, by sink. +pub fn dropped_lines() -> Vec<(&'static str, usize)> { + SINKS + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()) + .sinks + .iter() + .map(|sink| (sink.name, sink.counter.dropped_lines())) + .collect() +} + +/// Publish each sink's dropped-line count as `smg_log_dropped_lines{sink}` and +/// warn about the lines dropped since the last report. Returns `(sink, total, +/// since the last report)` per sink. +fn report_dropped_lines() -> Vec<(&'static str, usize, usize)> { + let mut sinks = SINKS + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()); + let mut reports = Vec::with_capacity(sinks.sinks.len()); + for sink in &mut sinks.sinks { + let total = sink.counter.dropped_lines(); + let since_last = total.saturating_sub(sink.reported); + sink.reported = total; + reports.push((sink.name, total, since_last)); + } + drop(sinks); + for &(name, total, since_last) in &reports { + Metrics::set_log_dropped_lines(name, total); + if since_last > 0 { + warn!( + sink = name, + dropped_total = total, + dropped_since_last_report = since_last, + "log lines dropped: the writer queue was full" + ); + } + } + reports +} + +/// Flush and stop the writer threads: the last lines queued are written +/// before the process exits. Called once, at the end of `main`; nothing logs +/// through the stopped writers afterwards. +pub fn close_logging() { + CLOSED.store(true, Ordering::Relaxed); + let sinks = std::mem::take( + &mut SINKS + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()) + .sinks, + ); + drop(sinks); +} + +/// Move `sink` to its own thread behind a bounded, lossy channel. +/// +/// Callers never block on the sink: a line is queued for the writer thread, or +/// counted and dropped when the queue is full. A sink that stalls therefore +/// slows nothing but its own thread. +#[inline] +fn non_blocking( + sink: W, + thread_name: &str, +) -> (NonBlocking, WorkerGuard) { + NonBlockingBuilder::default() + .lossy(true) + .thread_name(thread_name) + .finish(sink) } #[inline] @@ -207,11 +361,17 @@ pub fn init_logging(config: LoggingConfig, otel_layer_config: Option>, + } + + impl Write for StalledSink { + fn write(&mut self, buf: &[u8]) -> io::Result { + if let Some(gate) = self.gate.take() { + let _ = gate.recv(); + } + Ok(buf.len()) + } + + fn flush(&mut self) -> io::Result<()> { + Ok(()) + } + } + + /// A sink that forwards every write on a channel, so a test can see + /// whether a line reached the writer thread. + struct ChannelSink(mpsc::Sender>); + + impl Write for ChannelSink { + fn write(&mut self, buf: &[u8]) -> io::Result { + let _ = self.0.send(buf.to_vec()); + Ok(buf.len()) + } + + fn flush(&mut self) -> io::Result<()> { + Ok(()) + } + } + + /// The writer threads belong to the process, not to the value + /// `init_logging` returns: a router started after the first in one + /// process logs through the same subscriber, whose writer must still be + /// running once the first `startup` has returned and dropped its guard. + #[test] + fn the_writer_outlives_the_handle_init_returns() { + let (tx, rx) = mpsc::channel(); + let mut writer = install(ChannelSink(tx), "test-log-outlives"); + { + let _handle = LogGuard(()); + } + writer + .write_all(b"after the handle dropped\n") + .expect("a lossy writer reports every line as written"); + let line = rx + .recv_timeout(Duration::from_secs(5)) + .expect("the writer thread still runs after the handle dropped"); + assert_eq!(line, b"after the handle dropped\n"); + } + + /// Lines a full queue drops are counted per sink and reported as the + /// total and the delta since the last report. + #[test] + fn dropped_lines_are_counted_and_reported_per_sink() { + let (release, gate) = mpsc::channel(); + let mut writer = install(StalledSink { gate: Some(gate) }, "test-log-dropped"); + for _ in 0..DEFAULT_BUFFERED_LINES_LIMIT + 1_000 { + writer + .write_all(b"finished processing request\n") + .expect("a lossy writer reports every line as written"); + } + let counted = dropped_lines() + .into_iter() + .find(|(name, _)| *name == "test-log-dropped") + .map(|(_, dropped)| dropped) + .expect("the sink is counted"); + assert!(counted > 0, "overflow is counted"); + let (_, total, since_last) = report_dropped_lines() + .into_iter() + .find(|(name, _, _)| *name == "test-log-dropped") + .expect("the sink is reported"); + assert_eq!(total, counted); + assert_eq!(since_last, counted, "the first report carries everything"); + let (_, _, since_last) = report_dropped_lines() + .into_iter() + .find(|(name, _, _)| *name == "test-log-dropped") + .expect("the sink is reported"); + assert_eq!(since_last, 0, "a second report carries only new drops"); + let _ = release.send(()); + } + + #[test] + fn a_stalled_sink_never_blocks_the_emitting_thread() { + let (release, gate) = mpsc::channel(); + let (mut writer, guard) = + non_blocking(StalledSink { gate: Some(gate) }, "test-log-stalled"); + let (done, finished) = mpsc::channel(); + thread::spawn(move || { + // More lines than the channel holds, while the sink accepts none. + for _ in 0..DEFAULT_BUFFERED_LINES_LIMIT + 1_000 { + writer + .write_all(b"finished processing request\n") + .expect("a lossy writer reports every line as written"); + } + let _ = done.send(writer.error_counter().dropped_lines()); + }); + let dropped = finished.recv_timeout(Duration::from_secs(10)); + // Unblock the writer thread whatever happened, so a failure does not leak it. + let _ = release.send(()); + let dropped = dropped.expect("writes against a stalled sink must return, not wait for it"); + assert!( + dropped > 0, + "the queue is bounded: overflow is dropped, never waited for" + ); + drop(guard); } } diff --git a/model_gateway/src/observability/metrics.rs b/model_gateway/src/observability/metrics.rs index 40f7be82f8..cbf189f9b8 100644 --- a/model_gateway/src/observability/metrics.rs +++ b/model_gateway/src/observability/metrics.rs @@ -134,6 +134,31 @@ pub(crate) const UPKEEP_INTERVAL_SECS: u64 = 5 * 60; pub(crate) const CACHE_AWARE_MATCH_RATIO_BUCKETS: &[f64] = &[0.0, 0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9, 1.0]; +/// Histogram buckets for `smg_kv_index_lookup_seconds` and +/// `smg_kv_event_apply_seconds`: a lookup takes microseconds and a batch +/// apply tens of microseconds, below the request-latency buckets' first +/// edge, so these double from half a microsecond to a quarter of a second. +pub(crate) const KV_INDEX_MICRO_BUCKETS: &[f64] = &[ + 0.000_000_5, + 0.000_001, + 0.000_002, + 0.000_004, + 0.000_008, + 0.000_016, + 0.000_032, + 0.000_064, + 0.000_128, + 0.000_256, + 0.000_512, + 0.001_024, + 0.002_048, + 0.004_096, + 0.008_192, + 0.016_384, + 0.065_536, + 0.262_144, +]; + /// Marks jemalloc as the final artifact's Rust global allocator. /// /// Call this before [`start_prometheus`] only from a binary or extension that @@ -332,15 +357,148 @@ pub(crate) fn init_metrics() { "KV event subscription task failures by worker and reason \ (panic, join_error, intern_failed)" ); + describe_counter!( + "smg_kv_event_subscriptions_total", + "KV event streams connected, by worker; a reconnect counts again" + ); + describe_counter!( + "smg_kv_event_batches_total", + "KV event batches by worker and disposition (applied, stale, tail_overflow, snapshot)" + ); + describe_counter!( + "smg_engine_load_polls_total", + "Load-monitor tick decisions by worker and mode: poll (no pushed load record on file), \ + fallback (pushed records older than the tick interval), skipped_fresh_push (the \ + worker's KV-event stream pushed its load within the interval; no GetLoads RPC)" + ); + describe_counter!( + "smg_kv_event_gaps_total", + "KV event sequence gaps by worker and outcome (replay_requested, \ + unrecovered_kept, unrecovered_cleared)" + ); + describe_counter!( + "smg_kv_event_missed_batches_total", + "KV event batches the publisher skipped and could not replay, by worker" + ); + describe_counter!( + "smg_kv_event_resyncs_total", + "KV event rank resyncs by worker and reason (out_of_range, data_loss, \ + publisher_restart, gap_cleared, snapshot)" + ); + describe_histogram!( + "smg_kv_event_lag_seconds", + "Age of a KV event batch when applied: now minus the publisher timestamp, by worker" + ); + describe_counter!( + "smg_kv_event_parentless_stores_total", + "Stores whose parent block the index did not hold for the worker, placed as a \ + new chain from the root instead (a parent evicted, dropped, or cleared), by worker" + ); + describe_counter!( + "smg_kv_event_parentless_blocks_total", + "Blocks of the parent-less stores, by worker" + ); + describe_counter!( + "smg_kv_event_blocks_total", + "Blocks named by applied KV events, by worker and op (stored, removed)" + ); + describe_histogram!( + "smg_kv_event_apply_seconds", + "Time to apply one KV event batch to the index, by worker" + ); + describe_histogram!( + "smg_kv_index_lookup_seconds", + "Time of one KV index lookup (overlap scoring of a request's block hashes) \ + in cache-aware routing, by index kind (positional, chain)" + ); + describe_gauge!( + "smg_kv_event_degraded_ranks", + "KV event ranks whose index may be stale after an unreplayed gap, by worker" + ); + describe_gauge!( + "smg_kv_event_tail_depth", + "Live KV event batches held while a snapshot resync is in flight, by worker" + ); + describe_gauge!( + "smg_log_dropped_lines", + "Log lines the lossy writer queue dropped since the process started, by sink \ + (stdout, file)" + ); + describe_gauge!( + "smg_kv_index_memberships", + "Blocks the KV index holds across a model's workers (a block two workers hold \ + counts twice), by model; published every 30 s" + ); + describe_gauge!( + "smg_kv_index_entries", + "Distinct entries in the KV index, by model: (position, content hash) pairs in \ + the positional indexer, distinct blocks on a chain in the chain index" + ); + describe_gauge!( + "smg_kv_index_runs_live", + "Runs linked in the chain index, by model (blocks_live over runs_live is the \ + mean run length; a falling ratio under steady traffic is fragmentation)" + ); + describe_gauge!( + "smg_kv_index_blocks_live", + "Content hashes held by the chain index's live runs, by model" + ); + describe_gauge!( + "smg_kv_index_arena_bytes", + "Bytes the chain index's word arena has handed out (hash arrays, child tables, \ + free lists included), by model" + ); + describe_gauge!( + "smg_kv_index_arena_free_bytes", + "Bytes of the chain index's word arena sitting in free lists, by model" + ); + describe_gauge!( + "smg_kv_index_slab_bytes", + "Bytes the chain index's run slab holds from the allocator, by model" + ); + describe_gauge!( + "smg_kv_index_moved_hashes", + "Stores that moved a held engine hash to another place in the chain index \ + (cumulative), by model" + ); + describe_gauge!( + "smg_kv_index_engine_conflicts", + "Stored blocks whose engine hash differed from the one the chain index holds for \ + the block, a fleet whose engines do not name content alike (cumulative), by model" + ); + describe_gauge!( + "smg_kv_index_blocks", + "Blocks the positional index holds for a worker, as the index counts them; \ + set when a KV event batch is applied, when the worker's state is reset \ + and when the worker is removed" + ); describe_gauge!( "smg_workers_overloaded", "Workers currently flagged overloaded and excluded from routing, by model" ); + describe_gauge!( + "smg_worker_stalled", + "Liveness veto on a worker: 1 while routing skips it, by worker and reason" + ); + describe_counter!( + "smg_worker_stall_transitions_total", + "Times the liveness veto was set on a worker, by worker and reason" + ); describe_counter!( "smg_worker_overload_shed_total", "Requests shed because every worker for the model is overloaded, by stage \ (selection, dispatch)" ); + describe_counter!( + "smg_worker_overload_fallback_total", + "Requests routed to the least-loaded worker because every worker for the model is \ + overloaded and shedding is off, by stage" + ); + describe_counter!( + "smg_worker_liveness_fallback_total", + "Requests routed to the least-loaded ready worker because every worker for the model \ + is vetoed and the ready ones are vetoed by the liveness tracker, by stage" + ); describe_gauge!( "smg_manual_policy_cache_entries", "Number of routing entries in manual policy cache" @@ -366,6 +524,11 @@ pub(crate) fn init_metrics() { "Cache-aware tree-mode selection branch (tree_match, spill, expected_wait_fallback, \ first_healthy_fallback)" ); + describe_counter!( + "smg_policy_inflight_reconciled_total", + "Policy bookings released at the per-poll in-flight reconciliation because their \ + completion never arrived, by policy" + ); describe_histogram!( "smg_cache_aware_match_ratio", "Cache-aware tree-mode best prefix match ratio per request (matched/input, 0..1)" @@ -587,6 +750,11 @@ pub fn start_prometheus(config: PrometheusConfig) -> PrometheusHandle { // summary. let match_ratio_matcher = Matcher::Full(String::from("smg_cache_aware_match_ratio")); + // The KV index's lookup and apply times are microseconds: their own + // buckets, or the recorder renders them as summaries. + let kv_lookup_matcher = Matcher::Full(String::from("smg_kv_index_lookup_seconds")); + let kv_apply_matcher = Matcher::Full(String::from("smg_kv_event_apply_seconds")); + PrometheusBuilder::new() .upkeep_timeout(Duration::from_secs(UPKEEP_INTERVAL_SECS)) .set_buckets_for_metric(duration_matcher, &duration_bucket) @@ -602,6 +770,10 @@ pub fn start_prometheus(config: PrometheusConfig) -> PrometheusHandle { .expect("failed to set event loop delay buckets") .set_buckets_for_metric(match_ratio_matcher, CACHE_AWARE_MATCH_RATIO_BUCKETS) .expect("failed to set cache-aware match ratio buckets") + .set_buckets_for_metric(kv_lookup_matcher, KV_INDEX_MICRO_BUCKETS) + .expect("failed to set KV index lookup buckets") + .set_buckets_for_metric(kv_apply_matcher, KV_INDEX_MICRO_BUCKETS) + .expect("failed to set KV event apply buckets") .install_recorder() .inspect(|_| { #[cfg(all( @@ -1368,6 +1540,27 @@ impl Metrics { .increment(1); } + /// Record a request steered to the least-loaded worker because every + /// worker for the model is overloaded and shedding is off. + pub fn record_worker_overload_fallback(stage: &'static str) { + counter!( + "smg_worker_overload_fallback_total", + "stage" => stage + ) + .increment(1); + } + + /// Record a request steered to the least-loaded ready worker because + /// every worker for the model is vetoed and the ready ones are vetoed by + /// the liveness tracker. + pub fn record_worker_liveness_fallback(stage: &'static str) { + counter!( + "smg_worker_liveness_fallback_total", + "stage" => stage + ) + .increment(1); + } + /// Record manual policy execution branch for routing decisions pub fn record_worker_manual_policy_branch(branch: &'static str) { counter!( @@ -1424,6 +1617,16 @@ impl Metrics { .increment(1); } + /// Record bookings a policy released at the in-flight reconciliation + /// tick: dispatches whose completion never reached it. + pub fn record_policy_inflight_reconciled(policy: &'static str, released: usize) { + counter!( + "smg_policy_inflight_reconciled_total", + "policy" => policy + ) + .increment(released as u64); + } + /// Record the best prefix match ratio (matched/input, 0..1) of a cache-aware /// tree-mode routing decision pub fn record_cache_aware_match_ratio(ratio: f64) { @@ -1468,6 +1671,26 @@ impl Metrics { .set(count as f64); } + /// Flip the liveness veto gauge for a worker; count the transition when + /// the veto is set. + pub fn set_worker_stalled(worker_url: &str, reason: &'static str, stalled: bool) { + let worker = intern_string(worker_url); + gauge!( + "smg_worker_stalled", + "worker" => Arc::clone(&worker), + "reason" => reason + ) + .set(if stalled { 1.0 } else { 0.0 }); + if stalled { + counter!( + "smg_worker_stall_transitions_total", + "worker" => worker, + "reason" => reason + ) + .increment(1); + } + } + /// Set worker health status pub fn set_worker_health(worker_url: &str, healthy: bool) { let worker_interned = intern_string(worker_url); @@ -1499,6 +1722,149 @@ impl Metrics { .increment(1); } + /// Count a KV event stream connected: the first time and every reconnect, + /// so a drill can time a resubscription without reading the log. + pub fn record_kv_event_subscription(worker_url: &str) { + counter!( + "smg_kv_event_subscriptions_total", + "worker" => intern_string(worker_url) + ) + .increment(1); + } + + /// Count a KV event batch by what the subscriber did with it. + /// One load-monitor tick decision for a worker (see `PollMode`). + pub fn record_engine_load_poll(worker_url: &str, mode: &'static str) { + counter!( + "smg_engine_load_polls_total", + "worker" => intern_string(worker_url), + "mode" => mode + ) + .increment(1); + } + + pub fn record_kv_event_batch(worker_url: &str, disposition: &'static str) { + counter!( + "smg_kv_event_batches_total", + "worker" => intern_string(worker_url), + "disposition" => disposition + ) + .increment(1); + } + + /// Count a sequence gap and the batches it skipped. + pub fn record_kv_event_gap(worker_url: &str, outcome: &'static str, missed: u64) { + let worker_interned = intern_string(worker_url); + counter!( + "smg_kv_event_gaps_total", + "worker" => Arc::clone(&worker_interned), + "outcome" => outcome + ) + .increment(1); + if missed > 0 { + counter!("smg_kv_event_missed_batches_total", "worker" => worker_interned) + .increment(missed); + } + } + + /// Count a rank resync (its index state dropped) by reason. + pub fn record_kv_event_resync(worker_url: &str, reason: &'static str) { + counter!( + "smg_kv_event_resyncs_total", + "worker" => intern_string(worker_url), + "reason" => reason + ) + .increment(1); + } + + /// Observe how old a batch was when it was applied. + pub fn record_kv_event_lag(worker_url: &str, seconds: f64) { + histogram!("smg_kv_event_lag_seconds", "worker" => intern_string(worker_url)) + .record(seconds); + } + + /// Count the blocks an applied batch's events named (`stored` or + /// `removed`), before the monitor's tier and group filters: the rate the + /// index is offered. + pub fn record_kv_event_blocks(worker_url: &str, op: &'static str, blocks: usize) { + counter!( + "smg_kv_event_blocks_total", + "worker" => intern_string(worker_url), + "op" => op + ) + .increment(blocks as u64); + } + + /// Count the stores a batch placed without their parent (as new chains + /// from the root) and the blocks they carried. + pub fn record_kv_event_parentless(worker_url: &str, stores: u64, blocks: u64) { + let worker = intern_string(worker_url); + counter!("smg_kv_event_parentless_stores_total", "worker" => worker.clone()) + .increment(stores); + counter!("smg_kv_event_parentless_blocks_total", "worker" => worker).increment(blocks); + } + + /// Time to apply one batch to the index. + pub fn record_kv_event_apply(worker_url: &str, seconds: f64) { + histogram!("smg_kv_event_apply_seconds", "worker" => intern_string(worker_url)) + .record(seconds); + } + + /// Time of one index lookup on the routing path, recorded after the + /// lookup returned and outside any lock. + pub fn record_kv_index_lookup(index: &'static str, seconds: f64) { + histogram!("smg_kv_index_lookup_seconds", "index" => index).record(seconds); + } + + /// Ranks of this worker whose index may be stale. + pub fn set_kv_event_degraded_ranks(worker_url: &str, count: usize) { + gauge!("smg_kv_event_degraded_ranks", "worker" => intern_string(worker_url)) + .set(count as f64); + } + + /// Live batches held for a worker while a snapshot resync is in flight. + pub fn set_kv_event_tail_depth(worker_url: &str, depth: usize) { + gauge!("smg_kv_event_tail_depth", "worker" => intern_string(worker_url)).set(depth as f64); + } + + /// Publish the blocks the positional index holds for a worker. Called from + /// the KV event subscriber where it already counts applied batches, never + /// from the lookup path, so routing reads nothing that writes. + pub fn set_kv_index_blocks(worker_url: &str, blocks: usize) { + gauge!("smg_kv_index_blocks", "worker" => intern_string(worker_url)).set(blocks as f64); + } + + /// Publish the lines a log writer's queue has dropped so far, by sink. + pub fn set_log_dropped_lines(sink: &'static str, dropped: usize) { + gauge!("smg_log_dropped_lines", "sink" => sink).set(dropped as f64); + } + + /// Publish a model's KV index size: memberships across workers and + /// distinct entries. Called from the monitor's periodic stats task, never + /// from the lookup path. + pub fn set_kv_index_size(model_id: &str, memberships: usize, entries: usize) { + let model = intern_string(model_id); + gauge!("smg_kv_index_memberships", "model" => model.clone()).set(memberships as f64); + gauge!("smg_kv_index_entries", "model" => model).set(entries as f64); + } + + /// Publish the chain index's shape and memory for a model, from its own + /// counters: live runs and blocks, arena and slab bytes, moved hashes and + /// engine conflicts. + pub fn set_kv_index_chain_stats(model_id: &str, stats: &kv_index::ChainIndexStats) { + let model = intern_string(model_id); + gauge!("smg_kv_index_runs_live", "model" => model.clone()).set(stats.runs_live as f64); + gauge!("smg_kv_index_blocks_live", "model" => model.clone()).set(stats.blocks_live as f64); + gauge!("smg_kv_index_arena_bytes", "model" => model.clone()).set(stats.arena_bytes as f64); + gauge!("smg_kv_index_arena_free_bytes", "model" => model.clone()) + .set(stats.arena_free_bytes as f64); + gauge!("smg_kv_index_slab_bytes", "model" => model.clone()).set(stats.slab_bytes as f64); + gauge!("smg_kv_index_moved_hashes", "model" => model.clone()) + .set(stats.moved_hashes as f64); + gauge!("smg_kv_index_engine_conflicts", "model" => model) + .set(stats.engine_conflicts as f64); + } + // ======================================================================== // Layer 3: Worker resilience metrics (circuit breaker) // ======================================================================== @@ -1835,24 +2201,30 @@ impl Metrics { } } +/// Metrics helpers for tests elsewhere in the crate. #[cfg(test)] -mod tests { - use std::net::TcpListener; - +pub(crate) mod test_support { use metrics_exporter_prometheus::PrometheusBuilder; - use openai_protocol::worker::{SchedulerLoadSnapshot, WorkerLoadResponse}; - - use super::*; /// Run `f` under a thread-local Prometheus recorder and return the /// rendered `/metrics` text — the same scrape output the :29000 endpoint /// serves in production. - fn render_with_recorder(f: impl FnOnce()) -> String { + pub(crate) fn render_with_recorder(f: impl FnOnce()) -> String { let recorder = PrometheusBuilder::new().build_recorder(); let handle = recorder.handle(); metrics::with_local_recorder(&recorder, f); handle.render() } +} + +#[cfg(test)] +mod tests { + use std::net::TcpListener; + + use metrics_exporter_prometheus::PrometheusBuilder; + use openai_protocol::worker::{SchedulerLoadSnapshot, WorkerLoadResponse}; + + use super::{test_support::render_with_recorder, *}; #[test] fn tokenizer_activity_registers_both_layers_on_scrape() { diff --git a/model_gateway/src/policies/cache_aware.rs b/model_gateway/src/policies/cache_aware.rs index 9f20d8ce55..1d5291d998 100644 --- a/model_gateway/src/policies/cache_aware.rs +++ b/model_gateway/src/policies/cache_aware.rs @@ -10,10 +10,12 @@ 1. Event-Driven (gRPC + KV events) ------------------------------------------- - Uses PositionalIndexer overlap scoring from KvEventMonitor. Routes based + Uses the KV index's overlap scoring from KvEventMonitor. Routes based on actual backend KV cache state. Selects the worker with the highest overlap count; LeastLoad breaks equal-affinity ties atomically. - Falls back to LeastLoad when no cache overlap exists. + Falls back to LeastLoad when no cache overlap exists. Over a pool larger + than 32 workers the decision reads the holders of the request's blocks + and a sample of 16 others, not the pool (see `FLEET_SAMPLE`). 2. Approximate Token Tree (gRPC, no KV events) ------------------------------------------- @@ -68,23 +70,32 @@ use std::{ collections::{HashMap, HashSet}, sync::{ - atomic::{AtomicBool, Ordering}, + atomic::{AtomicBool, AtomicU64, Ordering}, Arc, }, time::{Duration, Instant}, }; +use arc_swap::ArcSwap; use dashmap::DashMap; -use kv_index::{compute_request_content_hashes, PositionalIndexer, TenantId, TokenTree, Tree}; +use kv_index::{ + compute_request_content_hashes, request_prefix_hashes, salt::request_content_hashes_with_seed, + ContentHash, OverlapScores, TenantId, TokenTree, Tree, +}; use openai_protocol::worker::WorkerLoadResponse; use parking_lot::RwLock; use rand::RngExt; use serde::{Deserialize, Serialize}; -use tracing::{debug, warn}; +use tracing::{debug, error, warn}; use super::{ - normalize_model_key, utils::PeriodicTask, CacheAwareConfig, CacheNamespace, LeastLoadPolicy, - LoadBalancingPolicy, SelectWorkerInfo, TEXT_MARKER_LEN, + cost::{ + self, CandidateInputs, OptimisticAccounting, Pick, RequestInputs, WorkerSelectionPolicy, + }, + normalize_model_key, + utils::PeriodicTask, + CacheAwareConfig, CacheNamespace, LeastLoadPolicy, LoadBalancingPolicy, SelectWorkerInfo, + TEXT_MARKER_LEN, }; /// Latest per-worker backend load snapshot stream, keyed by worker URL. pub(crate) use crate::worker::load_state::{LoadReceiver, LoadSnapshot}; @@ -92,9 +103,42 @@ use crate::{ config::CacheIndexKind, mesh::adapters::tree_sync::{RepairEntry, TreeDelta, TreeRepairPage, TreeSyncAdapter}, observability::metrics::Metrics, - worker::{KvEventMonitor, Worker}, + worker::{liveness, KvEventMonitor, KvIndex, Worker}, }; +/// An overlap of at most this many blocks counts as a miss for the warm-up +/// slice: a chat template's head is this long, shared by every holder, and +/// recomputing it costs less than keeping a returned worker idle. +const WARMUP_MISS_BLOCKS: f64 = 4.0; + +/// A hit whose best overlap is at most this share of the request is a +/// shallow one: diverting it to a thin worker recomputes little, so these go +/// first (see [`CacheAwarePolicy::warmup_divert`]). +const DIVERT_SHALLOW_SHARE: f64 = 0.5; + +/// A thin worker with requests in flight gets no second diversion sooner than +/// this after its last one. +const DIVERT_WINDOW_MS: u64 = 2_000; + +/// How many eligible workers an event-driven decision reads beside the +/// holders of the request's blocks once the pool is larger than twice this: +/// they are the miss path's choices, the spill targets, the rows an +/// all-workers selection policy ranks, and the sample the count-pressure +/// gate's fleet mean is taken from. Sixteen choices balance as a full scan +/// does (two already do, the power of d choices), and cost the same at 128 +/// workers as at 10,000. A pool of at most twice this many workers is read +/// whole, so nothing changes for a small fleet. +const FLEET_SAMPLE: usize = 16; + +/// How long a pool table serves before a decision rebuilds it: a worker the +/// index interned after the build gains affinity within this, and a worker +/// that entered the warm-up window is sliced to within this. +const POOL_TABLE_REFRESH_MS: u64 = 1_000; + +/// Pool tables kept at once: one per routing pool the policy sees (a model's +/// regular pool, the PD legs' pools); past this the oldest goes. +const POOL_TABLES_KEPT: usize = 8; + /// Cache-aware routing policy /// /// Routes requests based on cache affinity when load is balanced, @@ -113,6 +157,13 @@ pub struct CacheAwarePolicy { /// candidate set. It owns backend snapshots and since-poll dispatch credit, /// keeping selection and credit atomic under concurrent arrivals. load_scorer: LeastLoadPolicy, + /// Selection policy run over the per-worker inputs gathered for each + /// request (`cost` module); the default reproduces the affinity-group + /// decision exactly and the ports of published cost functions replace it. + selection: WorkerSelectionPolicy, + /// Optimistic self-accounting of dispatches the engines have not yet + /// reported; `None` unless `selection_accounting_ttl_ms > 0`. + accounting: Option, /// String-based trees for HTTP connections (text input) string_trees: Arc>>, /// Token-based trees for gRPC connections (pre-tokenized input) @@ -125,6 +176,19 @@ pub struct CacheAwarePolicy { /// trigger. `None` until wired by the registry (then the policy stays /// count-only, preserving current behavior). load_rx: RwLock>, + /// Misses seen by the warm-up slice; every `period`th goes to a warming + /// worker (see [`liveness::Warmup`]). + warmup_misses: AtomicU64, + /// Hits seen since the last diversion to a thin worker, and whether a + /// shallow one was among them (see [`Self::warmup_divert`]). + divert_hits: AtomicU64, + divert_shallow_seen: AtomicBool, + /// Each routing pool's workers by KV index id and position, built once + /// per pool snapshot (see [`PoolTable`]). + pool_tables: ArcSwap>>, + /// Set while one decision rebuilds a pool table so the others keep + /// serving the one they have. + pool_table_building: AtomicBool, /// Model-scoped hash indexes for resolving tenant delta hashes. /// Outer key is the normalized model_id; inner maps hold /// `hash → reconstructable prefix/tokens` per tree kind. @@ -364,14 +428,37 @@ impl CacheAwarePolicy { None }; + let selection_policy_name = config + .selection_policy + .as_deref() + .unwrap_or(cost::DEFAULT_POLICY); + let selection = cost::build(selection_policy_name, config.selection_temperature) + .unwrap_or_else(|err| { + // Configuration validation rejects this before a policy is built; + // a policy constructed outside that path still routes, with the + // default decision, rather than failing every request. + error!(%err, "Invalid selection policy; using {}", cost::DEFAULT_POLICY); + cost::default_policy(config.selection_temperature) + }); + let accounting = (config.selection_accounting_ttl_ms > 0).then(|| { + OptimisticAccounting::new(Duration::from_millis(config.selection_accounting_ttl_ms)) + }); + Self { config, load_scorer: LeastLoadPolicy::new(), + selection, + accounting, string_trees, token_trees, _eviction_task: eviction_task, kv_monitor: RwLock::new(None), load_rx: RwLock::new(None), + warmup_misses: AtomicU64::new(0), + divert_hits: AtomicU64::new(0), + divert_shallow_seen: AtomicBool::new(false), + pool_tables: ArcSwap::from_pointee(Vec::new()), + pool_table_building: AtomicBool::new(false), hash_index, populate_hash_index: AtomicBool::new(false), mesh_tree_sync: RwLock::new(None), @@ -542,9 +629,7 @@ impl CacheAwarePolicy { // Both defaults are 1.0, above the maximum clamped utilization and // spread. CacheAware now polls loads for LeastLoad even at defaults, // so return before touching the snapshot or scanning the fleet. - if self.config.balance_token_usage_threshold >= 1.0 - && self.config.overload_token_usage_threshold >= 1.0 - { + if !self.kv_pressure_gate_configured() { return false; } @@ -1093,12 +1178,179 @@ impl TreeHandle for CacheAwarePolicy { } } -/// One positive-overlap candidate: slice index and possibly decayed score. +/// One positive-overlap candidate: slice index, undecayed overlap in blocks +/// (the tree paths report matched units over the block size) and the possibly +/// decayed score the affinity-group decision ranks on. struct OverlapCandidate { idx: usize, + raw_score: f64, effective_score: f64, } +/// One routing pool's workers by their KV index ids, built once per pool +/// snapshot and refreshed every [`POOL_TABLE_REFRESH_MS`]. The registry hands +/// the policy the same slice until membership changes, so an overlap lookup's +/// `index id -> blocks` result maps to slice positions through this table +/// instead of a string lookup per healthy worker per request, and the warm-up +/// slice reads the workers inside its window from it instead of scanning the +/// pool on every thin miss. +#[derive(Debug)] +struct PoolTable { + /// Address and length of the slice the table was built from. + slice: (usize, usize), + /// Gateway clock at the build, for the refresh. + built_ms: u64, + /// Address of the worker at each position at the build: a pool rebuilt + /// at the same address with other workers fails this check per holder. + worker_at: Vec, + /// The index id of the worker at each position; `None` while it is not + /// interned. + id_at: Vec>, + /// Index id -> position. + position_of: HashMap, + /// The fleet's index level at the build: the upper median of the + /// workers' index sizes, what a thin worker is thin against. + fleet_level: usize, + /// Positions of the warm-up slice's candidates at the build: the workers + /// thinner than the fleet (an index emptied by a resync, or never fed) + /// whatever their age, else the workers inside the warm-up window + /// (admitted within it, index still growing) unless every worker was (a + /// young fleet slices nothing). + warming: Vec, + /// Positions of the thin workers alone, the targets of the hit diversion + /// (see [`CacheAwarePolicy::warmup_divert`]). + thin: Vec, +} + +impl PoolTable { + fn slice_key(workers: &[Arc]) -> (usize, usize) { + (workers.as_ptr() as usize, workers.len()) + } + + fn worker_address(worker: &Arc) -> usize { + Arc::as_ptr(worker).cast::<()>() as usize + } + + /// One pass over the pool: the string lookup per worker happens here, + /// once per snapshot and refresh, not per request. + fn build(workers: &[Arc], indexer: &KvIndex, now_ms: u64) -> Self { + let warmup = liveness::warmup(); + let mut worker_at = Vec::with_capacity(workers.len()); + let mut id_at = Vec::with_capacity(workers.len()); + let mut position_of = HashMap::with_capacity(workers.len()); + let mut sized: Vec> = Vec::with_capacity(workers.len()); + for (position, worker) in workers.iter().enumerate() { + worker_at.push(Self::worker_address(worker)); + let id = indexer.worker_id(worker.url()); + id_at.push(id); + if let Some(id) = id { + position_of.insert(id, position as u32); + } + // Index size, and its growth since the admission for the age rule + // (the baseline restarts when the count drops); thinness is read + // against the fleet's level below, whatever the growth. + sized.push(id.map(|id| { + let indexed = indexer.worker_block_count(id); + (indexed, worker.warmup_growth(indexed)) + })); + } + // The fleet's level: the upper median of the sizes, so one emptied + // worker does not drag it down. + let mut sizes: Vec = sized + .iter() + .map(|sized| sized.map_or(0, |(indexed, _)| indexed)) + .collect(); + sizes.sort_unstable(); + let fleet_level = sizes.get(sizes.len() / 2).copied().unwrap_or(0); + let mut thin = Vec::new(); + let mut young = Vec::new(); + if warmup.share > 0.0 { + for (position, (worker, sized)) in workers.iter().zip(&sized).enumerate() { + let (indexed, grown) = match sized { + Some((indexed, grown)) => (*indexed, Some(*grown)), + None => (0, None), + }; + if !warmup.applies(worker.admitted_age(), grown, indexed, fleet_level) { + continue; + } + if warmup.is_thin(indexed, fleet_level) { + thin.push(position as u32); + } else { + young.push(position as u32); + } + } + } + // A thin worker is a candidate whatever the fleet's age; workers + // warming only because they are young are candidates only when the + // fleet is not all young. + let warming = if !thin.is_empty() { + thin.clone() + } else if young.len() < workers.len() { + young + } else { + Vec::new() + }; + Self { + slice: Self::slice_key(workers), + built_ms: now_ms, + worker_at, + id_at, + position_of, + fleet_level, + warming, + thin, + } + } + + /// Whether the table was built from `workers` (the same slice). + fn describes(&self, workers: &[Arc]) -> bool { + self.slice == Self::slice_key(workers) + } + + /// Whether the table is within its refresh period. + fn fresh(&self, now_ms: u64) -> bool { + now_ms.saturating_sub(self.built_ms) < POOL_TABLE_REFRESH_MS + } + + /// The position holding index worker `id`, when the worker there is + /// still the one the table was built from. + fn position(&self, id: u32, workers: &[Arc]) -> Option { + let position = *self.position_of.get(&id)? as usize; + let same_worker = workers + .get(position) + .is_some_and(|worker| Self::worker_address(worker) == self.worker_at[position]); + same_worker.then_some(position) + } +} + +/// What the sampled decision made of a request. +enum Sampled { + /// The event-driven decision over the holders and a sample of the pool: + /// the selection, or `None` when it declined every candidate. + Decided(Option), + /// Not a request for the sampled decision: the pool is small, the request + /// carries no tokens or has no event index, the configuration wants the + /// KV-pressure gate (a reading of the whole fleet), or the sample held no + /// eligible worker. The scan decides. + Scan, +} + +/// Which holders a decision gathers from the overlap scores. A shared chat +/// template puts a shallow overlap on most of a fleet, so the scores name +/// nearly every worker; the decision must not pay per named worker. +#[derive(Clone, Copy)] +enum Gather { + /// The eligible holders of the deepest overlap only: what the exact + /// maximum-group decision reads (the default policy at temperature zero, + /// no decay in effect, no accounting). A tie wider than `cap` (a fleet + /// sharing a chat template's head, a prompt reaching no further) is a + /// uniform draw of `cap` of them: the expected-wait selector breaks the + /// tie among those, as a miss is placed among a sample of the pool. + TopGroup { cap: Option }, + /// Every eligible holder in slice order, the `cap` deepest when set. + All { cap: Option }, +} + /// Pressure-tuning inputs for [`CacheAwarePolicy::overlap_candidates`]: the two /// config knobs plus the immutable load snapshot captured from the load /// receiver at selection time, from which each worker's waiting-prefill @@ -1107,12 +1359,118 @@ struct OverlapCandidate { /// receiver is wired; workers absent from the snapshot are never decayed. struct OverlapTuning<'a> { overlap_decay: f32, + /// Zero means the exact maximum-overlap group; the policy carries the + /// temperature for its draw, the host reads it to know which holders to + /// gather (see [`CacheAwarePolicy::exact_maximum_decision`]). selection_temperature: f32, waiting_prefill_tokens: Option<&'a LoadSnapshot>, } impl LoadBalancingPolicy for CacheAwarePolicy { fn select_worker(&self, workers: &[Arc], info: &SelectWorkerInfo) -> Option { + // The event-driven decision over a large pool reads the holders of + // the request's blocks and a bounded sample of the rest; every other + // path, and a pool the sample would cover anyway, reads the pool. + if let Sampled::Decided(selected) = self.select_worker_sampled(workers, info) { + return selected; + } + self.select_worker_scanned(workers, info) + } + + fn on_request_complete(&self, worker_url: &str, success: bool) { + if let Some(accounting) = &self.accounting { + accounting.release(worker_url); + } + self.selection.on_request_complete(worker_url); + // Could track success rates per worker for more intelligent routing + if !success { + // Optionally reduce affinity for failed requests + tracing::debug!( + "Request to {} completed with success={}", + worker_url, + success + ); + } + } + + fn name(&self) -> &'static str { + "cache_aware" + } + + fn needs_request_text(&self) -> bool { + true // Cache-aware policy needs request text for cache affinity + } + + fn update_loads(&self, loads: &HashMap) { + // WorkerMonitor invokes this immediately before publishing its complete + // immutable snapshot (with no await between the two operations). Advance + // ExpectedWait here so each successful worker's new load and credit + // reset stay atomic. KV-pressure and overlap decay intentionally keep + // using the last fully published snapshot during that handoff; scoring + // ExpectedWait from its old value would instead pair an + // old queue with a new reset and reopen the incast race. + self.load_scorer.update_loads(loads); + } + + /// Expected-wait selection needs backend snapshots for every CacheAware + /// configuration; KV pressure and overlap decay consume them as well. + fn needs_backend_loads(&self) -> bool { + true + } + + fn remove_worker(&self, url: &str) { + // The CacheAware-specific removal path already prunes trees and hash + // placements before the registry invokes this generic load-aware + // hook. Keep this hook scoped to LeastLoad state so worker churn does + // not repeat the full all-model cache scan. + self.load_scorer.remove_worker(url); + if let Some(accounting) = &self.accounting { + accounting.forget_worker(url); + } + self.selection.on_worker_removed(url); + } + + fn reconcile_in_flight(&self, worker_url: &str, in_flight: usize) { + let released = self + .accounting + .as_ref() + .map_or(0, |accounting| accounting.reconcile(worker_url, in_flight)); + self.selection.reconcile_in_flight(worker_url, in_flight); + if released > 0 { + Metrics::record_policy_inflight_reconciled(self.name(), released); + debug!( + worker = worker_url, + in_flight, released, "Released bookings whose completion never arrived" + ); + } + } + + fn reset(&self) { + self.load_scorer.reset(); + } + + fn as_any(&self) -> &dyn std::any::Any { + self + } +} + +/// The event-driven index resolved once per request for the selection. +struct EventIndex<'a> { + indexer: Arc, + block_size: usize, + model_id: &'a str, +} + +// Private helper methods for select_worker +impl CacheAwarePolicy { + /// The decision over the whole pool: one eligibility read per worker, + /// then the hash, tree or event-driven path. The sampled decision covers + /// the event-driven path over a large pool; this is everything else. + fn select_worker_scanned( + &self, + workers: &[Arc], + info: &SelectWorkerInfo, + ) -> Option { let request_text = info.request_text; let request_tokens = info.tokens; @@ -1157,15 +1515,14 @@ impl LoadBalancingPolicy for CacheAwarePolicy { } // Cache-aware routing when balanced — three types (mutually exclusive): - // 1. Event-driven: PositionalIndexer overlap scoring (gRPC + KV events) + // 1. Event-driven: KV index overlap scoring (gRPC + KV events) // 2. Approximate token tree: TokenTree prefix matching (gRPC, no events) // 3. Approximate string tree: Tree prefix matching (HTTP) if let Some(tokens) = request_tokens { // Event-driven mode re-hashes engine-reported blocks from their - // token ids on both sides, so a namespace marker on the request - // side alone would break every same-namespace match. It stays - // unpartitioned here; the approximate and hash modes below key - // under the request's cache namespace. + // token ids on both sides, so it keys a namespace through the + // hash seed rather than a marker (see cache_namespace.rs); the + // approximate and hash modes below key under the marker. if let Some(index) = self.event_index_for(model_id) { self.select_worker_event_driven( workers, @@ -1191,69 +1548,152 @@ impl LoadBalancingPolicy for CacheAwarePolicy { } } - fn on_request_complete(&self, worker_url: &str, success: bool) { - // Could track success rates per worker for more intelligent routing - if !success { - // Optionally reduce affinity for failed requests - tracing::debug!( - "Request to {} completed with success={}", - worker_url, - success - ); + /// The event-driven decision for a pool larger than twice + /// [`FLEET_SAMPLE`]: eligibility is read for a uniform sample of the pool + /// (the miss path's choices, the spill targets, the count-pressure gate's + /// fleet mean) and for the holders the index names, never for the whole + /// pool. A sample without one eligible worker defers to the scan, which + /// finds any that remain. + fn select_worker_sampled( + &self, + workers: &[Arc], + info: &SelectWorkerInfo, + ) -> Sampled { + let Some(tokens) = info.tokens else { + return Sampled::Scan; + }; + if workers.len() <= 2 * FLEET_SAMPLE + || self.config.cache_index == CacheIndexKind::Hash + || self.kv_pressure_gate_configured() + { + return Sampled::Scan; } + let mut sample = [0usize; FLEET_SAMPLE]; + let mut sampled = 0usize; + let mut load_sum = 0usize; + let mut rng = rand::rng(); + // Draws with replacement, repeats dropped: above twice the sample + // size a repeat costs a slot and nothing else. + for _ in 0..2 * FLEET_SAMPLE { + if sampled == FLEET_SAMPLE { + break; + } + let idx = rng.random_range(0..workers.len()); + if sample[..sampled].contains(&idx) { + continue; + } + let state = workers[idx].routing_state(); + if state.eligible() { + sample[sampled] = idx; + sampled += 1; + load_sum += state.load; + } + } + if sampled == 0 { + return Sampled::Scan; + } + let sample = &mut sample[..sampled]; + sample.sort_unstable(); + let avg_load = load_sum as f64 / sampled as f64; + let model_id = normalize_model_key(workers[sample[0]].model_id()); + let Some(index) = self.event_index_for(model_id) else { + return Sampled::Scan; + }; + Sampled::Decided( + self.select_worker_event_driven(workers, tokens, sample, avg_load, &index, info), + ) } - fn name(&self) -> &'static str { - "cache_aware" - } - - fn needs_request_text(&self) -> bool { - true // Cache-aware policy needs request text for cache affinity - } - - fn update_loads(&self, loads: &HashMap) { - // WorkerMonitor invokes this immediately before publishing its complete - // immutable snapshot (with no await between the two operations). Advance - // ExpectedWait here so each successful worker's new load and credit - // reset stay atomic. KV-pressure and overlap decay intentionally keep - // using the last fully published snapshot during that handoff; scoring - // ExpectedWait from its old value would instead pair an - // old queue with a new reset and reopen the incast race. - self.load_scorer.update_loads(loads); - } - - /// Expected-wait selection needs backend snapshots for every CacheAware - /// configuration; KV pressure and overlap decay consume them as well. - fn needs_backend_loads(&self) -> bool { - true - } - - fn remove_worker(&self, url: &str) { - // The CacheAware-specific removal path already prunes trees and hash - // placements before the registry invokes this generic load-aware - // hook. Keep this hook scoped to LeastLoad state so worker churn does - // not repeat the full all-model cache scan. - self.load_scorer.remove_worker(url); + /// Whether the decision is the exact maximum-overlap group: the default + /// selection policy at temperature zero, no decay in effect and no + /// accounting. Then only the deepest holders are gathered; every other + /// configuration reads each holder. + fn exact_maximum_decision(&self, tuning: &OverlapTuning<'_>) -> bool { + self.selection.name() == cost::DEFAULT_POLICY + && tuning.selection_temperature <= 0.0 + && self.accounting.is_none() + && (tuning.overlap_decay <= 0.0 || tuning.waiting_prefill_tokens.is_none()) } - fn reset(&self) { - self.load_scorer.reset(); + /// Whether the KV-pressure gate is configured: it reads the whole + /// fleet's token usage, so the decision reads the pool. + fn kv_pressure_gate_configured(&self) -> bool { + self.config.balance_token_usage_threshold < 1.0 + || self.config.overload_token_usage_threshold < 1.0 } - fn as_any(&self) -> &dyn std::any::Any { - self + /// The pool table for `workers`: the one kept when it is fresh, a stale + /// one while another decision rebuilds, else built here and published. + fn pool_table( + &self, + workers: &[Arc], + indexer: &KvIndex, + now_ms: u64, + ) -> Arc { + { + let tables = self.pool_tables.load(); + if let Some(table) = tables.iter().find(|table| table.describes(workers)) { + if table.fresh(now_ms) || self.pool_table_building.swap(true, Ordering::AcqRel) { + return Arc::clone(table); + } + } + } + let table = Arc::new(PoolTable::build(workers, indexer, now_ms)); + self.pool_tables.rcu(|tables| { + let mut next: Vec> = tables + .iter() + .filter(|kept| kept.slice != table.slice) + .cloned() + .collect(); + if next.len() >= POOL_TABLES_KEPT { + if let Some(oldest) = next + .iter() + .enumerate() + .min_by_key(|(_, kept)| kept.built_ms) + .map(|(position, _)| position) + { + next.swap_remove(oldest); + } + } + next.push(Arc::clone(&table)); + next + }); + self.pool_table_building.store(false, Ordering::Release); + table + } + + /// The union of two ascending position lists, ascending: the eligible + /// workers a decision read and the holders, which a sampled decision's + /// list need not contain. + fn merge_rows(eligible: &[usize], candidates: &[OverlapCandidate]) -> Vec { + let mut rows = Vec::with_capacity(eligible.len() + candidates.len()); + let (mut i, mut j) = (0, 0); + loop { + let next = match (eligible.get(i), candidates.get(j)) { + (Some(&idx), Some(candidate)) if idx < candidate.idx => { + i += 1; + idx + } + (Some(&idx), Some(candidate)) if idx == candidate.idx => { + i += 1; + j += 1; + idx + } + (_, Some(candidate)) => { + j += 1; + candidate.idx + } + (Some(&idx), None) => { + i += 1; + idx + } + (None, None) => break, + }; + rows.push(next); + } + rows } -} -/// The event-driven index resolved once per request for the selection. -struct EventIndex<'a> { - indexer: Arc, - block_size: usize, - model_id: &'a str, -} - -// Private helper methods for select_worker -impl CacheAwarePolicy { /// The event-driven index for this model: its indexer and block size. /// `None` when there is no monitor or indexer, or the indexer is empty /// (startup, reconnect), so routing falls through to the approximate @@ -1264,7 +1704,7 @@ impl CacheAwarePolicy { let guard = self.kv_monitor.read(); let monitor = guard.as_ref()?; let indexer = monitor.get_indexer(model_id)?; - if indexer.current_size() == 0 { + if indexer.is_empty() { return None; } // Per-model block_size: learned from events > config default @@ -1296,6 +1736,19 @@ impl CacheAwarePolicy { guard.as_ref().map(|rx| rx.borrow().clone()) } + /// Request inputs for the tree paths: units are tokens (token tree) or + /// chars (string tree); the trees carry no prefix hashes. + fn tree_request(&self, units: usize, avg_load: f64) -> RequestInputs<'static> { + let block_size = self.config.block_size.max(1); + RequestInputs { + prompt_tokens: units, + block_size, + request_blocks: (units / block_size).max(1), + avg_load, + prefix_hashes: None, + } + } + /// Select and credit one final worker atomically with LeastLoad's /// expected-wait algorithm. CacheAware must call this exactly once per /// successful dispatch, after cache affinity and spill logic have fixed @@ -1345,6 +1798,169 @@ impl CacheAwarePolicy { && load > avg_load + self.config.balance_abs_threshold as f64 } + /// The warm-up slice: one cache miss in `1 / share` goes to the + /// least-loaded warming worker, so a returned, new or emptied worker + /// builds a cache instead of idling behind the fleet's affinity (on the + /// real fleet a restarted worker went a minute without a request; in the + /// soaks a worker whose index a publisher-restart resync had emptied never + /// saw a request again, every prompt holding an overlap elsewhere). The + /// candidates are the pool table's, chosen at its build (refreshed every + /// second): workers thinner than the fleet (see + /// [`liveness::Warmup::is_thin`]) whatever the fleet's age, else the + /// workers inside the warm-up window unless every worker is (a young + /// fleet slices nothing, a miss getting a load-balanced pick anyway); a + /// settled fleet pays nothing here. Each candidate is checked live before + /// the pick, and the pick is credited through the expected-wait selector + /// like any other. + fn warmup_slice( + &self, + workers: &[Arc], + table: &PoolTable, + indexer: &KvIndex, + info: &SelectWorkerInfo, + ) -> Option { + if table.warming.is_empty() { + return None; + } + let warmup = liveness::warmup(); + if warmup.share <= 0.0 { + return None; + } + let mut pick: Option<(usize, usize)> = None; + let mut tied = 0u32; + let mut rng = rand::rng(); + for &position in &table.warming { + let idx = position as usize; + let Some(worker) = workers.get(idx) else { + continue; + }; + let state = worker.routing_state(); + if !state.eligible() { + continue; + } + let indexed = table.id_at[idx].map(|id| indexer.worker_block_count(id)); + let grown = indexed.map(|indexed| worker.warmup_growth(indexed)); + if !warmup.applies( + worker.admitted_age(), + grown, + indexed.unwrap_or(0), + table.fleet_level, + ) { + continue; + } + match pick { + Some((_, load)) if state.load > load => {} + Some((_, load)) if state.load == load => { + // Equal loads draw uniformly: an idle warming fleet must + // not hand every slice to its lowest position. + tied += 1; + if rng.random_range(0..=tied) == 0 { + pick = Some((idx, state.load)); + } + } + _ => { + pick = Some((idx, state.load)); + tied = 0; + } + } + } + let (idx, _) = pick?; + if !self + .warmup_misses + .fetch_add(1, Ordering::Relaxed) + .is_multiple_of(warmup.period()) + { + return None; + } + self.select_expected_wait(workers, &[idx], info) + } + + /// The thin worker's share of hits. On a replay where every request has + /// a holder (the Mooncake trace: a system prompt or a chat template in + /// front of everything) the miss path never runs, so a worker whose + /// index a resync emptied gets nothing from the slice and idles for good + /// (the churn runs c1 and c2, the soaks s6 and s7: one worker lost per + /// publisher restart, for half an hour). One hit in + /// `Warmup::divert_every` therefore goes to the least-loaded thin worker + /// although another holds its prefix, the recompute accepted: shallow + /// overlaps first (at most [`DIVERT_SHALLOW_SHARE`] of the request), any + /// hit once a window of `divert_every` hits passed without a shallow + /// one; never while the thin worker has anything in flight unless its + /// last diversion is older than [`DIVERT_WINDOW_MS`] (the stamp lives on + /// the worker, so it survives the table's rebuilds), so a hot fleet is + /// not disturbed; and only until its index crosses the thinness ratio of + /// the fleet's level. Equal loads draw uniformly like the slice. A fleet + /// with no thin worker returns before any clock read or counter: this + /// costs the steady state nothing. Protection and the liveness vetoes + /// apply through the routing state as everywhere; the pick is credited + /// through the expected-wait selector like the slice's. + fn warmup_divert( + &self, + workers: &[Arc], + table: &PoolTable, + indexer: &KvIndex, + info: &SelectWorkerInfo, + overlap_share: f64, + ) -> Option { + if table.thin.is_empty() { + return None; + } + let warmup = liveness::warmup(); + if warmup.divert_every == 0 || warmup.share <= 0.0 { + return None; + } + let shallow = overlap_share <= DIVERT_SHALLOW_SHARE; + if shallow { + self.divert_shallow_seen.store(true, Ordering::Relaxed); + } + let hits = self.divert_hits.fetch_add(1, Ordering::Relaxed) + 1; + let shallow_seen = self.divert_shallow_seen.load(Ordering::Relaxed); + let due = (hits >= warmup.divert_every && (shallow || !shallow_seen)) + || hits >= 2 * warmup.divert_every; + if !due { + return None; + } + let now = liveness::now_ms(); + let mut pick: Option<(usize, usize)> = None; + let mut tied = 0u32; + let mut rng = rand::rng(); + for &position in &table.thin { + let idx = position as usize; + let Some(worker) = workers.get(idx) else { + continue; + }; + let state = worker.routing_state(); + if !state.eligible() { + continue; + } + let indexed = table.id_at[idx].map_or(0, |id| indexer.worker_block_count(id)); + if !warmup.is_thin(indexed, table.fleet_level) { + continue; + } + if state.load > 0 && now < worker.divert_until_ms() { + continue; + } + match pick { + Some((_, load)) if state.load > load => {} + Some((_, load)) if state.load == load => { + tied += 1; + if rng.random_range(0..=tied) == 0 { + pick = Some((idx, state.load)); + } + } + _ => { + pick = Some((idx, state.load)); + tied = 0; + } + } + } + let (idx, _) = pick?; + self.divert_hits.store(0, Ordering::Relaxed); + self.divert_shallow_seen.store(false, Ordering::Relaxed); + workers[idx].note_diverted(now + DIVERT_WINDOW_MS); + self.select_expected_wait(workers, &[idx], info) + } + /// Resolve an affinity score group to one final worker. Safe affinity /// candidates retain priority; only when every tied holder trips the /// pressure gate do we scan the healthy fleet for non-gated spill targets. @@ -1367,9 +1983,10 @@ impl CacheAwarePolicy { return self.select_expected_wait(workers, &safe_affinity, info); } - // A miss has no affinity candidates and is not a spill: every healthy - // worker participates. A true spill excludes other workers that trip - // the same gate, preventing a fallback from reselecting a hot holder. + // A miss has no affinity candidates and is not a spill: every + // eligible worker the decision read participates. A true spill + // excludes other workers that trip the same gate, preventing a + // fallback from reselecting a hot holder. if affinity_candidates.is_empty() { return self.select_expected_wait(workers, healthy_indices, info); } @@ -1386,11 +2003,139 @@ impl CacheAwarePolicy { self.select_expected_wait(workers, candidates, info) } + /// Run the selection policy over this request's inputs and resolve its + /// pick with the host's gate, expected-wait selector and credit. + /// + /// `candidates` are the positive-overlap workers with their decayed + /// scores. Policies that also want cold workers or backend loads declare + /// it in their `Needs`; the default policy declares neither, so with no + /// accounting its hot path gathers exactly what the affinity-group + /// decision did and reaches the same `select_final_from_affinity` call. + /// + /// - `Pick::Group`: the affinity group, resolved as before (pressure + /// gate, then expected wait, which credits the result). + /// - `Pick::None`: the miss path (expected wait over the healthy fleet). + /// - `Pick::Final`: that worker, credited through the expected-wait + /// selector; the policy owns the load trade-off, so the count-pressure + /// gate does not apply. If the waiting-queue veto drops it, the fleet + /// fallback runs. + fn resolve_selection( + &self, + workers: &[Arc], + healthy_indices: &[usize], + candidates: &[OverlapCandidate], + request: &RequestInputs<'_>, + avg_load: f64, + info: &SelectWorkerInfo, + ) -> Option { + let needs = self.selection.needs(); + let accounting = self.accounting.as_ref(); + let wants_all = needs.all_workers || accounting.is_some(); + if candidates.is_empty() && !wants_all { + return self.select_final_from_affinity(workers, &[], healthy_indices, avg_load, info); + } + + // With `all_workers` the rows are the eligible workers the decision + // read merged with the holders, in slice order (a sampled decision + // reads eligibility for a bounded sample of the pool beside the + // holders, so the two lists differ); otherwise the candidates are + // the rows and no list is built. + let all_rows: Vec = if wants_all { + Self::merge_rows(healthy_indices, candidates) + } else { + Vec::new() + }; + let predicted = match (accounting, request.prefix_hashes) { + (Some(accounting), Some(hashes)) => accounting.predicted_overlaps(hashes), + _ => Vec::new(), + }; + let gather = |idx: usize, raw: f64, effective: f64| { + let url = workers[idx].url(); + let predicted_blocks = predicted + .iter() + .find(|(predicted_url, _)| &**predicted_url == url) + .map_or(0.0, |(_, blocks)| *blocks); + // A prediction deeper than the index's view stands in for both + // scores: the blocks are expected to be resident by the time the + // request lands, so no decay is applied to them. + let (device_blocks, effective_score) = if predicted_blocks > raw { + (predicted_blocks, predicted_blocks) + } else { + (raw, effective) + }; + CandidateInputs { + idx, + url, + device_blocks, + effective_score, + } + }; + let inputs: Vec> = if wants_all { + // Rows and candidates are both in slice order, so one merge pass + // pairs them. + let mut next = candidates.iter().peekable(); + all_rows + .iter() + .map(|&idx| { + let (raw, effective) = next + .next_if(|candidate| candidate.idx == idx) + .map_or((0.0, 0.0), |candidate| { + (candidate.raw_score, candidate.effective_score) + }); + gather(idx, raw, effective) + }) + .collect() + } else { + candidates + .iter() + .map(|candidate| { + gather( + candidate.idx, + candidate.raw_score, + candidate.effective_score, + ) + }) + .collect() + }; + let selected = match self.selection.select(request, &inputs) { + Pick::None => { + self.select_final_from_affinity(workers, &[], healthy_indices, avg_load, info) + } + Pick::Group(rows) => { + let group: Vec = rows.iter().map(|&row| inputs[row].idx).collect(); + self.select_final_from_affinity(workers, &group, healthy_indices, avg_load, info) + } + Pick::Final(row) => { + let idx = inputs[row].idx; + self.select_expected_wait(workers, &[idx], info) + .or_else(|| self.select_expected_wait(workers, healthy_indices, info)) + } + }?; + + if let Some(dispatched) = inputs.iter().find(|candidate| candidate.idx == selected) { + self.selection.on_dispatch(request, dispatched); + if let Some(accounting) = accounting { + let uncached = dispatched.uncached_prompt_tokens(request) as u64; + accounting.record_dispatch( + dispatched.url, + uncached, + request.prefix_hashes.unwrap_or(&[]), + ); + } + } + Some(selected) + } + /// Pick an effective-affinity score group while preserving temperature. /// At temperature zero this is the exact maximum. With temperature, the /// existing softmax samples one worker; expanding that draw back to every /// equal-score worker preserves each score group's aggregate probability, /// then LeastLoad breaks the tie inside the sampled group. + /// + /// Kept as the reference for the default selection policy, which must + /// reproduce it (see the pin tests); production goes through + /// `resolve_selection`. + #[cfg(test)] fn affinity_score_group( candidates: &[OverlapCandidate], selection_temperature: f32, @@ -1428,16 +2173,20 @@ impl CacheAwarePolicy { workers: &[Arc], healthy_indices: &[usize], matched_tenants: &[TenantId], - request_units: usize, - avg_load: f64, + matched_units: usize, + request: &RequestInputs<'_>, info: &SelectWorkerInfo, ) -> Option { + let avg_load = request.avg_load; + let request_units = request.prompt_tokens; + let matched_blocks = (matched_units / self.config.block_size.max(1)) as f64; let mut candidates: Vec = Vec::new(); for &idx in healthy_indices { let url = workers[idx].url(); if matched_tenants.iter().any(|tenant| tenant.as_ref() == url) { candidates.push(OverlapCandidate { idx, + raw_score: matched_blocks, effective_score: 1.0, }); } @@ -1456,18 +2205,17 @@ impl CacheAwarePolicy { self.config.block_size, &tuning, ); - let affinity_candidates = - Self::affinity_score_group(&candidates, tuning.selection_temperature); - self.select_final_from_affinity( + self.resolve_selection( workers, - &affinity_candidates, healthy_indices, + &candidates, + request, avg_load, info, ) } - /// Event-driven routing: PositionalIndexer overlap scoring (Type 1). + /// Event-driven routing: KV index overlap scoring (Type 1). /// /// Self-contained — when overlap is found, selects the worker with the best /// cache match. When no overlap (cold start, novel tokens, short request), @@ -1495,101 +2243,343 @@ impl CacheAwarePolicy { waiting_prefill_tokens: waiting_prefill_tokens.as_deref(), }; - let candidates = Self::overlap_candidates( - workers, - tokens, - healthy_indices, + // The engines fold the LoRA name and cache salt into their block + // hashes and the monitor recomputes stored blocks under the same + // seed, so a request in a namespace is hashed under it and matches + // only its own blocks; a plain request keeps the plain hash. + let content_hashes = match info.cache_namespace { + Some(namespace) => { + request_content_hashes_with_seed(tokens, block_size, namespace.event_seed()) + } + None => compute_request_content_hashes(tokens, block_size), + }; + let table = self.pool_table(workers, indexer, liveness::now_ms()); + // Over a sampled pool the candidates read are bounded as the pool + // is: a tie among the deepest holders is drawn from, and an + // all-workers policy ranks the sample and the deepest holders; a + // shallow shared prefix is not a reason to read a thousand workers. + let cap = (workers.len() > 2 * FLEET_SAMPLE).then_some(FLEET_SAMPLE); + let gather = if self.exact_maximum_decision(&tuning) { + Gather::TopGroup { cap } + } else { + let wants_all = self.selection.needs().all_workers || self.accounting.is_some(); + Gather::All { + cap: cap.filter(|_| wants_all), + } + }; + let candidates = Self::overlap_candidates_for_hashes( + workers, + &content_hashes, + &table, indexer, block_size, &tuning, + gather, ); - let affinity_candidates = - Self::affinity_score_group(&candidates, tuning.selection_temperature); - if !affinity_candidates.is_empty() { - let idx = self.select_final_from_affinity( - workers, - &affinity_candidates, - healthy_indices, - avg_load, - info, - )?; + // Chain hashes are computed only for policies (or accounting) keyed + // on prefixes; the default decision never pays for them. + let prefix_hashes: Vec = + if self.selection.needs().prefix_hashes || self.accounting.is_some() { + request_prefix_hashes(&content_hashes) + .into_iter() + .map(|hash| hash.0) + .collect() + } else { + Vec::new() + }; + let request = RequestInputs { + prompt_tokens: tokens.len(), + block_size, + request_blocks: content_hashes.len().max(1), + avg_load, + prefix_hashes: (!prefix_hashes.is_empty()).then_some(prefix_hashes.as_slice()), + }; + // A miss for the warm-up slice is no overlap or a thin one: chat requests + // share their template's head with every holder, which is affinity in + // name only. Thin is at most the tree-mode `cache_threshold` share of + // the request, or a few blocks outright when they are under half of + // it (the head of a short request; a short request cached whole is a + // hit and stays with its holder). + let best_overlap = candidates + .iter() + .map(|candidate| candidate.raw_score) + .fold(0.0_f64, f64::max); + let request_blocks = request.request_blocks as f64; + let thin_overlap = best_overlap / request_blocks <= f64::from(self.config.cache_threshold) + || (best_overlap <= WARMUP_MISS_BLOCKS && best_overlap * 2.0 < request_blocks); + if thin_overlap { + if let Some(idx) = self.warmup_slice(workers, &table, indexer, info) { + Metrics::record_worker_cache_aware_policy_branch("warmup_slice"); + debug!( + worker = workers[idx].url(), + branch = "warmup_slice", + request_blocks = content_hashes.len(), + "Cache miss routed to a warming worker" + ); + return Some(idx); + } + } else if let Some(idx) = self.warmup_divert( + workers, + &table, + indexer, + info, + best_overlap / request_blocks, + ) { + Metrics::record_worker_cache_aware_policy_branch("warmup_divert"); debug!( worker = workers[idx].url(), - branch = if affinity_candidates.contains(&idx) { - "event_hit" - } else { - "event_spill" - }, - model_id, - "Event-driven routing: overlap match" + branch = "warmup_divert", + overlap_blocks = best_overlap as u64, + request_blocks = content_hashes.len(), + "Cache hit diverted to a thin worker" ); return Some(idx); } - - // No cache overlap — expected-wait fallback over the healthy fleet. - let selected = self.select_expected_wait(workers, healthy_indices, info)?; + let had_overlap = !candidates.is_empty(); + let idx = self.resolve_selection( + workers, + healthy_indices, + &candidates, + &request, + avg_load, + info, + )?; + // `overlap_blocks` is the chosen worker's undecayed overlap, so a + // per-decision join against the engine's `cached_tokens` compares + // blocks and not only the branch. + let overlap_blocks = candidates + .iter() + .find(|candidate| candidate.idx == idx) + .map_or(0, |candidate| candidate.raw_score as u64); + let branch = if !had_overlap { + "event_miss" + } else if overlap_blocks > 0 { + "event_hit" + } else { + "event_spill" + }; + Metrics::record_worker_cache_aware_policy_branch(branch); debug!( - worker = workers[selected].url(), - model_id, "Event-driven routing: no overlap, expected-wait fallback" + worker = workers[idx].url(), + branch, + overlap_blocks, + request_blocks = content_hashes.len(), + policy = self.selection.name(), + model_id, + "Event-driven routing" ); - Some(selected) + Some(idx) } /// Build positive-overlap candidates for event-driven routing. /// - /// Returns each healthy worker with a positive, optionally decayed overlap - /// score. This helper neither selects nor credits a worker: + /// Returns each eligible worker with a positive, optionally decayed + /// overlap score. This helper neither selects nor credits a worker: /// `affinity_score_group` chooses the effective-score group (exact maximum /// at temperature zero; a softmax-sampled group otherwise), and /// `select_final_from_affinity` uses LeastLoad to choose and credit one - /// final worker from that group. An empty result means no full-block overlap. + /// final worker from that group. An empty result means no full-block + /// overlap. The tests name the eligible set themselves. + #[cfg(test)] fn overlap_candidates( workers: &[Arc], tokens: &[u32], healthy_indices: &[usize], - indexer: &PositionalIndexer, + indexer: &KvIndex, block_size: usize, tuning: &OverlapTuning<'_>, ) -> Vec { let content_hashes = compute_request_content_hashes(tokens, block_size); + let table = PoolTable::build(workers, indexer, liveness::now_ms()); + let mut candidates = Self::overlap_holders( + workers, + &content_hashes, + &table, + indexer, + Gather::All { cap: None }, + ); + candidates.retain(|candidate| healthy_indices.contains(&candidate.idx)); + if !candidates.is_empty() { + Self::apply_overlap_decay( + workers, + &mut candidates, + content_hashes.len(), + block_size, + tuning, + ); + } + candidates + } + + /// `overlap_candidates` over already-computed block hashes. + fn overlap_candidates_for_hashes( + workers: &[Arc], + content_hashes: &[ContentHash], + table: &PoolTable, + indexer: &KvIndex, + block_size: usize, + tuning: &OverlapTuning<'_>, + gather: Gather, + ) -> Vec { + let mut candidates = Self::overlap_holders(workers, content_hashes, table, indexer, gather); + if candidates.is_empty() { + return candidates; + } + + Self::apply_overlap_decay( + workers, + &mut candidates, + content_hashes.len(), + block_size, + tuning, + ); + + candidates + } + + /// The eligible holders of the request's blocks with their undecayed + /// overlap, in slice order: the index names them by id, the pool table + /// places them, and eligibility is read for them alone, so the cost + /// follows the holders gathered and not the pool. + fn overlap_holders( + workers: &[Arc], + content_hashes: &[ContentHash], + table: &PoolTable, + indexer: &KvIndex, + gather: Gather, + ) -> Vec { if content_hashes.is_empty() { return Vec::new(); } - let overlap = indexer.find_matches(&content_hashes, false); + let started = Instant::now(); + let overlap = indexer.find_matches(content_hashes, false); + Metrics::record_kv_index_lookup(indexer.name(), started.elapsed().as_secs_f64()); if overlap.scores.is_empty() { return Vec::new(); } - // Gather the positive-overlap candidates once; both selection modes - // and the decay's fleet-floor computation need the full set. - let mut candidates: Vec = Vec::new(); - for &idx in healthy_indices { - let Some(score) = indexer - .worker_id(workers[idx].url()) - .and_then(|id| overlap.scores.get(&id)) - .copied() - .filter(|&s| s > 0) - else { + let mut candidates = match gather { + Gather::TopGroup { cap } => Self::top_group(workers, table, &overlap, cap), + Gather::All { cap } => Self::all_holders(workers, table, &overlap, cap), + }; + // Downstream reads candidates in slice order: the all-workers merge + // in `resolve_selection` and the tie-breaks. + candidates.sort_unstable_by_key(|candidate| candidate.idx); + candidates + } + + /// The eligible holders of the deepest overlap: one pass over the + /// scores for the depth, one to place and read those holders. When none + /// of them is in this pool and routable (a PD leg's pool sharing the + /// model's index with the other leg, a holder down), the full gather + /// decides, so the result is always the deepest group a scan of the pool + /// would have found. + fn top_group( + workers: &[Arc], + table: &PoolTable, + overlap: &OverlapScores, + cap: Option, + ) -> Vec { + let deepest = overlap.scores.values().copied().max().unwrap_or(0); + if deepest == 0 { + return Vec::new(); + } + let group = Self::holders_at(workers, table, overlap, deepest, cap); + if !group.is_empty() { + return group; + } + let mut all = Self::all_holders(workers, table, overlap, None); + let best = all + .iter() + .map(|candidate| candidate.raw_score) + .fold(0.0_f64, f64::max); + all.retain(|candidate| candidate.raw_score == best); + Self::draw(&mut all, cap); + all + } + + /// The eligible holders in this pool whose overlap is exactly `depth` + /// blocks; `cap` of the tied holders, drawn uniformly, when they are more. + fn holders_at( + workers: &[Arc], + table: &PoolTable, + overlap: &OverlapScores, + depth: u32, + cap: Option, + ) -> Vec { + let mut tied: Vec = overlap + .scores + .iter() + .filter(|(_, &score)| score == depth) + .map(|(&id, _)| id) + .collect(); + Self::draw(&mut tied, cap); + let mut group = Vec::with_capacity(tied.len()); + for id in tied { + let Some(idx) = table.position(id, workers) else { + continue; + }; + if !workers[idx].routing_state().eligible() { + continue; + } + group.push(OverlapCandidate { + idx, + raw_score: f64::from(depth), + effective_score: f64::from(depth), + }); + } + group + } + + /// Keep a uniform draw of `cap` items when there are more. + fn draw(items: &mut Vec, cap: Option) { + let Some(cap) = cap else { + return; + }; + if items.len() <= cap { + return; + } + let mut rng = rand::rng(); + for i in 0..cap { + let j = rng.random_range(i..items.len()); + items.swap(i, j); + } + items.truncate(cap); + } + + /// Every eligible holder in this pool with a positive overlap; the `cap` + /// deepest of them when set. + fn all_holders( + workers: &[Arc], + table: &PoolTable, + overlap: &OverlapScores, + cap: Option, + ) -> Vec { + let mut candidates: Vec = Vec::with_capacity(overlap.scores.len()); + for (&id, &score) in &overlap.scores { + if score == 0 { + continue; + } + let Some(idx) = table.position(id, workers) else { continue; }; + if !workers[idx].routing_state().eligible() { + continue; + } candidates.push(OverlapCandidate { idx, + raw_score: f64::from(score), effective_score: f64::from(score), }); } - if candidates.is_empty() { - return candidates; + if let Some(cap) = cap { + if candidates.len() > cap { + candidates + .select_nth_unstable_by(cap - 1, |a, b| b.raw_score.total_cmp(&a.raw_score)); + candidates.truncate(cap); + } } - - Self::apply_overlap_decay( - workers, - &mut candidates, - content_hashes.len(), - block_size, - tuning, - ); - candidates } @@ -1639,6 +2629,7 @@ impl CacheAwarePolicy { /// or 2000. The best candidate's exponent is exactly 0 (overflow-safe); /// a degenerate spread (all equal) is a uniform draw. Inverse-CDF /// sampling with a last-row fallback against floating-point drift. + #[cfg(test)] fn sample_by_temperature(candidates: &[OverlapCandidate], temperature: f32) -> Option { let first = candidates.first()?; let (min, max) = candidates.iter().fold( @@ -1701,6 +2692,15 @@ impl CacheAwarePolicy { }; Metrics::record_worker_cache_aware_policy_branch(branch); Metrics::record_cache_aware_match_ratio(histogram_ratio); + // `credited_units` is what the served worker is expected to have + // cached: the matched prefix only when it serves the request. On the + // fallback and spill branches the match belongs to another tenant, so + // a join against the engine's cached tokens must not count it. + let credited_units = if branch == "tree_match" { + matched_units + } else { + 0 + }; debug!( index = "tree", branch, @@ -1708,6 +2708,9 @@ impl CacheAwarePolicy { model_id, matched_ratio = f64::from(matched_ratio), threshold = f64::from(self.config.cache_threshold), + matched_units, + input_units, + credited_units, "Cache-aware selection" ); } @@ -2015,17 +3018,18 @@ impl CacheAwarePolicy { } else { matched as f32 / input as f32 }; + let request = self.tree_request(tokens.len(), avg_load); selected_idx = if match_rate > self.config.cache_threshold { self.select_matched_candidate( workers, healthy_indices, &result.matched_tenants, - tokens.len(), - avg_load, + matched, + &request, info, ) } else { - self.select_final_from_affinity(workers, &[], healthy_indices, avg_load, info) + self.resolve_selection(workers, healthy_indices, &[], &request, avg_load, info) }; selected_idx.map(|idx| workers[idx].url()) }); @@ -2092,17 +3096,18 @@ impl CacheAwarePolicy { } else { matched as f32 / input as f32 }; + let request = self.tree_request(input, avg_load); selected_idx = if match_rate > self.config.cache_threshold { self.select_matched_candidate( workers, healthy_indices, &result.matched_tenants, - input, - avg_load, + matched, + &request, info, ) } else { - self.select_final_from_affinity(workers, &[], healthy_indices, avg_load, info) + self.resolve_selection(workers, healthy_indices, &[], &request, avg_load, info) }; selected_idx.map(|idx| workers[idx].url()) }); @@ -2147,7 +3152,7 @@ impl Default for CacheAwarePolicy { #[cfg(test)] mod tests { - use kv_index::{compute_content_hash, SequenceHash, StoredBlock, WorkerBlockMap}; + use kv_index::{compute_content_hash, SequenceHash, StoredBlock}; use metrics_exporter_prometheus::{Matcher, PrometheusBuilder}; use openai_protocol::worker::{ HealthCheckConfig, SchedulerLoadSnapshot, WorkerLoadResponse, WorkerStatus, @@ -2173,7 +3178,7 @@ mod tests { workers: &[Arc], tokens: &[u32], healthy_indices: &[usize], - indexer: &PositionalIndexer, + indexer: &KvIndex, block_size: usize, tuning: &OverlapTuning<'_>, ) -> Vec { @@ -2189,7 +3194,7 @@ mod tests { } use crate::{ observability::metrics::CACHE_AWARE_MATCH_RATIO_BUCKETS, - worker::{BasicWorkerBuilder, WorkerType}, + worker::{BasicWorkerBuilder, WorkerBlocks, WorkerType}, }; fn no_health_check() -> HealthCheckConfig { @@ -3855,19 +4860,19 @@ mod tests { } // ----------------------------------------------------------------------- - // Event-driven routing tests (Type 1: PositionalIndexer overlap scoring) + // Event-driven routing tests (Type 1: KV index overlap scoring) // ----------------------------------------------------------------------- - /// Helper: create a PositionalIndexer and store blocks for a worker. + /// Helper: create a positional KV index and store blocks for a worker. /// `token_chunks` is a list of token-id slices — each becomes one block. fn setup_indexer_with_blocks( worker_url: &str, token_chunks: &[&[u32]], jump_size: usize, - ) -> Arc { - let indexer = Arc::new(PositionalIndexer::new(jump_size)); + ) -> Arc { + let indexer = Arc::new(KvIndex::positional(jump_size)); let worker_id = indexer.intern_worker(worker_url).unwrap(); - let mut wb = WorkerBlockMap::default(); + let mut wb = WorkerBlocks::default(); let blocks: Vec = token_chunks .iter() .enumerate() @@ -3958,7 +4963,7 @@ mod tests { let indexer = setup_indexer_with_blocks("http://w1:8000", &chunks, 4); // Same content cached on w2 under distinct backend seq hashes. let w2 = indexer.intern_worker("http://w2:8000").unwrap(); - let mut wb2 = WorkerBlockMap::default(); + let mut wb2 = WorkerBlocks::default(); let blocks: Vec = chunks .iter() .enumerate() @@ -3984,7 +4989,7 @@ mod tests { /// Two workers with identical cached blocks (the tie-test topology): both /// fully match the request. - fn equal_overlap_fixture() -> (Vec>, Arc) { + fn equal_overlap_fixture() -> (Vec>, Arc) { let workers: Vec> = vec![ Arc::new( BasicWorkerBuilder::new("http://w1:8000") @@ -4002,7 +5007,7 @@ mod tests { let chunks: [&[u32]; 2] = [&[1, 2, 3, 4], &[5, 6, 7, 8]]; let indexer = setup_indexer_with_blocks("http://w1:8000", &chunks, 4); let w2 = indexer.intern_worker("http://w2:8000").unwrap(); - let mut wb2 = WorkerBlockMap::default(); + let mut wb2 = WorkerBlocks::default(); let blocks: Vec = chunks .iter() .enumerate() @@ -4074,7 +5079,7 @@ mod tests { } /// w1 caches both request blocks (score 2), w2 only the first (score 1). - fn unequal_overlap_fixture() -> (Vec>, Arc) { + fn unequal_overlap_fixture() -> (Vec>, Arc) { let workers: Vec> = vec![ Arc::new( BasicWorkerBuilder::new("http://w1:8000") @@ -4092,7 +5097,7 @@ mod tests { let indexer = setup_indexer_with_blocks("http://w1:8000", &[&[1, 2, 3, 4], &[5, 6, 7, 8]], 4); let w2 = indexer.intern_worker("http://w2:8000").unwrap(); - let mut wb2 = WorkerBlockMap::default(); + let mut wb2 = WorkerBlocks::default(); let blocks = vec![StoredBlock { seq_hash: SequenceHash(100), content_hash: compute_content_hash(&[1, 2, 3, 4]), @@ -4218,11 +5223,11 @@ mod tests { policy.init_workers(&workers); // Store same blocks for both workers (equal overlap) - let indexer = Arc::new(PositionalIndexer::new(4)); + let indexer = Arc::new(KvIndex::positional(4)); let w1_id = indexer.intern_worker("http://w1:8000").unwrap(); let w2_id = indexer.intern_worker("http://w2:8000").unwrap(); - let mut wb1 = WorkerBlockMap::default(); - let mut wb2 = WorkerBlockMap::default(); + let mut wb1 = WorkerBlocks::default(); + let mut wb2 = WorkerBlocks::default(); let blocks = vec![StoredBlock { seq_hash: SequenceHash(1), content_hash: compute_content_hash(&[1, 2, 3, 4]), @@ -4270,11 +5275,11 @@ mod tests { ]; policy.init_workers(&workers); - let indexer = Arc::new(PositionalIndexer::new(4)); + let indexer = Arc::new(KvIndex::positional(4)); let w1_id = indexer.intern_worker("http://w1:8000").unwrap(); let w2_id = indexer.intern_worker("http://w2:8000").unwrap(); - let mut wb1 = WorkerBlockMap::default(); - let mut wb2 = WorkerBlockMap::default(); + let mut wb1 = WorkerBlocks::default(); + let mut wb2 = WorkerBlocks::default(); // Both workers have block [1,2,3,4] (equal overlap, equal load) let block = vec![StoredBlock { @@ -4350,11 +5355,11 @@ mod tests { ]; policy.init_workers(&workers); - let indexer = Arc::new(PositionalIndexer::new(4)); + let indexer = Arc::new(KvIndex::positional(4)); let w1_id = indexer.intern_worker("http://w1:8000").unwrap(); let w2_id = indexer.intern_worker("http://w2:8000").unwrap(); - let mut wb1 = WorkerBlockMap::default(); - let mut wb2 = WorkerBlockMap::default(); + let mut wb1 = WorkerBlocks::default(); + let mut wb2 = WorkerBlocks::default(); // w1 has 4 blocks cached let blocks_w1: Vec = (0..4) @@ -4391,43 +5396,384 @@ mod tests { // Query with all 4 blocks worth of tokens → w1 wins (higher overlap: 4 vs 2) let result = overlap_affinity_group( &workers, - &[1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16], - &[0, 1], + &[1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16], + &[0, 1], + &indexer, + 4, + &default_tuning(), + ); + assert_eq!(result, vec![0]); // w1 (higher overlap) + } + + // -- select_worker_event_driven integration tests -- + + #[test] + fn event_unique_hit_keeps_affinity_and_credits_expected_wait() { + let policy = CacheAwarePolicy::with_config(test_config()); + let workers = make_workers(&["http://w1:8000", "http://w2:8000"]); + policy.init_workers(&workers); + workers[1].increment_load(); + update_expected_wait_loads(&policy, &workers, &[10_000, 0]); + + let monitor = Arc::new(KvEventMonitor::new(Some(4))); + let indexer = + setup_indexer_with_blocks("http://w1:8000", &[&[1, 2, 3, 4], &[5, 6, 7, 8]], 4); + monitor.indexers.insert("unknown".to_string(), indexer); + policy.set_kv_event_monitor(Some(monitor)); + + let cached = [1, 2, 3, 4, 5, 6, 7, 8]; + assert_eq!( + policy.select_worker( + &workers, + &SelectWorkerInfo { + tokens: Some(&cached), + ..Default::default() + } + ), + Some(0) + ); + assert_only_final_worker_credited(&policy, &workers, 0, 8); + } + + #[test] + fn a_worker_with_an_emptied_index_gets_the_warm_up_slice_whatever_its_age() { + // The soaks' case in miniature: w1 holds a fleet-sized cache (1,100 + // blocks), w2's index was cleared by a resync (nothing indexed), and + // every request shares a head with w1, so affinity alone would never + // send w2 a request again. The slice does, on the thin overlap: w2 is + // thin against the fleet's level whatever its age. + let policy = CacheAwarePolicy::with_config(test_config()); + let workers = make_workers(&["http://w1:8000", "http://w2:8000"]); + policy.init_workers(&workers); + update_expected_wait_loads(&policy, &workers, &[0, 0]); + let held: Vec = (1..=4_400).collect(); + let monitor = Arc::new(KvEventMonitor::new(Some(4))); + let indexer = setup_indexer_with_blocks("http://w1:8000", &[&held], 4); + monitor.indexers.insert("unknown".to_string(), indexer); + policy.set_kv_event_monitor(Some(monitor)); + // w1's first block (the shared head) followed by eleven novel blocks. + let mut request: Vec = (1..=4).collect(); + request.extend(90_000..90_044); + let picks: Vec = (0..8) + .map(|_| { + policy + .select_worker( + &workers, + &SelectWorkerInfo { + tokens: Some(&request), + ..Default::default() + }, + ) + .unwrap() + }) + .collect(); + assert!( + picks.contains(&1), + "the emptied worker got its slice of the thin-overlap requests: {picks:?}" + ); + assert!(picks.contains(&0), "the holder kept the rest: {picks:?}"); + } + + /// Store `tokens` for `worker` as blocks of `block` tokens, with sequence + /// hashes from `seq_base` (distinct per call). + fn store_blocks(indexer: &KvIndex, worker: u32, tokens: &[u32], block: usize, seq_base: u64) { + let mut wb = WorkerBlocks::default(); + let blocks: Vec = tokens + .chunks(block) + .enumerate() + .map(|(i, chunk)| StoredBlock { + seq_hash: SequenceHash(seq_base + i as u64), + content_hash: compute_content_hash(chunk), + }) + .collect(); + indexer + .apply_stored(worker, &blocks, None, &mut wb) + .unwrap(); + } + + /// A fleet for the diversion tests: the policy, its eight workers, the + /// index they share and each holder's head (its first 240 tokens, 60 + /// blocks); a request made of a head and a four-block tail of its own is + /// a deep hit on that holder (60 of 64 blocks) that no other request + /// repeats. + struct HolderFleet { + policy: CacheAwarePolicy, + workers: Vec>, + indexer: Arc, + heads: Vec>, + } + + /// Eight workers; the first `holders` hold 1,100 blocks each of their own + /// 4,400-token sequence, the rest nothing. + fn fleet_with_holders(holders: usize) -> HolderFleet { + fleet_with_holders_of(holders, 1_100) + } + + /// Eight workers; the first `holders` hold `blocks` blocks each of their + /// own sequence, the rest nothing. + fn fleet_with_holders_of(holders: usize, blocks: usize) -> HolderFleet { + let policy = CacheAwarePolicy::with_config(test_config()); + let urls: Vec = (0..8).map(|i| format!("http://w{i}:8000")).collect(); + let refs: Vec<&str> = urls.iter().map(String::as_str).collect(); + let workers = make_workers(&refs); + policy.init_workers(&workers); + update_expected_wait_loads(&policy, &workers, &[0, 0, 0, 0, 0, 0, 0, 0]); + let indexer = Arc::new(KvIndex::positional(4)); + let mut prefixes = Vec::new(); + for (h, url) in urls.iter().enumerate() { + let id = indexer.intern_worker(url).unwrap(); + if h >= holders { + continue; + } + let tokens: Vec = (0..blocks as u32 * 4) + .map(|i| (h as u32 + 1) * 100_000 + i) + .collect(); + store_blocks(&indexer, id, &tokens, 4, (h as u64 + 1) * 1_000_000); + prefixes.push(tokens[..240].to_vec()); + } + let monitor = Arc::new(KvEventMonitor::new(Some(4))); + monitor + .indexers + .insert("unknown".to_string(), Arc::clone(&indexer)); + policy.set_kv_event_monitor(Some(monitor)); + HolderFleet { + policy, + workers, + indexer, + heads: prefixes, + } + } + + #[test] + fn an_emptied_worker_gets_its_share_of_hits_until_it_refills() { + // The churn runs' case: every request is cached whole on one of seven + // holders (a deep hit, never a miss), the eighth worker's index was + // emptied by a resync. Affinity alone never sends it anything; the + // diversion hands it one hit in eight (no shallow hit ever comes, so + // the fallback at the window's end), and each diverted request lands + // in its index, until it crosses half the fleet's level (1,100 -> 550). + let HolderFleet { + policy, + workers, + indexer, + heads, + } = fleet_with_holders(7); + let emptied = indexer.worker_id("http://w7:8000").unwrap(); + // A request: a holder's 60-block head and a four-block tail of its own. + let request = |i: usize| -> Vec { + let mut tokens = heads[i % 7].clone(); + tokens.extend((0..16).map(|t| 90_000_000 + i as u32 * 16 + t)); + tokens + }; + let mut decisions = 0usize; + let mut first = None; + let mut received = 0usize; + while decisions < 1_000 && indexer.worker_block_count(emptied) < 550 { + let tokens = request(decisions); + let idx = policy + .select_worker( + &workers, + &SelectWorkerInfo { + tokens: Some(&tokens), + ..Default::default() + }, + ) + .unwrap(); + decisions += 1; + if idx == 7 { + received += 1; + first.get_or_insert(decisions); + // The request prefills on w7: its blocks join w7's index. + store_blocks( + &indexer, + emptied, + &tokens, + 4, + 9_000_000 + decisions as u64 * 100, + ); + } + } + assert!( + first.is_some_and(|first| first <= 8), + "the emptied worker is served within the first eight hits, was {first:?}" + ); + assert!( + indexer.worker_block_count(emptied) >= 550, + "refilled to the ratio: {} blocks after {decisions} decisions", + indexer.worker_block_count(emptied) + ); + assert!( + decisions <= 40 * 8 + 16, + "about forty requests of 64 blocks (seven heads of 60, then tails of 4) at one in eight: \ + {decisions} decisions, {received} received" + ); + // At the ratio the diversion stops: the hit counter runs on without a + // reset (w7 now holds the heads too and may win a tie as a holder, + // which is affinity, not a diversion). + let hits_before = policy.divert_hits.load(Ordering::Relaxed); + for i in 0..32 { + let tokens = request(decisions + i); + policy + .select_worker( + &workers, + &SelectWorkerInfo { + tokens: Some(&tokens), + ..Default::default() + }, + ) + .unwrap(); + } + assert_eq!( + policy.divert_hits.load(Ordering::Relaxed), + hits_before + 32, + "no diversion once the worker is no longer thin" + ); + } + + #[test] + fn a_thin_worker_keeps_its_share_past_the_warm_up_blocks_until_the_ratio() { + // Churn c3 (6ce84c9e): the fleet's level was 32,767 blocks a worker, + // far above the warm-up blocks. Within a minute of the publisher + // restart the emptied worker's index had regrown to 1,765 from the + // decode blocks of its requests in flight, and the pool table, + // applying the age rule's growth cap to a thin worker, dropped it + // there after one diversion: idle for the 25 minutes to the next + // fault. Here the level is 4,096 (the ratio at 2,048), the worker has + // regrown to 1,100 before any diversion, every request is a deep hit + // (60 of 64 blocks, held elsewhere), each diverted request stays in + // flight for the rest of the test, and the table is rebuilt every + // eight decisions as the second's refresh would. + let HolderFleet { + policy, + workers, + indexer, + heads, + } = fleet_with_holders_of(7, 4_096); + let thin = indexer.worker_id("http://w7:8000").unwrap(); + // The growth baselines date from the first table, built on the empty + // fleet as on the run; the regrowth comes after. + policy.pool_table(&workers, &indexer, liveness::now_ms()); + let regrown: Vec = (0..4_400).map(|i| 80_000_000 + i).collect(); + store_blocks(&indexer, thin, ®rown, 4, 8_000_000); + assert_eq!(indexer.worker_block_count(thin), 1_100); + let request = |i: usize| -> Vec { + let mut tokens = heads[i % 7].clone(); + tokens.extend((0..16).map(|t| 90_000_000 + i as u32 * 16 + t)); + tokens + }; + let select = |tokens: &[u32]| -> usize { + policy + .select_worker( + &workers, + &SelectWorkerInfo { + tokens: Some(tokens), + ..Default::default() + }, + ) + .unwrap() + }; + let mut rebuilds = 1u64; + let mut decisions = 0usize; + let mut first = None; + let mut received = 0usize; + while decisions < 2_000 && indexer.worker_block_count(thin) < 2_048 { + if decisions.is_multiple_of(8) { + rebuilds += 1; + policy.pool_table( + &workers, + &indexer, + liveness::now_ms() + rebuilds * POOL_TABLE_REFRESH_MS, + ); + } + let tokens = request(decisions); + let idx = select(&tokens); + decisions += 1; + if idx != 7 { + continue; + } + received += 1; + store_blocks( + &indexer, + thin, + &tokens, + 4, + 9_000_000 + decisions as u64 * 100, + ); + // The diverted request runs on (tens of seconds on the trace). + workers[7].increment_load(); + if first.is_none() { + first = Some(decisions); + // Inside the window, with that request in flight, no second + // diversion: the hit counter runs on without a reset. + let hits = policy.divert_hits.load(Ordering::Relaxed); + for i in 0..16 { + select(&request(10_000 + i)); + } + assert_eq!( + policy.divert_hits.load(Ordering::Relaxed), + hits + 16, + "no second diversion inside the window while the first is in flight" + ); + } + // The window elapses; the request is still in flight. + workers[7].note_diverted(0); + } + assert!( + first.is_some_and(|first| first <= 8), + "the thin worker is served within the first eight hits although it \ + regrew past the warm-up blocks, was {first:?}" + ); + assert!( + indexer.worker_block_count(thin) >= 2_048, + "refilled to the ratio across the rebuilds: {} blocks after {decisions} decisions", + indexer.worker_block_count(thin) + ); + assert!( + received >= 15 && decisions <= 140 * 8 + 32, + "one hit in eight all the way (seven heads of 60 blocks, then tails of 4): \ + {decisions} decisions, {received} received, {} in flight", + workers[7].load() + ); + // At the ratio the table rebuilt lists no thin worker, and with no + // thin worker the diversion costs nothing: the hits are not even + // counted. + let table = policy.pool_table( + &workers, &indexer, - 4, - &default_tuning(), + liveness::now_ms() + (rebuilds + 1) * POOL_TABLE_REFRESH_MS, + ); + assert!(table.thin.is_empty(), "no longer thin: {:?}", table.thin); + let hits = policy.divert_hits.load(Ordering::Relaxed); + for i in 0..32 { + select(&request(20_000 + i)); + } + assert_eq!( + policy.divert_hits.load(Ordering::Relaxed), + hits, + "no hit counted, let alone diverted, once no worker is thin" ); - assert_eq!(result, vec![0]); // w1 (higher overlap) } - // -- select_worker_event_driven integration tests -- - #[test] - fn event_unique_hit_keeps_affinity_and_credits_expected_wait() { - let policy = CacheAwarePolicy::with_config(test_config()); - let workers = make_workers(&["http://w1:8000", "http://w2:8000"]); - policy.init_workers(&workers); - workers[1].increment_load(); - update_expected_wait_loads(&policy, &workers, &[10_000, 0]); - - let monitor = Arc::new(KvEventMonitor::new(Some(4))); - let indexer = - setup_indexer_with_blocks("http://w1:8000", &[&[1, 2, 3, 4], &[5, 6, 7, 8]], 4); - monitor.indexers.insert("unknown".to_string(), indexer); - policy.set_kv_event_monitor(Some(monitor)); - - let cached = [1, 2, 3, 4, 5, 6, 7, 8]; - assert_eq!( - policy.select_worker( - &workers, - &SelectWorkerInfo { - tokens: Some(&cached), - ..Default::default() - } - ), - Some(0) - ); - assert_only_final_worker_credited(&policy, &workers, 0, 8); + fn a_fleet_without_a_thin_worker_sees_no_diversion() { + let fleet = fleet_with_holders(8); + for i in 0..64 { + let holder = i % 8; + let idx = fleet + .policy + .select_worker( + &fleet.workers, + &SelectWorkerInfo { + tokens: Some(&fleet.heads[holder]), + ..Default::default() + }, + ) + .unwrap(); + assert_eq!( + idx, holder, + "every hit stays with its holder (decision {i})" + ); + } } #[test] @@ -4832,9 +6178,9 @@ mod tests { let monitor = Arc::new(KvEventMonitor::new(Some(4))); // Store blocks using block_size=8 (tokens chunked in groups of 8) - let indexer = Arc::new(PositionalIndexer::new(4)); + let indexer = Arc::new(KvIndex::positional(4)); let w1_id = indexer.intern_worker("http://w1:8000").unwrap(); - let mut wb = WorkerBlockMap::default(); + let mut wb = WorkerBlocks::default(); let block = vec![StoredBlock { seq_hash: SequenceHash(1), content_hash: compute_content_hash(&[1, 2, 3, 4, 5, 6, 7, 8]), @@ -4933,7 +6279,7 @@ mod tests { // Set up monitor with an empty indexer let monitor = Arc::new(KvEventMonitor::new(Some(4))); - let empty_indexer = Arc::new(PositionalIndexer::new(4)); + let empty_indexer = Arc::new(KvIndex::positional(4)); monitor .indexers .insert("unknown".to_string(), empty_indexer); @@ -5709,4 +7055,532 @@ mod tests { forged.extend_from_slice(&tokens); assert_eq!(route_namespaced(&policy, &workers, &forged, None), 0); } + + // -- selection policy layer -- + + /// Random positive-overlap candidate sets, as `overlap_candidates` would + /// produce them (healthy order, scores from a small set so ties occur). + fn random_overlap_candidates(rng: &mut impl RngExt, workers: usize) -> Vec { + let mut candidates = Vec::new(); + for idx in 0..workers { + if rng.random::() >= 0.6 { + continue; + } + let score = f64::from(rng.random_range(1..6u32)) * 1.5; + let decayed = if rng.random::() < 0.3 { + score / 1.5 + } else { + score + }; + candidates.push(OverlapCandidate { + idx, + raw_score: score, + effective_score: decayed, + }); + } + candidates + } + + fn policy_inputs<'a>( + candidates: &'a [OverlapCandidate], + urls: &'a [String], + ) -> Vec> { + candidates + .iter() + .map(|candidate| CandidateInputs { + idx: candidate.idx, + url: &urls[candidate.idx], + device_blocks: candidate.raw_score, + effective_score: candidate.effective_score, + }) + .collect() + } + + #[test] + fn default_policy_reproduces_affinity_score_group_at_zero_temperature() { + let policy = cost::build(cost::DEFAULT_POLICY, 0.0).unwrap(); + let urls: Vec = (0..16).map(|i| format!("http://w{i:02}:8000")).collect(); + let request = RequestInputs { + prompt_tokens: 64, + block_size: 4, + request_blocks: 16, + avg_load: 0.0, + prefix_hashes: None, + }; + let mut rng = rand::rng(); + for _ in 0..2_000 { + let candidates = random_overlap_candidates(&mut rng, urls.len()); + let mut reference = CacheAwarePolicy::affinity_score_group(&candidates, 0.0); + reference.sort_unstable(); + let inputs = policy_inputs(&candidates, &urls); + let mut group = match policy.select(&request, &inputs) { + Pick::Group(rows) => rows.iter().map(|&row| inputs[row].idx).collect::>(), + Pick::None => Vec::new(), + Pick::Final(_) => panic!("the default policy never returns a single pick"), + }; + group.sort_unstable(); + assert_eq!(group, reference); + } + } + + #[test] + fn default_policy_temperature_groups_are_score_groups() { + let policy = cost::build(cost::DEFAULT_POLICY, 0.5).unwrap(); + let urls: Vec = (0..8).map(|i| format!("http://w{i}:8000")).collect(); + let candidates: Vec = (0..8) + .map(|idx| OverlapCandidate { + idx, + raw_score: if idx < 4 { 8.0 } else { 2.0 }, + effective_score: if idx < 4 { 8.0 } else { 2.0 }, + }) + .collect(); + let inputs = policy_inputs(&candidates, &urls); + let request = RequestInputs { + prompt_tokens: 32, + block_size: 4, + request_blocks: 8, + avg_load: 0.0, + prefix_hashes: None, + }; + let mut saw_high = false; + let mut saw_low = false; + for _ in 0..400 { + let Pick::Group(rows) = policy.select(&request, &inputs) else { + panic!("expected a group"); + }; + let score = inputs[rows[0]].effective_score; + assert!(rows.iter().all(|&row| inputs[row].effective_score == score)); + assert_eq!(rows.len(), 4, "a score group is every equal-score worker"); + if score == 8.0 { + saw_high = true; + } else { + saw_low = true; + } + } + assert!( + saw_high && saw_low, + "temperature must reach both score groups" + ); + } + + #[test] + fn every_catalog_policy_routes_event_driven_hits_and_misses() { + for name in cost::POLICY_NAMES { + let policy = CacheAwarePolicy::with_config(CacheAwareConfig { + selection_policy: Some((*name).to_string()), + ..test_config() + }); + let workers = make_workers(&["http://w1:8000", "http://w2:8000"]); + policy.init_workers(&workers); + let monitor = Arc::new(KvEventMonitor::new(Some(4))); + let indexer = + setup_indexer_with_blocks("http://w2:8000", &[&[1, 2, 3, 4], &[5, 6, 7, 8]], 4); + monitor.indexers.insert("unknown".to_string(), indexer); + policy.set_kv_event_monitor(Some(monitor)); + + let hit = policy.select_worker( + &workers, + &SelectWorkerInfo { + tokens: Some(&[1, 2, 3, 4, 5, 6, 7, 8]), + ..Default::default() + }, + ); + assert_eq!(hit, Some(1), "{name}: the holder of every block must win"); + policy.on_request_complete("http://w2:8000", true); + + let miss = policy.select_worker( + &workers, + &SelectWorkerInfo { + tokens: Some(&[100, 200, 300, 400]), + ..Default::default() + }, + ); + assert!(miss.is_some(), "{name}: a miss must still route"); + } + } + + #[test] + fn every_catalog_policy_routes_tree_matches() { + for name in cost::POLICY_NAMES { + let policy = CacheAwarePolicy::with_config(CacheAwareConfig { + selection_policy: Some((*name).to_string()), + ..test_config() + }); + let workers = make_workers(&["http://w1:8000", "http://w2:8000"]); + policy.init_workers(&workers); + let text = "a shared system prompt that is long enough to match"; + let first = policy + .select_worker( + &workers, + &SelectWorkerInfo { + request_text: Some(text), + ..Default::default() + }, + ) + .unwrap(); + policy.on_request_complete(workers[first].url(), true); + let second = policy + .select_worker( + &workers, + &SelectWorkerInfo { + request_text: Some(text), + ..Default::default() + }, + ) + .unwrap(); + assert_eq!(first, second, "{name}: an idle holder keeps its prefix"); + } + } + + #[test] + fn accounting_books_the_dispatch_until_completion() { + let policy = CacheAwarePolicy::with_config(CacheAwareConfig { + selection_accounting_ttl_ms: 60_000, + ..test_config() + }); + let workers = make_workers(&["http://w1:8000", "http://w2:8000"]); + policy.init_workers(&workers); + let monitor = Arc::new(KvEventMonitor::new(Some(4))); + let indexer = setup_indexer_with_blocks("http://w1:8000", &[&[1, 2, 3, 4]], 4); + monitor.indexers.insert("unknown".to_string(), indexer); + policy.set_kv_event_monitor(Some(monitor)); + let accounting = policy.accounting.as_ref().expect("accounting enabled"); + + let tokens = [1, 2, 3, 4, 5, 6, 7, 8]; + let selected = policy + .select_worker( + &workers, + &SelectWorkerInfo { + tokens: Some(&tokens), + ..Default::default() + }, + ) + .unwrap(); + assert_eq!(selected, 0); + // One block cached, one block (4 tokens) booked as uncached prefill. + assert_eq!(accounting.pending_prefill_tokens("http://w1:8000"), 4); + let hashes: Vec = request_prefix_hashes(&compute_request_content_hashes(&tokens, 4)) + .into_iter() + .map(|hash| hash.0) + .collect(); + let predicted = accounting.predicted_overlaps(&hashes); + assert_eq!(predicted.len(), 1); + assert_eq!(&*predicted[0].0, "http://w1:8000"); + assert_eq!( + predicted[0].1, 2.0, + "the whole two-block prefix is predicted resident" + ); + + // A sibling that shares the first block but not the second is predicted one block on w1. + let sibling: Vec = request_prefix_hashes(&compute_request_content_hashes( + &[1, 2, 3, 4, 9, 9, 9, 9], + 4, + )) + .into_iter() + .map(|hash| hash.0) + .collect(); + assert_eq!(accounting.predicted_overlaps(&sibling)[0].1, 1.0); + + policy.on_request_complete("http://w1:8000", true); + assert_eq!(accounting.pending_prefill_tokens("http://w1:8000"), 0); + } + + #[test] + fn reconciliation_releases_bookings_whose_completion_never_arrived() { + let policy = CacheAwarePolicy::with_config(CacheAwareConfig { + selection_accounting_ttl_ms: 60_000, + ..test_config() + }); + let accounting = policy.accounting.as_ref().expect("accounting enabled"); + accounting.record_dispatch("http://w1:8000", 100, &[]); + accounting.record_dispatch("http://w1:8000", 200, &[]); + accounting.record_dispatch("http://w1:8000", 400, &[]); + + // The router still holds all three: nothing to release. + policy.reconcile_in_flight("http://w1:8000", 3); + assert_eq!(accounting.pending_prefill_tokens("http://w1:8000"), 700); + + // Two requests ended without a completion report: the two oldest + // bookings go, the one the router still holds stays. + policy.reconcile_in_flight("http://w1:8000", 1); + assert_eq!(accounting.pending_prefill_tokens("http://w1:8000"), 400); + + // A worker without bookings is a no-op, with or without accounting. + policy.reconcile_in_flight("http://w2:8000", 0); + CacheAwarePolicy::with_config(test_config()).reconcile_in_flight("http://w1:8000", 0); + } + + #[test] + fn invalid_selection_policy_falls_back_to_the_default() { + let policy = CacheAwarePolicy::with_config(CacheAwareConfig { + selection_policy: Some("no-such-policy".to_string()), + ..test_config() + }); + assert_eq!(policy.selection.name(), cost::DEFAULT_POLICY); + } + + #[test] + fn event_driven_salted_request_matches_only_its_namespace() { + use kv_index::salt::{content_hash_with_seed, namespace_seed}; + use openai_protocol::common::CachePartition; + + let policy = CacheAwarePolicy::with_config(test_config()); + let w1 = BasicWorkerBuilder::new("http://w1:8000") + .worker_type(WorkerType::Regular) + .health_config(no_health_check()) + .build(); + let w2 = BasicWorkerBuilder::new("http://w2:8000") + .worker_type(WorkerType::Regular) + .health_config(no_health_check()) + .build(); + // w1 carries more live load, so every miss resolves to w2. + for _ in 0..3 { + w1.increment_load(); + } + let workers: Vec> = vec![Arc::new(w1), Arc::new(w2)]; + policy.init_workers(&workers); + + // Blocks stored on w1 under (lora "adapter", salt "tenant-a"), as the + // monitor hashes a salted KvBlocksStored event. + let indexer = Arc::new(KvIndex::positional(4)); + let worker_id = indexer.intern_worker("http://w1:8000").unwrap(); + let mut wb = WorkerBlocks::default(); + let seed = namespace_seed(Some("adapter"), Some("tenant-a")); + let blocks: Vec = [[1u32, 2, 3, 4], [5, 6, 7, 8]] + .iter() + .enumerate() + .map(|(i, tokens)| StoredBlock { + seq_hash: SequenceHash(i as u64 + 1), + content_hash: content_hash_with_seed(tokens, seed), + }) + .collect(); + indexer + .apply_stored(worker_id, &blocks, None, &mut wb) + .unwrap(); + let monitor = Arc::new(KvEventMonitor::new(Some(4))); + monitor.indexers.insert("unknown".to_string(), indexer); + policy.set_kv_event_monitor(Some(monitor)); + + let namespace = |salt: &'static str| { + CacheNamespace::derive(&CachePartition { + cache_salt: Some(salt), + extra_key: None, + lora_path: Some("adapter"), + }) + }; + let tokens = [1, 2, 3, 4, 5, 6, 7, 8]; + let route = |cache_namespace: Option| { + policy + .select_worker( + &workers, + &SelectWorkerInfo { + tokens: Some(&tokens), + cache_namespace, + ..Default::default() + }, + ) + .unwrap() + }; + assert_eq!(route(namespace("tenant-a")), 0, "same namespace hits w1"); + assert_eq!(route(namespace("tenant-b")), 1, "another salt misses"); + assert_eq!(route(None), 1, "a plain request misses salted blocks"); + } + + // ----------------------------------------------------------------------- + // Pool tables and the sampled decision over a large pool + // ----------------------------------------------------------------------- + + /// `count` workers, all interned in an event index that holds `chunks` + /// for the worker at `holder`, behind a policy built from `config`. + fn large_pool( + count: usize, + holder: usize, + chunks: &[&[u32]], + config: CacheAwareConfig, + ) -> (CacheAwarePolicy, Vec>) { + let urls: Vec = (0..count).map(|i| format!("http://w{i}:8000")).collect(); + let refs: Vec<&str> = urls.iter().map(String::as_str).collect(); + let workers = make_workers(&refs); + let indexer = setup_indexer_with_blocks(&urls[holder], chunks, 4); + for (i, url) in urls.iter().enumerate() { + if i != holder { + indexer.intern_worker(url).unwrap(); + } + } + let policy = CacheAwarePolicy::with_config(config); + policy.init_workers(&workers); + let monitor = Arc::new(KvEventMonitor::new(Some(4))); + monitor.indexers.insert("unknown".to_string(), indexer); + monitor.set_block_size("unknown", 4); + policy.set_kv_event_monitor(Some(monitor)); + (policy, workers) + } + + fn tokens_info(tokens: &[u32]) -> SelectWorkerInfo<'_> { + SelectWorkerInfo { + tokens: Some(tokens), + ..Default::default() + } + } + + #[test] + fn a_large_pool_routes_to_the_holder_the_index_names() { + let (policy, workers) = large_pool(40, 37, &[&[1, 2, 3, 4], &[5, 6, 7, 8]], test_config()); + for _ in 0..8 { + let idx = policy + .select_worker(&workers, &tokens_info(&[1, 2, 3, 4, 5, 6, 7, 8])) + .unwrap(); + assert_eq!( + idx, 37, + "the holder sits outside any sample of {FLEET_SAMPLE}; the index places it" + ); + } + } + + #[test] + fn a_large_pool_miss_finds_the_one_eligible_worker_the_sample_missed() { + let (policy, workers) = large_pool(40, 0, &[&[1, 2, 3, 4]], test_config()); + for worker in &workers[..39] { + worker.set_status(WorkerStatus::NotReady); + } + let idx = policy.select_worker(&workers, &tokens_info(&[9, 9, 9, 9, 8, 8, 8, 8])); + assert_eq!(idx, Some(39)); + } + + #[test] + fn a_large_pool_miss_spreads_over_the_pool() { + let (policy, workers) = large_pool(40, 0, &[&[1, 2, 3, 4]], test_config()); + let mut picked = HashSet::new(); + for turn in 0..200u32 { + let tokens = [turn + 100, 9, 9, 9, 8, 8, 8, 8]; + let idx = policy + .select_worker(&workers, &tokens_info(&tokens)) + .unwrap(); + picked.insert(idx); + policy.on_request_complete(workers[idx].url(), true); + } + assert!( + picked.len() > 2 * FLEET_SAMPLE, + "200 misses reached {} workers; a fresh sample per decision reaches the pool", + picked.len() + ); + } + + #[test] + fn pool_table_places_ids_and_rejects_a_worker_swapped_in_at_the_position() { + let workers = make_workers(&["http://w1:8000", "http://w2:8000", "http://w3:8000"]); + let indexer = setup_indexer_with_blocks("http://w3:8000", &[&[1, 2, 3, 4]], 4); + let w3 = indexer.worker_id("http://w3:8000").unwrap(); + let table = PoolTable::build(&workers, &indexer, liveness::now_ms()); + assert!(table.describes(&workers)); + assert!(table.fresh(liveness::now_ms())); + assert_eq!(table.position(w3, &workers), Some(2)); + assert_eq!(table.id_at, vec![None, None, Some(w3)]); + assert_eq!( + table.position(w3 + 100, &workers), + None, + "an id never interned" + ); + // The same url behind a new worker object at the same position: the + // table was built from the old one and says so. + let mut swapped = workers.clone(); + swapped[2] = make_workers(&["http://w3:8000"]).remove(0); + assert_eq!(table.position(w3, &swapped), None); + // Every worker was admitted just now: a young fleet lists nothing to slice. + assert!(table.warming.is_empty()); + assert!(!table.fresh(liveness::now_ms() + POOL_TABLE_REFRESH_MS)); + } + + /// `count` workers that all hold the shared block `[1, 2, 3, 4]` (a chat + /// template's head), the one at `deep` also `[5, 6, 7, 8]`. + fn shared_prefix_pool( + count: usize, + deep: usize, + config: CacheAwareConfig, + ) -> (CacheAwarePolicy, Vec>) { + let urls: Vec = (0..count).map(|i| format!("http://w{i}:8000")).collect(); + let refs: Vec<&str> = urls.iter().map(String::as_str).collect(); + let workers = make_workers(&refs); + let indexer = Arc::new(KvIndex::positional(4)); + for (i, url) in urls.iter().enumerate() { + let id = indexer.intern_worker(url).unwrap(); + let mut wb = WorkerBlocks::default(); + let mut blocks = vec![StoredBlock { + seq_hash: SequenceHash(1), + content_hash: compute_content_hash(&[1, 2, 3, 4]), + }]; + if i == deep { + blocks.push(StoredBlock { + seq_hash: SequenceHash(2), + content_hash: compute_content_hash(&[5, 6, 7, 8]), + }); + } + indexer.apply_stored(id, &blocks, None, &mut wb).unwrap(); + } + let policy = CacheAwarePolicy::with_config(config); + policy.init_workers(&workers); + let monitor = Arc::new(KvEventMonitor::new(Some(4))); + monitor.indexers.insert("unknown".to_string(), indexer); + monitor.set_block_size("unknown", 4); + policy.set_kv_event_monitor(Some(monitor)); + (policy, workers) + } + + #[test] + fn the_default_decision_takes_the_deepest_holder_among_a_fleet_sharing_the_head() { + let (policy, workers) = shared_prefix_pool(64, 50, test_config()); + for _ in 0..4 { + let idx = policy + .select_worker(&workers, &tokens_info(&[1, 2, 3, 4, 5, 6, 7, 8])) + .unwrap(); + assert_eq!(idx, 50); + } + // The deepest holder down: the decision falls to the holders of the + // head, never to a worker without the blocks. + workers[50].set_status(WorkerStatus::NotReady); + let idx = policy + .select_worker(&workers, &tokens_info(&[1, 2, 3, 4, 5, 6, 7, 8])) + .unwrap(); + assert_ne!(idx, 50); + assert!(workers[idx].is_healthy()); + } + + #[test] + fn a_fleet_wide_tie_is_drawn_from_not_scored_whole() { + let (policy, workers) = shared_prefix_pool(64, 50, test_config()); + let mut picked = HashSet::new(); + for _ in 0..200 { + let idx = policy + .select_worker(&workers, &tokens_info(&[1, 2, 3, 4])) + .unwrap(); + picked.insert(idx); + policy.on_request_complete(workers[idx].url(), true); + } + assert!( + picked.len() > 2 * FLEET_SAMPLE, + "200 decisions on a head every worker holds reached {} workers", + picked.len() + ); + } + + #[test] + fn merge_rows_unions_the_eligible_sample_with_the_holders_in_slice_order() { + let candidates: Vec = [3usize, 7, 20] + .iter() + .map(|&idx| OverlapCandidate { + idx, + raw_score: 1.0, + effective_score: 1.0, + }) + .collect(); + assert_eq!( + CacheAwarePolicy::merge_rows(&[1, 3, 9, 30], &candidates), + vec![1, 3, 7, 9, 20, 30] + ); + assert_eq!( + CacheAwarePolicy::merge_rows(&[], &candidates), + vec![3, 7, 20] + ); + assert_eq!(CacheAwarePolicy::merge_rows(&[2, 4], &[]), vec![2, 4]); + } } diff --git a/model_gateway/src/policies/cache_namespace.rs b/model_gateway/src/policies/cache_namespace.rs index fcc0a708b6..7bd8a98410 100644 --- a/model_gateway/src/policies/cache_namespace.rs +++ b/model_gateway/src/policies/cache_namespace.rs @@ -17,6 +17,13 @@ //! gossip. Requests without any partition field get no namespace, and their //! routing keys are byte-identical to before. //! +//! The event-driven index is keyed differently: the engines fold the LoRA +//! name and the cache salt into their block hashes, and the KV event monitor +//! recomputes stored blocks under the same XXH3 seed (`kv_index::salt`). A +//! namespace therefore also carries that seed, so a salted or LoRA request +//! is hashed under it and matches only its own blocks; a plain request keeps +//! the plain hash. The extra key is not part of the engines' seed. +//! //! The marker is excluded from the match ratio, so it cannot by itself push //! two unrelated prompts of one namespace over the cache threshold. //! @@ -32,6 +39,7 @@ //! Size it for tenants × working set; a per-request salt makes every request //! a unique path. +use kv_index::salt::namespace_seed; use openai_protocol::common::CachePartition; use xxhash_rust::xxh3::Xxh3; @@ -58,9 +66,13 @@ const MARKER_PAD_TOKEN: u32 = MARKER_TOKEN_BIT; /// does not begin with. const TEXT_MARKER_DELIM: char = '\u{1}'; -/// Fixed-width identity of a request's cache partition. +/// Fixed-width identity of a request's cache partition, and the XXH3 seed +/// under which the engines' KV events for that partition are hashed. #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] -pub struct CacheNamespace(u64); +pub struct CacheNamespace { + marker: u64, + event_seed: u64, +} impl CacheNamespace { /// Derive the namespace from a request's partition fields; `None` when @@ -88,21 +100,31 @@ impl CacheNamespace { None => hasher.update(&[0u8]), } } - Some(Self(hasher.digest())) + Some(Self { + marker: hasher.digest(), + event_seed: namespace_seed(partition.lora_path, partition.cache_salt), + }) + } + + /// The XXH3 seed the engines' blocks for this partition are hashed under + /// (see `kv_index::salt::namespace_seed`); the plain seed when neither a + /// LoRA name nor a cache salt is set. + pub fn event_seed(self) -> u64 { + self.event_seed } /// The namespace as token ids for the token tree and hash mode: two ids /// with the reserved high bit set, so they never match prompt tokens. pub fn token_marker(self) -> [u32; TOKEN_MARKER_LEN] { [ - (self.0 >> 32) as u32 | MARKER_TOKEN_BIT, - self.0 as u32 | MARKER_TOKEN_BIT, + (self.marker >> 32) as u32 | MARKER_TOKEN_BIT, + self.marker as u32 | MARKER_TOKEN_BIT, ] } /// The namespace as a fixed-width text prefix for the string tree. pub fn text_marker(self) -> String { - format!("{d}{:016x}{d}", self.0, d = TEXT_MARKER_DELIM) + format!("{d}{:016x}{d}", self.marker, d = TEXT_MARKER_DELIM) } /// Length of the token-tree marker for a tree with `page_size`-token @@ -265,4 +287,22 @@ mod tests { assert!(keyed.ends_with("hello")); assert!(keyed.starts_with(&marker)); } + + #[test] + fn event_seed_follows_the_lora_name_and_cache_salt_only() { + let salted = + CacheNamespace::derive(&partition(Some("tenant-a"), None, Some("adapter"))).unwrap(); + assert_eq!( + salted.event_seed(), + namespace_seed(Some("adapter"), Some("tenant-a")) + ); + let extra_only = CacheNamespace::derive(&partition(None, Some("k"), None)).unwrap(); + assert_eq!(extra_only.event_seed(), kv_index::XXH3_SEED); + assert_ne!( + salted.event_seed(), + CacheNamespace::derive(&partition(Some("tenant-b"), None, Some("adapter"))) + .unwrap() + .event_seed() + ); + } } diff --git a/model_gateway/src/policies/cost/accounting.rs b/model_gateway/src/policies/cost/accounting.rs new file mode 100644 index 0000000000..8887b10e96 --- /dev/null +++ b/model_gateway/src/policies/cost/accounting.rs @@ -0,0 +1,423 @@ +//! Optimistic self-accounting: what the router has already decided but the engines have not yet +//! reported. +//! +//! Between a dispatch and the engine's first KV event (or load report) the index and the load +//! snapshot are stale by exactly that request. Under a burst of sibling requests (best-of-N, agent +//! fan-out, replayed sessions) that window is enough to herd them all onto one worker. This keeps +//! two short-lived views, in the spirit of a predict-on-route side index and an active-sequence +//! booking: +//! +//! - **booked prefill**: the predicted uncached tokens of each dispatched request, charged to the +//! chosen worker until it completes or the booking expires; +//! - **predicted placement**: the prefix of each dispatched request and the worker that will hold +//! its blocks. A prefix is keyed by its chain hash at block depths `1, 2, 4, 8, …` and at its full +//! length, so a later request sharing `d` leading blocks finds the placement at the deepest power +//! of two not above `d` with `O(log d)` lookups and no per-block index. +//! +//! Both expire after `ttl`, which should be a little longer than the engine's event lag. Output +//! blocks can be credited as generation crosses block boundaries through `on_output_blocks`. + +use std::{ + collections::{HashMap, VecDeque}, + sync::Arc, + time::{Duration, Instant}, +}; + +use parking_lot::Mutex; + +#[derive(Debug, Clone, Copy)] +struct Booking { + tokens: u64, + expires: Instant, +} + +#[derive(Debug, Clone)] +struct Placement { + url: Arc, + expires: Instant, +} + +#[derive(Debug, Default)] +struct State { + booked: HashMap, VecDeque>, + output_blocks: HashMap, f64>, + /// Prefix chain hash → workers predicted to hold that prefix, newest last. + predicted: HashMap>, + /// Insertion order of prediction keys, for bounded eviction. + predicted_order: VecDeque<(u64, Instant)>, +} + +/// Bound on remembered prediction keys; the oldest go first once reached. +const MAX_PREDICTED_KEYS: usize = 1 << 16; +/// Workers remembered per prefix key (a prefix spilled to several workers). +const MAX_PLACEMENTS_PER_KEY: usize = 4; + +#[derive(Debug)] +pub struct OptimisticAccounting { + ttl: Duration, + state: Mutex, +} + +/// Block positions (zero-based) at which a prefix of `blocks` blocks is keyed: block counts +/// `1, 2, 4, …` below `blocks`, then `blocks` itself. +fn key_positions(blocks: usize) -> impl Iterator { + let mut count = 1usize; + std::iter::from_fn(move || { + if count < blocks { + let position = count - 1; + count *= 2; + Some(position) + } else if count == usize::MAX { + None + } else { + count = usize::MAX; + blocks.checked_sub(1) + } + }) +} + +impl OptimisticAccounting { + pub fn new(ttl: Duration) -> Self { + Self { + ttl, + state: Mutex::new(State::default()), + } + } + + pub fn ttl(&self) -> Duration { + self.ttl + } + + fn intern(state: &State, url: &str) -> Arc { + state + .booked + .get_key_value(url) + .map(|(key, _)| Arc::clone(key)) + .unwrap_or_else(|| Arc::from(url)) + } + + /// Book a dispatch: `uncached_tokens` of prefill on `url`, and the prompt's prefix + /// (`prefix_hashes`, chain hashes by block) predicted to become resident there. + pub fn record_dispatch(&self, url: &str, uncached_tokens: u64, prefix_hashes: &[u64]) { + let now = Instant::now(); + let expires = now + self.ttl; + let mut state = self.state.lock(); + let key = Self::intern(&state, url); + let queue = state.booked.entry(Arc::clone(&key)).or_default(); + while queue.front().is_some_and(|b| b.expires <= now) { + queue.pop_front(); + } + queue.push_back(Booking { + tokens: uncached_tokens, + expires, + }); + if prefix_hashes.is_empty() { + return; + } + let State { + predicted, + predicted_order, + .. + } = &mut *state; + for position in key_positions(prefix_hashes.len()) { + let hash = prefix_hashes[position]; + let placements = predicted.entry(hash).or_default(); + if placements.is_empty() { + predicted_order.push_back((hash, expires)); + } + placements.retain(|p| p.expires > now); + match placements.iter_mut().find(|p| p.url == key) { + Some(existing) => existing.expires = expires, + None => { + if placements.len() == MAX_PLACEMENTS_PER_KEY { + placements.remove(0); + } + placements.push(Placement { + url: Arc::clone(&key), + expires, + }); + } + } + } + // Beyond the key bound the oldest keys go, live or not, so the walk + // shortens the queue; below it, an expired head whose placements a + // later dispatch re-armed goes back to the tail under its newest + // expiry, which is not yet due, so it is not seen again this walk. + while let Some(&(hash, expiry)) = predicted_order.front() { + let over_bound = predicted_order.len() > MAX_PREDICTED_KEYS; + if !over_bound && expiry > now { + break; + } + predicted_order.pop_front(); + if over_bound { + predicted.remove(&hash); + continue; + } + let Some(placements) = predicted.get_mut(&hash) else { + continue; + }; + placements.retain(|p| p.expires > now); + if placements.is_empty() { + predicted.remove(&hash); + } else { + let latest = placements.iter().map(|p| p.expires).max().unwrap_or(now); + predicted_order.push_back((hash, latest)); + } + } + } + + /// A request on `url` finished (or produced its first token): release its oldest booking. + pub fn release(&self, url: &str) { + let mut state = self.state.lock(); + if let Some(queue) = state.booked.get_mut(url) { + queue.pop_front(); + if queue.is_empty() { + state.booked.remove(url); + } + } + } + + /// Trim the live bookings on `url` to the router's in-flight count there, + /// oldest first, and return how many were released. A dispatch books once + /// and a completion releases once, so bookings beyond the live count are + /// completions that never arrived. Expired bookings are dropped on the way + /// and not counted: they were already out of every sum. + pub fn reconcile(&self, url: &str, in_flight: usize) -> usize { + let now = Instant::now(); + let mut state = self.state.lock(); + let Some(queue) = state.booked.get_mut(url) else { + return 0; + }; + queue.retain(|booking| booking.expires > now); + let excess = queue.len().saturating_sub(in_flight); + queue.drain(..excess); + if queue.is_empty() { + state.booked.remove(url); + } + excess + } + + /// Prefill tokens booked on `url` that have not expired or been released. + pub fn pending_prefill_tokens(&self, url: &str) -> u64 { + let now = Instant::now(); + let state = self.state.lock(); + state.booked.get(url).map_or(0, |queue| { + queue + .iter() + .filter(|b| b.expires > now) + .map(|b| b.tokens) + .sum() + }) + } + + /// Workers predicted to hold a prefix of the prompt described by `prefix_hashes`, each with the + /// number of leading blocks predicted resident there (the deepest keyed depth that matched). + pub fn predicted_overlaps(&self, prefix_hashes: &[u64]) -> Vec<(Arc, f64)> { + if prefix_hashes.is_empty() { + return Vec::new(); + } + let now = Instant::now(); + let positions: Vec = key_positions(prefix_hashes.len()).collect(); + let state = self.state.lock(); + let mut found: Vec<(Arc, f64)> = Vec::new(); + for &position in positions.iter().rev() { + let Some(placements) = state.predicted.get(&prefix_hashes[position]) else { + continue; + }; + for placement in placements.iter().filter(|p| p.expires > now) { + if !found.iter().any(|(url, _)| *url == placement.url) { + found.push((Arc::clone(&placement.url), (position + 1) as f64)); + } + } + } + found + } + + /// Credit `blocks` of generated output on `url` (decode-side growth the engine will report + /// later); `release_output` forgets it when the request ends. + pub fn on_output_blocks(&self, url: &str, blocks: f64) { + let mut state = self.state.lock(); + let key = Self::intern(&state, url); + *state.output_blocks.entry(key).or_default() += blocks; + } + + pub fn release_output(&self, url: &str, blocks: f64) { + let mut state = self.state.lock(); + if let Some(current) = state.output_blocks.get_mut(url) { + *current = (*current - blocks).max(0.0); + if *current == 0.0 { + state.output_blocks.remove(url); + } + } + } + + pub fn output_blocks(&self, url: &str) -> f64 { + self.state + .lock() + .output_blocks + .get(url) + .copied() + .unwrap_or(0.0) + } + + /// Drop every booking and prediction for a worker that left the fleet. + pub fn forget_worker(&self, url: &str) { + let mut state = self.state.lock(); + state.booked.remove(url); + state.output_blocks.remove(url); + for placements in state.predicted.values_mut() { + placements.retain(|p| &*p.url != url); + } + state + .predicted + .retain(|_, placements| !placements.is_empty()); + } +} + +#[cfg(test)] +mod tests { + use std::sync::mpsc; + + use super::*; + + #[test] + fn key_positions_are_powers_of_two_then_full_length() { + assert_eq!(key_positions(0).collect::>(), Vec::::new()); + assert_eq!(key_positions(1).collect::>(), vec![0]); + assert_eq!(key_positions(2).collect::>(), vec![0, 1]); + assert_eq!(key_positions(5).collect::>(), vec![0, 1, 3, 4]); + assert_eq!(key_positions(8).collect::>(), vec![0, 1, 3, 7]); + assert_eq!(key_positions(9).collect::>(), vec![0, 1, 3, 7, 8]); + } + + #[test] + fn bookings_sum_until_released_or_expired() { + let acc = OptimisticAccounting::new(Duration::from_millis(50)); + acc.record_dispatch("w1", 100, &[]); + acc.record_dispatch("w1", 200, &[]); + assert_eq!(acc.pending_prefill_tokens("w1"), 300); + acc.release("w1"); + assert_eq!(acc.pending_prefill_tokens("w1"), 200); + std::thread::sleep(Duration::from_millis(60)); + assert_eq!(acc.pending_prefill_tokens("w1"), 0); + } + + #[test] + fn reconcile_releases_bookings_beyond_the_live_count_oldest_first() { + let acc = OptimisticAccounting::new(Duration::from_secs(5)); + acc.record_dispatch("w1", 100, &[]); + acc.record_dispatch("w1", 200, &[]); + acc.record_dispatch("w1", 400, &[]); + assert_eq!(acc.reconcile("w1", 3), 0, "nothing beyond the live count"); + assert_eq!(acc.reconcile("w1", 1), 2, "two completions never arrived"); + assert_eq!( + acc.pending_prefill_tokens("w1"), + 400, + "the newest booking is the one still in flight" + ); + assert_eq!(acc.reconcile("w1", 0), 1); + assert_eq!(acc.pending_prefill_tokens("w1"), 0); + assert_eq!( + acc.reconcile("w1", 0), + 0, + "a worker with no bookings releases nothing" + ); + } + + #[test] + fn predicted_overlap_matches_shared_prefixes_at_keyed_depths() { + let acc = OptimisticAccounting::new(Duration::from_secs(5)); + let prefix: Vec = (1..=10).collect(); + acc.record_dispatch("w1", 0, &prefix); + + // A longer sibling sharing all ten blocks matches the full-length key. + let longer: Vec = (1..=12).collect(); + let found = acc.predicted_overlaps(&longer); + assert_eq!(found.len(), 1); + assert_eq!(&*found[0].0, "w1"); + assert_eq!( + found[0].1, 8.0, + "deepest keyed depth below 12 shared with 10 is 8" + ); + + // The identical prompt gets the full ten blocks. + let found = acc.predicted_overlaps(&prefix); + assert_eq!(found[0].1, 10.0); + + // A prompt sharing only three blocks matches the depth-2 key. + let short: Vec = vec![1, 2, 3, 99, 98]; + let found = acc.predicted_overlaps(&short); + assert_eq!(found[0].1, 2.0); + + // Nothing shared, nothing predicted. + assert!(acc.predicted_overlaps(&[70, 80, 90]).is_empty()); + + // Another worker gets the same prefix: both are reported. + acc.record_dispatch("w2", 0, &prefix); + let mut urls: Vec = acc + .predicted_overlaps(&prefix) + .into_iter() + .map(|(url, _)| url.to_string()) + .collect(); + urls.sort(); + assert_eq!(urls, vec!["w1", "w2"]); + + acc.forget_worker("w1"); + let found = acc.predicted_overlaps(&prefix); + assert_eq!(found.len(), 1); + assert_eq!(&*found[0].0, "w2"); + } + + /// More live keys than the bound under one long TTL: the eviction walk + /// must drop the oldest and return, not cycle a live head through the + /// queue for ever under the lock. + #[test] + fn more_live_keys_than_the_bound_evict_the_oldest_and_return() { + let acc = Arc::new(OptimisticAccounting::new(Duration::from_secs(3600))); + let extra = 100u64; + let total = MAX_PREDICTED_KEYS as u64 + extra; + let (done_tx, done_rx) = mpsc::channel(); + let worker = Arc::clone(&acc); + std::thread::spawn(move || { + for key in 1..=total { + worker.record_dispatch("w1", 0, &[key]); + } + let _ = done_tx.send(()); + }); + done_rx + .recv_timeout(Duration::from_secs(60)) + .expect("record_dispatch returns with more live keys than the bound"); + let state = acc.state.lock(); + assert!(state.predicted_order.len() <= MAX_PREDICTED_KEYS); + assert_eq!(state.predicted.len(), state.predicted_order.len()); + assert!( + state.predicted.contains_key(&total), + "the newest keys are the ones kept" + ); + assert!( + !state.predicted.contains_key(&1), + "the oldest keys are the ones evicted" + ); + } + + #[test] + fn predictions_expire() { + let acc = OptimisticAccounting::new(Duration::from_millis(30)); + acc.record_dispatch("w1", 0, &[1, 2, 3, 4]); + assert_eq!(acc.predicted_overlaps(&[1, 2, 3, 4]).len(), 1); + std::thread::sleep(Duration::from_millis(40)); + assert!(acc.predicted_overlaps(&[1, 2, 3, 4]).is_empty()); + } + + #[test] + fn output_blocks_accumulate_and_release() { + let acc = OptimisticAccounting::new(Duration::from_secs(1)); + acc.on_output_blocks("w1", 2.0); + acc.on_output_blocks("w1", 3.0); + assert_eq!(acc.output_blocks("w1"), 5.0); + acc.release_output("w1", 4.0); + assert_eq!(acc.output_blocks("w1"), 1.0); + acc.release_output("w1", 4.0); + assert_eq!(acc.output_blocks("w1"), 0.0); + } +} diff --git a/model_gateway/src/policies/cost/catalog.rs b/model_gateway/src/policies/cost/catalog.rs new file mode 100644 index 0000000000..b3d71e0692 --- /dev/null +++ b/model_gateway/src/policies/cost/catalog.rs @@ -0,0 +1,58 @@ +//! The policy catalog: names, parameters, construction. +//! +//! One selection policy ships, the cache-aware default; it takes no parameters. Any other name +//! is rejected at configuration time. + +use super::{default, policy::WorkerSelectionPolicy}; + +/// The policy used when none is configured: the pre-policy cache-aware decision. +pub const DEFAULT_POLICY: &str = default::POLICY_NAME; + +pub const POLICY_NAMES: &[&str] = &[default::POLICY_NAME]; + +#[derive(Debug, thiserror::Error)] +pub enum CatalogError { + #[error("unknown selection policy '{0}'; known: {known}", known = POLICY_NAMES.join(", "))] + Unknown(String), +} + +/// The default policy at the given cache-aware temperature; it takes no parameters and cannot +/// fail to build. +pub fn default_policy(selection_temperature: f32) -> WorkerSelectionPolicy { + default::policy(selection_temperature) +} + +/// Build a policy by name. `selection_temperature` is the cache-aware temperature the default +/// policy keeps using. +pub fn build( + name: &str, + selection_temperature: f32, +) -> Result { + match name { + default::POLICY_NAME => Ok(default::policy(selection_temperature)), + other => Err(CatalogError::Unknown(other.to_string())), + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn every_name_builds_with_defaults() { + for name in POLICY_NAMES { + let policy = build(name, 0.0).expect("builds"); + assert_eq!(policy.name(), *name); + } + } + + #[test] + fn unknown_names_are_rejected() { + for name in ["nope", "not-a-policy"] { + assert!( + matches!(build(name, 0.0), Err(CatalogError::Unknown(_))), + "{name}" + ); + } + } +} diff --git a/model_gateway/src/policies/cost/default.rs b/model_gateway/src/policies/cost/default.rs new file mode 100644 index 0000000000..f171f51ff2 --- /dev/null +++ b/model_gateway/src/policies/cost/default.rs @@ -0,0 +1,92 @@ +//! The cache-aware decision as it was before the policy layer, expressed as a policy. +//! +//! Only workers with a positive effective (decayed) overlap are candidates; the picker returns +//! the top effective-score group: the exact maximum at temperature zero, or the softmax-sampled +//! score group otherwise. The host then resolves that group exactly as it always has: the +//! per-request pressure gate, then LeastLoad's expected-wait selection and credit. An empty group +//! is the miss path. Decisions are therefore identical to the pre-policy code; the tests in +//! `cache_aware.rs` compare the two at zero temperature and under a temperature. + +use super::{ + inputs::{CandidateInputs, RequestInputs}, + policy::{Needs, Pick, WorkerFilter, WorkerPicker, WorkerScorer, WorkerSelectionPolicy}, + softmax::sample_by_score_temperature, +}; + +pub const POLICY_NAME: &str = "cache-aware-default"; + +#[derive(Debug)] +struct PositiveOverlap; + +impl WorkerFilter for PositiveOverlap { + fn keep(&self, _request: &RequestInputs<'_>, candidate: &CandidateInputs<'_>) -> bool { + candidate.effective_score > 0.0 + } +} + +/// Cost is the negated affinity, so the generic "lowest cost" reading of the costs agrees with +/// the picker below (which ranks on the score directly to keep the draw byte-identical). +#[derive(Debug)] +struct NegatedAffinity; + +impl WorkerScorer for NegatedAffinity { + fn score( + &self, + _request: &RequestInputs<'_>, + candidates: &[CandidateInputs<'_>], + costs: &mut [f64], + ) { + for (candidate, cost) in candidates.iter().zip(costs) { + *cost += -candidate.effective_score; + } + } +} + +#[derive(Debug)] +struct AffinityGroupPicker { + temperature: f32, +} + +impl WorkerPicker for AffinityGroupPicker { + fn pick( + &self, + _request: &RequestInputs<'_>, + candidates: &[CandidateInputs<'_>], + _costs: &[f64], + ) -> Pick { + if candidates.is_empty() { + return Pick::None; + } + let scores: Vec = candidates.iter().map(|c| c.effective_score).collect(); + let selected = if self.temperature > 0.0 { + sample_by_score_temperature(&scores, self.temperature) + } else { + // `max_by` keeps the last of equal maxima, as the pre-policy code did; the group + // expansion below makes which one immaterial. + scores + .iter() + .enumerate() + .max_by(|a, b| a.1.total_cmp(b.1)) + .map(|(i, _)| i) + }; + let Some(selected) = selected else { + return Pick::None; + }; + let selected_score = scores[selected]; + Pick::Group( + (0..scores.len()) + .filter(|&i| scores[i] == selected_score) + .collect(), + ) + } +} + +pub(super) fn policy(temperature: f32) -> WorkerSelectionPolicy { + WorkerSelectionPolicy::new( + POLICY_NAME, + Needs::default(), + vec![Box::new(PositiveOverlap)], + vec![Box::new(NegatedAffinity)], + Box::new(AffinityGroupPicker { temperature }), + ) +} diff --git a/model_gateway/src/policies/cost/inputs.rs b/model_gateway/src/policies/cost/inputs.rs new file mode 100644 index 0000000000..3bb78f93bd --- /dev/null +++ b/model_gateway/src/policies/cost/inputs.rs @@ -0,0 +1,46 @@ +//! Per-request and per-worker inputs a selection policy sees. +//! +//! The host (cache-aware routing) gathers these once per request from the index or tree and the +//! optimistic accounting, so a policy never touches a lock or a tree itself; nothing a policy +//! must not do is reachable from them (no handles to workers, no mutable shared state). + +/// The request being routed. +#[derive(Debug, Clone, Copy)] +pub struct RequestInputs<'a> { + /// Prompt length in tokens (the routing key's length for text routing). + pub prompt_tokens: usize, + /// Cache block size in tokens, at least 1. + pub block_size: usize, + /// Prompt length in blocks, at least 1 (an empty prompt still costs one block of accounting). + pub request_blocks: usize, + /// Mean in-flight request count over the healthy fleet. + pub avg_load: f64, + /// Chain hash of the prompt's blocks, position `i` covering blocks `0..=i`, when the host + /// computed them (event-driven and token routing). Policies that key on prefixes need them; + /// without them they fall back to load-only behaviour. + pub prefix_hashes: Option<&'a [u64]>, +} + +/// One eligible worker as the policy sees it. +#[derive(Debug, Clone)] +pub struct CandidateInputs<'a> { + /// Position in the host's worker slice; policies return picks by position in the candidate + /// slice, the host maps them back. + pub idx: usize, + /// Worker URL: the stable identity for deterministic tie-breaking and hashing. + pub url: &'a str, + /// Prefix blocks this worker holds on the device tier (GPU), undecayed, or the deeper + /// placement the optimistic accounting predicts for it. + pub device_blocks: f64, + /// The host's decayed affinity score (device overlap after the waiting-prefill decay); what + /// the cache-aware decision ranks on. + pub effective_score: f64, +} + +impl CandidateInputs<'_> { + /// Prompt tokens this worker would still have to prefill after its device-resident prefix. + pub fn uncached_prompt_tokens(&self, request: &RequestInputs<'_>) -> usize { + let cached = (self.device_blocks.max(0.0) * request.block_size as f64) as usize; + request.prompt_tokens.saturating_sub(cached) + } +} diff --git a/model_gateway/src/policies/cost/mod.rs b/model_gateway/src/policies/cost/mod.rs new file mode 100644 index 0000000000..372aac5a62 --- /dev/null +++ b/model_gateway/src/policies/cost/mod.rs @@ -0,0 +1,29 @@ +//! Cost-function worker selection. +//! +//! Cache-aware routing gathers, once per request, what it knows about every candidate worker (its +//! prefix overlap, the decayed affinity score the decision ranks on, and the placement the +//! optimistic accounting predicts for it) and hands that to a *selection policy*: a pipeline of +//! filters, additive cost scorers and one picker, registered by name. +//! +//! - [`catalog::DEFAULT_POLICY`] reproduces the pre-policy cache-aware decision exactly (an affinity +//! group that the host resolves with its pressure gate and expected-wait selector). +//! - [`accounting::OptimisticAccounting`] closes the window between a dispatch and the engine's +//! first event, when enabled. +//! +//! The cost of the stage itself is measured by `benches/policy_selection.rs`. + +pub mod accounting; +pub mod catalog; +pub mod inputs; +pub mod policy; +pub mod softmax; + +mod default; + +#[cfg(test)] +mod sim_tests; + +pub use accounting::OptimisticAccounting; +pub use catalog::{build, default_policy, CatalogError, DEFAULT_POLICY, POLICY_NAMES}; +pub use inputs::{CandidateInputs, RequestInputs}; +pub use policy::{Needs, Pick, WorkerFilter, WorkerPicker, WorkerScorer, WorkerSelectionPolicy}; diff --git a/model_gateway/src/policies/cost/policy.rs b/model_gateway/src/policies/cost/policy.rs new file mode 100644 index 0000000000..f586ad4aae --- /dev/null +++ b/model_gateway/src/policies/cost/policy.rs @@ -0,0 +1,161 @@ +//! The filter / score / pick pipeline and its result. + +use std::fmt::Debug; + +use super::inputs::{CandidateInputs, RequestInputs}; + +/// Drops candidates before scoring. Every filter must keep a candidate for it to be scored. +pub trait WorkerFilter: Send + Sync + Debug { + fn keep(&self, request: &RequestInputs<'_>, candidate: &CandidateInputs<'_>) -> bool; +} + +/// Adds to each candidate's cost. Lower is better; costs start at zero and scorers are additive, +/// so a policy can combine independent terms. A scorer sees all candidates at once so it can +/// normalise against the fleet (minimum backlog, warmest overlap). +pub trait WorkerScorer: Send + Sync + Debug { + fn score( + &self, + request: &RequestInputs<'_>, + candidates: &[CandidateInputs<'_>], + costs: &mut [f64], + ); +} + +/// Turns the scored candidates into a decision. A picker that keeps its own in-flight +/// accounting (size-weighted reservations released at completion) learns about the host's final +/// dispatch and about completions through the hooks below; all default to no-ops. +pub trait WorkerPicker: Send + Sync + Debug { + fn pick( + &self, + request: &RequestInputs<'_>, + candidates: &[CandidateInputs<'_>], + costs: &[f64], + ) -> Pick; + + /// The host dispatched `request` to `candidate`, after its own gate and credit. Not every + /// pick becomes a dispatch, so reservations belong here rather than in `pick`. + fn on_dispatch(&self, _request: &RequestInputs<'_>, _candidate: &CandidateInputs<'_>) {} + + /// A request on `url` finished, successfully or not. + fn on_request_complete(&self, _url: &str) {} + + /// The router holds `in_flight` requests on `url` right now; a picker + /// keeping per-dispatch reservations releases any beyond that count (see + /// `LoadBalancingPolicy::reconcile_in_flight`). + fn reconcile_in_flight(&self, _url: &str, _in_flight: usize) {} + + /// `url` left the fleet. + fn on_worker_removed(&self, _url: &str) {} +} + +/// A picker's decision, in positions of the candidate slice the policy was given. +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum Pick { + /// No candidate qualifies; the host runs its miss path (expected-wait over the fleet). + None, + /// One worker. The host credits it and dispatches. + Final(usize), + /// An affinity group the host resolves with its pressure gate and expected-wait selector; + /// this is how the pre-policy cache-aware decision is expressed exactly. + Group(Vec), +} + +/// Inputs a policy needs the host to gather before calling it. The default policy needs none of +/// them, so its hot path gathers exactly what the pre-policy decision did. +#[derive(Debug, Clone, Copy, Default)] +pub struct Needs { + /// Chain hashes of the prompt's blocks (`RequestInputs::prefix_hashes`). + pub prefix_hashes: bool, + /// Every eligible worker, not only those with a positive overlap. + pub all_workers: bool, +} + +/// A named selection policy: zero or more filters, zero or more scorers, one picker. +#[derive(Debug)] +pub struct WorkerSelectionPolicy { + name: &'static str, + needs: Needs, + filters: Vec>, + scorers: Vec>, + picker: Box, +} + +impl WorkerSelectionPolicy { + pub fn new( + name: &'static str, + needs: Needs, + filters: Vec>, + scorers: Vec>, + picker: Box, + ) -> Self { + Self { + name, + needs, + filters, + scorers, + picker, + } + } + + pub fn name(&self) -> &'static str { + self.name + } + + pub fn needs(&self) -> Needs { + self.needs + } + + /// Tell the picker which candidate the host finally dispatched to. + pub fn on_dispatch(&self, request: &RequestInputs<'_>, candidate: &CandidateInputs<'_>) { + self.picker.on_dispatch(request, candidate); + } + + /// Tell the picker a request on `url` completed. + pub fn on_request_complete(&self, url: &str) { + self.picker.on_request_complete(url); + } + + /// Tell the picker how many requests the router holds on `url`. + pub fn reconcile_in_flight(&self, url: &str, in_flight: usize) { + self.picker.reconcile_in_flight(url, in_flight); + } + + /// Tell the picker `url` left the fleet. + pub fn on_worker_removed(&self, url: &str) { + self.picker.on_worker_removed(url); + } + + /// Run the pipeline. Picks are positions in `candidates` as given by the caller, whatever + /// the filters removed. + pub fn select(&self, request: &RequestInputs<'_>, candidates: &[CandidateInputs<'_>]) -> Pick { + if candidates.is_empty() { + return Pick::None; + } + let kept: Vec = (0..candidates.len()) + .filter(|&i| { + self.filters + .iter() + .all(|filter| filter.keep(request, &candidates[i])) + }) + .collect(); + if kept.is_empty() { + return Pick::None; + } + let compact: Vec>; + let view: &[CandidateInputs<'_>] = if kept.len() == candidates.len() { + candidates + } else { + compact = kept.iter().map(|&i| candidates[i].clone()).collect(); + &compact + }; + let mut costs = vec![0.0f64; view.len()]; + for scorer in &self.scorers { + scorer.score(request, view, &mut costs); + } + match self.picker.pick(request, view, &costs) { + Pick::None => Pick::None, + Pick::Final(i) => Pick::Final(kept[i]), + Pick::Group(group) => Pick::Group(group.into_iter().map(|i| kept[i]).collect()), + } + } +} diff --git a/model_gateway/src/policies/cost/sim_tests.rs b/model_gateway/src/policies/cost/sim_tests.rs new file mode 100644 index 0000000000..5596b9fa7f --- /dev/null +++ b/model_gateway/src/policies/cost/sim_tests.rs @@ -0,0 +1,324 @@ +//! A toy fleet simulation that checks each policy does what it claims, before any engine is +//! involved: it beats a load-blind baseline when prefixes repeat (cache awareness) and it does not +//! collapse onto one worker under a hot prefix (load awareness). Only orderings are asserted; the +//! numbers are printed so a `--nocapture` run reports the relative TTFTs. +//! +//! The model: `N` workers, each with a FIFO block cache of `cache_blocks` and a prefill queue that +//! drains `drain_tokens_per_tick` per tick. A request's TTFT, in ticks, is the queue ahead of it +//! plus its own uncached tokens, over the drain rate. Dispatching inserts the prompt's blocks and +//! queues its uncached tokens. One tick is one simulated second for the policies' time parameters. +#![expect( + clippy::print_stderr, + reason = "the simulation reports its numbers under --nocapture" +)] + +use std::collections::{HashSet, VecDeque}; + +use rand::{rngs::StdRng, RngExt, SeedableRng}; + +use super::{build, CandidateInputs, Pick, RequestInputs, WorkerSelectionPolicy}; + +const BLOCK: usize = 16; + +struct SimWorker { + url: String, + cache: VecDeque, + cached: HashSet, + cache_blocks: usize, + /// Prefill tokens queued, oldest first; each entry is one request. + queue: VecDeque, +} + +impl SimWorker { + fn backlog(&self) -> f64 { + self.queue.iter().sum() + } + + fn overlap(&self, prompt: &[u64]) -> usize { + prompt + .iter() + .take_while(|hash| self.cached.contains(hash)) + .count() + } + + fn insert(&mut self, prompt: &[u64]) { + for &hash in prompt { + if self.cached.insert(hash) { + self.cache.push_back(hash); + if self.cache.len() > self.cache_blocks { + if let Some(evicted) = self.cache.pop_front() { + self.cached.remove(&evicted); + } + } + } + } + } + + /// Drain the queue by `tokens`; returns how many requests completed. + fn drain(&mut self, mut tokens: f64) -> usize { + let mut completed = 0; + while tokens > 0.0 { + let Some(front) = self.queue.front_mut() else { + break; + }; + if *front <= tokens { + tokens -= *front; + self.queue.pop_front(); + completed += 1; + } else { + *front -= tokens; + tokens = 0.0; + } + } + completed + } +} + +#[derive(Clone, Copy)] +struct Scenario { + name: &'static str, + workers: usize, + cache_blocks: usize, + drain_tokens_per_tick: f64, + sessions: usize, + prefix_blocks: usize, + suffix_blocks: usize, + arrivals_per_tick: usize, + /// Share of requests that go to session 0. + hot_share: f64, + requests: usize, +} + +enum Chooser<'a> { + Random, + Sticky, + Policy(&'a WorkerSelectionPolicy), +} + +struct Outcome { + mean_ttft: f64, + p99_ttft: f64, +} + +fn chain_hashes( + session: u64, + prefix_blocks: usize, + suffix_seed: u64, + suffix_blocks: usize, +) -> Vec { + (0..prefix_blocks as u64) + .map(|i| (session << 32) | i) + .chain((0..suffix_blocks as u64).map(|i| (1 << 63) | (suffix_seed << 16) | i)) + .collect() +} + +fn run(scenario: Scenario, chooser: Chooser<'_>, seed: u64) -> Outcome { + let mut rng = StdRng::seed_from_u64(seed); + let mut workers: Vec = (0..scenario.workers) + .map(|i| SimWorker { + url: format!("http://w{i:03}:8000"), + cache: VecDeque::new(), + cached: HashSet::new(), + cache_blocks: scenario.cache_blocks, + queue: VecDeque::new(), + }) + .collect(); + let mut ttfts: Vec = Vec::with_capacity(scenario.requests); + let warmup = scenario.requests / 5; + let mut issued = 0usize; + let mut suffix_seed = 0u64; + while issued < scenario.requests { + for worker in &mut workers { + let completed = worker.drain(scenario.drain_tokens_per_tick); + if let Chooser::Policy(policy) = &chooser { + for _ in 0..completed { + policy.on_request_complete(&worker.url); + } + } + } + for _ in 0..scenario.arrivals_per_tick { + if issued == scenario.requests { + break; + } + let session = if rng.random::() < scenario.hot_share { + 0 + } else { + rng.random_range(0..scenario.sessions as u64) + }; + suffix_seed += 1; + let prompt = chain_hashes( + session, + scenario.prefix_blocks, + suffix_seed, + scenario.suffix_blocks, + ); + let prompt_tokens = prompt.len() * BLOCK; + let overlaps: Vec = workers.iter().map(|w| w.overlap(&prompt)).collect(); + let backlogs: Vec = workers.iter().map(SimWorker::backlog).collect(); + let inflight: Vec = workers.iter().map(|w| w.queue.len()).collect(); + let avg_load = inflight.iter().sum::() as f64 / workers.len() as f64; + + let lowest_backlog = |rows: &[usize]| { + rows.iter() + .copied() + .min_by(|&a, &b| { + backlogs[a] + .total_cmp(&backlogs[b]) + .then_with(|| workers[a].url.cmp(&workers[b].url)) + }) + .expect("non-empty") + }; + let all: Vec = (0..workers.len()).collect(); + let chosen = match &chooser { + Chooser::Random => rng.random_range(0..workers.len()), + Chooser::Sticky => { + let best = overlaps.iter().copied().max().unwrap_or(0); + let tied: Vec = all + .iter() + .copied() + .filter(|&i| overlaps[i] == best) + .collect(); + lowest_backlog(&tied) + } + Chooser::Policy(policy) => { + let request = RequestInputs { + prompt_tokens, + block_size: BLOCK, + request_blocks: prompt.len().max(1), + avg_load, + prefix_hashes: Some(&prompt), + }; + let inputs: Vec> = workers + .iter() + .enumerate() + .map(|(i, w)| CandidateInputs { + idx: i, + url: &w.url, + device_blocks: overlaps[i] as f64, + effective_score: overlaps[i] as f64, + }) + .collect(); + let chosen = match policy.select(&request, &inputs) { + Pick::None => lowest_backlog(&all), + Pick::Final(row) => inputs[row].idx, + Pick::Group(rows) => { + let group: Vec = rows.iter().map(|&r| inputs[r].idx).collect(); + lowest_backlog(&group) + } + }; + policy.on_dispatch(&request, &inputs[chosen]); + chosen + } + }; + let uncached = (prompt.len() - overlaps[chosen]) as f64 * BLOCK as f64; + let ttft = (backlogs[chosen] + uncached) / scenario.drain_tokens_per_tick; + if issued >= warmup { + ttfts.push(ttft); + } + workers[chosen].queue.push_back(uncached); + workers[chosen].insert(&prompt); + issued += 1; + } + } + ttfts.sort_by(f64::total_cmp); + let mean_ttft = ttfts.iter().sum::() / ttfts.len() as f64; + let p99_ttft = ttfts[(ttfts.len() * 99 / 100).min(ttfts.len() - 1)]; + Outcome { + mean_ttft, + p99_ttft, + } +} + +/// Policies under test: the product's default. +fn policies() -> Vec<(&'static str, WorkerSelectionPolicy)> { + vec![( + "cache-aware-default", + build("cache-aware-default", 0.0).unwrap(), + )] +} + +const REPEAT_HEAVY: Scenario = Scenario { + name: "repeat_heavy", + workers: 8, + cache_blocks: 1024, + drain_tokens_per_tick: 512.0, + sessions: 48, + prefix_blocks: 128, + suffix_blocks: 8, + arrivals_per_tick: 4, + hot_share: 0.0, + requests: 1500, +}; + +/// Three quarters of eight arrivals per tick share one prefix: a worker that keeps all of them +/// drains 512 tokens per tick against 768 queued, so staying sticky means an unbounded queue. +const HOT_PREFIX: Scenario = Scenario { + name: "hot_prefix", + sessions: 24, + arrivals_per_tick: 8, + hot_share: 0.75, + requests: 2_400, + ..REPEAT_HEAVY +}; + +#[test] +fn every_policy_beats_random_when_prefixes_repeat() { + let random = run(REPEAT_HEAVY, Chooser::Random, 1); + let sticky = run(REPEAT_HEAVY, Chooser::Sticky, 1); + eprintln!( + "{}: random mean {:.2} p99 {:.2}; sticky mean {:.2} p99 {:.2}", + REPEAT_HEAVY.name, random.mean_ttft, random.p99_ttft, sticky.mean_ttft, sticky.p99_ttft + ); + let outcomes: Vec<(&str, Outcome)> = policies() + .iter() + .map(|(name, policy)| (*name, run(REPEAT_HEAVY, Chooser::Policy(policy), 1))) + .collect(); + for (name, outcome) in &outcomes { + eprintln!( + "{}: {name} mean {:.2} p99 {:.2}", + REPEAT_HEAVY.name, outcome.mean_ttft, outcome.p99_ttft + ); + } + for (name, outcome) in &outcomes { + assert!( + outcome.mean_ttft < 0.5 * random.mean_ttft, + "{name}: mean TTFT {:.2} is not under half of random's {:.2}", + outcome.mean_ttft, + random.mean_ttft + ); + } +} + +#[test] +fn every_policy_spreads_a_hot_prefix() { + let random = run(HOT_PREFIX, Chooser::Random, 2); + let sticky = run(HOT_PREFIX, Chooser::Sticky, 2); + eprintln!( + "{}: random mean {:.2} p99 {:.2}; sticky mean {:.2} p99 {:.2}", + HOT_PREFIX.name, random.mean_ttft, random.p99_ttft, sticky.mean_ttft, sticky.p99_ttft + ); + for (name, policy) in policies() { + let outcome = run(HOT_PREFIX, Chooser::Policy(&policy), 2); + eprintln!( + "{}: {name} mean {:.2} p99 {:.2}", + HOT_PREFIX.name, outcome.mean_ttft, outcome.p99_ttft + ); + // The cache-aware default is the sticky baseline by construction: its + // load relief is the host's spill gate and expected-wait selector, + // which this simulation does not model. + if name != "cache-aware-default" { + assert!( + outcome.p99_ttft < sticky.p99_ttft, + "{name}: p99 TTFT {:.2} does not beat sticky's {:.2}", + outcome.p99_ttft, + sticky.p99_ttft + ); + } + assert!( + outcome.mean_ttft < random.mean_ttft, + "{name}: mean TTFT {:.2} does not beat random's {:.2}", + outcome.mean_ttft, + random.mean_ttft + ); + } +} diff --git a/model_gateway/src/policies/cost/softmax.rs b/model_gateway/src/policies/cost/softmax.rs new file mode 100644 index 0000000000..0b09e1c697 --- /dev/null +++ b/model_gateway/src/policies/cost/softmax.rs @@ -0,0 +1,32 @@ +//! Shared selection arithmetic: the softmax draw over scores. + +use rand::RngExt; + +/// Softmax selection over min-max normalised *scores* (higher is better): the draw cache-aware +/// routing has always used. The normalisation makes temperature scale-free, the best candidate's +/// exponent is exactly 0 (overflow-safe), a degenerate spread is a uniform draw, and the +/// inverse-CDF walk falls back to the last row against floating-point drift. Returns a position. +pub fn sample_by_score_temperature(scores: &[f64], temperature: f32) -> Option { + let first = *scores.first()?; + let (min, max) = scores + .iter() + .fold((first, first), |(min, max), &s| (min.min(s), max.max(s))); + let range = max - min; + if range <= 0.0 { + return Some(rand::rng().random_range(0..scores.len())); + } + let weights: Vec = scores + .iter() + .map(|&s| (((s - min) / range - 1.0) / f64::from(temperature)).exp()) + .collect(); + let total: f64 = weights.iter().sum(); + let draw = rand::rng().random::() * total; + let mut cumulative = 0.0; + for (position, weight) in weights.iter().enumerate() { + cumulative += weight; + if cumulative >= draw { + return Some(position); + } + } + Some(scores.len() - 1) +} diff --git a/model_gateway/src/policies/factory.rs b/model_gateway/src/policies/factory.rs index e18213870e..4254235085 100755 --- a/model_gateway/src/policies/factory.rs +++ b/model_gateway/src/policies/factory.rs @@ -54,6 +54,8 @@ impl PolicyFactory { cache_index, cache_ttl_secs, cache_boundaries, + selection_policy, + selection_accounting_ttl_ms, } => { let config = CacheAwareConfig { cache_threshold: *cache_threshold, @@ -69,6 +71,8 @@ impl PolicyFactory { cache_index: *cache_index, cache_ttl_secs: *cache_ttl_secs, cache_boundaries: cache_boundaries.clone(), + selection_policy: selection_policy.clone(), + selection_accounting_ttl_ms: *selection_accounting_ttl_ms, }; Arc::new(CacheAwarePolicy::with_config(config)) } @@ -168,6 +172,8 @@ mod tests { cache_index: Default::default(), cache_ttl_secs: 180, cache_boundaries: Vec::new(), + selection_policy: None, + selection_accounting_ttl_ms: 0, }); assert_eq!(policy.name(), "cache_aware"); diff --git a/model_gateway/src/policies/least_load.rs b/model_gateway/src/policies/least_load.rs index c1f37c7a9e..9b39a19908 100644 --- a/model_gateway/src/policies/least_load.rs +++ b/model_gateway/src/policies/least_load.rs @@ -1,10 +1,11 @@ use std::{ - collections::HashMap, - sync::{Arc, RwLock}, + collections::{HashMap, VecDeque}, + sync::{Arc, Mutex, RwLock}, + time::Instant, }; use openai_protocol::worker::WorkerLoadResponse; -use rand::RngExt; +use rand::{rngs::StdRng, RngExt, SeedableRng}; use tracing::debug; use super::{get_healthy_worker_indices, LoadBalancingPolicy, SelectWorkerInfo}; @@ -14,11 +15,58 @@ pub use crate::worker::expected_wait::{ }; use crate::worker::{expected_wait::ExpectedWait, load_state::LoadSnapshot, Worker}; -/// Since-poll dispatch tally for one worker. -#[derive(Clone, Copy, Debug, Default)] +/// Dispatches kept per worker between reports. A worker that never reports +/// (a dark fleet is scored by its live in-flight count instead) would +/// otherwise accumulate them without end; past this the oldest are dropped. +const SINCE_POLL_DISPATCHES_KEPT: usize = 8_192; + +/// Expected waits within this of the minimum are one tie and are drawn from +/// uniformly. Idle workers with identical reports score exactly equal, but a +/// report's throughput or a KV digit sets two otherwise identical workers a +/// few nanoseconds apart, and an exact-equality tie then hands every request +/// to the lower index: a cold fleet never spreads that way. +const TIE_EPSILON_SECS: f64 = 1e-6; + +/// Since-poll dispatch tally for one worker: the token-work and request count +/// of the dispatches no report has reflected yet, with each dispatch's +/// instant, so a report that says when the engine was sampled releases +/// exactly the dispatches it saw and keeps the later ones as credit. +#[derive(Clone, Debug, Default)] struct SincePollDispatch { tokens: u64, requests: u64, + /// `(dispatched_at, tokens)` per dispatch, oldest first. + dispatches: VecDeque<(Instant, u64)>, +} + +impl SincePollDispatch { + fn record(&mut self, at: Instant, tokens: u64) { + self.tokens += tokens; + self.requests += 1; + self.dispatches.push_back((at, tokens)); + if self.dispatches.len() > SINCE_POLL_DISPATCHES_KEPT { + self.pop_oldest(); + } + } + + /// Release the dispatches a report sampled at `sampled_at` already + /// reflects: those made at or before it. + fn release_through(&mut self, sampled_at: Instant) { + while self + .dispatches + .front() + .is_some_and(|&(at, _)| at <= sampled_at) + { + self.pop_oldest(); + } + } + + fn pop_oldest(&mut self) { + if let Some((_, tokens)) = self.dispatches.pop_front() { + self.tokens -= tokens; + self.requests -= 1; + } + } } /// Least-(token-)work routing — route to the worker with the lowest estimated @@ -89,9 +137,11 @@ struct SincePollDispatch { pub struct LeastLoadPolicy { /// Cached load reports from the worker monitor (keyed by worker URL). cached_loads: RwLock>, - /// Per-worker dispatch tally since the last load poll (keyed by worker - /// URL); reset when a fresh report arrives. Token-work feeds the score's - /// in-flight term; the request count feeds the waiting-queue veto. + /// Per-worker dispatch tally since the last load report (keyed by worker + /// URL): a report that carries `sampled_at` releases the dispatches made + /// up to that instant, one without it resets the tally. Token-work feeds + /// the score's in-flight term; the request count feeds the waiting-queue + /// veto. inflight_tokens: RwLock>, /// KV-pressure weight `λ_t` (seconds). kv_pressure_weight: f64, @@ -103,6 +153,9 @@ pub struct LeastLoadPolicy { default_throughput: f64, /// Per-worker waiting-queue cap; `0` disables the veto. max_waiting_requests: u32, + /// Seeded source for the tie draw when set (tests reproduce a selection + /// sequence with it); the thread's generator otherwise. + tie_rng: Option>, } /// Everything one expected-wait score reads besides the worker itself. @@ -158,6 +211,24 @@ impl LeastLoadPolicy { DEFAULT_THROUGHPUT }, max_waiting_requests, + tie_rng: None, + } + } + + /// Draw ties from a seeded generator instead of the thread's. + pub fn with_tie_break_seed(mut self, seed: u64) -> Self { + self.tie_rng = Some(Mutex::new(StdRng::seed_from_u64(seed))); + self + } + + /// Uniform draw in `0..n`: the reservoir step of the argmin's tie-break. + fn tie_draw(&self, n: u32) -> u32 { + match &self.tie_rng { + Some(rng) => rng + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()) + .random_range(0..n), + None => rand::rng().random_range(0..n), } } @@ -182,14 +253,16 @@ impl LeastLoadPolicy { .read() .unwrap_or_else(|poisoned| poisoned.into_inner()) .contains_key(url); - let dispatch = self + let inflight = self .inflight_tokens .read() - .unwrap_or_else(|poisoned| poisoned.into_inner()) - .get(url) - .copied() - .unwrap_or_default(); - (has_load, dispatch.tokens, dispatch.requests) + .unwrap_or_else(|poisoned| poisoned.into_inner()); + let dispatch = inflight.get(url); + ( + has_load, + dispatch.map_or(0, |dispatch| dispatch.tokens), + dispatch.map_or(0, |dispatch| dispatch.requests), + ) } /// Expected-wait score for a worker (lower is better). @@ -212,17 +285,11 @@ impl LeastLoadPolicy { let url = worker.url(); match Self::fresh_load(loads, complete_snapshot, url) { Some(load) => { - let inflight_tokens = inflight.get(url).copied().unwrap_or_default().tokens; + let inflight_tokens = inflight.get(url).map_or(0, |dispatch| dispatch.tokens); let queued_tokens = self.queued_tokens(load); - let live_throughput = load.total_gen_throughput(); - let throughput = if live_throughput > 0.0 { - live_throughput - } else { - self.default_throughput - }; ExpectedWait::new( queued_tokens, - throughput, + self.drain_rate(load), load.effective_token_usage(), self.kv_pressure_weight, ) @@ -258,6 +325,74 @@ impl LeastLoadPolicy { loads.and_then(|map| map.get(url)) } + /// The drain rate a reporting worker's wait is priced at: its live + /// generation rate, else the configured default. + fn drain_rate(&self, load: &WorkerLoadResponse) -> f64 { + let live = load.total_gen_throughput(); + if live > 0.0 { + live + } else { + self.default_throughput + } + } + + /// The scoring inputs for one pass over `candidates`: the nominal drain + /// rate (mean of the positive reports) that stands in for a worker + /// missing a fresh snapshot; whether anyone reports at all, which + /// separates a partial gap (estimate the missing worker at the nominal + /// rate) from a dark fleet (join-shortest-queue on live in-flight); and + /// the best-known reporting peer's score, which a worker without a report + /// starts from: never better than a worker whose load is known, never + /// starved by one. The baseline is computed only when some candidate + /// lacks a report, so the common all-reporting case pays nothing for it. + fn score_inputs<'a>( + &self, + workers: &[Arc], + candidates: &[usize], + loads: Option<&'a HashMap>, + complete_snapshot: Option<&'a LoadSnapshot>, + inflight: &'a HashMap, + ) -> ScoreInputs<'a> { + let (tp_sum, tp_count) = candidates + .iter() + .filter_map(|&i| Self::fresh_load(loads, complete_snapshot, workers[i].url())) + .map(|l| l.total_gen_throughput()) + .filter(|t| *t > 0.0) + .fold((0.0, 0u32), |(s, n), t| (s + t, n + 1)); + let nominal_throughput = if tp_count > 0 { + tp_sum / tp_count as f64 + } else { + self.default_throughput + }; + let reporting = candidates + .iter() + .filter(|&&i| Self::fresh_load(loads, complete_snapshot, workers[i].url()).is_some()) + .count(); + let mut inputs = ScoreInputs { + loads, + complete_snapshot, + inflight, + nominal_throughput, + fleet_has_loads: reporting > 0, + peer_baseline: 0.0, + }; + if reporting < candidates.len() { + let best_known = candidates + .iter() + .filter(|&&i| { + Self::fresh_load(loads, complete_snapshot, workers[i].url()).is_some() + }) + .map(|&i| self.score(&workers[i], &inputs)) + .fold(f64::INFINITY, f64::min); + // Nobody reports: the dark-fleet arm scores by live in-flight and + // never reads the baseline; keep it neutral rather than infinite. + if best_known.is_finite() { + inputs.peer_baseline = best_known; + } + } + inputs + } + /// Waiting-queue token-work for a worker. /// /// Prefers the backend's own `num_waiting_uncached_tokens`. Backends that @@ -335,9 +470,7 @@ impl LeastLoadPolicy { Some(load) => { let since_poll = inflight_guard .get(url) - .copied() - .unwrap_or_default() - .requests; + .map_or(0, |dispatch| dispatch.requests); (load.total_waiting_reqs().max(0) as u64) + since_poll < cap } None => true, @@ -348,91 +481,33 @@ impl LeastLoadPolicy { }; let (&first, rest) = candidates.split_first()?; - // Nominal throughput (mean of positive reports) stands in for a - // worker missing a fresh snapshot; `fleet_has_loads` distinguishes a - // partial gap (estimate that worker's drain time at the nominal rate) - // from a fully dark fleet (fall back to join-shortest-queue). - let (tp_sum, tp_count) = candidates - .iter() - .filter_map(|&i| Self::fresh_load(loads, complete_snapshot, workers[i].url())) - .map(|l| l.total_gen_throughput()) - .filter(|t| *t > 0.0) - .fold((0.0, 0u32), |(s, n), t| (s + t, n + 1)); - let nominal_throughput = if tp_count > 0 { - tp_sum / tp_count as f64 - } else { - self.default_throughput - }; - let fleet_has_loads = candidates - .iter() - .any(|&i| Self::fresh_load(loads, complete_snapshot, workers[i].url()).is_some()); - // Held across selection so the in-flight estimate stays consistent and // the chosen worker can be credited before the guard is released. let mut inflight = self .inflight_tokens .write() .unwrap_or_else(|poisoned| poisoned.into_inner()); + let inputs = self.score_inputs(workers, candidates, loads, complete_snapshot, &inflight); - // Argmin with reservoir tie-breaking: equal-score workers (the common - // idle/homogeneous case scores exactly equal) are sampled uniformly - // instead of first-index-wins, which herded ties onto one worker. - // A worker without a fresh report scores as the best-known reporting - // peer plus its own in-flight: never better than a worker whose load - // is known, never starved by one. Computed only when some candidate - // lacks a report, so the common all-reporting case pays nothing. - let peer_baseline = if candidates - .iter() - .all(|&i| Self::fresh_load(loads, complete_snapshot, workers[i].url()).is_some()) - { - 0.0 - } else { - let known = ScoreInputs { - loads, - complete_snapshot, - inflight: &inflight, - nominal_throughput, - fleet_has_loads, - peer_baseline: 0.0, - }; - let best_known = candidates - .iter() - .filter(|&&i| { - Self::fresh_load(loads, complete_snapshot, workers[i].url()).is_some() - }) - .map(|&i| self.score(&workers[i], &known)) - .fold(f64::INFINITY, f64::min); - // Nobody reports: the dark-fleet arm scores by live in-flight and - // never reads the baseline; keep it neutral rather than infinite. - if best_known.is_finite() { - best_known - } else { - 0.0 - } - }; - let mut rng = rand::rng(); + // Argmin with reservoir tie-breaking: workers within + // `TIE_EPSILON_SECS` of the minimum (the common idle/homogeneous case + // scores equal to the digit) are sampled uniformly instead of + // first-index-wins, which herded ties onto one worker. let mut best = first; - let inputs = ScoreInputs { - loads, - complete_snapshot, - inflight: &inflight, - nominal_throughput, - fleet_has_loads, - peer_baseline, - }; let mut best_score = self.score(&workers[best], &inputs); let mut tied = 1u32; for &idx in rest { let s = self.score(&workers[idx], &inputs); - if s < best_score { + if s < best_score - TIE_EPSILON_SECS { best = idx; best_score = s; tied = 1; - } else if s == best_score { + } else if (s - best_score).abs() <= TIE_EPSILON_SECS { // Keep each tying candidate with probability 1/k so the final // pick is uniform over all ties without collecting them. tied += 1; - if rng.random_range(0..tied) == 0 { + best_score = best_score.min(s); + if self.tie_draw(tied) == 0 { best = idx; } } @@ -441,9 +516,10 @@ impl LeastLoadPolicy { // In-flight correction: credit the chosen worker with this request's // token-work until its next poll refreshes the snapshot. let req_tokens = self.request_tokens(info); - let tally = inflight.entry(workers[best].url().to_string()).or_default(); - tally.tokens += req_tokens; - tally.requests += 1; + inflight + .entry(workers[best].url().to_string()) + .or_default() + .record(Instant::now(), req_tokens); drop(inflight); debug!( @@ -471,10 +547,23 @@ impl LeastLoadPolicy { }; cached.extend(loads.iter().map(|(k, v)| (k.clone(), v.clone()))); after_publish(); - // A fresh snapshot already reflects work up to the poll, so reset the - // since-poll in-flight estimate for the workers it covers. - for url in loads.keys() { - inflight.insert(url.clone(), SincePollDispatch::default()); + // A report reflects the work dispatched up to the instant the engine + // was sampled: release those dispatches and keep the later ones as + // credit, so a record republished unchanged (the monitor republishes + // the shared snapshot when a pushed record arrives) releases nothing + // twice and a dispatch made after the sample is not lost. A report + // with no sample instant resets the tally as before. + for (url, load) in loads { + match load.sampled_at { + Some(sampled_at) => { + if let Some(dispatch) = inflight.get_mut(url) { + dispatch.release_through(sampled_at); + } + } + None => { + inflight.insert(url.clone(), SincePollDispatch::default()); + } + } } } } @@ -601,6 +690,71 @@ mod tests { ) } + #[test] + fn a_cold_fleet_of_128_spreads_a_thousand_misses() { + // No report anywhere: every worker scores its live in-flight count, + // zero for all, so the whole fleet ties on every request. + let policy = LeastLoadPolicy::new().with_tie_break_seed(7); + let workers: Vec> = + (0..128).map(|i| mk(&format!("http://w{i}:8000"))).collect(); + let info = SelectWorkerInfo::default(); + let mut hits = vec![0usize; workers.len()]; + for _ in 0..1000 { + hits[policy.select_worker(&workers, &info).unwrap()] += 1; + } + let mean = 1000.0 / workers.len() as f64; + assert!( + hits.iter().all(|&h| h >= 1), + "a worker never chosen: {hits:?}" + ); + let max = *hits.iter().max().unwrap(); + assert!( + max as f64 <= 3.0 * mean, + "max {max} over a mean of {mean:.1}: {hits:?}" + ); + } + + #[test] + fn a_strictly_cheaper_worker_wins_every_time() { + let policy = LeastLoadPolicy::new().with_tie_break_seed(7); + let workers: Vec> = + (0..128).map(|i| mk(&format!("http://w{i}:8000"))).collect(); + for (i, worker) in workers.iter().enumerate() { + if i != 77 { + worker.increment_load(); + } + } + let info = SelectWorkerInfo::default(); + for _ in 0..1000 { + assert_eq!(policy.select_worker(&workers, &info), Some(77)); + } + } + + #[test] + fn scores_within_epsilon_of_the_minimum_tie() { + // Four idle workers whose reports differ by a KV digit far below the + // epsilon; an exact-equality tie handed everything to the first. + let policy = LeastLoadPolicy::new().with_tie_break_seed(7); + let workers: Vec> = + (0..4).map(|i| mk(&format!("http://w{i}:8000"))).collect(); + let mut loads = HashMap::new(); + for (i, worker) in workers.iter().enumerate() { + loads.insert( + worker.url().to_string(), + make_load(0, i as f64 * 1e-9, 100.0), + ); + } + policy.update_loads(&loads); + let info = SelectWorkerInfo::default(); + let mut hits = [0usize; 4]; + for _ in 0..400 { + hits[policy.select_worker(&workers, &info).unwrap()] += 1; + // Release the winner's credit so every pick sees the same four scores. + policy.update_loads(&loads); + } + assert!(hits.iter().all(|&h| h >= 50), "{hits:?}"); + } + #[test] fn equal_score_ties_spread_across_workers() { // Three identically-loaded workers score exactly equal; the argmin @@ -750,6 +904,59 @@ mod tests { ); } + #[test] + fn a_report_releases_the_dispatches_made_up_to_its_sample_and_keeps_the_rest() { + let policy = LeastLoadPolicy::new(); + let workers = vec![mk("http://a:8000")]; + let info = SelectWorkerInfo::default(); + let mut loads = HashMap::new(); + loads.insert("http://a:8000".to_string(), make_load(0, 0.1, 100.0)); + policy.update_loads(&loads); + + // Two dispatches, then the engine is sampled, then one more. + policy.select_min_expected_wait(&workers, &[0], &info, "test"); + policy.select_min_expected_wait(&workers, &[0], &info, "test"); + std::thread::sleep(std::time::Duration::from_millis(2)); + let sampled_at = Instant::now(); + std::thread::sleep(std::time::Duration::from_millis(2)); + policy.select_min_expected_wait(&workers, &[0], &info, "test"); + assert_eq!(policy.load_state_for_test("http://a:8000"), (true, 3072, 3)); + + let mut sampled = make_load(0, 0.1, 100.0); + sampled.sampled_at = Some(sampled_at); + loads.insert("http://a:8000".to_string(), sampled.clone()); + policy.update_loads(&loads); + assert_eq!( + policy.load_state_for_test("http://a:8000"), + (true, 1024, 1), + "the dispatch after the sample stays as credit" + ); + + // The same record republished (the monitor's shared snapshot on a + // pushed record from another worker) releases nothing more. + policy.update_loads(&loads); + assert_eq!(policy.load_state_for_test("http://a:8000"), (true, 1024, 1)); + + // A report without a sample instant resets, as every poll did before. + loads.insert("http://a:8000".to_string(), make_load(0, 0.1, 100.0)); + policy.update_loads(&loads); + assert_eq!(policy.load_state_for_test("http://a:8000"), (true, 0, 0)); + } + + #[test] + fn the_tally_keeps_a_bounded_history_for_a_worker_that_never_reports() { + let mut dispatch = SincePollDispatch::default(); + let at = Instant::now(); + for _ in 0..(SINCE_POLL_DISPATCHES_KEPT + 100) { + dispatch.record(at, 7); + } + assert_eq!(dispatch.dispatches.len(), SINCE_POLL_DISPATCHES_KEPT); + assert_eq!(dispatch.requests as usize, SINCE_POLL_DISPATCHES_KEPT); + assert_eq!(dispatch.tokens as usize, 7 * SINCE_POLL_DISPATCHES_KEPT); + dispatch.release_through(at); + assert_eq!((dispatch.tokens, dispatch.requests), (0, 0)); + } + #[test] fn waiting_queue_veto_ignores_workers_without_snapshots() { // A dark fleet with a cap configured has no queue evidence to veto diff --git a/model_gateway/src/policies/mod.rs b/model_gateway/src/policies/mod.rs index 94c034af28..22dde6a22c 100644 --- a/model_gateway/src/policies/mod.rs +++ b/model_gateway/src/policies/mod.rs @@ -16,6 +16,7 @@ mod bucket; mod cache_aware; mod cache_namespace; mod consistent_hashing; +pub mod cost; mod dp_min_token; mod factory; mod least_load; @@ -99,6 +100,19 @@ pub trait LoadBalancingPolicy: Send + Sync + Debug { // Default: no-op for policies that don't cache per-worker state } + /// Reconcile the state this policy booked on `worker_url` against the + /// router's live count of requests in flight there. + /// + /// Every request end reaches [`Self::on_request_complete`] through the + /// worker's load guard; this is the safety net behind it. The worker + /// monitor calls it once per load poll with the requests the router still + /// holds on the worker. A dispatch books once and a completion releases + /// once, so whatever a policy holds beyond `in_flight` is a completion + /// that never arrived (a worker that moved under another policy, a sink + /// never installed) and is released here rather than kept for good. + /// Default: no-op for policies that book nothing. + fn reconcile_in_flight(&self, _worker_url: &str, _in_flight: usize) {} + /// Reset any internal state /// /// This is useful for policies that maintain state (e.g., round-robin counters). @@ -187,6 +201,14 @@ pub struct CacheAwareConfig { /// Ascending token positions at which serving engines retain reusable /// prefix state; the hash index keys request heads at these boundaries. pub cache_boundaries: Vec, + /// Worker selection policy run over the gathered per-worker inputs, by + /// name from [`cost::POLICY_NAMES`]. `None` is [`cost::DEFAULT_POLICY`], + /// the affinity-group decision this policy has always made. + pub selection_policy: Option, + /// Lifetime of optimistic dispatch bookings (predicted prefill and prefix + /// placement charged to the chosen worker before the engine reports it). + /// `0` disables (default); set a little above the engine's event lag. + pub selection_accounting_ttl_ms: u64, } impl Default for CacheAwareConfig { @@ -209,6 +231,8 @@ impl Default for CacheAwareConfig { cache_index: CacheIndexKind::Tree, cache_ttl_secs: 180, cache_boundaries: Vec::new(), + selection_policy: None, + selection_accounting_ttl_ms: 0, } } } diff --git a/model_gateway/src/policies/registry.rs b/model_gateway/src/policies/registry.rs index b09292dab0..90a32e1e57 100644 --- a/model_gateway/src/policies/registry.rs +++ b/model_gateway/src/policies/registry.rs @@ -1,6 +1,6 @@ use std::{ collections::{HashMap, HashSet}, - sync::{Arc, OnceLock}, + sync::{Arc, OnceLock, Weak}, }; use dashmap::DashMap; @@ -28,9 +28,29 @@ use crate::{ routers::common::header_utils::{ extract_routing_key_hint_named, parse_routing_tokens_hint, ROUTING_KEY_HINT_MAX_BYTES, }, - worker::{KvEventMonitor, Worker}, + worker::{KvEventMonitor, RequestCompletionSink, Worker, WorkerType}, }; +/// Routes a request completion to the policy that placed the request: the +/// prefill, decode or encode policy by worker type, else the model's policy. +/// Holds the registry weakly so a worker outliving its registry reports to +/// nobody instead of keeping the registry alive. +#[derive(Debug)] +struct PolicyCompletionSink { + registry: Weak, +} + +impl RequestCompletionSink for PolicyCompletionSink { + fn request_completed(&self, worker: &dyn Worker) { + let Some(registry) = self.registry.upgrade() else { + return; + }; + registry + .policy_for_worker(worker) + .on_request_complete(worker.url(), true); + } +} + /// Registry for managing model-to-policy mappings #[derive(Clone)] pub struct PolicyRegistry { @@ -495,6 +515,36 @@ impl PolicyRegistry { /// Called when a worker is added /// Returns the policy that should be used for this worker's model + /// The request-completion observer to install on every worker (see + /// [`RequestCompletionSink`]): policies that book state at dispatch + /// release it when the request's load guard drops. + pub fn completion_sink(self: &Arc) -> Arc { + Arc::new(PolicyCompletionSink { + registry: Arc::downgrade(self), + }) + } + + /// The policy that places requests on `worker`: the prefill, decode or + /// encode policy by worker type, else the model's policy. Completion + /// reports and in-flight reconciliation both go there. + fn policy_for_worker(&self, worker: &dyn Worker) -> Arc { + match worker.worker_type() { + WorkerType::Prefill => self.get_prefill_policy(), + WorkerType::Decode => self.get_decode_policy(), + WorkerType::Encode => self.get_encode_policy(), + WorkerType::Regular => self.get_policy_or_default(worker.model_id()), + } + } + + /// Reconcile the policy that places requests on `worker` against the + /// router's live in-flight count there (see + /// [`LoadBalancingPolicy::reconcile_in_flight`]). The worker monitor + /// calls this once per load poll for every polled worker. + pub fn reconcile_in_flight(&self, worker: &dyn Worker) { + self.policy_for_worker(worker) + .reconcile_in_flight(worker.url(), worker.load()); + } + pub fn on_worker_added( &self, model_id: &str, @@ -1260,6 +1310,8 @@ mod tests { cache_index: Default::default(), cache_ttl_secs: 180, cache_boundaries: Vec::new(), + selection_policy: None, + selection_accounting_ttl_ms: 0, } } @@ -1669,6 +1721,8 @@ mod tests { cache_index: Default::default(), cache_ttl_secs: 180, cache_boundaries: Vec::new(), + selection_policy: None, + selection_accounting_ttl_ms: 0, }, rid_override(ManualAssignmentMode::Delegate), ); @@ -2090,6 +2144,8 @@ mod tests { cache_index: Default::default(), cache_ttl_secs: 180, cache_boundaries: Vec::new(), + selection_policy: None, + selection_accounting_ttl_ms: 0, } } @@ -2138,6 +2194,8 @@ mod tests { cache_index: Default::default(), cache_ttl_secs: 180, cache_boundaries: Vec::new(), + selection_policy: None, + selection_accounting_ttl_ms: 0, }); // Hinted policy is a fresh per-model instance, not the shared default. @@ -2210,6 +2268,8 @@ mod tests { cache_index: Default::default(), cache_ttl_secs: 180, cache_boundaries: Vec::new(), + selection_policy: None, + selection_accounting_ttl_ms: 0, })); for round in 0..64 { @@ -2352,4 +2412,121 @@ mod tests { registry.remove_worker_from_pd_cache_aware("http://prefill-1:8000"); registry.remove_worker_from_pd_cache_aware("http://decode-1:8000"); } + + /// A policy that remembers completions and reconciliation ticks. + #[derive(Debug, Default)] + struct CompletionRecorder { + completed: std::sync::Mutex>, + reconciled: std::sync::Mutex>, + } + + impl LoadBalancingPolicy for CompletionRecorder { + fn select_worker( + &self, + workers: &[Arc], + _info: &SelectWorkerInfo, + ) -> Option { + (!workers.is_empty()).then_some(0) + } + + fn on_request_complete(&self, worker_url: &str, _success: bool) { + self.completed.lock().unwrap().push(worker_url.to_string()); + } + + fn reconcile_in_flight(&self, worker_url: &str, in_flight: usize) { + self.reconciled + .lock() + .unwrap() + .push((worker_url.to_string(), in_flight)); + } + + fn name(&self) -> &'static str { + "completion_recorder" + } + + fn as_any(&self) -> &dyn std::any::Any { + self + } + } + + #[test] + fn completion_sink_reports_to_the_policy_that_owns_the_worker_type() { + use crate::worker::{BasicWorkerBuilder, WorkerLoadGuard}; + + let registry = Arc::new(PolicyRegistry::new(PolicyConfig::Random)); + let prefill_policy = Arc::new(CompletionRecorder::default()); + let decode_policy = Arc::new(CompletionRecorder::default()); + registry.set_prefill_policy(Arc::clone(&prefill_policy) as Arc); + registry.set_decode_policy(Arc::clone(&decode_policy) as Arc); + let sink = registry.completion_sink(); + + let prefill: Arc = Arc::new( + BasicWorkerBuilder::new("grpc://prefill:9000") + .worker_type(WorkerType::Prefill) + .build(), + ); + let decode: Arc = Arc::new( + BasicWorkerBuilder::new("grpc://decode:9000") + .worker_type(WorkerType::Decode) + .build(), + ); + let regular: Arc = Arc::new( + BasicWorkerBuilder::new("grpc://regular:9000") + .worker_type(WorkerType::Regular) + .build(), + ); + for worker in [&prefill, &decode, ®ular] { + worker.set_completion_sink(Some(Arc::clone(&sink))); + } + + drop(WorkerLoadGuard::new(Arc::clone(&prefill), None)); + drop(WorkerLoadGuard::new(Arc::clone(&decode), None)); + // The regular worker's model policy (random) accepts the completion silently. + drop(WorkerLoadGuard::new(Arc::clone(®ular), None)); + + assert_eq!( + prefill_policy.completed.lock().unwrap().as_slice(), + ["grpc://prefill:9000"] + ); + assert_eq!( + decode_policy.completed.lock().unwrap().as_slice(), + ["grpc://decode:9000"] + ); + + // A worker that outlives its registry reports to nobody. + drop(registry); + drop(WorkerLoadGuard::new(Arc::clone(&prefill), None)); + assert_eq!(prefill_policy.completed.lock().unwrap().len(), 1); + } + + #[test] + fn reconcile_in_flight_reaches_the_policy_that_owns_the_worker_type() { + use crate::worker::{BasicWorkerBuilder, WorkerLoadGuard}; + + let registry = Arc::new(PolicyRegistry::new(PolicyConfig::Random)); + let decode_policy = Arc::new(CompletionRecorder::default()); + registry.set_decode_policy(Arc::clone(&decode_policy) as Arc); + let decode: Arc = Arc::new( + BasicWorkerBuilder::new("grpc://decode:9000") + .worker_type(WorkerType::Decode) + .build(), + ); + + // The count handed over is the router's live one: two guards held, + // then none. + let held = [ + WorkerLoadGuard::new(Arc::clone(&decode), None), + WorkerLoadGuard::new(Arc::clone(&decode), None), + ]; + registry.reconcile_in_flight(decode.as_ref()); + drop(held); + registry.reconcile_in_flight(decode.as_ref()); + assert_eq!( + decode_policy.reconciled.lock().unwrap().as_slice(), + [ + ("grpc://decode:9000".to_string(), 2), + ("grpc://decode:9000".to_string(), 0) + ] + ); + } } diff --git a/model_gateway/src/routers/common/overload.rs b/model_gateway/src/routers/common/overload.rs index 9de740b4e8..e91ee08f47 100644 --- a/model_gateway/src/routers/common/overload.rs +++ b/model_gateway/src/routers/common/overload.rs @@ -1,5 +1,10 @@ -//! Shed responses for the absolute worker-overload guard. Failure paths only — -//! nothing here runs for a served request. +//! The answers to a candidate pool whose every worker is vetoed, by the +//! overload thresholds or by the liveness tracker: by default the request +//! steers to the least-loaded of them ([`fallback_if_all_vetoed`]); under +//! `--worker-overload-shed` an all-overloaded pool is refused with a distinct +//! 503 ([`shed_if_all_overloaded`]) while a liveness veto still steers. +//! Nothing here runs while some worker is free of both vetoes: such workers +//! are simply left out of selection. //! //! The verdict is taken from the candidate pool the caller selected over, never //! from the model index: selection narrows by worker type and transport first, @@ -15,6 +20,7 @@ use axum::{ http::{header::RETRY_AFTER, HeaderValue}, response::Response, }; +use rand::RngExt; use tracing::debug; use crate::{ @@ -22,7 +28,8 @@ use crate::{ routers::{common::retry::mark_non_retryable, error}, worker::{ overload::{ - BRANCH_ALL_OVERLOADED_SHED, BRANCH_OVERLOADED_AT_DISPATCH, BRANCH_PD_ADMISSION_SHED, + BRANCH_ALL_OVERLOADED_FALLBACK, BRANCH_ALL_OVERLOADED_SHED, + BRANCH_ALL_STALLED_FALLBACK, BRANCH_OVERLOADED_AT_DISPATCH, BRANCH_PD_ADMISSION_SHED, STAGE_DISPATCH, STAGE_PD_ADMISSION, STAGE_SELECTION, }, Worker, @@ -56,12 +63,19 @@ pub fn all_overloaded(candidates: &[Arc]) -> bool { !candidates.is_empty() && candidates.iter().all(|w| w.is_overloaded()) } -/// Shed when every worker selection could have used is flagged overloaded. +/// Shed when every worker selection could have used is flagged overloaded +/// and shedding is on (`--worker-overload-shed`, `shedding`). /// -/// `None` means the empty pool has some other cause, and the caller's existing -/// not-found / unavailable answer stands. -pub fn shed_if_all_overloaded(candidates: &[Arc], model_id: &str) -> Option { - if !all_overloaded(candidates) { +/// `None` means the pool is not all-overloaded, or shedding is off and the +/// caller steers through [`fallback_if_all_vetoed`] instead; either way +/// the caller's existing not-found / unavailable answer stands when nothing +/// is routable. +pub fn shed_if_all_overloaded( + candidates: &[Arc], + model_id: &str, + shedding: bool, +) -> Option { + if !shedding || !all_overloaded(candidates) { return None; } Some(shed( @@ -72,12 +86,92 @@ pub fn shed_if_all_overloaded(candidates: &[Arc], model_id: &str) -> )) } -/// Dispatch-time re-check: one atomic read on the already-chosen worker, -/// covering the selection→dispatch window. Deliberately sheds rather than -/// re-selecting — the flag moves at the poll interval, so the window is rare — -/// and reports only what it knows: this worker went over, not the fleet. -pub fn shed_if_worker_overloaded(worker: &dyn Worker, model_id: &str) -> Option { - if !worker.is_overloaded() { +/// Whether every worker in a non-empty candidate pool is vetoed, by the +/// overload flag or by the liveness tracker (`Worker::stall_reason`). +/// +/// `candidates` is the pool *before* the `is_available()` filter, narrowed by +/// exactly the worker-type / connection-mode filter selection used. +pub fn all_vetoed(candidates: &[Arc]) -> bool { + !candidates.is_empty() + && candidates + .iter() + .all(|w| w.is_overloaded() || w.stall_reason().is_some()) +} + +/// The steering answer to a pool whose every worker is vetoed: the +/// least-loaded worker that is routable but for its veto (ready, circuit +/// closed). A fleet that is uniformly over the thresholds, or whose last +/// worker the liveness tracker doubts, is still a fleet; refusing the request +/// would turn a load signal or a suspicion into an outage, and a worker that +/// really is gone fails the request as fast as a refusal would. +/// +/// Workers vetoed by overload alone are taken first, so a liveness veto still +/// steers while any other worker can take the request. Under +/// `--worker-overload-shed` (`shedding`) an overloaded worker is never taken: +/// the shed is that flag's answer to overload, and the callers shed an +/// all-overloaded pool before asking here. `None` when the pool is not +/// all-vetoed, or when no worker in it is routable at all (then the caller's +/// unavailable answer stands). +pub fn fallback_if_all_vetoed( + candidates: &[Arc], + model_id: &str, + stage: &'static str, + shedding: bool, +) -> Option> { + if !all_vetoed(candidates) { + return None; + } + let mut rng = rand::rng(); + let ready = |w: &Arc| w.is_healthy() && w.circuit_breaker_can_execute(); + let overloaded_first = (!shedding) + .then(|| { + least_loaded_uniform( + candidates + .iter() + .filter(|w| ready(w) && w.stall_reason().is_none()), + &mut rng, + ) + }) + .flatten(); + let (worker, branch) = match overloaded_first { + Some(worker) => (worker, BRANCH_ALL_OVERLOADED_FALLBACK), + None => ( + least_loaded_uniform( + candidates + .iter() + .filter(|w| ready(w) && !(shedding && w.is_overloaded())), + &mut rng, + )?, + BRANCH_ALL_STALLED_FALLBACK, + ), + }; + if branch == BRANCH_ALL_OVERLOADED_FALLBACK { + Metrics::record_worker_overload_fallback(stage); + } else { + Metrics::record_worker_liveness_fallback(stage); + } + debug!( + branch, + stage, + worker = worker.url(), + model_id, + "Veto fallback" + ); + Some(Arc::clone(worker)) +} + +/// Dispatch-time re-check under `--worker-overload-shed` (`shedding`): one atomic +/// read on the already-chosen worker, covering the selection→dispatch window. +/// Deliberately sheds rather than re-selecting — the flag moves at the poll +/// interval, so the window is rare — and reports only what it knows: this +/// worker went over, not the fleet. Without shedding the dispatch stands: the +/// worker was the right choice when it was made. +pub fn shed_if_worker_overloaded( + worker: &dyn Worker, + model_id: &str, + shedding: bool, +) -> Option { + if !shedding || !worker.is_overloaded() { return None; } let url = worker.url(); @@ -134,6 +228,35 @@ fn shed(branch: &'static str, stage: &'static str, worker: &str, message: String response } +/// The least-loaded of `candidates`; one of them drawn uniformly when +/// several share the lowest load, so a fleet whose workers are all equally +/// loaded does not send every fallback to the first of them. +pub(crate) fn least_loaded_uniform<'a, R: RngExt>( + candidates: impl Iterator>, + rng: &mut R, +) -> Option<&'a Arc> { + let mut best: Option<(&'a Arc, usize)> = None; + let mut tied = 0u32; + for worker in candidates { + let load = worker.load(); + match best { + Some((_, best_load)) if load > best_load => {} + Some((_, best_load)) if load == best_load => { + // The k-th tied worker replaces the pick with probability 1/k. + tied += 1; + if rng.random_range(0..=tied) == 0 { + best = Some((worker, load)); + } + } + _ => { + best = Some((worker, load)); + tied = 0; + } + } + } + best.map(|(worker, _)| worker) +} + #[cfg(test)] mod tests { use axum::http::StatusCode; @@ -172,12 +295,16 @@ mod tests { a.set_overloaded(true); assert!( - shed_if_all_overloaded(&pool, "m").is_none(), + shed_if_all_overloaded(&pool, "m", true).is_none(), "one eligible worker left is not a shed" ); b.set_overloaded(true); - let response = shed_if_all_overloaded(&pool, "m").expect("shed"); + assert!( + shed_if_all_overloaded(&pool, "m", false).is_none(), + "without shedding an all-overloaded pool is steered, not refused" + ); + let response = shed_if_all_overloaded(&pool, "m", true).expect("shed"); assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE); assert_eq!( extract_error_code_from_response(&response), @@ -199,7 +326,7 @@ mod tests { let a = worker("http://127.0.0.1:9811", "m"); a.set_overloaded(true); - let selection = shed_if_all_overloaded(std::slice::from_ref(&a), "m").expect("shed"); + let selection = shed_if_all_overloaded(std::slice::from_ref(&a), "m", true).expect("shed"); assert!( is_retryable_status(selection.status()), "the status stays the retryable 503 clients already understand" @@ -209,7 +336,7 @@ mod tests { "but the retry layer must decline it" ); - let dispatch = shed_if_worker_overloaded(a.as_ref(), "m").expect("shed"); + let dispatch = shed_if_worker_overloaded(a.as_ref(), "m", true).expect("shed"); assert!(!is_retryable_response(&dispatch)); } @@ -220,8 +347,8 @@ mod tests { let a = worker("http://127.0.0.1:9821", "m"); a.set_overloaded(true); - let selection = shed_if_all_overloaded(std::slice::from_ref(&a), "m").expect("shed"); - let dispatch = shed_if_worker_overloaded(a.as_ref(), "m").expect("shed"); + let selection = shed_if_all_overloaded(std::slice::from_ref(&a), "m", true).expect("shed"); + let dispatch = shed_if_worker_overloaded(a.as_ref(), "m", true).expect("shed"); for response in [&selection, &dispatch] { let value = response .headers() @@ -237,20 +364,166 @@ mod tests { } } - /// An empty pool is a 404/unavailable question for the caller, not a shed. + /// An empty pool is a 404/unavailable question for the caller, not a shed + /// and not a fallback. #[test] fn empty_pool_is_not_a_shed() { - assert!(shed_if_all_overloaded(&[], "nobody").is_none()); + assert!(shed_if_all_overloaded(&[], "nobody", true).is_none()); + assert!(fallback_if_all_vetoed(&[], "nobody", STAGE_SELECTION, false).is_none()); assert!(!all_overloaded(&[])); } + /// The steering default: an all-overloaded pool routes to its least-loaded + /// routable worker; a pool with an eligible worker left is not touched, + /// and a worker that is also unhealthy is never the fallback. + #[test] + fn equally_loaded_workers_share_the_fallback_and_a_lighter_one_wins() { + use rand::{rngs::StdRng, SeedableRng}; + let pool: Vec> = (0..128) + .map(|i| worker(&format!("http://127.0.0.1:{}", 20000 + i), "m")) + .collect(); + let mut rng = StdRng::seed_from_u64(7); + let mut hits = vec![0usize; pool.len()]; + for _ in 0..1000 { + let picked = least_loaded_uniform(pool.iter(), &mut rng).expect("a worker"); + hits[pool.iter().position(|w| Arc::ptr_eq(w, picked)).unwrap()] += 1; + } + assert!(hits.iter().all(|&h| h >= 1), "{hits:?}"); + assert!(*hits.iter().max().unwrap() <= 24, "{hits:?}"); + for (i, w) in pool.iter().enumerate() { + if i != 5 { + w.increment_load(); + } + } + for _ in 0..100 { + let picked = least_loaded_uniform(pool.iter(), &mut rng).expect("a worker"); + assert!(Arc::ptr_eq(picked, &pool[5])); + } + } + + #[test] + fn all_overloaded_falls_back_to_the_least_loaded_routable_worker() { + let a = worker("http://127.0.0.1:9831", "m"); + let b = worker("http://127.0.0.1:9832", "m"); + let c = worker("http://127.0.0.1:9833", "m"); + let pool = vec![Arc::clone(&a), Arc::clone(&b), Arc::clone(&c)]; + for _ in 0..3 { + a.increment_load(); + } + b.increment_load(); + + a.set_overloaded(true); + b.set_overloaded(true); + assert!( + fallback_if_all_vetoed(&pool, "m", STAGE_SELECTION, false).is_none(), + "an eligible worker left means selection handles it" + ); + + c.set_overloaded(true); + let picked = fallback_if_all_vetoed(&pool, "m", STAGE_SELECTION, false).expect("fallback"); + assert_eq!( + picked.url(), + "http://127.0.0.1:9833", + "the least-loaded wins" + ); + + c.set_status(openai_protocol::worker::WorkerStatus::NotReady); + let picked = fallback_if_all_vetoed(&pool, "m", STAGE_SELECTION, false).expect("fallback"); + assert_eq!( + picked.url(), + "http://127.0.0.1:9832", + "an unhealthy worker is not routable even as the fallback" + ); + + a.set_status(openai_protocol::worker::WorkerStatus::NotReady); + b.set_status(openai_protocol::worker::WorkerStatus::NotReady); + assert!( + fallback_if_all_vetoed(&pool, "m", STAGE_SELECTION, false).is_none(), + "nothing routable: the caller's unavailable answer stands" + ); + } + + /// A liveness veto steers but never refuses: when every ready worker is + /// vetoed by the tracker the request goes to the least-loaded of them, + /// after any worker vetoed by overload alone, and under shedding never + /// to an overloaded one. + #[test] + fn all_stalled_falls_back_to_the_least_loaded_ready_worker() { + use crate::worker::worker::StallReason; + + let a = worker("http://127.0.0.1:9841", "m"); + let b = worker("http://127.0.0.1:9842", "m"); + let pool = vec![Arc::clone(&a), Arc::clone(&b)]; + for _ in 0..2 { + a.increment_load(); + } + + assert!(a.set_stall(Some(StallReason::Wedged))); + assert!( + fallback_if_all_vetoed(&pool, "m", STAGE_SELECTION, false).is_none(), + "a worker free of both vetoes means selection handles it" + ); + + assert!(b.set_stall(Some(StallReason::Unreachable))); + let picked = fallback_if_all_vetoed(&pool, "m", STAGE_SELECTION, false).expect("fallback"); + assert_eq!( + picked.url(), + "http://127.0.0.1:9842", + "the least-loaded of the vetoed is taken" + ); + let picked = + fallback_if_all_vetoed(&pool, "m", STAGE_SELECTION, true).expect("under shedding"); + assert_eq!( + picked.url(), + "http://127.0.0.1:9842", + "a liveness veto steers under shedding too" + ); + + assert!(b.set_stall(None)); + b.set_overloaded(true); + let picked = fallback_if_all_vetoed(&pool, "m", STAGE_SELECTION, false).expect("fallback"); + assert_eq!( + picked.url(), + "http://127.0.0.1:9842", + "overload alone outranks a liveness veto" + ); + let picked = + fallback_if_all_vetoed(&pool, "m", STAGE_SELECTION, true).expect("under shedding"); + assert_eq!( + picked.url(), + "http://127.0.0.1:9841", + "under shedding the overloaded worker is never taken" + ); + + a.set_overloaded(true); + assert!( + fallback_if_all_vetoed(&pool, "m", STAGE_SELECTION, true).is_none(), + "all overloaded under shedding: the caller sheds" + ); + + b.set_status(openai_protocol::worker::WorkerStatus::NotReady); + a.set_overloaded(false); + let picked = fallback_if_all_vetoed(&pool, "m", STAGE_SELECTION, false).expect("fallback"); + assert_eq!( + picked.url(), + "http://127.0.0.1:9841", + "an unhealthy worker is never the fallback; the stalled ready one is" + ); + + a.set_status(openai_protocol::worker::WorkerStatus::NotReady); + assert!( + fallback_if_all_vetoed(&pool, "m", STAGE_SELECTION, false).is_none(), + "nothing ready: the caller's unavailable answer stands" + ); + } + #[test] fn dispatch_recheck_sheds_only_for_a_flagged_worker() { let w = worker("http://127.0.0.1:9803", "m"); - assert!(shed_if_worker_overloaded(w.as_ref(), "m").is_none()); + assert!(shed_if_worker_overloaded(w.as_ref(), "m", true).is_none()); w.set_overloaded(true); - let response = shed_if_worker_overloaded(w.as_ref(), "m").expect("shed"); + let response = shed_if_worker_overloaded(w.as_ref(), "m", true).expect("shed"); assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE); assert_eq!( extract_error_code_from_response(&response), diff --git a/model_gateway/src/routers/common/placement.rs b/model_gateway/src/routers/common/placement.rs index 00e73c6afd..b752e00a4f 100644 --- a/model_gateway/src/routers/common/placement.rs +++ b/model_gateway/src/routers/common/placement.rs @@ -15,6 +15,7 @@ use std::sync::Arc; use axum::{http::HeaderMap, response::Response}; +use rand::RngExt; use tracing::{debug, warn}; use crate::{ @@ -25,6 +26,7 @@ use crate::{ }, routers::common::{header_utils, overload}, worker::{ + overload::{BRANCH_ALL_OVERLOADED_FALLBACK, BRANCH_ALL_STALLED_FALLBACK, STAGE_SELECTION}, ConnectionMode, ConnectionModeExt, PdPairIndex, PrefillCandidateError, PrefillSelectionContext, RoutingPool, RuntimeType, Worker, WorkerRegistry, }, @@ -229,7 +231,7 @@ pub(crate) fn select_from( &filtered }; if available.is_empty() { - return None; + return overload_fallback(registry, policy.name(), model_id, candidates); } // Cached hash ring for consistent hashing (O(log n) lookup). @@ -237,7 +239,7 @@ pub(crate) fn select_from( // The registry applies the routing-key sticky override when enabled and // otherwise delegates to the configured policy. - let idx = policies.select_worker_for_model( + let Some(idx) = policies.select_worker_for_model( &policy, model_id, available, @@ -251,7 +253,11 @@ pub(crate) fn select_from( hash_ring, leg: WorkerLeg::Single, }, - )?; + ) else { + // A policy that filters availability itself misses on an + // all-overloaded pool the same way the pre-filter empties. + return overload_fallback(registry, policy.name(), model_id, candidates); + }; let selected = available[idx].clone(); Metrics::record_worker_selection( @@ -264,6 +270,33 @@ pub(crate) fn select_from( Some(selected) } +/// The steering default for a pool whose every worker is vetoed, by the +/// overload thresholds or by the liveness tracker: its least-loaded routable +/// worker, recorded as a selection under `policy`. Under +/// `--worker-overload-shed` an overloaded worker is never taken (the caller +/// sheds an all-overloaded pool); `None` when the pool is not all-vetoed or +/// has nothing routable. +fn overload_fallback( + registry: &WorkerRegistry, + policy: &'static str, + model_id: &str, + candidates: &[Arc], +) -> Option> { + let selected = overload::fallback_if_all_vetoed( + candidates, + model_id, + STAGE_SELECTION, + registry.overload_shed_enabled(), + )?; + Metrics::record_worker_selection( + metrics_labels::WORKER_REGULAR, + selected.connection_mode().as_metric_label(), + model_id, + policy, + ); + Some(selected) +} + /// Classify a failed single-worker placement from the same pool it drew from. pub(crate) fn single_failure( registry: &WorkerRegistry, @@ -272,20 +305,133 @@ pub(crate) fn single_failure( wire: Option, ) -> PlacementFailure { let candidates = candidates(registry, model_id, pool, wire); - failure_from(candidates.as_slice(), model_id) + failure_from( + candidates.as_slice(), + model_id, + registry.overload_shed_enabled(), + ) } -/// Classify a failed placement from the candidates it drew from. -pub(crate) fn failure_from(candidates: &[Arc], model_id: &str) -> PlacementFailure { +/// Classify a failed placement from the candidates it drew from. `shed` is +/// `--worker-overload-shed`: without it an all-overloaded pool was already +/// steered by the placement, so reaching here means nothing was routable. +pub(crate) fn failure_from( + candidates: &[Arc], + model_id: &str, + shedding: bool, +) -> PlacementFailure { if candidates.is_empty() { return PlacementFailure::NoCandidates; } - if let Some(shed) = overload::shed_if_all_overloaded(candidates, model_id) { + if let Some(shed) = overload::shed_if_all_overloaded(candidates, model_id, shedding) { return PlacementFailure::AllOverloaded(shed); } PlacementFailure::Unavailable } +/// The steering default for a disaggregated placement whose failed leg is +/// vetoed on every worker, by the overload thresholds or by the liveness +/// tracker: the least-loaded routable prefill that has a routable partner, +/// paired with its least-loaded routable partner. Pairs vetoed by overload +/// alone come first; a liveness veto is crossed only when nothing else pairs, +/// and never onto an overloaded worker under `--worker-overload-shed`. +/// `None` when the failed leg is not all-vetoed (an unhealthy leg keeps its +/// unavailable answer), when it is all-overloaded under shedding (the caller +/// sheds), or when nothing pairs; a prefill admission gate still has the last +/// word. +fn relaxed_pair( + registry: &WorkerRegistry, + model_id: &str, + pairs: &PdPairIndex, + wire: Option, + prefill_capacity: Option<&PrefillSelectionContext<'_>>, + inputs: &PlacementInputs<'_>, + failed_leg: &[Arc], +) -> Option>> { + let shedding = registry.overload_shed_enabled(); + if !overload::all_vetoed(failed_leg) || (shedding && overload::all_overloaded(failed_leg)) { + return None; + } + let eligible = |w: &Arc| { + w.is_healthy() + && w.circuit_breaker_can_execute() + && wire.is_none_or(|wire| { + w.metadata().spec.runtime_type == wire.runtime + && *w.connection_mode() == wire.connection + }) + && inputs + .candidate_filter + .is_none_or(|accepts| accepts(w.as_ref())) + }; + // Equal loads draw uniformly (partner and pair alike), so an equally + // loaded leg does not send every fallback to its first pair. + let mut rng = rand::rng(); + let mut search = |routable: &dyn Fn(&Arc) -> bool| { + let mut best: Option<(Arc, Arc)> = None; + let mut tied = 0u32; + for (i, prefill) in pairs.prefill.iter().enumerate() { + if !routable(prefill) { + continue; + } + let Some(decode) = overload::least_loaded_uniform( + pairs.partners[i].iter().filter(|d| routable(d)), + &mut rng, + ) else { + continue; + }; + let loads = (prefill.load(), decode.load()); + match best.as_ref().map(|(p, d)| (p.load(), d.load())) { + Some(best_loads) if loads > best_loads => {} + Some(best_loads) if loads == best_loads => { + tied += 1; + if rng.random_range(0..=tied) == 0 { + best = Some((Arc::clone(prefill), Arc::clone(decode))); + } + } + _ => { + best = Some((Arc::clone(prefill), Arc::clone(decode))); + tied = 0; + } + } + } + best + }; + let overloaded_first = (!shedding) + .then(|| search(&|w| eligible(w) && w.stall_reason().is_none())) + .flatten(); + let ((prefill, decode), branch) = match overloaded_first { + Some(pair) => (pair, BRANCH_ALL_OVERLOADED_FALLBACK), + None => ( + search(&|w| eligible(w) && !(shedding && w.is_overloaded()))?, + BRANCH_ALL_STALLED_FALLBACK, + ), + }; + if prefill_capacity.is_some_and(|capacity| !capacity.has_capacity(&prefill)) { + return Some(Err(Box::new(PairFailure { + leg: WorkerLeg::Prefill, + verdict: PlacementFailure::PrefillAtCapacity, + }))); + } + if branch == BRANCH_ALL_OVERLOADED_FALLBACK { + Metrics::record_worker_overload_fallback(STAGE_SELECTION); + } else { + Metrics::record_worker_liveness_fallback(STAGE_SELECTION); + } + debug!( + branch, + prefill = %prefill.url(), + decode = %decode.url(), + model_id, + "Veto fallback" + ); + let runtime = prefill.metadata().spec.runtime_type; + Some(Ok(Pair { + prefill, + decode, + runtime, + })) +} + /// Pick a prefill/decode pair for `model_id`, one worker per leg, each under /// its own policy and sticky namespace. /// @@ -367,10 +513,21 @@ pub(crate) fn select_pair( pairs.partners[i].iter().any(|d| partner_open(d, runtime)) }; if !pairs.prefill.iter().any(&eligible) { + if let Some(pair) = relaxed_pair( + registry, + model_id, + pairs, + wire, + prefill_capacity, + &inputs, + &pairs.prefill, + ) { + return pair; + } debug!("No available prefill workers"); return Err(fail( WorkerLeg::Prefill, - failure_from(&pairs.prefill, model_id), + failure_from(&pairs.prefill, model_id, registry.overload_shed_enabled()), )); } let leg_runtime = homogeneous_runtime @@ -409,10 +566,25 @@ pub(crate) fn select_pair( ); } if open.is_empty() { + if let Some(pair) = relaxed_pair( + registry, + model_id, + pairs, + wire, + prefill_capacity, + &inputs, + &pairs.decode_pool, + ) { + return pair; + } debug!(?leg_runtime, "No available PD pair"); return Err(fail( WorkerLeg::Decode, - failure_from(&pairs.decode_pool, model_id), + failure_from( + &pairs.decode_pool, + model_id, + registry.overload_shed_enabled(), + ), )); } @@ -1059,7 +1231,7 @@ mod tests { .expect("the narrowed slice still has a worker"); assert_eq!(selected.url(), "http://h:2"); assert!(matches!( - failure_from(&[], MODEL), + failure_from(&[], MODEL, false), PlacementFailure::NoCandidates )); } diff --git a/model_gateway/src/routers/common/worker_selection.rs b/model_gateway/src/routers/common/worker_selection.rs index ceb82765db..dceb10d377 100644 --- a/model_gateway/src/routers/common/worker_selection.rs +++ b/model_gateway/src/routers/common/worker_selection.rs @@ -17,7 +17,7 @@ use crate::{ }, error, }, - worker::{ProviderType, RuntimeType, Worker, WorkerRegistry}, + worker::{overload::STAGE_SELECTION, ProviderType, RuntimeType, Worker, WorkerRegistry}, }; /// Holds references to shared infrastructure needed for worker selection. @@ -69,8 +69,8 @@ impl<'a> WorkerSelector<'a> { // each of these requests into three registry walks and up to 5 s of // network wait to reach a 503 that carries neither the shed error code // nor the shed counter. - if let Some(shed) = self.shed_if_all_overloaded(req) { - return Err(shed); + if let Some(verdict) = self.all_overloaded_verdict(req) { + return verdict; } tracing::debug!( @@ -130,11 +130,23 @@ impl<'a> WorkerSelector<'a> { .min_by_key(|w| w.load()) } - /// Shed when every worker this request could have selected is vetoed. - /// Runs only on the miss path. - fn shed_if_all_overloaded(&self, req: &SelectWorkerRequest<'_>) -> Option { + /// The verdict on a pool whose every worker is vetoed, judged on the miss + /// path only: a shed when the pool is all-overloaded under + /// `--worker-overload-shed`, else the least-loaded routable worker (an + /// overloaded one never under shedding; a liveness veto steers either + /// way). `None` when the pool is not all-vetoed (or nothing in it is + /// routable), so the usual miss handling continues. + fn all_overloaded_verdict( + &self, + req: &SelectWorkerRequest<'_>, + ) -> Option, Response>> { let candidates = self.candidate_pool(req, false); - overload::shed_if_all_overloaded(&candidates, req.model_id) + let shedding = self.registry.overload_shed_enabled(); + if let Some(shed) = overload::shed_if_all_overloaded(&candidates, req.model_id, shedding) { + return Some(Err(shed)); + } + overload::fallback_if_all_vetoed(&candidates, req.model_id, STAGE_SELECTION, shedding) + .map(Ok) } /// Check if any healthy worker supports the model (regardless of circuit breaker). @@ -427,4 +439,43 @@ mod tests { fn default_request_does_not_require_realtime() { assert!(!SelectWorkerRequest::default().require_realtime_capable); } + + /// Every candidate over the thresholds: by default the least-loaded one + /// serves; under `--worker-overload-shed` the request is refused. + #[tokio::test] + async fn all_overloaded_steers_to_the_least_loaded_unless_shedding() { + let registry = WorkerRegistry::new(); + let busy = worker("http://127.0.0.1:18180", false); + let quiet = worker("http://127.0.0.1:18181", false); + for _ in 0..4 { + busy.increment_load(); + } + quiet.increment_load(); + registry.register_or_replace(Arc::clone(&busy)); + registry.register_or_replace(Arc::clone(&quiet)); + registry.set_worker_overloaded(&busy, true); + registry.set_worker_overloaded(&quiet, true); + + let picked = WorkerSelector::new(®istry) + .select_worker(&SelectWorkerRequest { + model_id: "m", + ..Default::default() + }) + .await + .expect("steered to the least-loaded worker"); + assert_eq!(picked.url(), "http://127.0.0.1:18181"); + + registry.set_overload_shed(true); + let refused = WorkerSelector::new(®istry) + .select_worker(&SelectWorkerRequest { + model_id: "m", + ..Default::default() + }) + .await + .expect_err("shedding refuses the all-overloaded pool"); + assert_eq!( + error::extract_error_code_from_response(&refused), + "worker_overload_protection_shed" + ); + } } diff --git a/model_gateway/src/routers/grpc/common/stages/client_acquisition.rs b/model_gateway/src/routers/grpc/common/stages/client_acquisition.rs index 8fe3619598..1e5b449e6b 100644 --- a/model_gateway/src/routers/grpc/common/stages/client_acquisition.rs +++ b/model_gateway/src/routers/grpc/common/stages/client_acquisition.rs @@ -27,10 +27,12 @@ use crate::{ pub(crate) async fn acquire_clients( workers: &WorkerSelection, model_id: &str, + shed: bool, ) -> Result { match workers { WorkerSelection::Single { worker } => { - if let Some(shed) = overload::shed_if_worker_overloaded(worker.as_ref(), model_id) { + if let Some(shed) = overload::shed_if_worker_overloaded(worker.as_ref(), model_id, shed) + { return Err(shed); } let client = get_backend_client_from_worker(worker).await?; @@ -46,13 +48,20 @@ pub(crate) async fn acquire_clients( // vetoed at selection through the same filter, so leaving it out // of the re-check would be the one dispatch path that can send // to a worker known to be over the ceiling. - if let Some(shed) = overload::shed_if_worker_overloaded(prefill.as_ref(), model_id) - .or_else(|| overload::shed_if_worker_overloaded(decode.as_ref(), model_id)) - .or_else(|| { - encode_assignments.iter().flatten().find_map(|assignment| { - overload::shed_if_worker_overloaded(assignment.worker.as_ref(), model_id) + if let Some(shed) = + overload::shed_if_worker_overloaded(prefill.as_ref(), model_id, shed) + .or_else(|| { + overload::shed_if_worker_overloaded(decode.as_ref(), model_id, shed) + }) + .or_else(|| { + encode_assignments.iter().flatten().find_map(|assignment| { + overload::shed_if_worker_overloaded( + assignment.worker.as_ref(), + model_id, + shed, + ) + }) }) - }) { return Err(shed); } diff --git a/model_gateway/src/routers/grpc/common/stages/request_execution.rs b/model_gateway/src/routers/grpc/common/stages/request_execution.rs index ce476c0d4f..384d632d27 100644 --- a/model_gateway/src/routers/grpc/common/stages/request_execution.rs +++ b/model_gateway/src/routers/grpc/common/stages/request_execution.rs @@ -633,6 +633,8 @@ async fn execute_single( proto_request.set_data_parallel_rank(rank as i32); } + let prompt_tokens = u64::try_from(proto_request.prompt_len()).unwrap_or(u64::MAX); + let streaming = proto_request.stream(); let result = client.generate(proto_request).await; workers.record_outcome(result.cb_status_code()); @@ -644,6 +646,17 @@ async fn execute_single( "start_generation_failed", ) })?; + // Every generation to one worker is tracked: each response on it is + // progress for the worker (liveness), the one answer of a non-streaming + // generation included, and its prompt is pending prefill there until the + // first response. Only a streaming generation joins the worker's pile: a + // non-streaming one shows nothing between dispatch and completion, and + // the pile rule never judges a worker by requests it cannot see progress + // on. + let stream = match workers.single() { + Some(worker) => stream.tracked(Arc::clone(worker), prompt_tokens, streaming), + None => stream, + }; Ok(ExecutionResult::Single { stream }) } @@ -1131,10 +1144,10 @@ async fn execute_sequential_pd( // on the decode read). // // A decode worker started with `--language-model-only` (the production - // vLLM P/D shape — Dynamo pairs the same way) has no vision encoder and + // vLLM P/D shape) has no vision encoder and // an encoder-cache budget of 0, so even the identity payload fails to // schedule there. Its model info reports supports_vision=false; for such - // a worker the decode leg is stripped down to the Dynamo contract: the + // a worker the decode leg is stripped down to what such an engine takes: the // prefill-expanded input_ids, the KV handoff, and the per-image content // hashes that the servicer folds into cache_salt so different images // cannot alias in the decode prefix cache. diff --git a/model_gateway/src/routers/grpc/common/stages/worker_selection.rs b/model_gateway/src/routers/grpc/common/stages/worker_selection.rs index f0f9fbbf66..993a4cf31c 100644 --- a/model_gateway/src/routers/grpc/common/stages/worker_selection.rs +++ b/model_gateway/src/routers/grpc/common/stages/worker_selection.rs @@ -544,7 +544,11 @@ impl WorkerSelectionStage { if capable.is_empty() { return self.media_refs_shed(model_id); } - match placement::failure_from(&capable, model_id) { + match placement::failure_from( + &capable, + model_id, + self.worker_registry.overload_shed_enabled(), + ) { PlacementFailure::AllOverloaded(shed) => return shed, PlacementFailure::Unavailable | PlacementFailure::PolicyDeclined(_) @@ -668,7 +672,11 @@ impl WorkerSelectionStage { .filter(|w| wire.is_none_or(|c| w.metadata().spec.runtime_type == c.runtime)) .cloned() .collect(); - placement::failure_from(&candidates, model_id) + placement::failure_from( + &candidates, + model_id, + self.worker_registry.overload_shed_enabled(), + ) } #[expect( @@ -1358,6 +1366,7 @@ mod tests { fn a_fully_vetoed_prefill_leg_sheds_rather_than_404s() { let model_id = "test-model-prefill-veto"; let worker_registry = Arc::new(WorkerRegistry::new()); + worker_registry.set_overload_shed(true); let (prefill_urls, _) = register_pd_workers(&worker_registry, model_id, 4); let stage = WorkerSelectionStage::new( @@ -1658,6 +1667,7 @@ mod tests { let model_id = "test-model-overload-shed"; let worker_registry = Arc::new(WorkerRegistry::new()); + worker_registry.set_overload_shed(true); let mut workers = Vec::new(); for i in 0..2 { let worker: Arc = Arc::new( @@ -2012,6 +2022,9 @@ mod tests { } } worker_registry.set_worker_overloaded(&pinned.expect("vllm worker registered"), true); + // Under shedding; the steering default would route the retry to the + // vetoed worker instead, which is the other test's subject. + worker_registry.set_overload_shed(true); let stage = WorkerSelectionStage::new( worker_registry, @@ -2081,6 +2094,8 @@ mod tests { cache_index: Default::default(), cache_ttl_secs: 180, cache_boundaries: Vec::new(), + selection_policy: None, + selection_accounting_ttl_ms: 0, })), WorkerSelectionMode::Regular, None, @@ -2262,12 +2277,75 @@ mod tests { assert_eq!(response.status(), StatusCode::NOT_FOUND); } + /// The steering default: with every worker over the thresholds the gRPC + /// path routes to the least-loaded one instead of refusing, on the single + /// path and on both PD legs. + #[test] + fn grpc_all_overloaded_steers_to_the_least_loaded_by_default() { + let model_id = "test-model-overload-steer"; + let worker_registry = Arc::new(WorkerRegistry::new()); + let mut workers = Vec::new(); + for i in 0..2 { + let worker: Arc = Arc::new( + BasicWorkerBuilder::new(format!("grpc://127.0.0.1:{}", 8450 + i)) + .model(ModelCard::new(model_id)) + .worker_type(WorkerType::Regular) + .connection_mode(ConnectionMode::Grpc) + .health_config(no_health_check()) + .build(), + ); + worker_registry.register(Arc::clone(&worker)).unwrap(); + workers.push(worker); + } + workers[0].increment_load(); + workers[0].increment_load(); + workers[1].increment_load(); + let stage = WorkerSelectionStage::new( + Arc::clone(&worker_registry), + Arc::new(PolicyRegistry::new(PolicyConfig::RoundRobin)), + WorkerSelectionMode::Regular, + None, + ); + for worker in &workers { + worker_registry.set_worker_overloaded(worker, true); + } + let selected = stage + .select_single_worker(model_id, None, None, None, None, None, None, false) + .expect("the all-overloaded pool is steered, not refused"); + assert_eq!( + selected.url(), + "grpc://127.0.0.1:8451", + "the least-loaded serves" + ); + + // PD: both legs vetoed still pair up. + let pd_model = "test-model-overload-steer-pd"; + let pd_registry = Arc::new(WorkerRegistry::new()); + let (prefill_urls, decode_urls) = register_pd_workers(&pd_registry, pd_model, 2); + let pd_stage = WorkerSelectionStage::new( + Arc::clone(&pd_registry), + Arc::new(PolicyRegistry::new(PolicyConfig::RoundRobin)), + WorkerSelectionMode::PrefillDecode, + None, + ); + for url in prefill_urls.iter().chain(decode_urls.iter()) { + let worker = pd_registry.get_by_url(url).expect("registered"); + pd_registry.set_worker_overloaded(&worker, true); + } + let (prefill, decode, _) = pd_stage + .select_pd_pair(pd_model, PlacementInputs::default(), None, None) + .expect("both legs over the thresholds still pair"); + assert!(prefill_urls.contains(&prefill.url().to_string())); + assert!(decode_urls.contains(&decode.url().to_string())); + } + /// Capable but overloaded workers keep the overload shed: its code, /// Retry-After and non-retryable marking survive worker mode. #[test] fn media_refs_overloaded_capable_workers_keep_the_overload_shed() { let model_id = "test-model-media-refs-overload"; let worker_registry = Arc::new(WorkerRegistry::new()); + worker_registry.set_overload_shed(true); let mut workers = Vec::new(); for i in 0..2 { let worker = vllm_grpc_worker( diff --git a/model_gateway/src/routers/grpc/context.rs b/model_gateway/src/routers/grpc/context.rs index 2d3cf228c4..3a89a3176b 100644 --- a/model_gateway/src/routers/grpc/context.rs +++ b/model_gateway/src/routers/grpc/context.rs @@ -1533,4 +1533,74 @@ mod tests { assert_eq!(plan.request_type(), "generate"); assert_eq!(plan.mode_label(), "prefill_decode"); } + + #[test] + fn load_guards_report_completion_for_each_leg_they_hold() { + use crate::worker::{BasicWorkerBuilder, RequestCompletionSink, WorkerType}; + #[derive(Debug, Default)] + struct CompletionSpy(std::sync::Mutex>); + impl RequestCompletionSink for CompletionSpy { + fn request_completed(&self, worker: &dyn Worker) { + self.0.lock().unwrap().push(worker.url().to_string()); + } + } + + let spy = Arc::new(CompletionSpy::default()); + let sink: Arc = spy.clone(); + let single: Arc = Arc::new( + BasicWorkerBuilder::new("grpc://single") + .worker_type(WorkerType::Regular) + .build(), + ); + let prefill: Arc = Arc::new( + BasicWorkerBuilder::new("grpc://prefill") + .worker_type(WorkerType::Prefill) + .build(), + ); + let decode: Arc = Arc::new( + BasicWorkerBuilder::new("grpc://decode") + .worker_type(WorkerType::Decode) + .build(), + ); + for worker in [&single, &prefill, &decode] { + worker.set_completion_sink(Some(Arc::clone(&sink))); + } + + // The regular gRPC path: one guard, one completion when the stream ends. + drop(LoadGuards::new( + &WorkerSelection::Single { + worker: Arc::clone(&single), + }, + None, + )); + assert_eq!(spy.0.lock().unwrap().as_slice(), ["grpc://single"]); + + // The disaggregated path: the decode leg completes with the guards, + // the prefill leg with its own guard when the prefill phase ends. + let guards = LoadGuards::new(&pd_selection(&prefill, &decode), None); + let prefill_guard = PrefillLoadGuard::Unbounded { + _guard: WorkerLoadGuard::new(Arc::clone(&prefill), None), + }; + drop(prefill_guard); + assert_eq!( + spy.0.lock().unwrap().last().map(String::as_str), + Some("grpc://prefill") + ); + drop(guards); + assert_eq!( + spy.0.lock().unwrap().last().map(String::as_str), + Some("grpc://decode") + ); + + // A batched fan-out reports once per sub-request. + let before = spy.0.lock().unwrap().len(); + drop(LoadGuards::scaled( + &WorkerSelection::Single { + worker: Arc::clone(&single), + }, + None, + 3, + )); + assert_eq!(spy.0.lock().unwrap().len(), before + 3); + } } diff --git a/model_gateway/src/routers/grpc/pipeline.rs b/model_gateway/src/routers/grpc/pipeline.rs index 6c39ad489f..9e53d218de 100644 --- a/model_gateway/src/routers/grpc/pipeline.rs +++ b/model_gateway/src/routers/grpc/pipeline.rs @@ -245,6 +245,8 @@ pub(crate) struct RequestPipeline { backend_type: &'static str, /// Disaggregation mode, for per-leg retry metric labels. mode: Mode, + /// The registry, read at dispatch for `--worker-overload-shed`. + worker_registry: Arc, } /// Outcome of one full pipeline run. @@ -408,6 +410,7 @@ impl RequestPipeline { stages: Arc::new(stages), backend_type: backend, mode, + worker_registry: deps.worker_registry.clone(), }) } @@ -460,7 +463,12 @@ impl RequestPipeline { )?; ctx.state.clients = Some(step!( "ClientAcquisition", - acquire_clients(workers, &ctx.input.model_id).await + acquire_clients( + workers, + &ctx.input.model_id, + self.worker_registry.overload_shed_enabled(), + ) + .await )?); if let Some(encode) = &stages.encode { step!(encode.name(), encode.execute(ctx).await)?; @@ -498,7 +506,14 @@ impl RequestPipeline { "Worker selection not completed", ) })?; - dctx.clients = Some(acquire_clients(workers, &dctx.model_id).await?); + dctx.clients = Some( + acquire_clients( + workers, + &dctx.model_id, + self.worker_registry.overload_shed_enabled(), + ) + .await?, + ); let retained = plan.as_mut().ok_or_else(|| { error!(function = "run_attempt", "Execution plan already consumed"); error::internal_error("execution_plan_consumed", "Execution plan already consumed") @@ -1526,7 +1541,10 @@ mod request_release_tests { routers::grpc::multimodal::{ MultimodalComponents, MultimodalConfigRegistry, MultimodalSettings, }, - worker::{BasicWorkerBuilder, ConnectionMode, RuntimeType, WorkerType}, + worker::{ + BasicWorkerBuilder, ConnectionMode, RequestCompletionSink, RuntimeType, Worker, + WorkerType, + }, }; const MODEL: &str = "request-release-test-model"; @@ -1801,20 +1819,27 @@ mod request_release_tests { panic!("release-test stub on port {port} never came up"); } - fn register_worker(registry: &WorkerRegistry, port: u16, worker_type: WorkerType) { - let worker = BasicWorkerBuilder::new(format!("grpc://127.0.0.1:{port}")) - .worker_type(worker_type) - .connection_mode(ConnectionMode::Grpc) - .runtime_type(RuntimeType::TokenSpeed) - .model(ModelCard::new(MODEL)) - .health_config(HealthCheckConfig { - disable_health_check: true, - ..Default::default() - }) - .build(); + fn register_worker( + registry: &WorkerRegistry, + port: u16, + worker_type: WorkerType, + ) -> Arc { + let worker: Arc = Arc::new( + BasicWorkerBuilder::new(format!("grpc://127.0.0.1:{port}")) + .worker_type(worker_type) + .connection_mode(ConnectionMode::Grpc) + .runtime_type(RuntimeType::TokenSpeed) + .model(ModelCard::new(MODEL)) + .health_config(HealthCheckConfig { + disable_health_check: true, + ..Default::default() + }) + .build(), + ); registry - .register(Arc::new(worker)) + .register(Arc::clone(&worker)) .expect("register release-test worker"); + worker } async fn components(worker_registry: Arc) -> Arc { @@ -1851,10 +1876,18 @@ mod request_release_tests { } fn completion_pipeline(worker_registry: &Arc, mode: Mode) -> RequestPipeline { + completion_pipeline_with_admission(worker_registry, mode, None) + } + + fn completion_pipeline_with_admission( + worker_registry: &Arc, + mode: Mode, + prefill_admission: Option>, + ) -> RequestPipeline { let deps = PipelineDeps::pair( worker_registry.clone(), Arc::new(PolicyRegistry::new(PolicyConfig::Random)), - None, + prefill_admission, None, ); RequestPipeline::build(Endpoint::Completion, mode, &deps).expect("completion pipeline") @@ -2200,6 +2233,203 @@ mod request_release_tests { ); } + // ------------------------------------------------------------------ + // Request completion on a client cancel. Policies learn that a request + // ended through the worker's completion sink, which the load guard + // drives; these pin that a client leaving mid-request reaches it on the + // gRPC paths, promptly, while the engine still owes its first token. + // ------------------------------------------------------------------ + + /// Records the worker URLs that reported a request completion. + #[derive(Debug, Default)] + struct CompletionSpy(Mutex>); + + impl RequestCompletionSink for CompletionSpy { + fn request_completed(&self, worker: &dyn Worker) { + self.0 + .lock() + .unwrap_or_else(PoisonError::into_inner) + .push(worker.url().to_string()); + } + } + + impl CompletionSpy { + fn urls(&self) -> Vec { + self.0 + .lock() + .unwrap_or_else(PoisonError::into_inner) + .clone() + } + + /// Wait for `expected` completions, well inside the stub's five-second + /// gate: a completion that only arrives with the engine's tokens is a + /// failure, not a late pass. + async fn wait_for(&self, expected: usize) -> Vec { + let deadline = Instant::now() + Duration::from_secs(2); + while self.urls().len() < expected { + assert!( + Instant::now() < deadline, + "only {:?} completed in time, expected {expected}", + self.urls() + ); + tokio::time::sleep(Duration::from_millis(10)).await; + } + self.urls() + } + } + + /// Install one spy as the completion sink of every registered worker. + fn install_spy(registry: &WorkerRegistry) -> Arc { + let spy = Arc::new(CompletionSpy::default()); + for worker in registry.get_all() { + worker.set_completion_sink(Some(Arc::clone(&spy) as Arc)); + } + spy + } + + /// A stub that accepts the generate RPC at once and withholds every token + /// while `hold` lives: the engine has the request and owes its first + /// token, the window a client cancel is hardest to account for. + fn stalled_while(hold: &Arc) -> GatedScheduler { + GatedScheduler { + probe: Some(Arc::downgrade(hold)), + ..Default::default() + } + } + + async fn start_stream( + pipeline: &RequestPipeline, + components: Arc, + ) -> Response { + let started = Instant::now(); + let response = pipeline + .execute_completion( + completion_request(true), + None, + MODEL.to_string(), + components, + None, + None, + None, + ) + .await; + assert_eq!(response.status(), http::StatusCode::OK); + assert!( + started.elapsed() < Duration::from_secs(4), + "the streaming response must open before the engine's first token" + ); + response + } + + /// gRPC regular path: the client drops the stream while the engine is + /// still prefilling. The load guard rides the response body, so the drop + /// releases the worker's load and reports the completion at once. + #[tokio::test] + async fn cancelled_grpc_stream_reports_the_completion_during_prefill() { + let hold = completion_request(true); + let port = spawn_stub(stalled_while(&hold)).await; + let worker_registry = Arc::new(WorkerRegistry::new()); + let worker = register_worker(&worker_registry, port, WorkerType::Regular); + let spy = install_spy(&worker_registry); + let pipeline = completion_pipeline(&worker_registry, Mode::Regular); + let components = components(Arc::clone(&worker_registry)).await; + + let response = start_stream(&pipeline, components).await; + assert_eq!(worker.load(), 1, "the dispatch holds the worker's load"); + assert!( + spy.urls().is_empty(), + "nothing completes while the client is connected" + ); + + drop(response); + + assert_eq!(spy.wait_for(1).await, [worker.url().to_string()]); + assert_eq!(worker.load(), 0); + drop(hold); + } + + /// gRPC PD path: the prefill leg answers on its own and reports when its + /// phase ends; the decode leg is still owed its first token when the + /// client leaves, and reports with the dropped body. Both legs end with + /// their load released, exactly once each. + #[tokio::test] + async fn cancelled_pd_stream_reports_the_prefill_and_decode_completions() { + let hold = completion_request(true); + let prefill_port = spawn_stub(GatedScheduler::default()).await; + let decode_port = spawn_stub(stalled_while(&hold)).await; + let worker_registry = Arc::new(WorkerRegistry::new()); + let prefill = register_worker(&worker_registry, prefill_port, WorkerType::Prefill); + let decode = register_worker(&worker_registry, decode_port, WorkerType::Decode); + let spy = install_spy(&worker_registry); + let pipeline = completion_pipeline(&worker_registry, Mode::PrefillDecode); + let components = components(Arc::clone(&worker_registry)).await; + + let response = start_stream(&pipeline, components).await; + assert_eq!( + decode.load(), + 1, + "the decode leg holds its load until the stream ends" + ); + + drop(response); + + let mut urls = spy.wait_for(2).await; + urls.sort(); + let mut expected = vec![prefill.url().to_string(), decode.url().to_string()]; + expected.sort(); + assert_eq!(urls, expected, "each leg reports exactly once"); + assert_eq!(prefill.load(), 0); + assert_eq!(decode.load(), 0); + drop(hold); + } + + /// gRPC PD path under prefill admission: the prefill leg's load is a + /// reservation in the admission gate rather than a bare guard, and it + /// still reports the completion when the phase ends, so a cancelled + /// request frees its slot for the next one. + #[tokio::test] + async fn cancelled_pd_stream_with_admitted_prefill_reports_and_frees_the_slot() { + let hold = completion_request(true); + let prefill_port = spawn_stub(GatedScheduler::default()).await; + let decode_port = spawn_stub(stalled_while(&hold)).await; + let worker_registry = Arc::new(WorkerRegistry::new()); + let prefill = register_worker(&worker_registry, prefill_port, WorkerType::Prefill); + let decode = register_worker(&worker_registry, decode_port, WorkerType::Decode); + let spy = install_spy(&worker_registry); + let admission = Arc::new(PrefillAdmission::new(1, 0, Duration::from_secs(1))); + let pipeline = completion_pipeline_with_admission( + &worker_registry, + Mode::PrefillDecode, + Some(Arc::clone(&admission)), + ); + let components = components(Arc::clone(&worker_registry)).await; + + let response = start_stream(&pipeline, components).await; + drop(response); + + let mut urls = spy.wait_for(2).await; + urls.sort(); + let mut expected = vec![prefill.url().to_string(), decode.url().to_string()]; + expected.sort(); + assert_eq!(urls, expected, "each leg reports exactly once"); + assert_eq!( + prefill.load(), + 0, + "the admitted prefill's reservation is released" + ); + assert_eq!(decode.load(), 0); + // The one slot is free again: a new admission does not queue. + let admitted = tokio::time::timeout( + Duration::from_millis(500), + admission.admit(None, |capacity| capacity.select(Arc::clone(&prefill), ())), + ) + .await + .expect("the freed slot must admit at once") + .expect("admitted"); + drop(admitted); + drop(hold); + } + // ------------------------------------------------------------------ // DeepSeek-V4.1 parity through the pipeline with the real checkpoint // tokenizer. Needs `DEEPSEEK_V41_MODEL_DIR` (or the tokenizer crate's diff --git a/model_gateway/src/routers/grpc/proto_wrapper.rs b/model_gateway/src/routers/grpc/proto_wrapper.rs index 20f7ad9723..c873d37af4 100644 --- a/model_gateway/src/routers/grpc/proto_wrapper.rs +++ b/model_gateway/src/routers/grpc/proto_wrapper.rs @@ -14,7 +14,7 @@ use std::{ process, sync::{ atomic::{AtomicU64, Ordering}, - OnceLock, + Arc, OnceLock, }, time::{Instant, SystemTime, UNIX_EPOCH}, }; @@ -42,9 +42,12 @@ use smg_grpc_client::{ }; use smg_mm_rdma::RdmaExporter; -use crate::routers::grpc::{ - multimodal::{log_mm_timing_enabled, mm_rdma_exporter}, - zmq_client::ZmqGenerateStream, +use crate::{ + routers::grpc::{ + multimodal::{log_mm_timing_enabled, mm_rdma_exporter}, + zmq_client::ZmqGenerateStream, + }, + worker::{liveness, Worker}, }; /// How a streaming response's per-token payloads (token ids, sampled @@ -1466,6 +1469,18 @@ impl ProtoGenerateRequest { } } + /// Whether the engine streams this request's responses as it produces + /// them, or answers once at the end. + pub fn stream(&self) -> bool { + match self { + Self::Vllm(req) => req.stream, + Self::Sglang(req) => req.stream, + Self::Trtllm(req) => req.streaming, + Self::Mlx(req) => req.stream, + Self::TokenSpeed(req) => req.stream, + } + } + /// Serialized wire size, for the release metric. pub fn wire_len(&self) -> usize { use prost::Message; @@ -1536,7 +1551,7 @@ impl ProtoGenerateRequest { /// every multimodal payload beyond the content hashes goes: pixels, /// placeholders, grid tensors and media references. The decode engine /// then sees a pure-text TokensPrompt and never touches its (zero-budget) - /// encoder cache, mirroring the Dynamo P/D contract; the kept hashes ride + /// encoder cache, as the language-model-only P/D contract requires; the kept hashes ride /// into `cache_salt` servicer-side so different images cannot alias in /// the decode prefix cache. Non-vLLM backends have no language-model-only /// mode, so they take the pixel-stripping clone. @@ -1647,6 +1662,34 @@ impl ProtoGenerateRequest { } } + /// Prompt tokens the engine has to prefill for this request: the + /// tokenized input's length (zero for a text input, which the engine + /// tokenizes itself). + pub fn prompt_len(&self) -> usize { + match self { + Self::Sglang(req) => req + .tokenized + .as_ref() + .map_or(0, |input| input.input_ids.len()), + Self::Vllm(req) => match &req.input { + Some(vllm::generate_request::Input::Tokenized(input)) => input.input_ids.len(), + _ => 0, + }, + Self::Trtllm(req) => req + .tokenized + .as_ref() + .map_or(0, |input| input.input_token_ids.len()), + Self::Mlx(req) => match &req.input { + Some(mlx::generate_request::Input::Tokenized(input)) => input.input_ids.len(), + _ => 0, + }, + Self::TokenSpeed(req) => req + .tokenized + .as_ref() + .map_or(0, |input| input.input_ids.len()), + } + } + /// Attach media references for worker-side multimodal processing (vLLM only). pub fn set_vllm_media_refs(&mut self, refs: vllm::MediaRefs) -> Result<(), String> { match self { @@ -2519,6 +2562,118 @@ pub enum ProtoStream { /// An n>1 fan-out over rendezvous-room PD pairs: one child per sample, /// each response stamped with its sample's index (see [`FanoutStream`]). Fanout(FanoutStream), + /// Any of the above, with the worker it came from: every response is a + /// sign of life for that worker (see [`crate::worker::liveness`]). + Tracked(Box), +} + +/// A [`ProtoStream`] paired with the worker serving it, so each response it +/// yields counts as token progress for that worker, and its prompt counts as +/// pending prefill on the worker until the first response. +pub struct TrackedStream { + inner: ProtoStream, + worker: Arc, + prefill: Option, + /// The request's place in the worker's pile; a non-streaming generation + /// holds none. + progress: Option, +} + +/// The request's place in its worker's tracked pile (see +/// [`Worker::tracked_load`]), held for the stream's life. +struct ProgressTicket { + worker: Arc, +} + +impl Drop for ProgressTicket { + fn drop(&mut self) { + self.worker.note_tracked_ended(); + } +} + +/// The prompt tokens a dispatched request adds to its worker's prefill +/// backlog (see [`Worker::prefill_backlog`]), released by the first response +/// or, failing that, when the stream is dropped. +struct PrefillTicket { + worker: Arc, + tokens: u64, + released: bool, +} + +impl PrefillTicket { + fn first_response(&mut self) { + if !self.released { + self.released = true; + self.worker.note_prefill_ended(self.tokens, true); + } + } +} + +impl Drop for PrefillTicket { + fn drop(&mut self) { + if !self.released { + self.worker.note_prefill_ended(self.tokens, false); + } + } +} + +#[cfg(test)] +mod tracked_tests { + use std::sync::Arc; + + use super::{FanoutStream, ProtoStream}; + use crate::worker::{BasicWorkerBuilder, Worker}; + + #[test] + fn a_tracked_stream_is_one_of_its_workers_pile_for_its_life() { + let worker: Arc = Arc::new(BasicWorkerBuilder::new("http://w1:8000").build()); + let stream = ProtoStream::Fanout(FanoutStream::::new(Vec::new())).tracked( + Arc::clone(&worker), + 0, + true, + ); + assert_eq!(worker.tracked_load(), 1, "counted at dispatch"); + let stream = stream.defer_abort_until_first_item(); + assert_eq!(worker.tracked_load(), 1, "and across the stream's rewraps"); + drop(stream); + assert_eq!(worker.tracked_load(), 0, "released with the stream"); + } + + #[test] + fn a_non_streaming_generation_joins_no_pile_but_keeps_its_prefill_and_its_progress() { + // Review finding on the pushed head: left untracked altogether, a + // non-streaming generation's one answer no longer counted as progress + // and its prompt left the prefill backlog, so a saturated engine + // finishing such requests while streaming ones waited was wedged. It + // is tracked like any other, minus the pile slot. + let worker: Arc = Arc::new(BasicWorkerBuilder::new("http://w1:8000").build()); + let stream = ProtoStream::Fanout(FanoutStream::::new(Vec::new())).tracked( + Arc::clone(&worker), + 64, + false, + ); + assert_eq!( + worker.tracked_load(), + 0, + "nothing to see progress on until it answers" + ); + match &stream { + ProtoStream::Tracked(tracked) => { + assert!(tracked.progress.is_none(), "no pile slot"); + assert!( + tracked.prefill.is_some(), + "its prompt is pending prefill there like any other" + ); + // Its one answer goes through the same `Tracked` arm of + // `next` as a token, which is where progress is recorded. + } + _ => panic!("a tracked stream"), + } + let stream = stream.defer_abort_until_first_item(); + assert_eq!(worker.tracked_load(), 0); + drop(stream); + assert_eq!(worker.tracked_load(), 0); + } } /// Surface an engine-side failure (`finish_reason == "error"`) as a stream error, like the ZMQ lane. @@ -2534,6 +2689,38 @@ fn reject_engine_error( } impl ProtoStream { + /// Count this stream's responses as progress for `worker` (a response is + /// progress whether it is a token of a streaming generation or the one + /// answer of a non-streaming one), its `prompt_tokens` as pending prefill + /// there until the first response, and, when the engine streams the + /// generation (`streaming`), the request as one of the worker's tracked + /// pile while the stream lives. A non-streaming generation shows nothing + /// between dispatch and completion, so it joins no pile: the pile rule + /// never judges a worker by requests it cannot see progress on. + #[must_use] + pub fn tracked(self, worker: Arc, prompt_tokens: u64, streaming: bool) -> Self { + let progress = streaming.then(|| { + worker.note_tracked_started(); + ProgressTicket { + worker: Arc::clone(&worker), + } + }); + let prefill = (prompt_tokens > 0).then(|| { + worker.note_prefill_started(prompt_tokens); + PrefillTicket { + worker: Arc::clone(&worker), + tokens: prompt_tokens, + released: false, + } + }); + Self::Tracked(Box::new(TrackedStream { + inner: self, + worker, + prefill, + progress, + })) + } + /// Get next item from stream pub async fn next(&mut self) -> Option> { let item = match self { @@ -2570,6 +2757,17 @@ impl ProtoStream { .await .map(|result| result.map(|r| ProtoGenerateResponse::Vllm(Box::new(r)))), Self::Fanout(stream) => stream.next().await, + Self::Tracked(tracked) => { + let item = Box::pin(tracked.inner.next()).await; + if matches!(item, Some(Ok(_))) { + liveness::on_token_progress(&tracked.worker); + if let Some(ticket) = tracked.prefill.as_mut() { + ticket.first_response(); + } + tracked.prefill = None; + } + return item; + } }; item.map(reject_engine_error) } @@ -2584,6 +2782,7 @@ impl ProtoStream { Self::TokenSpeed(stream) => stream.mark_completed(), Self::Zmq(stream) => stream.mark_completed(), Self::Fanout(stream) => stream.mark_completed(), + Self::Tracked(stream) => stream.inner.mark_completed(), } } @@ -2605,6 +2804,20 @@ impl ProtoStream { Self::TokenSpeed(stream) => Self::TokenSpeed(stream.defer_abort_until_first_item()), Self::Zmq(stream) => Self::Zmq(stream), Self::Fanout(stream) => Self::Fanout(stream.defer_abort_until_first_item()), + Self::Tracked(stream) => { + let TrackedStream { + inner, + worker, + prefill, + progress, + } = *stream; + Self::Tracked(Box::new(TrackedStream { + inner: inner.defer_abort_until_first_item(), + worker, + prefill, + progress, + })) + } } } } diff --git a/model_gateway/src/routers/grpc/regular/streaming/eof_tests.rs b/model_gateway/src/routers/grpc/regular/streaming/eof_tests.rs index 984e6821f9..22ae303ea7 100644 --- a/model_gateway/src/routers/grpc/regular/streaming/eof_tests.rs +++ b/model_gateway/src/routers/grpc/regular/streaming/eof_tests.rs @@ -118,6 +118,8 @@ async fn scripted_stream( .expect("bind mock worker"); let port = listener.local_addr().expect("mock worker address").port(); let config = Arc::new(mock_worker::config::Config { + admin_port: None, + context_length: 32768, host: "127.0.0.1".to_string(), http_base_port: 0, http_count: 0, @@ -132,7 +134,7 @@ async fn scripted_stream( output_tokens: 0, realistic: false, engine: mock_worker::engine::EngineParams::default(), - replay: Default::default(), + ..mock_worker::config::Config::default() }); let server = tokio::spawn(mock_worker::grpc::serve_with_listener(config, listener)); let client = VllmEngineClient::connect(&format!("http://127.0.0.1:{port}")) diff --git a/model_gateway/src/routers/http/pd_router.rs b/model_gateway/src/routers/http/pd_router.rs index 6892756faa..af8a5f9a8b 100644 --- a/model_gateway/src/routers/http/pd_router.rs +++ b/model_gateway/src/routers/http/pd_router.rs @@ -572,8 +572,12 @@ impl PDRouter { // Dispatch-time re-check of both legs, the same one the regular HTTP // and gRPC paths take just before their load guards. - if let Some(shed) = overload::shed_if_worker_overloaded(prefill.as_ref(), context.model_id) - .or_else(|| overload::shed_if_worker_overloaded(decode.as_ref(), context.model_id)) + let shedding = self.worker_registry.overload_shed_enabled(); + if let Some(shed) = + overload::shed_if_worker_overloaded(prefill.as_ref(), context.model_id, shedding) + .or_else(|| { + overload::shed_if_worker_overloaded(decode.as_ref(), context.model_id, shedding) + }) { return shed; } @@ -2711,7 +2715,11 @@ impl RouterTrait for PDRouter { Ok(bytes) => Bytes::from(bytes), Err(e) => return Self::handle_serialization_error(e), }; - if let Some(response) = overload::shed_if_worker_overloaded(worker.as_ref(), model_id) { + if let Some(response) = overload::shed_if_worker_overloaded( + worker.as_ref(), + model_id, + self.worker_registry.overload_shed_enabled(), + ) { return response; } let _load_guard = WorkerLoadGuard::with_key( diff --git a/model_gateway/src/routers/http/router.rs b/model_gateway/src/routers/http/router.rs index 891830bf50..f5d4835be6 100644 --- a/model_gateway/src/routers/http/router.rs +++ b/model_gateway/src/routers/http/router.rs @@ -543,7 +543,11 @@ impl Router { // Dispatch-time re-check of the one chosen worker: O(1), and the only // thing that closes the window between selection and dispatch in which // a load report can flip the veto. - if let Some(shed) = overload::shed_if_worker_overloaded(worker.as_ref(), model_id) { + if let Some(shed) = overload::shed_if_worker_overloaded( + worker.as_ref(), + model_id, + self.worker_registry.overload_shed_enabled(), + ) { return shed; } @@ -841,7 +845,11 @@ impl Router { // Judged from the same candidates whether the pre-filter emptied // or a self-filtering policy missed on an all-overloaded pool, so // a shed keeps its Retry-After, retryability and metric. - let resp = match placement::failure_from(&non_dp_workers, model_id) { + let resp = match placement::failure_from( + &non_dp_workers, + model_id, + self.worker_registry.overload_shed_enabled(), + ) { PlacementFailure::AllOverloaded(shed) => shed, PlacementFailure::NoCandidates | PlacementFailure::Unavailable @@ -866,7 +874,11 @@ impl Router { // occupies its worker for far longer than a chat completion, so a // report landing in the selection→dispatch window is the one case where // dispatching anyway is measurably worse. - if let Some(resp) = overload::shed_if_worker_overloaded(worker.as_ref(), model_id) { + if let Some(resp) = overload::shed_if_worker_overloaded( + worker.as_ref(), + model_id, + self.worker_registry.overload_shed_enabled(), + ) { record_pre_send_error(&resp); return resp; } @@ -2611,6 +2623,8 @@ mod tests { cache_index: Default::default(), cache_ttl_secs: 180, cache_boundaries: Vec::new(), + selection_policy: None, + selection_accounting_ttl_ms: 0, } } diff --git a/model_gateway/src/worker/builder.rs b/model_gateway/src/worker/builder.rs index 987c6aecd3..c3445530f6 100644 --- a/model_gateway/src/worker/builder.rs +++ b/model_gateway/src/worker/builder.rs @@ -380,6 +380,7 @@ impl BasicWorkerBuilder { models_override: Arc::new(ArcSwap::from_pointee(WorkerModels::Wildcard)), http_client, resilience, + completion_sink: Arc::new(std::sync::RwLock::new(None)), } } } diff --git a/model_gateway/src/worker/kv_event_monitor.rs b/model_gateway/src/worker/kv_event_monitor.rs index b272ae4813..6fd580661f 100644 --- a/model_gateway/src/worker/kv_event_monitor.rs +++ b/model_gateway/src/worker/kv_event_monitor.rs @@ -1,7 +1,8 @@ //! Per-worker KV cache event subscription manager. //! //! `KvEventMonitor` spawns a background tokio task per gRPC worker that subscribes -//! to KV cache events and feeds them into a shared `PositionalIndexer` (one per model). +//! to KV cache events and feeds them into a shared [`KvIndex`] (one per model; +//! the positional indexer or the chain index, per `--kv-index`). //! This enables event-driven cache-aware routing as an alternative to the approximate //! radix tree approach. //! @@ -9,121 +10,183 @@ //! - `on_worker_added` — spawns streaming task, creates indexer if needed //! - `on_worker_removed` — signals graceful shutdown, task cleans up indexer //! - `stop` — signals shutdown to all tasks, clears state - -use std::{collections::HashMap, fmt, sync::Arc, time::Duration}; +//! +//! The work of a subscription is split by stage: `subscription.rs` runs the +//! task per worker (connecting, reconnecting, the stream read loop, the +//! pushed load records, the worker's departure); `admission.rs` keeps one +//! stream's state (per-rank cursors, gap recovery, relay snapshots, and the +//! metrics around an applied batch); `apply.rs` takes one event into the +//! index (tiers, cache groups, namespaces, and the physical copies of a +//! block, counted per worker). + +use std::{ + collections::HashMap, + fmt, + sync::{ + atomic::{AtomicU64, Ordering}, + Arc, OnceLock, Weak, + }, +}; use dashmap::DashMap; use futures::FutureExt as _; -use kv_index::{ - compute_content_hash, ApplyError, PositionalIndexer, SequenceHash, StoredBlock, WorkerBlockMap, -}; -use smg_grpc_client::common_proto::{ - kv_cache_event, KvBlock, KvBlocksRemoved, KvBlocksStored, KvCacheEvent, KvEventBatch, -}; use tokio::{ - sync::{oneshot, Mutex, Semaphore}, + sync::{oneshot, watch, Mutex}, task::JoinHandle, }; use tracing::{debug, error, info, warn}; +use super::{ + kv_index_backend::{KvIndex, KvIndexKind}, + monitor::WorkerMonitor, +}; use crate::{ observability::metrics::Metrics, policies::utils::PeriodicTask, worker::{ConnectionMode, Worker, UNKNOWN_MODEL_ID}, }; -/// Default jump size for new `PositionalIndexer` instances. +mod admission; +mod apply; +mod subscription; + +pub(crate) use apply::WorkerIndexState; + +/// Default jump size for new positional indexers. const DEFAULT_JUMP_SIZE: usize = 64; /// Interval between positional-indexer prune cycles (matches the routing /// policies' default eviction cadence). const PRUNE_INTERVAL_SECS: u64 = 30; -/// Initial reconnection delay after stream failure. -const INITIAL_RECONNECT_DELAY_MS: u64 = 100; - -/// Maximum reconnection delay (caps exponential backoff). -const MAX_RECONNECT_DELAY_MS: u64 = 30_000; - -/// Positional-index cleanup is CPU-bound and can touch many blocks. Keep it -/// off Tokio workers and bound concurrent purges during fleet-wide drains. -const MAX_CONCURRENT_INDEX_REMOVALS: usize = 4; -static INDEX_REMOVAL_PERMITS: Semaphore = Semaphore::const_new(MAX_CONCURRENT_INDEX_REMOVALS); +/// Interval between publications of each model's index size and, for the chain +/// index, its shape and memory (`smg_kv_index_*` gauges by model). +const STATS_INTERVAL_SECS: u64 = 30; /// Manages per-worker KV cache event subscriptions. /// /// Each gRPC worker gets a dedicated tokio task that subscribes to the backend's -/// KV cache event stream and feeds events into a shared `PositionalIndexer` +/// KV cache event stream and feeds events into a shared [`KvIndex`] /// (one per `model_id`). Workers serving the same model share the same indexer. pub struct KvEventMonitor { - /// Per-model positional indexers: model_id → shared indexer. + /// Per-model KV indexes: model_id → shared indexer. /// Arc-wrapped so the prune task can share the map WITHOUT holding (even /// weakly) the monitor itself: `PeriodicTask` joins its thread on drop, so /// a task that could ever own the last monitor reference would run the /// monitor's drop — and thus its own join — on its own thread. - pub(crate) indexers: Arc>>, + pub(crate) indexers: Arc>>, /// Per-model block sizes learned from KV events or set via WorkerSpec. /// Used by CacheAwarePolicy to chunk request tokens at query time. /// Arc-wrapped so subscription tasks can update it from events. block_sizes: Arc>, - /// Per-worker subscription handles: worker_url → subscription info. - /// Mutex matches LoadMonitor pattern for atomic abort + remove. - worker_handles: Mutex>, - /// Jump size for new PositionalIndexer instances. + /// Per-worker subscription slots: worker_url → the live subscription, or + /// the reservation a removal leaves until the old task has cleaned up. + /// Mutex matches LoadMonitor pattern for atomic abort + remove; Arc so a + /// task can lift its own reservation at the end of its exit path. + worker_handles: Arc>>, + /// Which index new models get. + kind: KvIndexKind, + /// Jump size for new positional indexers. jump_size: usize, /// Periodic indexer prune, held so it aborts when the monitor drops. /// Set once by [`start_prune_task`](Self::start_prune_task); sync mutex /// because it is touched only at startup, never on event paths. prune_task: parking_lot::Mutex>, + /// Periodic index-shape publication, held like the prune task. + stats_task: parking_lot::Mutex>, + /// Where the load records on the event streams go (`KvEventBatch.load`): + /// the worker monitor, which treats them as polls. Weak: the monitors + /// are peers in the app context, neither owns the other. + load_sink: OnceLock>, +} + +/// A worker URL's slot in the subscription map. +enum Slot { + Active(WorkerSubscription), + /// The URL's old task has been told to shut down and is taking its + /// blocks out of the index. The index hands a URL its existing id until + /// that cleanup has released it, so a subscription interned in this + /// window would write under the old id and lose its blocks to the + /// cleanup; `on_worker_added` waits for the old task to end instead. + Removing(Reservation), +} + +/// What a removal leaves in a URL's slot until the old task has ended. +/// +/// The task lifts it itself at the end of its exit path +/// ([`KvEventMonitor::complete_removal`]), and `done` closes when the task +/// is gone for any reason (an abort, a panic past its guard). Nothing about +/// it depends on the `on_worker_removed` future living on: the registry's +/// removal step runs it under a timeout, and a cleanup waiting on the +/// removal permits can outlast that; a reservation only the remover could +/// lift would then hold every later add of the URL for good. +struct Reservation { + /// The subscription the reservation belongs to; a task lifts only its own. + id: u64, + /// Closes when the task has ended. + done: watch::Receiver<()>, } /// Tracks a single worker's subscription state. struct WorkerSubscription { + /// Tells this subscription from a later one under the same URL. + id: u64, handle: JoinHandle<()>, model_id: String, - /// Signals the subscription task to shut down gracefully. - /// The task owns its `WorkerBlockMap` and cleans up the indexer on exit. + /// Signals the subscription task to shut down gracefully. The task owns + /// the worker's index state and takes its blocks out of the index on exit. shutdown_tx: oneshot::Sender<()>, + /// Closes when the task has ended; a removal moves it into the + /// reservation it leaves in the slot. + done: watch::Receiver<()>, } -/// Result of processing a stream connection to completion. -enum StreamResult { - /// Stream closed normally (server-side). - Ended, - /// Stream produced an error. - Error(tonic::Status), - /// Detected a gap in sequence numbers. - GapDetected { expected: u64, received: u64 }, -} +/// Ids for [`WorkerSubscription`]s, unique within the process. +static NEXT_SUBSCRIPTION_ID: AtomicU64 = AtomicU64::new(1); impl KvEventMonitor { - /// Create a new `KvEventMonitor`. + /// A monitor whose models get positional indexers. /// - /// `jump_size` controls the `PositionalIndexer` jump search stride. + /// `jump_size` is the positional indexer's historical tuning knob. /// Pass `None` for the default (64). pub fn new(jump_size: Option) -> Self { + Self::with_kind(KvIndexKind::Positional, jump_size) + } + + /// A monitor whose models get indexes of `kind` (see `--kv-index`). + pub fn with_kind(kind: KvIndexKind, jump_size: Option) -> Self { let jump_size = jump_size.unwrap_or(DEFAULT_JUMP_SIZE).max(1); Self { indexers: Arc::new(DashMap::new()), block_sizes: Arc::new(DashMap::new()), - worker_handles: Mutex::new(HashMap::new()), + worker_handles: Arc::new(Mutex::new(HashMap::new())), + kind, jump_size, prune_task: parking_lot::Mutex::new(None), + stats_task: parking_lot::Mutex::new(None), + load_sink: OnceLock::new(), } } + /// The kind of index this monitor builds per model. + pub fn kind(&self) -> KvIndexKind { + self.kind + } + + /// Route the load records on every worker's event stream to the worker + /// monitor (once; a later call is ignored). + pub fn set_load_sink(&self, monitor: &Arc) { + let _ = self.load_sink.set(Arc::downgrade(monitor)); + } + /// Prune every model's positional indexer with the given bounds. /// `ttl_secs`/`max_entries` of 0 disable the respective pass — see - /// [`PositionalIndexer::prune`]. + /// [`KvIndex::prune`]. The chain index has no prune and is left alone. pub fn prune_all(&self, ttl_secs: u64, max_entries: usize) { Self::prune_indexers(&self.indexers, ttl_secs, max_entries); } - fn prune_indexers( - indexers: &DashMap>, - ttl_secs: u64, - max_entries: usize, - ) { + fn prune_indexers(indexers: &DashMap>, ttl_secs: u64, max_entries: usize) { let ttl = u32::try_from(ttl_secs).unwrap_or(u32::MAX - 1); let ttl = (ttl > 0).then_some(ttl); let max = (max_entries > 0).then_some(max_entries); @@ -131,7 +194,9 @@ impl KvEventMonitor { return; } for entry in indexers { - let stats = entry.value().prune(ttl, max); + let Some(stats) = entry.value().prune(ttl, max) else { + continue; + }; if stats.evicted_ttl + stats.evicted_capacity > 0 { info!( model_id = %entry.key(), @@ -144,6 +209,32 @@ impl KvEventMonitor { } } + /// Publish every model's index size, and the chain index's shape and + /// memory, as gauges: what a soak reads to tell fragmentation or growth + /// from load. `current_size` and `entry_count` are counter reads; the chain + /// index's `stats` walks its run headers, which is why this runs on a + /// 30 s cadence and never on a request. + fn publish_stats(indexers: &DashMap>) { + for entry in indexers { + let index = entry.value(); + Metrics::set_kv_index_size(entry.key(), index.current_size(), index.entry_count()); + if let Some(stats) = index.chain_stats() { + Metrics::set_kv_index_chain_stats(entry.key(), &stats); + } + } + } + + /// Start the periodic index-shape publication. Like the prune task, it + /// shares only the indexer map, never the monitor, and stops when the + /// monitor drops. + pub fn start_stats_task(&self) { + let indexers = Arc::clone(&self.indexers); + let task = PeriodicTask::spawn(STATS_INTERVAL_SECS, "KvIndexStats", move || { + Self::publish_stats(&indexers); + }); + *self.stats_task.lock() = Some(task); + } + /// Start the periodic indexer prune. No-op when both bounds are 0/unset. /// The task shares only the indexer map — never a reference to the monitor /// itself — so it can never be the one to run the monitor's drop (and with @@ -153,6 +244,16 @@ impl KvEventMonitor { if ttl_secs == 0 && max_entries == 0 { return; } + if self.kind == KvIndexKind::Chain { + warn!( + ttl_secs, + max_entries, + "The chain index has no prune: it holds what the engines report and \ + shrinks with their removals; --kv-indexer-ttl-secs and \ + --kv-indexer-max-entries apply to --kv-index positional only" + ); + return; + } let indexers = Arc::clone(&self.indexers); let task = PeriodicTask::spawn(PRUNE_INTERVAL_SECS, "KvIndexerPrune", move || { Self::prune_indexers(&indexers, ttl_secs, max_entries); @@ -169,7 +270,7 @@ impl KvEventMonitor { /// Start a KV event subscription for a worker. /// /// Spawns a background tokio task that subscribes to KV cache events via - /// server-streaming gRPC and applies them to the model's `PositionalIndexer`. + /// server-streaming gRPC and applies them to the model's `KvIndex`. /// Duplicate calls for the same worker URL are no-ops. pub async fn on_worker_added(&self, worker: &Arc) { let url = worker.url().to_string(); @@ -184,16 +285,33 @@ impl KvEventMonitor { return; } - let mut handles = self.worker_handles.lock().await; - if handles.contains_key(&url) { - debug!(worker_url = %url, "KV event subscription already active, skipping"); - return; - } + let mut handles = loop { + let mut handles = self.worker_handles.lock().await; + let mut done = match handles.get(&url) { + Some(Slot::Active(_)) => { + debug!(worker_url = %url, "KV event subscription already active, skipping"); + return; + } + Some(Slot::Removing(reservation)) => reservation.done.clone(), + None => break handles, + }; + // The previous subscription for this URL is still being taken + // out of the index: intern it again only once its task has ended. + if done.has_changed().is_err() { + // The task is gone without lifting its reservation (aborted, + // or a panic past its guard): nothing is left to wait for. + handles.remove(&url); + break handles; + } + drop(handles); + // Nothing is ever sent on the channel; this returns at its close. + let _ = done.changed().await; + }; let indexer = self .indexers .entry(model_id.clone()) - .or_insert_with(|| Arc::new(PositionalIndexer::new(self.jump_size))) + .or_insert_with(|| Arc::new(KvIndex::new(self.kind, self.jump_size))) .clone(); // Seed block_size provisionally from WorkerSpec. The event stream will // overwrite this with the backend's actual page size once received. @@ -216,8 +334,15 @@ impl KvEventMonitor { ); let (shutdown_tx, shutdown_rx) = oneshot::channel(); + let (done_tx, done) = watch::channel(()); + let id = NEXT_SUBSCRIPTION_ID.fetch_add(1, Ordering::Relaxed); + let load_sink = self.load_sink.get().cloned(); let loop_model_id = model_id.clone(); let task_url = url.clone(); + let task_model_id = model_id.clone(); + let slots = Arc::clone(&self.worker_handles); + let indexers = Arc::clone(&self.indexers); + let model_block_sizes = Arc::clone(&self.block_sizes); #[expect( clippy::disallowed_methods, @@ -235,6 +360,7 @@ impl KvEventMonitor { block_sizes, loop_model_id, shutdown_rx, + load_sink, )) .catch_unwind() .await; @@ -253,28 +379,91 @@ impl KvEventMonitor { ); Metrics::record_kv_event_subscription_failure(&task_url, "panic"); } + Self::complete_removal( + &slots, + &indexers, + &model_block_sizes, + &task_url, + id, + &task_model_id, + ) + .await; + // Closes `done` last: an add waiting on it finds the slot free. + drop(done_tx); }); handles.insert( url, - WorkerSubscription { + Slot::Active(WorkerSubscription { + id, handle, model_id, shutdown_tx, - }, + done, + }), ); } + /// The end of a removal, run by the subscription task once its blocks + /// are out of the index: lift the URL's reservation if it is this task's + /// (a later subscription under the URL has its own), and drop the + /// model's index with the model's last subscription. The task does this + /// rather than `on_worker_removed`, whose future a caller may drop + /// before the task has ended; a slot that is `Active` (the task ended on + /// its own, or `stop` took the subscription) is left as it is. + async fn complete_removal( + slots: &Mutex>, + indexers: &DashMap>, + block_sizes: &DashMap, + worker_url: &str, + id: u64, + model_id: &str, + ) { + let mut handles = slots.lock().await; + if !matches!(handles.get(worker_url), Some(Slot::Removing(held)) if held.id == id) { + return; + } + handles.remove(worker_url); + // Under the slot lock, which `on_worker_added` holds from its lookup + // of the model's index to its insert: an add of the model racing + // with this either keeps the index (its slot is already in) or + // creates the next one. + let last_of_model = !handles + .values() + .any(|slot| matches!(slot, Slot::Active(other) if other.model_id == model_id)); + if last_of_model { + indexers.remove(model_id); + block_sizes.remove(model_id); + } + } + /// Stop the KV event subscription for a worker. /// - /// Sends a graceful shutdown signal. The subscription task cleans up its - /// own `WorkerBlockMap` using the indexer's per-worker reverse map; that + /// Sends a graceful shutdown signal. The subscription task takes the + /// worker's blocks out of the index from its own per-worker map; that /// CPU-bound cleanup runs on the bounded blocking pool rather than a Tokio /// runtime worker. pub async fn on_worker_removed(&self, worker_url: &str) { let subscription = { let mut handles = self.worker_handles.lock().await; - handles.remove(worker_url) + match handles.remove(worker_url) { + Some(Slot::Active(sub)) => { + handles.insert( + worker_url.to_string(), + Slot::Removing(Reservation { + id: sub.id, + done: sub.done.clone(), + }), + ); + Some(sub) + } + Some(removing @ Slot::Removing(_)) => { + // Another removal of the same URL is already under way. + handles.insert(worker_url.to_string(), removing); + None + } + None => None, + } }; let Some(sub) = subscription else { @@ -283,6 +472,10 @@ impl KvEventMonitor { info!(worker_url = %worker_url, "Stopping KV event subscription"); // Signal graceful shutdown — task cleans up its worker_blocks in the indexer. let _ = sub.shutdown_tx.send(()); + // The task lifts the reservation and drops the model's last index on + // its own (`complete_removal`); this wait is for the caller, who has + // the worker out of the index on return. A caller that drops the + // future here loses nothing but the join error below. // Panics are caught inside the task; a JoinError here (abort or a // panic that escaped the guard) must still be surfaced, not discarded. if let Err(e) = sub.handle.await { @@ -293,27 +486,24 @@ impl KvEventMonitor { ); Metrics::record_kv_event_subscription_failure(worker_url, "join_error"); } - - // Re-check under lock whether this was the last worker for the model. - // Must re-acquire lock after shutdown to avoid TOCTOU with concurrent - // on_worker_added that may have added a new worker for the same model - // between our first lock release and this point. - let should_remove_indexer = { - let handles = self.worker_handles.lock().await; - !handles.values().any(|other| other.model_id == sub.model_id) - }; - - if should_remove_indexer { - self.indexers.remove(&sub.model_id); - self.block_sizes.remove(&sub.model_id); - } } /// Stop all subscriptions and clean up. pub async fn stop(&self) { - let subscriptions: HashMap = { + let subscriptions: Vec<(String, WorkerSubscription)> = { let mut handles = self.worker_handles.lock().await; - std::mem::take(&mut *handles) + let mut active = Vec::new(); + for (url, slot) in std::mem::take(&mut *handles) { + match slot { + Slot::Active(sub) => active.push((url, sub)), + // The task behind a removal in flight lifts its own + // reservation. + reserved @ Slot::Removing(_) => { + handles.insert(url, reserved); + } + } + } + active }; if !subscriptions.is_empty() { @@ -340,7 +530,7 @@ impl KvEventMonitor { } /// Get the indexer for a model (used by `CacheAwarePolicy` for queries). - pub fn get_indexer(&self, model_id: &str) -> Option> { + pub fn get_indexer(&self, model_id: &str) -> Option> { self.indexers.get(model_id).map(|r| Arc::clone(&r)) } @@ -357,11 +547,6 @@ impl KvEventMonitor { .or_insert(block_size); } - /// Check if any subscription is running. - pub async fn is_running(&self) -> bool { - !self.worker_handles.lock().await.is_empty() - } - /// Normalize model_id to match routing's `normalize_model_key`. /// Empty model IDs map to UNKNOWN_MODEL_ID for consistent keying. fn normalize_model_id(model_id: &str) -> String { @@ -371,419 +556,69 @@ impl KvEventMonitor { model_id.to_string() } } +} - // ----------------------------------------------------------------------- - // Subscription loop - // ----------------------------------------------------------------------- +/// Hooks for benchmarks and integration tests that put an index in front of +/// the policy and feed it events the way a subscription does, without a +/// stream: the subscriber's per-worker state and the same apply path. +#[cfg(any(test, feature = "test-util"))] +pub mod bench_support { + use std::sync::Arc; - /// Learn `block_size` from the first `KvBlock` in a stored event. - /// - /// Called once per model when the first stored event arrives, providing - /// ground truth from the backend. `CacheAwarePolicy` uses this to chunk - /// request tokens into blocks for overlap scoring. - /// - /// Overwrites any provisional value seeded from `WorkerSpec` since the - /// event stream reflects the backend's actual page size. - fn learn_block_size( - block_sizes: &DashMap, - model_id: &str, - learned: &mut bool, - batch: &KvEventBatch, - ) { - if *learned { - return; - } - for event in &batch.events { - if let Some(kv_cache_event::Data::Stored(stored)) = &event.data { - if let Some(block) = stored.blocks.first() { - if block.block_size > 0 { - let bs = block.block_size as usize; - block_sizes.insert(model_id.to_string(), bs); - info!( - model_id = %model_id, - block_size = bs, - "Learned block_size from KV event" - ); - *learned = true; - return; - } - } - } - } - } + use kv_index::WorkerIdExhausted; + use smg_grpc_client::common_proto::KvEventBatch; - /// Main subscription loop for a single worker. - /// - /// Owns the `WorkerBlockMap` for this worker and cleans it up on exit. - /// Exits when `shutdown_rx` fires or the backend returns `Unimplemented`. - async fn remove_indexer_worker( - indexer: Arc, - worker_id: u32, - worker_blocks: WorkerBlockMap, - ) { - let Ok(permit) = INDEX_REMOVAL_PERMITS.acquire().await else { - error!(worker_id, "Positional-index cleanup semaphore closed"); - return; - }; - let result = tokio::task::spawn_blocking(move || { - let _permit = permit; - indexer.remove_worker(worker_id, worker_blocks); - }) - .await; + use super::{KvEventMonitor, KvIndex, WorkerIndexState}; - if let Err(error) = result { - error!(worker_id, %error, "Positional-index worker cleanup task failed"); + impl KvEventMonitor { + /// Make `index` the model's index, as a worker's first subscription + /// would; a later subscription for the model shares it. + pub fn set_index(&self, model_id: &str, index: Arc) { + self.indexers.insert(model_id.to_string(), index); } } - async fn subscription_loop( - worker: Arc, - worker_url: String, - indexer: Arc, - block_sizes: Arc>, - model_id: String, - mut shutdown_rx: oneshot::Receiver<()>, - ) { - let worker_id = match indexer.intern_worker(&worker_url) { - Ok(id) => id, - Err(e) => { - error!( - worker_url = %worker_url, - error = %e, - "Failed to intern worker; KV events from this worker will \ - not feed cache-aware routing" - ); - Metrics::record_kv_event_subscription_failure(&worker_url, "intern_failed"); - return; - } - }; - let mut worker_blocks = WorkerBlockMap::default(); - let mut last_seq: u64 = 0; - let mut reconnect_delay_ms = INITIAL_RECONNECT_DELAY_MS; - let mut block_size_learned = false; - - /// Sleep with shutdown check. Returns `true` if shutdown was signaled. - macro_rules! sleep_or_shutdown { - ($delay:expr, $rx:expr) => { - tokio::select! { - _ = tokio::time::sleep($delay) => false, - _ = &mut *$rx => true, - } - }; - } - - loop { - let backend_client = match worker.get_backend_client().await { - Ok(Some(client)) => client, - Ok(None) => { - // HTTP workers are filtered in on_worker_added, so this should - // be unreachable. Retry defensively rather than exiting and - // leaving a stale entry in worker_handles. - warn!( - worker_url = %worker_url, - delay_ms = reconnect_delay_ms, - "Worker has no backend client yet, retrying" - ); - if sleep_or_shutdown!( - Duration::from_millis(reconnect_delay_ms), - &mut shutdown_rx - ) { - Self::remove_indexer_worker(Arc::clone(&indexer), worker_id, worker_blocks) - .await; - return; - } - reconnect_delay_ms = (reconnect_delay_ms * 2).min(MAX_RECONNECT_DELAY_MS); - continue; - } - Err(e) => { - warn!( - worker_url = %worker_url, - error = %e, - delay_ms = reconnect_delay_ms, - "Failed to get backend client, retrying" - ); - if sleep_or_shutdown!( - Duration::from_millis(reconnect_delay_ms), - &mut shutdown_rx - ) { - Self::remove_indexer_worker(Arc::clone(&indexer), worker_id, worker_blocks) - .await; - return; - } - reconnect_delay_ms = (reconnect_delay_ms * 2).min(MAX_RECONNECT_DELAY_MS); - continue; - } - }; - - let stream = match backend_client.subscribe_kv_events(last_seq).await { - Ok(stream) => { - info!( - worker_url = %worker_url, - start_seq = last_seq, - "KV event stream connected" - ); - reconnect_delay_ms = INITIAL_RECONNECT_DELAY_MS; - stream - } - Err(e) => { - // If the backend doesn't implement SubscribeKvEvents (e.g. vLLM), - // stop retrying — this RPC will never succeed. - if e.code() == tonic::Code::Unimplemented { - warn!( - worker_url = %worker_url, - "Backend does not implement SubscribeKvEvents, \ - disabling KV event subscription for this worker" - ); - Self::remove_indexer_worker(Arc::clone(&indexer), worker_id, worker_blocks) - .await; - return; - } - if e.code() == tonic::Code::OutOfRange { - warn!( - worker_url = %worker_url, - last_seq = last_seq, - "KV event replay cursor expired; clearing worker state and requesting a current snapshot" - ); - indexer.apply_cleared(worker_id, &mut worker_blocks); - last_seq = 0; - reconnect_delay_ms = INITIAL_RECONNECT_DELAY_MS; - continue; - } - warn!( - worker_url = %worker_url, - error = %e, - delay_ms = reconnect_delay_ms, - "Failed to subscribe to KV events, retrying" - ); - if sleep_or_shutdown!( - Duration::from_millis(reconnect_delay_ms), - &mut shutdown_rx - ) { - Self::remove_indexer_worker(Arc::clone(&indexer), worker_id, worker_blocks) - .await; - return; - } - reconnect_delay_ms = (reconnect_delay_ms * 2).min(MAX_RECONNECT_DELAY_MS); - continue; - } - }; - - let on_batch = |batch: &KvEventBatch| { - Self::learn_block_size(&block_sizes, &model_id, &mut block_size_learned, batch); - }; - let stream_result = tokio::select! { - result = Self::process_stream( - stream, &worker_url, worker_id, &indexer, - &mut worker_blocks, &mut last_seq, on_batch, - ) => result, - _ = &mut shutdown_rx => { - Self::remove_indexer_worker( - Arc::clone(&indexer), - worker_id, - worker_blocks, - ) - .await; - return; - } - }; - - match stream_result { - StreamResult::Ended => { - info!( - worker_url = %worker_url, - last_seq = last_seq, - delay_ms = reconnect_delay_ms, - "KV event stream ended, reconnecting" - ); - // Backoff to avoid tight reconnect loop if server keeps - // closing the stream cleanly (e.g., rolling connections). - if sleep_or_shutdown!( - Duration::from_millis(reconnect_delay_ms), - &mut shutdown_rx - ) { - Self::remove_indexer_worker(Arc::clone(&indexer), worker_id, worker_blocks) - .await; - return; - } - reconnect_delay_ms = (reconnect_delay_ms * 2).min(MAX_RECONNECT_DELAY_MS); - } - StreamResult::Error(e) => { - if e.code() == tonic::Code::DataLoss { - warn!( - worker_url = %worker_url, - error = %e, - last_seq = last_seq, - "KV event subscriber fell behind; clearing worker state and requesting a current snapshot" - ); - indexer.apply_cleared(worker_id, &mut worker_blocks); - last_seq = 0; - reconnect_delay_ms = INITIAL_RECONNECT_DELAY_MS; - continue; - } - warn!( - worker_url = %worker_url, - error = %e, - last_seq = last_seq, - delay_ms = reconnect_delay_ms, - "KV event stream error, reconnecting" - ); - if sleep_or_shutdown!( - Duration::from_millis(reconnect_delay_ms), - &mut shutdown_rx - ) { - Self::remove_indexer_worker(Arc::clone(&indexer), worker_id, worker_blocks) - .await; - return; - } - reconnect_delay_ms = (reconnect_delay_ms * 2).min(MAX_RECONNECT_DELAY_MS); - } - StreamResult::GapDetected { expected, received } => { - warn!( - worker_url = %worker_url, - expected = expected, - received = received, - "Sequence gap detected, reconnecting for replay from seq {last_seq}" - ); - // No backoff — gap replay is a normal recovery path. - } - } - } + /// One worker's feed into an index. + pub struct IndexFeed { + worker: u32, + state: WorkerIndexState, } - // ----------------------------------------------------------------------- - // Stream processing + proto conversion - // ----------------------------------------------------------------------- - - /// Process batches from a single stream connection. - async fn process_stream( - mut stream: tonic::Streaming, - worker_url: &str, - worker_id: u32, - indexer: &PositionalIndexer, - worker_blocks: &mut WorkerBlockMap, - last_seq: &mut u64, - mut on_batch: impl FnMut(&KvEventBatch), - ) -> StreamResult { - use tokio_stream::StreamExt; - - while let Some(result) = stream.next().await { - let batch = match result { - Ok(batch) => batch, - Err(e) => return StreamResult::Error(e), - }; - - // Skip stale/duplicate batches (can occur after reconnect replay). - if *last_seq > 0 && batch.sequence_number <= *last_seq { - debug!( - worker_url = %worker_url, - last_seq = *last_seq, - received = batch.sequence_number, - "Skipping stale KV event batch" - ); - continue; - } - - // Gap detection. - if *last_seq > 0 && batch.sequence_number > *last_seq + 1 { - return StreamResult::GapDetected { - expected: *last_seq + 1, - received: batch.sequence_number, - }; - } - - on_batch(&batch); - - for event in &batch.events { - Self::apply_event(event, worker_id, indexer, worker_blocks); - } - - *last_seq = batch.sequence_number; + impl IndexFeed { + /// Intern `worker_url` in `index` and start with empty state. + pub fn new(index: &KvIndex, worker_url: &str) -> Result { + Ok(Self { + worker: index.intern_worker(worker_url)?, + state: WorkerIndexState::default(), + }) } - StreamResult::Ended - } - - /// Apply a single KV cache event to the indexer. - fn apply_event( - event: &KvCacheEvent, - worker_id: u32, - indexer: &PositionalIndexer, - worker_blocks: &mut WorkerBlockMap, - ) { - let Some(ref data) = event.data else { - return; - }; - - match data { - kv_cache_event::Data::Stored(stored) => { - Self::apply_stored(stored, worker_id, indexer, worker_blocks); - } - kv_cache_event::Data::Removed(removed) => { - Self::apply_removed(removed, worker_id, indexer, worker_blocks); - } - kv_cache_event::Data::Cleared(_) => { - indexer.apply_cleared(worker_id, worker_blocks); - } + pub fn worker_id(&self) -> u32 { + self.worker } - } - - /// Convert proto `KvBlocksStored` and apply to the indexer. - fn apply_stored( - stored: &KvBlocksStored, - worker_id: u32, - indexer: &PositionalIndexer, - worker_blocks: &mut WorkerBlockMap, - ) { - let blocks: Vec = stored.blocks.iter().map(convert_kv_block).collect(); - - let parent_seq_hash = stored.parent_block_hash.map(SequenceHash::from); - match indexer.apply_stored(worker_id, &blocks, parent_seq_hash, worker_blocks) { - Ok(()) => {} - Err(ApplyError::WorkerNotTracked | ApplyError::ParentBlockNotFound) => { - // Cold start or parent evicted — retry without parent to start a new chain. - if let Err(e) = indexer.apply_stored(worker_id, &blocks, None, worker_blocks) { - warn!( - worker_id = worker_id, - error = %e, - "Failed to apply stored event after fallback" - ); - } + /// Apply every event of `batch`, as an admitted batch is applied. + pub fn apply(&mut self, index: &KvIndex, batch: &KvEventBatch) { + for event in &batch.events { + KvEventMonitor::apply_event(event, self.worker, index, &mut self.state); } } - } - - /// Convert proto `KvBlocksRemoved` and apply to the indexer. - fn apply_removed( - removed: &KvBlocksRemoved, - worker_id: u32, - indexer: &PositionalIndexer, - worker_blocks: &mut WorkerBlockMap, - ) { - let seq_hashes: Vec = removed - .block_hashes - .iter() - .map(|&h| SequenceHash::from(h)) - .collect(); - - indexer.apply_removed(worker_id, &seq_hashes, worker_blocks); - } -} -/// Convert a proto `KvBlock` to a kv-index `StoredBlock`. -fn convert_kv_block(block: &KvBlock) -> StoredBlock { - StoredBlock { - seq_hash: SequenceHash::from(block.block_hash), - content_hash: compute_content_hash(&block.token_ids), + /// The worker leaves: its blocks go with it. + pub fn remove(self, index: &KvIndex) { + index.remove_worker(self.worker, self.state.blocks); + } } } impl Drop for KvEventMonitor { fn drop(&mut self) { if let Ok(mut handles) = self.worker_handles.try_lock() { - for (_, sub) in handles.drain() { - let _ = sub.shutdown_tx.send(()); - sub.handle.abort(); // Can't await in Drop, abort as fallback + for (_, slot) in handles.drain() { + if let Slot::Active(sub) = slot { + let _ = sub.shutdown_tx.send(()); + sub.handle.abort(); // Can't await in Drop, abort as fallback + } } } } @@ -794,6 +629,7 @@ impl fmt::Debug for KvEventMonitor { f.debug_struct("KvEventMonitor") .field("models", &self.indexers.len()) .field("block_sizes", &self.block_sizes.len()) + .field("kind", &self.kind) .field("jump_size", &self.jump_size) .finish() } @@ -801,367 +637,221 @@ impl fmt::Debug for KvEventMonitor { #[cfg(test)] mod tests { - use super::*; - - // ----------------------------------------------------------------------- - // Proto → kv-index conversion - // ----------------------------------------------------------------------- - - #[test] - fn test_convert_kv_block() { - let block = KvBlock { - block_hash: 42, - token_ids: vec![1, 2, 3, 4], - block_size: 4, - lora_id: None, - cache_level: None, - }; - let stored = convert_kv_block(&block); - assert_eq!(stored.seq_hash, SequenceHash::from(42i64)); - assert_eq!(stored.content_hash, compute_content_hash(&[1, 2, 3, 4])); - } - - #[test] - fn test_convert_kv_block_negative_hash() { - let block = KvBlock { - block_hash: -1, - token_ids: vec![10, 20], - block_size: 2, - lora_id: None, - cache_level: None, - }; - let stored = convert_kv_block(&block); - assert_eq!(stored.seq_hash, SequenceHash(u64::MAX)); - } - - #[test] - fn test_convert_kv_block_empty_tokens() { - let block = KvBlock { - block_hash: 100, - token_ids: vec![], - block_size: 0, - lora_id: None, - cache_level: None, - }; - let stored = convert_kv_block(&block); - assert_eq!(stored.seq_hash, SequenceHash::from(100i64)); - assert_eq!(stored.content_hash, compute_content_hash(&[])); - } - - // ----------------------------------------------------------------------- - // apply_event integration with PositionalIndexer - // ----------------------------------------------------------------------- - - #[test] - fn test_apply_stored_no_parent() { - let indexer = PositionalIndexer::new(64); - let w1 = indexer.intern_worker("http://w1:8000").unwrap(); - let mut wb = WorkerBlockMap::default(); - let stored = KvBlocksStored { - blocks: vec![ - KvBlock { - block_hash: 1, - token_ids: vec![10, 20, 30, 40], - block_size: 4, - lora_id: None, - cache_level: None, - }, - KvBlock { - block_hash: 2, - token_ids: vec![50, 60, 70, 80], - block_size: 4, - lora_id: None, - cache_level: None, - }, - ], - parent_block_hash: None, - }; - - KvEventMonitor::apply_stored(&stored, w1, &indexer, &mut wb); - assert_eq!(indexer.current_size(), 2); - } + use std::time::Duration; - #[test] - fn test_apply_stored_with_parent() { - let indexer = PositionalIndexer::new(64); - let w1 = indexer.intern_worker("http://w1:8000").unwrap(); - let mut wb = WorkerBlockMap::default(); - - let stored1 = KvBlocksStored { - blocks: vec![KvBlock { - block_hash: 1, - token_ids: vec![10, 20, 30, 40], - block_size: 4, - lora_id: None, - cache_level: None, - }], - parent_block_hash: None, - }; - KvEventMonitor::apply_stored(&stored1, w1, &indexer, &mut wb); - - let stored2 = KvBlocksStored { - blocks: vec![KvBlock { - block_hash: 2, - token_ids: vec![50, 60, 70, 80], - block_size: 4, - lora_id: None, - cache_level: None, - }], - parent_block_hash: Some(1), - }; - KvEventMonitor::apply_stored(&stored2, w1, &indexer, &mut wb); - assert_eq!(indexer.current_size(), 2); - } - - #[test] - fn test_apply_stored_fallback_on_worker_not_tracked() { - let indexer = PositionalIndexer::new(64); - let w1 = indexer.intern_worker("http://new-worker:8000").unwrap(); - let mut wb = WorkerBlockMap::default(); - - // Pass parent_block_hash for an untracked worker — should fallback to no parent. - let stored = KvBlocksStored { - blocks: vec![KvBlock { - block_hash: 1, - token_ids: vec![10, 20, 30, 40], - block_size: 4, - lora_id: None, - cache_level: None, - }], - parent_block_hash: Some(999), - }; - KvEventMonitor::apply_stored(&stored, w1, &indexer, &mut wb); - assert_eq!(indexer.current_size(), 1); - } - - #[test] - fn test_apply_removed() { - let indexer = PositionalIndexer::new(64); - let w1 = indexer.intern_worker("http://w1:8000").unwrap(); - let mut wb = WorkerBlockMap::default(); - - let stored = KvBlocksStored { - blocks: vec![ - KvBlock { - block_hash: 1, - token_ids: vec![10, 20, 30, 40], - block_size: 4, - lora_id: None, - cache_level: None, - }, - KvBlock { - block_hash: 2, - token_ids: vec![50, 60, 70, 80], - block_size: 4, - lora_id: None, - cache_level: None, - }, - ], - parent_block_hash: None, - }; - KvEventMonitor::apply_stored(&stored, w1, &indexer, &mut wb); - - let removed = KvBlocksRemoved { - block_hashes: vec![2], - cache_level: None, - }; - KvEventMonitor::apply_removed(&removed, w1, &indexer, &mut wb); - assert_eq!(indexer.current_size(), 1); - } + use kv_index::{ContentHash, SequenceHash, StoredBlock}; + use openai_protocol::worker::{ConnectionMode, HealthCheckConfig, RuntimeType, WorkerType}; - #[test] - fn test_apply_cleared_event() { - let indexer = PositionalIndexer::new(64); - let w1 = indexer.intern_worker("http://w1:8000").unwrap(); - let mut wb = WorkerBlockMap::default(); - - let stored = KvBlocksStored { - blocks: vec![KvBlock { - block_hash: 1, - token_ids: vec![10, 20, 30, 40], - block_size: 4, - lora_id: None, - cache_level: None, - }], - parent_block_hash: None, - }; - KvEventMonitor::apply_stored(&stored, w1, &indexer, &mut wb); - assert_eq!(indexer.current_size(), 1); + use super::{ + subscription::{INDEX_REMOVAL_PERMITS, MAX_CONCURRENT_INDEX_REMOVALS}, + *, + }; + use crate::worker::{kv_index_backend::WorkerBlocks, BasicWorkerBuilder}; - indexer.apply_cleared(w1, &mut wb); - assert_eq!(indexer.current_size(), 0); - } + async fn wait_until(what: &str, mut condition: impl FnMut() -> bool) { + for _ in 0..200 { + if condition() { + return; + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + panic!("timed out waiting until {what}"); + } + + /// A gRPC worker at a port that refuses connections at once: its task + /// loops on reconnects and answers shutdown, with no engine involved. + fn refusing_grpc_worker() -> Arc { + Arc::new( + BasicWorkerBuilder::new("grpc://127.0.0.1:1") + .worker_type(WorkerType::Regular) + .connection_mode(ConnectionMode::Grpc) + .runtime_type(RuntimeType::TokenSpeed) + .health_config(HealthCheckConfig { + disable_health_check: true, + ..Default::default() + }) + .build(), + ) + } + + /// A worker removed and added again under the same URL while the old + /// subscription's cleanup is still running: the index hands the URL its + /// old id until that cleanup has released it, so the re-add must wait + /// for the removal to finish rather than intern into the old id and + /// lose its blocks and its id to the cleanup. + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + async fn a_same_url_re_add_waits_for_the_old_subscriptions_cleanup() { + // The chain index releases a removed worker's id and name, which is + // what makes the window observable: the positional index keeps both. + let monitor = Arc::new(KvEventMonitor::with_kind(KvIndexKind::Chain, None)); + let worker = refusing_grpc_worker(); + let url = worker.url().to_string(); + monitor.on_worker_added(&worker).await; + let index = monitor + .get_indexer(UNKNOWN_MODEL_ID) + .expect("the model's index"); + wait_until("the subscription interns its worker", || { + index.worker_id(&url).is_some() + }) + .await; - #[test] - fn test_apply_event_dispatch_stored() { - let indexer = PositionalIndexer::new(64); - let w1 = indexer.intern_worker("http://w1:8000").unwrap(); - let mut wb = WorkerBlockMap::default(); - let event = KvCacheEvent { - event_id: 1, - data: Some(kv_cache_event::Data::Stored(KvBlocksStored { - blocks: vec![KvBlock { - block_hash: 42, - token_ids: vec![1, 2, 3, 4], - block_size: 4, - lora_id: None, - cache_level: None, - }], - parent_block_hash: None, - })), - }; + // Hold every cleanup permit: the old task's cleanup cannot finish. + let permits = INDEX_REMOVAL_PERMITS + .acquire_many(MAX_CONCURRENT_INDEX_REMOVALS as u32) + .await + .expect("the removal semaphore is open"); + #[expect( + clippy::disallowed_methods, + reason = "the removal and the re-add are awaited below" + )] + let removal = tokio::spawn({ + let monitor = Arc::clone(&monitor); + let url = url.clone(); + async move { monitor.on_worker_removed(&url).await } + }); + tokio::time::sleep(Duration::from_millis(200)).await; + assert!(!removal.is_finished(), "the removal waits for the cleanup"); + #[expect( + clippy::disallowed_methods, + reason = "the removal and the re-add are awaited below" + )] + let re_add = tokio::spawn({ + let monitor = Arc::clone(&monitor); + let worker = Arc::clone(&worker); + async move { monitor.on_worker_added(&worker).await } + }); + tokio::time::sleep(Duration::from_millis(200)).await; + assert!( + !re_add.is_finished(), + "the re-add waits for the old subscription's cleanup" + ); - KvEventMonitor::apply_event(&event, w1, &indexer, &mut wb); - assert_eq!(indexer.current_size(), 1); + drop(permits); + removal.await.expect("removal task"); + assert!( + index.worker_id(&url).is_none(), + "the removal returns once the old subscription has released its id" + ); + re_add.await.expect("re-add task"); + // The worker was the model's last, so the removal dropped the model's + // index and the re-add created a new one; the new subscription + // interned the URL there, after the release, with an id of its own. + let current = monitor + .get_indexer(UNKNOWN_MODEL_ID) + .expect("the re-added worker's model has an index"); + wait_until("the new subscription interns its worker", || { + current.worker_id(&url).is_some() + }) + .await; + monitor.stop().await; + assert!( + current.worker_id(&url).is_none(), + "stop takes the new subscription out of the index" + ); } - #[test] - fn test_apply_event_dispatch_removed() { - let indexer = PositionalIndexer::new(64); - let w1 = indexer.intern_worker("http://w1:8000").unwrap(); - let mut wb = WorkerBlockMap::default(); - - let stored_event = KvCacheEvent { - event_id: 1, - data: Some(kv_cache_event::Data::Stored(KvBlocksStored { - blocks: vec![KvBlock { - block_hash: 1, - token_ids: vec![1, 2, 3, 4], - block_size: 4, - lora_id: None, - cache_level: None, - }], - parent_block_hash: None, - })), - }; - KvEventMonitor::apply_event(&stored_event, w1, &indexer, &mut wb); - - let removed_event = KvCacheEvent { - event_id: 2, - data: Some(kv_cache_event::Data::Removed(KvBlocksRemoved { - block_hashes: vec![1], - cache_level: None, - })), - }; - KvEventMonitor::apply_event(&removed_event, w1, &indexer, &mut wb); - assert_eq!(indexer.current_size(), 0); - } + /// A caller may drop `on_worker_removed` before the old task has ended: + /// the registry's removal step runs under a timeout, and a cleanup + /// waiting on the removal permits can outlast it. The reservation the + /// removal left must not outlive the task: once the task is gone, the + /// model's last index is gone with it and an add of the URL goes + /// through, with nobody awaiting the removal. + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + async fn a_dropped_removal_leaves_no_reservation_behind() { + let monitor = Arc::new(KvEventMonitor::with_kind(KvIndexKind::Chain, None)); + let worker = refusing_grpc_worker(); + let url = worker.url().to_string(); + monitor.on_worker_added(&worker).await; + let index = monitor + .get_indexer(UNKNOWN_MODEL_ID) + .expect("the model's index"); + wait_until("the subscription interns its worker", || { + index.worker_id(&url).is_some() + }) + .await; - #[test] - fn test_apply_event_dispatch_cleared() { - let indexer = PositionalIndexer::new(64); - let w1 = indexer.intern_worker("http://w1:8000").unwrap(); - let mut wb = WorkerBlockMap::default(); - - KvEventMonitor::apply_event( - &KvCacheEvent { - event_id: 1, - data: Some(kv_cache_event::Data::Stored(KvBlocksStored { - blocks: vec![KvBlock { - block_hash: 1, - token_ids: vec![1, 2, 3, 4], - block_size: 4, - lora_id: None, - cache_level: None, - }], - parent_block_hash: None, - })), - }, - w1, - &indexer, - &mut wb, + // Hold every cleanup permit: the old task's cleanup cannot finish. + let permits = INDEX_REMOVAL_PERMITS + .acquire_many(MAX_CONCURRENT_INDEX_REMOVALS as u32) + .await + .expect("the removal semaphore is open"); + // The removal is dropped while it awaits the old task, as a timeout + // around it would drop it. + let removal = + tokio::time::timeout(Duration::from_millis(200), monitor.on_worker_removed(&url)).await; + assert!( + removal.is_err(), + "the removal was still waiting for the old task's cleanup" + ); + // The cleanup is still held: a re-add waits, and nothing has moved. + let re_add = + tokio::time::timeout(Duration::from_millis(200), monitor.on_worker_added(&worker)) + .await; + assert!( + re_add.is_err(), + "the re-add waits while the old task's cleanup is held" + ); + assert!( + monitor.get_indexer(UNKNOWN_MODEL_ID).is_some(), + "the model's index stays until the old task has finished" ); - // Clear - KvEventMonitor::apply_event( - &KvCacheEvent { - event_id: 2, - data: Some(kv_cache_event::Data::Cleared( - smg_grpc_client::common_proto::KvCacheCleared {}, - )), - }, - w1, - &indexer, - &mut wb, + drop(permits); + // Nobody awaits the removal any more: the old task finishes it on + // its own, releasing its id and dropping the model's last index. + wait_until("the old task drops the model's index", || { + monitor.get_indexer(UNKNOWN_MODEL_ID).is_none() + }) + .await; + assert!( + index.worker_id(&url).is_none(), + "the old task released its id" ); - assert_eq!(indexer.current_size(), 0); + tokio::time::timeout(Duration::from_secs(5), monitor.on_worker_added(&worker)) + .await + .expect("the re-add completes once the old task has ended"); + let current = monitor + .get_indexer(UNKNOWN_MODEL_ID) + .expect("the re-added worker's model has an index"); + wait_until("the new subscription interns its worker", || { + current.worker_id(&url).is_some() + }) + .await; + monitor.stop().await; } #[test] - fn test_apply_event_no_data() { - let indexer = PositionalIndexer::new(64); - let w1 = indexer.intern_worker("http://w1:8000").unwrap(); - let mut wb = WorkerBlockMap::default(); - let event = KvCacheEvent { - event_id: 1, - data: None, - }; - KvEventMonitor::apply_event(&event, w1, &indexer, &mut wb); - assert_eq!(indexer.current_size(), 0); - } - - // ----------------------------------------------------------------------- - // Lifecycle - // ----------------------------------------------------------------------- - - #[tokio::test] - async fn test_monitor_new() { + fn a_new_monitor_holds_no_index() { let monitor = KvEventMonitor::new(None); - assert!(!monitor.is_running().await); + assert!(monitor.indexers.is_empty()); } #[tokio::test] - async fn test_monitor_new_clamps_zero_jump_size() { + async fn a_zero_jump_size_is_clamped_to_one() { let monitor = KvEventMonitor::new(Some(0)); assert_eq!(monitor.jump_size, 1); } #[tokio::test] - async fn test_get_indexer_nonexistent() { + async fn an_unknown_model_has_no_index() { let monitor = KvEventMonitor::new(None); assert!(monitor.get_indexer("nonexistent").is_none()); } #[tokio::test] - async fn test_stop_empty_monitor() { + async fn stopping_an_empty_monitor_is_a_no_op() { let monitor = KvEventMonitor::new(None); monitor.stop().await; } #[tokio::test] - async fn test_on_worker_removed_nonexistent() { + async fn removing_an_unknown_worker_is_a_no_op() { let monitor = KvEventMonitor::new(None); monitor.on_worker_removed("http://nonexistent:8000").await; } - #[tokio::test] - async fn test_remove_indexer_worker_runs_cleanup_off_runtime() { - let indexer = Arc::new(PositionalIndexer::new(64)); - let worker_id = indexer.intern_worker("http://w1:8000").unwrap(); - let mut worker_blocks = WorkerBlockMap::default(); - indexer - .apply_stored( - worker_id, - &[StoredBlock { - seq_hash: SequenceHash(1), - content_hash: compute_content_hash(&[1, 2, 3]), - }], - None, - &mut worker_blocks, - ) - .unwrap(); - - KvEventMonitor::remove_indexer_worker(Arc::clone(&indexer), worker_id, worker_blocks).await; - - assert_eq!(indexer.current_size(), 0); - } - - // ----------------------------------------------------------------------- - // block_size learning - // ----------------------------------------------------------------------- - #[test] - fn test_set_block_size() { + fn set_block_size_keeps_the_first_value() { let monitor = KvEventMonitor::new(None); // Initially no block_size @@ -1177,7 +867,7 @@ mod tests { } #[tokio::test] - async fn test_stop_clears_block_sizes() { + async fn stop_forgets_the_block_sizes() { let monitor = KvEventMonitor::new(None); monitor.set_block_size("llama", 16); assert_eq!(monitor.block_size("llama"), Some(16)); @@ -1187,21 +877,21 @@ mod tests { } #[test] - fn test_prune_all_enforces_capacity_ceiling() { + fn prune_all_enforces_the_capacity_ceiling_after_the_grace() { let monitor = KvEventMonitor::new(None); let indexer = monitor .indexers .entry("llama".to_string()) - .or_insert_with(|| Arc::new(PositionalIndexer::new(64))) + .or_insert_with(|| Arc::new(KvIndex::positional(64))) .clone(); let worker = indexer.intern_worker("http://w1:8000").unwrap(); - let mut worker_blocks = WorkerBlockMap::default(); + let mut worker_blocks = WorkerBlocks::default(); // Ten independent single-block chains → ten index entries. for i in 0u64..10 { let block = StoredBlock { seq_hash: SequenceHash(1000 + i), - content_hash: kv_index::ContentHash(2000 + i), + content_hash: ContentHash(2000 + i), }; indexer .apply_stored(worker, &[block], None, &mut worker_blocks) @@ -1227,7 +917,7 @@ mod tests { } #[tokio::test] - async fn test_start_prune_task_noop_when_disabled() { + async fn the_prune_task_starts_only_with_a_bound() { let monitor = Arc::new(KvEventMonitor::new(None)); monitor.start_prune_task(0, 0); assert!(monitor.prune_task.lock().is_none()); @@ -1235,4 +925,83 @@ mod tests { monitor.start_prune_task(60, 0); assert!(monitor.prune_task.lock().is_some()); } + + /// The index-shape gauges come from the indexes' own counters, by model: + /// memberships and entries for both kinds, the chain index's runs, blocks, + /// arena and slab bytes and moved hashes for the chain kind. + #[test] + fn index_shape_gauges_follow_the_indexes() { + use metrics_exporter_prometheus::{PrometheusBuilder, PrometheusHandle}; + + fn gauge(handle: &PrometheusHandle, name: &str, model: &str) -> Option { + let prefix = format!("{name}{{model=\"{model}\"}}"); + handle + .render() + .lines() + .find(|line| line.starts_with(&prefix)) + .and_then(|line| line.rsplit(' ').next()) + .and_then(|value| value.parse().ok()) + } + + let recorder = PrometheusBuilder::new().build_recorder(); + let handle = recorder.handle(); + metrics::with_local_recorder(&recorder, || { + let monitor = KvEventMonitor::new(None); + for (model, kind) in [ + ("pos", KvIndexKind::Positional), + ("chain", KvIndexKind::Chain), + ] { + let index = Arc::new(KvIndex::new(kind, 8)); + let worker = index.intern_worker("grpc://w1:9000").unwrap(); + let mut held = WorkerBlocks::default(); + let blocks: Vec = (1..=3u64) + .map(|i| StoredBlock { + seq_hash: SequenceHash(i), + content_hash: ContentHash(100 + i), + }) + .collect(); + index + .apply_stored(worker, &blocks, None, &mut held) + .unwrap(); + monitor.indexers.insert(model.to_string(), index); + } + KvEventMonitor::publish_stats(&monitor.indexers); + for model in ["pos", "chain"] { + assert_eq!( + gauge(&handle, "smg_kv_index_memberships", model), + Some(3.0), + "{model}" + ); + assert_eq!( + gauge(&handle, "smg_kv_index_entries", model), + Some(3.0), + "{model}" + ); + } + assert_eq!(gauge(&handle, "smg_kv_index_runs_live", "chain"), Some(1.0)); + assert_eq!( + gauge(&handle, "smg_kv_index_blocks_live", "chain"), + Some(3.0) + ); + assert_eq!( + gauge(&handle, "smg_kv_index_moved_hashes", "chain"), + Some(0.0) + ); + assert_eq!( + gauge(&handle, "smg_kv_index_engine_conflicts", "chain"), + Some(0.0) + ); + assert!(gauge(&handle, "smg_kv_index_arena_bytes", "chain").is_some_and(|b| b > 0.0)); + assert!(gauge(&handle, "smg_kv_index_slab_bytes", "chain").is_some_and(|b| b > 0.0)); + assert_eq!(gauge(&handle, "smg_kv_index_runs_live", "pos"), None); + }); + } + + #[test] + fn start_stats_task_holds_the_task() { + let monitor = KvEventMonitor::new(None); + assert!(monitor.stats_task.lock().is_none()); + monitor.start_stats_task(); + assert!(monitor.stats_task.lock().is_some()); + } } diff --git a/model_gateway/src/worker/kv_event_monitor/admission.rs b/model_gateway/src/worker/kv_event_monitor/admission.rs new file mode 100644 index 0000000000..ebcb32a7c0 --- /dev/null +++ b/model_gateway/src/worker/kv_event_monitor/admission.rs @@ -0,0 +1,1214 @@ +//! One stream's state: the admission cursor per data-parallel rank, gap +//! recovery, relay snapshots, and the metrics around an applied batch. + +use std::{ + collections::HashMap, + time::{Instant, SystemTime, UNIX_EPOCH}, +}; + +use smg_grpc_client::common_proto::{kv_cache_event, KvEventBatch, KvSnapshotChunk}; +use tracing::{debug, info, warn}; + +use super::{ + apply::{WorkerIndexCounters, WorkerIndexState}, + KvEventMonitor, +}; +use crate::{ + observability::metrics::Metrics, + worker::{ + kv_event_recovery::{Admission, RankState, ResyncReason}, + kv_index_backend::KvIndex, + }, +}; + +/// A worker's subscription state: one admission cursor per data-parallel +/// rank (every publisher numbers its own batches) over the worker's one +/// index state, whose copies are pooled across ranks (see +/// [`WorkerIndexState`]). Cursors and copy counts never mix. +#[derive(Default)] +pub(super) struct WorkerStreamState { + pub(super) ranks: HashMap, + pub(super) index: WorkerIndexState, + /// A relay snapshot whose chunks are still arriving on the stream. + pub(super) snapshot: Option, +} + +/// Where an in-band relay snapshot stands (`KvSnapshotChunk`). +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub(super) struct SnapshotProgress { + count: u32, + applied: u32, + blocks: u64, +} + +impl WorkerStreamState { + /// The cursor to send when resubscribing: rank 0's last applied sequence + /// (the servicers replay rank 0's publisher). Other ranks keep their own + /// cursors and dedup what arrives. + pub(super) fn resume_sequence(&self) -> u64 { + self.ranks.get(&0).map_or(0, RankState::resume_from) + } + + /// A new stream connection was made: every rank's next batch is its + /// first on it. + pub(super) fn reconnected(&mut self) { + for rank in self.ranks.values_mut() { + rank.reconnected(); + } + } + + pub(super) fn degraded_ranks(&self) -> usize { + self.ranks + .values() + .filter(|rank| rank.is_degraded()) + .count() + } + + /// The stream ended with a snapshot still arriving: what was applied is + /// a partial live set, so the next subscription asks from zero (and gets + /// a whole snapshot) instead of resuming after the last chunk's stamp. + /// Returns whether a snapshot was abandoned. + pub(super) fn abandon_snapshot(&mut self) -> bool { + let Some(progress) = self.snapshot.take() else { + return false; + }; + for cursor in self.ranks.values_mut() { + cursor.reset(); + } + debug_assert!(progress.applied < progress.count); + true + } +} + +/// What `admit_batch` did with a batch. +#[derive(Debug, PartialEq, Eq)] +pub(super) enum BatchOutcome { + /// Applied to the index. + Applied, + /// Dropped (duplicate) or held (snapshot tail). + Skipped, + /// A gap: the caller reconnects asking for a replay from `expected`. + Gap { expected: u64, received: u64 }, +} + +impl KvEventMonitor { + /// Run one batch through its rank's admission cursor and apply it to the + /// worker's index state. + pub(super) fn admit_batch( + batch: &KvEventBatch, + worker_url: &str, + worker_id: u32, + indexer: &KvIndex, + state: &mut WorkerStreamState, + on_batch: &mut impl FnMut(&KvEventBatch), + ) -> BatchOutcome { + if let Some(chunk) = &batch.snapshot { + return Self::admit_snapshot_chunk( + batch, chunk, worker_url, worker_id, indexer, state, on_batch, + ); + } + let rank = batch.dp_rank.unwrap_or(0); + let seq = batch.sequence_number; + let clears = batch.events.iter().any(|event| { + matches!(&event.data, Some(kv_cache_event::Data::Cleared(cleared)) if cleared.ownership.is_none()) + }); + let admission = state.ranks.entry(rank).or_default().admit(seq, clears); + let mut degraded_changed = false; + match admission { + Admission::Apply => {} + Admission::Stale => { + debug!(worker_url = %worker_url, rank, received = seq, "Skipping stale KV event batch"); + Metrics::record_kv_event_batch(worker_url, "stale"); + return BatchOutcome::Skipped; + } + Admission::Restart => { + warn!( + worker_url = %worker_url, + rank, + received = seq, + "KV event publisher restarted; clearing the worker's index state" + ); + Self::clear_worker(worker_id, indexer, state, rank); + Metrics::record_kv_event_resync( + worker_url, + ResyncReason::PublisherRestart.as_str(), + ); + degraded_changed = true; + } + Admission::Replay { expected } => { + Metrics::record_kv_event_gap(worker_url, "replay_requested", seq - expected); + return BatchOutcome::Gap { + expected, + received: seq, + }; + } + Admission::Unrecovered { missed, cleared } => { + warn!( + worker_url = %worker_url, + rank, + missed, + cleared, + "KV event gap could not be replayed; continuing from the live stream" + ); + if cleared { + Self::clear_worker(worker_id, indexer, state, rank); + Metrics::record_kv_event_resync(worker_url, ResyncReason::GapCleared.as_str()); + } + Metrics::record_kv_event_gap( + worker_url, + if cleared { + "unrecovered_cleared" + } else { + "unrecovered_kept" + }, + missed, + ); + degraded_changed = true; + } + Admission::Buffered => { + if let Some(cursor) = state.ranks.get_mut(&rank) { + if !cursor.buffer_live(batch.clone()) { + Metrics::record_kv_event_batch(worker_url, "tail_overflow"); + } + Metrics::set_kv_event_tail_depth(worker_url, cursor.tail_len()); + } + return BatchOutcome::Skipped; + } + } + + on_batch(batch); + let started = Instant::now(); + let counters_before = state.index.counters; + let (mut stored_blocks, mut removed_blocks) = (0usize, 0usize); + for event in &batch.events { + match &event.data { + Some(kv_cache_event::Data::Stored(stored)) => stored_blocks += stored.blocks.len(), + Some(kv_cache_event::Data::Removed(removed)) => { + removed_blocks += removed.block_hashes.len(); + } + _ => {} + } + Self::apply_event(event, worker_id, indexer, &mut state.index); + } + Metrics::record_kv_event_apply(worker_url, started.elapsed().as_secs_f64()); + Self::record_parentless(worker_url, &state.index.counters, &counters_before); + if stored_blocks > 0 { + Metrics::record_kv_event_blocks(worker_url, "stored", stored_blocks); + } + if removed_blocks > 0 { + Metrics::record_kv_event_blocks(worker_url, "removed", removed_blocks); + } + Self::record_lag(worker_url, batch.timestamp); + Metrics::record_kv_event_batch(worker_url, "applied"); + Metrics::set_kv_index_blocks(worker_url, indexer.worker_block_count(worker_id)); + if degraded_changed { + Metrics::set_kv_event_degraded_ranks(worker_url, state.degraded_ranks()); + } + BatchOutcome::Applied + } + + /// A chunk of a relay state snapshot (`KvSnapshotChunk`): the live set the + /// relay recorded, replacing the worker's state. Chunk 0 clears the worker + /// and every cursor and counts a `snapshot` resync; every chunk is applied + /// outside the admission rules and moves its rank's cursor to its stamp, + /// which the relay chose so that live events continue after the last one. + fn admit_snapshot_chunk( + batch: &KvEventBatch, + chunk: &KvSnapshotChunk, + worker_url: &str, + worker_id: u32, + indexer: &KvIndex, + state: &mut WorkerStreamState, + on_batch: &mut impl FnMut(&KvEventBatch), + ) -> BatchOutcome { + let rank = batch.dp_rank.unwrap_or(0); + if chunk.index == 0 || state.snapshot.is_none() { + if chunk.index != 0 { + warn!( + worker_url = %worker_url, + rank, + index = chunk.index, + "KV event relay snapshot arrived without its first chunk; taking it as a resync" + ); + } + info!( + worker_url = %worker_url, + rank, + chunks = chunk.count, + blocks = chunk.blocks, + unknown_before = chunk.unknown_before, + through = batch.sequence_number + u64::from(chunk.count.saturating_sub(chunk.index + 1)), + "KV event relay served a state snapshot; replacing the worker's index state" + ); + if chunk.unknown_before > 0 { + warn!( + worker_url = %worker_url, + rank, + unknown_before = chunk.unknown_before, + "KV event relay snapshot starts late: the engine's blocks from before the \ + relay's record are unknown; the worker's index is partial until they leave \ + the engine (rank degraded)" + ); + } + Self::apply_cleared(worker_id, indexer, &mut state.index); + for cursor in state.ranks.values_mut() { + cursor.reset(); + } + Metrics::record_kv_event_resync(worker_url, ResyncReason::Snapshot.as_str()); + state.snapshot = Some(SnapshotProgress { + count: chunk.count.max(1), + applied: 0, + blocks: chunk.blocks, + }); + } + let cursor = state.ranks.entry(rank).or_default(); + cursor.resync_to(batch.sequence_number); + if chunk.unknown_before > 0 { + cursor.mark_degraded(); + } + on_batch(batch); + let counters_before = state.index.counters; + for event in &batch.events { + Self::apply_event(event, worker_id, indexer, &mut state.index); + } + Self::record_parentless(worker_url, &state.index.counters, &counters_before); + Self::record_lag(worker_url, batch.timestamp); + Metrics::record_kv_event_batch(worker_url, "snapshot"); + if let Some(progress) = &mut state.snapshot { + progress.applied += 1; + if progress.applied >= progress.count { + info!( + worker_url = %worker_url, + chunks = progress.count, + blocks = progress.blocks, + through = batch.sequence_number, + "KV event relay snapshot applied; continuing with live events" + ); + state.snapshot = None; + Metrics::set_kv_event_degraded_ranks(worker_url, state.degraded_ranks()); + } + } + BatchOutcome::Applied + } + + /// One rank's publisher lost its history (a restart, or an unreplayable + /// gap too large to keep): the worker's pooled index state goes with it, + /// and the other ranks' cursors start over so their next batch is taken + /// as a first one instead of clearing the index a second time. + fn clear_worker(worker_id: u32, indexer: &KvIndex, state: &mut WorkerStreamState, rank: i32) { + Self::apply_cleared(worker_id, indexer, &mut state.index); + for (other, cursor) in &mut state.ranks { + if *other != rank { + cursor.reset(); + } + } + } + + /// Publish the parent-less stores a batch produced, if any. + fn record_parentless( + worker_url: &str, + after: &WorkerIndexCounters, + before: &WorkerIndexCounters, + ) { + let (stores, blocks) = after.parentless_since(before); + if stores > 0 { + Metrics::record_kv_event_parentless(worker_url, stores, blocks); + } + } + + /// The server declared its history gone: drop the worker's index state + /// and every cursor, so the next stream is taken from wherever it starts. + pub(super) fn reset_worker( + indexer: &KvIndex, + worker_id: u32, + state: &mut WorkerStreamState, + worker_url: &str, + reason: ResyncReason, + ) { + Self::apply_cleared(worker_id, indexer, &mut state.index); + for cursor in state.ranks.values_mut() { + cursor.reset(); + } + state.snapshot = None; + Metrics::record_kv_event_resync(worker_url, reason.as_str()); + Metrics::set_kv_event_degraded_ranks(worker_url, 0); + Metrics::set_kv_index_blocks(worker_url, indexer.worker_block_count(worker_id)); + } + + /// Age of a batch when applied, from the publisher's wall-clock stamp. + fn record_lag(worker_url: &str, published_at: f64) { + if published_at <= 0.0 { + return; + } + let now = SystemTime::now() + .duration_since(UNIX_EPOCH) + .map(|d| d.as_secs_f64()) + .unwrap_or(0.0); + let lag = now - published_at; + if lag.is_finite() && lag >= 0.0 { + Metrics::record_kv_event_lag(worker_url, lag); + } + } +} + +#[cfg(test)] +mod tests { + use std::collections::BTreeSet; + + use kv_index::{ + salt::namespace_seed, ContentHash, ReferenceIndexer, SequenceHash, StoredBlock, + }; + use smg_grpc_client::common_proto::{ + EngineLoad, KvBlock, KvBlocksRemoved, KvBlocksStored, KvCacheCleared, KvCacheEvent, + KvCacheTier, + }; + + use super::{super::apply::convert_kv_block, *}; + use crate::worker::{ + kv_event_recovery::{Cursor, RESTART_WINDOW}, + kv_index_backend::KvIndexKind, + }; + + /// Token ids for engine block `id`: distinct content per id. + fn tokens_for(id: i64) -> Vec { + (0..4u32).map(|i| (id as u32) * 16 + i).collect() + } + + fn kv_block(id: i64) -> KvBlock { + KvBlock { + block_hash: id, + token_ids: tokens_for(id), + block_size: 4, + ..Default::default() + } + } + + fn stored(parent: Option, ids: &[i64]) -> KvCacheEvent { + KvCacheEvent { + event_id: 0, + data: Some(kv_cache_event::Data::Stored(KvBlocksStored { + blocks: ids.iter().map(|&id| kv_block(id)).collect(), + parent_block_hash: parent, + ..Default::default() + })), + } + } + + fn removed(ids: &[i64]) -> KvCacheEvent { + KvCacheEvent { + event_id: 0, + data: Some(kv_cache_event::Data::Removed(KvBlocksRemoved { + block_hashes: ids.to_vec(), + ..Default::default() + })), + } + } + + fn cleared() -> KvCacheEvent { + KvCacheEvent { + event_id: 0, + data: Some(kv_cache_event::Data::Cleared(KvCacheCleared::default())), + } + } + + fn batch(seq: u64, rank: Option, events: Vec) -> KvEventBatch { + KvEventBatch { + sequence_number: seq, + timestamp: 0.0, + events, + dp_rank: rank, + snapshot: None, + load: None, + } + } + + /// `smg_kv_index_blocks{worker}` is set where applied batches are counted, + /// from the index's own per-worker counter: it follows stores, removals and + /// a clear, and costs the lookup path nothing. Both indexes keep that + /// counter, so the gauge reads the same under `--kv-index chain`. + #[test] + fn index_block_gauge_follows_stores_removals_and_a_clear() { + for kind in [KvIndexKind::Positional, KvIndexKind::Chain] { + index_block_gauge_follows_the_index(kind); + } + } + + fn index_block_gauge_follows_the_index(kind: KvIndexKind) { + use metrics_exporter_prometheus::{PrometheusBuilder, PrometheusHandle}; + + fn index_blocks(handle: &PrometheusHandle) -> Option { + handle + .render() + .lines() + .find(|line| line.starts_with("smg_kv_index_blocks{worker=\"grpc://w1:9000\"}")) + .and_then(|line| line.rsplit(' ').next()) + .and_then(|value| value.parse().ok()) + } + + let recorder = PrometheusBuilder::new().build_recorder(); + let handle = recorder.handle(); + metrics::with_local_recorder(&recorder, || { + let mut sim = Sim::with_kind(kind); + assert_eq!(index_blocks(&handle), None, "nothing applied yet"); + + assert_eq!( + sim.feed(&batch(1, None, vec![stored(None, &[1, 2, 3])])), + BatchOutcome::Applied + ); + assert_eq!(sim.indexer.worker_block_count(sim.worker), 3); + assert_eq!( + index_blocks(&handle), + Some(3.0), + "{kind:?}: three blocks stored" + ); + + assert_eq!( + sim.feed(&batch(2, None, vec![removed(&[3])])), + BatchOutcome::Applied + ); + assert_eq!(index_blocks(&handle), Some(2.0), "{kind:?}: one removed"); + + assert_eq!( + sim.feed(&batch(3, None, vec![stored(Some(2), &[4, 5])])), + BatchOutcome::Applied + ); + assert_eq!( + index_blocks(&handle), + Some(4.0), + "{kind:?}: two more stored" + ); + + assert_eq!( + sim.feed(&batch(4, None, vec![cleared()])), + BatchOutcome::Applied + ); + assert_eq!(sim.indexer.worker_block_count(sim.worker), 0); + assert_eq!(index_blocks(&handle), Some(0.0), "{kind:?}: cleared"); + }); + } + + /// The production subscriber's per-worker state next to a reference that + /// sees the stream the subscriber *should* have applied. + struct Sim { + indexer: KvIndex, + worker: u32, + state: WorkerStreamState, + reference: ReferenceIndexer, + } + + impl Sim { + fn new() -> Self { + Self::with_kind(KvIndexKind::Positional) + } + + /// The same harness over the index selected by `--kv-index`. + fn with_kind(kind: KvIndexKind) -> Self { + let indexer = KvIndex::new(kind, 8); + let worker = indexer.intern_worker("grpc://w1:9000").unwrap(); + Self { + indexer, + worker, + state: WorkerStreamState::default(), + reference: ReferenceIndexer::new(), + } + } + + /// Feed a batch through the real admission path. + fn feed(&mut self, b: &KvEventBatch) -> BatchOutcome { + KvEventMonitor::admit_batch( + b, + "grpc://w1:9000", + self.worker, + &self.indexer, + &mut self.state, + &mut |_: &KvEventBatch| {}, + ) + } + + /// Apply a batch to the reference with the subscriber's semantics + /// (a store whose parent is unknown starts a new chain). + fn reference_apply(&mut self, b: &KvEventBatch) { + let seed = namespace_seed(None, None); + for event in &b.events { + match event.data.as_ref().unwrap() { + kv_cache_event::Data::Stored(st) => { + let blocks: Vec = st + .blocks + .iter() + .map(|block| convert_kv_block(block, seed)) + .collect(); + let parent = st.parent_block_hash.map(SequenceHash::from); + if self + .reference + .apply_stored(self.worker, &blocks, parent) + .is_err() + { + self.reference + .apply_stored(self.worker, &blocks, None) + .unwrap(); + } + } + kv_cache_event::Data::Removed(rm) => { + let hashes: Vec = rm + .block_hashes + .iter() + .map(|&h| SequenceHash::from(h)) + .collect(); + self.reference.apply_removed(self.worker, &hashes); + } + kv_cache_event::Data::Cleared(_) => self.reference.apply_cleared(self.worker), + } + } + } + + /// Feed to both: what a correctly delivered batch does. + fn deliver(&mut self, b: &KvEventBatch) -> BatchOutcome { + self.reference_apply(b); + self.feed(b) + } + + /// Apply a batch straight into the worker's index state, bypassing + /// admission (how a snapshot or a buffered tail lands). + fn apply_direct(&mut self, b: &KvEventBatch) { + for event in &b.events { + KvEventMonitor::apply_event( + event, + self.worker, + &self.indexer, + &mut self.state.index, + ); + } + } + + fn assert_matches_reference(&self) { + let production: BTreeSet<(u32, usize, ContentHash, SequenceHash)> = + self.indexer.debug_blocks().into_iter().collect(); + assert_eq!(production, self.reference.blocks(), "index content"); + for query in self.queries() { + let scores = self.indexer.find_matches(&query, false).scores; + let expected = self.reference.find_matches(&query); + let got: Vec<(u32, u32)> = scores.into_iter().collect(); + let want: Vec<(u32, u32)> = expected.into_iter().filter(|(_, s)| *s > 0).collect(); + assert_eq!(got, want, "scores for {query:?}"); + } + } + + /// Lookups: every stored chain, plus a mutated copy of each. + fn queries(&self) -> Vec> { + let mut chains: Vec> = Vec::new(); + let blocks = self.reference.blocks(); + let max_pos = blocks.iter().map(|b| b.1).max().unwrap_or(0); + // Rebuild chains by walking positions from the reference content. + let mut by_pos: Vec> = vec![Vec::new(); max_pos + 1]; + for (_, pos, content, _) in &blocks { + if !by_pos[*pos].contains(content) { + by_pos[*pos].push(*content); + } + } + let mut chain = Vec::new(); + for level in &by_pos { + if let Some(c) = level.first() { + chain.push(*c); + chains.push(chain.clone()); + } + } + if let Some(full) = chains.last().cloned() { + let mut mutated = full.clone(); + if mutated.len() > 1 { + mutated[1] = ContentHash(0xDEAD_BEEF); + chains.push(mutated); + } + } + chains.push(vec![ContentHash(1), ContentHash(2)]); + chains + } + } + + #[test] + fn recovery_in_order_stream_matches_reference() { + let mut sim = Sim::new(); + assert_eq!( + sim.deliver(&batch(1, None, vec![stored(None, &[1, 2, 3])])), + BatchOutcome::Applied + ); + assert_eq!( + sim.deliver(&batch(2, None, vec![stored(Some(3), &[4, 5])])), + BatchOutcome::Applied + ); + assert_eq!( + sim.deliver(&batch(3, None, vec![removed(&[5])])), + BatchOutcome::Applied + ); + assert_eq!( + sim.deliver(&batch(4, None, vec![stored(Some(4), &[6])])), + BatchOutcome::Applied + ); + sim.assert_matches_reference(); + } + + #[test] + fn recovery_stream_numbered_from_zero_matches_reference() { + // vLLM and SGLang publishers count from 0; 0 must not be treated as + // "no cursor" once a batch carried it. + let mut sim = Sim::new(); + assert_eq!( + sim.deliver(&batch(0, None, vec![cleared(), stored(None, &[1, 2])])), + BatchOutcome::Applied + ); + assert_eq!( + sim.deliver(&batch(1, None, vec![stored(Some(2), &[3])])), + BatchOutcome::Applied + ); + assert_eq!( + sim.feed(&batch(1, None, vec![removed(&[1])])), + BatchOutcome::Skipped + ); + assert_eq!(sim.state.resume_sequence(), 1); + sim.assert_matches_reference(); + } + + #[test] + fn recovery_gap_filled_by_replay_matches_reference() { + let mut sim = Sim::new(); + sim.deliver(&batch(1, None, vec![stored(None, &[1, 2])])); + sim.deliver(&batch(2, None, vec![stored(Some(2), &[3])])); + // Batch 3 is lost on the wire; 4 arrives: one replay is asked for. + let b3 = batch(3, None, vec![removed(&[3])]); + let b4 = batch(4, None, vec![stored(Some(2), &[7])]); + assert_eq!( + sim.feed(&b4), + BatchOutcome::Gap { + expected: 3, + received: 4 + } + ); + assert_eq!(sim.state.resume_sequence(), 2); + // Resuming after 2, the server replays 3 and continues live. + assert_eq!(sim.deliver(&b3), BatchOutcome::Applied); + assert_eq!(sim.deliver(&b4), BatchOutcome::Applied); + assert_eq!( + sim.deliver(&batch(5, None, vec![stored(Some(7), &[8])])), + BatchOutcome::Applied + ); + assert!(!sim.state.ranks[&0].is_degraded()); + sim.assert_matches_reference(); + } + + #[test] + fn recovery_gap_without_replay_keeps_state_and_continues() { + let mut sim = Sim::new(); + sim.deliver(&batch(1, None, vec![stored(None, &[1, 2])])); + sim.deliver(&batch(2, None, vec![stored(Some(2), &[3])])); + // Batches 3 and 4 are lost; the server has no history (the Rust + // relay): after the replay request it streams 5 again. + let b5 = batch(5, None, vec![stored(Some(3), &[9])]); + assert_eq!( + sim.feed(&b5), + BatchOutcome::Gap { + expected: 3, + received: 5 + } + ); + assert_eq!(sim.deliver(&b5), BatchOutcome::Applied); + assert!(sim.state.ranks[&0].is_degraded()); + // Everything after keeps applying; no reconnect loop. + assert_eq!( + sim.deliver(&batch(6, None, vec![stored(Some(9), &[10])])), + BatchOutcome::Applied + ); + // The reference saw the same stream minus the lost batches. + sim.assert_matches_reference(); + // A later gap starts a fresh single replay attempt. + assert_eq!( + sim.feed(&batch(8, None, vec![])), + BatchOutcome::Gap { + expected: 7, + received: 8 + } + ); + } + + #[test] + fn recovery_duplicates_and_out_of_order_batches_are_skipped() { + let mut sim = Sim::new(); + let b1 = batch(1, None, vec![stored(None, &[1, 2])]); + let b2 = batch(2, None, vec![stored(Some(2), &[3])]); + let b3 = batch(3, None, vec![removed(&[3])]); + let b4 = batch(4, None, vec![stored(Some(2), &[4])]); + sim.deliver(&b1); + sim.deliver(&b2); + sim.deliver(&b3); + // Replayed overlap after a reconnect: already applied, must not re-apply. + // (A duplicate of sequence 1 would be taken as a new publisher: see + // `a_counter_at_its_start_below_the_cursor_is_a_restart`.) + assert_eq!(sim.feed(&b2), BatchOutcome::Skipped); + assert_eq!(sim.feed(&b3), BatchOutcome::Skipped); + sim.deliver(&b4); + sim.assert_matches_reference(); + } + + #[test] + fn recovery_publisher_restart_clears_the_rank() { + let mut sim = Sim::new(); + sim.deliver(&batch(1, None, vec![stored(None, &[1, 2, 3])])); + for seq in 2..=(RESTART_WINDOW + 5) { + sim.deliver(&batch(seq, None, vec![])); + } + // The engine restarts: its cache is empty and it counts from 0 again. + let fresh = batch(0, None, vec![stored(None, &[21, 22])]); + sim.reference.apply_cleared(sim.worker); + sim.reference_apply(&fresh); + assert_eq!(sim.feed(&fresh), BatchOutcome::Applied); + assert_eq!( + sim.deliver(&batch(1, None, vec![stored(Some(22), &[23])])), + BatchOutcome::Applied + ); + assert!(!sim.state.ranks[&0].is_degraded()); + sim.assert_matches_reference(); + } + + #[test] + fn recovery_restart_seen_on_a_fresh_connection_clears_the_rank() { + let mut sim = Sim::new(); + sim.deliver(&batch(1, None, vec![stored(None, &[1, 2])])); + sim.deliver(&batch(2, None, vec![stored(Some(2), &[3])])); + // The worker died and came back: the stream reconnected and the new + // publisher counts from 1 with an empty cache (the relay and the mock + // engine stream live; neither can replay). + sim.state.ranks.get_mut(&0).unwrap().reconnected(); + let fresh = batch(1, None, vec![stored(None, &[7])]); + sim.reference.apply_cleared(sim.worker); + sim.reference_apply(&fresh); + assert_eq!(sim.feed(&fresh), BatchOutcome::Applied); + sim.assert_matches_reference(); + assert_eq!(sim.state.resume_sequence(), 1); + } + + #[test] + fn recovery_clear_below_the_cursor_is_a_restart() { + let mut sim = Sim::new(); + sim.deliver(&batch(1, None, vec![stored(None, &[1, 2])])); + sim.deliver(&batch(2, None, vec![stored(Some(2), &[3])])); + // SGLang's first batch after a restart carries AllBlocksCleared and + // its counter starts over, with the servicer's stream still up. + let first = batch(0, None, vec![cleared(), stored(None, &[9])]); + sim.reference.apply_cleared(sim.worker); + sim.reference_apply(&first); + assert_eq!(sim.feed(&first), BatchOutcome::Applied); + sim.assert_matches_reference(); + // A residency agent's clear below the cursor is just a duplicate. + let agent = KvCacheEvent { + event_id: 0, + data: Some(kv_cache_event::Data::Cleared(KvCacheCleared { + ownership: Some("kvcr".to_string()), + })), + }; + assert_eq!( + sim.feed(&batch(0, None, vec![agent])), + BatchOutcome::Skipped + ); + sim.assert_matches_reference(); + } + + #[test] + fn recovery_cleared_event_matches_reference() { + let mut sim = Sim::new(); + sim.deliver(&batch(1, None, vec![stored(None, &[1, 2, 3])])); + sim.deliver(&batch(2, None, vec![cleared()])); + sim.deliver(&batch(3, None, vec![stored(None, &[4])])); + sim.assert_matches_reference(); + assert_eq!(sim.reference.worker_block_count(sim.worker), 1); + } + + #[test] + fn recovery_dp_ranks_keep_independent_cursors() { + let mut sim = Sim::new(); + // Two publishers, each numbering from 1, interleaved on one stream. + sim.deliver(&batch(1, Some(0), vec![stored(None, &[1, 2])])); + sim.deliver(&batch(1, Some(1), vec![stored(None, &[101, 102])])); + sim.deliver(&batch(2, Some(1), vec![stored(Some(102), &[103])])); + sim.deliver(&batch(2, Some(0), vec![stored(Some(2), &[3])])); + assert_eq!( + sim.deliver(&batch(3, Some(0), vec![removed(&[3])])), + BatchOutcome::Applied + ); + // Each rank dedups against its own cursor. + assert_eq!( + sim.feed(&batch(2, Some(1), vec![stored(None, &[200])])), + BatchOutcome::Skipped + ); + assert_eq!( + sim.deliver(&batch(3, Some(1), vec![stored(Some(103), &[104])])), + BatchOutcome::Applied + ); + assert_eq!(sim.state.ranks.len(), 2); + assert_eq!(sim.state.degraded_ranks(), 0); + assert_eq!(sim.state.resume_sequence(), 3); + sim.assert_matches_reference(); + // Rank 1's publisher restarts: the worker's pooled index state is + // cleared and rank 0's cursor starts over, so its next batch is taken + // as a first one rather than clearing the index again. + sim.state.reconnected(); + let fresh = batch(1, Some(1), vec![stored(None, &[111])]); + sim.reference.apply_cleared(sim.worker); + sim.reference_apply(&fresh); + assert_eq!(sim.feed(&fresh), BatchOutcome::Applied); + assert_eq!( + sim.deliver(&batch(4, Some(0), vec![stored(None, &[5])])), + BatchOutcome::Applied + ); + sim.assert_matches_reference(); + assert_eq!(sim.state.resume_sequence(), 4); + } + + #[test] + fn recovery_resync_forgets_copy_counts() { + // Copy counts live in the pooled index state and cursors in the rank + // state; a restart clears the counts with the blocks, so no stale + // host copy keeps a block routable afterwards. + let mut sim = Sim::new(); + let mut on_host = stored(None, &[1, 2]); + if let Some(kv_cache_event::Data::Stored(st)) = on_host.data.as_mut() { + st.tier = Some(KvCacheTier::Host as i32); + } + sim.deliver(&batch(1, None, vec![stored(None, &[1, 2])])); + sim.deliver(&batch(2, None, vec![on_host])); + sim.state.reconnected(); + let fresh = batch(1, None, vec![stored(None, &[1, 2])]); + sim.reference.apply_cleared(sim.worker); + sim.reference_apply(&fresh); + assert_eq!(sim.feed(&fresh), BatchOutcome::Applied); + assert_eq!( + sim.deliver(&batch(2, None, vec![removed(&[2])])), + BatchOutcome::Applied + ); + sim.assert_matches_reference(); + assert_eq!(sim.reference.worker_block_count(sim.worker), 1); + } + + /// The chunks a relay sends for a live set, stamped up to `through`: + /// chunk 0 begins with the clear, every chunk is marked. + fn snapshot_chunks( + through: u64, + rank: Option, + stores_per_chunk: &[Vec], + blocks: u64, + ) -> Vec { + let count = stores_per_chunk.len() as u32; + stores_per_chunk + .iter() + .enumerate() + .map(|(index, stores)| { + let mut events = Vec::new(); + if index == 0 { + events.push(cleared()); + } + events.extend(stores.iter().cloned()); + KvEventBatch { + sequence_number: through + 1 - u64::from(count) + index as u64, + timestamp: 0.0, + events, + dp_rank: rank, + snapshot: Some(KvSnapshotChunk { + index: index as u32, + count, + blocks, + unknown_before: 0, + }), + load: None, + } + }) + .collect() + } + + fn index_blocks(sim: &Sim) -> BTreeSet<(u32, usize, ContentHash, SequenceHash)> { + sim.indexer.debug_blocks().into_iter().collect() + } + + /// A gateway that starts after the relay's history rolled receives the + /// live set as a snapshot: applied as one `snapshot` resync, it leaves + /// the index equal to the one that saw the whole stream (and to the + /// reference), and the live stream continues from the cut with nothing + /// skipped or repeated. + #[test] + fn a_relay_snapshot_is_applied_as_a_resync_and_rebuilds_the_live_set() { + let mut seen = Sim::new(); + seen.deliver(&batch(1, None, vec![stored(None, &[1, 2, 3])])); + seen.deliver(&batch(2, None, vec![stored(Some(3), &[4, 5])])); + seen.deliver(&batch(3, None, vec![stored(None, &[10])])); + seen.deliver(&batch(4, None, vec![removed(&[5, 10])])); + seen.deliver(&batch(5, None, vec![stored(Some(2), &[6])])); + seen.assert_matches_reference(); + + // Live set at the cut (sequence 5): 1, 2, 3, 4 and the branch 6. + let chunks = snapshot_chunks( + 5, + None, + &[ + vec![stored(None, &[1, 2, 3])], + vec![stored(Some(3), &[4]), stored(Some(2), &[6])], + ], + 5, + ); + let mut fresh = Sim::new(); + for chunk in &chunks { + assert_eq!(fresh.feed(chunk), BatchOutcome::Applied); + } + assert!(fresh.state.snapshot.is_none(), "both chunks arrived"); + assert_eq!(fresh.state.ranks[&0].cursor(), Cursor::Live(5)); + assert_eq!(fresh.state.resume_sequence(), 5); + assert_eq!( + index_blocks(&fresh), + index_blocks(&seen), + "index after the snapshot" + ); + assert_eq!( + index_blocks(&fresh), + seen.reference.blocks(), + "against the reference" + ); + + // The cut's own sequence is behind the cursor; the next one applies. + assert_eq!( + fresh.feed(&batch(5, None, vec![stored(None, &[99])])), + BatchOutcome::Skipped + ); + let live = batch(6, None, vec![stored(Some(6), &[7]), removed(&[4])]); + assert_eq!(fresh.feed(&live), BatchOutcome::Applied); + assert_eq!(seen.deliver(&live), BatchOutcome::Applied); + seen.assert_matches_reference(); + assert_eq!( + index_blocks(&fresh), + index_blocks(&seen), + "index after the live batch" + ); + assert_eq!(seen.reference.worker_block_count(seen.worker), 5); + } + + /// The snapshot lists a block once per physical copy the engine still + /// holds, so the copy counts after it match a gateway that saw every + /// store and removal: the same removals evict the same blocks. + #[test] + fn a_relay_snapshot_carries_the_copies_the_engine_still_holds() { + let mut seen = Sim::new(); + seen.feed(&batch(1, None, vec![stored(None, &[1, 2])])); + // A second copy of 1 and of 2, one copy of 1 removed since. + seen.feed(&batch( + 2, + None, + vec![stored(None, &[1]), stored(Some(1), &[2])], + )); + seen.feed(&batch(3, None, vec![removed(&[1])])); + let chunks = snapshot_chunks( + 3, + None, + &[vec![stored(None, &[1, 2]), stored(Some(1), &[2])]], + 3, + ); + let mut fresh = Sim::new(); + assert_eq!(fresh.feed(&chunks[0]), BatchOutcome::Applied); + assert_eq!(index_blocks(&fresh), index_blocks(&seen)); + for (seq, hashes) in [(4, &[1][..]), (5, &[2][..]), (6, &[2][..])] { + let removal = batch(seq, None, vec![removed(hashes)]); + assert_eq!(fresh.feed(&removal), BatchOutcome::Applied); + assert_eq!(seen.feed(&removal), BatchOutcome::Applied); + assert_eq!(index_blocks(&fresh), index_blocks(&seen), "after {seq}"); + } + assert!(index_blocks(&fresh).is_empty(), "every copy is gone"); + } + + /// A snapshot chunk is a resync wherever the rank's cursor stands: the + /// gap and duplicate rules do not apply to it, the worker's old blocks + /// go, the cursor moves to the chunk's stamp. + #[test] + fn a_snapshot_replaces_whatever_the_rank_held_and_moves_its_cursor() { + let mut sim = Sim::new(); + sim.deliver(&batch(1, None, vec![stored(None, &[1, 2])])); + sim.deliver(&batch(2, None, vec![stored(Some(2), &[3])])); + assert_eq!( + sim.feed(&batch(9, None, vec![stored(None, &[9])])), + BatchOutcome::Gap { + expected: 3, + received: 9 + } + ); + // The relay's answer to the resubscription: its window had rolled. + let chunks = snapshot_chunks(40, None, &[vec![stored(None, &[7, 8])]], 2); + assert_eq!(sim.feed(&chunks[0]), BatchOutcome::Applied); + sim.reference_apply(&chunks[0]); + sim.assert_matches_reference(); + assert_eq!(sim.state.ranks[&0].cursor(), Cursor::Live(40)); + assert!(sim.state.ranks[&0].replay_pending().is_none()); + assert!(!sim.state.ranks[&0].is_degraded()); + assert_eq!(sim.reference.worker_block_count(sim.worker), 2); + let next = batch(41, None, vec![removed(&[7])]); + assert_eq!(sim.deliver(&next), BatchOutcome::Applied); + sim.assert_matches_reference(); + assert_eq!(sim.reference.worker_block_count(sim.worker), 1); + } + + /// A batch that carries only a load record repeats the last sequence and + /// has no events: it is recognised before admission, so it is neither a + /// duplicate nor a restart to the rank's cursor. + #[test] + fn a_load_only_batch_is_recognised_before_admission() { + let record = EngineLoad { + running_requests: 3, + load_only: true, + ..Default::default() + }; + let mut only = batch(7, None, vec![]); + only.load = Some(record.clone()); + assert!(KvEventMonitor::is_load_only(&only)); + let mut carrying = batch(8, None, vec![stored(None, &[1])]); + carrying.load = Some(EngineLoad { + load_only: false, + ..record + }); + assert!(!KvEventMonitor::is_load_only(&carrying)); + assert!(!KvEventMonitor::is_load_only(&batch(9, None, vec![]))); + } + + /// A snapshot whose relay joined the publisher late and could not replay + /// the start leaves the rank degraded, as an unrecovered gap would; a + /// whole one lifts it. + #[test] + fn a_snapshot_that_starts_late_marks_the_rank_degraded() { + let mut sim = Sim::new(); + let mut late = snapshot_chunks(9, None, &[vec![stored(None, &[1])]], 1); + late[0].snapshot.as_mut().unwrap().unknown_before = 40; + assert_eq!(sim.feed(&late[0]), BatchOutcome::Applied); + assert!(sim.state.ranks[&0].is_degraded()); + assert_eq!(sim.state.degraded_ranks(), 1); + assert_eq!( + sim.feed(&batch(10, None, vec![stored(Some(1), &[2])])), + BatchOutcome::Applied + ); + assert!(sim.state.ranks[&0].is_degraded(), "until the next resync"); + let whole = snapshot_chunks(12, None, &[vec![stored(None, &[1])]], 1); + assert_eq!(sim.feed(&whole[0]), BatchOutcome::Applied); + assert!(!sim.state.ranks[&0].is_degraded()); + assert_eq!(sim.state.degraded_ranks(), 0); + } + + /// A stream that ends with chunks still owed leaves a partial live set: + /// the cursors are forgotten so the next subscription asks from zero and + /// gets a whole snapshot; a complete one leaves nothing to abandon. + #[test] + fn a_stream_that_ends_mid_snapshot_starts_the_next_subscription_from_zero() { + let mut sim = Sim::new(); + let chunks = snapshot_chunks( + 10, + None, + &[vec![stored(None, &[1])], vec![stored(Some(1), &[2])]], + 2, + ); + assert_eq!(sim.feed(&chunks[0]), BatchOutcome::Applied); + assert_eq!( + sim.state.snapshot, + Some(SnapshotProgress { + count: 2, + applied: 1, + blocks: 2 + }) + ); + assert_eq!(sim.state.resume_sequence(), 9); + assert!(sim.state.abandon_snapshot()); + assert_eq!(sim.state.resume_sequence(), 0); + assert!(!sim.state.abandon_snapshot()); + for chunk in &chunks { + assert_eq!(sim.feed(chunk), BatchOutcome::Applied); + } + assert!(sim.state.snapshot.is_none()); + assert_eq!(sim.state.resume_sequence(), 10); + assert!(!sim.state.abandon_snapshot()); + } + + #[test] + fn recovery_snapshot_resync_applies_the_tail_in_order() { + let mut sim = Sim::new(); + sim.deliver(&batch(1, None, vec![stored(None, &[1, 2])])); + sim.deliver(&batch(2, None, vec![stored(Some(2), &[3])])); + // Lost history: the subscriber asks for a snapshot out of band and + // holds live batches meanwhile. + sim.state.ranks.get_mut(&0).unwrap().begin_snapshot(); + let live3 = batch(3, None, vec![removed(&[3])]); + let live4 = batch(4, None, vec![stored(Some(2), &[5])]); + assert_eq!(sim.feed(&live3), BatchOutcome::Skipped); + assert_eq!(sim.feed(&live4), BatchOutcome::Skipped); + assert_eq!(sim.state.ranks[&0].tail_len(), 2); + // The snapshot (engine state through sequence 2) replaces the rank. + let snapshot = batch( + 2, + None, + vec![cleared(), stored(None, &[1, 2]), stored(Some(2), &[3])], + ); + KvEventMonitor::apply_cleared(sim.worker, &sim.indexer, &mut sim.state.index); + sim.apply_direct(&snapshot); + let tail = sim + .state + .ranks + .get_mut(&0) + .unwrap() + .finish_snapshot(2) + .expect("tail intact"); + assert_eq!(tail.len(), 2); + for b in &tail { + sim.apply_direct(b); + } + sim.reference_apply(&live3); + sim.reference_apply(&live4); + assert_eq!( + sim.deliver(&batch(5, None, vec![stored(Some(5), &[6])])), + BatchOutcome::Applied + ); + assert!(!sim.state.ranks[&0].is_degraded()); + sim.assert_matches_reference(); + } + + /// A store whose parent the index does not hold is placed from the root + /// and counted, per worker, so a soak can see how much of a chain's + /// content arrives detached (where the chain index fragments first). + #[test] + fn parentless_stores_are_counted_per_worker() { + use metrics_exporter_prometheus::{PrometheusBuilder, PrometheusHandle}; + + fn counter(handle: &PrometheusHandle, name: &str) -> Option { + handle + .render() + .lines() + .find(|line| line.starts_with(&format!("{name}{{worker=\"grpc://w1:9000\"}}"))) + .and_then(|line| line.rsplit(' ').next()) + .and_then(|value| value.parse().ok()) + } + + let recorder = PrometheusBuilder::new().build_recorder(); + let handle = recorder.handle(); + metrics::with_local_recorder(&recorder, || { + let mut sim = Sim::new(); + sim.deliver(&batch(1, None, vec![stored(None, &[1, 2])])); + assert_eq!( + counter(&handle, "smg_kv_event_parentless_stores_total"), + None + ); + // The parent of this store was never seen (dropped or evicted). + sim.deliver(&batch(2, None, vec![stored(Some(99), &[3, 4, 5])])); + assert_eq!( + counter(&handle, "smg_kv_event_parentless_stores_total"), + Some(1.0) + ); + assert_eq!( + counter(&handle, "smg_kv_event_parentless_blocks_total"), + Some(3.0) + ); + assert_eq!(sim.state.index.counters.parentless_stores, 1); + // A chained store after a held parent is not counted. + sim.deliver(&batch(3, None, vec![stored(Some(2), &[6])])); + assert_eq!( + counter(&handle, "smg_kv_event_parentless_stores_total"), + Some(1.0) + ); + sim.assert_matches_reference(); + }); + } +} diff --git a/model_gateway/src/worker/kv_event_monitor/apply.rs b/model_gateway/src/worker/kv_event_monitor/apply.rs new file mode 100644 index 0000000000..c7e30135ae --- /dev/null +++ b/model_gateway/src/worker/kv_event_monitor/apply.rs @@ -0,0 +1,1037 @@ +//! One event into the index: which tiers, cache groups, localities and +//! owners are indexed, content hashing under the event's namespace, and the +//! physical copies of a block counted per worker so a removal evicts the +//! block only when its last copy goes. + +use std::collections::{hash_map::Entry, HashMap}; + +use kv_index::{ + salt::{content_hash_with_seed, namespace_seed}, + ApplyError, SequenceHash, StoredBlock, +}; +use smg_grpc_client::common_proto::{ + kv_cache_event, KvBlock, KvBlocksRemoved, KvBlocksStored, KvCacheEvent, KvCacheLocality, + KvCacheTier, +}; +use tracing::warn; + +use super::KvEventMonitor; +use crate::worker::kv_index_backend::{KvIndex, WorkerBlocks}; + +impl KvEventMonitor { + /// Apply a single KV cache event to the indexer. + pub(crate) fn apply_event( + event: &KvCacheEvent, + worker_id: u32, + indexer: &KvIndex, + worker_blocks: &mut WorkerIndexState, + ) { + let Some(ref data) = event.data else { + return; + }; + + match data { + kv_cache_event::Data::Stored(stored) => { + Self::apply_stored(stored, worker_id, indexer, worker_blocks); + } + kv_cache_event::Data::Removed(removed) => { + Self::apply_removed(removed, worker_id, indexer, worker_blocks); + } + kv_cache_event::Data::Cleared(cleared) => { + if worker_blocks.admits(None, None, None, cleared.ownership.as_deref()) { + Self::apply_cleared(worker_id, indexer, worker_blocks); + } + } + } + } + + /// Convert proto `KvBlocksStored` and apply to the indexer. + /// + /// Blocks on the disk and external tiers, in cache groups other than main + /// attention, not local to the worker, or owned by a residency agent are + /// counted and skipped. Content hashes are computed under the event's + /// LoRA name and cache salt, so a salted block matches only a request + /// hashed under the same namespace. + fn apply_stored( + stored: &KvBlocksStored, + worker_id: u32, + indexer: &KvIndex, + worker_blocks: &mut WorkerIndexState, + ) { + if !worker_blocks.admits( + stored.kv_cache_spec_kind.as_deref(), + stored.group_idx, + stored.locality, + stored.ownership.as_deref(), + ) { + return; + } + let first_level = stored.blocks.first().and_then(|block| block.cache_level); + let Some(tier) = indexed_tier(stored.tier, first_level) else { + worker_blocks.counters.untracked_tier += 1; + return; + }; + + let seed = namespace_seed(stored.lora_name.as_deref(), stored.cache_salt.as_deref()); + let blocks: Vec = stored + .blocks + .iter() + .map(|block| convert_kv_block(block, seed)) + .collect(); + worker_blocks.note_stored(&blocks, tier); + + let parent_seq_hash = stored.parent_block_hash.map(SequenceHash::from); + + match indexer.apply_stored( + worker_id, + &blocks, + parent_seq_hash, + &mut worker_blocks.blocks, + ) { + Ok(()) => {} + Err(ApplyError::WorkerNotTracked | ApplyError::ParentBlockNotFound) => { + // Cold start or parent evicted — retry without parent to start a new chain. + worker_blocks.counters.parentless_stores += 1; + worker_blocks.counters.parentless_blocks += blocks.len() as u64; + if let Err(e) = + indexer.apply_stored(worker_id, &blocks, None, &mut worker_blocks.blocks) + { + warn!( + worker_id = worker_id, + error = %e, + "Failed to apply stored event after fallback" + ); + } + } + } + } + + /// Convert proto `KvBlocksRemoved` and apply to the indexer. + /// + /// A removal names one tier; a block leaves the index only when no + /// indexed copy remains on another tier. + fn apply_removed( + removed: &KvBlocksRemoved, + worker_id: u32, + indexer: &KvIndex, + worker_blocks: &mut WorkerIndexState, + ) { + if !worker_blocks.admits( + None, + removed.group_idx, + removed.locality, + removed.ownership.as_deref(), + ) { + return; + } + let Some(tier) = indexed_tier(removed.tier, removed.cache_level) else { + worker_blocks.counters.untracked_tier += 1; + return; + }; + + let hashes = removed.block_hashes.iter().map(|&h| SequenceHash::from(h)); + let seq_hashes: Vec = + if tier == IndexedTier::Device && worker_blocks.copies.is_empty() { + hashes.collect() + } else { + hashes + .filter(|&seq_hash| worker_blocks.release(seq_hash, tier)) + .collect() + }; + + indexer.apply_removed(worker_id, &seq_hashes, &mut worker_blocks.blocks); + } + + /// Drop every block of a worker from the indexer and forget its copies. + pub(super) fn apply_cleared( + worker_id: u32, + indexer: &KvIndex, + worker_blocks: &mut WorkerIndexState, + ) { + indexer.apply_cleared(worker_id, &mut worker_blocks.blocks); + worker_blocks.copies.clear(); + } +} + +/// Convert a proto `KvBlock` to a kv-index `StoredBlock`, hashing its tokens +/// under `seed` (see [`namespace_seed`]). +pub(super) fn convert_kv_block(block: &KvBlock, seed: u64) -> StoredBlock { + StoredBlock { + seq_hash: SequenceHash::from(block.block_hash), + content_hash: content_hash_with_seed(&block.token_ids, seed), + } +} + +/// The residency tiers the index tracks: the device, and the host cache the +/// engine restores from without recompute. Disk and external copies are not +/// indexed; a hit there costs an engine-side fetch the router cannot price. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +enum IndexedTier { + Device, + Host, +} + +/// The tier an event names: its `tier` when set, else a block's +/// `cache_level` (absent means the device). `None` when the index does not +/// track that tier. +fn indexed_tier(tier: Option, cache_level: Option) -> Option { + let tier = match tier.and_then(|tier| KvCacheTier::try_from(tier).ok()) { + Some(KvCacheTier::Unspecified) | None => match cache_level.unwrap_or(0) { + 0 => KvCacheTier::Device, + 1 => KvCacheTier::Host, + 2 => KvCacheTier::Disk, + _ => KvCacheTier::External, + }, + Some(tier) => tier, + }; + match tier { + KvCacheTier::Unspecified | KvCacheTier::Device => Some(IndexedTier::Device), + KvCacheTier::Host => Some(IndexedTier::Host), + KvCacheTier::Disk | KvCacheTier::External => None, + } +} + +/// Cache-group kinds whose blocks hold the main attention KV, the ones +/// prefix matching is about. Sliding-window and Mamba groups are skipped. +const MAIN_ATTENTION_KINDS: [&str; 3] = ["full_attention", "mla_attention", "sink_full_attention"]; + +/// The most physical copies of one block counted per tier. vLLM's opt-in +/// `kv_cache_report_mode: full` re-announces whole hit chains without +/// removals, which would grow a count without bound; the cap turns that into +/// at most this many extra removals before the block leaves the index. +const COPIES_CAP: u8 = 8; + +/// What the positional index did not take at face value, by reason; logged +/// when the worker's subscription ends. +#[derive(Default, Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) struct WorkerIndexCounters { + /// Stores and removals on tiers the index does not track. + pub(crate) untracked_tier: u64, + /// Events for cache groups other than main attention. + pub(crate) non_main_group: u64, + /// Events for blocks not local to the worker. + pub(crate) remote: u64, + /// Events owned by a residency agent rather than the engine. + pub(crate) foreign_owner: u64, + /// Removals on a tier that held no counted copy of the block. + pub(crate) unknown_copy: u64, + /// Stores of a block already indexed on that tier (a second physical copy). + pub(crate) duplicate_copies: u64, + /// Copy counts that hit [`COPIES_CAP`]. + pub(crate) capped_copies: u64, + /// Stores whose parent the index did not hold, placed as a new chain from + /// the root instead, and the blocks they carried. Every one duplicates + /// content the chain may already hold further down, so a rising count is + /// where fragmentation of the chain index is looked for first. + pub(crate) parentless_stores: u64, + pub(crate) parentless_blocks: u64, +} + +impl WorkerIndexCounters { + /// Parent-less stores and blocks this state saw since `before`. + pub(super) fn parentless_since(&self, before: &Self) -> (u64, u64) { + ( + self.parentless_stores - before.parentless_stores, + self.parentless_blocks - before.parentless_blocks, + ) + } +} + +/// Physical copies of one block per tier, all of the worker's ranks pooled: +/// the worker URL is the routing target and its copies are interchangeable +/// for a prefix hit. +#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)] +struct Copies { + device: u8, + host: u8, +} + +impl Copies { + fn on(&mut self, tier: IndexedTier) -> &mut u8 { + match tier { + IndexedTier::Device => &mut self.device, + IndexedTier::Host => &mut self.host, + } + } + + fn none(self) -> bool { + self.device == 0 && self.host == 0 + } +} + +/// A worker's share of the positional index: the indexer's reverse map plus +/// the physical copies of each block per tier, once the worker reports more +/// than one copy or a tier other than the device. +/// +/// The engines do not deduplicate physical blocks: vLLM recomputes the last +/// block of an exact resend into a second copy with the same hash and emits +/// `BlockRemoved` per copy, and SGLang's HiCache keeps a host copy next to +/// the device one. The relay forwards every store and removal, so this state +/// counts copies and lets a removal evict the block only when none remains. +/// +/// Counting is sparse. A block gets an entry in `copies` only on a host +/// store or on a second store of an indexed hash; an indexed block without +/// an entry is a single device copy. A worker that never duplicates and +/// never offloads pays one lookup per stored block and no memory. +#[derive(Default)] +pub(crate) struct WorkerIndexState { + /// The indexer's caller-owned reverse map for this worker. + pub(crate) blocks: WorkerBlocks, + /// Copies per tier of blocks with a host copy or more than one copy. + copies: HashMap, + /// Cache groups whose kind is not main attention. + non_main_groups: Vec, + pub(crate) counters: WorkerIndexCounters, +} + +impl WorkerIndexState { + /// Whether an event with these attributes belongs in the index. A store + /// names its group's kind; later events for that group may omit it, so + /// non-main groups are remembered. + fn admits( + &mut self, + kind: Option<&str>, + group_idx: Option, + locality: Option, + ownership: Option<&str>, + ) -> bool { + if ownership.is_some_and(|owner| owner.eq_ignore_ascii_case("kvcr")) { + self.counters.foreign_owner += 1; + return false; + } + if locality.is_some_and(|locality| locality == KvCacheLocality::Remote as i32) { + self.counters.remote += 1; + return false; + } + let main = match kind { + Some(kind) => { + let main = MAIN_ATTENTION_KINDS.contains(&kind); + if let Some(group) = group_idx { + if main { + self.non_main_groups.retain(|&known| known != group); + } else if !self.non_main_groups.contains(&group) { + self.non_main_groups.push(group); + } + } + main + } + None => !group_idx.is_some_and(|group| self.non_main_groups.contains(&group)), + }; + if !main { + self.counters.non_main_group += 1; + } + main + } + + /// Count a store on `tier`, before the indexer applies it. A second copy + /// of an indexed block, or any host copy, opens the block's entry; the + /// implicit single device copy is credited when it does. + fn note_stored(&mut self, blocks: &[StoredBlock], tier: IndexedTier) { + for block in blocks { + let indexed = self.blocks.contains_key(block.seq_hash); + let entry = match self.copies.entry(block.seq_hash) { + Entry::Occupied(entry) => entry.into_mut(), + Entry::Vacant(vacant) => match tier { + IndexedTier::Device if !indexed => continue, + _ => vacant.insert(Copies { + device: u8::from(indexed), + host: 0, + }), + }, + }; + let count = entry.on(tier); + if *count > 0 { + self.counters.duplicate_copies += 1; + } + if *count < COPIES_CAP { + *count += 1; + } else { + self.counters.capped_copies += 1; + } + } + } + + /// Drop one copy of a block on `tier`; `true` when no counted copy + /// remains and the block should leave the index. + fn release(&mut self, seq_hash: SequenceHash, tier: IndexedTier) -> bool { + match self.copies.entry(seq_hash) { + Entry::Occupied(mut entry) => { + let count = entry.get_mut().on(tier); + if *count == 0 { + self.counters.unknown_copy += 1; + return false; + } + *count -= 1; + if entry.get().none() { + entry.remove(); + true + } else { + false + } + } + Entry::Vacant(_) => match tier { + IndexedTier::Device => true, + IndexedTier::Host => { + self.counters.unknown_copy += 1; + false + } + }, + } + } +} + +#[cfg(test)] +mod tests { + use kv_index::{ + compute_content_hash, compute_request_content_hashes, + salt::namespaced_request_content_hashes, ContentHash, XXH3_SEED, + }; + use smg_grpc_client::common_proto::KvCacheCleared; + + use super::*; + + #[test] + fn a_proto_block_becomes_its_engine_and_content_hashes() { + let block = KvBlock { + block_hash: 42, + token_ids: vec![1, 2, 3, 4], + block_size: 4, + lora_id: None, + cache_level: None, + ..Default::default() + }; + let stored = convert_kv_block(&block, XXH3_SEED); + assert_eq!(stored.seq_hash, SequenceHash::from(42i64)); + assert_eq!(stored.content_hash, compute_content_hash(&[1, 2, 3, 4])); + } + + #[test] + fn a_negative_engine_hash_keeps_its_bits() { + let block = KvBlock { + block_hash: -1, + token_ids: vec![10, 20], + block_size: 2, + lora_id: None, + cache_level: None, + ..Default::default() + }; + let stored = convert_kv_block(&block, XXH3_SEED); + assert_eq!(stored.seq_hash, SequenceHash(u64::MAX)); + } + + #[test] + fn a_block_without_tokens_hashes_the_empty_content() { + let block = KvBlock { + block_hash: 100, + token_ids: vec![], + block_size: 0, + lora_id: None, + cache_level: None, + ..Default::default() + }; + let stored = convert_kv_block(&block, XXH3_SEED); + assert_eq!(stored.seq_hash, SequenceHash::from(100i64)); + assert_eq!(stored.content_hash, compute_content_hash(&[])); + } + + #[test] + fn a_root_store_indexes_its_blocks() { + let indexer = KvIndex::positional(64); + let w1 = indexer.intern_worker("http://w1:8000").unwrap(); + let mut wb = WorkerIndexState::default(); + let stored = KvBlocksStored { + blocks: vec![ + KvBlock { + block_hash: 1, + token_ids: vec![10, 20, 30, 40], + block_size: 4, + lora_id: None, + cache_level: None, + ..Default::default() + }, + KvBlock { + block_hash: 2, + token_ids: vec![50, 60, 70, 80], + block_size: 4, + lora_id: None, + cache_level: None, + ..Default::default() + }, + ], + parent_block_hash: None, + ..Default::default() + }; + + KvEventMonitor::apply_stored(&stored, w1, &indexer, &mut wb); + assert_eq!(indexer.current_size(), 2); + } + + #[test] + fn a_chained_store_extends_its_parent() { + let indexer = KvIndex::positional(64); + let w1 = indexer.intern_worker("http://w1:8000").unwrap(); + let mut wb = WorkerIndexState::default(); + + let stored1 = KvBlocksStored { + blocks: vec![KvBlock { + block_hash: 1, + token_ids: vec![10, 20, 30, 40], + block_size: 4, + lora_id: None, + cache_level: None, + ..Default::default() + }], + parent_block_hash: None, + ..Default::default() + }; + KvEventMonitor::apply_stored(&stored1, w1, &indexer, &mut wb); + + let stored2 = KvBlocksStored { + blocks: vec![KvBlock { + block_hash: 2, + token_ids: vec![50, 60, 70, 80], + block_size: 4, + lora_id: None, + cache_level: None, + ..Default::default() + }], + parent_block_hash: Some(1), + ..Default::default() + }; + KvEventMonitor::apply_stored(&stored2, w1, &indexer, &mut wb); + assert_eq!(indexer.current_size(), 2); + } + + #[test] + fn a_store_with_an_unknown_parent_starts_a_new_chain() { + let indexer = KvIndex::positional(64); + let w1 = indexer.intern_worker("http://new-worker:8000").unwrap(); + let mut wb = WorkerIndexState::default(); + + // Pass parent_block_hash for an untracked worker — should fallback to no parent. + let stored = KvBlocksStored { + blocks: vec![KvBlock { + block_hash: 1, + token_ids: vec![10, 20, 30, 40], + block_size: 4, + lora_id: None, + cache_level: None, + ..Default::default() + }], + parent_block_hash: Some(999), + ..Default::default() + }; + KvEventMonitor::apply_stored(&stored, w1, &indexer, &mut wb); + assert_eq!(indexer.current_size(), 1); + } + + #[test] + fn a_removal_takes_the_block_out() { + let indexer = KvIndex::positional(64); + let w1 = indexer.intern_worker("http://w1:8000").unwrap(); + let mut wb = WorkerIndexState::default(); + + let stored = KvBlocksStored { + blocks: vec![ + KvBlock { + block_hash: 1, + token_ids: vec![10, 20, 30, 40], + block_size: 4, + lora_id: None, + cache_level: None, + ..Default::default() + }, + KvBlock { + block_hash: 2, + token_ids: vec![50, 60, 70, 80], + block_size: 4, + lora_id: None, + cache_level: None, + ..Default::default() + }, + ], + parent_block_hash: None, + ..Default::default() + }; + KvEventMonitor::apply_stored(&stored, w1, &indexer, &mut wb); + + let removed = KvBlocksRemoved { + block_hashes: vec![2], + cache_level: None, + ..Default::default() + }; + KvEventMonitor::apply_removed(&removed, w1, &indexer, &mut wb); + assert_eq!(indexer.current_size(), 1); + } + + #[test] + fn a_clear_empties_the_worker() { + let indexer = KvIndex::positional(64); + let w1 = indexer.intern_worker("http://w1:8000").unwrap(); + let mut wb = WorkerIndexState::default(); + + let stored = KvBlocksStored { + blocks: vec![KvBlock { + block_hash: 1, + token_ids: vec![10, 20, 30, 40], + block_size: 4, + lora_id: None, + cache_level: None, + ..Default::default() + }], + parent_block_hash: None, + ..Default::default() + }; + KvEventMonitor::apply_stored(&stored, w1, &indexer, &mut wb); + assert_eq!(indexer.current_size(), 1); + + KvEventMonitor::apply_cleared(w1, &indexer, &mut wb); + assert_eq!(indexer.current_size(), 0); + } + + #[test] + fn apply_event_dispatches_a_store() { + let indexer = KvIndex::positional(64); + let w1 = indexer.intern_worker("http://w1:8000").unwrap(); + let mut wb = WorkerIndexState::default(); + let event = KvCacheEvent { + event_id: 1, + data: Some(kv_cache_event::Data::Stored(KvBlocksStored { + blocks: vec![KvBlock { + block_hash: 42, + token_ids: vec![1, 2, 3, 4], + block_size: 4, + lora_id: None, + cache_level: None, + ..Default::default() + }], + parent_block_hash: None, + ..Default::default() + })), + }; + + KvEventMonitor::apply_event(&event, w1, &indexer, &mut wb); + assert_eq!(indexer.current_size(), 1); + } + + #[test] + fn apply_event_dispatches_a_removal() { + let indexer = KvIndex::positional(64); + let w1 = indexer.intern_worker("http://w1:8000").unwrap(); + let mut wb = WorkerIndexState::default(); + + let stored_event = KvCacheEvent { + event_id: 1, + data: Some(kv_cache_event::Data::Stored(KvBlocksStored { + blocks: vec![KvBlock { + block_hash: 1, + token_ids: vec![1, 2, 3, 4], + block_size: 4, + lora_id: None, + cache_level: None, + ..Default::default() + }], + parent_block_hash: None, + ..Default::default() + })), + }; + KvEventMonitor::apply_event(&stored_event, w1, &indexer, &mut wb); + + let removed_event = KvCacheEvent { + event_id: 2, + data: Some(kv_cache_event::Data::Removed(KvBlocksRemoved { + block_hashes: vec![1], + cache_level: None, + ..Default::default() + })), + }; + KvEventMonitor::apply_event(&removed_event, w1, &indexer, &mut wb); + assert_eq!(indexer.current_size(), 0); + } + + #[test] + fn apply_event_dispatches_a_clear() { + let indexer = KvIndex::positional(64); + let w1 = indexer.intern_worker("http://w1:8000").unwrap(); + let mut wb = WorkerIndexState::default(); + + KvEventMonitor::apply_event( + &KvCacheEvent { + event_id: 1, + data: Some(kv_cache_event::Data::Stored(KvBlocksStored { + blocks: vec![KvBlock { + block_hash: 1, + token_ids: vec![1, 2, 3, 4], + block_size: 4, + lora_id: None, + cache_level: None, + ..Default::default() + }], + parent_block_hash: None, + ..Default::default() + })), + }, + w1, + &indexer, + &mut wb, + ); + + // Clear + KvEventMonitor::apply_event( + &KvCacheEvent { + event_id: 2, + data: Some(kv_cache_event::Data::Cleared(KvCacheCleared::default())), + }, + w1, + &indexer, + &mut wb, + ); + assert_eq!(indexer.current_size(), 0); + } + + #[test] + fn an_event_without_data_is_ignored() { + let indexer = KvIndex::positional(64); + let w1 = indexer.intern_worker("http://w1:8000").unwrap(); + let mut wb = WorkerIndexState::default(); + let event = KvCacheEvent { + event_id: 1, + data: None, + }; + KvEventMonitor::apply_event(&event, w1, &indexer, &mut wb); + assert_eq!(indexer.current_size(), 0); + } + + const TOKENS: [u32; 4] = [10, 20, 30, 40]; + + fn stored_event(hash: i64, tokens: &[u32]) -> KvBlocksStored { + KvBlocksStored { + blocks: vec![KvBlock { + block_hash: hash, + token_ids: tokens.to_vec(), + block_size: tokens.len() as i32, + ..Default::default() + }], + ..Default::default() + } + } + + fn removed_event(hash: i64) -> KvBlocksRemoved { + KvBlocksRemoved { + block_hashes: vec![hash], + ..Default::default() + } + } + + fn routable(indexer: &KvIndex, worker: u32, hashes: &[ContentHash]) -> bool { + indexer + .find_matches(hashes, false) + .scores + .get(&worker) + .is_some_and(|&depth| depth > 0) + } + + #[test] + fn salted_stores_match_only_their_namespace() { + let indexer = KvIndex::positional(64); + let w1 = indexer.intern_worker("http://w1:8000").unwrap(); + let mut wb = WorkerIndexState::default(); + let mut stored = stored_event(1, &TOKENS); + stored.lora_name = Some("adapter".to_string()); + stored.cache_salt = Some("tenant-a".to_string()); + KvEventMonitor::apply_stored(&stored, w1, &indexer, &mut wb); + + let same = namespaced_request_content_hashes(&TOKENS, 4, Some("adapter"), Some("tenant-a")); + assert!(routable(&indexer, w1, &same)); + let plain = compute_request_content_hashes(&TOKENS, 4); + assert!(!routable(&indexer, w1, &plain)); + let lora_only = namespaced_request_content_hashes(&TOKENS, 4, Some("adapter"), None); + assert!(!routable(&indexer, w1, &lora_only)); + } + + #[test] + fn device_removal_keeps_a_block_still_on_the_host() { + let indexer = KvIndex::positional(64); + let w1 = indexer.intern_worker("http://w1:8000").unwrap(); + let mut wb = WorkerIndexState::default(); + let hashes = compute_request_content_hashes(&TOKENS, 4); + + KvEventMonitor::apply_stored(&stored_event(1, &TOKENS), w1, &indexer, &mut wb); + let mut on_host = stored_event(1, &TOKENS); + on_host.tier = Some(KvCacheTier::Host as i32); + KvEventMonitor::apply_stored(&on_host, w1, &indexer, &mut wb); + assert_eq!(indexer.current_size(), 1); + + KvEventMonitor::apply_removed(&removed_event(1), w1, &indexer, &mut wb); + assert!( + routable(&indexer, w1, &hashes), + "the host copy keeps the block routable" + ); + + let mut from_host = removed_event(1); + from_host.tier = Some(KvCacheTier::Host as i32); + KvEventMonitor::apply_removed(&from_host, w1, &indexer, &mut wb); + assert!(!routable(&indexer, w1, &hashes)); + assert_eq!(indexer.current_size(), 0); + assert!(wb.copies.is_empty()); + } + + #[test] + fn cache_level_stands_in_for_the_tier() { + let indexer = KvIndex::positional(64); + let w1 = indexer.intern_worker("http://w1:8000").unwrap(); + let mut wb = WorkerIndexState::default(); + let hashes = compute_request_content_hashes(&TOKENS, 4); + + let mut on_host = stored_event(1, &TOKENS); + on_host.blocks[0].cache_level = Some(1); + KvEventMonitor::apply_stored(&on_host, w1, &indexer, &mut wb); + KvEventMonitor::apply_stored(&stored_event(1, &TOKENS), w1, &indexer, &mut wb); + + let mut from_host = removed_event(1); + from_host.cache_level = Some(1); + KvEventMonitor::apply_removed(&from_host, w1, &indexer, &mut wb); + assert!( + routable(&indexer, w1, &hashes), + "the device copy keeps the block routable" + ); + + KvEventMonitor::apply_removed(&removed_event(1), w1, &indexer, &mut wb); + assert!(!routable(&indexer, w1, &hashes)); + } + + #[test] + fn host_removal_without_a_host_copy_evicts_nothing() { + let indexer = KvIndex::positional(64); + let w1 = indexer.intern_worker("http://w1:8000").unwrap(); + let mut wb = WorkerIndexState::default(); + let hashes = compute_request_content_hashes(&TOKENS, 4); + + KvEventMonitor::apply_stored(&stored_event(1, &TOKENS), w1, &indexer, &mut wb); + let mut from_host = removed_event(1); + from_host.tier = Some(KvCacheTier::Host as i32); + KvEventMonitor::apply_removed(&from_host, w1, &indexer, &mut wb); + assert!(routable(&indexer, w1, &hashes)); + assert_eq!(wb.counters.unknown_copy, 1); + } + + #[test] + fn disk_and_external_tiers_are_counted_not_indexed() { + let indexer = KvIndex::positional(64); + let w1 = indexer.intern_worker("http://w1:8000").unwrap(); + let mut wb = WorkerIndexState::default(); + + let mut on_disk = stored_event(1, &TOKENS); + on_disk.tier = Some(KvCacheTier::Disk as i32); + KvEventMonitor::apply_stored(&on_disk, w1, &indexer, &mut wb); + let mut external = stored_event(2, &TOKENS); + external.blocks[0].cache_level = Some(3); + KvEventMonitor::apply_stored(&external, w1, &indexer, &mut wb); + assert_eq!(indexer.current_size(), 0); + + KvEventMonitor::apply_stored(&stored_event(1, &TOKENS), w1, &indexer, &mut wb); + let mut from_disk = removed_event(1); + from_disk.tier = Some(KvCacheTier::Disk as i32); + KvEventMonitor::apply_removed(&from_disk, w1, &indexer, &mut wb); + assert_eq!( + indexer.current_size(), + 1, + "a disk removal does not touch the device copy" + ); + assert_eq!(wb.counters.untracked_tier, 3); + } + + #[test] + fn non_main_attention_groups_are_skipped_once_their_kind_is_known() { + let indexer = KvIndex::positional(64); + let w1 = indexer.intern_worker("http://w1:8000").unwrap(); + let mut wb = WorkerIndexState::default(); + + let mut sliding = stored_event(1, &TOKENS); + sliding.group_idx = Some(1); + sliding.kv_cache_spec_kind = Some("sliding_window".to_string()); + KvEventMonitor::apply_stored(&sliding, w1, &indexer, &mut wb); + assert_eq!(indexer.current_size(), 0); + + let mut full = stored_event(2, &TOKENS); + full.group_idx = Some(0); + full.kv_cache_spec_kind = Some("full_attention".to_string()); + KvEventMonitor::apply_stored(&full, w1, &indexer, &mut wb); + assert_eq!(indexer.current_size(), 1); + + // Later events for group 1 omit the kind; the group is remembered. + let mut later = stored_event(3, &[50, 60, 70, 80]); + later.group_idx = Some(1); + KvEventMonitor::apply_stored(&later, w1, &indexer, &mut wb); + let mut removal = removed_event(2); + removal.group_idx = Some(1); + KvEventMonitor::apply_removed(&removal, w1, &indexer, &mut wb); + assert_eq!(indexer.current_size(), 1); + assert_eq!(wb.counters.non_main_group, 3); + + removal.group_idx = Some(0); + KvEventMonitor::apply_removed(&removal, w1, &indexer, &mut wb); + assert_eq!(indexer.current_size(), 0); + } + + #[test] + fn remote_and_residency_agent_events_are_skipped() { + let indexer = KvIndex::positional(64); + let w1 = indexer.intern_worker("http://w1:8000").unwrap(); + let mut wb = WorkerIndexState::default(); + + let mut remote = stored_event(1, &TOKENS); + remote.locality = Some(KvCacheLocality::Remote as i32); + KvEventMonitor::apply_stored(&remote, w1, &indexer, &mut wb); + let mut agent = stored_event(1, &TOKENS); + agent.ownership = Some("kvcr".to_string()); + KvEventMonitor::apply_stored(&agent, w1, &indexer, &mut wb); + assert_eq!(indexer.current_size(), 0); + + KvEventMonitor::apply_stored(&stored_event(1, &TOKENS), w1, &indexer, &mut wb); + let cleared = KvCacheEvent { + event_id: 1, + data: Some(kv_cache_event::Data::Cleared(KvCacheCleared { + ownership: Some("kvcr".to_string()), + })), + }; + KvEventMonitor::apply_event(&cleared, w1, &indexer, &mut wb); + assert_eq!( + indexer.current_size(), + 1, + "an agent's clear leaves the engine's blocks" + ); + assert_eq!(wb.counters.remote, 1); + assert_eq!(wb.counters.foreign_owner, 2); + } + + #[test] + fn clearing_forgets_copies_and_residency() { + let indexer = KvIndex::positional(64); + let w1 = indexer.intern_worker("http://w1:8000").unwrap(); + let mut wb = WorkerIndexState::default(); + let hashes = compute_request_content_hashes(&TOKENS, 4); + + let mut on_host = stored_event(1, &TOKENS); + on_host.tier = Some(KvCacheTier::Host as i32); + KvEventMonitor::apply_stored(&on_host, w1, &indexer, &mut wb); + assert!(!wb.copies.is_empty()); + + KvEventMonitor::apply_cleared(w1, &indexer, &mut wb); + assert_eq!(indexer.current_size(), 0); + assert!(wb.copies.is_empty()); + + KvEventMonitor::apply_stored(&stored_event(1, &TOKENS), w1, &indexer, &mut wb); + KvEventMonitor::apply_removed(&removed_event(1), w1, &indexer, &mut wb); + assert!( + !routable(&indexer, w1, &hashes), + "no stale host bit survives a clear" + ); + } + + #[test] + fn two_copies_need_two_removals() { + let indexer = KvIndex::positional(64); + let w1 = indexer.intern_worker("http://w1:8000").unwrap(); + let mut wb = WorkerIndexState::default(); + let hashes = compute_request_content_hashes(&TOKENS, 4); + + // vLLM recomputes the last block of an exact resend into a second + // physical copy with the same hash and removes the copies one by one. + KvEventMonitor::apply_stored(&stored_event(1, &TOKENS), w1, &indexer, &mut wb); + assert!(wb.copies.is_empty(), "a single device copy costs no entry"); + KvEventMonitor::apply_stored(&stored_event(1, &TOKENS), w1, &indexer, &mut wb); + assert_eq!(wb.counters.duplicate_copies, 1); + + KvEventMonitor::apply_removed(&removed_event(1), w1, &indexer, &mut wb); + assert!( + routable(&indexer, w1, &hashes), + "the other copy is still cached" + ); + assert_eq!(indexer.current_size(), 1); + + KvEventMonitor::apply_removed(&removed_event(1), w1, &indexer, &mut wb); + assert!(!routable(&indexer, w1, &hashes)); + assert_eq!(indexer.current_size(), 0); + assert!(wb.copies.is_empty()); + } + + #[test] + fn device_and_host_copies_are_counted_per_tier() { + let indexer = KvIndex::positional(64); + let w1 = indexer.intern_worker("http://w1:8000").unwrap(); + let mut wb = WorkerIndexState::default(); + let hashes = compute_request_content_hashes(&TOKENS, 4); + let mut on_host = stored_event(1, &TOKENS); + on_host.tier = Some(KvCacheTier::Host as i32); + let mut from_host = removed_event(1); + from_host.tier = Some(KvCacheTier::Host as i32); + + KvEventMonitor::apply_stored(&stored_event(1, &TOKENS), w1, &indexer, &mut wb); + KvEventMonitor::apply_stored(&stored_event(1, &TOKENS), w1, &indexer, &mut wb); + KvEventMonitor::apply_stored(&on_host, w1, &indexer, &mut wb); + assert_eq!( + wb.copies[&SequenceHash::from(1i64)], + Copies { device: 2, host: 1 } + ); + + // A host removal consumes the host copy only. + KvEventMonitor::apply_removed(&from_host, w1, &indexer, &mut wb); + assert!(routable(&indexer, w1, &hashes)); + // A second host removal has nothing to take and evicts nothing. + KvEventMonitor::apply_removed(&from_host, w1, &indexer, &mut wb); + assert!(routable(&indexer, w1, &hashes)); + assert_eq!(wb.counters.unknown_copy, 1); + + KvEventMonitor::apply_removed(&removed_event(1), w1, &indexer, &mut wb); + assert!(routable(&indexer, w1, &hashes), "one device copy left"); + KvEventMonitor::apply_removed(&removed_event(1), w1, &indexer, &mut wb); + assert!(!routable(&indexer, w1, &hashes)); + } + + #[test] + fn clearing_forgets_copy_counts() { + let indexer = KvIndex::positional(64); + let w1 = indexer.intern_worker("http://w1:8000").unwrap(); + let mut wb = WorkerIndexState::default(); + let hashes = compute_request_content_hashes(&TOKENS, 4); + + KvEventMonitor::apply_stored(&stored_event(1, &TOKENS), w1, &indexer, &mut wb); + KvEventMonitor::apply_stored(&stored_event(1, &TOKENS), w1, &indexer, &mut wb); + KvEventMonitor::apply_cleared(w1, &indexer, &mut wb); + assert!(wb.copies.is_empty()); + + KvEventMonitor::apply_stored(&stored_event(1, &TOKENS), w1, &indexer, &mut wb); + KvEventMonitor::apply_removed(&removed_event(1), w1, &indexer, &mut wb); + assert!( + !routable(&indexer, w1, &hashes), + "no count survives a clear" + ); + } + + #[test] + fn copy_counts_are_capped() { + let indexer = KvIndex::positional(64); + let w1 = indexer.intern_worker("http://w1:8000").unwrap(); + let mut wb = WorkerIndexState::default(); + let hashes = compute_request_content_hashes(&TOKENS, 4); + + // vLLM's `kv_cache_report_mode: full` re-announces a hit chain on + // every lookup without removals; the count stops at the cap. + for _ in 0..20 { + KvEventMonitor::apply_stored(&stored_event(1, &TOKENS), w1, &indexer, &mut wb); + } + assert_eq!(wb.copies[&SequenceHash::from(1i64)].device, COPIES_CAP); + assert_eq!(wb.counters.capped_copies, 20 - u64::from(COPIES_CAP)); + + for _ in 1..COPIES_CAP { + KvEventMonitor::apply_removed(&removed_event(1), w1, &indexer, &mut wb); + } + assert!(routable(&indexer, w1, &hashes)); + KvEventMonitor::apply_removed(&removed_event(1), w1, &indexer, &mut wb); + assert!( + !routable(&indexer, w1, &hashes), + "the cap bounds the extra removals" + ); + } +} diff --git a/model_gateway/src/worker/kv_event_monitor/subscription.rs b/model_gateway/src/worker/kv_event_monitor/subscription.rs new file mode 100644 index 0000000000..a5fb2a5f00 --- /dev/null +++ b/model_gateway/src/worker/kv_event_monitor/subscription.rs @@ -0,0 +1,734 @@ +//! The subscription task of one worker: connect, read the stream, reconnect +//! with backoff, hand the pushed load records on, and take the worker's +//! blocks out of the index when it leaves. + +use std::{ + sync::{Arc, Weak}, + time::{Duration, Instant}, +}; + +use dashmap::DashMap; +use smg_grpc_client::common_proto::{kv_cache_event, EngineLoad, KvEventBatch}; +use tokio::sync::{oneshot, Semaphore}; +use tracing::{debug, error, info, warn}; + +use super::{ + admission::{BatchOutcome, WorkerStreamState}, + apply::{WorkerIndexCounters, WorkerIndexState}, + KvEventMonitor, +}; +use crate::{ + observability::metrics::Metrics, + worker::{ + kv_event_recovery::ResyncReason, kv_index_backend::KvIndex, liveness, + monitor::WorkerMonitor, Worker, + }, +}; + +/// Initial reconnection delay after stream failure. +const INITIAL_RECONNECT_DELAY_MS: u64 = 100; + +/// Resolves when the worker is heard from (see `Worker::contact_wake`), or +/// never for a worker without a notifier. +async fn woken(wake: Option<&Arc>) { + match wake { + Some(notify) => notify.notified().await, + None => std::future::pending().await, + } +} + +/// How long a subscription call may take to answer with headers. A port that +/// accepts but does not serve yet (an engine still starting) hangs the call. +const SUBSCRIBE_DEADLINE: Duration = Duration::from_secs(2); + +/// A contact with the worker while a subscription call is pending (a poll +/// answered, a probe passed) says the worker serves now: a call still without +/// an answer this long after the contact is abandoned and retried. A live +/// server answers in milliseconds, so the contacts of a healthy worker (every +/// token it streams is one) never cut a call short. +const SUBSCRIBE_RETRY_GRACE: Duration = Duration::from_millis(500); + +/// The least a reconnect waits, contact or not: a server that keeps closing +/// the stream of a worker that is otherwise talking must not be hammered. +const RECONNECT_FLOOR: Duration = Duration::from_millis(INITIAL_RECONNECT_DELAY_MS); + +/// Maximum backoff between subscription attempts. Kept short: a worker that +/// restarts is healthy again within a few seconds, and until the stream is +/// back the blocks it stores are invisible to routing (the servicers resume +/// after the cursor and never resend them). A connect attempt is cheap. +const MAX_RECONNECT_DELAY_MS: u64 = 5_000; + +/// Positional-index cleanup is CPU-bound and can touch many blocks. Keep it +/// off Tokio workers and bound concurrent purges during fleet-wide drains. +pub(super) const MAX_CONCURRENT_INDEX_REMOVALS: usize = 4; +pub(super) static INDEX_REMOVAL_PERMITS: Semaphore = + Semaphore::const_new(MAX_CONCURRENT_INDEX_REMOVALS); + +/// Result of processing a stream connection to completion. +enum StreamResult { + /// Stream closed normally (server-side). + Ended, + /// Stream produced an error. + Error(tonic::Status), + /// Detected a gap in sequence numbers. + GapDetected { expected: u64, received: u64 }, +} + +impl KvEventMonitor { + /// Learn `block_size` from the first `KvBlock` in a stored event. + /// + /// Called once per model when the first stored event arrives, providing + /// ground truth from the backend. `CacheAwarePolicy` uses this to chunk + /// request tokens into blocks for overlap scoring. + /// + /// Overwrites any provisional value seeded from `WorkerSpec` since the + /// event stream reflects the backend's actual page size. + fn learn_block_size( + block_sizes: &DashMap, + model_id: &str, + learned: &mut bool, + batch: &KvEventBatch, + ) { + if *learned { + return; + } + for event in &batch.events { + if let Some(kv_cache_event::Data::Stored(stored)) = &event.data { + if let Some(block) = stored.blocks.first() { + if block.block_size > 0 { + let bs = block.block_size as usize; + block_sizes.insert(model_id.to_string(), bs); + info!( + model_id = %model_id, + block_size = bs, + "Learned block_size from KV event" + ); + *learned = true; + return; + } + } + } + } + } + + /// Take the worker's blocks out of the index when its subscription ends, + /// on the blocking pool and under a fleet-wide bound, and log what its + /// stream carried that the index did not take as is. + async fn remove_indexer_worker( + indexer: Arc, + worker_id: u32, + worker_url: &str, + worker_blocks: WorkerIndexState, + ) { + let Ok(permit) = INDEX_REMOVAL_PERMITS.acquire().await else { + error!(worker_id, "Positional-index cleanup semaphore closed"); + return; + }; + let WorkerIndexState { + blocks, counters, .. + } = worker_blocks; + if counters != WorkerIndexCounters::default() { + debug!( + worker_id, + ?counters, + "KV events the positional index did not take as is" + ); + } + let result = tokio::task::spawn_blocking(move || { + let _permit = permit; + indexer.remove_worker(worker_id, blocks); + }) + .await; + + if let Err(error) = result { + error!(worker_id, %error, "Positional-index worker cleanup task failed"); + } + Metrics::set_kv_index_blocks(worker_url, 0); + } + + /// Main subscription loop for a single worker. + /// + /// Owns the worker's index state and takes its blocks out of the index on + /// exit. Exits when `shutdown_rx` fires or the backend returns + /// `Unimplemented`. + pub(super) async fn subscription_loop( + worker: Arc, + worker_url: String, + indexer: Arc, + block_sizes: Arc>, + model_id: String, + mut shutdown_rx: oneshot::Receiver<()>, + load_sink: Option>, + ) { + let worker_id = match indexer.intern_worker(&worker_url) { + Ok(id) => id, + Err(e) => { + error!( + worker_url = %worker_url, + error = %e, + "Failed to intern worker; KV events from this worker will \ + not feed cache-aware routing" + ); + Metrics::record_kv_event_subscription_failure(&worker_url, "intern_failed"); + return; + } + }; + let mut state = WorkerStreamState::default(); + // A contact with the worker (a poll answered, a probe passed) ends the + // reconnect backoff early: a worker that is back gets its stream back + // at once instead of after the remaining delay. + let wake = worker.contact_wake(); + let mut reconnect_delay_ms = INITIAL_RECONNECT_DELAY_MS; + let mut block_size_learned = false; + + /// Sleep with shutdown check. Returns `true` if shutdown was signaled. + /// A contact with the worker ends the sleep early, but never before + /// `RECONNECT_FLOOR`. + macro_rules! sleep_or_shutdown { + ($delay:expr, $rx:expr) => {{ + let delay: Duration = $delay; + tokio::select! { + _ = tokio::time::sleep(delay) => false, + () = async { + tokio::time::sleep(delay.min(RECONNECT_FLOOR)).await; + woken(wake.as_ref()).await; + } => false, + _ = &mut *$rx => true, + } + }}; + } + + loop { + let backend_client = match worker.get_backend_client().await { + Ok(Some(client)) => client, + Ok(None) => { + // HTTP workers are filtered in on_worker_added, so this should + // be unreachable. Retry defensively rather than exiting and + // leaving a stale entry in worker_handles. + warn!( + worker_url = %worker_url, + delay_ms = reconnect_delay_ms, + "Worker has no backend client yet, retrying" + ); + if sleep_or_shutdown!( + Duration::from_millis(reconnect_delay_ms), + &mut shutdown_rx + ) { + Self::remove_indexer_worker( + Arc::clone(&indexer), + worker_id, + &worker_url, + state.index, + ) + .await; + return; + } + reconnect_delay_ms = (reconnect_delay_ms * 2).min(MAX_RECONNECT_DELAY_MS); + continue; + } + Err(e) => { + warn!( + worker_url = %worker_url, + error = %e, + delay_ms = reconnect_delay_ms, + "Failed to get backend client, retrying" + ); + if sleep_or_shutdown!( + Duration::from_millis(reconnect_delay_ms), + &mut shutdown_rx + ) { + Self::remove_indexer_worker( + Arc::clone(&indexer), + worker_id, + &worker_url, + state.index, + ) + .await; + return; + } + reconnect_delay_ms = (reconnect_delay_ms * 2).min(MAX_RECONNECT_DELAY_MS); + continue; + } + }; + + let start_seq = state.resume_sequence(); + // A live server answers a subscription with headers at once; one + // that does not within the deadline is unreachable (a half-open + // connection behind a partition), and must not hang the loop. + let subscribed = tokio::select! { + attempt = tokio::time::timeout( + SUBSCRIBE_DEADLINE, + backend_client.subscribe_kv_events(start_seq), + ) => attempt.unwrap_or_else(|_| { + Err(tonic::Status::unavailable(format!( + "SubscribeKvEvents did not answer within {SUBSCRIBE_DEADLINE:?}" + ))) + }), + // The worker was heard from on another path (a poll, a probe) + // and this call still hangs: a fresh one will land. + () = async { + woken(wake.as_ref()).await; + tokio::time::sleep(SUBSCRIBE_RETRY_GRACE).await; + } => continue, + // The worker left while the call was pending. Without this arm + // a call that never answers, retried at every contact, would + // keep the task alive past the worker's removal. + _ = &mut shutdown_rx => { + Self::remove_indexer_worker( + Arc::clone(&indexer), + worker_id, + &worker_url, + state.index, + ) + .await; + return; + } + }; + let stream = match subscribed { + Ok(stream) => { + info!( + worker_url = %worker_url, + start_seq, + "KV event stream connected" + ); + Metrics::record_kv_event_subscription(&worker_url); + reconnect_delay_ms = INITIAL_RECONNECT_DELAY_MS; + state.reconnected(); + liveness::on_contact(&worker); + stream + } + Err(e) => { + // If the backend doesn't implement SubscribeKvEvents (e.g. vLLM), + // stop retrying — this RPC will never succeed. + if e.code() == tonic::Code::Unimplemented { + warn!( + worker_url = %worker_url, + "Backend does not implement SubscribeKvEvents, \ + disabling KV event subscription for this worker" + ); + Self::remove_indexer_worker( + Arc::clone(&indexer), + worker_id, + &worker_url, + state.index, + ) + .await; + return; + } + if e.code() == tonic::Code::OutOfRange { + warn!( + worker_url = %worker_url, + start_seq, + "KV event replay cursor expired; clearing worker state and requesting a current snapshot" + ); + Self::reset_worker( + &indexer, + worker_id, + &mut state, + &worker_url, + ResyncReason::OutOfRange, + ); + reconnect_delay_ms = INITIAL_RECONNECT_DELAY_MS; + continue; + } + if liveness::is_transport_failure(e.code()) { + liveness::on_contact_failed(&worker, "kv subscribe"); + } + warn!( + worker_url = %worker_url, + error = %e, + delay_ms = reconnect_delay_ms, + "Failed to subscribe to KV events, retrying" + ); + if sleep_or_shutdown!( + Duration::from_millis(reconnect_delay_ms), + &mut shutdown_rx + ) { + Self::remove_indexer_worker( + Arc::clone(&indexer), + worker_id, + &worker_url, + state.index, + ) + .await; + return; + } + reconnect_delay_ms = (reconnect_delay_ms * 2).min(MAX_RECONNECT_DELAY_MS); + continue; + } + }; + + let on_batch = |batch: &KvEventBatch| { + liveness::on_contact(&worker); + Self::learn_block_size(&block_sizes, &model_id, &mut block_size_learned, batch); + }; + // The load record on a batch is a poll of this worker, received now. + let on_load = |batch: &KvEventBatch, load: &EngineLoad| { + liveness::on_contact(&worker); + if let Some(monitor) = load_sink.as_ref().and_then(Weak::upgrade) { + monitor.apply_pushed_load( + &worker, + batch.dp_rank.unwrap_or(0), + load, + Instant::now(), + ); + } + }; + let stream_result = tokio::select! { + result = Self::process_stream( + stream, &worker_url, worker_id, &indexer, &mut state, on_batch, on_load, + ) => result, + _ = &mut shutdown_rx => { + Self::remove_indexer_worker( + Arc::clone(&indexer), + worker_id, + &worker_url, + state.index, + ) + .await; + return; + } + }; + + if state.abandon_snapshot() { + warn!( + worker_url = %worker_url, + "KV event stream ended during a relay snapshot; the next subscription \ + starts over from zero" + ); + } + match stream_result { + StreamResult::Ended => { + info!( + worker_url = %worker_url, + resume_from = state.resume_sequence(), + delay_ms = reconnect_delay_ms, + "KV event stream ended, reconnecting" + ); + // Backoff to avoid tight reconnect loop if server keeps + // closing the stream cleanly (e.g., rolling connections). + if sleep_or_shutdown!( + Duration::from_millis(reconnect_delay_ms), + &mut shutdown_rx + ) { + Self::remove_indexer_worker( + Arc::clone(&indexer), + worker_id, + &worker_url, + state.index, + ) + .await; + return; + } + reconnect_delay_ms = (reconnect_delay_ms * 2).min(MAX_RECONNECT_DELAY_MS); + } + StreamResult::Error(e) => { + if e.code() == tonic::Code::DataLoss { + warn!( + worker_url = %worker_url, + error = %e, + resume_from = state.resume_sequence(), + "KV event subscriber fell behind; clearing worker state and requesting a current snapshot" + ); + Self::reset_worker( + &indexer, + worker_id, + &mut state, + &worker_url, + ResyncReason::DataLoss, + ); + reconnect_delay_ms = INITIAL_RECONNECT_DELAY_MS; + continue; + } + if liveness::is_transport_failure(e.code()) { + liveness::on_contact_failed(&worker, "kv stream"); + } + warn!( + worker_url = %worker_url, + error = %e, + resume_from = state.resume_sequence(), + delay_ms = reconnect_delay_ms, + "KV event stream error, reconnecting" + ); + if sleep_or_shutdown!( + Duration::from_millis(reconnect_delay_ms), + &mut shutdown_rx + ) { + Self::remove_indexer_worker( + Arc::clone(&indexer), + worker_id, + &worker_url, + state.index, + ) + .await; + return; + } + reconnect_delay_ms = (reconnect_delay_ms * 2).min(MAX_RECONNECT_DELAY_MS); + } + StreamResult::GapDetected { expected, received } => { + warn!( + worker_url = %worker_url, + expected, + received, + "Sequence gap detected, reconnecting for replay from seq {expected}" + ); + // No backoff: gap replay is a normal recovery path, and the + // rank state asks for it once; if the server skips ahead + // again the gap is settled instead of retried. + } + } + } + } + + /// Process batches from a single stream connection. + async fn process_stream( + mut stream: tonic::Streaming, + worker_url: &str, + worker_id: u32, + indexer: &KvIndex, + state: &mut WorkerStreamState, + mut on_batch: impl FnMut(&KvEventBatch), + mut on_load: impl FnMut(&KvEventBatch, &EngineLoad), + ) -> StreamResult { + use tokio_stream::StreamExt; + + while let Some(result) = stream.next().await { + let batch = match result { + Ok(batch) => batch, + Err(e) => return StreamResult::Error(e), + }; + if let Some(load) = &batch.load { + on_load(&batch, load); + if Self::is_load_only(&batch) { + Metrics::record_kv_event_batch(worker_url, "load_only"); + continue; + } + } + if let BatchOutcome::Gap { expected, received } = + Self::admit_batch(&batch, worker_url, worker_id, indexer, state, &mut on_batch) + { + return StreamResult::GapDetected { expected, received }; + } + } + + StreamResult::Ended + } + + /// A batch that carries only a load record (`EngineLoad.load_only`): no + /// events, the last sequence repeated; it never enters admission. + pub(super) fn is_load_only(batch: &KvEventBatch) -> bool { + batch.load.as_ref().is_some_and(|load| load.load_only) + } +} + +#[cfg(test)] +mod tests { + use std::{net::SocketAddr, pin::Pin}; + + use futures::Stream; + use kv_index::{compute_content_hash, SequenceHash, StoredBlock}; + use openai_protocol::worker::{ConnectionMode, HealthCheckConfig, RuntimeType, WorkerType}; + use smg_grpc_client::{ + common_proto as common, + tokenspeed_scheduler::tokenspeed_proto::{ + self as ts, + token_speed_scheduler_server::{TokenSpeedScheduler, TokenSpeedSchedulerServer}, + }, + }; + use tonic::{transport::Server, Request, Response, Status}; + + use super::*; + use crate::worker::BasicWorkerBuilder; + + /// A scheduler whose `SubscribeKvEvents` accepts the call and never + /// answers it, as an engine still starting behind an open port does. + struct HangingScheduler; + + type Never = Pin> + Send>>; + + #[tonic::async_trait] + impl TokenSpeedScheduler for HangingScheduler { + type GenerateStream = Never; + type SubscribeKvEventsStream = Never; + type GetTokenizerStream = Never; + + async fn generate( + &self, + _: Request, + ) -> Result, Status> { + Err(Status::unimplemented("test scheduler")) + } + + async fn health_check( + &self, + _: Request, + ) -> Result, Status> { + Err(Status::unimplemented("test scheduler")) + } + + async fn abort( + &self, + _: Request, + ) -> Result, Status> { + Err(Status::unimplemented("test scheduler")) + } + + async fn get_model_info( + &self, + _: Request, + ) -> Result, Status> { + Err(Status::unimplemented("test scheduler")) + } + + async fn get_server_info( + &self, + _: Request, + ) -> Result, Status> { + Err(Status::unimplemented("test scheduler")) + } + + async fn get_loads( + &self, + _: Request, + ) -> Result, Status> { + Err(Status::unimplemented("test scheduler")) + } + + async fn subscribe_kv_events( + &self, + _: Request, + ) -> Result, Status> { + std::future::pending().await + } + + async fn flush_cache( + &self, + _: Request, + ) -> Result, Status> { + Err(Status::unimplemented("test scheduler")) + } + + async fn start_profile( + &self, + _: Request, + ) -> Result, Status> { + Err(Status::unimplemented("test scheduler")) + } + + async fn stop_profile( + &self, + _: Request, + ) -> Result, Status> { + Err(Status::unimplemented("test scheduler")) + } + + async fn get_tokenizer( + &self, + _: Request, + ) -> Result, Status> { + Err(Status::unimplemented("test scheduler")) + } + } + + async fn wait_until_listening(addr: SocketAddr) { + for _ in 0..100 { + if tokio::net::TcpStream::connect(addr).await.is_ok() { + return; + } + tokio::time::sleep(Duration::from_millis(20)).await; + } + panic!("test scheduler did not start listening on {addr}"); + } + + fn grpc_worker(addr: SocketAddr) -> Arc { + Arc::new( + BasicWorkerBuilder::new(format!("grpc://{addr}")) + .worker_type(WorkerType::Regular) + .connection_mode(ConnectionMode::Grpc) + .runtime_type(RuntimeType::TokenSpeed) + .health_config(HealthCheckConfig { + disable_health_check: true, + ..Default::default() + }) + .build(), + ) + } + + /// A subscribe call that never answers, retried at every contact with + /// the worker, must still see the worker's removal: with no shutdown arm + /// on the call the task lived on and `on_worker_removed` never returned. + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + async fn removal_ends_a_subscribe_call_that_never_answers() { + let port = portpicker::pick_unused_port().expect("a free port"); + let addr: SocketAddr = format!("127.0.0.1:{port}").parse().unwrap(); + #[expect( + clippy::disallowed_methods, + reason = "the test scheduler lives as long as the test" + )] + let server = tokio::spawn( + Server::builder() + .add_service(TokenSpeedSchedulerServer::new(HangingScheduler)) + .serve(addr), + ); + wait_until_listening(addr).await; + + let worker = grpc_worker(addr); + let monitor = KvEventMonitor::new(None); + monitor.on_worker_added(&worker).await; + // Let the task connect and park in the subscribe call. + tokio::time::sleep(Duration::from_millis(500)).await; + // The worker keeps being heard from on other paths; every contact + // retries the parked call well inside its deadline. + let pinger = { + let worker = Arc::clone(&worker); + #[expect( + clippy::disallowed_methods, + reason = "the contact source is stopped at the end of the test" + )] + tokio::spawn(async move { + loop { + tokio::time::sleep(Duration::from_millis(100)).await; + worker.note_contact(); + } + }) + }; + + tokio::time::timeout( + Duration::from_secs(5), + monitor.on_worker_removed(worker.url()), + ) + .await + .expect("removal returns while the subscribe call hangs"); + + pinger.abort(); + server.abort(); + } + + #[tokio::test] + async fn a_removed_workers_blocks_leave_the_index_off_the_runtime() { + let indexer = Arc::new(KvIndex::positional(64)); + let worker_id = indexer.intern_worker("http://w1:8000").unwrap(); + let mut worker_blocks = WorkerIndexState::default(); + indexer + .apply_stored( + worker_id, + &[StoredBlock { + seq_hash: SequenceHash(1), + content_hash: compute_content_hash(&[1, 2, 3]), + }], + None, + &mut worker_blocks.blocks, + ) + .unwrap(); + + KvEventMonitor::remove_indexer_worker( + Arc::clone(&indexer), + worker_id, + "grpc://w1:9000", + worker_blocks, + ) + .await; + + assert_eq!(indexer.current_size(), 0); + } +} diff --git a/model_gateway/src/worker/kv_event_recovery.rs b/model_gateway/src/worker/kv_event_recovery.rs new file mode 100644 index 0000000000..2a0ca6334c --- /dev/null +++ b/model_gateway/src/worker/kv_event_recovery.rs @@ -0,0 +1,581 @@ +//! Admission state for one KV-event stream rank: cursors, gap recovery and +//! the live tail held during an out-of-band resync. +//! +//! Every engine publisher (one per data-parallel rank) numbers its batches +//! with its own monotonic sequence. The subscriber keeps one [`RankState`] per +//! `(worker, dp_rank)` and asks it what to do with each batch it receives: +//! +//! - contiguous and first batches are applied; +//! - duplicates (a replay overlapping what was already applied) are skipped; +//! - a publisher restart clears the rank's state and applies the batch: the +//! engine's cache is empty again. It shows as a sequence below the cursor on +//! a fresh connection (servers resume strictly after the cursor, so nothing +//! legitimate sits there), a cache clear carried below the cursor (SGLang's +//! first batch after startup), a counter at its start (0 or 1) below the +//! cursor, or a sequence far below the cursor mid-stream; +//! - a gap is answered once by asking the server to replay from the expected +//! sequence; if the next batch still skips ahead the server kept no history +//! (the Rust relay never replays, the SGLang servicer only within its +//! buffer) and the rank is marked degraded: a small gap keeps the existing +//! state (an engine still holds most of those blocks and will never re-send +//! stores for them), a large one clears it, mirroring the servicers' own +//! `OUT_OF_RANGE` / `DATA_LOSS` signal. +//! +//! The Rust relay serves a state snapshot in band: once its history no longer +//! starts at the publisher's first batch, a subscription from zero begins with +//! chunks marked `KvSnapshotChunk` (chunk 0 carries the engine's clear, the +//! stores follow parents first) stamped with consecutive sequence numbers +//! ending at the sequence the state was cut at, and live events continue +//! from the next one. The monitor applies the chunks as a `snapshot` resync +//! outside the admission rules above and sets the rank's cursor to each +//! chunk's stamp ([`RankState::resync_to`]); nothing is buffered because the +//! stream is ordered. +//! +//! The tail buffer serves a snapshot resync delivered on a side channel: live +//! batches arriving while the snapshot is in flight are held, bounded, and +//! applied in order after it. No servicer sends one out of band today, so the +//! monitor never enters that mode; the logic is here, tested, for the +//! protocol in `crates/kv_index/docs/recovery-protocol.md`. + +use std::collections::VecDeque; + +use smg_grpc_client::common_proto::KvEventBatch; + +/// A sequence this far below the cursor is a restarted publisher, not a late +/// duplicate: replays start at the sequence we asked for, so legitimate +/// duplicates sit just below the cursor. +pub(crate) const RESTART_WINDOW: u64 = 1_024; + +/// Missed batches beyond this clear the rank instead of keeping its state: +/// the same decision the servicers make when their replay buffer overflows. +const GAP_CLEAR_THRESHOLD: u64 = 1_024; + +/// Live batches held while a snapshot resync is in flight. +const TAIL_LIMIT: usize = 1_024; + +/// Where the rank's cursor stands. +#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)] +pub(crate) enum Cursor { + /// Nothing admitted yet: the first batch sets the cursor wherever it is. + #[default] + Initial, + /// The last sequence number applied. + Live(u64), +} + +/// A replay the rank asked the server for. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub(crate) struct PendingReplay { + /// The sequence the server was asked to resume from. + expected: u64, + /// The sequence that revealed the gap. + received: u64, +} + +/// What the subscriber must do with a batch. +#[derive(Debug, PartialEq, Eq)] +pub(crate) enum Admission { + /// Apply the batch's events; the cursor advanced to it. + Apply, + /// A duplicate of something already applied; drop it. + Stale, + /// The publisher restarted: clear the rank's blocks, then apply. + Restart, + /// A gap with no replay asked yet: reconnect with `expected` as the start + /// sequence; the batch itself is not applied (the replay will resend it). + Replay { expected: u64 }, + /// A gap the server could not fill: `missed` batches are lost. The rank's + /// state is kept (`cleared == false`) or dropped (`cleared == true`); then + /// apply the batch. + Unrecovered { missed: u64, cleared: bool }, + /// A snapshot resync is in flight: queue the batch in the tail. + Buffered, +} + +/// How the rank got where it is, for metrics. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub(crate) enum ResyncReason { + OutOfRange, + DataLoss, + PublisherRestart, + GapCleared, + /// The relay served a state snapshot in place of the history it no longer + /// had: the worker's state is replaced by the live set it carries. + Snapshot, +} + +impl ResyncReason { + pub(crate) fn as_str(self) -> &'static str { + match self { + Self::OutOfRange => "out_of_range", + Self::DataLoss => "data_loss", + Self::PublisherRestart => "publisher_restart", + Self::GapCleared => "gap_cleared", + Self::Snapshot => "snapshot", + } + } +} + +/// Admission state for one `(worker, dp_rank)`. +#[derive(Debug, Default)] +pub(crate) struct RankState { + cursor: Cursor, + replay: Option, + /// Set when a gap could not be recovered; cleared by a resync. + degraded: bool, + /// Live batches held during a snapshot resync, oldest first. + tail: VecDeque, + snapshot_inflight: bool, + tail_overflowed: bool, + /// No batch received yet on the current connection. + fresh: bool, +} + +impl RankState { + #[cfg(test)] + pub(crate) fn cursor(&self) -> Cursor { + self.cursor + } + + /// The cursor to send when (re)subscribing: the last sequence applied, or + /// 0 for none. Servers resume strictly after it (the SGLang servicer asks + /// its engine for `cursor + 1`, the mock engine streams `> cursor`), so a + /// pending replay needs nothing more: its `expected` is `cursor + 1`. + pub(crate) fn resume_from(&self) -> u64 { + match self.cursor { + Cursor::Live(last) => last, + Cursor::Initial => 0, + } + } + + /// A new stream connection was made. The first batch on it shows where + /// the server stands: below the cursor there means a restarted publisher, + /// not a late duplicate. + pub(crate) fn reconnected(&mut self) { + self.fresh = true; + } + + pub(crate) fn is_degraded(&self) -> bool { + self.degraded + } + + #[cfg(test)] + pub(crate) fn replay_pending(&self) -> Option { + self.replay + } + + pub(crate) fn tail_len(&self) -> usize { + self.tail.len() + } + + /// Decide what to do with a batch carrying `seq`; `clears` says whether + /// it carries the engine's own cache clear. + pub(crate) fn admit(&mut self, seq: u64, clears: bool) -> Admission { + let fresh = std::mem::replace(&mut self.fresh, false); + if self.snapshot_inflight { + return Admission::Buffered; + } + match self.cursor { + Cursor::Initial => { + self.cursor = Cursor::Live(seq); + self.replay = None; + Admission::Apply + } + Cursor::Live(last) if seq == last + 1 => { + self.cursor = Cursor::Live(seq); + self.replay = None; + Admission::Apply + } + Cursor::Live(last) if seq <= last => { + // A counter at its start is a new publisher too: replays + // never reach below `cursor + 1`, so 0 or 1 under a cursor of + // 2 or more cannot be a late duplicate. + let restarted = clears + || (fresh && seq < last) + || (seq <= 1 && seq < last) + || seq + RESTART_WINDOW <= last; + if restarted { + // A fresh publisher counting from the start again. + self.cursor = Cursor::Live(seq); + self.replay = None; + self.degraded = false; + Admission::Restart + } else { + Admission::Stale + } + } + Cursor::Live(last) => { + let expected = last + 1; + match self.replay { + None => { + self.replay = Some(PendingReplay { + expected, + received: seq, + }); + Admission::Replay { expected } + } + Some(pending) => { + // We already asked to resume from `expected` and the + // server skipped ahead anyway: no history there. + debug_assert_eq!(pending.expected, expected); + let missed = seq - expected; + let cleared = missed > GAP_CLEAR_THRESHOLD; + self.replay = None; + self.degraded = !cleared; + self.cursor = Cursor::Live(seq); + Admission::Unrecovered { missed, cleared } + } + } + } + } + } + + /// Hold a live batch while a snapshot is in flight. Returns `false` when + /// the tail is full; the batch is dropped and the resync must be redone. + pub(crate) fn buffer_live(&mut self, batch: KvEventBatch) -> bool { + if self.tail.len() >= TAIL_LIMIT { + self.tail_overflowed = true; + return false; + } + self.tail.push_back(batch); + true + } + + /// Enter snapshot mode: live batches are buffered until + /// [`finish_snapshot`](Self::finish_snapshot). + #[cfg_attr( + not(test), + expect( + dead_code, + reason = "entered once a servicer serves snapshots out of band (docs/recovery-protocol.md)" + ) + )] + pub(crate) fn begin_snapshot(&mut self) { + self.snapshot_inflight = true; + self.tail.clear(); + self.tail_overflowed = false; + self.replay = None; + } + + /// The snapshot (complete through `through_seq`) has been applied. Returns + /// the buffered live batches that come after it, in order and without + /// duplicates, and sets the cursor to the last of them. `None` means the + /// tail overflowed while waiting and the snapshot has to be taken again. + #[cfg_attr( + not(test), + expect( + dead_code, + reason = "entered once a servicer serves snapshots out of band (docs/recovery-protocol.md)" + ) + )] + pub(crate) fn finish_snapshot(&mut self, through_seq: u64) -> Option> { + self.snapshot_inflight = false; + self.degraded = false; + if self.tail_overflowed { + self.tail.clear(); + self.tail_overflowed = false; + self.cursor = Cursor::Live(through_seq); + return None; + } + let mut cursor = through_seq; + let mut out = Vec::new(); + for batch in self.tail.drain(..) { + if batch.sequence_number <= cursor { + continue; + } + cursor = batch.sequence_number; + out.push(batch); + } + self.cursor = Cursor::Live(cursor); + Some(out) + } + + /// A snapshot chunk stamped `seq` was applied: the cursor stands there, + /// whatever it was, and nothing is pending or degraded any more. + pub(crate) fn resync_to(&mut self, seq: u64) { + self.cursor = Cursor::Live(seq); + self.replay = None; + self.degraded = false; + self.fresh = false; + self.tail.clear(); + self.snapshot_inflight = false; + self.tail_overflowed = false; + } + + /// The relay's snapshot could not cover the engine's whole life + /// (`KvSnapshotChunk.unknown_before`): the rank's blocks are a partial + /// set until the next resync, as after an unrecovered gap. + pub(crate) fn mark_degraded(&mut self) { + self.degraded = true; + } + + /// The server declared its history gone (`OUT_OF_RANGE` / `DATA_LOSS`) or + /// the subscriber decided to drop the rank: forget the cursor so the next + /// stream is taken from wherever it starts. + pub(crate) fn reset(&mut self) { + self.cursor = Cursor::Initial; + self.replay = None; + self.degraded = false; + self.tail.clear(); + self.snapshot_inflight = false; + self.tail_overflowed = false; + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn batch(seq: u64) -> KvEventBatch { + KvEventBatch { + sequence_number: seq, + timestamp: 0.0, + events: vec![], + dp_rank: None, + snapshot: None, + load: None, + } + } + + #[test] + fn first_batch_sets_the_cursor_wherever_it_is() { + let mut rank = RankState::default(); + assert_eq!(rank.admit(17, false), Admission::Apply); + assert_eq!(rank.cursor(), Cursor::Live(17)); + assert_eq!(rank.resume_from(), 17); + } + + #[test] + fn contiguous_batches_apply_and_duplicates_are_stale() { + let mut rank = RankState::default(); + for seq in 1..=5 { + assert_eq!(rank.admit(seq, false), Admission::Apply); + } + assert_eq!(rank.admit(4, false), Admission::Stale); + assert_eq!(rank.admit(5, false), Admission::Stale); + assert_eq!(rank.admit(6, false), Admission::Apply); + assert_eq!(rank.cursor(), Cursor::Live(6)); + } + + /// A servicer that pushes load records sends `load_only` heartbeats + /// with no events and the last sequence it sent. A gateway that predates + /// the field (prost drops it) sees a repeat of its cursor: a late + /// duplicate, dropped. It is never a restart (that needs a clear, a + /// fresh connection below the cursor, a count back at 0 or 1, or a + /// sequence a window below) and never a gap, however many arrive. + #[test] + fn a_repeat_of_the_cursor_is_stale_never_a_restart_or_a_gap() { + let mut rank = RankState::default(); + for seq in 0..=7 { + assert_eq!(rank.admit(seq, false), Admission::Apply); + } + for _ in 0..50 { + assert_eq!(rank.admit(7, false), Admission::Stale); + } + assert_eq!(rank.cursor(), Cursor::Live(7)); + assert!(rank.replay_pending().is_none()); + assert_eq!(rank.admit(8, false), Admission::Apply); + // After a reconnect too: a repeat of the cursor is not "below" it. + rank.reconnected(); + assert_eq!(rank.admit(8, false), Admission::Stale); + assert_eq!(rank.admit(9, false), Admission::Apply); + } + + #[test] + fn a_gap_asks_for_one_replay_then_resumes_when_the_server_fills_it() { + let mut rank = RankState::default(); + for seq in 1..=5 { + rank.admit(seq, false); + } + assert_eq!(rank.admit(8, false), Admission::Replay { expected: 6 }); + assert_eq!(rank.resume_from(), 5); + assert!(rank.replay_pending().is_some()); + // The server resends from 6: everything is contiguous again. + assert_eq!(rank.admit(6, false), Admission::Apply); + assert!(rank.replay_pending().is_none()); + assert_eq!(rank.admit(7, false), Admission::Apply); + assert_eq!(rank.admit(8, false), Admission::Apply); + assert!(!rank.is_degraded()); + } + + #[test] + fn a_gap_the_server_does_not_fill_keeps_state_and_marks_the_rank_degraded() { + let mut rank = RankState::default(); + for seq in 1..=5 { + rank.admit(seq, false); + } + assert_eq!(rank.admit(9, false), Admission::Replay { expected: 6 }); + // Reconnected from 6, but the server streams live from 9 again. + assert_eq!( + rank.admit(9, false), + Admission::Unrecovered { + missed: 3, + cleared: false + } + ); + assert!(rank.is_degraded()); + assert_eq!(rank.cursor(), Cursor::Live(9)); + assert_eq!(rank.admit(10, false), Admission::Apply); + // Never loops: the next gap starts a new, single replay attempt. + assert_eq!(rank.admit(12, false), Admission::Replay { expected: 11 }); + } + + #[test] + fn a_large_unfilled_gap_clears_the_rank() { + let mut rank = RankState::default(); + rank.admit(1, false); + let far = 2 + GAP_CLEAR_THRESHOLD + 1; + assert_eq!(rank.admit(far, false), Admission::Replay { expected: 2 }); + assert_eq!( + rank.admit(far, false), + Admission::Unrecovered { + missed: far - 2, + cleared: true + } + ); + assert!(!rank.is_degraded()); + assert_eq!(rank.cursor(), Cursor::Live(far)); + } + + #[test] + fn a_sequence_far_below_the_cursor_is_a_publisher_restart() { + let mut rank = RankState::default(); + for seq in 1..=3_000 { + rank.admit(seq, false); + } + assert_eq!(rank.admit(2_999, false), Admission::Stale); + assert_eq!( + rank.admit(3_000 - RESTART_WINDOW + 1, false), + Admission::Stale + ); + assert_eq!( + rank.admit(3_000 - RESTART_WINDOW, false), + Admission::Restart + ); + assert_eq!(rank.cursor(), Cursor::Live(3_000 - RESTART_WINDOW)); + let mut fresh = RankState::default(); + for seq in 1..=3_000 { + fresh.admit(seq, false); + } + assert_eq!(fresh.admit(1, false), Admission::Restart); + assert_eq!(fresh.admit(2, false), Admission::Apply); + } + + #[test] + fn a_lower_sequence_on_a_fresh_connection_is_a_restart() { + let mut rank = RankState::default(); + for seq in 1..=5 { + rank.admit(seq, false); + } + rank.reconnected(); + assert_eq!(rank.admit(1, false), Admission::Restart); + assert_eq!(rank.admit(2, false), Admission::Apply); + assert_eq!(rank.admit(3, false), Admission::Apply); + // Mid-stream, a sequence just below the cursor is a duplicate. + assert_eq!(rank.admit(2, false), Admission::Stale); + // A replay that starts exactly at the cursor is a duplicate too. + rank.reconnected(); + assert_eq!(rank.admit(3, false), Admission::Stale); + assert_eq!(rank.admit(4, false), Admission::Apply); + // A gap on a fresh connection is still a gap. + rank.reconnected(); + assert_eq!(rank.admit(9, false), Admission::Replay { expected: 5 }); + } + + #[test] + fn a_counter_at_its_start_below_the_cursor_is_a_restart() { + let mut rank = RankState::default(); + for seq in 1..=5 { + rank.admit(seq, false); + } + // vLLM and SGLang count from 0, the mock engine from 1; neither + // number can be a replayed duplicate under a cursor of 5. + assert_eq!(rank.admit(0, false), Admission::Restart); + assert_eq!(rank.admit(1, false), Admission::Apply); + for seq in 2..=5 { + rank.admit(seq, false); + } + assert_eq!(rank.admit(1, false), Admission::Restart); + assert_eq!(rank.cursor(), Cursor::Live(1)); + } + + #[test] + fn a_clear_below_the_cursor_is_a_restart() { + let mut rank = RankState::default(); + for seq in 1..=5 { + rank.admit(seq, false); + } + assert_eq!(rank.admit(0, true), Admission::Restart); + assert_eq!(rank.cursor(), Cursor::Live(0)); + assert_eq!(rank.admit(1, false), Admission::Apply); + // A clear at or above the cursor is an ordinary event. + assert_eq!(rank.admit(2, true), Admission::Apply); + } + + #[test] + fn a_snapshot_chunk_moves_the_cursor_wherever_it_is_stamped() { + let mut rank = RankState::default(); + for seq in 1..=5 { + rank.admit(seq, false); + } + assert_eq!(rank.admit(9, false), Admission::Replay { expected: 6 }); + rank.resync_to(40); + assert_eq!(rank.cursor(), Cursor::Live(40)); + assert!(rank.replay_pending().is_none()); + assert_eq!(rank.admit(41, false), Admission::Apply); + // Below the cursor as well: a stale cursor is replaced, not restarted. + rank.reconnected(); + rank.resync_to(3); + assert_eq!(rank.admit(4, false), Admission::Apply); + assert_eq!(rank.resume_from(), 4); + } + + #[test] + fn reset_forgets_the_cursor() { + let mut rank = RankState::default(); + rank.admit(40, false); + rank.admit(42, false); + rank.reset(); + assert_eq!(rank.cursor(), Cursor::Initial); + assert_eq!(rank.resume_from(), 0); + assert_eq!(rank.admit(7, false), Admission::Apply); + } + + #[test] + fn snapshot_mode_buffers_live_batches_and_replays_the_tail_in_order() { + let mut rank = RankState::default(); + for seq in 1..=5 { + rank.admit(seq, false); + } + rank.begin_snapshot(); + assert_eq!(rank.admit(6, false), Admission::Buffered); + assert!(rank.buffer_live(batch(6))); + assert!(rank.buffer_live(batch(7))); + assert!(rank.buffer_live(batch(7))); + assert!(rank.buffer_live(batch(8))); + assert_eq!(rank.tail_len(), 4); + // The snapshot covered everything through 6. + let tail = rank.finish_snapshot(6).expect("tail intact"); + let seqs: Vec = tail.iter().map(|b| b.sequence_number).collect(); + assert_eq!(seqs, vec![7, 8]); + assert_eq!(rank.cursor(), Cursor::Live(8)); + assert_eq!(rank.admit(9, false), Admission::Apply); + assert!(!rank.is_degraded()); + } + + #[test] + fn a_tail_overflow_is_bounded_and_forces_another_snapshot() { + let mut rank = RankState::default(); + rank.admit(1, false); + rank.begin_snapshot(); + for seq in 2..=(TAIL_LIMIT as u64 + 1) { + assert!(rank.buffer_live(batch(seq))); + } + assert_eq!(rank.tail_len(), TAIL_LIMIT); + assert!(!rank.buffer_live(batch(TAIL_LIMIT as u64 + 2))); + assert_eq!(rank.tail_len(), TAIL_LIMIT); + assert!(rank.finish_snapshot(1).is_none()); + assert_eq!(rank.tail_len(), 0); + assert_eq!(rank.cursor(), Cursor::Live(1)); + } +} diff --git a/model_gateway/src/worker/kv_index_backend.rs b/model_gateway/src/worker/kv_index_backend.rs new file mode 100644 index 0000000000..f29c282440 --- /dev/null +++ b/model_gateway/src/worker/kv_index_backend.rs @@ -0,0 +1,701 @@ +//! The event-driven KV index behind cache-aware routing, as one type over two +//! implementations: the [`PositionalIndexer`] the gateway has routed with so +//! far and the chain index ([`ShardedChainIndex`], one shard here: chains stored as runs), +//! selected at startup by `--kv-index {positional,chain}` ([`KvIndexKind`]). +//! +//! Both indexers are fed the same engine events and answer the same question +//! (how many leading blocks of a request each worker holds), and their +//! write-path methods already share a shape: every call takes a caller-owned +//! per-worker reverse map. What differs is the map's type, so [`WorkerBlocks`] +//! carries whichever map the index needs, created on first use. The monitor and +//! the policy see only [`KvIndex`] and [`WorkerBlocks`]. +//! +//! An enum rather than a trait object: the lookup, the one call on the request +//! path, dispatches on a discriminant the branch predictor learns at startup, +//! and each variant's method is called directly, so the positional path costs +//! what it did before this type existed. +//! +//! Under `cfg(test)` a third variant wraps the crate's single-threaded +//! [`ReferenceIndexer`](kv_index::ReferenceIndexer), so the exactness tests +//! feed the production apply path to all three and compare answers. + +use std::{collections::BTreeSet, fmt}; + +use kv_index::{ + ApplyError, ChainBlockMap, ChainIndexStats, ContentHash, OverlapScores, PositionalIndexer, + PruneStats, SequenceHash, ShardedChainIndex, StoredBlock, WorkerBlockMap, WorkerIdExhausted, +}; + +pub use crate::config::KvIndexKind; + +/// Workers one chain index interns at once: the index's own ceiling per shard. +/// Ids are handed back when a worker is removed, so this bounds the workers +/// of one model that stream events at the same time, not the churn over a +/// lifetime. Coverage costs one bit per worker per run, rounded up to 64, +/// and the lookup ANDs that many words per run on the matched path. +const CHAIN_INDEX_MAX_WORKERS: usize = 1024; + +/// Shards of the chain index: one. The sharded type keeps one chain index per +/// shard and merges lookups across them, so that a gateway whose runtime +/// spans both sockets can later place each worker's index on the socket its +/// event lane runs on; until that placement exists, one shard is the plain +/// chain index with the same ids. +const CHAIN_INDEX_SHARDS: usize = 1; + +/// The gateway's KV index: one per model, shared by the model's workers' +/// event subscriptions (writers) and the cache-aware policy (readers). +#[expect( + clippy::large_enum_variant, + reason = "one value per model behind an Arc: boxing the positional indexer would put a \ + second pointer hop on the lookup path to save bytes nobody pays for" +)] +pub enum KvIndex { + /// One index entry per `(position, content hash)`, probed per block. + Positional(PositionalIndexer), + /// The chain index: chains as runs with per-run worker coverage; + /// lock-free, store-free lookups. + Chain(ChainBackend), + /// The single-threaded reference the other two are checked against. + #[cfg(test)] + Reference(reference::ReferenceBackend), +} + +/// The chain index with what the gateway keeps beside it. Every call the +/// gateway makes into the chain index goes through this type: a plain +/// `ChainIndex` and the sharded one share their call surface, so the payload +/// is the sharded type at one shard, and a per-socket placement later is a +/// matter of which shard a worker is interned into. +pub struct ChainBackend { + index: ShardedChainIndex, +} + +impl ChainBackend { + fn new() -> Self { + Self { + index: ShardedChainIndex::new(CHAIN_INDEX_SHARDS, CHAIN_INDEX_MAX_WORKERS), + } + } + + fn intern_worker(&self, worker: &str) -> Result { + self.index.intern_worker(worker) + } + + fn worker_id(&self, worker: &str) -> Option { + self.index.worker_id(worker) + } + + fn apply_stored( + &self, + worker: u32, + blocks: &[StoredBlock], + parent: Option, + held: &mut WorkerBlocks, + ) -> Result<(), ApplyError> { + self.index + .apply_stored(worker, blocks, parent, held.chain()) + } + + fn apply_removed(&self, worker: u32, hashes: &[SequenceHash], held: &mut WorkerBlocks) { + self.index.apply_removed(worker, hashes, held.chain()); + } + + fn apply_cleared(&self, worker: u32, held: &mut WorkerBlocks) { + self.index.apply_cleared(worker, held.chain()); + } + + fn remove_worker(&self, worker: u32, held: WorkerBlocks) { + self.index + .remove_worker(worker, held.chain.unwrap_or_default()); + } + + #[inline] + fn find_matches(&self, content_hashes: &[ContentHash], early_exit: bool) -> OverlapScores { + self.index.find_matches(content_hashes, early_exit) + } + + fn worker_block_count(&self, worker: u32) -> usize { + self.index.worker_block_count(worker) + } + + #[inline] + fn is_empty(&self) -> bool { + self.index.is_empty() + } + + fn current_size(&self) -> usize { + self.index.current_size() + } + + fn entry_count(&self) -> usize { + self.index.entry_count() + } + + fn stats(&self) -> ChainIndexStats { + self.index.stats() + } + + fn debug_blocks(&self) -> BTreeSet<(u32, usize, ContentHash, SequenceHash)> { + self.index.debug_blocks() + } +} + +/// One worker's share of a [`KvIndex`]: the per-worker reverse map the index +/// needs, in the shape the index needs it, created the first time the worker +/// stores a block. A worker's state only ever meets the one index its +/// subscription writes to, so at most one map is live. +#[derive(Default)] +pub struct WorkerBlocks { + positional: Option, + chain: Option, + #[cfg(test)] + reference: Option, +} + +impl WorkerBlocks { + /// Whether the worker holds the block the engine named `seq_hash`. + pub fn contains_key(&self, seq_hash: SequenceHash) -> bool { + if let Some(map) = &self.positional { + return map.contains_key(&seq_hash); + } + if let Some(map) = &self.chain { + return map.contains_key(seq_hash); + } + #[cfg(test)] + if let Some(held) = &self.reference { + return held.contains(seq_hash); + } + false + } + + /// Blocks the worker holds. + pub fn len(&self) -> usize { + if let Some(map) = &self.positional { + return map.len(); + } + if let Some(map) = &self.chain { + return map.len(); + } + #[cfg(test)] + if let Some(held) = &self.reference { + return held.len(); + } + 0 + } + + /// Whether the worker holds no block. + pub fn is_empty(&self) -> bool { + self.len() == 0 + } + + fn positional(&mut self) -> &mut WorkerBlockMap { + self.positional.get_or_insert_with(WorkerBlockMap::default) + } + + fn chain(&mut self) -> &mut ChainBlockMap { + self.chain.get_or_insert_with(ChainBlockMap::default) + } +} + +impl KvIndex { + /// An index of `kind`. `jump_size` is the positional indexer's historical + /// tuning knob and is ignored by the chain index. + pub fn new(kind: KvIndexKind, jump_size: usize) -> Self { + match kind { + KvIndexKind::Positional => Self::positional(jump_size), + KvIndexKind::Chain => Self::chain(), + } + } + + pub fn positional(jump_size: usize) -> Self { + Self::Positional(PositionalIndexer::new(jump_size)) + } + + pub fn chain() -> Self { + Self::Chain(ChainBackend::new()) + } + + #[cfg(test)] + pub fn reference() -> Self { + Self::Reference(reference::ReferenceBackend::default()) + } + + /// The variant's name, for logs. + pub fn name(&self) -> &'static str { + match self { + Self::Positional(_) => "positional", + Self::Chain(_) => "chain", + #[cfg(test)] + Self::Reference(_) => "reference", + } + } + + /// Intern a worker name; the same name maps to the same id until the + /// worker is removed. + pub fn intern_worker(&self, worker: &str) -> Result { + match self { + Self::Positional(index) => index.intern_worker(worker), + Self::Chain(chain) => chain.intern_worker(worker), + #[cfg(test)] + Self::Reference(reference) => Ok(reference.intern_worker(worker)), + } + } + + /// The id a worker name was interned to, if it is interned now. + pub fn worker_id(&self, worker: &str) -> Option { + match self { + Self::Positional(index) => index.worker_id(worker), + Self::Chain(chain) => chain.worker_id(worker), + #[cfg(test)] + Self::Reference(reference) => reference.worker_id(worker), + } + } + + /// Store `blocks` for `worker` after `parent` (position 0 when `None`). + pub fn apply_stored( + &self, + worker: u32, + blocks: &[StoredBlock], + parent: Option, + held: &mut WorkerBlocks, + ) -> Result<(), ApplyError> { + match self { + Self::Positional(index) => { + index.apply_stored(worker, blocks, parent, held.positional()) + } + Self::Chain(chain) => chain.apply_stored(worker, blocks, parent, held), + #[cfg(test)] + Self::Reference(reference) => reference.apply_stored(worker, blocks, parent, held), + } + } + + /// Forget the named blocks of `worker`; unknown hashes are ignored. + pub fn apply_removed(&self, worker: u32, hashes: &[SequenceHash], held: &mut WorkerBlocks) { + match self { + Self::Positional(index) => index.apply_removed(worker, hashes, held.positional()), + Self::Chain(chain) => chain.apply_removed(worker, hashes, held), + #[cfg(test)] + Self::Reference(reference) => reference.apply_removed(worker, hashes, held), + } + } + + /// Forget every block of `worker`; the caller keeps the emptied state. + pub fn apply_cleared(&self, worker: u32, held: &mut WorkerBlocks) { + match self { + Self::Positional(index) => index.apply_cleared(worker, held.positional()), + Self::Chain(chain) => chain.apply_cleared(worker, held), + #[cfg(test)] + Self::Reference(reference) => reference.apply_cleared(worker, held), + } + } + + /// Forget every block of `worker` and the worker itself; proportional to + /// the worker's blocks, not to the index. + pub fn remove_worker(&self, worker: u32, held: WorkerBlocks) { + match self { + Self::Positional(index) => { + index.remove_worker(worker, held.positional.unwrap_or_default()); + } + Self::Chain(chain) => chain.remove_worker(worker, held), + #[cfg(test)] + Self::Reference(reference) => reference.remove_worker(worker), + } + } + + /// Score every worker by how many leading blocks of the request it holds. + /// With `early_exit`, report the workers holding the first block, each + /// scored 1. The request path's one call into the index. + #[inline] + pub fn find_matches(&self, content_hashes: &[ContentHash], early_exit: bool) -> OverlapScores { + match self { + Self::Positional(index) => index.find_matches(content_hashes, early_exit), + Self::Chain(chain) => chain.find_matches(content_hashes, early_exit), + #[cfg(test)] + Self::Reference(reference) => reference.find_matches(content_hashes, early_exit), + } + } + + /// Blocks the index holds for `worker`: one counter read. + pub fn worker_block_count(&self, worker: u32) -> usize { + match self { + Self::Positional(index) => index.worker_block_count(worker), + Self::Chain(chain) => chain.worker_block_count(worker), + #[cfg(test)] + Self::Reference(reference) => reference.worker_block_count(worker), + } + } + + /// Whether no worker holds a block. Read once per request before the + /// lookup, so it must stay cheap: the positional indexer keeps a running + /// total; the chain index reads its root's child table under one seqlock. + #[inline] + pub fn is_empty(&self) -> bool { + match self { + Self::Positional(index) => index.current_size() == 0, + Self::Chain(chain) => chain.is_empty(), + #[cfg(test)] + Self::Reference(reference) => reference.is_empty(), + } + } + + /// Blocks held across all workers (a block two workers hold counts + /// twice). Not for the request path: see [`is_empty`](Self::is_empty). + pub fn current_size(&self) -> usize { + match self { + Self::Positional(index) => index.current_size(), + Self::Chain(chain) => chain.current_size(), + #[cfg(test)] + Self::Reference(reference) => reference.current_size(), + } + } + + /// Distinct index entries: `(position, content hash)` pairs in the + /// positional indexer, distinct blocks on a chain in the chain index. + pub fn entry_count(&self) -> usize { + match self { + Self::Positional(index) => index.entry_count(), + Self::Chain(chain) => chain.entry_count(), + #[cfg(test)] + Self::Reference(reference) => reference.entry_count(), + } + } + + /// Evict stale and excess entries. Only the positional indexer has a prune + /// (its entries carry a last-touch stamp); the chain index holds exactly + /// what the engines report and shrinks with their removals, so for it + /// this is `None` and the bounds do not apply. + pub fn prune(&self, ttl_secs: Option, max_entries: Option) -> Option { + match self { + Self::Positional(index) => Some(index.prune(ttl_secs, max_entries)), + Self::Chain(_) => None, + #[cfg(test)] + Self::Reference(_) => None, + } + } + + /// The chain index's shape and memory counters; `None` for the others. + pub fn chain_stats(&self) -> Option { + match self { + Self::Chain(chain) => Some(chain.stats()), + Self::Positional(_) => None, + #[cfg(test)] + Self::Reference(_) => None, + } + } + + /// Every membership as `(worker, position, content hash, prefix hash)`: + /// a full walk, for the exactness tests only. + #[doc(hidden)] + pub fn debug_blocks(&self) -> BTreeSet<(u32, usize, ContentHash, SequenceHash)> { + match self { + Self::Positional(index) => index.debug_blocks().into_iter().collect(), + Self::Chain(chain) => chain.debug_blocks(), + #[cfg(test)] + Self::Reference(reference) => reference.blocks(), + } + } +} + +impl fmt::Debug for KvIndex { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("KvIndex") + .field("kind", &self.name()) + .field("blocks", &self.current_size()) + .finish() + } +} + +#[cfg(test)] +mod exactness; + +#[cfg(test)] +mod mock_streams; + +#[cfg(test)] +mod reference { + //! The reference indexer behind the [`KvIndex`](super::KvIndex) surface: + //! interning and the per-worker membership set the other variants keep in + //! their maps, over `kv_index`'s single-threaded model. + + use std::collections::{BTreeSet, HashMap, HashSet}; + + use kv_index::{ + ApplyError, ContentHash, OverlapScores, ReferenceIndexer, SequenceHash, StoredBlock, + }; + use parking_lot::Mutex; + + use super::WorkerBlocks; + + /// The engine hashes a worker holds, mirrored from the reference so the + /// monitor's copy counting sees the same `contains_key` answers. + #[derive(Default)] + pub struct ReferenceBlocks(HashSet); + + impl ReferenceBlocks { + pub fn contains(&self, seq_hash: SequenceHash) -> bool { + self.0.contains(&seq_hash) + } + + pub fn len(&self) -> usize { + self.0.len() + } + } + + #[derive(Default)] + struct Inner { + index: ReferenceIndexer, + names: HashMap, + next: u32, + } + + #[derive(Default)] + pub struct ReferenceBackend { + inner: Mutex, + } + + impl ReferenceBackend { + pub fn intern_worker(&self, worker: &str) -> u32 { + let mut inner = self.inner.lock(); + if let Some(&id) = inner.names.get(worker) { + return id; + } + let id = inner.next; + inner.next += 1; + inner.names.insert(worker.to_string(), id); + id + } + + pub fn worker_id(&self, worker: &str) -> Option { + self.inner.lock().names.get(worker).copied() + } + + pub fn apply_stored( + &self, + worker: u32, + blocks: &[StoredBlock], + parent: Option, + held: &mut WorkerBlocks, + ) -> Result<(), ApplyError> { + self.inner + .lock() + .index + .apply_stored(worker, blocks, parent)?; + let mirror = held.reference.get_or_insert_with(ReferenceBlocks::default); + mirror.0.extend(blocks.iter().map(|block| block.seq_hash)); + Ok(()) + } + + pub fn apply_removed(&self, worker: u32, hashes: &[SequenceHash], held: &mut WorkerBlocks) { + self.inner.lock().index.apply_removed(worker, hashes); + if let Some(mirror) = held.reference.as_mut() { + for hash in hashes { + mirror.0.remove(hash); + } + } + } + + pub fn apply_cleared(&self, worker: u32, held: &mut WorkerBlocks) { + self.inner.lock().index.apply_cleared(worker); + held.reference = None; + } + + pub fn remove_worker(&self, worker: u32) { + let mut inner = self.inner.lock(); + inner.index.remove_worker(worker); + inner.names.retain(|_, id| *id != worker); + } + + pub fn find_matches( + &self, + content_hashes: &[ContentHash], + early_exit: bool, + ) -> OverlapScores { + let inner = self.inner.lock(); + let scored = if early_exit { + inner + .index + .find_matches(&content_hashes[..content_hashes.len().min(1)]) + } else { + inner.index.find_matches(content_hashes) + }; + let mut out = OverlapScores::default(); + for (worker, score) in scored { + out.scores.insert(worker, score); + } + out + } + + pub fn worker_block_count(&self, worker: u32) -> usize { + self.inner.lock().index.worker_block_count(worker) + } + + pub fn is_empty(&self) -> bool { + self.inner.lock().index.blocks().is_empty() + } + + pub fn current_size(&self) -> usize { + self.inner.lock().index.blocks().len() + } + + pub fn entry_count(&self) -> usize { + let inner = self.inner.lock(); + inner + .index + .blocks() + .into_iter() + .map(|(_, position, content, prefix)| (position, content, prefix)) + .collect::>() + .len() + } + + pub fn blocks(&self) -> BTreeSet<(u32, usize, ContentHash, SequenceHash)> { + self.inner.lock().index.blocks() + } + } +} + +#[cfg(test)] +mod tests { + use kv_index::compute_content_hash; + + use super::*; + + fn chain(contents: &[&[u32]]) -> Vec { + let hashes: Vec = contents.iter().map(|c| compute_content_hash(c)).collect(); + hashes + .iter() + .zip(kv_index::request_prefix_hashes(&hashes)) + .map(|(&content_hash, seq_hash)| StoredBlock { + seq_hash, + content_hash, + }) + .collect() + } + + fn backends() -> Vec { + vec![ + KvIndex::positional(64), + KvIndex::chain(), + KvIndex::reference(), + ] + } + + #[test] + fn every_backend_scores_stores_removals_and_clears_the_same_way() { + let blocks = chain(&[&[1, 2, 3, 4], &[5, 6, 7, 8], &[9, 10, 11, 12]]); + let request: Vec = blocks.iter().map(|b| b.content_hash).collect(); + for index in backends() { + let name = index.name(); + assert!(index.is_empty(), "{name}: empty at start"); + let w1 = index.intern_worker("grpc://w1").unwrap(); + let w2 = index.intern_worker("grpc://w2").unwrap(); + assert_eq!(index.worker_id("grpc://w2"), Some(w2), "{name}"); + let mut held1 = WorkerBlocks::default(); + let mut held2 = WorkerBlocks::default(); + index.apply_stored(w1, &blocks, None, &mut held1).unwrap(); + index + .apply_stored(w2, &blocks[..2], None, &mut held2) + .unwrap(); + assert!(!index.is_empty(), "{name}: not empty after stores"); + assert!(held1.contains_key(blocks[2].seq_hash), "{name}"); + assert!(!held2.contains_key(blocks[2].seq_hash), "{name}"); + assert_eq!(held1.len(), 3, "{name}"); + assert_eq!(index.worker_block_count(w1), 3, "{name}"); + assert_eq!(index.worker_block_count(w2), 2, "{name}"); + assert_eq!(index.current_size(), 5, "{name}"); + + let scores = index.find_matches(&request, false).scores; + assert_eq!(scores.get(&w1), Some(&3), "{name}: w1 holds the chain"); + assert_eq!(scores.get(&w2), Some(&2), "{name}: w2 holds two blocks"); + let early = index.find_matches(&request, true).scores; + assert_eq!(early.get(&w1), Some(&1), "{name}: early exit scores 1"); + assert_eq!(early.get(&w2), Some(&1), "{name}: early exit scores 1"); + + // A middle removal ends w1's match at the hole; w2 is untouched. + index.apply_removed(w1, &[blocks[1].seq_hash], &mut held1); + let scores = index.find_matches(&request, false).scores; + assert_eq!(scores.get(&w1), Some(&1), "{name}: hole ends the match"); + assert_eq!(scores.get(&w2), Some(&2), "{name}"); + assert_eq!(index.worker_block_count(w1), 2, "{name}"); + + // A store after the parent heals the hole. + index + .apply_stored(w1, &blocks[1..2], Some(blocks[0].seq_hash), &mut held1) + .unwrap(); + let scores = index.find_matches(&request, false).scores; + assert_eq!(scores.get(&w1), Some(&3), "{name}: healed"); + + // An unknown parent is reported, as the monitor's fallback expects. + assert!( + matches!( + index.apply_stored(w2, &blocks[2..], Some(SequenceHash(0xdead)), &mut held2), + Err(ApplyError::ParentBlockNotFound) + ), + "{name}" + ); + let mut fresh = WorkerBlocks::default(); + let w3 = index.intern_worker("grpc://w3").unwrap(); + assert!( + matches!( + index.apply_stored(w3, &blocks[1..], Some(blocks[0].seq_hash), &mut fresh), + Err(ApplyError::WorkerNotTracked) + ), + "{name}" + ); + + index.apply_cleared(w1, &mut held1); + assert!(held1.is_empty(), "{name}: cleared state is empty"); + assert_eq!(index.worker_block_count(w1), 0, "{name}"); + assert!( + !index.find_matches(&request, false).scores.contains_key(&w1), + "{name}: cleared worker scores nothing" + ); + index.remove_worker(w2, held2); + assert_eq!(index.worker_block_count(w2), 0, "{name}"); + assert!(index.is_empty(), "{name}: empty again"); + assert_eq!(index.current_size(), 0, "{name}"); + assert!(index.debug_blocks().is_empty(), "{name}"); + } + } + + #[test] + fn the_chain_index_hands_back_removed_worker_ids_and_knows_when_it_is_empty() { + let index = KvIndex::chain(); + let first = index.intern_worker("grpc://a").unwrap(); + let mut held = WorkerBlocks::default(); + let blocks = chain(&[&[1, 2, 3, 4]]); + index.apply_stored(first, &blocks, None, &mut held).unwrap(); + index.remove_worker(first, held); + assert_eq!(index.worker_id("grpc://a"), None); + assert!(index.is_empty()); + // The freed id is reused; the emptiness check still covers it. + let again = index.intern_worker("grpc://b").unwrap(); + assert_eq!(again, first); + let mut held = WorkerBlocks::default(); + index.apply_stored(again, &blocks, None, &mut held).unwrap(); + assert!(!index.is_empty()); + assert_eq!(index.chain_stats().map(|s| s.blocks_live), Some(1)); + } + + #[test] + fn only_the_positional_index_prunes() { + let positional = KvIndex::positional(64); + assert!(positional.prune(Some(1), None).is_some()); + assert!(KvIndex::chain().prune(Some(1), Some(1)).is_none()); + assert!(KvIndex::chain().chain_stats().is_some()); + assert!(positional.chain_stats().is_none()); + } + + #[test] + fn kind_selects_the_backend() { + assert_eq!( + KvIndex::new(KvIndexKind::Positional, 8).name(), + "positional" + ); + assert_eq!(KvIndexKind::Chain.as_str(), "chain"); + assert_eq!(KvIndex::new(KvIndexKind::Chain, 8).name(), "chain"); + assert_eq!( + format!("{:?}", KvIndex::chain()), + "KvIndex { kind: \"chain\", blocks: 0 }" + ); + } +} diff --git a/model_gateway/src/worker/kv_index_backend/exactness.rs b/model_gateway/src/worker/kv_index_backend/exactness.rs new file mode 100644 index 0000000000..71e577d724 --- /dev/null +++ b/model_gateway/src/worker/kv_index_backend/exactness.rs @@ -0,0 +1,1275 @@ +//! Exactness of the gateway's index backends (`crates/kv_index/docs/kv-router-leap.md`, +//! guardrails 1 to 3): the positional indexer and the chain index, fed through +//! the monitor's own apply path, must answer every lookup exactly as the +//! reference indexer does, after every batch, and hold exactly the same blocks. +//! +//! Two corpora. Engine streams generated in-process by the mock engine +//! (`mock_streams`: a vLLM-shaped run with a restart, an SGLang-shaped run, +//! two data-parallel ranks merged by receive order, and a run with a host +//! tier) leave the engine on its publisher wire and go through the relay's +//! decoder and normalizer, then through `KvEventMonitor::apply_event`, so +//! what reaches the index is what reaches it in production; at every +//! checkpoint the backends must also predict for every prompt the hit the +//! engine itself holds. A seeded synthetic corpus then does what an engine's +//! own stream does too little of: holes (evictions in the middle or at the +//! head of a chain the engine keeps using), heals (the missing block stored +//! again after its parent), divergent siblings sharing a prefix, duplicate +//! physical copies, host-tier copies, clears and worker removals, over +//! several workers at once. +//! +//! The request set replayed against the indexes is built from the chains the +//! events describe: every chain in full, its first half, a copy with the middle +//! block replaced, one with the last block replaced, one extended by a block +//! nobody stored, and a request nobody stored at all. + +use std::collections::{BTreeMap, BTreeSet, HashMap, HashSet}; + +use kv_index::{ + compute_content_hash, compute_request_content_hashes, request_prefix_hashes, ContentHash, + SequenceHash, +}; +use smg_grpc_client::common_proto::{ + kv_cache_event, KvBlock, KvBlocksRemoved, KvBlocksStored, KvCacheCleared, KvCacheEvent, + KvCacheTier, KvEventBatch, +}; + +use super::{ + mock_streams::{self, Shape, Stream}, + KvIndex, +}; +use crate::worker::kv_event_monitor::{KvEventMonitor, WorkerIndexState}; + +// --------------------------------------------------------------------------- +// The three backends side by side +// --------------------------------------------------------------------------- + +struct Backend { + index: KvIndex, + workers: BTreeMap, +} + +impl Backend { + fn new(index: KvIndex) -> Self { + Self { + index, + workers: BTreeMap::new(), + } + } + + fn worker(&mut self, name: &str) -> &mut (u32, WorkerIndexState) { + if !self.workers.contains_key(name) { + let id = self.index.intern_worker(name).expect("worker id"); + self.workers + .insert(name.to_string(), (id, WorkerIndexState::default())); + } + self.workers.get_mut(name).expect("just inserted") + } + + fn names(&self) -> BTreeMap { + self.workers + .iter() + .map(|(name, (id, _))| (*id, name.clone())) + .collect() + } + + /// Scores by worker name, so backends with different interned ids compare. + fn scores(&self, query: &[ContentHash], early_exit: bool) -> BTreeMap { + let names = self.names(); + self.index + .find_matches(query, early_exit) + .scores + .into_iter() + .map(|(id, score)| (names[&id].clone(), score)) + .collect() + } + + fn blocks(&self) -> BTreeSet<(String, usize, ContentHash, SequenceHash)> { + let names = self.names(); + self.index + .debug_blocks() + .into_iter() + .map(|(id, position, content, prefix)| (names[&id].clone(), position, content, prefix)) + .collect() + } + + fn counts(&self) -> BTreeMap { + self.workers + .iter() + .map(|(name, (id, _))| (name.clone(), self.index.worker_block_count(*id))) + .collect() + } +} + +/// The production backends next to the reference, fed identically. +struct Trio { + backends: Vec, + lookups: usize, +} + +impl Trio { + fn new() -> Self { + Self { + backends: vec![ + Backend::new(KvIndex::reference()), + Backend::new(KvIndex::positional(64)), + Backend::new(KvIndex::chain()), + ], + lookups: 0, + } + } + + fn add_worker(&mut self, name: &str) { + for backend in &mut self.backends { + backend.worker(name); + } + } + + /// One batch through the monitor's apply path, for every backend. + fn apply(&mut self, worker: &str, batch: &KvEventBatch) { + for backend in &mut self.backends { + let Backend { index, workers } = backend; + if !workers.contains_key(worker) { + let id = index.intern_worker(worker).expect("worker id"); + workers.insert(worker.to_string(), (id, WorkerIndexState::default())); + } + let (id, state) = workers.get_mut(worker).expect("present"); + for event in &batch.events { + KvEventMonitor::apply_event(event, *id, index, state); + } + } + } + + fn remove_worker(&mut self, worker: &str) { + for backend in &mut self.backends { + let Some((id, state)) = backend.workers.remove(worker) else { + continue; + }; + backend.index.remove_worker(id, state.blocks); + } + } + + /// Every backend answers every query as the reference does, holds the + /// reference's blocks and counts them the same way. + fn check(&mut self, queries: &[Vec], label: &str) { + let (reference, production) = self.backends.split_first().expect("three backends"); + for backend in production { + let name = backend.index.name(); + for query in queries { + for early_exit in [false, true] { + let want = reference.scores(query, early_exit); + let got = backend.scores(query, early_exit); + assert_eq!( + got, want, + "{label}: {name} scores (early_exit {early_exit}) for {query:?}" + ); + self.lookups += 1; + } + } + assert_eq!( + backend.counts(), + reference.counts(), + "{label}: {name} block counts" + ); + assert_eq!( + backend.blocks(), + reference.blocks(), + "{label}: {name} index content" + ); + } + } + + /// Every backend predicts for every prompt the hit the engine holds: the + /// prompt's leading blocks the worker has, by the engine's own count at + /// the checkpoint. Returns the prompts checked. + fn agree( + &mut self, + worker: &str, + prompts: &BTreeSet>, + held: &HashSet, + label: &str, + ) -> usize { + for prompt in prompts { + let want = Stream::prefix_match(held, prompt) as u32; + let hashes = compute_request_content_hashes(prompt, mock_streams::BLOCK); + for backend in &self.backends { + let got = backend + .scores(&hashes, false) + .get(worker) + .copied() + .unwrap_or(0); + assert_eq!( + got, + want, + "{label}: {} predicts {got} cached blocks of a prompt the engine holds {want} of", + backend.index.name() + ); + } + self.lookups += self.backends.len(); + } + prompts.len() + } +} + +// --------------------------------------------------------------------------- +// Queries from chains +// --------------------------------------------------------------------------- + +/// A request nobody stored: distinct per call, never a real content hash. +fn novel(counter: &mut u64) -> ContentHash { + *counter += 1; + ContentHash(0xdead_beef_0000_0000 | *counter) +} + +/// The request set for a chain: itself, its first half, the middle block +/// replaced, the last block replaced, extended by a block nobody stored. +fn variants(chain: &[ContentHash], counter: &mut u64) -> Vec> { + let mut out = vec![chain.to_vec()]; + if chain.len() > 1 { + out.push(chain[..chain.len() / 2].to_vec()); + let mut middle = chain.to_vec(); + middle[chain.len() / 2] = novel(counter); + out.push(middle); + let mut suffix = chain.to_vec(); + *suffix.last_mut().expect("non-empty") = novel(counter); + out.push(suffix); + } + let mut extended = chain.to_vec(); + extended.push(novel(counter)); + out.push(extended); + out +} + +/// The query set over every chain seen so far, plus one nobody stored. +fn query_set(chains: &BTreeSet>) -> Vec> { + let mut counter = 0u64; + let mut queries: Vec> = chains + .iter() + .flat_map(|chain| variants(chain, &mut counter)) + .collect(); + queries.push(vec![novel(&mut counter), novel(&mut counter)]); + queries +} + +// --------------------------------------------------------------------------- +// Engine streams +// --------------------------------------------------------------------------- + +/// Chains the stores of a stream describe, as a request would hash them: +/// each store's blocks appended to the chain of its parent block. Only plain +/// (unsalted, non-LoRA) stores, which is all the engine's stores carry. +#[derive(Default)] +struct Chains { + /// Engine hash of a block -> the token chain from the root through it. + tokens: BTreeMap>, + block_size: usize, + seen: BTreeSet>, +} + +impl Chains { + fn note(&mut self, batch: &KvEventBatch) { + for event in &batch.events { + let Some(kv_cache_event::Data::Stored(stored)) = &event.data else { + continue; + }; + assert!( + stored.lora_name.is_none() && stored.cache_salt.is_none(), + "the engine's stores carry no namespace" + ); + let mut chain = stored + .parent_block_hash + .and_then(|parent| self.tokens.get(&parent).cloned()) + .unwrap_or_default(); + for block in &stored.blocks { + if self.block_size == 0 { + self.block_size = block.block_size as usize; + } + chain.extend_from_slice(&block.token_ids); + self.tokens.insert(block.block_hash, chain.clone()); + } + if self.block_size > 0 { + self.seen + .insert(compute_request_content_hashes(&chain, self.block_size)); + } + } + } +} + +/// What a stream carried, by its normalized batches: each shape's claims are +/// checked against these counts. +#[derive(Debug, Default)] +struct Carried { + stored: usize, + removed: usize, + cleared: usize, + /// The ranks that stored blocks. + ranks: BTreeSet>, + /// Stores of a block its rank already held: second physical copies. + second_copies: usize, + /// Device removals that left a copy on the same rank: the block stays. + pinned_by_copy: usize, + /// Stores of a block another rank held at the time. + shared_across_ranks: usize, + host_stored: usize, + host_removed: usize, + /// Device removals of a block the host still held: the block stays. + pinned_by_host: usize, + /// Host removals of a block no device copy was left of: the block goes. + last_copy: usize, +} + +/// The copies a stream has announced so far: per rank and block on the +/// device, and the blocks on the host. +#[derive(Default)] +struct Copies { + device: HashMap, HashMap>, + host: HashSet, +} + +impl Copies { + fn on_device(&self, hash: i64) -> bool { + self.device.values().any(|held| held.contains_key(&hash)) + } +} + +impl Carried { + fn count(batches: &[KvEventBatch]) -> Self { + let mut carried = Self::default(); + let mut copies = Copies::default(); + for batch in batches { + for event in &batch.events { + match &event.data { + Some(kv_cache_event::Data::Stored(stored)) => { + carried.note_stored(&mut copies, batch.dp_rank, stored); + } + Some(kv_cache_event::Data::Removed(removed)) => { + carried.note_removed(&mut copies, batch.dp_rank, removed); + } + Some(kv_cache_event::Data::Cleared(_)) => { + carried.cleared += 1; + copies = Copies::default(); + } + None => {} + } + } + } + carried + } + + fn note_stored(&mut self, copies: &mut Copies, rank: Option, stored: &KvBlocksStored) { + let hashes = stored.blocks.iter().map(|block| block.block_hash); + if is_host(stored.tier) { + self.host_stored += 1; + copies.host.extend(hashes); + return; + } + self.stored += 1; + self.ranks.insert(rank); + for hash in hashes { + if copies + .device + .iter() + .any(|(other, held)| *other != rank && held.contains_key(&hash)) + { + self.shared_across_ranks += 1; + } + let count = copies + .device + .entry(rank) + .or_default() + .entry(hash) + .or_insert(0); + if *count > 0 { + self.second_copies += 1; + } + *count += 1; + } + } + + fn note_removed(&mut self, copies: &mut Copies, rank: Option, removed: &KvBlocksRemoved) { + if is_host(removed.tier) { + self.host_removed += 1; + for hash in &removed.block_hashes { + copies.host.remove(hash); + if !copies.on_device(*hash) { + self.last_copy += 1; + } + } + return; + } + self.removed += 1; + let held = copies.device.entry(rank).or_default(); + for hash in &removed.block_hashes { + if let Some(count) = held.get_mut(hash) { + *count -= 1; + if *count > 0 { + self.pinned_by_copy += 1; + } else { + held.remove(hash); + } + } + if copies.host.contains(hash) { + self.pinned_by_host += 1; + } + } + } +} + +fn is_host(tier: Option) -> bool { + tier == Some(KvCacheTier::Host as i32) +} + +/// What a replayed stream established. +struct Replayed { + carried: Carried, + /// Ranks that published payloads. + publishers: usize, + checkpoints: usize, + lookups: usize, + agreements: usize, +} + +impl Replayed { + /// Enough of the stream was checked: the engine was caught up at enough + /// checkpoints, and every one compared the whole query set and every + /// prompt. + fn assert_coverage(&self) { + assert!( + self.checkpoints >= 16, + "{} checkpoints with the engine caught up", + self.checkpoints + ); + assert!(self.lookups > 100_000, "{} lookups compared", self.lookups); + assert!( + self.agreements > 5_000, + "{} hits predicted", + self.agreements + ); + } +} + +/// One engine stream through the monitor into every backend. At every +/// checkpoint the backends answer the query set as the reference does, hold +/// its blocks, and predict for every prompt the hit the engine holds. +fn replay_stream(shape: Shape, seed: u64, requests: usize) -> Replayed { + let stream = Stream::generate(shape, seed, requests); + let batches = stream.normalized(); + let label = shape.name(); + assert_eq!( + batches.len(), + stream.payloads.len(), + "{label}: every payload decodes" + ); + let worker = format!("grpc://{label}"); + let mut trio = Trio::new(); + trio.add_worker(&worker); + let prompts: BTreeSet> = stream.prompts.iter().cloned().collect(); + let mut chains = Chains::default(); + let mut checkpoints = stream.checkpoints.iter().peekable(); + let mut agreements = 0; + for (index, batch) in batches.iter().enumerate() { + chains.note(batch); + trio.apply(&worker, batch); + while let Some(checkpoint) = checkpoints.next_if(|checkpoint| checkpoint.after == index + 1) + { + let at = format!("{label} batch {index}"); + trio.check(&query_set(&chains.seen), &at); + agreements += trio.agree(&worker, &prompts, &checkpoint.held, &at); + } + } + assert!( + checkpoints.next().is_none(), + "{label}: a checkpoint past the stream" + ); + let publishers: BTreeSet = stream.payloads.iter().map(|payload| payload.rank).collect(); + Replayed { + carried: Carried::count(&batches), + publishers: publishers.len(), + checkpoints: stream.checkpoints.len(), + lookups: trio.lookups, + agreements, + } +} + +#[test] +fn vllm_stream_scores_as_the_reference_and_predicts_the_engine_hit() { + let replayed = replay_stream(Shape::Vllm, 1, 320); + let carried = &replayed.carried; + assert!( + carried.stored >= 500 && carried.removed >= 300, + "{carried:?}" + ); + assert_eq!(carried.cleared, 1, "the restart's clear: {carried:?}"); + assert!( + carried.second_copies >= 100 && carried.pinned_by_copy >= 30, + "second physical copies and the removals they survived: {carried:?}" + ); + assert_eq!(replayed.publishers, 1); + replayed.assert_coverage(); +} + +#[test] +fn sglang_stream_scores_as_the_reference_and_predicts_the_engine_hit() { + let replayed = replay_stream(Shape::Sglang, 2, 320); + let carried = &replayed.carried; + assert!( + carried.stored >= 500 && carried.removed >= 300, + "{carried:?}" + ); + assert_eq!(carried.cleared, 0, "{carried:?}"); + assert_eq!(replayed.publishers, 1); + replayed.assert_coverage(); +} + +#[test] +fn two_rank_stream_scores_as_the_reference_and_predicts_the_engine_hit() { + let replayed = replay_stream(Shape::TwoRank, 3, 400); + let carried = &replayed.carried; + assert_eq!( + carried.ranks, + BTreeSet::from([Some(0), Some(1)]), + "{carried:?}" + ); + assert!( + carried.shared_across_ranks >= 200, + "blocks held by both ranks: {carried:?}" + ); + assert!( + carried.stored >= 500 && carried.removed >= 300, + "{carried:?}" + ); + assert_eq!(replayed.publishers, 2, "both ranks published"); + replayed.assert_coverage(); +} + +#[test] +fn host_tier_stream_scores_as_the_reference_and_predicts_the_engine_hit() { + let replayed = replay_stream(Shape::HostTier, 4, 320); + let carried = &replayed.carried; + assert!( + carried.host_stored >= 1000 && carried.host_removed >= 1000, + "{carried:?}" + ); + assert!( + carried.pinned_by_host >= 1000, + "device removals the host copy survived: {carried:?}" + ); + assert!( + carried.last_copy >= 1000, + "host removals that took the last copy: {carried:?}" + ); + assert_eq!(replayed.publishers, 1); + replayed.assert_coverage(); +} + +// --------------------------------------------------------------------------- +// Synthetic corpus with holes +// --------------------------------------------------------------------------- + +/// xorshift64*, as in `kv_index`'s exactness tests. +struct Rng(u64); + +impl Rng { + fn new(seed: u64) -> Self { + Self(seed.max(1)) + } + + fn next(&mut self) -> u64 { + let mut x = self.0; + x ^= x >> 12; + x ^= x << 25; + x ^= x >> 27; + self.0 = x; + x.wrapping_mul(0x2545_F491_4F6C_DD1D) + } + + fn below(&mut self, n: usize) -> usize { + (self.next() % n as u64) as usize + } + + fn range(&mut self, lo: usize, hi_inclusive: usize) -> usize { + lo + self.below(hi_inclusive - lo + 1) + } + + fn chance(&mut self, numerator: u64, denominator: u64) -> bool { + self.next() % denominator < numerator + } +} + +const BLOCK: usize = 4; + +/// Token ids of block `position` of content stream `stream`: distinct per +/// (stream, position), so a chain's content hashes are distinct too. +fn tokens(stream: u64, position: usize) -> Vec { + vec![ + (stream & 0xffff_ffff) as u32, + (stream >> 32) as u32, + position as u32, + 7, + ] +} + +/// One block of a chain as the engine names it: its tokens and its engine +/// hash, the chain hash of the contents so far (shared by every worker that +/// stores the same prefix, as the engines' hashes are). +#[derive(Clone)] +struct Block { + tokens: Vec, + hash: i64, +} + +/// A chain a worker stored, with what the generator believes is still held. +struct Chain { + blocks: Vec, + alive: Vec, +} + +impl Chain { + fn contents(&self) -> Vec { + self.blocks + .iter() + .map(|block| compute_content_hash(&block.tokens)) + .collect() + } +} + +fn chain_of(token_blocks: &[Vec]) -> Vec { + let contents: Vec = token_blocks + .iter() + .map(|tokens| compute_content_hash(tokens)) + .collect(); + token_blocks + .iter() + .zip(request_prefix_hashes(&contents)) + .map(|(tokens, prefix)| Block { + tokens: tokens.clone(), + hash: prefix.0 as i64, + }) + .collect() +} + +fn kv_block(block: &Block) -> KvBlock { + KvBlock { + block_hash: block.hash, + token_ids: block.tokens.clone(), + block_size: BLOCK as i32, + ..Default::default() + } +} + +fn stored(parent: Option, blocks: &[Block], tier: Option) -> KvCacheEvent { + KvCacheEvent { + event_id: 0, + data: Some(kv_cache_event::Data::Stored(KvBlocksStored { + blocks: blocks.iter().map(kv_block).collect(), + parent_block_hash: parent, + tier: tier.map(|tier| tier as i32), + ..Default::default() + })), + } +} + +fn removed(hashes: Vec, tier: Option) -> KvCacheEvent { + KvCacheEvent { + event_id: 0, + data: Some(kv_cache_event::Data::Removed(KvBlocksRemoved { + block_hashes: hashes, + tier: tier.map(|tier| tier as i32), + ..Default::default() + })), + } +} + +fn cleared() -> KvCacheEvent { + KvCacheEvent { + event_id: 0, + data: Some(kv_cache_event::Data::Cleared(KvCacheCleared::default())), + } +} + +fn batch(events: Vec) -> KvEventBatch { + KvEventBatch { + events, + ..Default::default() + } +} + +struct Corpus { + rng: Rng, + workers: Vec, + held: BTreeMap>, + prompts: Vec>>, + chains: BTreeSet>, + next_stream: u64, + holes: usize, + heals: usize, +} + +impl Corpus { + fn new(seed: u64, workers: usize) -> Self { + let mut rng = Rng::new(seed); + let mut next_stream = 1u64; + let prompts = (0..rng.range(3, 6)) + .map(|_| { + next_stream += 1; + (0..rng.range(4, 24)) + .map(|position| tokens(next_stream, position)) + .collect() + }) + .collect(); + Self { + rng, + workers: (0..workers).map(|i| format!("grpc://w{i}:9000")).collect(), + held: BTreeMap::new(), + prompts, + chains: BTreeSet::new(), + next_stream, + holes: 0, + heals: 0, + } + } + + fn fresh_blocks(&mut self, count: usize) -> Vec> { + self.next_stream += 1; + let stream = self.next_stream; + (0..count) + .map(|position| tokens(stream, position)) + .collect() + } + + fn worker(&mut self) -> String { + self.workers[self.rng.below(self.workers.len())].clone() + } + + fn remember(&mut self, chain: &Chain) { + self.chains.insert(chain.contents()); + } + + /// Store a chain: a shared prompt prefix with a novel suffix (the common + /// case) or something entirely new, in one event or split at a parent. + fn store_new(&mut self, trio: &mut Trio) { + let worker = self.worker(); + let mut token_blocks = if self.rng.chance(3, 4) { + let prompt = &self.prompts[self.rng.below(self.prompts.len())]; + let keep = self.rng.range(1, prompt.len()); + prompt[..keep].to_vec() + } else { + Vec::new() + }; + let suffix = self.rng.range(1, 12); + token_blocks.extend(self.fresh_blocks(suffix)); + let blocks = chain_of(&token_blocks); + let events = if blocks.len() > 2 && self.rng.chance(1, 2) { + let cut = self.rng.range(1, blocks.len() - 1); + vec![ + stored(None, &blocks[..cut], None), + stored(Some(blocks[cut - 1].hash), &blocks[cut..], None), + ] + } else { + vec![stored(None, &blocks, None)] + }; + trio.apply(&worker, &batch(events)); + let chain = Chain { + alive: vec![true; blocks.len()], + blocks, + }; + self.remember(&chain); + self.held.entry(worker).or_default().push(chain); + } + + fn pick_chain(&mut self) -> Option<(String, usize)> { + let worker = self.worker(); + let count = self.held.get(&worker).map_or(0, Vec::len); + (count > 0).then(|| (worker, self.rng.below(count))) + } + + /// Decode extends the chain after its last block, while that block is + /// held (an engine extends what it holds; a store naming an evicted parent + /// is the gateway's fallback regime, which has its own test below). + fn extend(&mut self, trio: &mut Trio) { + let Some((worker, at)) = self.pick_chain() else { + return; + }; + if self.held[&worker][at].alive.last() != Some(&true) { + return; + } + let more = self.rng.range(1, 3); + let fresh = self.fresh_blocks(more); + let chain = &self.held[&worker][at]; + let mut token_blocks: Vec> = chain + .blocks + .iter() + .map(|block| block.tokens.clone()) + .collect(); + let parent = chain.blocks.last().map(|block| block.hash); + token_blocks.extend(fresh); + let blocks = chain_of(&token_blocks); + let new = &blocks[blocks.len() - more..]; + trio.apply(&worker, &batch(vec![stored(parent, new, None)])); + let held = &mut self.held.get_mut(&worker).expect("held")[at]; + held.blocks.extend_from_slice(new); + held.alive.resize(held.blocks.len(), true); + let contents = held.contents(); + self.chains.insert(contents); + } + + /// A sibling diverges from a chain's prefix, possibly on another worker. + fn sibling(&mut self, trio: &mut Trio) { + let Some((worker, at)) = self.pick_chain() else { + return; + }; + let source = &self.held[&worker][at]; + if source.blocks.len() < 2 { + return; + } + let keep = self.rng.range(1, source.blocks.len() - 1); + let mut token_blocks: Vec> = source.blocks[..keep] + .iter() + .map(|block| block.tokens.clone()) + .collect(); + let more = self.rng.range(1, 6); + token_blocks.extend(self.fresh_blocks(more)); + let blocks = chain_of(&token_blocks); + let target = if self.rng.chance(1, 2) { + worker.clone() + } else { + self.worker() + }; + // The target stores the shared prefix first unless it holds it already + // (the parent must be its own block), then the divergent tail. + let holds_prefix = self.held.get(&target).is_some_and(|chains| { + chains.iter().any(|chain| { + chain.blocks.len() >= keep + && chain.blocks[..keep] + .iter() + .zip(&blocks[..keep]) + .all(|(a, b)| a.hash == b.hash) + && chain.alive[..keep].iter().all(|&alive| alive) + }) + }); + let mut events = Vec::new(); + if !holds_prefix { + events.push(stored(None, &blocks[..keep], None)); + } + events.push(stored(Some(blocks[keep - 1].hash), &blocks[keep..], None)); + trio.apply(&target, &batch(events)); + let chain = Chain { + alive: vec![true; blocks.len()], + blocks, + }; + self.remember(&chain); + self.held.entry(target).or_default().push(chain); + } + + /// Evict from the tail, or punch a hole in the middle or at the head + /// while the blocks after it stay (what block-level LRU eviction does). + fn evict(&mut self, trio: &mut Trio) { + let Some((worker, at)) = self.pick_chain() else { + return; + }; + let chain = &mut self.held.get_mut(&worker).expect("held")[at]; + let alive: Vec = (0..chain.blocks.len()) + .filter(|&i| chain.alive[i]) + .collect(); + if alive.is_empty() { + return; + } + let kind = self.rng.below(4); + let victims: Vec = match kind { + 0 | 1 => { + let count = self.rng.range(1, alive.len()); + alive[alive.len() - count..].to_vec() + } + 2 if alive.len() > 2 => vec![alive[self.rng.range(1, alive.len() - 2)]], + _ => vec![alive[0]], + }; + if kind >= 2 { + self.holes += 1; + } + let hashes: Vec = victims.iter().map(|&i| chain.blocks[i].hash).collect(); + for &i in &victims { + chain.alive[i] = false; + } + trio.apply(&worker, &batch(vec![removed(hashes, None)])); + } + + /// Store a dead block again after its parent (or at the head), as an + /// engine does when the request comes back. + fn heal(&mut self, trio: &mut Trio) { + let Some((worker, at)) = self.pick_chain() else { + return; + }; + let chain = &mut self.held.get_mut(&worker).expect("held")[at]; + let Some(dead) = (0..chain.blocks.len()).find(|&i| !chain.alive[i]) else { + return; + }; + if dead > 0 && !chain.alive[dead - 1] { + // The parent is gone too: healing starts from the head. + return; + } + let parent = (dead > 0).then(|| chain.blocks[dead - 1].hash); + let block = chain.blocks[dead].clone(); + chain.alive[dead] = true; + self.heals += 1; + trio.apply(&worker, &batch(vec![stored(parent, &[block], None)])); + } + + /// A second physical copy of a held prefix (vLLM re-prefills the last + /// block of an exact resend) followed later by its own removal. + fn duplicate(&mut self, trio: &mut Trio) { + let Some((worker, at)) = self.pick_chain() else { + return; + }; + let chain = &self.held[&worker][at]; + let keep = chain.alive.iter().take_while(|&&alive| alive).count(); + if keep == 0 { + return; + } + let blocks = chain.blocks[..keep].to_vec(); + let last = blocks[keep - 1].hash; + trio.apply(&worker, &batch(vec![stored(None, &blocks, None)])); + // One copy of the last block goes away: the block must stay indexed. + trio.apply(&worker, &batch(vec![removed(vec![last], None)])); + } + + /// HiCache write-through: a host copy of a held prefix, the device copy + /// evicted (the block stays through the host copy), then the host copy. + fn host_copy(&mut self, trio: &mut Trio) { + let Some((worker, at)) = self.pick_chain() else { + return; + }; + let chain = &mut self.held.get_mut(&worker).expect("held")[at]; + let keep = chain.alive.iter().take_while(|&&alive| alive).count(); + if keep == 0 { + return; + } + let blocks = chain.blocks[..keep].to_vec(); + let hashes: Vec = blocks.iter().map(|block| block.hash).collect(); + trio.apply( + &worker, + &batch(vec![stored(None, &blocks, Some(KvCacheTier::Host))]), + ); + trio.apply( + &worker, + &batch(vec![removed(hashes.clone(), Some(KvCacheTier::Device))]), + ); + if self.rng.chance(1, 2) { + trio.apply( + &worker, + &batch(vec![removed(hashes, Some(KvCacheTier::Host))]), + ); + for alive in chain.alive.iter_mut().take(keep) { + *alive = false; + } + } + } + + fn clear(&mut self, trio: &mut Trio) { + let worker = self.worker(); + trio.apply(&worker, &batch(vec![cleared()])); + self.held.remove(&worker); + } + + fn replace_worker(&mut self, trio: &mut Trio) { + let worker = self.worker(); + trio.remove_worker(&worker); + self.held.remove(&worker); + trio.add_worker(&worker); + } + + fn step(&mut self, trio: &mut Trio) { + match self.rng.below(100) { + 0..=24 => self.store_new(trio), + 25..=39 => self.extend(trio), + 40..=54 => self.sibling(trio), + 55..=74 => self.evict(trio), + 75..=84 => self.heal(trio), + 85..=89 => self.duplicate(trio), + 90..=94 => self.host_copy(trio), + 95..=97 => self.clear(trio), + _ => self.replace_worker(trio), + } + } +} + +fn run_corpus(seed: u64, workers: usize, steps: usize, checkpoint: usize) -> (Corpus, usize) { + let mut trio = Trio::new(); + let mut corpus = Corpus::new(seed, workers); + for worker in corpus.workers.clone() { + trio.add_worker(&worker); + } + for step in 0..steps { + corpus.step(&mut trio); + if step % checkpoint == checkpoint - 1 { + trio.check( + &query_set(&corpus.chains), + &format!("seed {seed} step {step}"), + ); + } + } + trio.check(&query_set(&corpus.chains), &format!("seed {seed} end")); + (corpus, trio.lookups) +} + +#[test] +fn synthetic_corpus_with_holes_scores_identically_on_every_backend() { + for seed in [1, 2, 3] { + let (corpus, lookups) = run_corpus(seed, 5, 400, 10); + assert!(corpus.holes > 20, "seed {seed}: {} holes", corpus.holes); + assert!(corpus.heals > 5, "seed {seed}: {} heals", corpus.heals); + assert!(lookups > 10_000, "seed {seed}: {lookups} lookups compared"); + } +} + +/// The gateway's fallback regime: a store names a parent the index no longer +/// holds, so the monitor stores it again without a parent and the blocks land +/// at position 0 under engine hashes that belong further down the chain. When +/// the engine later stores the same hashes at their true positions, the +/// latest store wins in every index: a hash is held at one place per worker, +/// the mislaid pair scores nothing, and nothing is left for the next worker +/// interned into a freed id to inherit. +#[test] +fn a_hash_stored_again_at_its_true_position_after_a_fallback() { + let mut trio = Trio::new(); + let worker = "grpc://w0:9000"; + trio.add_worker(worker); + let token_blocks = Corpus::new(7, 1).fresh_blocks(6); + let blocks = chain_of(&token_blocks); + let contents: Vec = token_blocks + .iter() + .map(|tokens| compute_content_hash(tokens)) + .collect(); + // b0..b3 held, b2 and b3 evicted, then the engine extends after b3: the + // parent is unknown to the index, the fallback puts b4 and b5 at 0 and 1. + trio.apply(worker, &batch(vec![stored(None, &blocks[..4], None)])); + trio.apply( + worker, + &batch(vec![removed(vec![blocks[2].hash, blocks[3].hash], None)]), + ); + trio.apply( + worker, + &batch(vec![stored(Some(blocks[3].hash), &blocks[4..], None)]), + ); + let mislaid: Vec = contents[4..].to_vec(); + trio.check(&[mislaid.clone(), contents.clone()], "after the fallback"); + assert_eq!( + trio.backends[0].scores(&mislaid, false).get(worker), + Some(&2), + "the fallback placed the pair at the head" + ); + // The engine recomputes the chain and announces it whole: b4 and b5 move + // to positions 4 and 5. + trio.apply(worker, &batch(vec![stored(None, &blocks, None)])); + trio.check( + &[mislaid.clone(), contents.clone()], + "after the true-position store", + ); + for backend in &trio.backends { + let name = backend.index.name(); + assert_eq!( + backend.scores(&contents, false).get(worker), + Some(&6), + "{name}: the whole chain is held" + ); + assert!( + backend.scores(&mislaid, false).is_empty(), + "{name}: the mislaid pair scores nothing" + ); + assert_eq!( + backend.counts()[worker], + 6, + "{name}: six blocks, one place each" + ); + } + // The worker leaves and another takes its id: nothing comes with it. + trio.remove_worker(worker); + trio.add_worker("grpc://w9:9000"); + for backend in &trio.backends { + let name = backend.index.name(); + assert!( + backend.scores(&contents, false).is_empty(), + "{name}: a fresh worker holds nothing" + ); + assert_eq!(backend.counts()["grpc://w9:9000"], 0, "{name}"); + } +} + +/// A second physical copy of a block arrives inside a longer store under the +/// same parent (its first blocks second copies, the rest first copies), as +/// vLLM publishes one; a removal takes one copy. The block stays, and a +/// lookup through it still hits, until its last copy goes. +#[test] +fn a_second_copy_inside_a_longer_store_pins_the_block_until_its_last_removal() { + let mut trio = Trio::new(); + let worker = "grpc://w0:9000"; + trio.add_worker(worker); + let token_blocks = Corpus::new(9, 1).fresh_blocks(7); + let blocks = chain_of(&token_blocks); + let contents: Vec = token_blocks + .iter() + .map(|tokens| compute_content_hash(tokens)) + .collect(); + // The preamble b0..b3; b4 b5 under b3; then b4 b5 b6 under b3: second + // copies of b4 and b5 and the first copy of b6. + trio.apply(worker, &batch(vec![stored(None, &blocks[..4], None)])); + trio.apply( + worker, + &batch(vec![stored(Some(blocks[3].hash), &blocks[4..6], None)]), + ); + trio.apply( + worker, + &batch(vec![stored(Some(blocks[3].hash), &blocks[4..7], None)]), + ); + let queries = [contents[..6].to_vec(), contents.clone()]; + trio.check(&queries, "two copies"); + // One copy each of b5 and b4 goes, tail first: the chain still hits + // through b6. + trio.apply( + worker, + &batch(vec![removed(vec![blocks[5].hash, blocks[4].hash], None)]), + ); + trio.check(&queries, "one copy left"); + for backend in &trio.backends { + let name = backend.index.name(); + assert_eq!( + backend.scores(&contents, false).get(worker), + Some(&7), + "{name}: the whole chain hits with one copy of b4 and b5 left" + ); + assert_eq!(backend.counts()[worker], 7, "{name}"); + } + // The last copies go: the chain is cut at b4; b6 stays, unreachable. + trio.apply( + worker, + &batch(vec![removed(vec![blocks[5].hash, blocks[4].hash], None)]), + ); + trio.check(&queries, "no copy left"); + for backend in &trio.backends { + let name = backend.index.name(); + assert_eq!( + backend.scores(&contents, false).get(worker), + Some(&4), + "{name}: the chain is cut at b4" + ); + assert_eq!(backend.counts()[worker], 5, "{name}: the preamble and b6"); + } +} + +// --------------------------------------------------------------------------- +// Guardrails 2 and 3 on the chain index +// --------------------------------------------------------------------------- + +/// Lookups leave the chain index as they found it: the same runs, arena words +/// and blocks before and after a burst of queries (the store-free property of +/// the read path is established in `kv_index`; this checks the gateway's use +/// of it adds nothing). +#[test] +fn chain_index_lookups_change_nothing() { + let mut trio = Trio::new(); + let mut corpus = Corpus::new(11, 4); + for worker in corpus.workers.clone() { + trio.add_worker(&worker); + } + for _ in 0..200 { + corpus.step(&mut trio); + } + let chain = trio + .backends + .iter() + .find(|backend| backend.index.name() == "chain") + .expect("chain backend"); + let before = chain.index.chain_stats().expect("chain stats"); + let blocks = chain.blocks(); + let queries = query_set(&corpus.chains); + for _ in 0..20 { + for query in &queries { + let _ = chain.index.find_matches(query, false); + let _ = chain.index.find_matches(query, true); + } + } + let after = chain.index.chain_stats().expect("chain stats"); + assert_eq!( + format!("{after:?}"), + format!("{before:?}"), + "lookups changed the index's counters" + ); + assert_eq!( + chain.blocks(), + blocks, + "lookups changed the index's content" + ); +} + +/// A clear and a worker removal give the chain index's memory back: nothing +/// stays live, every arena word is back in a free list, and the same corpus +/// stored again after the release is served from what was freed, so neither +/// the run slab nor the arena grows across fill-release cycles (guardrail 3: +/// bounded, recycled). +#[test] +fn chain_index_releases_state_on_clear_and_worker_removal() { + let mut trio = Trio { + backends: vec![Backend::new(KvIndex::chain())], + lookups: 0, + }; + let mut first_fill: Option<(usize, usize)> = None; + for cycle in 0..3 { + let mut corpus = Corpus::new(100, 4); + for worker in corpus.workers.clone() { + trio.add_worker(&worker); + } + for _ in 0..300 { + corpus.step(&mut trio); + } + let chain = &mut trio.backends[0]; + assert!( + chain.index.current_size() > 0, + "cycle {cycle}: nothing indexed" + ); + let filled = chain.index.chain_stats().expect("stats"); + match first_fill { + None => first_fill = Some((filled.runs_allocated, filled.arena_bytes)), + // The root's child table may be rebuilt once more; nothing else may grow. + Some((runs, arena)) => assert!( + filled.runs_allocated <= runs + 8 && filled.arena_bytes <= arena + 16 * 1024, + "cycle {cycle}: the refill grew the index: {filled:?} after {runs} runs, {arena} B" + ), + } + // Half the workers clear and leave, the rest just leave. + let names: Vec = chain.workers.keys().cloned().collect(); + for (n, name) in names.iter().enumerate() { + let (id, mut state) = chain.workers.remove(name).expect("state"); + if n % 2 == 0 { + chain.index.apply_cleared(id, &mut state.blocks); + assert!( + state.blocks.is_empty(), + "cycle {cycle}: {name} cleared state" + ); + assert_eq!( + chain.index.worker_block_count(id), + 0, + "cycle {cycle}: {name}" + ); + } + chain.index.remove_worker(id, state.blocks); + assert_eq!( + chain.index.worker_block_count(id), + 0, + "cycle {cycle}: {name}" + ); + } + let stats = chain.index.chain_stats().expect("stats"); + assert_eq!(stats.runs_live, 0, "cycle {cycle}: live runs after release"); + assert_eq!( + stats.blocks_live, 0, + "cycle {cycle}: live blocks after release" + ); + assert_eq!(chain.index.current_size(), 0, "cycle {cycle}"); + assert_eq!(chain.index.entry_count(), 0, "cycle {cycle}"); + assert!(chain.index.is_empty(), "cycle {cycle}"); + assert!(chain.index.debug_blocks().is_empty(), "cycle {cycle}"); + // Everything but the root's own child table is back in a free list. + assert!( + stats.arena_free_bytes + 16 * 1024 >= stats.arena_bytes, + "cycle {cycle}: arena words not back in free lists: {stats:?}" + ); + } +} diff --git a/model_gateway/src/worker/kv_index_backend/mock_streams.rs b/model_gateway/src/worker/kv_index_backend/mock_streams.rs new file mode 100644 index 0000000000..050642c788 --- /dev/null +++ b/model_gateway/src/worker/kv_index_backend/mock_streams.rs @@ -0,0 +1,806 @@ +//! Engine streams for the index tests and the decision bench, generated +//! in-process by the mock engine (`mock_worker::engine`, the simulator the +//! mock fleet runs). A seeded mix of shared chat-template prefixes, +//! multi-turn sessions and prompts nobody shares drives one or two engines +//! whose prefix cache fills, evicts and announces every change as KV events: +//! stores chained to their parents, prefixes shared across requests and +//! ranks, a cache that turns over many times in a few hundred requests. Each +//! batch leaves its engine on a publisher wire, vLLM's or SGLang's event +//! layout by the mock worker's own encoder, and comes back through the +//! relay's decoder and normalizer, so what reaches the index is what reaches +//! it in production, with nothing recorded and checked in. +//! +//! Beside the payloads the generator keeps the engine's own truth: at +//! checkpoints, taken when the engine has published everything it produced, +//! the keys of the blocks any rank holds on any tier. The index's prefix +//! match for a prompt must equal the engine's at every one of them: the hit +//! a routing decision predicts against the hit the engine serves. + +use std::{ + collections::{BTreeMap, HashMap, HashSet, VecDeque}, + time::Duration, +}; + +use engine_servicer::kv_wire::{Normalizer, WireBatch}; +use engine_zmq_client::codec::TrailingTolerant; +use futures::{ + future::LocalBoxFuture, + stream::{select_all, FuturesUnordered}, + FutureExt, StreamExt, +}; +use mock_worker::{ + engine::{Engine, EngineParams, GenEvent, NewRequest}, + kv_zmq::{encode_batch, Wire}, +}; +use serde_json::{json, Value}; +use smg_grpc_client::common_proto::{ + kv_cache_event, KvBlock, KvBlocksRemoved, KvBlocksStored, KvCacheEvent, KvEventBatch, +}; +use tokio::{runtime::Builder, sync::mpsc, time::timeout}; + +/// The engines' page size in tokens. +pub const BLOCK: usize = 16; +/// KV capacity in blocks: a few hundred requests turn it over many times. +const CAPACITY_BLOCKS: u64 = 256; +/// Blocks the host tier holds before it evicts, oldest write-back first. +const HOST_BLOCKS: usize = 96; +/// Requests in flight at once. +const IN_FLIGHT: usize = 6; +/// Checkpoints asked for per stream, spread over the completed requests. +const CHECKPOINTS: usize = 32; +/// Sessions kept at once; a new one beyond this replaces one at random. +const SESSIONS: usize = 48; +/// A session this long, in blocks, ends; the next turn starts a new one. +const SESSION_BLOCKS: usize = 96; +/// Simulated time the engines get to finish every request. +const SIMULATED_LIMIT: Duration = Duration::from_secs(3600); + +/// The engines and the wire a stream comes from. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum Shape { + /// One vLLM engine on vLLM's event layout (unsigned hashes; medium, + /// cache group and spec kind on every event; the rank in the batch), + /// with vLLM's second physical copies of blocks two requests computed + /// at once (see [`SecondCopies`]), restarted once part-way + /// (`AllBlocksCleared`). + Vllm, + /// One SGLang engine on SGLang's layout (signed hashes, one removal per + /// node, no medium, no rank), scheduling prefill first. + Sglang, + /// Two data-parallel ranks of one vLLM worker, their streams merged by + /// receive order. Sessions stick to a rank but a quarter of their turns + /// cross over, so both ranks hold the same blocks and a block's last + /// copy may leave from either. + TwoRank, + /// One SGLang engine with a host tier: a block leaving the device is + /// written back to the host first and evicted from there later, so the + /// index must hold a block until its last copy on any tier goes. + HostTier, +} + +impl Shape { + pub fn name(self) -> &'static str { + match self { + Shape::Vllm => "vllm", + Shape::Sglang => "sglang", + Shape::TwoRank => "vllm-dp2", + Shape::HostTier => "sglang-hicache", + } + } + + /// The publisher layout; the mock's SGLang layout carries no rank, so + /// the two-rank run rides vLLM's `data_parallel_rank`. + fn wire(self) -> Wire { + match self { + Shape::Vllm | Shape::TwoRank => Wire::Vllm, + Shape::Sglang | Shape::HostTier => Wire::Sglang, + } + } + + fn ranks(self) -> usize { + match self { + Shape::TwoRank => 2, + Shape::Vllm | Shape::Sglang | Shape::HostTier => 1, + } + } +} + +/// One publisher payload as the wire carries it, and the rank it came from. +pub struct Payload { + pub rank: usize, + pub bytes: Vec, +} + +/// The engine's truth once `after` payloads are applied: the keys of the +/// blocks any rank holds on any tier, as [`Engine::block_keys`] names them. +pub struct Checkpoint { + pub after: usize, + pub held: HashSet, +} + +/// A generated stream: its payloads in receive order, every prompt the +/// requests carried, and the checkpoints. +pub struct Stream { + pub payloads: Vec, + pub prompts: Vec>, + pub checkpoints: Vec, +} + +impl Stream { + /// `requests` seeded requests through the shape's engines, on a paused + /// clock: the engines' pass times are simulated, not slept. + pub fn generate(shape: Shape, seed: u64, requests: usize) -> Self { + let runtime = Builder::new_current_thread() + .enable_all() + .start_paused(true) + .build() + .expect("a current-thread runtime"); + runtime.block_on(async { + timeout(SIMULATED_LIMIT, Generator::new(shape, seed, requests).run()) + .await + .expect("the engines finish every request within the simulated limit") + }) + } + + /// The stream as the relay forwards it to one worker: every payload + /// decoded and normalized in order, one normalizer for every rank. + pub fn normalized(&self) -> Vec { + let mut normalizer = Normalizer::new(); + let mut event_id = 0; + self.payloads + .iter() + .enumerate() + .map(|(seq, payload)| { + let batch = rmp_serde::from_slice::>(&payload.bytes) + .expect("a msgpack event batch") + .0; + normalizer.normalize_batch(batch, seq as u64, &mut event_id) + }) + .collect() + } + + /// The engine's prefix match for a prompt against a checkpoint's blocks: + /// its leading full blocks held, as the engine counts a cache hit. + pub fn prefix_match(held: &HashSet, prompt: &[u32]) -> usize { + Engine::block_keys(prompt, BLOCK) + .iter() + .take_while(|key| held.contains(key)) + .count() + } +} + +/// SplitMix64: one seed, one request mix. +struct Rng(u64); + +impl Rng { + fn next(&mut self) -> u64 { + self.0 = self.0.wrapping_add(0x9e37_79b9_7f4a_7c15); + let mut z = self.0; + z = (z ^ (z >> 30)).wrapping_mul(0xbf58_476d_1ce4_e5b9); + z = (z ^ (z >> 27)).wrapping_mul(0x94d0_49bb_1331_11eb); + z ^ (z >> 31) + } + + fn below(&mut self, n: usize) -> usize { + (self.next() % n as u64) as usize + } + + /// Uniform in `[lo, hi]`. + fn range(&mut self, lo: usize, hi: usize) -> usize { + lo + self.below(hi - lo + 1) + } + + fn chance(&mut self, numerator: u64, denominator: u64) -> bool { + self.next() % denominator < numerator + } +} + +/// Token ids of block `position` of token stream `stream`: distinct per pair. +fn tokens(stream: u64, position: usize) -> Vec { + (0..BLOCK as u32) + .map(|i| { + let word = stream + .wrapping_mul(0x2545_f491_4f6c_dd1d) + .wrapping_add(position as u64 * 0x9e37_79b9) + .wrapping_add(u64::from(i)); + (word >> 7) as u32 & 0x3_ffff + }) + .collect() +} + +/// A conversation: the tokens its next turn continues from, its rank, and a +/// generation so a turn finishing after the slot was reused does not land. +struct Session { + tokens: Vec, + rank: usize, + generation: u64, +} + +/// A request the engine is running: the session (slot, generation) it +/// extends, and its prompt, which with the output becomes the session's +/// next starting point. +struct InFlight { + session: Option<(usize, u64)>, + prompt: Vec, +} + +/// A drawn request: its rank, its prompt, and the session (slot, generation) +/// it extends. +struct Draw { + rank: usize, + prompt: Vec, + session: Option<(usize, u64)>, +} + +/// A request's index and output tokens, once it is done. +type Completion = LocalBoxFuture<'static, (usize, Vec)>; + +struct Generator { + shape: Shape, + rng: Rng, + requests: usize, + prefixes: Vec>, + sessions: Vec, + generations: u64, + /// The next token stream nobody has used (see [`tokens`]). + next_stream: u64, + engines: Vec, + in_flight: HashMap, + submitted: usize, + completed: usize, + /// The last sequence number received from each rank. + received: Vec, + checkpoint_due: bool, + host: Option, + copies: Option, + stream: Stream, +} + +impl Generator { + fn new(shape: Shape, seed: u64, requests: usize) -> Self { + let mut rng = Rng(seed); + let prefixes = (0..6) + .map(|p| { + let len = rng.range(2, 6); + (0..len) + .flat_map(|position| tokens(1_000 + p, position)) + .collect() + }) + .collect(); + let params = EngineParams { + kv_capacity_tokens: CAPACITY_BLOCKS * BLOCK as u64, + block_size: BLOCK as u32, + max_running: 8, + max_batched_tokens: 2048, + prefill_first: matches!(shape, Shape::Sglang | Shape::HostTier), + kv_broadcast_capacity: 4096, + ..Default::default() + }; + let engines = (0..shape.ranks()) + .map(|rank| { + Engine::spawn_named( + params.clone(), + format!("{}-rank{rank}", shape.name()), + false, + ) + }) + .collect(); + let host = (shape == Shape::HostTier).then(|| HostTier::new(seed)); + let copies = (shape == Shape::Vllm).then(|| SecondCopies::new(seed)); + Self { + shape, + rng, + requests, + prefixes, + sessions: Vec::new(), + generations: 0, + next_stream: 1_000_000, + engines, + in_flight: HashMap::new(), + submitted: 0, + completed: 0, + received: vec![0; shape.ranks()], + checkpoint_due: false, + host, + copies, + stream: Stream { + payloads: Vec::new(), + prompts: Vec::new(), + checkpoints: Vec::new(), + }, + } + } + + async fn run(mut self) -> Stream { + let mut kv = select_all( + self.engines + .iter() + .enumerate() + .map(|(rank, engine)| engine.subscribe_kv(0).map(move |item| (rank, item))), + ); + let mut completions: FuturesUnordered = FuturesUnordered::new(); + self.fill(&mut completions); + while self.completed < self.requests { + tokio::select! { + biased; + next = kv.next() => { + let Some((rank, Ok(batch))) = next else { break }; + self.receive(rank, batch); + } + Some((index, output)) = completions.next(), if !completions.is_empty() => { + self.finish(index, output); + self.fill(&mut completions); + } + } + } + // The final truth is read once everything produced has been received. + while !self.caught_up() { + let Some((rank, Ok(batch))) = kv.next().await else { + break; + }; + self.receive(rank, batch); + } + self.checkpoint(); + // Dropping the handles ends the engines; their streams end once the + // publishers have released what they still hold. + self.engines.clear(); + while let Some((rank, Ok(batch))) = kv.next().await { + self.receive(rank, batch); + } + self.stream + } + + /// A batch from `rank`: the second copies folded in, the host tier's + /// write-back ahead of it, then the batch on the shape's wire; a due + /// checkpoint once the engines are caught up. + fn receive(&mut self, rank: usize, mut batch: KvEventBatch) { + self.received[rank] = batch.sequence_number; + if let Some(copies) = &mut self.copies { + copies.transform(&mut batch); + } + if let Some(bytes) = self.host.as_mut().and_then(|host| host.write_back(&batch)) { + self.stream.payloads.push(Payload { rank, bytes }); + } + let bytes = encode_batch(&batch, rank as i32, self.shape.wire()); + self.stream.payloads.push(Payload { rank, bytes }); + if self.checkpoint_due && self.caught_up() { + self.checkpoint(); + self.checkpoint_due = false; + } + } + + /// Whether every batch the engines produced has been received. The + /// actor publishes a pass's batch, its cache mirror and its load + /// snapshot together after the pass's simulated time, so an equal count + /// means the engine's cache is exactly what the received batches + /// describe. + fn caught_up(&self) -> bool { + self.engines + .iter() + .zip(&self.received) + .all(|(engine, &received)| { + u64::try_from(engine.load().num_kv_batches).unwrap_or(0) == received + }) + } + + fn checkpoint(&mut self) { + let after = self.stream.payloads.len(); + if self + .stream + .checkpoints + .last() + .is_some_and(|last| last.after == after) + { + return; + } + let mut held: HashSet = self.engines.iter().flat_map(Engine::cache_keys).collect(); + if let Some(host) = &self.host { + held.extend(host.keys()); + } + if let Some(copies) = &self.copies { + held.extend(copies.keys()); + } + self.stream.checkpoints.push(Checkpoint { after, held }); + } + + /// A request finished: its session continues from prompt and output; a + /// checkpoint falls due every so many completions; the vLLM run restarts + /// its engine once part-way. + fn finish(&mut self, index: usize, output: Vec) { + self.completed += 1; + let flight = self.in_flight.remove(&index).expect("a request in flight"); + if let Some((slot, generation)) = flight.session { + if let Some(session) = self + .sessions + .get_mut(slot) + .filter(|session| session.generation == generation) + { + session.tokens = flight.prompt; + session.tokens.extend(output); + } + } + if self + .completed + .is_multiple_of((self.requests / CHECKPOINTS).max(1)) + { + self.checkpoint_due = true; + } + if self.shape == Shape::Vllm && self.completed == self.requests * 3 / 5 { + self.engines[0].reset(); + } + } + + /// Keep [`IN_FLIGHT`] requests running until every request is submitted. + fn fill(&mut self, completions: &mut FuturesUnordered) { + while self.submitted < self.requests && self.in_flight.len() < IN_FLIGHT { + let index = self.submitted; + self.submitted += 1; + let Draw { + rank, + prompt, + session, + } = self.draw(); + let max_new = self.rng.range(8, 64) as u32; + let (events, mut receiver) = mpsc::unbounded_channel(); + self.engines[rank].submit(NewRequest { + request_id: format!("r{index}"), + prompt_token_ids: prompt.clone(), + max_new, + events, + }); + self.stream.prompts.push(prompt.clone()); + self.in_flight.insert(index, InFlight { session, prompt }); + completions.push( + async move { + let mut output = Vec::new(); + while let Some(event) = receiver.recv().await { + match event { + GenEvent::Token { token_id, .. } => output.push(token_id), + GenEvent::Done { .. } => break, + } + } + (index, output) + } + .boxed_local(), + ); + } + } + + /// The next request: a turn on a session (55 in 100), a new session (30), + /// or a prompt nobody shares (15); with its rank and the session it + /// extends. + fn draw(&mut self) -> Draw { + let roll = self.rng.below(100); + if roll < 55 { + if let Some(turn) = self.next_turn() { + return turn; + } + } + if roll < 85 { + return self.new_session(); + } + self.novel() + } + + /// A new turn on a session: its tokens so far (the previous turns and + /// their outputs), one to six blocks of new user tokens and a partial + /// block. A session that has grown long is replaced by a new one in its + /// slot. On two ranks a quarter of the turns go to the other rank. + fn next_turn(&mut self) -> Option { + if self.sessions.is_empty() { + return None; + } + let slot = self.rng.below(self.sessions.len()); + if self.sessions[slot].tokens.len() >= SESSION_BLOCKS * BLOCK { + return Some(self.new_session_at(slot)); + } + let mut rank = self.sessions[slot].rank; + if self.shape.ranks() > 1 && self.rng.chance(1, 4) { + rank = (rank + 1) % self.shape.ranks(); + } + let mut prompt = self.sessions[slot].tokens.clone(); + self.append_blocks(&mut prompt, 1, 6); + Some(Draw { + rank, + prompt, + session: Some((slot, self.sessions[slot].generation)), + }) + } + + /// A new session in a free slot, or in place of one at random. + fn new_session(&mut self) -> Draw { + let slot = if self.sessions.len() < SESSIONS { + self.sessions.len() + } else { + self.rng.below(SESSIONS) + }; + self.new_session_at(slot) + } + + /// A new session on a shared prefix (three in four) or on its own: one + /// to twelve blocks of body and a partial block, on a random rank. + fn new_session_at(&mut self, slot: usize) -> Draw { + let mut prompt = if self.rng.chance(3, 4) { + self.prefixes[self.rng.below(self.prefixes.len())].clone() + } else { + Vec::new() + }; + self.append_blocks(&mut prompt, 1, 12); + let rank = self.rng.below(self.shape.ranks()); + self.generations += 1; + let session = Session { + tokens: prompt.clone(), + rank, + generation: self.generations, + }; + if slot < self.sessions.len() { + self.sessions[slot] = session; + } else { + self.sessions.push(session); + } + Draw { + rank, + prompt, + session: Some((slot, self.generations)), + } + } + + /// A prompt nobody shares: one to sixteen blocks and a partial block. + fn novel(&mut self) -> Draw { + let mut prompt = Vec::new(); + self.append_blocks(&mut prompt, 1, 16); + Draw { + rank: self.rng.below(self.shape.ranks()), + prompt, + session: None, + } + } + + /// `lo` to `hi` fresh blocks and a partial block (none to all but one of + /// its tokens) onto `prompt`. + fn append_blocks(&mut self, prompt: &mut Vec, lo: usize, hi: usize) { + let stream = self.next_stream; + self.next_stream += 1; + let blocks = self.rng.range(lo, hi); + for position in 0..blocks { + prompt.extend(tokens(stream, position)); + } + let partial = self.rng.below(BLOCK); + prompt.extend(&tokens(stream, blocks)[..partial]); + } +} + +/// A host tier behind an engine, as a hierarchical cache keeps one: a block +/// the device evicts is written back to the host first (two in three, when +/// its parent is held somewhere or it is a root), on a batch of its own +/// ahead of the device's; the host evicts its oldest write-backs beyond its +/// capacity. Its events ride SGLang's layout with `medium: "CPU"`. +struct HostTier { + rng: Rng, + /// Tokens and parent of every block the device stored, by engine hash. + blocks: HashMap, Option)>, + /// What the device holds, by the stream so far. + device: HashSet, + /// What the host holds, and the order it was written back in. + held: HashSet, + order: VecDeque, +} + +impl HostTier { + fn new(seed: u64) -> Self { + Self { + rng: Rng(seed ^ 0x5eed_0000_0000_0000), + blocks: HashMap::new(), + device: HashSet::new(), + held: HashSet::new(), + order: VecDeque::new(), + } + } + + /// The host batch a device batch calls for, to go ahead of it. + fn write_back(&mut self, batch: &KvEventBatch) -> Option> { + let mut events = Vec::new(); + for event in &batch.events { + match &event.data { + Some(kv_cache_event::Data::Stored(stored)) => { + let mut parent = stored.parent_block_hash; + for block in &stored.blocks { + self.blocks + .insert(block.block_hash, (block.token_ids.clone(), parent)); + self.device.insert(block.block_hash); + parent = Some(block.block_hash); + } + } + Some(kv_cache_event::Data::Removed(removed)) => { + for &hash in &removed.block_hashes { + self.device.remove(&hash); + events.extend(self.write_back_one(hash)); + } + } + Some(kv_cache_event::Data::Cleared(_)) => { + self.device.clear(); + self.held.clear(); + self.order.clear(); + } + None => {} + } + } + while self.held.len() > HOST_BLOCKS { + let Some(hash) = self.order.pop_front() else { + break; + }; + self.held.remove(&hash); + events.push(json!({"type": "BlockRemoved", "block_hashes": [hash], "medium": "CPU"})); + } + if events.is_empty() { + return None; + } + let batch = json!([batch.timestamp, events, null]); + Some(rmp_serde::to_vec_named(&batch).expect("a msgpack event batch")) + } + + /// The host store for a block the device is evicting, when it is + /// written back. + fn write_back_one(&mut self, hash: i64) -> Option { + if self.held.contains(&hash) || !self.rng.chance(2, 3) { + return None; + } + let (tokens, parent) = self.blocks.get(&hash)?; + let chained = match parent { + None => true, + Some(parent) => self.held.contains(parent) || self.device.contains(parent), + }; + if !chained { + return None; + } + let store = json!({ + "type": "BlockStored", + "block_hashes": [hash], + "parent_block_hash": parent, + "token_ids": tokens, + "block_size": BLOCK, + "lora_id": null, + "medium": "CPU", + }); + self.held.insert(hash); + self.order.push_back(hash); + Some(store) + } + + fn keys(&self) -> impl Iterator + '_ { + self.held.iter().map(|&hash| hash as u64) + } +} + +/// vLLM's second physical copies: when two requests in flight compute the +/// same prefix the engine keeps two blocks under one hash, publishes the +/// second inside the later request's longer store (its first blocks second +/// copies, the rest first copies) and removes each copy on its own. The mock +/// engine keeps one block per hash, so this layer gives a quarter of the +/// chained stores one to four second copies up the parent chain and removes +/// each some batches later. A block with a copy here is held whatever the +/// device did with the first. +struct SecondCopies { + rng: Rng, + /// Tokens and parent of every block the device stored, by engine hash. + blocks: HashMap, Option)>, + /// What the device holds, by the stream so far. + device: HashSet, + /// Blocks whose second copy is held, with the batch count it goes at. + second: BTreeMap, + batches: usize, +} + +impl SecondCopies { + fn new(seed: u64) -> Self { + Self { + rng: Rng(seed ^ 0x2c0b_0000_0000_0000), + blocks: HashMap::new(), + device: HashSet::new(), + second: BTreeMap::new(), + batches: 0, + } + } + + /// Rewrite a device batch: second copies at the head of some stores, + /// the due removals appended. + fn transform(&mut self, batch: &mut KvEventBatch) { + self.batches += 1; + for event in &mut batch.events { + match &mut event.data { + Some(kv_cache_event::Data::Stored(stored)) => { + let mut parent = stored.parent_block_hash; + for block in &stored.blocks { + self.blocks + .insert(block.block_hash, (block.token_ids.clone(), parent)); + self.device.insert(block.block_hash); + parent = Some(block.block_hash); + } + if self.rng.chance(1, 4) { + self.prepend_second_copies(stored); + } + } + Some(kv_cache_event::Data::Removed(removed)) => { + for hash in &removed.block_hashes { + self.device.remove(hash); + } + } + Some(kv_cache_event::Data::Cleared(_)) => { + self.device.clear(); + self.second.clear(); + } + None => {} + } + } + let due: Vec = self + .second + .iter() + .filter(|(_, &at)| at <= self.batches) + .map(|(&hash, _)| hash) + .collect(); + if due.is_empty() { + return; + } + for hash in &due { + self.second.remove(hash); + } + batch.events.push(KvCacheEvent { + event_id: 0, + data: Some(kv_cache_event::Data::Removed(KvBlocksRemoved { + block_hashes: due, + ..Default::default() + })), + }); + } + + /// One to four blocks up from the store's parent, all on the device and + /// without a second copy yet, become the store's first blocks; the store + /// then hangs off the block above them. + fn prepend_second_copies(&mut self, stored: &mut KvBlocksStored) { + let Some(parent) = stored.parent_block_hash else { + return; + }; + let wanted = self.rng.range(1, 4); + let mut copies = Vec::new(); + let mut cursor = Some(parent); + while copies.len() < wanted { + let Some(hash) = cursor else { break }; + if !self.device.contains(&hash) || self.second.contains_key(&hash) { + break; + } + let Some((_, above)) = self.blocks.get(&hash) else { + break; + }; + copies.push(hash); + cursor = *above; + } + if copies.is_empty() { + return; + } + // Collected child to parent; a store lists them root side first. + copies.reverse(); + let mut blocks: Vec = copies + .iter() + .map(|hash| { + let (tokens, _) = self.blocks.get(hash).expect("a stored block"); + KvBlock { + block_hash: *hash, + token_ids: tokens.clone(), + block_size: BLOCK as i32, + ..Default::default() + } + }) + .collect(); + stored.parent_block_hash = self.blocks.get(&copies[0]).and_then(|(_, above)| *above); + blocks.append(&mut stored.blocks); + stored.blocks = blocks; + for hash in copies { + let at = self.batches + self.rng.range(8, 64); + self.second.insert(hash, at); + } + } + + fn keys(&self) -> impl Iterator + '_ { + self.second.keys().map(|&hash| hash as u64) + } +} diff --git a/model_gateway/src/worker/liveness.rs b/model_gateway/src/worker/liveness.rs new file mode 100644 index 0000000000..1f7d777adc --- /dev/null +++ b/model_gateway/src/worker/liveness.rs @@ -0,0 +1,994 @@ +//! Progress-based liveness beside the health check. +//! +//! The health checker needs `failure_threshold * check_interval` to exclude a +//! worker, tens of seconds with the defaults. The transport often knows +//! sooner: a worker that dies or restarts resets its connection and fails its +//! KV event stream, its load poll and its in-flight streams at once; one that +//! falls silent fails them at the keepalive timeout or the poll deadline. +//! This module turns those failures into a routing veto and clears it on the +//! first successful contact. +//! +//! Two vetoes, both read by routing through [`Worker::stall_reason`]: +//! +//! - **unreachable**: a connection failure (load poll, KV event stream) while +//! nothing has been heard from the worker for the stall threshold. Any +//! successful contact clears it. Health probes count as contact when they +//! pass; a failed probe is left to the health state machine, because probe +//! timeouts on a slow but streaming worker are not unreachability. +//! - **wedged**: the worker still answers polls but has produced no token or +//! completion for the wedge bound while it holds in-flight requests and +//! its waiting queue grows. Progress clears it; so does an empty pile (the +//! sweep lifts the veto once nothing is in flight, and a stuck engine +//! re-arms it within the bound on the next requests: steering, never a +//! refusal); and a transport failure followed by silence turns it into an +//! unreachable veto, which contact clears, so a wedged worker whose +//! connection the keepalive tears down is not left waiting for progress +//! from streams that no longer exist. A paused engine that keeps +//! answering health looks exactly like this. The requests counted are the +//! tracked ones, streaming generations to the worker over gRPC, whose +//! responses the gateway sees one by one ([`Worker::tracked_load`]): an +//! HTTP worker, a PD leg or a non-streaming generation gives no signal +//! between dispatch and completion, so it never forms a pile and is never +//! judged by one; a non-streaming generation's one answer is progress all +//! the same, and its prompt counts in the prefill backlog. The clock +//! starts at the first dispatch of a run of +//! tracked requests, never at registration, and the bound is the +//! configured threshold or, if longer, the time the engine may still need +//! to prefill what is in flight ([`Worker::prefill_backlog`]): a batch that +//! is all in prefill streams nothing and is not wedged. A zero threshold +//! turns the rule off. +//! +//! Neither veto touches the worker's health status: the health checker keeps +//! its own state machine, and the veto is simply gone once the worker talks. + +use std::{ + sync::{Arc, OnceLock}, + time::{Duration, Instant}, +}; + +use openai_protocol::worker::WorkerStatus; +use tracing::{info, warn}; + +use super::{worker::StallReason, Worker}; +use crate::observability::metrics::Metrics; + +const DEFAULT_STALL: Duration = Duration::from_secs(2); +const DEFAULT_WEDGE: Duration = Duration::from_secs(3); + +/// How often the sweep runs. +pub(crate) const SWEEP_INTERVAL: Duration = Duration::from_millis(250); + +/// In-flight requests this many deep with no token for the wedge threshold +/// are a wedged engine even when nothing new arrives (a saturated client +/// stops adding to the pile); fewer could be a long prefill. +const WEDGE_PILE: usize = 4; + +/// The prefill-aware wedge bound never exceeds this (unless the configured +/// threshold itself does): a leak in the prefill books must not blind the +/// rule for good. +const WEDGE_BOUND_CAP: Duration = Duration::from_secs(120); + +/// The wedge bound for `worker` now: the configured threshold, or the time +/// the engine may still need before the first token of its in-flight prompts +/// is due, whichever is longer. +fn wedge_bound(worker: &Arc, wedge: Duration) -> Duration { + wedge + .max(worker.prefill_backlog()) + .min(wedge.max(WEDGE_BOUND_CAP)) +} + +static THRESHOLDS: OnceLock<(Duration, Duration)> = OnceLock::new(); +static WARMUP: OnceLock = OnceLock::new(); +static EPOCH: OnceLock = OnceLock::new(); + +/// The warm-up slice: one cache miss in `1 / share` is routed to a warming +/// worker (the least-loaded one) so it builds a cache instead of idling behind +/// the fleet's affinity. A worker is warming for `secs` after it became +/// routable, until its index has gained `blocks` blocks since (a new worker's +/// first cache), or, whatever its age, while its index is thin: holding less +/// than `thin_ratio` of the fleet's level (the median over healthy workers) or +/// nothing at all. A resync after a publisher restart, an `OUT_OF_RANGE` or +/// `DATA_LOSS`, or an engine that came back empty leaves a worker whose every +/// prompt has an overlap elsewhere, so no miss would ever reach it otherwise; +/// it stays warming until it crosses the ratio, however many blocks it regains +/// on the way (the `blocks` cap bounds the age rule only: on the churn run a +/// worker that regrew to 1,765 of a 32,767 level was dropped at the cap after +/// one diversion and idled for 25 minutes). `share == 0` disables; +/// `thin_ratio == 0` keeps the age rule alone. +#[derive(Clone, Copy, Debug, PartialEq)] +pub(crate) struct Warmup { + pub secs: Duration, + pub share: f32, + pub blocks: usize, + pub thin_ratio: f32, + /// One hit in this many goes to a thin worker although another worker + /// holds its prefix (0 disables): on a replay where every request has a + /// holder the miss path never runs, and the slice alone would leave an + /// emptied worker idle for good (see `CacheAwarePolicy::warmup_divert`). + pub divert_every: u64, +} + +impl Warmup { + /// Whether a worker admitted `age` ago, holding `indexed` blocks against a + /// fleet level of `fleet_level`, whose index has gained `growth` blocks + /// since its admission (`None` when the index has never seen it; the + /// baseline restarts when the count drops), is warming up: thin against + /// the fleet, or young and short of its first `blocks`. + pub(crate) fn applies( + &self, + age: Duration, + growth: Option, + indexed: usize, + fleet_level: usize, + ) -> bool { + self.share > 0.0 + && (self.is_thin(indexed, fleet_level) + || (age < self.secs && growth.is_none_or(|blocks| blocks < self.blocks))) + } + + /// Whether an index of `indexed` blocks is thin against the fleet's + /// `fleet_level`: empty, or below `thin_ratio` of the level, once the + /// fleet holds a cache worth catching up to (a level of at least the + /// warm-up `blocks`). A fleet below that (young, or tiny) makes nobody + /// thin; the age rule decides there. + pub(crate) fn is_thin(&self, indexed: usize, fleet_level: usize) -> bool { + self.thin_ratio > 0.0 + && fleet_level >= self.blocks + && (indexed == 0 || (indexed as f64) < fleet_level as f64 * f64::from(self.thin_ratio)) + } + + /// Every how many misses one goes to a warming worker. + pub(crate) fn period(&self) -> u64 { + if self.share <= 0.0 { + return u64::MAX; + } + ((1.0 / f64::from(self.share)).round() as u64).max(1) + } +} + +const DEFAULT_WARMUP: Warmup = Warmup { + secs: Duration::from_secs(60), + share: 0.25, + blocks: 1024, + thin_ratio: 0.5, + divert_every: 8, +}; + +/// Set the warm-up slice from the gateway configuration; the first call wins. +pub(crate) fn configure_warmup(warmup: Warmup) { + let _ = WARMUP.set(warmup); +} + +pub(crate) fn warmup() -> Warmup { + WARMUP.get().copied().unwrap_or(DEFAULT_WARMUP) +} + +/// Set the stall and wedge thresholds from the gateway configuration. The +/// first call wins; the defaults are two and three seconds. +pub(crate) fn configure(stall: Duration, wedge: Duration) { + let _ = THRESHOLDS.set((stall, wedge)); +} + +fn thresholds() -> (Duration, Duration) { + THRESHOLDS + .get() + .copied() + .unwrap_or((DEFAULT_STALL, DEFAULT_WEDGE)) +} + +/// Milliseconds since the gateway started: the clock behind the workers' +/// contact and progress stamps. +pub(crate) fn now_ms() -> u64 { + u64::try_from(EPOCH.get_or_init(Instant::now).elapsed().as_millis()).unwrap_or(u64::MAX) +} + +/// Whether a gRPC status describes the connection rather than the request: +/// what a dead peer, a reset or a failed keepalive produce. A deadline or an +/// engine-side error says the worker is slow or wrong, not gone, and a slow +/// worker that still streams tokens must not flap in and out of routing. +pub(crate) fn is_transport_failure(code: tonic::Code) -> bool { + matches!( + code, + tonic::Code::Unavailable + | tonic::Code::Unknown + | tonic::Code::Cancelled + | tonic::Code::Aborted + ) +} + +/// A successful interaction over the transport (a poll answered, an event +/// batch, a response): the worker is reachable, so an unreachable veto ends +/// here. A worker the health checker demoted while it was gone (its probes +/// failed with its transport) is promoted now rather than at its next +/// scheduled probe; any other demotion is the health checker's to lift, at +/// its success threshold. A reachable transport is not health: an engine +/// whose `/health` answers 503 while it drains still relays KV batches and +/// answers `GetLoads`, and promoting it on each of those would flap it back +/// to Ready within milliseconds of every demotion. +pub(crate) fn on_contact(worker: &Arc) { + worker.note_contact(); + if worker.stall_reason() != Some(StallReason::Unreachable) { + return; + } + set(worker, None, "contact"); + if matches!( + worker.status(), + WorkerStatus::NotReady | WorkerStatus::Failed + ) { + worker.signal_connected(); + } +} + +/// A health probe passed: contact, and the end of an unreachable veto, but +/// no promotion. The health state machine promotes on its success threshold; +/// a probe that short-circuited it would make one passing probe out of a +/// flapping engine's run a return to service. +pub(crate) fn on_probe_passed(worker: &Arc) { + worker.note_contact(); + if worker.stall_reason() == Some(StallReason::Unreachable) { + set(worker, None, "probe"); + } +} + +/// A transport failure from `what`. Vetoes the worker when nothing has been +/// heard from it for the stall threshold; a failure right after a contact is +/// remembered, and [`sweep`] vetoes the worker once the threshold passes +/// without a contact. A worker vetoed as wedged is no exception: the failure +/// took the connection its pile was on, so the streams that were to show +/// progress are gone and only a contact can say the worker is back; its veto +/// becomes `Unreachable`, which the first contact clears. On the one-second +/// keepalive this could not arise, the transport failed before any pile +/// wedged; on the thirty-second profile a partitioned or frozen worker is +/// wedged at 3 s and loses its connection at 40 s, and a wedged veto that +/// waited for progress from streams the keepalive had already failed held +/// the worker out of routing for good. +pub(crate) fn on_contact_failed(worker: &Arc, what: &'static str) { + on_contact_failed_with(worker, what, thresholds().0); +} + +/// [`on_contact_failed`] at the given stall threshold. +fn on_contact_failed_with(worker: &Arc, what: &'static str, stall: Duration) { + worker.note_transport_failure(); + if worker.stall_reason() == Some(StallReason::Unreachable) { + return; + } + if stalled(worker.contact_age(), stall) { + set(worker, Some(StallReason::Unreachable), what); + } +} + +/// Periodic check: a worker whose transport failed and that has then stayed +/// silent for the stall threshold is vetoed now, not at its next failed poll, +/// so the exclusion lands at the threshold itself. +pub(crate) fn sweep(worker: &Arc) { + let (stall, wedge) = thresholds(); + sweep_with(worker, stall, wedge); +} + +/// [`sweep`] at the given thresholds. +fn sweep_with(worker: &Arc, stall: Duration, wedge: Duration) { + let load = worker.tracked_load(); + let previous_load = worker.swap_load_sample(load); + let silent_after_failure = + worker.transport_failure_pending() && stalled(worker.contact_age(), stall); + match worker.stall_reason() { + Some(StallReason::Unreachable) => return, + Some(StallReason::Wedged) => { + // A wedged veto ends with its connection or with its pile. The + // connection failed and the worker has been silent since: the + // streams that were to show progress are gone, so the veto is + // unreachable, which the first contact clears. Nothing in flight: + // the pile that was the evidence is gone, and if the engine is + // still stuck the next requests re-arm the veto within the wedge + // bound; that is steering. + if silent_after_failure { + set( + worker, + Some(StallReason::Unreachable), + "silent since a transport failure", + ); + } else if load == 0 { + set(worker, None, "drained"); + } + return; + } + None => {} + } + if silent_after_failure { + set( + worker, + Some(StallReason::Unreachable), + "silent since a transport failure", + ); + return; + } + // The gateway's own view of a wedged engine: tracked requests pile up on + // it and none has produced a response for the wedge bound. Needs no poll; + // off at a zero threshold. + if !wedge.is_zero() + && wedged_by_pile( + load, + previous_load, + worker.token_progress_age(), + wedge_bound(worker, wedge), + ) + { + set( + worker, + Some(StallReason::Wedged), + "in-flight requests pile up without progress", + ); + } +} + +/// A token or a completion from the worker: progress clears any veto. +pub(crate) fn on_token_progress(worker: &Arc) { + worker.note_token_progress(); + if worker.stall_reason().is_some() { + set(worker, None, "progress"); + } +} + +/// A load report: the engine answers, but does it move? Wedged when tracked +/// requests are in flight, the engine reports a waiting queue (or one that +/// grew since the previous report) and no token or completion arrived within +/// the wedge threshold. +pub(crate) fn on_load_report(worker: &Arc, waiting: i64) { + let (_, wedge) = thresholds(); + on_load_report_with(worker, waiting, wedge); +} + +/// [`on_load_report`] at the given wedge threshold. +fn on_load_report_with(worker: &Arc, waiting: i64, wedge: Duration) { + let previous = worker.swap_waiting_reqs(waiting); + match worker.stall_reason() { + Some(StallReason::Wedged) => { + if worker.token_progress_age() < wedge { + set(worker, None, "progress"); + } + } + Some(StallReason::Unreachable) => {} + None => { + if !wedge.is_zero() + && wedged_by_queue( + worker.tracked_load(), + waiting, + previous, + worker.token_progress_age(), + wedge_bound(worker, wedge), + ) + { + set( + worker, + Some(StallReason::Wedged), + "no progress with a waiting queue", + ); + } + } + } +} + +/// The unreachable rule on its inputs. +fn stalled(contact_age: Duration, stall: Duration) -> bool { + contact_age >= stall +} + +/// The wedged rule from an engine's load report: work in flight, a waiting +/// queue (or one that grew), and silence for the wedge threshold. +fn wedged_by_queue( + in_flight: usize, + waiting: i64, + previous: i64, + token_age: Duration, + wedge: Duration, +) -> bool { + in_flight > 0 && (waiting > 0 || waiting > previous) && token_age >= wedge +} + +/// The wedged rule from the gateway's own counters: the pile of in-flight +/// requests grew, or is already deep, and nothing moved for the threshold. +fn wedged_by_pile( + in_flight: usize, + previous_in_flight: usize, + token_age: Duration, + wedge: Duration, +) -> bool { + token_age >= wedge + && in_flight > 0 + && (in_flight > previous_in_flight || in_flight >= WEDGE_PILE) +} + +fn set(worker: &Arc, reason: Option, cause: &'static str) { + let previous = worker.stall_reason(); + if !worker.set_stall(reason) { + return; + } + match reason { + Some(reason) => { + warn!( + worker_url = %worker.url(), + reason = reason.as_str(), + cause, + contact_age_ms = u64::try_from(worker.contact_age().as_millis()).unwrap_or(u64::MAX), + "Worker vetoed by liveness" + ); + // A veto that changes reason (wedged, then unreachable once its + // connection failed) leaves one gauge up, not two. + if let Some(previous) = previous.filter(|previous| *previous != reason) { + Metrics::set_worker_stalled(worker.url(), previous.as_str(), false); + } + Metrics::set_worker_stalled(worker.url(), reason.as_str(), true); + } + None => { + info!(worker_url = %worker.url(), cause, "Worker re-admitted by liveness"); + worker.note_admitted(); + if let Some(previous) = previous { + // The outage's connection failures opened the circuit breaker + // as well, and it would hold the worker out of routing for + // its timeout (30 s by default) and reopen on one error from + // a stale channel. The contact that clears the veto says the + // worker is back; fresh failures reopen the breaker as usual. + if previous == StallReason::Unreachable { + worker.reset_circuit_breaker(); + } + Metrics::set_worker_stalled(worker.url(), previous.as_str(), false); + } + } + } +} + +#[cfg(test)] +mod tests { + use std::thread; + + use tokio::sync::mpsc; + + use super::*; + use crate::worker::{ + circuit_breaker::CircuitBreakerConfig, event::WorkerConnected, BasicWorkerBuilder, + }; + + fn worker() -> Arc { + Arc::new(BasicWorkerBuilder::new("http://w1:8000").build()) + } + + #[test] + fn unreachable_needs_a_stall_not_just_a_failure() { + assert!(!stalled(Duration::from_millis(500), DEFAULT_STALL)); + assert!(stalled(Duration::from_secs(2), DEFAULT_STALL)); + } + + #[test] + fn wedged_by_queue_needs_in_flight_work_a_queue_and_silence() { + let wedge = DEFAULT_WEDGE; + assert!(wedged_by_queue(3, 5, 2, Duration::from_secs(4), wedge)); + assert!( + wedged_by_queue(3, 5, 5, Duration::from_secs(4), wedge), + "a standing queue counts once the client stops adding to it" + ); + assert!( + !wedged_by_queue(0, 5, 2, Duration::from_secs(4), wedge), + "nothing in flight" + ); + assert!( + !wedged_by_queue(3, 0, 0, Duration::from_secs(4), wedge), + "no queue: a long prefill, not a wedge" + ); + assert!( + !wedged_by_queue(3, 5, 2, Duration::from_secs(1), wedge), + "tokens still flowing" + ); + } + + #[test] + fn wedged_by_pile_needs_growth_or_depth_and_silence() { + let wedge = DEFAULT_WEDGE; + assert!( + wedged_by_pile(2, 1, Duration::from_secs(4), wedge), + "growing" + ); + assert!(wedged_by_pile(4, 4, Duration::from_secs(4), wedge), "deep"); + assert!( + !wedged_by_pile(1, 1, Duration::from_secs(4), wedge), + "one quiet request could be prefilling" + ); + assert!( + !wedged_by_pile(8, 4, Duration::from_secs(1), wedge), + "tokens still flowing" + ); + } + + #[test] + fn the_first_requests_after_registration_are_not_a_wedge() { + // The policy lane saw every mock worker vetoed 3.5 s after + // registration when the first requests arrived: the no-progress clock + // counted from registration. It counts from the dispatch now. + let w = worker(); + thread::sleep(Duration::from_millis(15)); + for _ in 0..3 { + w.increment_load(); + w.note_tracked_started(); + } + let age = w.token_progress_age(); + assert!( + age < Duration::from_millis(10), + "clock started at the dispatch" + ); + sweep(&w); + assert!(w.stall_reason().is_none()); + on_load_report(&w, 3); + assert!(w.stall_reason().is_none()); + // The same pile with a clock that had run from registration would be + // the false positive. + assert!(wedged_by_pile( + 3, + 0, + Duration::from_millis(3_500), + DEFAULT_WEDGE + )); + assert!(!wedged_by_pile(3, 0, age, DEFAULT_WEDGE)); + } + + #[test] + fn a_batch_still_in_prefill_is_not_a_wedge() { + // The GPU lane saw the veto fire with no token for 3 to 5 s while a + // running batch was all in prefill: 128 prompts of 1,152 tokens. + let w = worker(); + for _ in 0..128 { + w.increment_load(); + } + w.note_prefill_started(128 * 1_152); + let bound = wedge_bound(&w, DEFAULT_WEDGE); + assert_eq!( + bound, + Duration::from_millis(14_745), + "cold prior: 10k tokens/s" + ); + let silent = Duration::from_millis(5_300); + assert!(!wedged_by_pile(128, 0, silent, bound)); + assert!(!wedged_by_queue(128, 16, 0, silent, bound)); + assert!( + wedged_by_pile(128, 0, silent, DEFAULT_WEDGE), + "the bare threshold would have fired" + ); + // The first tokens arrive: the backlog is gone and the bound is the + // configured threshold again. + w.note_prefill_ended(128 * 1_152, true); + assert_eq!(wedge_bound(&w, DEFAULT_WEDGE), DEFAULT_WEDGE); + // The bound is capped against leaking books. + w.note_prefill_started(100_000_000); + assert_eq!(wedge_bound(&w, DEFAULT_WEDGE), WEDGE_BOUND_CAP); + assert_eq!( + wedge_bound(&w, Duration::from_secs(600)), + Duration::from_secs(600), + "a longer configured threshold stands" + ); + } + + #[test] + fn a_paused_engine_with_short_prompts_is_still_a_wedge() { + // The pause drill: eight chat prompts of ~150 tokens in flight, the + // engine frozen, health and load polls still answering. + let w = worker(); + for _ in 0..8 { + w.increment_load(); + } + w.note_prefill_started(8 * 150); + let bound = wedge_bound(&w, DEFAULT_WEDGE); + assert_eq!( + bound, DEFAULT_WEDGE, + "0.12 s of prefill is inside the threshold" + ); + let silent = Duration::from_millis(3_200); + assert!(wedged_by_queue(8, 8, 8, silent, bound), "standing queue"); + assert!(wedged_by_pile(8, 8, silent, bound), "deep pile"); + assert!(!wedged_by_pile(8, 8, Duration::from_millis(2_900), bound)); + } + + #[test] + fn a_veto_removes_the_worker_from_routing_and_contact_restores_it() { + let w = worker(); + assert!(w.stall_reason().is_none()); + set(&w, Some(StallReason::Unreachable), "test"); + assert_eq!(w.stall_reason(), Some(StallReason::Unreachable)); + assert!(w.routing_state().stalled); + assert!(!w.routing_state().eligible()); + assert!(!w.is_healthy_and_eligible()); + on_contact(&w); + assert!(w.stall_reason().is_none()); + assert!(!w.routing_state().stalled); + } + + #[test] + fn a_contact_that_clears_the_unreachable_veto_closes_the_breaker() { + let w: Arc = Arc::new( + BasicWorkerBuilder::new("http://w1:8000") + .circuit_breaker_config(CircuitBreakerConfig::default()) + .build(), + ); + for _ in 0..8 { + w.record_circuit_breaker_outcome(false); + } + assert!( + !w.circuit_breaker_can_execute(), + "the outage's failures opened the breaker" + ); + set(&w, Some(StallReason::Unreachable), "test"); + on_contact(&w); + assert!(w.stall_reason().is_none()); + assert!( + w.circuit_breaker_can_execute(), + "back in routing at once, breaker closed" + ); + } + + /// A pile of `n` tracked requests dispatched now. + fn tracked_pile(w: &Arc, n: usize) { + for _ in 0..n { + w.increment_load(); + w.note_tracked_started(); + } + } + + #[test] + fn an_untracked_pile_is_never_a_wedge() { + // An HTTP worker under steady traffic (review finding): four in + // flight for longer than the wedge bound, no token signal because + // HTTP responses are not tracked, and before the fix the sweep vetoed + // it as wedged 3 s into its busy run with nothing to clear it. The + // same for a PD leg and a non-streaming gRPC generation. + let w = worker(); + for _ in 0..4 { + w.increment_load(); + } + thread::sleep(Duration::from_millis(5)); + sweep_with(&w, DEFAULT_STALL, Duration::from_millis(1)); + assert!(w.stall_reason().is_none(), "no signal, no pile, no veto"); + on_load_report_with(&w, 8, Duration::from_millis(1)); + assert!( + w.stall_reason().is_none(), + "a waiting queue without tracked requests is not a wedge either" + ); + assert_eq!(w.tracked_load(), 0); + } + + #[test] + fn a_tracked_pile_without_progress_is_a_wedge_until_a_response() { + let w = worker(); + tracked_pile(&w, 4); + thread::sleep(Duration::from_millis(5)); + sweep_with(&w, DEFAULT_STALL, Duration::from_millis(1)); + assert_eq!( + w.stall_reason(), + Some(StallReason::Wedged), + "four streams and no response past the bound" + ); + on_token_progress(&w); + assert!(w.stall_reason().is_none(), "a response clears it"); + for _ in 0..4 { + w.note_tracked_ended(); + w.decrement_load(); + } + assert_eq!(w.tracked_load(), 0); + } + + #[test] + fn a_completion_inside_the_bound_keeps_a_waiting_pile_routable() { + // Review finding on the pushed head: a saturated engine steadily + // finishing non-streaming requests while four streaming ones wait + // for their first token. The completions are progress (the tracked + // stream records its one answer like a token), so the pile is not a + // wedge; without them the veto still lands. + let w = worker(); + tracked_pile(&w, 4); + thread::sleep(Duration::from_millis(5)); + on_token_progress(&w); + sweep_with(&w, DEFAULT_STALL, Duration::from_millis(50)); + assert!( + w.stall_reason().is_none(), + "finishing work is progress, whatever the request's shape" + ); + thread::sleep(Duration::from_millis(55)); + sweep_with(&w, DEFAULT_STALL, Duration::from_millis(50)); + assert_eq!( + w.stall_reason(), + Some(StallReason::Wedged), + "and without a completion inside the bound the pile is a wedge" + ); + } + + #[test] + fn the_wedged_clock_starts_with_the_first_tracked_request() { + // An untracked request has been in flight for a while when the first + // stream is dispatched: the stream's run starts now. + let w = worker(); + w.increment_load(); + thread::sleep(Duration::from_millis(15)); + assert!(w.token_progress_age() >= Duration::from_millis(15)); + w.increment_load(); + w.note_tracked_started(); + assert!( + w.token_progress_age() < Duration::from_millis(10), + "the clock restarted with the tracked run" + ); + } + + #[test] + fn a_zero_wedge_threshold_turns_the_rule_off() { + let w = worker(); + tracked_pile(&w, 4); + thread::sleep(Duration::from_millis(5)); + sweep_with(&w, DEFAULT_STALL, Duration::ZERO); + assert!(w.stall_reason().is_none(), "--worker-wedge-secs 0"); + on_load_report_with(&w, 8, Duration::ZERO); + assert!(w.stall_reason().is_none()); + assert!( + wedged_by_pile(4, 4, Duration::from_millis(5), Duration::ZERO), + "the bare predicate would fire at once: the guard is in the callers" + ); + } + + #[test] + fn a_transport_failure_turns_a_silent_wedged_worker_unreachable_and_contact_clears_it() { + // The 45 s partition on the 30 s keepalive: wedged at 3 s, the + // connection torn down at 40 s with the streams on it, the link back + // at 48 s, and before this the worker never returned. + let w = worker(); + tracked_pile(&w, 4); + set(&w, Some(StallReason::Wedged), "test"); + on_contact_failed_with(&w, "kv stream", Duration::from_secs(600)); + assert_eq!( + w.stall_reason(), + Some(StallReason::Wedged), + "heard from within the threshold: remembered, the wedge stands" + ); + assert!(w.transport_failure_pending()); + thread::sleep(Duration::from_millis(5)); + on_contact_failed_with(&w, "kv stream", Duration::from_millis(1)); + assert_eq!( + w.stall_reason(), + Some(StallReason::Unreachable), + "silent past the threshold: the pile's connection is gone" + ); + on_contact(&w); + assert!( + w.stall_reason().is_none(), + "the first contact is the return" + ); + assert!(w.circuit_breaker_can_execute(), "with the breaker closed"); + } + + #[test] + fn the_sweep_turns_a_wedged_worker_unreachable_once_silent_after_a_failure() { + let w = worker(); + tracked_pile(&w, 4); + set(&w, Some(StallReason::Wedged), "test"); + on_contact_failed_with(&w, "load poll", Duration::from_secs(600)); + sweep_with(&w, Duration::from_secs(600), DEFAULT_WEDGE); + assert_eq!( + w.stall_reason(), + Some(StallReason::Wedged), + "not silent for the threshold yet" + ); + thread::sleep(Duration::from_millis(5)); + sweep_with(&w, Duration::from_millis(1), DEFAULT_WEDGE); + assert_eq!(w.stall_reason(), Some(StallReason::Unreachable)); + on_contact(&w); + assert!(w.stall_reason().is_none()); + } + + #[test] + fn the_sweep_clears_a_wedged_veto_once_the_pile_is_gone_and_a_new_pile_re_arms_it() { + // Load returns to zero without any contact or progress: the streams + // ended in error, or the clients gave up. + let w = worker(); + tracked_pile(&w, 4); + set(&w, Some(StallReason::Wedged), "test"); + sweep(&w); + assert_eq!(w.stall_reason(), Some(StallReason::Wedged), "still piled"); + for _ in 0..4 { + w.note_tracked_ended(); + w.decrement_load(); + } + assert!(!w.transport_failure_pending()); + sweep(&w); + assert!( + w.stall_reason().is_none(), + "nothing in flight: the evidence is gone" + ); + // The engine is still stuck: the next requests pile up and re-arm + // the veto within the bound. + tracked_pile(&w, 4); + thread::sleep(Duration::from_millis(5)); + sweep_with(&w, DEFAULT_STALL, Duration::from_millis(1)); + assert_eq!(w.stall_reason(), Some(StallReason::Wedged), "re-armed"); + } + + #[test] + fn a_wedged_veto_survives_polls_and_ends_with_progress() { + let w = worker(); + set(&w, Some(StallReason::Wedged), "test"); + on_contact(&w); + assert_eq!( + w.stall_reason(), + Some(StallReason::Wedged), + "answering a poll is not progress" + ); + on_token_progress(&w); + assert!(w.stall_reason().is_none()); + } + + #[test] + fn warm_up_ends_with_time_or_blocks_and_slices_by_share() { + let warmup = DEFAULT_WARMUP; + let s = Duration::from_secs; + // A young fleet: every index empty, nobody thin by comparison, the + // age rule decides. + assert!(warmup.applies(s(10), None, 0, 0), "never indexed"); + assert!(warmup.applies(s(10), Some(100), 100, 100)); + assert!(!warmup.applies(s(61), Some(100), 100, 100), "too old"); + assert!( + !warmup.applies(s(10), Some(2048), 2048, 2048), + "warm already" + ); + assert_eq!(warmup.period(), 4); + let off = Warmup { + share: 0.0, + ..warmup + }; + assert!(!off.applies(Duration::ZERO, None, 0, 0)); + assert_eq!(off.period(), u64::MAX); + } + + #[test] + fn a_worker_emptied_by_a_resync_is_thin_until_it_regrows() { + // The soaks' case: hours after admission a publisher restart cleared + // the worker's index (32,739 -> 180 blocks) while the fleet held + // ~30,000 per worker; every prompt had an overlap elsewhere, so no + // miss ever reached it and it idled for good. + let warmup = DEFAULT_WARMUP; + let old = Duration::from_secs(3_600); + assert!(warmup.is_thin(180, 30_000)); + assert!( + warmup.applies(old, Some(180), 180, 30_000), + "emptied: thin, 180 blocks regrown" + ); + assert!(warmup.applies(old, Some(0), 0, 30_000), "empty outright"); + // Churn c3: the index regrew to 1,765 of a 32,767 level within a + // minute of the clear (the decode blocks of the requests in flight) + // and the cap meant for a new worker's first cache ended the warm-up + // there, thin or not; one diversion, then idle for 25 minutes. + assert!( + warmup.applies(old, Some(1_765), 1_765, 32_767), + "regrown past the warm-up blocks but still thin: served on" + ); + assert!( + warmup.applies(old, Some(16_383), 16_383, 32_767), + "thin until the ratio" + ); + assert!( + !warmup.applies(old, Some(16_384), 16_384, 32_767), + "at the ratio it is back to affinity" + ); + assert!( + !warmup.applies(old, Some(0), 20_000, 30_000), + "two thirds of the fleet's level is not thin at a half" + ); + assert!( + warmup.applies(Duration::from_secs(10), Some(2_000), 2_000, 30_000), + "a young worker in an old fleet is served past its first blocks while thin" + ); + assert!( + !warmup.applies(Duration::from_secs(10), Some(2_000), 2_000, 2_000), + "the cap bounds the age rule: warm at the fleet's level" + ); + assert!( + !warmup.applies(old, None, 0, 0), + "an empty fleet has no level: the age rule alone, and this worker is old" + ); + assert!( + !warmup.applies(old, Some(0), 0, 100), + "a fleet holding less than the warm-up blocks is not worth catching up to" + ); + assert!( + warmup.applies(old, Some(0), 0, 1_024), + "at the warm-up blocks it is" + ); + let age_only = Warmup { + thin_ratio: 0.0, + ..warmup + }; + assert!( + !age_only.applies(old, Some(0), 0, 30_000), + "thin_ratio 0 keeps the age rule alone" + ); + } + + #[test] + fn re_admission_restarts_the_warm_up_clock() { + let w = worker(); + thread::sleep(Duration::from_millis(20)); + let before = w.admitted_age(); + set(&w, Some(StallReason::Unreachable), "test"); + on_contact(&w); + assert!( + w.admitted_age() < before, + "cleared veto counts as an admission" + ); + } + + fn signalling_worker() -> (Arc, mpsc::UnboundedReceiver) { + let (tx, rx) = mpsc::unbounded_channel(); + let w: Arc = Arc::new( + BasicWorkerBuilder::new("http://w1:8000") + .connect_signal_tx(tx) + .build(), + ); + (w, rx) + } + + #[test] + fn contact_asks_for_promotion_only_of_a_worker_that_was_unreachable() { + let (w, mut rx) = signalling_worker(); + on_contact(&w); + assert!(rx.try_recv().is_err(), "a Ready worker needs no promotion"); + // Demoted by its probes while its transport was gone: the first + // contact is its return, promoted now. + set(&w, Some(StallReason::Unreachable), "test"); + w.set_status(WorkerStatus::NotReady); + on_contact(&w); + assert!(w.stall_reason().is_none()); + let signal = rx + .try_recv() + .expect("a demoted worker back from a transport outage is signalled"); + assert_eq!(signal.url, "http://w1:8000"); + // Demoted by its probes alone (a draining engine answering 503 on + // /health while it still relays KV batches and answers GetLoads): + // every batch is a contact and none of them is health. + w.set_status(WorkerStatus::NotReady); + on_contact(&w); + on_contact(&w); + assert!( + rx.try_recv().is_err(), + "a reachable transport does not lift a health demotion" + ); + w.set_status(WorkerStatus::Failed); + on_contact(&w); + assert!(rx.try_recv().is_err(), "nor a Failed one"); + } + + #[test] + fn a_passing_probe_is_contact_but_never_a_promotion() { + let (w, mut rx) = signalling_worker(); + set(&w, Some(StallReason::Unreachable), "test"); + w.set_status(WorkerStatus::NotReady); + thread::sleep(Duration::from_millis(15)); + let before = w.contact_age(); + on_probe_passed(&w); + assert!(w.contact_age() < before, "a passing probe is a contact"); + assert!(w.stall_reason().is_none(), "and ends an unreachable veto"); + assert!( + rx.try_recv().is_err(), + "the health state machine promotes at its success threshold, not the probe" + ); + } + + #[test] + fn a_failure_right_after_contact_is_a_blip() { + let w = worker(); + w.note_contact(); + on_contact_failed(&w, "test"); + assert!(w.stall_reason().is_none()); + assert!( + w.transport_failure_pending(), + "but it is remembered for the sweep" + ); + sweep(&w); + assert!( + w.stall_reason().is_none(), + "the sweep waits for the threshold" + ); + w.note_contact(); + assert!(!w.transport_failure_pending(), "a contact forgets it"); + } +} diff --git a/model_gateway/src/worker/manager.rs b/model_gateway/src/worker/manager.rs index 5fa2a32a44..6f4f878274 100644 --- a/model_gateway/src/worker/manager.rs +++ b/model_gateway/src/worker/manager.rs @@ -28,6 +28,7 @@ use crate::{ observability::metrics::{metrics_labels, Metrics}, worker::{ event::{WorkerConnected, WorkerEvent}, + liveness, load_state::LoadSnapshot, metrics_aggregator::{self, MetricPack}, registry::{WorkerDescriptor, WorkerId}, @@ -543,6 +544,9 @@ async fn apply_probe_completion( metrics_labels::CB_FAILURE }, ); + if probe_ok { + liveness::on_probe_passed(&worker); + } let Some(((), transition)) = registry.apply_if_revision(&worker_id, expected_revision, |current_worker| { @@ -599,10 +603,12 @@ async fn recv_connect_signal( } } -/// Promote a worker whose backend handshake just completed, without waiting -/// for its next scheduled poll. Resolves the URL to a live worker id and flips -/// the status through the revision-checked setter, so a signal that lost a race -/// with a same-URL replacement — or a worker already removed — is discarded. +/// Promote a worker whose backend handshake just completed, or that the +/// health checker demoted and that has just answered a contact (see +/// `worker::liveness::on_contact`), without waiting for its next scheduled +/// poll. Resolves the URL to a live worker id and flips the status through the +/// revision-checked setter, so a signal that lost a race with a same-URL +/// replacement — or a worker already removed — is discarded. fn apply_connect_signal( registry: &Arc, connected: WorkerConnected, @@ -614,6 +620,16 @@ fn apply_connect_signal( debug!(worker_url = %url, "Connect signal for an unknown worker; ignoring"); return; }; + // A Failed worker under removal is on its way out: the contact that + // signalled it was a last poll, not a return. + if config.remove_unhealthy + && registry + .get(&worker_id) + .is_some_and(|worker| worker.status() == WorkerStatus::Failed) + { + debug!(worker_url = %url, "Connect signal for a worker being removed; ignoring"); + return; + } match registry.transition_status_if_revision(&worker_id, revision, WorkerStatus::Ready) { Some((old, new)) => { debug!(worker_url = %url, ?old, ?new, "Promoted worker on connect signal"); @@ -1209,6 +1225,7 @@ impl WorkerManager { dp_rank_count: loads.len() as i32, aggregate: EngineAggregateMetricsSnapshot::from_ranks(&loads), loads, + sampled_at: None, } } @@ -2181,6 +2198,7 @@ mod tests { dp_rank_count: loads.len() as i32, aggregate: None, loads, + sampled_at: None, } } diff --git a/model_gateway/src/worker/mod.rs b/model_gateway/src/worker/mod.rs index fd7ae2d11c..c9f63981b2 100644 --- a/model_gateway/src/worker/mod.rs +++ b/model_gateway/src/worker/mod.rs @@ -10,6 +10,9 @@ pub mod expected_wait; pub mod hash_ring; pub mod http_client; pub mod kv_event_monitor; +mod kv_event_recovery; +pub mod kv_index_backend; +pub(crate) mod liveness; pub(crate) mod load_state; pub mod manager; pub mod metrics_aggregator; @@ -42,6 +45,7 @@ pub use error::{WorkerError, WorkerResult}; pub use hash_ring::HashRing; pub use http_client::WorkerHttpClientCache; pub use kv_event_monitor::KvEventMonitor; +pub use kv_index_backend::{KvIndex, WorkerBlocks}; pub use manager::WorkerManager; pub use monitor::{WorkerLoadManager, WorkerMonitor}; // Re-export UNKNOWN_MODEL_ID from protocols @@ -66,7 +70,7 @@ pub use sampling_defaults::DEFAULT_SAMPLING_PARAMS_LABEL; pub use service::WorkerService; pub(crate) use worker::ConnectionModeExt; pub use worker::{ - AttachedBody, BasicWorker, ConnectionMode, RuntimeType, Worker, WorkerLoadGuard, WorkerType, - DEFAULT_BOOTSTRAP_PORT, MOONCAKE_CONNECTOR, MORIIO_CONNECTOR, MORIIO_MODE_LABEL, - NIXL_CONNECTOR, + AttachedBody, BasicWorker, ConnectionMode, RequestCompletionSink, RuntimeType, Worker, + WorkerLoadGuard, WorkerType, DEFAULT_BOOTSTRAP_PORT, MOONCAKE_CONNECTOR, MORIIO_CONNECTOR, + MORIIO_MODE_LABEL, NIXL_CONNECTOR, }; diff --git a/model_gateway/src/worker/monitor.rs b/model_gateway/src/worker/monitor.rs index 9039bde55a..40a679098c 100644 --- a/model_gateway/src/worker/monitor.rs +++ b/model_gateway/src/worker/monitor.rs @@ -48,10 +48,10 @@ //! channel and merge in the fresh loads. use std::{ - collections::HashMap, + collections::{BTreeMap, HashMap}, fmt::Debug, sync::{Arc, Weak}, - time::Duration, + time::{Duration, Instant}, }; use dashmap::DashMap; @@ -61,7 +61,11 @@ use openai_protocol::worker::{ }; use parking_lot::{Mutex, RwLock}; use reqwest::StatusCode; -use tokio::{sync::broadcast, task::JoinHandle}; +use smg_grpc_client::common_proto::EngineLoad; +use tokio::{ + sync::{broadcast, Notify}, + task::JoinHandle, +}; use tracing::{debug, info, warn}; use crate::{ @@ -69,6 +73,7 @@ use crate::{ policies::PolicyRegistry, worker::{ event::WorkerEvent, + liveness, load_state::{LoadReceiver, LoadSnapshot, LoadState}, ConnectionMode, Worker, WorkerRegistry, }, @@ -337,8 +342,91 @@ pub struct WorkerMonitor { group_handles: Mutex>, event_task: Mutex>>, eviction_flush_task: Mutex>>, + liveness_sweep_task: Mutex>>, + /// Load records pushed on the KV-event streams, per worker and rank, + /// merged into one report per worker the way a poll reports every rank. + pushed_ranks: Mutex>>, + /// When each rank of a worker last pushed a load record, by receipt: + /// the tick's poll decision reads it (see [`PollMode`]). + pushed_at: Mutex>>, + /// Pushed reports waiting for the next coalesced snapshot publish. + pending_pushed: Mutex>, + pushed_notify: Arc, + pushed_flush_task: Mutex>>, +} + +/// A worker's report built from pushed records, with the worker it came from +/// for the publish fence. +type PushedReport = (Arc, Arc); + +/// Where a worker's load comes from at a tick. An engine that reports its +/// load on its KV-event stream is not asked for it again: in gRPC mode the +/// `GetLoads` poll is the fallback, for servicers that predate the pushed +/// record, for a stream that is down and for heartbeats that stopped. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub(crate) enum PollMode { + /// No pushed record on file: the poll is the only source. + Poll, + /// Pushed records on file, but some pushing rank's newest is older than + /// the tick interval: the poll takes over until the records resume. + Fallback, + /// Every pushing rank delivered a record within the tick interval: the + /// stream is the source, no RPC. + Suppressed, +} + +impl PollMode { + /// The `mode` label of `smg_engine_load_polls_total`. + pub(crate) fn as_str(self) -> &'static str { + match self { + Self::Poll => "poll", + Self::Fallback => "fallback", + Self::Suppressed => "skipped_fresh_push", + } + } +} + +/// Telemetry a record does not carry (an event batch's record is the core +/// only; a heartbeat's and a stream's first carry it all) stays what the +/// last record that carried it said, so the gauges and `GET /loads` keep +/// the engine's figures between heartbeats. Absent is absent: a section the +/// engine never reports stays unset. +fn keep_telemetry( + fresh: &mut SchedulerLoadSnapshot, + previous: &SchedulerLoadSnapshot, + record: &EngineLoad, +) { + if record.cache_hit_rate.is_none() { + fresh.cache_hit_rate = previous.cache_hit_rate; + } + if record.num_used_tokens.is_none() { + fresh.num_used_tokens = previous.num_used_tokens; + } + if record.max_total_num_tokens.is_none() { + fresh.max_total_num_tokens = previous.max_total_num_tokens; + } + if record.memory.is_none() { + fresh.memory.clone_from(&previous.memory); + } + if record.queues.is_none() { + fresh.queues.clone_from(&previous.queues); + } + if record.disaggregation.is_none() { + fresh.kv_transfer_latency_ms = previous.kv_transfer_latency_ms; + fresh.kv_transfer_speed_gb_s = previous.kv_transfer_speed_gb_s; + fresh.prefill_queue_reqs = previous.prefill_queue_reqs; + fresh.decode_queue_reqs = previous.decode_queue_reqs; + fresh.disagg_mode.clone_from(&previous.disagg_mode); + } } +/// A pushed record is sampled this much before its receipt on top of its own +/// `age_ms`: the servicer-to-gateway one-way latency. +const PUSHED_ONE_WAY_MARGIN: Duration = Duration::from_millis(20); +/// Pushed reports are published into the shared snapshot in one rebuild per +/// window: the rebuild is O(fleet), the records arrive per scheduler step. +const PUSHED_PUBLISH_WINDOW: Duration = Duration::from_millis(100); + /// Debounce window for batching worker evictions into one snapshot rebuild. /// Registry churn (a rollout, a scale-down) emits removals as a gradual /// stream; waiting a beat lets the whole wave land in a single publish while @@ -379,9 +467,165 @@ impl WorkerMonitor { group_handles: Mutex::new(HashMap::new()), event_task: Mutex::new(None), eviction_flush_task: Mutex::new(None), + liveness_sweep_task: Mutex::new(None), + pushed_ranks: Mutex::new(HashMap::new()), + pushed_at: Mutex::new(HashMap::new()), + pending_pushed: Mutex::new(HashMap::new()), + pushed_notify: Arc::new(Notify::new()), + pushed_flush_task: Mutex::new(None), + } + } + + /// A load record pushed on a worker's KV-event stream (`KvEventBatch.load`, + /// the servicer's `GetLoads` figures for `dp_rank`, received at + /// `received_at`): the same input as a poll of that worker. It replaces + /// the rank's entry in the worker's report, goes to the overload verdict, + /// the wedged rule, the load-aware policies and the `smg_engine_*` gauges + /// at once, and into the shared snapshot in the next coalesced publish. + /// The report's `sampled_at` is the receipt less the record's age and + /// the one-way margin, what a policy resets its in-flight tally against. + pub fn apply_pushed_load( + self: &Arc, + worker: &Arc, + dp_rank: i32, + record: &EngineLoad, + received_at: Instant, + ) { + let age = Duration::from_millis(u64::from(record.age_ms)) + PUSHED_ONE_WAY_MARGIN; + let sampled_at = received_at.checked_sub(age).unwrap_or(received_at); + let url = worker.url().to_string(); + // The poll decision reads the receipt, not the sample's age: a record + // that arrives is a stream that works, whatever the engine's clock. + self.pushed_at + .lock() + .entry(url.clone()) + .or_default() + .insert(dp_rank, received_at); + let response = { + let mut pushed = self.pushed_ranks.lock(); + let ranks = pushed.entry(url.clone()).or_default(); + let previous = ranks.get(&dp_rank).cloned().unwrap_or_default(); + // What a poll would have put in the rank's entry, with the + // telemetry the record left out (an event batch's record is the + // core only) kept from the last record that carried it. + let mut fresh = SchedulerLoadSnapshot::from(record); + fresh.dp_rank = dp_rank; + keep_telemetry(&mut fresh, &previous, record); + ranks.insert(dp_rank, fresh); + Arc::new(WorkerLoadResponse { + timestamp: chrono::Utc::now().to_rfc3339(), + version: "pushed".to_string(), + dp_rank_count: i32::try_from(ranks.len()).unwrap_or(i32::MAX), + loads: ranks.values().cloned().collect(), + aggregate: None, + sampled_at: Some(sampled_at), + }) + }; + // What a poll does with a fresh report, minus the DP-rank token cache + // (a record carries no absolute token counts). + let overload = worker.metadata().overload; + if overload.is_enabled() { + self.worker_registry + .set_worker_overloaded(worker, overload.is_overloaded(&response)); + } + let waiting: i64 = response + .loads + .iter() + .map(|rank| i64::from(rank.num_waiting_reqs)) + .sum(); + liveness::on_load_report(worker, waiting); + let single: HashMap = + HashMap::from([(url.clone(), (*response).clone())]); + for policy in self.policy_registry.get_all_load_aware_policies() { + policy.update_loads(&single); + } + Metrics::record_engine_load(&url, worker.model_id(), &response); + // The DP-rank token cache takes a report with absolute token counts, + // as it takes a poll; a core-only record leaves its last-known-good + // entry alone. + if response.has_absolute_token_data() && response.ranks_are_dp_ranks() { + self.worker_load_manager + .update_dp_loads(&HashMap::from([(url.clone(), response.dp_rank_loads())])); + } + self.pending_pushed + .lock() + .insert(url, (Arc::clone(worker), response)); + self.start_pushed_flusher(); + self.pushed_notify.notify_one(); + } + + /// The poll decision for `url` at `now` from the receipt times of its + /// pushed records: none on file polls, a record within `interval` on + /// every pushing rank suppresses the poll, an older one on any rank + /// falls back to it. + pub(crate) fn poll_mode(&self, url: &str, now: Instant, interval: Duration) -> PollMode { + let pushed = self.pushed_at.lock(); + match pushed.get(url) { + Some(ranks) if !ranks.is_empty() => { + let fresh = ranks + .values() + .all(|&at| now.saturating_duration_since(at) < interval); + if fresh { + PollMode::Suppressed + } else { + PollMode::Fallback + } + } + _ => PollMode::Poll, } } + /// The coalescing publisher of pushed reports, started on the first one: + /// every window it moves what accumulated into the shared snapshot in one + /// rebuild. Holds the monitor weakly, like the eviction flusher. + fn start_pushed_flusher(self: &Arc) { + let mut task = self.pushed_flush_task.lock(); + if task.as_ref().is_some_and(|task| !task.is_finished()) { + return; + } + let monitor = Arc::downgrade(self); + let notify = Arc::clone(&self.pushed_notify); + #[expect( + clippy::disallowed_methods, + reason = "ends when the monitor is dropped; nothing waits on it at shutdown" + )] + let handle = tokio::spawn(async move { + loop { + notify.notified().await; + tokio::time::sleep(PUSHED_PUBLISH_WINDOW).await; + let Some(monitor) = monitor.upgrade() else { + return; + }; + monitor.flush_pushed(); + } + }); + *task = Some(handle); + } + + /// Publish every pushed report that accumulated since the last flush. + pub(crate) fn flush_pushed(&self) { + let pending: Vec = self + .pending_pushed + .lock() + .drain() + .map(|(_, entry)| entry) + .collect(); + if pending.is_empty() { + return; + } + let urls: Vec = pending + .iter() + .map(|(worker, _)| worker.url().to_string()) + .collect(); + self.load_state.publish_group(&urls, pending); + } + + /// Pushed reports waiting to be published (tests). + #[cfg(test)] + pub(crate) fn pending_pushed_len(&self) -> usize { + self.pending_pushed.lock().len() + } + /// The current published load snapshot — what routing is acting on, and /// what `GET /loads` serves. pub(crate) fn load_snapshot(&self) -> Arc { @@ -448,6 +692,18 @@ impl WorkerMonitor { }); *self.eviction_flush_task.lock() = Some(flush_handle); + + // The liveness sweep: vetoes a worker that stayed silent for the stall + // threshold after a transport failure (see `worker::liveness::sweep`). + let monitor = Arc::downgrade(self); + #[expect( + clippy::disallowed_methods, + reason = "liveness sweep runs for the monitor's lifetime; the JoinHandle is stored on the monitor and aborted in Drop" + )] + let sweep_handle = tokio::spawn(async move { + liveness_sweep_loop(monitor).await; + }); + *self.liveness_sweep_task.lock() = Some(sweep_handle); } /// Stop every per-group polling loop and clear the shared load @@ -603,6 +859,10 @@ impl WorkerMonitor { self.worker_registry.set_worker_overloaded(worker, false); self.worker_load_manager.remove_worker(url); self.native_loads_memo.remove(url); + // The pushed-record state goes with the worker: a replacement is + // polled until its own stream pushes. + self.pushed_ranks.lock().remove(url); + self.pushed_at.lock().remove(url); self.load_state.enqueue_eviction(Arc::clone(worker)); } @@ -693,6 +953,7 @@ impl WorkerMonitor { } } if response.is_some() { + liveness::on_contact(worker); return response; } @@ -844,11 +1105,31 @@ impl WorkerMonitor { } }; - match backend_client.get_loads().await { - Ok(load) if !load.loads.is_empty() => Some(load), - Ok(_) => None, - Err(e) => { + // A deadline, so one half-open connection cannot stall the whole + // group's poll (the group awaits every worker's poll together). A + // poll that times out is a slow worker, not a dead one: the keepalive + // on the connection is what reports those. + match tokio::time::timeout(LOAD_POLL_DEADLINE, backend_client.get_loads()).await { + Ok(Ok(load)) if !load.loads.is_empty() => { + liveness::on_contact(worker); + Some(load) + } + Ok(Ok(_)) => { + liveness::on_contact(worker); + None + } + Ok(Err(e)) => { debug!("backend GetLoads failed for {}: {e}", worker.url()); + if liveness::is_transport_failure(e.code()) { + liveness::on_contact_failed(worker, "load poll"); + } + None + } + Err(_) => { + debug!( + "backend GetLoads for {} did not answer within {LOAD_POLL_DEADLINE:?}", + worker.url() + ); None } } @@ -863,6 +1144,9 @@ impl Drop for WorkerMonitor { if let Some(handle) = self.eviction_flush_task.get_mut().take() { handle.abort(); } + if let Some(handle) = self.liveness_sweep_task.get_mut().take() { + handle.abort(); + } for (_, state) in self.group_handles.get_mut().drain() { state.handle.abort(); } @@ -1013,6 +1297,26 @@ async fn run_event_loop( /// as `run_event_loop`. The temporary `Arc` is upgraded after the /// timer tick and dropped before the next tick so the monitor's /// `Drop` is reachable. +/// How long a backend load poll may take before the tick moves on without it. +const LOAD_POLL_DEADLINE: Duration = Duration::from_secs(3); + +/// Run [`liveness::sweep`] over every registered worker every quarter second. +async fn liveness_sweep_loop(monitor: Weak) { + let mut ticker = tokio::time::interval(liveness::SWEEP_INTERVAL); + loop { + ticker.tick().await; + let Some(monitor) = monitor.upgrade() else { + return; + }; + for worker in monitor + .worker_registry + .get_workers_filtered(None, None, None, None, false) + { + liveness::sweep(&worker); + } + } +} + async fn group_monitor_loop( monitor: Weak, group_key: WorkerGroupKey, @@ -1027,46 +1331,76 @@ async fn group_monitor_loop( debug!("WorkerMonitor was dropped; exiting group loop for {group_key}"); return; }; + poll_group_once(&monitor, &group_key, interval, Instant::now()).await; + // Drop the temporary strong reference so we do not keep the + // monitor alive across the next `interval_timer.tick().await`. + drop(monitor); + } +} - // Only poll Ready workers — Pending/NotReady/Failed do not - // serve traffic and should not contribute load samples. - let workers: Vec> = monitor - .worker_registry - .get_workers_filtered( - Some(&group_key.model_id), - Some(group_key.worker_type), - Some(group_key.connection_mode), - None, - false, - ) - .into_iter() - .filter(|w| w.status() == WorkerStatus::Ready) - .collect(); - - if workers.is_empty() { - debug!("No Ready workers in group {group_key}, skipping"); - drop(monitor); - continue; +/// One tick of a group at `now`: the poll decision per worker, the polls of +/// those not served by their stream, and the publish of what they returned. +async fn poll_group_once( + monitor: &Arc, + group_key: &WorkerGroupKey, + interval: Duration, + now: Instant, +) { + // Only poll Ready workers — Pending/NotReady/Failed do not + // serve traffic and should not contribute load samples. + let workers: Vec> = monitor + .worker_registry + .get_workers_filtered( + Some(&group_key.model_id), + Some(group_key.worker_type), + Some(group_key.connection_mode), + None, + false, + ) + .into_iter() + .filter(|w| w.status() == WorkerStatus::Ready) + .collect(); + + if workers.is_empty() { + debug!("No Ready workers in group {group_key}, skipping"); + return; + } + + // Polling is unconditional by default so registration alone gives a + // worker live load state. `--disable-load-monitoring` restores the + // conditional gate: a load-aware policy, engine-metrics re-export, or + // overload protection on any group member still forces the poll — + // never "never poll". + let load_aware_policies = monitor.policy_registry.get_all_load_aware_policies(); + if monitor.conditional_polling { + let routing_needs_load = !load_aware_policies.is_empty() + || monitor.policy_registry.get_dp_rank_policy().is_some(); + let overload_needs_load = workers.iter().any(|w| w.metadata().overload.is_enabled()); + if !routing_needs_load && !monitor.engine_metrics && !overload_needs_load { + debug!("Load monitoring disabled and nothing needs the data, skipping load fetch for group {group_key}"); + return; } + } - // Polling is unconditional by default so registration alone gives a - // worker live load state. `--disable-load-monitoring` restores the - // conditional gate: a load-aware policy, engine-metrics re-export, or - // overload protection on any group member still forces the poll — - // never "never poll". - let load_aware_policies = monitor.policy_registry.get_all_load_aware_policies(); - if monitor.conditional_polling { - let routing_needs_load = !load_aware_policies.is_empty() - || monitor.policy_registry.get_dp_rank_policy().is_some(); - let overload_needs_load = workers.iter().any(|w| w.metadata().overload.is_enabled()); - if !routing_needs_load && !monitor.engine_metrics && !overload_needs_load { - debug!("Load monitoring disabled and nothing needs the data, skipping load fetch for group {group_key}"); - drop(monitor); - continue; - } + // A worker whose stream pushed a load record on every pushing rank + // within the interval is served by that stream: its report, verdicts + // and gauges are already current from `apply_pushed_load`, and nothing + // here touches its entry. The others are polled: no record ever + // (`poll`) or records gone quiet (`fallback`). + let mut polled: Vec> = Vec::with_capacity(workers.len()); + for worker in &workers { + let mode = monitor.poll_mode(worker.url(), now, interval); + Metrics::record_engine_load_poll(worker.url(), mode.as_str()); + if mode != PollMode::Suppressed { + polled.push(Arc::clone(worker)); } + } - let futures: Vec<_> = workers + let results = if polled.is_empty() { + debug!("Every worker in group {group_key} pushed its load within {interval:?}; no poll"); + Vec::new() + } else { + let futures: Vec<_> = polled .iter() .map(|worker| { let native_loads_memo = Arc::clone(&monitor.native_loads_memo); @@ -1085,110 +1419,124 @@ async fn group_monitor_loop( } }) .collect(); + future::join_all(futures).await + }; - let results = future::join_all(futures).await; - - let mut group_loads: HashMap = HashMap::new(); - let mut group_dp_loads: HashMap> = HashMap::new(); - let mut dp_evict: Vec = Vec::new(); - for (worker, response) in results { - let url = worker.url().to_string(); - // The overload predicate runs exactly here, once per report, never - // on a request path, against the worker's effective thresholds - // (resolved at registration). A failed fetch means no fresh - // signal, which clears the flag — absent means no opinion. - let overload = worker.metadata().overload; - if overload.is_enabled() { - let verdict = response - .as_ref() - .is_some_and(|load| overload.is_overloaded(load)); - monitor - .worker_registry - .set_worker_overloaded(&worker, verdict); - } - if let Some(load) = response { - // Only feed the DP-rank cache from responses that carry real - // absolute per-rank token counts. Ratio-only snapshots, - // which would otherwise poison with a fake `{0: 0}` - // entry and collapse DP routing onto rank 0. - // - // A fleet rollup from a gateway worker is keyed by downstream - // worker, not by rank, so its repeated `dp_rank: 0` entries - // would overwrite each other down to a single bogus rank. - if load.has_absolute_token_data() && load.ranks_are_dp_ranks() { - group_dp_loads.insert(url.clone(), load.dp_rank_loads()); - } else { - dp_evict.push(url.clone()); - } - group_loads.insert(url, load); + let mut group_loads: HashMap = HashMap::new(); + let mut group_dp_loads: HashMap> = HashMap::new(); + let mut dp_evict: Vec = Vec::new(); + for (worker, response) in results { + let url = worker.url().to_string(); + // The overload predicate runs exactly here, once per report, never + // on a request path, against the worker's effective thresholds + // (resolved at registration). A failed fetch means no fresh + // signal, which clears the flag — absent means no opinion. + let overload = worker.metadata().overload; + if overload.is_enabled() { + let verdict = response + .as_ref() + .is_some_and(|load| overload.is_overloaded(load)); + monitor + .worker_registry + .set_worker_overloaded(&worker, verdict); + } + if let Some(mut load) = response { + load.sampled_at = Some(Instant::now()); + // The wedged rule reads the queue depth off every report. + let waiting: i64 = load + .loads + .iter() + .map(|rank| i64::from(rank.num_waiting_reqs)) + .sum(); + liveness::on_load_report(&worker, waiting); + // Only feed the DP-rank cache from responses that carry real + // absolute per-rank token counts. Ratio-only snapshots, + // which would otherwise poison with a fake `{0: 0}` + // entry and collapse DP routing onto rank 0. + // + // A fleet rollup from a gateway worker is keyed by downstream + // worker, not by rank, so its repeated `dp_rank: 0` entries + // would overwrite each other down to a single bogus rank. + if load.has_absolute_token_data() && load.ranks_are_dp_ranks() { + group_dp_loads.insert(url.clone(), load.dp_rank_loads()); + } else { + dp_evict.push(url.clone()); } + group_loads.insert(url, load); } + } - // Compute the URL set up front so both the success and - // empty-fetch branches can prune stale entries from the watch - // snapshot. Without the empty-fetch prune, a group that - // starts timing out keeps publishing its previous tick's - // loads forever — subscribers see a stale snapshot indefinitely. - let all_group_urls: Vec = workers.iter().map(|w| w.url().to_string()).collect(); + // The poll is also the reconciliation tick: each policy trims what it + // booked on a worker to the requests the router still holds there, + // the safety net behind the load guard's completion report. It runs + // for every worker, polled or served by its stream, whether or not a + // fetch succeeded; the live count is the router's. + for worker in &workers { + monitor.policy_registry.reconcile_in_flight(worker.as_ref()); + } - if group_loads.is_empty() { - debug!("No loads fetched for group {group_key}, pruning stale entries"); - monitor - .load_state - .publish_group(&all_group_urls, Vec::new()); - // The DP cache deliberately keeps last-known-good entries - // so routing decisions still have a hint to fall back to - // when the upstream is briefly unreachable. - drop(monitor); - continue; - } + if polled.is_empty() { + return; + } - debug!( - "Fetched loads from {}/{} workers in group {group_key}", - group_loads.len(), - workers.len() - ); + // The URL set the publish prunes: the polled workers only, so a worker + // served by its stream keeps its entry (the pushed path owns it). + // Without the empty-fetch prune, a group that starts timing out keeps + // publishing its previous tick's loads forever — subscribers see a + // stale snapshot indefinitely. + let polled_urls: Vec = polled.iter().map(|w| w.url().to_string()).collect(); - for policy in &load_aware_policies { - policy.update_loads(&group_loads); - } - monitor.worker_load_manager.update_dp_loads(&group_dp_loads); + if group_loads.is_empty() { + debug!("No loads fetched for group {group_key}, pruning stale entries"); + monitor.load_state.publish_group(&polled_urls, Vec::new()); + // The DP cache deliberately keeps last-known-good entries + // so routing decisions still have a hint to fall back to + // when the upstream is briefly unreachable. + return; + } - if !dp_evict.is_empty() { - monitor.worker_load_manager.remove_workers(&dp_evict); - } + debug!( + "Fetched loads from {}/{} polled workers in group {group_key}", + group_loads.len(), + polled.len() + ); - // Every successful load poll is also the canonical observability - // sample. Load-aware policies already require this poll; the explicit - // engine-metrics option only forces polling when routing does not. - // Reusing the response avoids a second Engine RPC. - for (url, load) in &group_loads { - Metrics::record_engine_load(url, &group_key.model_id, load); - } + for policy in &load_aware_policies { + policy.update_loads(&group_loads); + } + monitor.worker_load_manager.update_dp_loads(&group_dp_loads); - // Merge into the shared snapshot in one rebuild: clear stale entries - // for *this group's* URLs first, then insert the fresh loads — each - // paired with the worker that produced it so the incarnation fence - // can drop reports that raced a removal or replacement. Workers that - // failed this tick get their stale entries pruned along with the - // rest. The responses move (no deep clones): the policy push and the - // metrics pass above already took their references. - let worker_by_url: HashMap<&str, &Arc> = - workers.iter().map(|w| (w.url(), w)).collect(); - let fresh: Vec<(Arc, Arc)> = group_loads - .into_iter() - .filter_map(|(url, load)| { - worker_by_url - .get(url.as_str()) - .map(|worker| (Arc::clone(worker), Arc::new(load))) - }) - .collect(); - monitor.load_state.publish_group(&all_group_urls, fresh); + if !dp_evict.is_empty() { + monitor.worker_load_manager.remove_workers(&dp_evict); + } - // Drop the temporary strong reference so we do not keep the - // monitor alive across the next `interval_timer.tick().await`. - drop(monitor); + // Every successful load poll is also the canonical observability + // sample. Load-aware policies already require this poll; the explicit + // engine-metrics option only forces polling when routing does not. + // Reusing the response avoids a second Engine RPC. (A pushed record + // feeds the same gauges from `apply_pushed_load`.) + for (url, load) in &group_loads { + Metrics::record_engine_load(url, &group_key.model_id, load); } + + // Merge into the shared snapshot in one rebuild: clear stale entries + // for the polled URLs first, then insert the fresh loads — each + // paired with the worker that produced it so the incarnation fence + // can drop reports that raced a removal or replacement. Workers that + // failed this tick get their stale entries pruned along with the + // rest. The responses move (no deep clones): the policy push and the + // metrics pass above already took their references. + let worker_by_url: HashMap<&str, &Arc> = + polled.iter().map(|w| (w.url(), w)).collect(); + let fresh: Vec<(Arc, Arc)> = group_loads + .into_iter() + .filter_map(|(url, load)| { + worker_by_url + .get(url.as_str()) + .map(|worker| (Arc::clone(worker), Arc::new(load))) + }) + .collect(); + monitor.load_state.publish_group(&polled_urls, fresh); } #[cfg(test)] @@ -1274,6 +1622,7 @@ mod worker_monitor_tests { use super::*; use crate::{ config::types::PolicyConfig, + observability::metrics::test_support::render_with_recorder, policies::PolicyRegistry, worker::{BasicWorkerBuilder, ConnectionMode, WorkerType}, }; @@ -1305,6 +1654,119 @@ mod worker_monitor_tests { (registry, monitor) } + /// A pushed record is a poll of its worker: the rank's entry is replaced + /// (other ranks kept), `sampled_at` is the receipt less the record's age + /// and the one-way margin, and the shared snapshot sees it at the next + /// coalesced publish. + #[tokio::test] + async fn a_pushed_load_record_becomes_the_workers_report_at_the_next_publish() { + let (registry, monitor) = build_monitor(); + let worker = ready_worker("grpc://w1:9000", "llama-3"); + registry.register(Arc::clone(&worker)).unwrap(); + let received = Instant::now(); + let record = EngineLoad { + running_requests: 7, + waiting_requests: 2, + waiting_uncached_tokens: Some(3_000), + token_usage: 0.4, + gen_throughput: 900.0, + max_running_requests: 64, + age_ms: 30, + sample: 1, + load_only: false, + ..Default::default() + }; + monitor.apply_pushed_load(&worker, 0, &record, received); + assert_eq!(monitor.pending_pushed_len(), 1); + assert!( + monitor + .load_state + .snapshot() + .get("grpc://w1:9000") + .is_none(), + "not published until the window closes" + ); + monitor.flush_pushed(); + let snapshot = monitor.load_state.snapshot(); + let report = snapshot.get("grpc://w1:9000").expect("published"); + assert_eq!(report.version, "pushed"); + assert_eq!(report.loads.len(), 1); + let rank0 = &report.loads[0]; + assert_eq!( + ( + rank0.dp_rank, + rank0.num_running_reqs, + rank0.num_waiting_reqs, + rank0.num_waiting_uncached_tokens, + rank0.num_total_reqs, + rank0.max_running_requests + ), + (0, 7, 2, 3_000, 9, 64) + ); + assert!((rank0.token_usage - 0.4).abs() < f64::EPSILON); + assert!((rank0.gen_throughput - 900.0).abs() < f64::EPSILON); + let sampled_at = report.sampled_at.expect("sampled_at"); + assert_eq!( + received.duration_since(sampled_at), + Duration::from_millis(30) + PUSHED_ONE_WAY_MARGIN + ); + assert_eq!(monitor.pending_pushed_len(), 0); + + // A second rank joins the report; a later record replaces only its rank. + monitor.apply_pushed_load( + &worker, + 1, + &EngineLoad { + running_requests: 1, + waiting_uncached_tokens: None, + ..record.clone() + }, + Instant::now(), + ); + monitor.apply_pushed_load( + &worker, + 0, + &EngineLoad { + running_requests: 8, + ..record.clone() + }, + Instant::now(), + ); + monitor.flush_pushed(); + let snapshot = monitor.load_state.snapshot(); + let report = snapshot.get("grpc://w1:9000").expect("published"); + let ranks: Vec<(i32, i32, i32)> = report + .loads + .iter() + .map(|rank| { + ( + rank.dp_rank, + rank.num_running_reqs, + rank.num_waiting_uncached_tokens, + ) + }) + .collect(); + assert_eq!(ranks, vec![(0, 8, 3_000), (1, 1, 0)]); + assert_eq!(report.dp_rank_count, 2); + } + + /// A record for a worker the registry no longer holds as this incarnation + /// is dropped at publish by the fence, like a late poll. + #[tokio::test] + async fn a_pushed_record_from_a_replaced_worker_is_fenced_at_publish() { + let (registry, monitor) = build_monitor(); + let worker = ready_worker("grpc://w1:9000", "llama-3"); + registry.register(Arc::clone(&worker)).unwrap(); + registry.remove_by_url("grpc://w1:9000"); + monitor.apply_pushed_load(&worker, 0, &EngineLoad::default(), Instant::now()); + monitor.flush_pushed(); + assert!(monitor + .load_state + .snapshot() + .get("grpc://w1:9000") + .is_none()); + } + #[tokio::test] async fn bootstrap_reconcile_starts_loops_for_existing_workers() { let (registry, monitor) = build_monitor(); @@ -1524,6 +1986,206 @@ mod worker_monitor_tests { "expected at least one Weak from the event task and one from the group loop" ); } + + fn pushed(running: u32) -> EngineLoad { + EngineLoad { + running_requests: running, + waiting_requests: 1, + token_usage: 0.2, + gen_throughput: 10.0, + max_running_requests: 32, + sample: 1, + ..Default::default() + } + } + + /// A record's telemetry feeds the `smg_engine_*` gauges, the PD gauges + /// included, as a poll's does; a core-only record after it (an event + /// batch between heartbeats) keeps the telemetry the worker last sent + /// and updates the core, so suppression blanks no gauge and `GET /loads` + /// shows the engine's sections between heartbeats. + #[tokio::test] + async fn a_pushed_records_telemetry_feeds_the_gauges_and_outlives_core_only_records() { + let (registry, monitor) = build_monitor(); + let worker = ready_worker("grpc://w1:9000", "llama-3"); + registry.register(Arc::clone(&worker)).unwrap(); + let full = EngineLoad { + cache_hit_rate: Some(0.75), + num_used_tokens: Some(2_048), + max_total_num_tokens: Some(8_192), + memory: Some(smg_grpc_client::common_proto::EngineMemory { + weight_gb: 15.0, + kv_cache_gb: 40.0, + graph_gb: 1.5, + token_capacity: 8_192, + }), + disaggregation: Some(smg_grpc_client::common_proto::EngineDisaggregation { + mode: "prefill".to_string(), + prefill_prealloc_queue_reqs: 4, + prefill_inflight_queue_reqs: 5, + kv_transfer_latency_ms: 3.5, + ..Default::default() + }), + ..pushed(7) + }; + let rendered = render_with_recorder(|| { + monitor.apply_pushed_load(&worker, 0, &full, Instant::now()); + }); + let gauge = |name: &str| { + rendered + .lines() + .find(|l| { + l.starts_with(&format!("{name}{{")) && l.contains("worker=\"grpc://w1:9000\"") + }) + .map(|l| l.rsplit(' ').next().unwrap_or("").to_string()) + .unwrap_or_else(|| panic!("{name} missing:\n{rendered}")) + }; + assert_eq!(gauge("smg_engine_cache_hit_rate"), "0.75"); + assert_eq!(gauge("smg_engine_pd_kv_transfer_latency_ms"), "3.5"); + assert_eq!(gauge("smg_engine_pd_prefill_queue_reqs"), "9"); + + monitor.apply_pushed_load(&worker, 0, &pushed(8), Instant::now()); + monitor.flush_pushed(); + let snapshot = monitor.load_state.snapshot(); + let rank0 = &snapshot.get("grpc://w1:9000").expect("published").loads[0]; + assert_eq!(rank0.num_running_reqs, 8); + assert!((rank0.cache_hit_rate - 0.75).abs() < f64::EPSILON); + assert_eq!(rank0.max_total_num_tokens, 8_192); + assert_eq!(rank0.memory.as_ref().map(|m| m.token_capacity), Some(8_192)); + assert_eq!(rank0.disagg_mode.as_deref(), Some("prefill")); + assert_eq!(rank0.prefill_queue_reqs, Some(9)); + } + + /// The poll decision reads the receipt of the worker's pushed records + /// against the tick interval: none on file polls as always, a record + /// within the interval on every pushing rank suppresses the poll, and a + /// rank whose newest record is older falls back to the poll; the + /// record's own age (its `sampled_at`) plays no part. + #[tokio::test] + async fn the_poll_decision_follows_the_receipt_of_pushed_records() { + let (registry, monitor) = build_monitor(); + let worker = ready_worker("grpc://w1:9000", "llama-3"); + registry.register(Arc::clone(&worker)).unwrap(); + let url = "grpc://w1:9000"; + let interval = Duration::from_secs(10); + let t0 = Instant::now(); + // An old servicer: no record, ever. + assert_eq!(monitor.poll_mode(url, t0, interval), PollMode::Poll); + // Heartbeats flowing: no poll while the newest record is younger + // than the interval. + monitor.apply_pushed_load(&worker, 0, &pushed(3), t0); + assert_eq!( + monitor.poll_mode(url, t0 + interval / 2, interval), + PollMode::Suppressed + ); + // Heartbeats stopped: the poll is back within one interval. + assert_eq!( + monitor.poll_mode(url, t0 + interval, interval), + PollMode::Fallback + ); + // A second rank pushes; the oldest pushing rank decides. + monitor.apply_pushed_load(&worker, 1, &pushed(4), t0 + interval); + assert_eq!( + monitor.poll_mode(url, t0 + interval + interval / 2, interval), + PollMode::Fallback + ); + monitor.apply_pushed_load(&worker, 0, &pushed(5), t0 + interval + interval / 2); + assert_eq!( + monitor.poll_mode(url, t0 + interval + interval / 2, interval), + PollMode::Suppressed + ); + // A stale sample in a fresh record is still a fresh record. + let late = t0 + 2 * interval; + monitor.apply_pushed_load( + &worker, + 0, + &EngineLoad { + age_ms: 60_000, + ..pushed(6) + }, + late, + ); + monitor.apply_pushed_load(&worker, 1, &pushed(6), late); + assert_eq!( + monitor.poll_mode(url, late + Duration::from_millis(1), interval), + PollMode::Suppressed + ); + // A worker that leaves takes its records with it. + monitor.evict_worker_loads(&worker); + assert_eq!(monitor.poll_mode(url, late, interval), PollMode::Poll); + } + + /// A tick does not poll a worker whose stream pushed within the + /// interval, leaves the pushed report in the shared snapshot, and + /// counts the skip; the record fed the `smg_engine_*` gauges on its + /// way in, so suppression blanks nothing. Once the records are older + /// than the interval the tick polls again (here a URL nobody serves: + /// the failed poll prunes the entry, as a failed poll always has). + #[test] + fn a_tick_skips_a_worker_whose_stream_pushed_within_the_interval() { + let runtime = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .unwrap(); + let (registry, monitor) = build_monitor(); + let worker = ready_worker("grpc://w1:9000", "llama-3"); + registry.register(Arc::clone(&worker)).unwrap(); + let key = WorkerGroupKey { + model_id: "llama-3".to_string(), + worker_type: WorkerType::Regular, + connection_mode: ConnectionMode::Http, + }; + let interval = Duration::from_secs(10); + let t0 = Instant::now(); + let rendered = render_with_recorder(|| { + runtime.block_on(async { + monitor.apply_pushed_load(&worker, 0, &pushed(7), t0); + monitor.flush_pushed(); + poll_group_once(&monitor, &key, interval, t0 + interval / 2).await; + }); + }); + let snapshot = monitor.load_state.snapshot(); + let report = snapshot + .get("grpc://w1:9000") + .expect("kept across the tick"); + assert_eq!(report.version, "pushed"); + assert_eq!(report.loads[0].num_running_reqs, 7); + let line = |name: &str, label: &str| { + rendered + .lines() + .find(|l| l.starts_with(&format!("{name}{{")) && l.contains(label)) + .map(str::to_string) + .unwrap_or_else(|| panic!("{name} with {label} missing:\n{rendered}")) + }; + assert!( + line("smg_engine_running_requests", "worker=\"grpc://w1:9000\"").ends_with(" 7"), + "{rendered}" + ); + assert!( + line("smg_engine_load_polls_total", "mode=\"skipped_fresh_push\"").ends_with(" 1"), + "{rendered}" + ); + + let rendered = render_with_recorder(|| { + runtime.block_on(poll_group_once(&monitor, &key, interval, t0 + 2 * interval)); + }); + assert!( + monitor + .load_state + .snapshot() + .get("grpc://w1:9000") + .is_none(), + "the fallback poll of an unreachable worker prunes its entry" + ); + assert!( + rendered + .lines() + .any(|l| l.starts_with("smg_engine_load_polls_total{") + && l.contains("mode=\"fallback\"") + && l.ends_with(" 1")), + "{rendered}" + ); + } } #[cfg(test)] @@ -1727,7 +2389,8 @@ mod native_loads_tests { use super::*; use crate::{ config::types::PolicyConfig, - worker::{BasicWorkerBuilder, ConnectionMode, WorkerType}, + policies::{LoadBalancingPolicy, SelectWorkerInfo}, + worker::{BasicWorkerBuilder, ConnectionMode, WorkerLoadGuard, WorkerType}, }; const VLLM_METRICS: &str = "vllm:num_requests_running{m=\"a\"} 7.0\n\ @@ -1860,6 +2523,8 @@ mod native_loads_tests { cache_index: Default::default(), cache_ttl_secs: 180, cache_boundaries: Vec::new(), + selection_policy: None, + selection_accounting_ttl_ms: 0, } } @@ -1895,6 +2560,71 @@ mod native_loads_tests { .unwrap_or_else(|_| panic!("timed out waiting for {label}")); } + /// A policy that records the reconciliation ticks it receives. + #[derive(Debug, Default)] + struct ReconcileRecorder(Mutex>); + + impl LoadBalancingPolicy for ReconcileRecorder { + fn select_worker( + &self, + workers: &[Arc], + _info: &SelectWorkerInfo, + ) -> Option { + (!workers.is_empty()).then_some(0) + } + + fn reconcile_in_flight(&self, worker_url: &str, in_flight: usize) { + self.0.lock().push((worker_url.to_string(), in_flight)); + } + + fn name(&self) -> &'static str { + "reconcile_recorder" + } + + fn as_any(&self) -> &dyn std::any::Any { + self + } + } + + /// Every poll hands each polled worker's policy the router's live + /// in-flight count on it, the safety net for a completion that never + /// reached the policy. + #[tokio::test] + async fn load_poll_reconciles_in_flight_with_the_workers_policy() { + let stub = spawn_engine(StatusCode::OK, NATIVE_BODY).await; + let (registry, monitor) = monitor_with(PolicyConfig::RoundRobin, false); + let recorder = Arc::new(ReconcileRecorder::default()); + monitor + .policy_registry + .set_decode_policy(Arc::clone(&recorder) as Arc); + let worker: Arc = Arc::new( + BasicWorkerBuilder::new(stub.url.as_str()) + .worker_type(WorkerType::Decode) + .connection_mode(ConnectionMode::Http) + .runtime_type(RuntimeType::Vllm) + .model(ModelCard::new("a")) + .health_config(HealthCheckConfig { + disable_health_check: true, + ..Default::default() + }) + .build(), + ); + worker.set_status(WorkerStatus::Ready); + let held = WorkerLoadGuard::new(Arc::clone(&worker), None); + registry.register(Arc::clone(&worker)).unwrap(); + monitor.start_event_loop(); + + wait_until("the poll to reconcile the worker's policy", || { + recorder + .0 + .lock() + .iter() + .any(|(url, in_flight)| *url == stub.url && *in_flight == 1) + }) + .await; + drop(held); + } + /// `NATIVE_BODY` reports 4 waiting requests, so a threshold of 4 vetoes and /// a threshold of 5 does not — the ingestion path must latch the verdict on /// the worker itself, under a policy (round-robin) that reads no loads at diff --git a/model_gateway/src/worker/overload.rs b/model_gateway/src/worker/overload.rs index 88c22f5245..466e036def 100644 --- a/model_gateway/src/worker/overload.rs +++ b/model_gateway/src/worker/overload.rs @@ -2,16 +2,27 @@ //! //! The predicate is evaluated once per ingested load report, never per request; //! the verdict is latched into routing state, which selection already reads. +//! +//! On by default, steering only: a worker over either threshold is left out +//! of selection while another worker is under them, and when every worker is +//! over them the request goes to the least-loaded of them (see +//! `routers::common::overload`). Refusing requests with a 503 instead is the +//! opt-in `--worker-overload-shed`. use openai_protocol::worker::{OverloadUpdate, WorkerLoadResponse}; use crate::config::types::RouterConfig; -/// Token-usage ceiling applied when protection is enabled by -/// `--worker-overload-protection` alone. KV token usage means the same thing -/// on every engine; a waiting-requests default would be workload-dependent, so -/// that signal stays unset unless configured. -pub const DEFAULT_TOKEN_USAGE_CEILING: f64 = 0.9; +/// Default KV token-usage ceiling (mean across DP ranks). On the replay +/// harness this is the setting that keeps `cache_aware`'s hit rate while +/// cutting its tail; the hardware fleet sits near 0.9 under load, so a +/// higher ceiling never fires there. +pub const DEFAULT_TOKEN_USAGE_CEILING: f64 = 0.8; + +/// Default waiting-requests ceiling (summed across DP ranks): the signal that +/// trips first on the hardware fleet, where the KV pools fill slowly and the +/// queue forms in front of the engine. +pub const DEFAULT_WAITING_REQUESTS: usize = 8; /// Decision-log branch: every worker selection could have used is overloaded, /// so the request is shed immediately instead of queued. @@ -22,6 +33,17 @@ pub const BRANCH_ALL_OVERLOADED_SHED: &str = "all_overloaded_shed"; /// here every other worker may well be idle. pub const BRANCH_OVERLOADED_AT_DISPATCH: &str = "overloaded_at_dispatch"; +/// Decision-log branch: every worker selection could have used is over the +/// thresholds and shedding is off, so the request goes to the least-loaded +/// of them instead of being refused. +pub const BRANCH_ALL_OVERLOADED_FALLBACK: &str = "all_overloaded_fallback"; + +/// Decision-log branch: every worker selection could have used is vetoed and +/// the only ones left ready are vetoed by the liveness tracker, so the request +/// goes to the least-loaded of them instead of being refused: a veto steers, +/// it never empties the pool. +pub const BRANCH_ALL_STALLED_FALLBACK: &str = "all_stalled_fallback"; + /// Decision-log branch: the decode leg of a disaggregated pair was already /// running its full engine window and no slot freed inside the admission wait. /// The pair is not overloaded by the threshold predicate above — the gateway @@ -59,19 +81,18 @@ impl OverloadThresholds { self.waiting_requests.is_some() || self.token_usage.is_some() } - /// Gateway-level thresholds: the explicit `--worker-overload-*` values, - /// with the token ceiling defaulted to [`DEFAULT_TOKEN_USAGE_CEILING`] - /// when `--worker-overload-protection` enables the feature without one. - /// Everything unset with the flag off disables protection — exact #2220 - /// behavior. + /// Gateway-level thresholds: the `--worker-overload-*` values, which + /// default to [`DEFAULT_WAITING_REQUESTS`] and + /// [`DEFAULT_TOKEN_USAGE_CEILING`]. `--disable-worker-overload-protection` + /// (`worker_overload_protection: false`) switches both off whatever the + /// thresholds say; per-worker `overload` blocks still apply on top. pub fn from_gateway_config(config: &RouterConfig) -> Self { + if !config.worker_overload_protection { + return Self::default(); + } Self { waiting_requests: config.worker_overload_waiting_requests, - token_usage: config.worker_overload_token_usage.or_else(|| { - config - .worker_overload_protection - .then_some(DEFAULT_TOKEN_USAGE_CEILING) - }), + token_usage: config.worker_overload_token_usage, } } @@ -175,41 +196,42 @@ mod tests { } } - /// The flag alone enables protection with the engine-universal token - /// ceiling and no waiting-requests threshold. + /// The gateway defaults enable protection on both signals. #[test] - fn protection_flag_alone_defaults_the_token_ceiling() { - let thresholds = OverloadThresholds::from_gateway_config(&gateway_config(true, None, None)); + fn protection_is_on_by_default_on_both_signals() { + let thresholds = OverloadThresholds::from_gateway_config(&RouterConfig::default()); assert_eq!( thresholds, OverloadThresholds { - waiting_requests: None, + waiting_requests: Some(DEFAULT_WAITING_REQUESTS), token_usage: Some(DEFAULT_TOKEN_USAGE_CEILING), } ); assert!(thresholds.is_enabled()); + assert!(!RouterConfig::default().worker_overload_shed); } - /// Explicit thresholds override the flag's default; without either, the - /// feature stays exactly off (#2220 back-compat). + /// Explicit thresholds replace the defaults; disabling protection switches + /// both signals off whatever the thresholds say. #[test] - fn explicit_thresholds_override_the_flag_default() { + fn explicit_thresholds_replace_the_defaults_and_disabling_wins() { let explicit = - OverloadThresholds::from_gateway_config(&gateway_config(true, Some(8), Some(0.5))); + OverloadThresholds::from_gateway_config(&gateway_config(true, Some(64), Some(0.5))); assert_eq!(explicit.token_usage, Some(0.5)); - assert_eq!(explicit.waiting_requests, Some(8)); + assert_eq!(explicit.waiting_requests, Some(64)); - let no_flag = - OverloadThresholds::from_gateway_config(&gateway_config(false, Some(8), None)); + let one_signal = + OverloadThresholds::from_gateway_config(&gateway_config(true, None, Some(0.5))); assert_eq!( - no_flag, + one_signal, OverloadThresholds { - waiting_requests: Some(8), - token_usage: None, + waiting_requests: None, + token_usage: Some(0.5), } ); - let off = OverloadThresholds::from_gateway_config(&gateway_config(false, None, None)); + let off = + OverloadThresholds::from_gateway_config(&gateway_config(false, Some(8), Some(0.8))); assert!(!off.is_enabled()); } diff --git a/model_gateway/src/worker/prefill_admission.rs b/model_gateway/src/worker/prefill_admission.rs index 0b5b362bb3..852a962dde 100644 --- a/model_gateway/src/worker/prefill_admission.rs +++ b/model_gateway/src/worker/prefill_admission.rs @@ -1186,4 +1186,37 @@ mod tests { )); drop(occupied); } + + #[tokio::test] + async fn an_admitted_prefill_reports_completion_when_its_reservation_drops() { + use crate::worker::{BasicWorkerBuilder, RequestCompletionSink, Worker, WorkerType}; + #[derive(Debug, Default)] + struct CompletionSpy(std::sync::Mutex>); + impl RequestCompletionSink for CompletionSpy { + fn request_completed(&self, worker: &dyn Worker) { + self.0.lock().unwrap().push(worker.url().to_string()); + } + } + + let spy = Arc::new(CompletionSpy::default()); + let worker: Arc = Arc::new( + BasicWorkerBuilder::new("grpc://prefill-admitted") + .worker_type(WorkerType::Prefill) + .build(), + ); + worker.set_completion_sink(Some(spy.clone() as Arc)); + let admission = PrefillAdmission::new(1, 0, Duration::from_secs(1)); + let admitted = admission + .admit(None, |capacity| capacity.select(Arc::clone(&worker), ())) + .await + .expect("admitted"); + assert_eq!(worker.load(), 1); + assert!(spy.0.lock().unwrap().is_empty()); + drop(admitted); + assert_eq!(worker.load(), 0); + assert_eq!( + spy.0.lock().unwrap().as_slice(), + ["grpc://prefill-admitted"] + ); + } } diff --git a/model_gateway/src/worker/registry.rs b/model_gateway/src/worker/registry.rs index 018d13a42c..6fcf7994cc 100644 --- a/model_gateway/src/worker/registry.rs +++ b/model_gateway/src/worker/registry.rs @@ -18,7 +18,7 @@ use std::{ collections::{BTreeSet, HashSet}, ops::Deref, sync::{ - atomic::{AtomicUsize, Ordering}, + atomic::{AtomicBool, AtomicUsize, Ordering}, Arc, OnceLock, }, }; @@ -339,6 +339,11 @@ pub struct WorkerRegistry { /// reading it costs one map probe and no worker walk. model_overloaded: Arc>, + /// `--worker-overload-shed`: refuse a request with a 503 when every + /// candidate is overloaded instead of steering it to the least-loaded + /// one. Read by every router at its overload decision points. + overload_shed: Arc, + /// Serializes overload *edges* so a flag flip and its counter adjustment /// land as one step. Two group loops flipping the same worker in opposite /// directions would otherwise be free to apply their deltas in the reverse @@ -407,6 +412,7 @@ impl WorkerRegistry { url_to_id: Arc::new(DashMap::new()), worker_mutation_locks: Arc::new(DashMap::new()), model_overloaded: Arc::new(DashMap::new()), + overload_shed: Arc::new(AtomicBool::new(false)), overload_transitions: Arc::new(parking_lot::Mutex::new(())), model_retry_configs: Arc::new(DashMap::new()), worker_origins: Arc::new(DashMap::new()), @@ -689,6 +695,17 @@ impl WorkerRegistry { Arc::clone(&self.current_global_routing_snapshot().all) } + /// Whether an all-overloaded candidate pool is refused (`true`) or + /// steered to its least-loaded worker (the default). + pub fn overload_shed_enabled(&self) -> bool { + self.overload_shed.load(Ordering::Relaxed) + } + + /// Set `--worker-overload-shed` for every router reading this registry. + pub fn set_overload_shed(&self, shed: bool) { + self.overload_shed.store(shed, Ordering::Relaxed); + } + /// Apply the absolute overload veto to `worker`, returning `true` when the /// flag transitioned. The only sanctioned writer of /// [`Worker::set_overloaded`]: counters and gauge move once per edge. @@ -2351,6 +2368,16 @@ impl WorkerRegistry { } worker.set_status(new_status); + if new_status == WorkerStatus::Ready { + // The warm-up slice (cache_aware) counts from here. + worker.note_admitted(); + // A worker promoted back from a demotion returns with a closed + // circuit breaker: the failures that opened it belong to the + // outage its probes just ended. + if matches!(old_status, WorkerStatus::NotReady | WorkerStatus::Failed) { + worker.reset_circuit_breaker(); + } + } let _ = self.event_tx.send(WorkerEvent::StatusChanged { worker_id: worker_id.clone(), @@ -3776,6 +3803,35 @@ mod tests { assert_eq!(current.revision(), stale_revision + 1); } + #[test] + fn test_promotion_back_to_ready_closes_the_circuit_breaker() { + let registry = WorkerRegistry::new(); + let worker: Arc = Arc::new( + BasicWorkerBuilder::new("http://w1:8080") + .worker_type(WorkerType::Regular) + .health_config(no_health_check()) + .circuit_breaker_config(CircuitBreakerConfig::default()) + .build(), + ); + let worker_id = registry.register(worker.clone()).unwrap(); + let revision = worker.revision(); + assert!(registry + .transition_status_if_revision(&worker_id, revision, WorkerStatus::NotReady) + .is_some()); + for _ in 0..8 { + worker.record_circuit_breaker_outcome(false); + } + assert!(!worker.circuit_breaker_can_execute()); + + assert!(registry + .transition_status_if_revision(&worker_id, revision, WorkerStatus::Ready) + .is_some()); + assert!( + worker.circuit_breaker_can_execute(), + "a worker that returns starts with a closed breaker" + ); + } + #[test] fn test_multi_model_worker_is_indexed_for_each_model() { let registry = WorkerRegistry::new(); diff --git a/model_gateway/src/worker/worker.rs b/model_gateway/src/worker/worker.rs index 3ebad3046a..88c95c7552 100644 --- a/model_gateway/src/worker/worker.rs +++ b/model_gateway/src/worker/worker.rs @@ -2,7 +2,7 @@ use std::{ any::Any, fmt, sync::{ - atomic::{AtomicBool, AtomicU64, AtomicU8, AtomicUsize, Ordering}, + atomic::{AtomicBool, AtomicI64, AtomicU64, AtomicU8, AtomicUsize, Ordering}, Arc, OnceLock, }, time::Duration, @@ -20,7 +20,7 @@ use openai_protocol::{ }; use smg_grpc_client::common_proto; use tokio::{ - sync::{mpsc, OnceCell}, + sync::{mpsc, Notify, OnceCell}, task::AbortHandle, time, }; @@ -412,6 +412,15 @@ pub trait Worker: Send + Sync + fmt::Debug + 'static { /// Decrement the load counter fn decrement_load(&self); + /// The request-completion observer installed on this worker, if any + /// (see [`RequestCompletionSink`]). + fn completion_sink(&self) -> Option> { + None + } + + /// Install (or clear) the request-completion observer. + fn set_completion_sink(&self, _sink: Option>) {} + /// Claim `count` PD bootstrap rooms while the worker's claimed total /// stays within `window`; a refusal claims nothing. /// @@ -517,14 +526,17 @@ pub trait Worker: Send + Sync + fmt::Debug + 'static { /// Check if the worker is available (healthy + circuit closed/half-open + /// not vetoed by the absolute overload guard). fn is_available(&self) -> bool { - self.is_healthy() && self.circuit_breaker_can_execute() && !self.is_overloaded() + self.is_healthy() + && self.circuit_breaker_can_execute() + && !self.is_overloaded() + && self.stall_reason().is_none() } /// [`Self::is_healthy`] fused with the overload veto. For the hash policies, /// which route on health alone and never consult the circuit breaker; /// `BasicWorker` overrides it to read both under a single runtime guard. fn is_healthy_and_eligible(&self) -> bool { - self.is_healthy() && !self.is_overloaded() + self.is_healthy() && !self.is_overloaded() && self.stall_reason().is_none() } /// Whether the absolute overload guard currently vetoes this worker. @@ -544,6 +556,131 @@ pub trait Worker: Send + Sync + fmt::Debug + 'static { false } + /// Why the liveness tracker vetoes this worker, if it does. + fn stall_reason(&self) -> Option { + None + } + + /// Set or clear the liveness veto, returning `true` when it changed. + /// Route writes through [`super::liveness`], which logs and counts the + /// transition. + fn set_stall(&self, _reason: Option) -> bool { + false + } + + /// Record a successful interaction with the worker (a poll answered, a + /// probe passed, an event batch, a response). + fn note_contact(&self) {} + + /// Record a token or a completion from the worker. + fn note_token_progress(&self) {} + + /// Record a transport failure (a poll, probe or stream that failed on the + /// connection); cleared by the next contact. + fn note_transport_failure(&self) {} + + /// Whether a transport failure happened since the last contact. + fn transport_failure_pending(&self) -> bool { + false + } + + /// Time since the last successful contact. + fn contact_age(&self) -> Duration { + Duration::ZERO + } + + /// Time since the last token or completion. + fn token_progress_age(&self) -> Duration { + Duration::ZERO + } + /// Record the start of a request whose responses the gateway sees one by + /// one (a streaming generation to this worker over gRPC): the pile the + /// wedged rule counts, and the start of its clock when a run begins. + /// Requests without that signal (HTTP, a PD leg, a non-streaming + /// generation) never form a pile. + fn note_tracked_started(&self) {} + /// Record the end of a tracked request. + fn note_tracked_ended(&self) {} + /// Requests in flight whose responses the gateway sees one by one. + fn tracked_load(&self) -> usize { + 0 + } + + /// Store the engine's reported waiting-queue depth, returning the previous + /// value. + fn swap_waiting_reqs(&self, _waiting: i64) -> i64 { + 0 + } + + /// Store the in-flight count seen by the liveness sweep, returning the + /// previous sample. + fn swap_load_sample(&self, _load: usize) -> usize { + 0 + } + + /// A notifier fired on every contact with the worker, for loops that back + /// off from it and should retry as soon as it is heard from again. + fn contact_wake(&self) -> Option> { + None + } + + /// Ask the health manager to promote the worker now, on the strength of a + /// successful contact, instead of at its next scheduled probe. + fn signal_connected(&self) {} + + /// Close the circuit breaker: the worker has returned (a liveness veto + /// cleared by a contact, or health promoted it back to Ready) and the + /// failures that opened it belong to the outage. It reopens on fresh + /// failures like any closed breaker. + fn reset_circuit_breaker(&self) {} + + /// Record that the worker just became routable (promotion to Ready, or a + /// liveness veto cleared); the warm-up slice counts from here. + fn note_admitted(&self) {} + + /// Time since the worker last became routable; `Duration::MAX` for a + /// worker that does not track it (never warming). + fn admitted_age(&self) -> Duration { + Duration::MAX + } + + /// Blocks the worker's index has gained since its current admission, for + /// the warm-up slice; `indexed` is what the index holds for it now. A + /// worker that does not track admissions reports the count itself. + fn warmup_growth(&self, indexed: usize) -> usize { + indexed + } + + /// Gateway clock (`liveness::now_ms`) before which a thin worker with + /// requests in flight takes no further diverted hit (see + /// `CacheAwarePolicy::warmup_divert`); zero for a worker that does not + /// track it. + fn divert_until_ms(&self) -> u64 { + 0 + } + + /// A hit was diverted to this thin worker; the next one waits until + /// `until_ms` while it has anything in flight. + fn note_diverted(&self, _until_ms: u64) {} + + /// A request of `tokens` prompt tokens was dispatched: work the engine + /// still has to prefill before its first token (see + /// [`Self::prefill_backlog`]). + fn note_prefill_started(&self, _tokens: u64) {} + + /// The request's first token came back (`prefilled`), or its stream ended + /// or was dropped before one (`!prefilled`): its prompt is no longer + /// pending prefill. + fn note_prefill_ended(&self, _tokens: u64, _prefilled: bool) {} + + /// How long the engine may still need before the first token of its + /// in-flight work is due: the prompt tokens dispatched and not yet + /// answered, over the prefill rate observed on this worker (a cold prior + /// until one is). Zero for a worker that does not track it. + fn prefill_backlog(&self) -> Duration { + Duration::ZERO + } + /// One-shot routing snapshot for the per-request O(workers) selection loops: /// reads status, load, processed and the overload veto together so the hot /// path takes one `ArcSwap` guard per backing cell per worker instead of one @@ -556,6 +693,7 @@ pub trait Worker: Send + Sync + fmt::Debug + 'static { load: self.load(), processed: self.processed_requests(), overloaded: self.is_overloaded(), + stalled: self.stall_reason().is_some(), } } @@ -1137,6 +1275,35 @@ impl WorkerMetadata { } /// One-shot routing snapshot — see [`Worker::routing_state`]. +/// Why the liveness tracker vetoes a worker (see [`super::liveness`]). +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +#[repr(u8)] +pub enum StallReason { + /// Its transport fails and nothing has been heard from it for the stall + /// threshold. + Unreachable = 1, + /// It answers polls but holds requests that make no progress while its + /// queue grows. + Wedged = 2, +} + +impl StallReason { + pub const fn as_str(self) -> &'static str { + match self { + Self::Unreachable => "unreachable", + Self::Wedged => "wedged", + } + } + + const fn from_u8(value: u8) -> Option { + match value { + 1 => Some(Self::Unreachable), + 2 => Some(Self::Wedged), + _ => None, + } + } +} + #[derive(Clone, Copy, Debug)] pub struct RoutingState { /// `status == Ready`. @@ -1149,6 +1316,8 @@ pub struct RoutingState { pub processed: usize, /// Absolute overload veto, set by the load monitor at ingestion time. pub overloaded: bool, + /// Liveness veto: unreachable or wedged (see [`super::liveness`]). + pub stalled: bool, } impl RoutingState { @@ -1156,7 +1325,7 @@ impl RoutingState { /// gather pass already performed: every field rides the one guard /// [`Worker::routing_state`] took. pub const fn eligible(self) -> bool { - self.healthy && self.can_execute && !self.overloaded + self.healthy && self.can_execute && !self.overloaded && !self.stalled } } @@ -1182,8 +1351,59 @@ pub struct WorkerRuntime { /// so selection reads it under the guard it already holds, and so a /// same-URL replacement inherits it with the rest of the shared runtime. overloaded: AtomicBool, + /// Liveness veto, a [`StallReason`] as `u8` (0 = none). + stall: AtomicU8, + /// Last successful contact and last token or completion, in + /// [`super::liveness::now_ms`] milliseconds. + last_contact_ms: AtomicU64, + last_token_ms: AtomicU64, + /// Waiting-queue depth from the previous load report. + last_waiting_reqs: AtomicI64, + /// A transport failure happened since the last contact. + transport_failed: AtomicBool, + /// In-flight count at the previous liveness sweep. + last_load_sample: AtomicUsize, + /// Woken on every contact, so a loop backing off from this worker (the + /// KV event subscriber) retries the moment the worker is heard from. + contact_wake: Arc, + /// When the worker last became routable (promoted to Ready, or a liveness + /// veto cleared), in [`super::liveness::now_ms`] milliseconds. + admitted_at_ms: AtomicU64, + /// The warm-up slice's baseline: the admission it was taken for and the + /// blocks the index held for the worker then (see [`Self::warmup_growth`]). + warmup_base_admitted_ms: AtomicU64, + warmup_base_blocks: AtomicUsize, + /// Clock before which this thin worker, with requests in flight, takes + /// no further diverted hit. + divert_until_ms: AtomicU64, + /// When the current run of in-flight requests began (the load counter + /// left zero), in [`super::liveness::now_ms`] milliseconds; zero while + /// idle. The no-progress clock of the wedged rule starts here, not at + /// registration or the last token before an idle spell. + busy_since_ms: AtomicU64, + /// Requests in flight whose responses the gateway sees one by one (see + /// [`Worker::tracked_load`]). + tracked_in_flight: AtomicUsize, + /// Prompt tokens dispatched whose first token has not come back. + prefill_tokens_pending: AtomicU64, + /// Observed aggregate prefill rate, tokens per second; zero until a + /// window of first tokens has been seen (see [`Self::observe_prefill_at`]). + prefill_rate_tps: AtomicU64, + prefill_window_start_ms: AtomicU64, + prefill_window_tokens: AtomicU64, } +/// Prefill rate assumed for a worker that has not shown one yet: slow enough +/// that a cold engine handed a burst of long prompts is not called wedged +/// before it can have answered (10k tokens/s is a small model on a modest +/// GPU; a current datacenter GPU prefills 8B weights at ~90k). +const COLD_PREFILL_TOKENS_PER_SEC: u64 = 10_000; + +/// First tokens are summed over windows at least this long before they make +/// a rate sample, so the sample is the engine's aggregate throughput and not +/// one request's time to first token (which includes its wait in the batch). +const PREFILL_WINDOW_MS: u64 = 1_000; + impl WorkerRuntime { pub fn new(url: &str, initial_status: WorkerStatus) -> Self { Self { @@ -1197,9 +1417,207 @@ impl WorkerRuntime { worker_routing_key_load: WorkerRoutingKeyLoad::new(url), revision: AtomicU64::new(0), overloaded: AtomicBool::new(false), + stall: AtomicU8::new(0), + last_contact_ms: AtomicU64::new(super::liveness::now_ms()), + last_token_ms: AtomicU64::new(super::liveness::now_ms()), + last_waiting_reqs: AtomicI64::new(0), + transport_failed: AtomicBool::new(false), + last_load_sample: AtomicUsize::new(0), + contact_wake: Arc::new(Notify::new()), + admitted_at_ms: AtomicU64::new(super::liveness::now_ms()), + warmup_base_admitted_ms: AtomicU64::new(u64::MAX), + warmup_base_blocks: AtomicUsize::new(0), + divert_until_ms: AtomicU64::new(0), + busy_since_ms: AtomicU64::new(0), + tracked_in_flight: AtomicUsize::new(0), + prefill_tokens_pending: AtomicU64::new(0), + prefill_rate_tps: AtomicU64::new(0), + prefill_window_start_ms: AtomicU64::new(0), + prefill_window_tokens: AtomicU64::new(0), } } + pub fn note_admitted(&self) { + self.admitted_at_ms + .store(super::liveness::now_ms(), Ordering::Relaxed); + } + + pub fn admitted_age(&self) -> Duration { + Self::age_of(self.admitted_at_ms.load(Ordering::Relaxed)) + } + + pub fn divert_until_ms(&self) -> u64 { + self.divert_until_ms.load(Ordering::Relaxed) + } + + pub fn note_diverted(&self, until_ms: u64) { + self.divert_until_ms.store(until_ms, Ordering::Relaxed); + } + + /// Blocks the index has gained for the worker since its current admission. + /// The index size itself is no measure of a returned worker's cache: a + /// restarted engine's stale blocks stay indexed until its first batch + /// reveals the restart, and an engine that gets no request sends no batch, + /// so the stale count would end the warm-up meant to break that circle. + /// The baseline is the count first seen for the current admission, and + /// drops to zero when the index was cleared underneath (the count went + /// down). Two racing callers may both take the same baseline; nothing + /// worse. + pub fn warmup_growth(&self, indexed: usize) -> usize { + let admitted = self.admitted_at_ms.load(Ordering::Relaxed); + if self.warmup_base_admitted_ms.load(Ordering::Relaxed) != admitted { + self.warmup_base_admitted_ms + .store(admitted, Ordering::Relaxed); + self.warmup_base_blocks.store(indexed, Ordering::Relaxed); + return 0; + } + let base = self.warmup_base_blocks.load(Ordering::Relaxed); + if indexed < base { + self.warmup_base_blocks.store(0, Ordering::Relaxed); + return indexed; + } + indexed - base + } + + // ── Liveness ──────────────────────────────────────────────────── + + pub fn note_transport_failure(&self) { + self.transport_failed.store(true, Ordering::Relaxed); + } + + pub fn transport_failure_pending(&self) -> bool { + self.transport_failed.load(Ordering::Relaxed) + } + + pub fn stall_reason(&self) -> Option { + StallReason::from_u8(self.stall.load(Ordering::Acquire)) + } + + /// Set or clear the veto; `true` when it changed. + pub fn set_stall(&self, reason: Option) -> bool { + let next = reason.map_or(0, |reason| reason as u8); + self.stall.swap(next, Ordering::AcqRel) != next + } + + pub fn note_contact(&self) { + self.last_contact_ms + .store(super::liveness::now_ms(), Ordering::Relaxed); + self.transport_failed.store(false, Ordering::Relaxed); + self.contact_wake.notify_one(); + } + + pub fn note_token_progress(&self) { + let now = super::liveness::now_ms(); + self.last_token_ms.store(now, Ordering::Relaxed); + self.last_contact_ms.store(now, Ordering::Relaxed); + self.transport_failed.store(false, Ordering::Relaxed); + self.contact_wake.notify_one(); + } + + /// The notifier [`Self::note_contact`] fires; `notify_one` semantics, so + /// a contact that happens before anyone waits is not lost. + pub fn contact_wake(&self) -> Arc { + Arc::clone(&self.contact_wake) + } + + pub fn contact_age(&self) -> Duration { + Self::age_of(self.last_contact_ms.load(Ordering::Relaxed)) + } + + /// Time without a token or completion, counted from the later of the last + /// one and the start of the current run of in-flight requests: a worker + /// that was idle (or just registered) has nothing to show progress on + /// until something is dispatched to it. + pub fn token_progress_age(&self) -> Duration { + Self::age_of(self.progress_reference_ms()) + } + + /// The stamp [`Self::token_progress_age`] counts from. + pub fn progress_reference_ms(&self) -> u64 { + self.last_token_ms + .load(Ordering::Relaxed) + .max(self.busy_since_ms.load(Ordering::Relaxed)) + } + + // ── Prefill backlog (the wedged rule's bound) ─────────────────── + + pub fn note_prefill_started(&self, tokens: u64) { + self.prefill_tokens_pending + .fetch_add(tokens, Ordering::Relaxed); + } + + pub fn note_prefill_ended(&self, tokens: u64, prefilled: bool) { + let _ = self.prefill_tokens_pending.fetch_update( + Ordering::Relaxed, + Ordering::Relaxed, + |pending| Some(pending.saturating_sub(tokens)), + ); + if prefilled && tokens > 0 { + self.observe_prefill_at(tokens, super::liveness::now_ms().max(1)); + } + } + + /// Fold `tokens` prefilled at `now_ms` into the observed rate: first + /// tokens are summed over windows of at least [`PREFILL_WINDOW_MS`] and + /// each closed window halves into the running estimate. Racing callers + /// may lose a few tokens of a window; the estimate is a bound, not a book. + pub fn observe_prefill_at(&self, tokens: u64, now_ms: u64) { + let start = self.prefill_window_start_ms.load(Ordering::Relaxed); + if start == 0 { + self.prefill_window_start_ms + .store(now_ms, Ordering::Relaxed); + self.prefill_window_tokens.store(tokens, Ordering::Relaxed); + return; + } + let total = self + .prefill_window_tokens + .fetch_add(tokens, Ordering::Relaxed) + + tokens; + let span_ms = now_ms.saturating_sub(start); + if span_ms < PREFILL_WINDOW_MS { + return; + } + let sample = total.saturating_mul(1000) / span_ms; + let rate = self.prefill_rate_tps.load(Ordering::Relaxed); + let next = if rate == 0 { + sample + } else { + rate.midpoint(sample) + }; + self.prefill_rate_tps.store(next.max(1), Ordering::Relaxed); + self.prefill_window_start_ms + .store(now_ms, Ordering::Relaxed); + self.prefill_window_tokens.store(0, Ordering::Relaxed); + } + + pub fn prefill_rate_tps(&self) -> u64 { + self.prefill_rate_tps.load(Ordering::Relaxed) + } + + pub fn prefill_backlog(&self) -> Duration { + let pending = self.prefill_tokens_pending.load(Ordering::Relaxed); + if pending == 0 { + return Duration::ZERO; + } + let rate = match self.prefill_rate_tps.load(Ordering::Relaxed) { + 0 => COLD_PREFILL_TOKENS_PER_SEC, + rate => rate, + }; + Duration::from_millis(pending.saturating_mul(1000) / rate) + } + + pub fn swap_waiting_reqs(&self, waiting: i64) -> i64 { + self.last_waiting_reqs.swap(waiting, Ordering::Relaxed) + } + + pub fn swap_load_sample(&self, load: usize) -> usize { + self.last_load_sample.swap(load, Ordering::Relaxed) + } + + fn age_of(stamp_ms: u64) -> Duration { + Duration::from_millis(super::liveness::now_ms().saturating_sub(stamp_ms)) + } + // ── Lifecycle status ──────────────────────────────────────────── pub fn status(&self) -> WorkerStatus { @@ -1255,25 +1673,66 @@ impl WorkerRuntime { } pub fn increment_load(&self) { - self.load_counter.fetch_add(1, Ordering::Relaxed); + if self.load_counter.fetch_add(1, Ordering::Relaxed) == 0 { + self.note_busy_started(); + } } pub fn try_increment_load(&self, max: usize) -> bool { - self.load_counter + match self + .load_counter .fetch_update(Ordering::Relaxed, Ordering::Relaxed, |current| { current.checked_add(1).filter(|next| *next <= max) - }) - .is_ok() + }) { + Ok(0) => { + self.note_busy_started(); + true + } + Ok(_) => true, + Err(_) => false, + } + } + + /// A tracked request started. The wedged rule's clock starts with the + /// first of a run, so an untracked request dispatched earlier (an HTTP or + /// a non-streaming one) cannot age the run before it begins. + pub fn note_tracked_started(&self) { + if self.tracked_in_flight.fetch_add(1, Ordering::Relaxed) == 0 { + self.note_busy_started(); + } + } + pub fn note_tracked_ended(&self) { + let _ = + self.tracked_in_flight + .fetch_update(Ordering::Relaxed, Ordering::Relaxed, |tracked| { + Some(tracked.saturating_sub(1)) + }); + } + pub fn tracked_load(&self) -> usize { + self.tracked_in_flight.load(Ordering::Relaxed) + } + /// The load counter left zero: the wedged rule's clock starts now. + fn note_busy_started(&self) { + self.busy_since_ms + .store(super::liveness::now_ms().max(1), Ordering::Relaxed); } /// Saturating decrement. Returns `true` if the counter was decremented, /// `false` if it was already zero — callers can log when that happens. pub fn try_decrement_load(&self) -> bool { - self.load_counter + match self + .load_counter .fetch_update(Ordering::Relaxed, Ordering::Relaxed, |current| { current.checked_sub(1) - }) - .is_ok() + }) { + Ok(1) => { + // Idle again: nothing in flight to wait for. + self.busy_since_ms.store(0, Ordering::Relaxed); + true + } + Ok(_) => true, + Err(_) => false, + } } // ── PD admission claims ───────────────────────────────────────── @@ -1340,6 +1799,16 @@ impl WorkerRuntime { } } +/// Observer of request completions on a worker. The request's load guard +/// notifies it when it drops, which is the one point every path that holds +/// worker load (HTTP, gRPC, both PD legs, admitted prefill) passes through on +/// success, error and client disconnect alike. The policy registry installs +/// one on each worker it learns about, so policies that book state at +/// dispatch (reservations, bookings) see the request end. +pub trait RequestCompletionSink: Send + Sync + fmt::Debug { + fn request_completed(&self, worker: &dyn Worker); +} + /// Basic worker implementation pub struct BasicWorker { pub metadata: WorkerMetadata, @@ -1381,12 +1850,16 @@ pub struct BasicWorker { pub http_client: Arc, /// Resolved resilience config (retry + circuit breaker settings). pub resilience: ResolvedResilience, + /// Request-completion observer, installed by the policy registry; shared + /// by clones so a DP-rank view reports to the same sink. + pub completion_sink: Arc>>>, } impl Clone for BasicWorker { fn clone(&self) -> Self { Self { metadata: self.metadata.clone(), + completion_sink: Arc::clone(&self.completion_sink), runtime: ArcSwap::from(self.runtime.load_full()), circuit_breaker: ArcSwap::from(self.circuit_breaker.load_full()), backend_client: Arc::clone(&self.backend_client), @@ -1749,18 +2222,134 @@ impl Worker for BasicWorker { self.runtime.load().set_overloaded(overloaded) } + fn stall_reason(&self) -> Option { + self.runtime.load().stall_reason() + } + + fn set_stall(&self, reason: Option) -> bool { + self.runtime.load().set_stall(reason) + } + + fn note_contact(&self) { + self.runtime.load().note_contact(); + } + + fn note_token_progress(&self) { + self.runtime.load().note_token_progress(); + } + + fn note_transport_failure(&self) { + self.runtime.load().note_transport_failure(); + } + + fn transport_failure_pending(&self) -> bool { + self.runtime.load().transport_failure_pending() + } + + fn contact_age(&self) -> Duration { + self.runtime.load().contact_age() + } + + fn token_progress_age(&self) -> Duration { + self.runtime.load().token_progress_age() + } + + fn swap_waiting_reqs(&self, waiting: i64) -> i64 { + self.runtime.load().swap_waiting_reqs(waiting) + } + + fn swap_load_sample(&self, load: usize) -> usize { + self.runtime.load().swap_load_sample(load) + } + + fn contact_wake(&self) -> Option> { + Some(self.runtime.load().contact_wake()) + } + + fn signal_connected(&self) { + if let Some(tx) = &self.connect_signal_tx { + let _ = tx.send(WorkerConnected { + url: self.url().to_string(), + revision: self.revision(), + }); + } + } + + fn reset_circuit_breaker(&self) { + self.circuit_breaker.load().reset(); + } + + fn note_admitted(&self) { + self.runtime.load().note_admitted(); + } + + fn admitted_age(&self) -> Duration { + self.runtime.load().admitted_age() + } + + fn warmup_growth(&self, indexed: usize) -> usize { + self.runtime.load().warmup_growth(indexed) + } + + fn divert_until_ms(&self) -> u64 { + self.runtime.load().divert_until_ms() + } + + fn note_diverted(&self, until_ms: u64) { + self.runtime.load().note_diverted(until_ms); + } + + fn note_prefill_started(&self, tokens: u64) { + self.runtime.load().note_prefill_started(tokens); + } + + fn note_prefill_ended(&self, tokens: u64, prefilled: bool) { + self.runtime.load().note_prefill_ended(tokens, prefilled); + } + + fn prefill_backlog(&self) -> Duration { + self.runtime.load().prefill_backlog() + } + + fn note_tracked_started(&self) { + self.runtime.load().note_tracked_started(); + } + + fn note_tracked_ended(&self) { + self.runtime.load().note_tracked_ended(); + } + + fn tracked_load(&self) -> usize { + self.runtime.load().tracked_load() + } + fn is_available(&self) -> bool { // Same two guards the pre-veto version took (`is_healthy` + - // `circuit_breaker_can_execute`): the veto rides the runtime guard. + // `circuit_breaker_can_execute`): the vetoes ride the runtime guard. let rt = self.runtime.load(); rt.status() == WorkerStatus::Ready && !rt.is_overloaded() + && rt.stall_reason().is_none() && self.circuit_breaker.load().can_execute() } fn is_healthy_and_eligible(&self) -> bool { let rt = self.runtime.load(); - rt.status() == WorkerStatus::Ready && !rt.is_overloaded() + rt.status() == WorkerStatus::Ready && !rt.is_overloaded() && rt.stall_reason().is_none() + } + + fn completion_sink(&self) -> Option> { + self.completion_sink + .read() + .unwrap_or_else(|poisoned| poisoned.into_inner()) + .clone() + } + + fn set_completion_sink(&self, sink: Option>) { + *self + .completion_sink + .write() + .unwrap_or_else(|poisoned| poisoned.into_inner()) = sink; } fn routing_state(&self) -> RoutingState { @@ -1774,6 +2363,7 @@ impl Worker for BasicWorker { load: rt.load(), processed: rt.processed_requests(), overloaded: rt.is_overloaded(), + stalled: rt.stall_reason().is_some(), } } @@ -2096,6 +2686,11 @@ impl Drop for WorkerLoadGuard { if let Some(ref key) = self.routing_key { self.worker.decrement_routing_key_load(key); } + // Every request path releases its worker load here, so this is where + // policies learn that the request ended, whatever the outcome. + if let Some(sink) = self.worker.completion_sink() { + sink.request_completed(self.worker.as_ref()); + } } } @@ -2173,6 +2768,9 @@ pub fn worker_to_info(worker: &Arc) -> WorkerInfo { load: worker.load(), http2: metadata.http2, pd_pairing, + stalled: worker + .stall_reason() + .map(|reason| reason.as_str().to_string()), engine_load: None, job_status: None, } @@ -3611,4 +4209,112 @@ mod tests { "handshake guard must reset so a later probe can retry" ); } + + /// Records which worker URLs reported a completion. + #[derive(Debug, Default)] + struct CompletionSpy(std::sync::Mutex>); + + impl RequestCompletionSink for CompletionSpy { + fn request_completed(&self, worker: &dyn Worker) { + self.0 + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()) + .push(worker.url().to_string()); + } + } + + #[test] + fn load_guard_drop_reports_the_completion_to_the_installed_sink() { + let spy = Arc::new(CompletionSpy::default()); + let worker: Arc = Arc::new( + BasicWorkerBuilder::new("http://w1:8000") + .worker_type(WorkerType::Regular) + .build(), + ); + worker.set_completion_sink(Some(Arc::clone(&spy) as Arc)); + + let guard = WorkerLoadGuard::with_key(Arc::clone(&worker), Some("session-a")); + let twin = guard.replicate(); + assert_eq!(worker.load(), 2); + assert!( + spy.0.lock().unwrap().is_empty(), + "nothing completes while guards live" + ); + drop(guard); + assert_eq!(worker.load(), 1); + assert_eq!(spy.0.lock().unwrap().as_slice(), ["http://w1:8000"]); + drop(twin); + assert_eq!(worker.load(), 0); + assert_eq!( + spy.0.lock().unwrap().len(), + 2, + "each held load reports once" + ); + + // Without a sink the guard is exactly what it was. + worker.set_completion_sink(None); + drop(WorkerLoadGuard::with_key(Arc::clone(&worker), None)); + assert_eq!(spy.0.lock().unwrap().len(), 2); + } + + #[test] + fn progress_is_counted_from_the_first_dispatch_not_registration() { + let runtime = WorkerRuntime::new("http://w1:8000", WorkerStatus::Ready); + let registered = runtime.progress_reference_ms(); + thread::sleep(Duration::from_millis(15)); + runtime.increment_load(); + assert!( + runtime.progress_reference_ms() > registered, + "the clock starts at the dispatch, not at registration" + ); + assert!(runtime.token_progress_age() < Duration::from_millis(10)); + assert!(runtime.try_decrement_load()); + assert_eq!( + runtime.progress_reference_ms(), + registered, + "idle again: back to the last token" + ); + assert!(runtime.try_increment_load(4)); + assert!(runtime.progress_reference_ms() > registered); + } + + #[test] + fn prefill_backlog_follows_pending_tokens_and_the_observed_rate() { + let runtime = WorkerRuntime::new("http://w1:8000", WorkerStatus::Ready); + assert_eq!(runtime.prefill_backlog(), Duration::ZERO); + // 128 prompts of 1,152 tokens on a worker that has shown no rate yet: + // the cold prior, 10k tokens/s, gives the engine 14.7 s. + runtime.note_prefill_started(128 * 1_152); + assert_eq!(runtime.prefill_backlog(), Duration::from_millis(14_745)); + // Their first tokens come back within 1.5 s: an observed rate of + // ~98k tokens/s replaces the prior, and nothing is pending. + runtime.observe_prefill_at(64 * 1_152, 1_000); + runtime.observe_prefill_at(64 * 1_152, 2_500); + runtime.note_prefill_ended(128 * 1_152, false); + assert_eq!(runtime.prefill_backlog(), Duration::ZERO); + assert_eq!(runtime.prefill_rate_tps(), 98_304); + runtime.note_prefill_started(1_000_000); + assert_eq!(runtime.prefill_backlog(), Duration::from_millis(10_172)); + // A dropped stream releases its tokens without a rate sample. + runtime.note_prefill_ended(2_000_000, false); + assert_eq!(runtime.prefill_backlog(), Duration::ZERO); + } + + #[test] + fn warm_up_growth_is_measured_since_admission_and_survives_a_clear() { + let runtime = WorkerRuntime::new("http://w1:8000", WorkerStatus::Ready); + // A returned worker is first seen with its stale blocks still indexed: + // those are the baseline, not growth. + assert_eq!(runtime.warmup_growth(2000), 0); + assert_eq!(runtime.warmup_growth(2100), 100); + // The restart is detected and the index cleared: growth restarts from + // zero, not from a baseline the index no longer holds. + assert_eq!(runtime.warmup_growth(50), 50); + assert_eq!(runtime.warmup_growth(700), 700); + // Admitted again: a fresh baseline. + thread::sleep(Duration::from_millis(2)); + runtime.note_admitted(); + assert_eq!(runtime.warmup_growth(900), 0); + assert_eq!(runtime.warmup_growth(1300), 400); + } } diff --git a/model_gateway/src/workflow/steps/local/create_worker.rs b/model_gateway/src/workflow/steps/local/create_worker.rs index e01a8b1735..19bc4f08f1 100644 --- a/model_gateway/src/workflow/steps/local/create_worker.rs +++ b/model_gateway/src/workflow/steps/local/create_worker.rs @@ -277,13 +277,13 @@ impl StepExecutor for CreateLocalWorkerStep { if let Some(group_size) = zmq_engine_group { builder = builder.zmq_engine_group(group_size); } - // ZMQ promotion is event-driven: the worker signals the manager - // the instant its handshake completes, so wire the registry's - // connect signal. Other transports promote via polling. - if *connection_mode == ConnectionMode::Zmq { - builder = builder - .connect_signal_tx(app_context.worker_registry.connect_signal_sender()); - } + // Promotion can be event-driven: a ZMQ worker signals the + // manager the instant its handshake completes, and any worker + // demoted by health is signalled on its first successful + // contact (see `worker::liveness`), so every worker gets the + // registry's connect signal. + builder = + builder.connect_signal_tx(app_context.worker_registry.connect_signal_sender()); // Builder sets initial status: Pending if health-checked, Ready if not. Arc::new(builder.build()) as Arc diff --git a/model_gateway/src/workflow/steps/local/update_policies_for_worker.rs b/model_gateway/src/workflow/steps/local/update_policies_for_worker.rs index 99beccc40d..df76ff90e1 100644 --- a/model_gateway/src/workflow/steps/local/update_policies_for_worker.rs +++ b/model_gateway/src/workflow/steps/local/update_policies_for_worker.rs @@ -61,6 +61,9 @@ impl StepExecutor for UpdatePoliciesForWorkerStep { // Notify policy registry of the update app_context.policy_registry.on_worker_added(model_id, None); + for worker in workers.iter() { + worker.set_completion_sink(Some(app_context.policy_registry.completion_sink())); + } } let prefill_workers = app_context.worker_registry.get_prefill_workers(); diff --git a/model_gateway/src/workflow/steps/local/update_worker_properties.rs b/model_gateway/src/workflow/steps/local/update_worker_properties.rs index 7deff04ea7..229ff393da 100644 --- a/model_gateway/src/workflow/steps/local/update_worker_properties.rs +++ b/model_gateway/src/workflow/steps/local/update_worker_properties.rs @@ -8,9 +8,7 @@ use tracing::{debug, info}; use wfaas::{StepExecutor, StepResult, WorkflowContext, WorkflowError, WorkflowResult}; use crate::{ - worker::{ - overload::OverloadThresholds, BasicWorker, BasicWorkerBuilder, ConnectionMode, Worker, - }, + worker::{overload::OverloadThresholds, BasicWorker, BasicWorkerBuilder, Worker}, workflow::data::WorkerUpdateWorkflowData, }; @@ -132,10 +130,8 @@ impl StepExecutor for UpdateWorkerPropertiesStep { // The spec retains the ZMQ address and engine count. The connect // signal still has to come from the registry. - if *worker.connection_mode() == ConnectionMode::Zmq { - builder = - builder.connect_signal_tx(app_context.worker_registry.connect_signal_sender()); - } + builder = + builder.connect_signal_tx(app_context.worker_registry.connect_signal_sender()); let mut new_worker = builder.build(); if let Some(previous) = worker.as_any().downcast_ref::() { @@ -208,7 +204,7 @@ mod tests { }, worker::{ circuit_breaker::{CircuitBreakerConfig, CircuitState}, - BasicWorker, RuntimeType, + BasicWorker, ConnectionMode, RuntimeType, }, }; diff --git a/model_gateway/src/workflow/steps/shared/update_policies.rs b/model_gateway/src/workflow/steps/shared/update_policies.rs index 4e8681db23..2f25b52a15 100644 --- a/model_gateway/src/workflow/steps/shared/update_policies.rs +++ b/model_gateway/src/workflow/steps/shared/update_policies.rs @@ -137,6 +137,9 @@ impl StepExecutor for UpdatePolicie app_context .policy_registry .on_worker_added(&model_id, policy_hint); + // Policies learn of the request's end through the worker's load + // guard; the registry routes it to the policy that placed it. + worker.set_completion_sink(Some(app_context.policy_registry.completion_sink())); // Initialize cache-aware policy if configured let all_workers = app_context.worker_registry.get_by_model(&model_id); diff --git a/model_gateway/tests/allocator_artifact_test.rs b/model_gateway/tests/allocator_artifact_test.rs index 7fa161e566..1eb3d08db7 100644 --- a/model_gateway/tests/allocator_artifact_test.rs +++ b/model_gateway/tests/allocator_artifact_test.rs @@ -1,4 +1,5 @@ -//! Verifies that the shipped SMG executable owns the expected jemalloc. +//! Verifies that the shipped SMG executable owns the expected jemalloc and +//! runs it with the gateway's long-running-server options. #![cfg(all( feature = "jemalloc-stats", @@ -8,20 +9,35 @@ use std::process::Command; -#[test] -fn smg_binary_uses_global_jemalloc() { +/// Runs the executable with jemalloc's JSON statistics printed at exit under +/// the given `_RJEM_MALLOC_CONF` and returns the statistics text (stderr). +fn jemalloc_stats(malloc_conf: &str) -> Result { let output = Command::new(env!("CARGO_BIN_EXE_smg")) .arg("--version") - .env("_RJEM_MALLOC_CONF", "stats_print:true,stats_print_opts:J") + .env("_RJEM_MALLOC_CONF", malloc_conf) .output() - .expect("run the SMG executable"); - assert!( - output.status.success(), - "SMG executable failed: {}", - String::from_utf8_lossy(&output.stderr) - ); + .map_err(|e| format!("run the SMG executable: {e}"))?; + if !output.status.success() { + return Err(format!( + "SMG executable failed: {}", + String::from_utf8_lossy(&output.stderr) + )); + } + Ok(String::from_utf8_lossy(&output.stderr).into_owned()) +} + +/// The `"opt"` object of jemalloc's JSON statistics: the options in effect. +/// It holds scalars only, so it ends at the first closing brace. +fn opt_section(stats: &str) -> Option<&str> { + let start = stats.find("\"opt\":{")?; + let end = start + stats[start..].find('}')?; + Some(&stats[start..end]) +} - let stats = String::from_utf8_lossy(&output.stderr); +#[test] +fn smg_binary_uses_global_jemalloc() { + let stats = jemalloc_stats("stats_print:true,stats_print_opts:J") + .expect("jemalloc statistics from the SMG executable"); assert!( stats.contains("\"jemalloc\""), "SMG executable did not emit jemalloc statistics: {stats}" @@ -33,3 +49,41 @@ fn smg_binary_uses_global_jemalloc() { "aarch64 SMG executable was not built for 64 KiB page compatibility: {stats}" ); } + +#[test] +fn smg_binary_applies_the_server_malloc_conf() { + // The executable's `_rjem_malloc_conf` symbol: purge on a background + // thread, dirty pages back to the OS after 10 s, muzzy pages at once. + let stats = jemalloc_stats("stats_print:true,stats_print_opts:J") + .expect("jemalloc statistics from the SMG executable"); + let opt = opt_section(&stats).expect("an \"opt\" section in the jemalloc statistics"); + for expected in [ + "\"background_thread\":true", + "\"dirty_decay_ms\":10000", + "\"muzzy_decay_ms\":0", + ] { + assert!( + opt.contains(expected), + "jemalloc did not take the executable's malloc_conf ({expected} missing): {opt}" + ); + } +} + +#[test] +fn environment_overrides_the_server_malloc_conf_entry_by_entry() { + let stats = jemalloc_stats( + "stats_print:true,stats_print_opts:J,background_thread:false,dirty_decay_ms:2000", + ) + .expect("jemalloc statistics from the SMG executable"); + let opt = opt_section(&stats).expect("an \"opt\" section in the jemalloc statistics"); + for expected in [ + "\"background_thread\":false", + "\"dirty_decay_ms\":2000", + "\"muzzy_decay_ms\":0", + ] { + assert!( + opt.contains(expected), + "_RJEM_MALLOC_CONF did not override the executable's malloc_conf ({expected} missing): {opt}" + ); + } +} diff --git a/model_gateway/tests/api/api_endpoints_test.rs b/model_gateway/tests/api/api_endpoints_test.rs index 5e58779ded..a542c9440e 100644 --- a/model_gateway/tests/api/api_endpoints_test.rs +++ b/model_gateway/tests/api/api_endpoints_test.rs @@ -1487,11 +1487,9 @@ mod cache_tests { .map(|entry| entry["worker"].as_str().expect("worker")) .collect(); reporters.sort_unstable(); - assert_eq!( - reporters, - vec!["http://127.0.0.1:18502", "http://127.0.0.1:18503"], - "{uri}" - ); + let mut expected: Vec<&str> = ctx.worker_urls.iter().map(String::as_str).collect(); + expected.sort_unstable(); + assert_eq!(reporters, expected, "{uri}"); for entry in loads { assert_eq!(entry["worker_type"], "regular", "{uri}"); assert_eq!(entry["num_running_reqs"], 2, "{uri}"); diff --git a/model_gateway/tests/common/mock_worker.rs b/model_gateway/tests/common/mock_worker.rs index 4fb3cf958a..10b57f7a47 100755 --- a/model_gateway/tests/common/mock_worker.rs +++ b/model_gateway/tests/common/mock_worker.rs @@ -290,9 +290,11 @@ pub struct MockWorker { config: Arc>, shutdown_handle: Option>, shutdown_tx: Option>, - /// Resolved bind port, cached so sync `Drop` can prune this worker's entry - /// from the global scheduler-controls table. - bound_port: Option, + /// The port the worker is known by: the configured one, or the bound one + /// when the config left it 0. The per-port test tables (scheduler + /// controls, request recorders) are keyed by it, and sync `Drop` prunes + /// this worker's entries with it. + named_port: Option, } impl MockWorker { @@ -301,11 +303,21 @@ impl MockWorker { config: Arc::new(RwLock::new(config)), shutdown_handle: None, shutdown_tx: None, - bound_port: None, + named_port: None, } } - /// Start the mock worker server + /// Start the mock worker server and return the URL it answers on. + /// + /// The worker binds the configured port when it is free (a test that + /// announces the port elsewhere, as the discovery tests do through a pod + /// annotation, needs the engine on exactly that port) and an ephemeral + /// one when the config left it 0 or another process on the host holds it + /// (a fleet running beside the tests, say the replay harness on 19500 and + /// up per lane, must not make a test fail to bind). Either way the + /// configured port stays the worker's *name*: the per-port test tables + /// are keyed by it, and the app test context rewrites URLs a test spells + /// out with it to the bound address. #[expect( clippy::disallowed_methods, clippy::print_stderr, @@ -313,19 +325,19 @@ impl MockWorker { )] pub async fn start(&mut self) -> Result> { let config = self.config.clone(); - let port = config.read().await.port; - - // If port is 0, find an available port - let port = if port == 0 { - let listener = std::net::TcpListener::bind("127.0.0.1:0")?; - let port = listener.local_addr()?.port(); - drop(listener); - config.write().await.port = port; - port - } else { - port + let named_port = config.read().await.port; + let listener = match named_port { + 0 => tokio::net::TcpListener::bind(("127.0.0.1", 0)).await?, + named => match tokio::net::TcpListener::bind(("127.0.0.1", named)).await { + Ok(listener) => listener, + Err(_) => tokio::net::TcpListener::bind(("127.0.0.1", 0)).await?, + }, }; - self.bound_port = Some(port); + let port = listener.local_addr()?.port(); + if named_port == 0 { + config.write().await.port = port; + } + self.named_port = Some(if named_port == 0 { port } else { named_port }); let app = Router::new() .route("/health", get(health_handler)) @@ -369,16 +381,9 @@ impl MockWorker { let (shutdown_tx, shutdown_rx) = oneshot::channel::<()>(); self.shutdown_tx = Some(shutdown_tx); - // Spawn the server in a separate task + // The listener is bound and queuing connections before this returns, + // so no start-up wait is needed. let handle = tokio::spawn(async move { - let listener = match tokio::net::TcpListener::bind(("127.0.0.1", port)).await { - Ok(l) => l, - Err(e) => { - eprintln!("Failed to bind to port {port}: {e}"); - return; - } - }; - let server = axum::serve(listener, app).with_graceful_shutdown(async move { let _ = shutdown_rx.await; }); @@ -390,11 +395,7 @@ impl MockWorker { self.shutdown_handle = Some(handle); - // Wait for the server to start - tokio::time::sleep(tokio::time::Duration::from_millis(100)).await; - - let url = format!("http://127.0.0.1:{port}"); - Ok(url) + Ok(format!("http://127.0.0.1:{port}")) } /// Stop the mock worker server @@ -417,8 +418,8 @@ impl Drop for MockWorker { let _ = shutdown_tx.send(()); } // Prune our scheduler controls and recorder so a later worker reusing - // this port doesn't inherit stale state. - if let Some(port) = self.bound_port { + // this name doesn't inherit stale state. + if let Some(port) = self.named_port { clear_scheduler_controls(port); clear_request_recorder(port); } diff --git a/model_gateway/tests/common/mod.rs b/model_gateway/tests/common/mod.rs index b40c7f7f22..a8239374a3 100644 --- a/model_gateway/tests/common/mod.rs +++ b/model_gateway/tests/common/mod.rs @@ -13,6 +13,7 @@ pub mod tls_mock_worker; // Re-export commonly used test builders use std::{ + collections::HashMap, fs, future::Future, path::PathBuf, @@ -85,7 +86,10 @@ impl WorkerTestContext { endpoint: &str, body: serde_json::Value, ) -> Result { - let client = reqwest::Client::new(); + let client = reqwest::Client::builder() + .no_proxy() + .build() + .map_err(|e| format!("test client: {e}"))?; let worker_url = self .first_worker_url() .ok_or_else(|| "No workers available".to_string())?; @@ -114,7 +118,10 @@ impl WorkerTestContext { ) -> Result, String> { use futures_util::StreamExt; - let client = reqwest::Client::new(); + let client = reqwest::Client::builder() + .no_proxy() + .build() + .map_err(|e| format!("test client: {e}"))?; let worker_url = self .first_worker_url() .ok_or_else(|| "No workers available".to_string())?; @@ -172,6 +179,23 @@ pub struct AppTestContext { pub router: Arc, pub config: RouterConfig, pub app_context: Arc, + /// The URL each worker bound, in the order the configs were given. + pub worker_urls: Vec, + /// The port each worker's config named, parallel to `worker_urls`. + named_ports: Vec, +} + +/// Rewrite a URL that names one of the started workers by its configured port +/// to the address that worker actually bound; any other URL is left alone. +fn rebind_url(url: &mut String, bound_by_named_port: &HashMap) { + let named = url + .trim_end_matches('/') + .rsplit(':') + .next() + .and_then(|port| port.parse::().ok()); + if let Some(bound) = named.and_then(|port| bound_by_named_port.get(&port)) { + url.clone_from(bound); + } } impl AppTestContext { @@ -212,10 +236,21 @@ impl AppTestContext { Box::pin(async move { let mut workers = Vec::new(); let mut worker_urls = Vec::new(); + let mut named_ports = Vec::new(); + // Workers bind ephemeral ports (see `MockWorker::start`); the + // ports the configs name are identities. URLs the router config + // spells out with those names are rewritten to the bound + // addresses below. + let mut bound_by_named_port: HashMap = HashMap::new(); for worker_config in worker_configs { + let named_port = worker_config.port; let mut worker = MockWorker::new(worker_config); let url = worker.start().await.unwrap(); + if named_port != 0 { + bound_by_named_port.insert(named_port, url.clone()); + } + named_ports.push(named_port); worker_urls.push(url); workers.push(worker); } @@ -230,10 +265,46 @@ impl AppTestContext { } | RoutingMode::OpenAI { worker_urls: ref mut urls, - } if urls.is_empty() => { - urls.clone_from(&worker_urls); } - _ => {} + | RoutingMode::Anthropic { + worker_urls: ref mut urls, + } + | RoutingMode::Gemini { + worker_urls: ref mut urls, + } => { + if urls.is_empty() { + urls.clone_from(&worker_urls); + } else { + for url in urls.iter_mut() { + rebind_url(url, &bound_by_named_port); + } + } + } + RoutingMode::PrefillDecode { + prefill_urls, + decode_urls, + .. + } => { + for (url, _) in prefill_urls.iter_mut() { + rebind_url(url, &bound_by_named_port); + } + for url in decode_urls.iter_mut() { + rebind_url(url, &bound_by_named_port); + } + } + RoutingMode::EncodePrefillDecode { + encode_urls, + prefill_urls, + decode_urls, + .. + } => { + for (url, _) in encode_urls.iter_mut().chain(prefill_urls.iter_mut()) { + rebind_url(url, &bound_by_named_port); + } + for url in decode_urls.iter_mut() { + rebind_url(url, &bound_by_named_port); + } + } } let app_context = create_test_context(config.clone()).await; @@ -287,10 +358,27 @@ impl AppTestContext { router, config, app_context, + worker_urls, + named_ports, } }) } + /// The URL bound by the worker whose config named `port` (the port a + /// test spelled in its `MockWorkerConfig`, not the one actually bound). + #[expect( + clippy::expect_used, + reason = "test helper - panicking on failure is intentional" + )] + pub fn worker_url_for(&self, port: u16) -> &str { + let index = self + .named_ports + .iter() + .position(|named| *named == port) + .expect("a worker config named this port"); + &self.worker_urls[index] + } + pub fn create_app(&self) -> axum::Router { test_app::create_test_app_with_context( Arc::clone(&self.router), @@ -326,7 +414,11 @@ async fn build_test_app_context( ) -> Arc { use smg_mcp::McpOrchestrator; - let client = reqwest::Client::new(); + // See `test_app::create_test_app_context`: no environment proxy in a test's path. + let client = reqwest::Client::builder() + .no_proxy() + .build() + .expect("test client"); // Initialize rate limiter let rate_limiter = match config.max_concurrent_requests { diff --git a/model_gateway/tests/common/test_app.rs b/model_gateway/tests/common/test_app.rs index dbd8f2b419..4bc1aaffea 100644 --- a/model_gateway/tests/common/test_app.rs +++ b/model_gateway/tests/common/test_app.rs @@ -218,7 +218,9 @@ pub fn create_test_app_with_context( )] pub async fn create_test_app_context() -> Arc { let router_config = RouterConfig::default(); - let client = Client::new(); + // A test's upstreams are its own loopback servers; the environment's + // proxy would otherwise sit in the path of every request. + let client = Client::builder().no_proxy().build().expect("test client"); // Initialize empty OnceLocks let worker_job_queue = Arc::new(OnceLock::new()); diff --git a/model_gateway/tests/grpc_context_length_test.rs b/model_gateway/tests/grpc_context_length_test.rs index 6fad2b926e..db78a8f8dc 100644 --- a/model_gateway/tests/grpc_context_length_test.rs +++ b/model_gateway/tests/grpc_context_length_test.rs @@ -54,6 +54,8 @@ async fn start_mock_grpc_worker() -> u16 { .expect("mock gRPC worker address") .port(); let cfg = Arc::new(mock_worker::config::Config { + admin_port: None, + context_length: 32768, host: "127.0.0.1".to_string(), http_base_port: 0, http_count: 0, @@ -68,7 +70,7 @@ async fn start_mock_grpc_worker() -> u16 { output_tokens: 2, realistic: false, engine: mock_worker::engine::EngineParams::default(), - replay: Default::default(), + ..mock_worker::config::Config::default() }); tokio::spawn(mock_worker::grpc::serve_with_listener(cfg, listener)); port diff --git a/model_gateway/tests/grpc_pd_fanout_test.rs b/model_gateway/tests/grpc_pd_fanout_test.rs index 5a5954678e..a0be204baa 100644 --- a/model_gateway/tests/grpc_pd_fanout_test.rs +++ b/model_gateway/tests/grpc_pd_fanout_test.rs @@ -115,6 +115,8 @@ async fn start_mock_grpc_worker() -> u16 { .expect("mock gRPC worker address") .port(); let cfg = Arc::new(mock_worker::config::Config { + admin_port: None, + context_length: 32768, host: "127.0.0.1".to_string(), http_base_port: 0, http_count: 0, @@ -129,7 +131,7 @@ async fn start_mock_grpc_worker() -> u16 { output_tokens: OUTPUT_TOKENS, realistic: false, engine: mock_worker::engine::EngineParams::default(), - replay: Default::default(), + ..mock_worker::config::Config::default() }); tokio::spawn(mock_worker::grpc::serve_with_listener(cfg, listener)); port diff --git a/model_gateway/tests/pushed_loads_test.rs b/model_gateway/tests/pushed_loads_test.rs new file mode 100644 index 0000000000..603b3d7b65 --- /dev/null +++ b/model_gateway/tests/pushed_loads_test.rs @@ -0,0 +1,193 @@ +//! A fleet of eight realistic mock gRPC workers reporting loads the way the +//! vLLM servicer does (no queued token-work): with load polling held off, +//! their loads still reach the gateway through the KV-event streams, as the +//! `EngineLoad` record on every batch and as `load_only` heartbeats while +//! the engines are idle, and land in the same `smg_engine_*` gauges a poll +//! fills. + +mod common; + +use std::{ + sync::Arc, + time::{Duration, Instant}, +}; + +use llm_tokenizer::{mock::MockTokenizer, traits::Tokenizer, TokenizerRegistry}; +use openai_protocol::worker::{ConnectionMode, HealthCheckConfig, RuntimeType, WorkerType}; +use smg::{ + config::RouterConfig, + observability::{ + metrics::{start_prometheus, PrometheusConfig}, + metrics_server::start_metrics_server, + }, + worker::{BasicWorkerBuilder, KvEventMonitor, ModelCard}, +}; +use tokio::net::TcpListener; + +const MODEL: &str = "mock-model"; +const WORKERS: usize = 8; + +/// A realistic mock gRPC worker with prefix caching (so it streams KV events) +/// and vLLM-like load reports. +#[expect( + clippy::expect_used, + clippy::disallowed_methods, + reason = "test helper - panicking on failure is intentional; the mock server task is fire-and-forget" +)] +async fn start_realistic_worker() -> u16 { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("bind mock gRPC worker"); + let port = listener + .local_addr() + .expect("mock gRPC worker address") + .port(); + let cfg = Arc::new(mock_worker::config::Config { + host: "127.0.0.1".to_string(), + http_count: 0, + grpc_base_port: port, + grpc_count: 1, + model_id: MODEL.to_string(), + tokenizer_path: MODEL.to_string(), + realistic: true, + engine: mock_worker::engine::EngineParams { + prefix_cache: true, + ..mock_worker::engine::EngineParams::default() + }, + loads_like: mock_worker::engine::LoadsLike::Vllm, + ..mock_worker::config::Config::default() + }); + tokio::spawn(mock_worker::grpc::serve_with_listener(cfg, listener)); + port +} + +#[tokio::test] +async fn eight_workers_loads_reach_the_gateway_through_the_event_streams() { + let handle = start_prometheus(PrometheusConfig { + port: 0, + host: "127.0.0.1".to_string(), + duration_buckets: None, + }); + let (metrics_addr, _metrics_server) = start_metrics_server(handle, "127.0.0.1".to_string(), 0) + .await + .expect("metrics server binds an ephemeral port"); + + let mut ports = Vec::with_capacity(WORKERS); + for _ in 0..WORKERS { + ports.push(start_realistic_worker().await); + } + + // Cache-aware routing (the KV-event subscriptions) with the load poll held + // off: the monitor's loops are never started, so every load the gateway + // learns came through an event stream. + let mut config = RouterConfig::builder() + .grpc_connection() + .cache_aware_policy(0.5, 32, 1.1, 60, 1_000_000) + .load_monitor_interval_secs(3600) + .host("127.0.0.1") + .port(0) + .build_unchecked(); + config.health_check.disable_health_check = true; + let tokenizer_registry = Arc::new(TokenizerRegistry::new()); + let tokenizer = Arc::new(MockTokenizer::new()) as Arc; + tokenizer_registry + .load( + "tokenizer-id", + MODEL, + "test", + || async move { Ok(tokenizer) }, + ) + .await + .unwrap(); + let app_context = + common::create_test_context_with_tokenizer_registry(config, tokenizer_registry).await; + let worker_monitor = app_context + .worker_monitor + .clone() + .expect("the test context builds a worker monitor"); + let kv_monitor = Arc::new(KvEventMonitor::new(None)); + kv_monitor.set_load_sink(&worker_monitor); + + let mut urls = Vec::with_capacity(WORKERS); + for port in &ports { + let url = format!("grpc://127.0.0.1:{port}"); + let worker = BasicWorkerBuilder::new(url.clone()) + .worker_type(WorkerType::Regular) + .connection_mode(ConnectionMode::Grpc) + .runtime_type(RuntimeType::TokenSpeed) + .model(ModelCard::new(MODEL)) + .health_config(HealthCheckConfig { + disable_health_check: true, + ..Default::default() + }) + .build(); + let worker: Arc = Arc::new(worker); + worker.set_status(openai_protocol::worker::WorkerStatus::Ready); + app_context + .worker_registry + .register(Arc::clone(&worker)) + .unwrap(); + kv_monitor.on_worker_added(&worker).await; + urls.push(url); + } + + // Idle engines publish no KV events; their streams still send the load + // record as heartbeats within a second or so of the subscription. + let client = reqwest::Client::builder() + .no_proxy() + .build() + .expect("test client"); + let deadline = Instant::now() + Duration::from_secs(15); + let body = loop { + let body = client + .get(format!("http://{metrics_addr}/metrics")) + .send() + .await + .expect("metrics endpoint reachable") + .text() + .await + .expect("metrics body"); + let reported = urls + .iter() + .filter(|url| { + body.lines().any(|line| { + line.starts_with("smg_engine_running_requests{") + && line.contains(&format!("worker=\"{url}\"")) + }) + }) + .count(); + // The records come as load_only batches (the engines are idle) that + // the gateway counts as such, admitting none into the index. A + // worker's first load may reach the gateway by its poll before its + // stream's heartbeat, so the batch count is waited for as well. + let load_only: f64 = body + .lines() + .filter(|line| { + line.starts_with("smg_kv_event_batches_total{") + && line.contains("disposition=\"load_only\"") + }) + .filter_map(|line| line.rsplit(' ').next()?.parse::().ok()) + .sum(); + if reported == WORKERS && load_only >= WORKERS as f64 { + break body; + } + assert!( + Instant::now() < deadline, + "{reported} of {WORKERS} workers reported a load and {load_only} load_only \ + batches arrived through the event streams:\n{body}" + ); + tokio::time::sleep(Duration::from_millis(100)).await; + }; + for url in &urls { + let usage = body + .lines() + .find(|line| { + line.starts_with("smg_engine_token_usage{") + && line.contains(&format!("worker=\"{url}\"")) + }) + .unwrap_or_else(|| panic!("token usage gauge for {url} missing:\n{body}")); + let value: f64 = usage.rsplit(' ').next().unwrap().parse().unwrap(); + assert!((0.0..=1.0).contains(&value), "{usage}"); + } + kv_monitor.stop().await; +} diff --git a/model_gateway/tests/routing/mod.rs b/model_gateway/tests/routing/mod.rs index 73c2703ec2..6b7fedb3d1 100644 --- a/model_gateway/tests/routing/mod.rs +++ b/model_gateway/tests/routing/mod.rs @@ -10,6 +10,7 @@ pub mod manual_routing_test; pub mod model_alias_test; pub mod payload_size_test; pub mod pd_routing_test; +pub mod policy_completion_test; pub mod policy_registry_integration; pub mod power_of_two_test; pub mod prefix_hash_test; diff --git a/model_gateway/tests/routing/model_alias_test.rs b/model_gateway/tests/routing/model_alias_test.rs index 34662746da..99667fe777 100644 --- a/model_gateway/tests/routing/model_alias_test.rs +++ b/model_gateway/tests/routing/model_alias_test.rs @@ -76,7 +76,7 @@ mod model_alias_tests { let ctx = AppTestContext::new(vec![TestWorkerConfig::healthy(WORKER_PORT)]).await; let app = ctx.create_app(); - declare_alias(&ctx, &format!("http://127.0.0.1:{WORKER_PORT}")); + declare_alias(&ctx, ctx.worker_url_for(WORKER_PORT)); let response = app .clone() @@ -103,7 +103,7 @@ mod model_alias_tests { let ctx = AppTestContext::new(vec![TestWorkerConfig::healthy(WORKER_PORT + 1)]).await; let app = ctx.create_app(); - declare_alias(&ctx, &format!("http://127.0.0.1:{}", WORKER_PORT + 1)); + declare_alias(&ctx, ctx.worker_url_for(WORKER_PORT + 1)); let response = app .clone() @@ -130,7 +130,7 @@ mod model_alias_tests { let ctx = AppTestContext::new(vec![TestWorkerConfig::healthy(WORKER_PORT + 3)]).await; let app = ctx.create_app(); - declare_alias(&ctx, &format!("http://127.0.0.1:{}", WORKER_PORT + 3)); + declare_alias(&ctx, ctx.worker_url_for(WORKER_PORT + 3)); let request = Request::builder() .method("POST") @@ -175,7 +175,7 @@ mod model_alias_tests { let ctx = AppTestContext::new(vec![TestWorkerConfig::healthy(WORKER_PORT + 4)]).await; let app = ctx.create_app(); - declare_alias(&ctx, &format!("http://127.0.0.1:{}", WORKER_PORT + 4)); + declare_alias(&ctx, ctx.worker_url_for(WORKER_PORT + 4)); let boundary = "alias-test-boundary"; let form = format!( @@ -217,7 +217,7 @@ mod model_alias_tests { let ctx = AppTestContext::new(vec![TestWorkerConfig::healthy(WORKER_PORT + 5)]).await; let app = ctx.create_app(); - declare_alias(&ctx, &format!("http://127.0.0.1:{}", WORKER_PORT + 5)); + declare_alias(&ctx, ctx.worker_url_for(WORKER_PORT + 5)); for (endpoint, payload) in [ ( @@ -264,7 +264,7 @@ mod model_alias_tests { let ctx = AppTestContext::new(vec![TestWorkerConfig::healthy(WORKER_PORT + 2)]).await; let app = ctx.create_app(); - declare_alias(&ctx, &format!("http://127.0.0.1:{}", WORKER_PORT + 2)); + declare_alias(&ctx, ctx.worker_url_for(WORKER_PORT + 2)); let response = app .clone() diff --git a/model_gateway/tests/routing/pd_routing_test.rs b/model_gateway/tests/routing/pd_routing_test.rs index 50dae88302..d92b23af72 100644 --- a/model_gateway/tests/routing/pd_routing_test.rs +++ b/model_gateway/tests/routing/pd_routing_test.rs @@ -197,8 +197,9 @@ mod pd_routing_tests { .await; let app = ctx.create_app(); + // The workers bound ephemeral ports; the registry knows them by those. let registry = &ctx.app_context.worker_registry; - for url in [&prefill_url, &decode_url] { + for url in [ctx.worker_url_for(19840), ctx.worker_url_for(19841)] { let worker_id = registry.get_id_by_url(url).unwrap(); let worker = registry.get(&worker_id).unwrap(); let mut spec = worker.metadata().spec.as_ref().clone(); diff --git a/model_gateway/tests/routing/policy_completion_test.rs b/model_gateway/tests/routing/policy_completion_test.rs new file mode 100644 index 0000000000..ff52d78ce4 --- /dev/null +++ b/model_gateway/tests/routing/policy_completion_test.rs @@ -0,0 +1,129 @@ +//! Policies learn that a request ended through the worker's load guard, on +//! the HTTP router's streaming path too: a client that drops the response +//! mid-stream must release the policy's booking for that worker. +use std::{ + convert::Infallible, + sync::{Arc, Mutex}, + time::Duration, +}; + +use axum::{ + response::sse::{Event, Sse}, + routing::post, + Router as AxumRouter, +}; +use futures_util::{stream, StreamExt}; +use http_body_util::BodyExt; +use openai_protocol::chat::ChatCompletionRequest; +use serde_json::json; +use smg::{ + routers::{router::Router as HttpRouter, RouterTrait}, + tenant::{RouteRequestMeta, TenantKey}, + worker::{BasicWorkerBuilder, ModelCard, RequestCompletionSink, Worker}, +}; +use tokio::{net::TcpListener, time::timeout}; + +use crate::common::test_app::create_test_app_context; + +/// Records the worker URLs that reported a completion. +#[derive(Debug, Default)] +struct CompletionSpy(Mutex>); + +impl RequestCompletionSink for CompletionSpy { + fn request_completed(&self, worker: &dyn Worker) { + self.0 + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()) + .push(worker.url().to_string()); + } +} + +/// Upstream that answers with SSE headers, one chunk, then stalls forever. +#[expect( + clippy::disallowed_methods, + clippy::unwrap_used, + reason = "test infrastructure - panicking on failure is intentional" +)] +async fn spawn_stalling_upstream() -> String { + let handler = || async { + let body = stream::iter([Ok::(Event::default().data("head"))]) + .chain(stream::pending::>()); + Sse::new(body) + }; + let app = AxumRouter::new().route("/v1/chat/completions", post(handler)); + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + tokio::spawn(async move { + axum::serve(listener, app).await.unwrap(); + }); + format!("http://{addr}") +} + +#[expect( + clippy::unwrap_used, + reason = "test infrastructure - panicking on failure is intentional" +)] +async fn router_with_spy(upstream_url: &str) -> (HttpRouter, Arc) { + let ctx = create_test_app_context().await; + let spy = Arc::new(CompletionSpy::default()); + let worker: Arc = Arc::new( + BasicWorkerBuilder::new(upstream_url) + .models(vec![ModelCard::new("mock-model")]) + .health_config(openai_protocol::worker::HealthCheckConfig { + disable_health_check: true, + ..Default::default() + }) + .build(), + ); + worker.set_completion_sink(Some(Arc::clone(&spy) as Arc)); + ctx.worker_registry.register(worker); + (HttpRouter::new(&ctx).await.unwrap(), spy) +} + +#[expect( + clippy::unwrap_used, + reason = "test infrastructure - panicking on failure is intentional" +)] +fn streaming_chat_request() -> ChatCompletionRequest { + serde_json::from_value(json!({ + "model": "mock-model", + "messages": [{"role": "user", "content": "Hello"}], + "stream": true + })) + .unwrap() +} + +/// A stream the client abandons after the first chunk still reports exactly +/// one completion for the worker that served it, once the relay lets go. +#[tokio::test] +async fn http_stream_cancelled_mid_way_reports_one_completion() { + let upstream_url = spawn_stalling_upstream().await; + let (router, spy) = router_with_spy(&upstream_url).await; + let meta = RouteRequestMeta::new(TenantKey::from("test-tenant")); + let response = router + .route_chat(None, &meta, streaming_chat_request(), "mock-model") + .await; + assert_eq!(response.status(), 200); + let mut body = response.into_body(); + let first = body.frame().await; + assert!( + matches!(first, Some(Ok(_))), + "expected a first chunk, got {first:?}" + ); + assert!( + spy.0.lock().unwrap().is_empty(), + "no completion while the client is still reading" + ); + drop(body); + timeout(Duration::from_secs(5), async { + loop { + if spy.0.lock().unwrap().len() == 1 { + break; + } + tokio::time::sleep(Duration::from_millis(20)).await; + } + }) + .await + .expect("the cancelled stream never reported its completion"); + assert_eq!(spy.0.lock().unwrap().as_slice(), [upstream_url.as_str()]); +} diff --git a/model_gateway/tests/routing/stream_request_body_test.rs b/model_gateway/tests/routing/stream_request_body_test.rs index 54b436da22..1fa6452862 100644 --- a/model_gateway/tests/routing/stream_request_body_test.rs +++ b/model_gateway/tests/routing/stream_request_body_test.rs @@ -87,6 +87,8 @@ fn cache_aware_policy() -> PolicyConfig { cache_index: Default::default(), cache_ttl_secs: 180, cache_boundaries: Vec::new(), + selection_policy: None, + selection_accounting_ttl_ms: 0, } } diff --git a/model_gateway/tests/routing/test_openai_routing.rs b/model_gateway/tests/routing/test_openai_routing.rs index 7542c773a8..0d3b34c9bc 100644 --- a/model_gateway/tests/routing/test_openai_routing.rs +++ b/model_gateway/tests/routing/test_openai_routing.rs @@ -915,7 +915,16 @@ async fn assert_streaming_json_error_content_type(status: StatusCode) { #[tokio::test] async fn test_openai_router_circuit_breaker() { let ctx = create_test_app_context().await; - register_external_worker(&ctx, "http://invalid-url-that-will-fail", None); + // A loopback port that refuses connections: bound once to be handed out, + // then released. Every attempt fails at connect, with no resolver and no + // proxy involved, whatever the environment. + let refused = { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let port = listener.local_addr().unwrap().port(); + drop(listener); + format!("http://127.0.0.1:{port}") + }; + register_external_worker(&ctx, &refused, None); let router = OpenAIRouter::new(&external_context(&ctx)).await.unwrap(); let chat_request = create_minimal_chat_request(); diff --git a/model_gateway/tests/routing/test_pd_routing.rs b/model_gateway/tests/routing/test_pd_routing.rs index 2d563ba2c4..043e8756bc 100644 --- a/model_gateway/tests/routing/test_pd_routing.rs +++ b/model_gateway/tests/routing/test_pd_routing.rs @@ -173,6 +173,8 @@ mod pd_routing_unit_tests { cache_index: Default::default(), cache_ttl_secs: 180, cache_boundaries: Vec::new(), + selection_policy: None, + selection_accounting_ttl_ms: 0, }, ), ( diff --git a/model_gateway/tests/tenant_rate_limiting_grpc_test.rs b/model_gateway/tests/tenant_rate_limiting_grpc_test.rs index d5775eca98..b46d33dea5 100644 --- a/model_gateway/tests/tenant_rate_limiting_grpc_test.rs +++ b/model_gateway/tests/tenant_rate_limiting_grpc_test.rs @@ -64,6 +64,8 @@ async fn start_mock_grpc_worker(output_tokens: u32) -> u16 { .expect("mock gRPC worker address") .port(); let cfg = Arc::new(mock_worker::config::Config { + admin_port: None, + context_length: 32768, host: "127.0.0.1".to_string(), http_base_port: 0, http_count: 0, @@ -78,7 +80,7 @@ async fn start_mock_grpc_worker(output_tokens: u32) -> u16 { output_tokens, realistic: false, engine: mock_worker::engine::EngineParams::default(), - replay: Default::default(), + ..mock_worker::config::Config::default() }); tokio::spawn(mock_worker::grpc::serve_with_listener(cfg, listener)); port diff --git a/model_gateway/tests/zmq_backend_test.rs b/model_gateway/tests/zmq_backend_test.rs index 24679d19c2..1b88d5377a 100644 --- a/model_gateway/tests/zmq_backend_test.rs +++ b/model_gateway/tests/zmq_backend_test.rs @@ -73,6 +73,8 @@ fn zmq_fixture() -> ZmqFixture { )] fn start_mock_zmq_engines(handshake: &str, count: u16) { let cfg = Arc::new(mock_worker::config::Config { + admin_port: None, + context_length: 32768, host: "127.0.0.1".to_string(), http_base_port: 0, http_count: 0, @@ -87,7 +89,7 @@ fn start_mock_zmq_engines(handshake: &str, count: u16) { output_tokens: OUTPUT_TOKENS, realistic: false, engine: mock_worker::engine::EngineParams::default(), - replay: Default::default(), + ..mock_worker::config::Config::default() }); for rank in 0..u32::from(count) { tokio::spawn(mock_worker::zmq::serve(