diff --git a/Cargo.lock b/Cargo.lock index cef1e055e2..f05c12d5a6 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -4397,6 +4397,7 @@ dependencies = [ "engine-servicer", "engine-zmq-client", "futures", + "prost-types", "reqwest 0.13.4", "rmp-serde", "rmpv", diff --git a/bindings/python/src/smg/rl.py b/bindings/python/src/smg/rl.py index 82fede79eb..986095ce82 100644 --- a/bindings/python/src/smg/rl.py +++ b/bindings/python/src/smg/rl.py @@ -8,9 +8,10 @@ SGLang requires a JSON body on `pause_generation` and `continue_generation` (a bodyless POST is a 400), so bodyless routes are sent as `{}`. -Only HTTP workers can be proxied. A gRPC or ZMQ worker matched by a selector is -reported in `failed[]` as `unsupported_connection_mode`, which makes `fanout` -raise `FanoutError` unless `allow_partial=True`. +A worker with no control endpoint (a gRPC or ZMQ worker without an +`rl.control_url` label) matched by a selector is reported in `failed[]` as +`no_control_endpoint`, which makes `fanout` raise `FanoutError` unless +`allow_partial=True`. """ from __future__ import annotations @@ -52,6 +53,7 @@ class Worker: role: str | None health: str weight_version: str | None + control_url: str | None = None labels: dict[str, str] = field(default_factory=dict) capabilities: dict[str, Any] = field(default_factory=dict) diff --git a/bindings/python/src/smg/serve.py b/bindings/python/src/smg/serve.py index 6a20667b6d..a012bf3db0 100644 --- a/bindings/python/src/smg/serve.py +++ b/bindings/python/src/smg/serve.py @@ -16,6 +16,7 @@ import subprocess import sys import time +import urllib.request from abc import ABC, abstractmethod from smg.launch_router import launch_router @@ -615,8 +616,6 @@ def build_command( def _http_health_check(url: str, timeout: float) -> bool: """GET the URL and return True on HTTP 200.""" try: - import urllib.request - req = urllib.request.Request(url, method="GET") with urllib.request.urlopen(req, timeout=timeout) as resp: return resp.status == 200 diff --git a/bindings/python/tests/test_rl_client.py b/bindings/python/tests/test_rl_client.py index e34aa9a77a..be9a73de6f 100644 --- a/bindings/python/tests/test_rl_client.py +++ b/bindings/python/tests/test_rl_client.py @@ -19,6 +19,7 @@ "model_id": "m", "worker_type": "regular", "connection_mode": "http", + "control_url": "http://e:1", "tp_size": 1, "dp_size": 1, "pp_size": 1, @@ -89,6 +90,7 @@ def test_workers_and_worker(stub): ws = rl.workers() assert len(ws) == 1 and ws[0].id == "w1" and ws[0].engine == "sglang" assert ws[0].capabilities["pause_modes"] == ["abort"] + assert ws[0].control_url == "http://e:1" assert rl.worker("w1").weight_version == "7" assert _Stub.seen[0]["auth"] == "Bearer k" @@ -155,11 +157,13 @@ def test_worker_from_json_defaults_missing_dicts(): del d["labels"] del d["capabilities"] del d["role"] + del d["control_url"] d["future_field"] = 1 w = Worker.from_json(d) assert w.labels == {} assert w.capabilities == {} assert w.role is None + assert w.control_url is None def test_call_raises_on_smg_error_envelope(stub): diff --git a/crates/mock_worker/Cargo.toml b/crates/mock_worker/Cargo.toml index e1318cc412..7ded00414f 100644 --- a/crates/mock_worker/Cargo.toml +++ b/crates/mock_worker/Cargo.toml @@ -33,6 +33,7 @@ serde_json.workspace = true futures.workspace = true tracing.workspace = true tracing-subscriber.workspace = true +prost-types.workspace = true rmpv.workspace = true zeromq.workspace = true diff --git a/crates/mock_worker/src/config.rs b/crates/mock_worker/src/config.rs index 10149afb06..972713c60a 100644 --- a/crates/mock_worker/src/config.rs +++ b/crates/mock_worker/src/config.rs @@ -1,6 +1,6 @@ //! Runtime configuration for the mock worker fleet, parsed from CLI flags. -use std::{path::PathBuf, time::Duration}; +use std::{collections::BTreeMap, path::PathBuf, time::Duration}; use crate::engine::{Calibration, EngineParams, LoadsLike, TimingModel}; @@ -38,6 +38,12 @@ pub struct Config { pub realistic: bool, /// Engine-simulator parameters (only used when `realistic`). pub engine: EngineParams, + /// Extra `server_args` entries the gRPC `GetServerInfo` advertises, as + /// string values (e.g. `rl.control_url`). Empty leaves `server_args` unset. + pub server_args: BTreeMap, + /// Weight version stamped on every gRPC generate chunk and completion; + /// `None` leaves the field unset, like an engine that predates it. + pub weight_version: Option, /// Port of the process-wide admin API (fleet, request records, cache /// dumps, resets); off when `None`. pub admin_port: Option, @@ -151,6 +157,8 @@ impl Default for Config { output_tokens: 8, realistic: false, engine: EngineParams::default(), + server_args: BTreeMap::new(), + weight_version: None, admin_port: None, context_length: 32768, kv_events_zmq_base_port: None, diff --git a/crates/mock_worker/src/grpc.rs b/crates/mock_worker/src/grpc.rs index 362ee1e221..b900ca430f 100644 --- a/crates/mock_worker/src/grpc.rs +++ b/crates/mock_worker/src/grpc.rs @@ -141,6 +141,7 @@ impl TokenSpeedScheduler for MockScheduler { .and_then(|s| s.max_new_tokens) .unwrap_or(self.cfg.output_tokens); let stream_chunks = req.stream; + let weight_version = self.cfg.weight_version.clone(); let (tx, rx) = mpsc::unbounded_channel(); engine.submit(NewRequest { request_id: request_id.clone(), @@ -152,6 +153,7 @@ impl TokenSpeedScheduler for MockScheduler { rx, stream_chunks, request_id, + weight_version, ))); } @@ -173,7 +175,7 @@ impl TokenSpeedScheduler for MockScheduler { cached_tokens: 0, output_logprobs: None, index: 0, - weight_version: None, + weight_version: self.cfg.weight_version.clone(), })), })); } @@ -188,6 +190,7 @@ impl TokenSpeedScheduler for MockScheduler { output_logprobs: None, matched_stop: None, index: 0, + weight_version: self.cfg.weight_version.clone(), ..Default::default() })), })); @@ -242,8 +245,23 @@ impl TokenSpeedScheduler for MockScheduler { &self, _request: Request, ) -> Result, Status> { + let server_args = (!self.cfg.server_args.is_empty()).then(|| prost_types::Struct { + fields: self + .cfg + .server_args + .iter() + .map(|(k, v)| { + ( + k.clone(), + prost_types::Value { + kind: Some(prost_types::value::Kind::StringValue(v.clone())), + }, + ) + }) + .collect(), + }); Ok(Response::new(ts::GetServerInfoResponse { - server_args: None, + server_args, scheduler_info: None, active_requests: 0, is_paused: false, @@ -351,11 +369,18 @@ fn generate_stream( rx: mpsc::UnboundedReceiver, stream_chunks: bool, request_id: String, + weight_version: Option, ) -> GenStream { - let init = (rx, Vec::::new(), stream_chunks, request_id); + let init = ( + rx, + Vec::::new(), + stream_chunks, + request_id, + weight_version, + ); Box::pin(stream::unfold( init, - |(mut rx, mut output_ids, stream_chunks, request_id)| async move { + |(mut rx, mut output_ids, stream_chunks, request_id, weight_version)| async move { loop { match rx.recv().await { Some(engine::GenEvent::Token { @@ -374,10 +399,13 @@ fn generate_stream( cached_tokens, output_logprobs: None, index: 0, - weight_version: None, + weight_version: weight_version.clone(), })), }; - return Some((Ok(resp), (rx, output_ids, stream_chunks, request_id))); + return Some(( + Ok(resp), + (rx, output_ids, stream_chunks, request_id, weight_version), + )); } // Non-streaming: keep accumulating until Done. } @@ -398,10 +426,14 @@ fn generate_stream( output_logprobs: None, matched_stop: None, index: 0, + weight_version: weight_version.clone(), ..Default::default() })), }; - return Some((Ok(resp), (rx, output_ids, stream_chunks, request_id))); + return Some(( + Ok(resp), + (rx, output_ids, stream_chunks, request_id, weight_version), + )); } None => return None, } diff --git a/crates/protocols/src/rl.rs b/crates/protocols/src/rl.rs index e5275fdfd9..56828eef36 100644 --- a/crates/protocols/src/rl.rs +++ b/crates/protocols/src/rl.rs @@ -48,7 +48,7 @@ pub struct RlWorkerEntry { pub id: String, /// Registry URL; carries an `@` suffix for DP-aware workers. pub url: String, - /// The address control calls are sent to. + /// `Worker::base_url()`; see `control_url` for where control calls go. pub base_url: String, /// Engine name (`sglang`, `vllm`, ...); `unknown` when undetected. pub engine: String, @@ -56,6 +56,10 @@ pub struct RlWorkerEntry { pub model_id: String, pub worker_type: String, pub connection_mode: String, + /// Base URL of the worker's RL control routes: the worker itself for an + /// HTTP worker, the engine-advertised or operator-supplied + /// `rl.control_url` otherwise; `null` when the worker has none. + pub control_url: Option, pub tp_size: Option, pub dp_size: Option, pub pp_size: Option, @@ -110,13 +114,13 @@ pub struct RlFailedCall { pub worker_id: String, pub url: String, /// `upstream_error`, `upstream_unreachable`, `upstream_timeout`, or - /// `unsupported_connection_mode`. + /// `no_control_endpoint`. pub error: String, pub message: String, /// The engine's HTTP status, for `upstream_error`. #[serde(default, skip_serializing_if = "Option::is_none")] pub status: Option, - /// The worker's connection mode, for `unsupported_connection_mode`. + /// The worker's connection mode, for `no_control_endpoint`. #[serde(default, skip_serializing_if = "Option::is_none")] pub connection_mode: Option, } @@ -150,6 +154,7 @@ mod tests { model_id: "m".to_string(), worker_type: "regular".to_string(), connection_mode: "http".to_string(), + control_url: Some("http://a:1".to_string()), tp_size: Some(1), dp_size: None, pp_size: None, diff --git a/crates/protocols/src/worker.rs b/crates/protocols/src/worker.rs index 488cfc7bba..416bb79d11 100644 --- a/crates/protocols/src/worker.rs +++ b/crates/protocols/src/worker.rs @@ -1108,7 +1108,9 @@ impl HealthCheckUpdate { /// Per-worker HTTP connection configuration. /// All fields optional — `None` means "use router/global default". #[serde_with::skip_serializing_none] -#[derive(Debug, Clone, Default, Serialize, Deserialize, schemars::JsonSchema)] +#[derive( + Debug, Clone, Default, PartialEq, Eq, Hash, Serialize, Deserialize, schemars::JsonSchema, +)] pub struct HttpPoolConfig { /// Max idle connections per host (default: 500). pub pool_max_idle_per_host: Option, diff --git a/crates/rl/COUPLING.md b/crates/rl/COUPLING.md index 5d8c7d0d5b..fc6d24b050 100644 --- a/crates/rl/COUPLING.md +++ b/crates/rl/COUPLING.md @@ -5,7 +5,7 @@ that adds a surface must update this file. | # | Surface | model_gateway file | Notes | |---|---|---|---| -| (a) | `RlWorkerView` read-only registry view | `src/rl_adapter.rs`, `src/lib.rs` | `RegistryRlView` over `WorkerRegistry::{get_all,get,get_id_by_url}`; hands the RL crate each HTTP worker's negotiated client through `Worker::{http_client_handle_if_initialized,http_client}`, the same client the gateway's admin ops use, so control calls inherit the worker's HTTP version, TLS identity and roots, and pool tuning; `lib.rs` gains `pub mod rl_adapter;` | +| (a) | `RlWorkerView` read-only registry view | `src/rl_adapter.rs`, `src/lib.rs` | `RegistryRlView` over `WorkerRegistry::{get_all,get,get_id_by_url}`; hands the RL crate each HTTP worker's negotiated client through `Worker::{http_client_handle_if_initialized,http_client}`, the same client the gateway's admin ops use, so control calls inherit the worker's HTTP version, TLS identity and roots, and pool tuning; `lib.rs` gains `pub mod rl_adapter;`; for gRPC/ZMQ workers it borrows a control client from `WorkerHttpClientCache::get(&spec.http_pool, false)` (`AppContext.worker_client_cache`) and reads the `rl.control_url` label (`CONTROL_URL_LABEL`), resolving wildcard hosts with `smg_rl::resolve_control_url` | | (d) | `AppContext.rl: Option>` | `src/app_context.rs` | built in `AppContextBuilder::build()` when `router_config.rl.enabled` | | (d) | route mount | `src/server.rs` `build_app` | `nest("/v1/rl", smg_rl::router(..))` under `apply_control_plane_auth` | | (d) | metrics HELP registration | `src/observability/metrics.rs` | `smg_rl::init_rl_metrics()` | @@ -22,6 +22,20 @@ because the struct grew a field. The gateway-level test relies on `TestRouterConfig` disabling health checks, so the mock stopped mid-test stays registered and the fan-out still targets it. +Gateway-side changes that carry RL data but are not crate couplings (the RL +crate does not call into them; they exist so the data plane reports what the +control plane changed): `src/routers/grpc/client.rs` lifts `rl.control_url` +and the `rl.*` capability keys from TokenSpeed gRPC server info into worker +labels (`TOKENSPEED_GRPC_KEYS`); `src/routers/grpc/proto_wrapper.rs`, +`src/routers/grpc/common/response_formatting.rs` (`effective_weight_version`), +`src/routers/grpc/regular/{processor,streaming}.rs` and +`src/routers/grpc/pipeline.rs` carry the engine-reported +`meta_info.weight_version`. +`model_gateway/tests/rl_tokenspeed_control_endpoint_test.rs` registers a +TokenSpeed gRPC mock through `POST /workers`, drives it through `/v1/rl` via +the control endpoint its server info advertised, and checks the version the +gRPC `/generate` path reports. + Wire types are not a gateway coupling: they live in `crates/protocols/src/rl.rs` (`openai_protocol::rl`) next to the `/workers` types, and `clients/openapi-gen/src/main.rs` registers the `/v1/rl/*` paths. diff --git a/crates/rl/NOTES.md b/crates/rl/NOTES.md index 62caeb4c02..74804e719f 100644 --- a/crates/rl/NOTES.md +++ b/crates/rl/NOTES.md @@ -7,3 +7,10 @@ planning docs, with date, engine version, and what was done. |---|---|---|---|---| | 2026-09-03 | SGLang 0.5.15.post1 | `/get_weight_version` | 404; version is in `/model_info` | discovery reads the `weight_version` registration label | | 2026-09-03 | SGLang 0.5.15.post1 | `POST /pause_generation` with no body succeeds | 400 Bad Request; FastAPI requires a JSON body, `{}` is enough | callers must send `{}` on bodyless control routes; `examples/rl/refit_from_disk.py` needs `rl.fanout("pause_generation", {}, ...)` (the e2e test already passes `json={}`) | +| 2026-09-21 | TokenSpeed d1464d8f | `flush_cache` on GET and POST like SGLang | GET only on the in-engine app | TokenSpeed PR adds POST; older builds need `method="GET"` | +| 2026-09-21 | TokenSpeed d1464d8f | bodyless POST → 400 | six control routes answer 500 on an empty/malformed body | TokenSpeed PR moves body parsing inside the guard; older builds: always send a JSON object | +| 2026-09-21 | TokenSpeed d1464d8f | `pause_generation {"mode": …}` | mode ignored, always `wait`; `keep` unreachable | TokenSpeed PR honors `mode`; static row advertises `wait,abort` only | +| 2026-09-21 | TokenSpeed d1464d8f | `update_weights_from_tensor` end to end | route exists, CUDA-IPC receive path not implemented | engine advertises `rl.update_from=distributed` (newer builds add `mooncake`); the static row says `distributed` | +| 2026-09-22 | TokenSpeed d1464d8f | `update_weights_from_disk`/`_tensor` end to end | the routes exist but the scheduler raises `NotImplementedError` and dies (engine down); the companion branch now answers 501 and advertises `distributed` only | static row lists `distributed`; refit TokenSpeed with `update_weights_from_distributed` (see `examples/rl` once the trainer-side example lands) | +| 2026-10-06 | TokenSpeed 0.1.0.post20261006 (`smg serve --connection-mode zmq`) | the in-engine control app serves a ZMQ worker too | headless engines (`launch_scheduler_headless`) build no `AsyncLLM`, and only `AsyncLLM` starts the control app, so `--rl-control-host`/`--rl-control-port` are inert there | RL control on TokenSpeed is gRPC-only; `smg serve` wires no control endpoint for ZMQ workers (an operator may still label `rl.control_url` for an app they run) | +| 2026-10-06 | TokenSpeed 0.1.0.post20261006 | `attn_tp_size` in server args, folded into `tp_size` by discovery | parallelism nested under `mapping.*`; no top-level `tp_size`/`attn_tp_size`, so discovery reports `tp_size: null` | main-side discovery to read `mapping.*`; trainers pass the width explicitly until then | diff --git a/crates/rl/README.md b/crates/rl/README.md index 9a7641275a..f15942e5e5 100644 --- a/crates/rl/README.md +++ b/crates/rl/README.md @@ -14,13 +14,31 @@ Response bodies are the `openai_protocol::rl` types (`RlWorkersResponse`, `clients/openapi-gen` so the generated SDKs carry them. `GET /v1/rl/workers` reports `protocol_version` (currently 1), bumped only for incompatible changes. -Only HTTP workers can be proxied. A gRPC or ZMQ worker that matches a selector is -reported in `failed[]` as `unsupported_connection_mode` (HTTP 422 on the per-worker -route), so a fan-out over a mixed fleet answers 207 and `smg.rl.RL.fanout` raises -`FanoutError` unless `allow_partial=True`. +## Control endpoints + +Control calls go to a worker's **control endpoint**, not necessarily its data +transport. An HTTP worker is controlled through itself; an `rl.control_url` +label on an HTTP worker is ignored. A gRPC or ZMQ worker needs the +`rl.control_url` label: TokenSpeed engines advertise it in server +info (SMG's discovery turns it into the label), and any worker can be given one +at `POST /workers` or through the worker update route. The engine advertises +the URL only for a concrete `--rl-control-host`: a wildcard bind (`0.0.0.0`, +`::`) is not advertised, and the default loopback is right only when SMG runs +on the same machine. A wildcard host in an operator-supplied label is replaced +by the worker's own host. The +worker's `api_key` is sent as the bearer to the control endpoint. The label is +trusted: whatever host `rl.control_url` names receives the worker's `api_key` +as a bearer, so only engines and operators you trust may set it. A worker with +no control endpoint is reported in `failed[]` as `no_control_endpoint` (HTTP +422 on the per-worker route), so a fan-out over such a fleet answers 207 and +`smg.rl.RL.fanout` raises `FanoutError` unless `allow_partial=True`. + +Capabilities come from the `rl.*` labels when an engine (or operator) supplies +them (`"source": "label"`), else from the built-in table (`"source": "static"`). Flags: `--enable-rl`, `--rl-control-timeout-secs` (600), `--rl-fanout-concurrency` (32). Recommended RL launch profile: `--enable-rl --disable-health-check --disable-circuit-breaker --request-timeout-secs 14400`. +TokenSpeed rollout engines: `python3 -m smg_grpc_servicer.tokenspeed --model … --port 30000 --rl-control-host --rl-control-port 30400 [--rl-control-api-key …]`, registered as `grpc://host:30000`. TokenSpeed refits are trainer-driven over NCCL (`init_weights_update_group` → `update_weights_from_distributed` → `destroy_weights_update_group`, each proxied per worker through `/v1/rl`); `update_weights_from_disk` and `update_weights_from_tensor` answer HTTP 501 on TokenSpeed. See `docs/guides/rl-tokenspeed.md`. ## Python client diff --git a/crates/rl/src/capability.rs b/crates/rl/src/capability.rs index 93c2893404..383e6e8df7 100644 --- a/crates/rl/src/capability.rs +++ b/crates/rl/src/capability.rs @@ -31,6 +31,15 @@ fn static_for(runtime: RuntimeType) -> RlCapabilities { sleep_wake: true, reports_weight_version: false, }, + RuntimeType::TokenSpeed => RlCapabilities { + source: RlCapabilitySource::Static, + pause_modes: strings(&["wait", "abort"]), + update_from: strings(&["distributed"]), + abort: true, + flush_cache: true, + sleep_wake: true, + reports_weight_version: true, + }, _ => RlCapabilities { source: RlCapabilitySource::Static, pause_modes: Vec::new(), @@ -124,7 +133,6 @@ mod tests { fn other_runtimes_have_no_capabilities() { for rt in [ RuntimeType::Trtllm, - RuntimeType::TokenSpeed, RuntimeType::Mlx, RuntimeType::Generic, RuntimeType::External, @@ -167,4 +175,32 @@ mod tests { assert!(c.flush_cache); assert_eq!(c.pause_modes, ["abort", "retract", "in_place"]); } + + /// The fallback for TokenSpeed builds that predate advertisement. `keep` + /// is absent on purpose: pre-advertisement builds cannot reach it over + /// HTTP, and an advertising engine overrides this row anyway. `disk` and + /// `tensor` are absent too: the scheduler has no receive path for + /// either, only the trainer-driven NCCL broadcast (`distributed`) works. + #[test] + fn tokenspeed_static_row() { + let t = capabilities_for(RuntimeType::TokenSpeed, &HashMap::new()); + assert_eq!(t.source, RlCapabilitySource::Static); + assert_eq!(t.pause_modes, ["wait", "abort"]); + assert_eq!(t.update_from, ["distributed"]); + assert!(t.abort && t.flush_cache && t.sleep_wake && t.reports_weight_version); + } + + #[test] + fn engine_advertised_labels_override_the_tokenspeed_row() { + let t = capabilities_for( + RuntimeType::TokenSpeed, + &labels(&[ + ("rl.pause_modes", "wait,abort,keep"), + ("rl.update_from", "distributed"), + ]), + ); + assert_eq!(t.source, RlCapabilitySource::Label); + assert_eq!(t.pause_modes, ["wait", "abort", "keep"]); + assert_eq!(t.update_from, ["distributed"]); + } } diff --git a/crates/rl/src/control.rs b/crates/rl/src/control.rs new file mode 100644 index 0000000000..4aaba2977d --- /dev/null +++ b/crates/rl/src/control.rs @@ -0,0 +1,112 @@ +//! Where a worker's RL control routes live. For an HTTP worker that is the +//! worker itself; for a gRPC or ZMQ worker it is the URL the engine advertised +//! (label `rl.control_url`), with a wildcard bind host swapped for the host the +//! gateway already reaches the worker on. + +/// Resolve an advertised control URL against the worker's own base URL. +/// +/// An engine does not advertise a wildcard bind, but an operator may label +/// one (`http://0.0.0.0:P`); the only host the gateway knows is reachable is +/// the worker's own, so use it. A worker with no host of its own (an `ipc://` +/// ZMQ worker) leaves the label as it is. +pub fn resolve_control_url(advertised: &str, worker_url: &str) -> String { + let advertised = advertised.trim().trim_end_matches('/'); + let Some((scheme, rest)) = advertised.split_once("://") else { + return advertised.to_string(); + }; + let (authority, path) = rest.split_once('/').map_or((rest, ""), |(a, p)| (a, p)); + let (host, port) = split_host_port(authority); + if !is_wildcard_host(host) { + return advertised.to_string(); + } + let worker_host = worker_url + .split_once("://") + .map_or(worker_url, |(_, r)| r) + .split('/') + .next() + .map(|authority| split_host_port(authority).0) + .unwrap_or(""); + if worker_host.is_empty() { + return advertised.to_string(); + } + let mut out = format!("{scheme}://{worker_host}"); + if let Some(port) = port { + out.push(':'); + out.push_str(port); + } + if !path.is_empty() { + out.push('/'); + out.push_str(path); + } + out +} + +fn is_wildcard_host(host: &str) -> bool { + matches!(host, "" | "0.0.0.0" | "::" | "[::]") +} + +/// `host:port` → (`host`, `Some(port)`), with IPv6 brackets kept on the host. +fn split_host_port(authority: &str) -> (&str, Option<&str>) { + if let Some(end) = authority + .strip_prefix('[') + .and_then(|_| authority.find(']')) + { + let host = &authority[..=end]; + let port = authority[end + 1..].strip_prefix(':'); + return (host, port); + } + match authority.rsplit_once(':') { + Some((host, port)) if !host.contains(':') => (host, Some(port)), + _ => (authority, None), + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn concrete_hosts_pass_through() { + assert_eq!( + resolve_control_url("http://10.0.0.5:40100", "grpc://10.0.0.5:30000"), + "http://10.0.0.5:40100" + ); + assert_eq!( + resolve_control_url("http://ctl.internal:8001/", "grpc://a:1"), + "http://ctl.internal:8001" + ); + } + + #[test] + fn wildcard_hosts_take_the_worker_host() { + assert_eq!( + resolve_control_url("http://0.0.0.0:40100", "grpc://10.0.0.5:30000"), + "http://10.0.0.5:40100" + ); + assert_eq!( + resolve_control_url("http://[::]:40100", "grpc://[fd00::5]:30000"), + "http://[fd00::5]:40100" + ); + assert_eq!( + resolve_control_url("http://:40100", "grpc://10.0.0.5:30000"), + "http://10.0.0.5:40100" + ); + } + + #[test] + fn a_worker_without_a_host_keeps_the_label_as_advertised() { + assert_eq!( + resolve_control_url("http://0.0.0.0:40100", "ipc:///tmp/engine.sock"), + "http://0.0.0.0:40100", + "there is no worker host to substitute" + ); + } + + #[test] + fn schemeless_or_odd_input_is_returned_unchanged() { + assert_eq!( + resolve_control_url("10.0.0.5:40100", "grpc://a:1"), + "10.0.0.5:40100" + ); + } +} diff --git a/crates/rl/src/discovery.rs b/crates/rl/src/discovery.rs index 63fa66b611..b912fa6572 100644 --- a/crates/rl/src/discovery.rs +++ b/crates/rl/src/discovery.rs @@ -108,6 +108,7 @@ pub fn entry(info: &RlWorkerInfo, dp_ranks: usize) -> RlWorkerEntry { model_id: info.model_id.clone(), worker_type: enum_str(&info.worker_type), connection_mode: enum_str(&info.connection_mode), + control_url: info.control_url.clone(), tp_size: int_label(&info.labels, "tp_size"), dp_size: int_label(&info.labels, "dp_size"), pp_size: int_label(&info.labels, "pp_size"), @@ -166,7 +167,8 @@ mod tests { use axum::{body::Body, http::Request}; use http_body_util::BodyExt; - use openai_protocol::worker::RuntimeType; + use openai_protocol::worker::{ConnectionMode, RuntimeType}; + use serde_json::Value; use tower::ServiceExt; use super::*; @@ -183,6 +185,10 @@ mod tests { )) } + async fn json_body(resp: Response) -> Value { + serde_json::from_slice(&resp.into_body().collect().await.unwrap().to_bytes()).unwrap() + } + #[test] fn collapse_groups_dp_ranks_and_sorts() { let ws = vec![ @@ -221,7 +227,7 @@ mod tests { .await .unwrap(); assert_eq!(resp.status(), StatusCode::OK); - let body: serde_json::Value = + let body: Value = serde_json::from_slice(&resp.into_body().collect().await.unwrap().to_bytes()).unwrap(); assert_eq!(body["dp_ranks"], 2); @@ -230,7 +236,7 @@ mod tests { .await .unwrap(); assert_eq!(resp.status(), StatusCode::OK); - let body: serde_json::Value = + let body: Value = serde_json::from_slice(&resp.into_body().collect().await.unwrap().to_bytes()).unwrap(); assert_eq!(body["dp_ranks"], 1); } @@ -258,7 +264,7 @@ mod tests { .oneshot(Request::get("/workers").body(Body::empty()).unwrap()) .await .unwrap(); - let body: serde_json::Value = + let body: Value = serde_json::from_slice(&resp.into_body().collect().await.unwrap().to_bytes()).unwrap(); assert_eq!(body["protocol_version"], 1); } @@ -276,18 +282,14 @@ mod tests { .await .unwrap(); assert_eq!(resp.status(), StatusCode::OK); - let body: serde_json::Value = + let body: Value = serde_json::from_slice(&resp.into_body().collect().await.unwrap().to_bytes()).unwrap(); assert_eq!(body["total"], 1); let e = &body["workers"][0]; assert_eq!(e["id"], "w1"); assert_eq!(e["engine"], "sglang"); assert_eq!(e["engine_version"], "0.5.15"); - assert_eq!( - e["tp_size"], - serde_json::Value::Null, - "garbage label -> null" - ); + assert_eq!(e["tp_size"], Value::Null, "garbage label -> null"); assert_eq!(e["dp_size"], 1); assert_eq!(e["dp_ranks"], 1); assert_eq!(e["health"], "ready"); @@ -308,4 +310,32 @@ mod tests { .unwrap(); assert_eq!(resp.status(), StatusCode::NOT_FOUND); } + + #[tokio::test] + async fn entry_reports_the_control_url() { + let mut g = worker("g1", "grpc://a:1", RuntimeType::TokenSpeed); + g.connection_mode = ConnectionMode::Grpc; + g.control_url = Some("http://a:40100".to_string()); + let mut n = worker("n1", "grpc://b:1", RuntimeType::TokenSpeed); + n.connection_mode = ConnectionMode::Grpc; + n.control_url = None; + let app = crate::router::<()>(state(vec![g, n])); + let resp = app + .oneshot(Request::get("/workers").body(Body::empty()).unwrap()) + .await + .unwrap(); + let body = json_body(resp).await; + let by_id = |id: &str| { + body["workers"] + .as_array() + .unwrap() + .iter() + .find(|w| w["id"] == id) + .cloned() + .unwrap() + }; + assert_eq!(by_id("g1")["control_url"], "http://a:40100"); + assert_eq!(by_id("n1")["control_url"], Value::Null); + assert_eq!(by_id("g1")["engine"], "tokenspeed"); + } } diff --git a/crates/rl/src/error.rs b/crates/rl/src/error.rs index 5b07ed45f4..122f8c6891 100644 --- a/crates/rl/src/error.rs +++ b/crates/rl/src/error.rs @@ -21,11 +21,14 @@ pub enum RlError { NoWorkersMatch(String), #[error("worker `{0}` not found")] WorkerNotFound(String), - #[error("worker `{worker_id}` uses connection mode `{mode}`, which cannot be proxied")] - UnsupportedConnectionMode { + #[error( + "worker `{worker_id}` has no control endpoint ({connection_mode} transport): {message}" + )] + NoControlEndpoint { worker_id: String, url: String, - mode: String, + connection_mode: String, + message: String, }, #[error("upstream `{url}` unreachable: {message}")] UpstreamUnreachable { @@ -50,7 +53,7 @@ impl RlError { Self::InvalidSelector { .. } => "invalid_selector", Self::NoWorkersMatch(_) => "no_workers_match", Self::WorkerNotFound(_) => "worker_not_found", - Self::UnsupportedConnectionMode { .. } => "unsupported_connection_mode", + Self::NoControlEndpoint { .. } => "no_control_endpoint", Self::UpstreamUnreachable { .. } => "upstream_unreachable", Self::UpstreamTimeout { .. } => "upstream_timeout", } @@ -63,7 +66,7 @@ impl RlError { | Self::InvalidSelector { .. } | Self::NoWorkersMatch(_) => StatusCode::BAD_REQUEST, Self::WorkerNotFound(_) => StatusCode::NOT_FOUND, - Self::UnsupportedConnectionMode { .. } => StatusCode::UNPROCESSABLE_ENTITY, + Self::NoControlEndpoint { .. } => StatusCode::UNPROCESSABLE_ENTITY, Self::UpstreamUnreachable { .. } => StatusCode::BAD_GATEWAY, Self::UpstreamTimeout { .. } => StatusCode::GATEWAY_TIMEOUT, } @@ -76,14 +79,15 @@ impl RlError { Self::InvalidSelector { offset, .. } => body["offset"] = json!(offset), Self::NoWorkersMatch(selector) => body["selector"] = json!(selector), Self::WorkerNotFound(id) => body["id"] = json!(id), - Self::UnsupportedConnectionMode { + Self::NoControlEndpoint { worker_id, url, - mode, + connection_mode, + .. } => { body["worker_id"] = json!(worker_id); body["url"] = json!(url); - body["connection_mode"] = json!(mode); + body["connection_mode"] = json!(connection_mode); } Self::UpstreamUnreachable { worker_id, url, .. } | Self::UpstreamTimeout { worker_id, url, .. } => { @@ -127,5 +131,16 @@ mod tests { }; assert_eq!(e.status(), StatusCode::GATEWAY_TIMEOUT); assert_eq!(e.to_json()["url"], "http://x"); + + let e = RlError::NoControlEndpoint { + worker_id: "g".into(), + url: "grpc://x:1".into(), + connection_mode: "grpc".into(), + message: "no `rl.control_url` label".into(), + }; + assert_eq!(e.status(), StatusCode::UNPROCESSABLE_ENTITY); + assert_eq!(e.to_json()["error"], "no_control_endpoint"); + assert_eq!(e.to_json()["connection_mode"], "grpc"); + assert!(e.to_string().contains("no `rl.control_url` label")); } } diff --git a/crates/rl/src/fanout.rs b/crates/rl/src/fanout.rs index 0ba22557f5..c836571c3d 100644 --- a/crates/rl/src/fanout.rs +++ b/crates/rl/src/fanout.rs @@ -88,7 +88,9 @@ pub async fn run_fanout( } Err(e) => { let connection_mode = match &e { - RlError::UnsupportedConnectionMode { mode, .. } => Some(mode.clone()), + RlError::NoControlEndpoint { + connection_mode, .. + } => Some(connection_mode.clone()), _ => None, }; report.failed.push(RlFailedCall { @@ -225,6 +227,7 @@ mod tests { let bad = FakeEngine::start(StatusCode::INTERNAL_SERVER_ERROR, json!({"e": 1}), 0).await; let mut grpc = worker("g1", &format!("{}/grpc", good.url), RuntimeType::Sglang); grpc.connection_mode = ConnectionMode::Grpc; + grpc.control_url = None; let app = crate::router::<()>(state( vec![ worker("w1", &good.url, RuntimeType::Sglang), @@ -261,7 +264,8 @@ mod tests { assert_eq!(by_id("w2")["error"], "upstream_error"); assert_eq!(by_id("w2")["status"], 500); assert_eq!(by_id("w3")["error"], "upstream_unreachable"); - assert_eq!(by_id("g1")["error"], "unsupported_connection_mode"); + assert_eq!(by_id("g1")["error"], "no_control_endpoint"); + assert_eq!(by_id("g1")["connection_mode"], "grpc"); assert_eq!( body["results"]["w2"]["status"], 500, "failed also in results" diff --git a/crates/rl/src/lib.rs b/crates/rl/src/lib.rs index a4b1895e53..12f8a00e1b 100644 --- a/crates/rl/src/lib.rs +++ b/crates/rl/src/lib.rs @@ -4,6 +4,7 @@ pub mod capability; pub mod config; +mod control; pub mod discovery; pub mod error; pub mod fanout; @@ -20,6 +21,7 @@ use std::sync::Arc; use axum::{routing::get, Router}; pub use config::RlConfig; +pub use control::resolve_control_url; pub use error::RlError; pub use metrics::init_rl_metrics; pub use state::RlState; diff --git a/crates/rl/src/proxy.rs b/crates/rl/src/proxy.rs index 898e118237..1405d1bbac 100644 --- a/crates/rl/src/proxy.rs +++ b/crates/rl/src/proxy.rs @@ -15,7 +15,7 @@ use axum::{ response::{IntoResponse, Response}, Json, }; -use openai_protocol::{rl::RlCallOutcome, worker::ConnectionMode}; +use openai_protocol::rl::RlCallOutcome; use serde_json::Value; use tracing::info; @@ -76,8 +76,8 @@ impl ProxyRequest { }) } - fn url_for(&self, worker: &RlWorkerInfo) -> String { - let base = worker.base_url.trim_end_matches('/'); + fn url_for(&self, base: &str) -> String { + let base = base.trim_end_matches('/'); match &self.query { Some(q) => format!("{base}/{}?{q}", self.path), None => format!("{base}/{}", self.path), @@ -117,6 +117,16 @@ pub(crate) fn parse_body( ) } +/// The worker has nowhere to send control calls, or nothing to send them with. +fn no_control_endpoint(worker: &RlWorkerInfo, message: &str) -> RlError { + RlError::NoControlEndpoint { + worker_id: worker.id.clone(), + url: worker.url.clone(), + connection_mode: enum_str(&worker.connection_mode), + message: message.to_string(), + } +} + /// Send `req` to `worker`. `Err` only for transport failures; an upstream /// 4xx/5xx is a successful proxy with that status in the outcome. pub async fn call_worker( @@ -124,18 +134,21 @@ pub async fn call_worker( worker: &RlWorkerInfo, req: &ProxyRequest, ) -> Result { - // A worker the gateway never speaks HTTP to has no client to borrow. - let client = match &worker.http_client { - Some(client) if worker.connection_mode == ConnectionMode::Http => client, - _ => { - return Err(RlError::UnsupportedConnectionMode { - worker_id: worker.id.clone(), - url: worker.url.clone(), - mode: enum_str(&worker.connection_mode), - }) - } + // Control is independent of the data transport: what matters is that the + // gateway knows an HTTP control endpoint and has a client to reach it. + let Some(base) = worker.control_url.as_deref() else { + return Err(no_control_endpoint( + worker, + "no `rl.control_url` label: set it at registration or through the worker update route (a TokenSpeed gRPC engine advertises it when launched with a routable --rl-control-host)", + )); }; - let url = req.url_for(worker); + let client = worker.control_client.as_ref().map_err(|e| { + no_control_endpoint( + worker, + &format!("no HTTP client for the control endpoint: {e}"), + ) + })?; + let url = req.url_for(base); // The worker's client carries the gateway's request timeout; a refit // can outlive it, so the control deadline is set per request, as the // gateway's own flush/profile admin calls do. @@ -231,8 +244,9 @@ pub async fn call_worker( ); info!( target: "smg_rl", - worker_id = %worker.id, url = %worker.url, method = %req.method, - path = %req.path, status, latency_ms = outcome.latency_ms, "rl.proxy" + worker_id = %worker.id, url = %worker.url, control_url = %base, + method = %req.method, path = %req.path, status, + latency_ms = outcome.latency_ms, "rl.proxy" ); Ok(outcome) } @@ -403,6 +417,7 @@ mod tests { let engine = FakeEngine::start(StatusCode::OK, json!({}), 0).await; let mut grpc = worker("g1", &engine.url, RuntimeType::Sglang); grpc.connection_mode = ConnectionMode::Grpc; + grpc.control_url = None; let app = crate::router::<()>(state( vec![worker("w1", &engine.url, RuntimeType::Sglang), grpc], 5, @@ -430,7 +445,7 @@ mod tests { .await .unwrap(); assert_eq!(r.status(), StatusCode::UNPROCESSABLE_ENTITY); - assert_eq!(json_body(r).await["error"], "unsupported_connection_mode"); + assert_eq!(json_body(r).await["error"], "no_control_endpoint"); let r = app .clone() @@ -456,13 +471,63 @@ mod tests { assert!(engine.seen().is_empty()); } + /// The transport SMG uses for data is irrelevant to control: a gRPC + /// worker with a control endpoint is proxied exactly like an HTTP one. + #[tokio::test] + async fn grpc_worker_with_a_control_endpoint_is_proxied() { + let control_app = FakeEngine::start(StatusCode::OK, json!({"success": true}), 0).await; + let mut w = worker("g1", "grpc://engine:30000", RuntimeType::TokenSpeed); + w.connection_mode = ConnectionMode::Grpc; + w.control_url = Some(control_app.url.clone()); + w.api_key = Some("ts-secret".to_string()); + let app = crate::router::<()>(state(vec![w], 5)); + + let resp = app + .oneshot( + Request::post("/workers/g1/engine/pause_generation") + .header("content-type", "application/json") + .body(Body::from(r#"{"mode":"wait"}"#)) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(resp.status(), StatusCode::OK); + let seen = control_app.seen(); + assert_eq!(seen.len(), 1); + assert_eq!(seen[0].path, "/pause_generation"); + assert_eq!(&seen[0].body[..], br#"{"mode":"wait"}"#); + assert_eq!(seen[0].headers["authorization"], "Bearer ts-secret"); + } + + #[tokio::test] + async fn control_url_without_a_client_is_reported_not_routed() { + let mut w = worker("g2", "grpc://engine:30000", RuntimeType::TokenSpeed); + w.connection_mode = ConnectionMode::Grpc; + w.control_url = Some("http://ctl:1".to_string()); + w.control_client = Err("bad CA bundle".to_string()); + let app = crate::router::<()>(state(vec![w], 5)); + let r = app + .oneshot( + Request::post("/workers/g2/engine/pause") + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(r.status(), StatusCode::UNPROCESSABLE_ENTITY); + let body = json_body(r).await; + assert_eq!(body["error"], "no_control_endpoint"); + let message = body["message"].as_str().unwrap(); + assert!(message.contains("bad CA bundle"), "{message}"); + } + /// The control deadline is applied per request, so it holds even when /// the worker's client carries no client-level timeout of its own. #[tokio::test] async fn control_timeout_is_applied_per_request_on_the_worker_client() { let slow_engine = FakeEngine::start(StatusCode::OK, json!({}), 1500).await; let mut w = worker("slow", &slow_engine.url, RuntimeType::Sglang); - w.http_client = Some(Arc::new(reqwest::Client::new())); + w.control_client = Ok(Arc::new(reqwest::Client::new())); let app = crate::router::<()>(state(vec![w], 1)); let r = app diff --git a/crates/rl/src/state.rs b/crates/rl/src/state.rs index 6789e4db40..a1f66f2f2c 100644 --- a/crates/rl/src/state.rs +++ b/crates/rl/src/state.rs @@ -12,11 +12,11 @@ pub struct RlState { impl RlState { /// Build the state. The control plane owns no HTTP client: every call - /// goes through the client the gateway negotiated for the target worker - /// (see [`crate::RlWorkerInfo::http_client`]), with the control deadline - /// applied per request. Breakers and load counters live on the gateway's - /// worker objects, not on the client, so control calls leave them - /// untouched. + /// goes through the client the gateway resolved for the target worker's + /// control endpoint (see [`crate::RlWorkerInfo::control_client`]), with + /// the control deadline applied per request. Breakers and load counters + /// live on the gateway's worker objects, not on the client, so control + /// calls leave them untouched. pub fn new(view: Arc, config: RlConfig) -> Self { Self { view, config } } diff --git a/crates/rl/src/testing.rs b/crates/rl/src/testing.rs index 31e43d56ce..ac6f038828 100644 --- a/crates/rl/src/testing.rs +++ b/crates/rl/src/testing.rs @@ -37,10 +37,11 @@ pub fn worker(id: &str, url: &str, runtime: RuntimeType) -> RlWorkerInfo { labels.insert("dp_size".to_string(), "1".to_string()); labels.insert("pp_size".to_string(), "1".to_string()); labels.insert("weight_version".to_string(), "default".to_string()); + let base_url = url.split('@').next().unwrap_or(url).to_string(); RlWorkerInfo { id: id.to_string(), url: url.to_string(), - base_url: url.split('@').next().unwrap_or(url).to_string(), + base_url: base_url.clone(), api_key: None, model_id: "mock-model".to_string(), runtime, @@ -50,7 +51,8 @@ pub fn worker(id: &str, url: &str, runtime: RuntimeType) -> RlWorkerInfo { is_dp_aware: url.contains('@'), dp_size: None, labels, - http_client: Some(test_client()), + control_url: Some(base_url), + control_client: Ok(test_client()), } } diff --git a/crates/rl/src/view.rs b/crates/rl/src/view.rs index 1b08fe3add..fae825f745 100644 --- a/crates/rl/src/view.rs +++ b/crates/rl/src/view.rs @@ -13,7 +13,7 @@ pub struct RlWorkerInfo { pub id: String, /// `Worker::url()`; carries an `@` suffix for DP-aware workers. pub url: String, - /// `Worker::base_url()`; the address control calls are sent to. + /// `Worker::base_url()`; see `control_url` for where control calls go. pub base_url: String, pub api_key: Option, pub model_id: String, @@ -25,10 +25,17 @@ pub struct RlWorkerInfo { pub dp_size: Option, /// `WorkerSpec.labels`: discovered metadata merged with caller labels. pub labels: HashMap, - /// The client the gateway negotiated for this worker (HTTP version, - /// TLS identity and roots, pool tuning), shared with its data-plane and - /// admin calls. `None` when the gateway does not speak HTTP to it. - pub http_client: Option>, + /// Base URL of the worker's RL control routes: the worker itself for an + /// HTTP worker, the engine-advertised `rl.control_url` (wildcard host + /// resolved) for a gRPC or ZMQ worker, `None` when it has neither. + pub control_url: Option, + /// The client for control calls. For an HTTP worker it is the client the + /// gateway negotiated for the worker (HTTP version, TLS identity and + /// roots, pool tuning); for other transports it is a client with the same + /// TLS settings on HTTP/1.1, shared by the workers with the same pool + /// config. `Err` says why the gateway has none; the 422 for the worker + /// repeats it. + pub control_client: Result, String>, } impl fmt::Debug for RlWorkerInfo { @@ -46,7 +53,11 @@ impl fmt::Debug for RlWorkerInfo { .field("is_dp_aware", &self.is_dp_aware) .field("dp_size", &self.dp_size) .field("labels", &self.labels) - .field("http_client", &self.http_client.as_ref().map(|_| "..")) + .field("control_url", &self.control_url) + .field( + "control_client", + &self.control_client.as_ref().map(|_| ".."), + ) .finish() } } @@ -63,9 +74,8 @@ pub trait RlWorkerView: Send + Sync { mod tests { use super::*; - #[test] - fn debug_output_redacts_the_api_key() { - let info = RlWorkerInfo { + fn sample() -> RlWorkerInfo { + RlWorkerInfo { id: "w".to_string(), url: "http://a:1".to_string(), base_url: "http://a:1".to_string(), @@ -78,11 +88,24 @@ mod tests { is_dp_aware: false, dp_size: None, labels: HashMap::new(), - http_client: None, - }; + control_url: None, + control_client: Err("no client".to_string()), + } + } + + #[test] + fn debug_output_redacts_the_api_key() { + let info = sample(); let dbg = format!("{info:?}"); assert!(!dbg.contains("hunter2-secret"), "{dbg}"); assert!(dbg.contains("api_key: Some(\"\")"), "{dbg}"); assert!(dbg.contains("id: \"w\""), "{dbg}"); } + + #[test] + fn debug_output_shows_the_control_url() { + let mut info = sample(); + info.control_url = Some("http://ctl:1".to_string()); + assert!(format!("{info:?}").contains("control_url: Some(\"http://ctl:1\")")); + } } diff --git a/docs/guides/rl-tokenspeed.md b/docs/guides/rl-tokenspeed.md new file mode 100644 index 0000000000..1f843d7482 --- /dev/null +++ b/docs/guides/rl-tokenspeed.md @@ -0,0 +1,91 @@ +# RL rollouts on TokenSpeed behind SMG + +SMG's RL control plane (`--enable-rl`, `/v1/rl/*`) drives TokenSpeed engines +through their in-engine control app (the SGLang-style routes slime speaks) +while the data plane stays on gRPC. + +## Launch + +Each engine, on its own GPUs: + + python3 -m smg_grpc_servicer.tokenspeed --model /ckpt/policy --host 0.0.0.0 --port 30000 \ + --rl-control-host 10.0.0.11 --rl-control-port 30400 --rl-control-api-key "$RL_KEY" \ + --enable-output-logprobs + +`--rl-control-host` is the address the gateway reaches the engine on. The +engine advertises `rl.control_url` only for a concrete host: a wildcard bind +(`0.0.0.0`, `::`) is not advertised, so the worker would have no control +endpoint and every control call would answer 422 `no_control_endpoint`. The +default, loopback, is right only when SMG runs on the same machine. + +The gateway, with the engines as startup workers (the router's connection +mode comes from these URLs, so they have to be the gRPC ones): + + smg launch --worker-urls grpc://rollout-1:30000 grpc://rollout-2:30000 \ + --policy cache_aware --enable-rl --disable-health-check --disable-circuit-breaker \ + --request-timeout-secs 14400 + +Startup workers register without a key, so give each one its control key +through the worker update route (the ids are in `GET /workers`); the proxy +sends it as the bearer to the control app: + + curl -X PATCH http://smg:30000/workers/ -H 'content-type: application/json' \ + -d '{"api_key":"'"$RL_KEY"'"}' + +The key travels in that request body, so reach the gateway's admin routes +over TLS (`--tls-cert-path`/`--tls-key-path` on `smg launch`) or over a +network you trust. + +A gateway started with `--enable-igw` also serves gRPC workers registered +later through `POST /workers`; there the key goes into the registration +body (`{"url":"grpc://rollout-1:30000","api_key":"..."}`). + +`GET /v1/rl/workers` then shows `engine: tokenspeed`, `connection_mode: grpc`, +`control_url: http://10.0.0.11:30400` and the engine's advertised capabilities. + +## Refit + +TokenSpeed's scheduler has no receive path for a disk or tensor refit: +`update_weights_from_disk` and `update_weights_from_tensor` both answer HTTP +501 (`{"success": false, "message": "... is not implemented by this build's +scheduler; supported sources: distributed"}`) and the engine keeps serving. +`examples/rl/refit_from_disk.py` is for SGLang only — do not point it at a +TokenSpeed selector. + +The only refit path TokenSpeed implements is the trainer-driven NCCL +broadcast (slime's path): the trainer calls, per worker, through `/v1/rl`, + + init_weights_update_group (once, to join the trainer's process group) + update_weights_from_distributed (once per refit) + destroy_weights_update_group (on teardown) + +with `pause_generation` / `continue_generation` fanned out around the +`update_weights_from_distributed` call, same as a disk refit. The next +`/generate` through SMG reports the new `meta_info.weight_version`, stamped +by the engine on the gRPC response. A worked trainer-side example will land +under `examples/rl` once available; until then, drive the three calls +directly with `smg.rl.RL.call`/`fanout` as shown in `crates/rl/README.md`. +Ranks come from each worker's `tp_size`, which discovery reads from the +engine's server args (TokenSpeed's own spelling, `attn_tp_size`, is folded +into it). The newest TokenSpeed builds nest their parallelism under +`mapping.*`, which discovery does not read yet, so such an engine reports +`tp_size: null` and the trainer must be told the width (or assume 1). + +## Security + +The control app accepts weight updates from anyone who can reach it, and a +remote gateway needs it on a routable host, so always set +`--rl-control-api-key` and give SMG the same key as the worker's `api_key`. + +## Older engines + +Engines that predate advertisement get SMG's static capability row (`wait` +and `abort`, `distributed` only) and need the `rl.control_url` label +supplied by the operator. The worker update route merges labels, so it +rides on the same PATCH as the key: + + curl -X PATCH http://smg:30000/workers/ -H 'content-type: application/json' \ + -d '{"api_key":"'"$RL_KEY"'","labels":{"rl.control_url":"http://rollout-1:30400"}}' + +(on an `--enable-igw` gateway, in the `POST /workers` body instead). See +`crates/rl/NOTES.md` for their route-level drift. diff --git a/examples/rl/refit_from_disk.py b/examples/rl/refit_from_disk.py index 07c4f9eec6..5a63226719 100755 --- a/examples/rl/refit_from_disk.py +++ b/examples/rl/refit_from_disk.py @@ -8,8 +8,8 @@ each as one fan-out, then one /generate through SMG to confirm the engine reports the new meta_info.weight_version. The pause/resume pair is `smg.rl.paused`, so a failure at any stage still resumes the engines that did -pause. Only HTTP workers can be proxied: a gRPC or ZMQ worker matched by -`--selector` fails the fan-out with `unsupported_connection_mode`. +pause. A worker with no control endpoint (a gRPC or ZMQ worker without an +`rl.control_url` label) fails the fan-out with `no_control_endpoint`. """ from __future__ import annotations diff --git a/grpc_servicer/smg_grpc_servicer/tokenspeed/redact.py b/grpc_servicer/smg_grpc_servicer/tokenspeed/redact.py index c4f206e8c4..ebcccf3556 100644 --- a/grpc_servicer/smg_grpc_servicer/tokenspeed/redact.py +++ b/grpc_servicer/smg_grpc_servicer/tokenspeed/redact.py @@ -16,7 +16,11 @@ def is_secret_key(key: str) -> bool: """Whether a server-args key names a credential that must not leave the engine.""" lowered = key.lower() - return any(fragment in lowered for fragment in SECRET_FRAGMENTS) or lowered.endswith("_token") + return ( + lowered == "token" + or lowered.endswith("_token") + or any(fragment in lowered for fragment in SECRET_FRAGMENTS) + ) def redact_secrets(value: Any) -> Any: diff --git a/grpc_servicer/smg_grpc_servicer/tokenspeed/servicer.py b/grpc_servicer/smg_grpc_servicer/tokenspeed/servicer.py index 2664eee38e..cecd1e5fe8 100644 --- a/grpc_servicer/smg_grpc_servicer/tokenspeed/servicer.py +++ b/grpc_servicer/smg_grpc_servicer/tokenspeed/servicer.py @@ -12,6 +12,7 @@ import asyncio import dataclasses +import enum import functools import hashlib import json @@ -1764,14 +1765,39 @@ def _version_str(version: Any) -> str | None: return None if version is None else str(version) -def _make_json_serializable(obj: Any) -> Any: - """Flatten an arbitrary dataclass/config graph into JSON-safe primitives.""" +def _make_json_serializable(obj: Any, _path: frozenset[int] = frozenset()) -> Any: + """Flatten an arbitrary dataclass/config graph into JSON-safe primitives. + + Anything that carries attributes (a nested dataclass, a plain config + object) becomes a dict, so :func:`redact_secrets` sees its keys instead of + a ``str`` rendering that could carry a credential past it. Enums, paths, + dtypes and the like still render with ``str``. ``_path`` holds the ids of + the containers on the current descent: a back-reference (a sub-config + pointing at its parent) becomes a ```` placeholder instead of + recursing forever (never ``str(obj)``, whose repr could carry a credential), + while the same object reached twice by different paths is expanded twice. + """ if obj is None or isinstance(obj, str | int | float | bool): return obj + if id(obj) in _path: + return f"" + path = _path | {id(obj)} if isinstance(obj, list | tuple | set): - return [_make_json_serializable(x) for x in obj] + return [_make_json_serializable(x, path) for x in obj] if isinstance(obj, dict): - return {str(k): _make_json_serializable(v) for k, v in obj.items()} + return {str(k): _make_json_serializable(v, path) for k, v in obj.items()} + if dataclasses.is_dataclass(obj) and not isinstance(obj, type): + return { + f.name: _make_json_serializable(getattr(obj, f.name), path) + for f in dataclasses.fields(obj) + } + attrs = getattr(obj, "__dict__", None) + if attrs and not isinstance(obj, enum.Enum | type): + return { + str(k): _make_json_serializable(v, path) + for k, v in attrs.items() + if not str(k).startswith("_") + } return str(obj) diff --git a/grpc_servicer/tests/test_tokenspeed_rl_provenance.py b/grpc_servicer/tests/test_tokenspeed_rl_provenance.py index ec4b0f2a42..80196cbef6 100644 --- a/grpc_servicer/tests/test_tokenspeed_rl_provenance.py +++ b/grpc_servicer/tests/test_tokenspeed_rl_provenance.py @@ -9,6 +9,7 @@ """ import asyncio +import enum from types import SimpleNamespace import pytest @@ -155,6 +156,7 @@ def test_is_secret_key_targets_credentials_only(): is_secret = redact.is_secret_key assert is_secret("rl_control_api_key") and is_secret("api_key") and is_secret("hf_token") assert is_secret("some_secret") and is_secret("db_password") + assert is_secret("token") and is_secret("TOKEN") assert not is_secret("max_total_tokens") and not is_secret("tokenizer") assert not is_secret("weight_version") @@ -169,6 +171,7 @@ def test_redaction_reaches_nested_configs_and_lists(): "kv_store": {"endpoint": "redis://cache", "password": "p", "ttl": 5}, "providers": [{"name": "p", "api_key": "k"}, "plain"], "nested": {"deeper": {"api_key": "k", "keep": 1}}, + "hub": {"token": "t", "repo": "r"}, } ) assert redacted == { @@ -176,9 +179,64 @@ def test_redaction_reaches_nested_configs_and_lists(): "kv_store": {"endpoint": "redis://cache", "ttl": 5}, "providers": [{"name": "p"}, "plain"], "nested": {"deeper": {"keep": 1}}, + "hub": {"repo": "r"}, } +def test_nested_objects_are_flattened_before_redaction(): + """A plain config object nested in the server args used to be stringified + whole, which could carry a credential past the redaction. It is flattened + to a dict first, so the key is dropped; enums still render as strings.""" + + class Hub: + def __init__(self): + self.endpoint = "https://hub" + self.api_key = "k" + self._cache = "private, skipped" + + class Mode(enum.Enum): + PREFILL = "prefill" + + flat = servicer_mod._make_json_serializable({"hub": Hub(), "mode": Mode.PREFILL}) + assert flat == {"hub": {"endpoint": "https://hub", "api_key": "k"}, "mode": "Mode.PREFILL"} + assert redact.redact_secrets(flat) == { + "hub": {"endpoint": "https://hub"}, + "mode": "Mode.PREFILL", + } + + +def test_back_references_do_not_recurse_forever(): + """A sub-config that points back at its parent must stop the descent with a + placeholder that carries no value: a ``str(obj)`` leaf would print the + parent's repr, credential included. A shared object on two separate paths + is still expanded on both.""" + + class Node: + def __init__(self, name): + self.name = name + self.api_key = f"sk-{name}" + self.parent = None + + def __repr__(self): + return f"Node(name={self.name!r}, api_key={self.api_key!r})" + + root, child = Node("root"), Node("child") + child.parent = root + root.child = child + flat = servicer_mod._make_json_serializable({"cfg": root}) + assert flat["cfg"]["child"]["name"] == "child" + assert flat["cfg"]["child"]["parent"] == "", ( + "the back-reference is a value-free leaf" + ) + assert "sk-root" not in str(redact.redact_secrets(flat)), ( + "nothing of the parent leaks through the leaf" + ) + + shared = Node("shared") + twice = servicer_mod._make_json_serializable({"a": shared, "b": shared}) + assert twice["a"] == twice["b"] == {"name": "shared", "api_key": "sk-shared", "parent": None} + + class TestGetModelInfo: def test_weight_version_is_the_live_server_args_value(self): s = _servicer() diff --git a/model_gateway/src/app_context.rs b/model_gateway/src/app_context.rs index cc1d0db332..19564618cb 100644 --- a/model_gateway/src/app_context.rs +++ b/model_gateway/src/app_context.rs @@ -418,7 +418,11 @@ impl AppContextBuilder { &router_config.tenant_api_keys, ); - let rl = crate::rl_adapter::build_rl_state(&worker_registry, &router_config); + let rl = crate::rl_adapter::build_rl_state( + &worker_registry, + &worker_client_cache, + &router_config, + ); Ok(AppContext { gateway_auth, diff --git a/model_gateway/src/rl_adapter.rs b/model_gateway/src/rl_adapter.rs index 66a61ae55d..32886585f5 100644 --- a/model_gateway/src/rl_adapter.rs +++ b/model_gateway/src/rl_adapter.rs @@ -1,36 +1,109 @@ //! Glue between the RL control plane crate and the gateway registry. This //! file is the whole of coupling touchpoint (a); see `crates/rl/COUPLING.md`. -use std::sync::Arc; +use std::{ + collections::HashMap, + sync::{Arc, Mutex, PoisonError}, +}; -use openai_protocol::worker::ConnectionMode; -use smg_rl::{RlState, RlWorkerInfo, RlWorkerView}; +use openai_protocol::worker::{ConnectionMode, HttpPoolConfig}; +use smg_rl::{resolve_control_url, RlState, RlWorkerInfo, RlWorkerView}; +use tracing::warn; use crate::{ config::RouterConfig, - worker::{registry::WorkerId, Worker, WorkerRegistry}, + worker::{registry::WorkerId, Worker, WorkerHttpClientCache, WorkerRegistry}, }; +/// Label an engine (or an operator) sets to name the worker's RL control app. +pub const CONTROL_URL_LABEL: &str = "rl.control_url"; + /// Read-only registry view for the RL crate. pub struct RegistryRlView { registry: Arc, + client_cache: Arc, + /// Control clients for gRPC and ZMQ workers, by pool config. The gateway + /// cache holds weak entries that live only while some HTTP worker uses the + /// same config, and these workers hold no client of their own, so without + /// a strong handle here every control call would build a fresh client. A + /// failed build is kept as well: logged once, reported on each 422. + /// Pruned in [`RlWorkerView::list`] to the configs still registered. + control_clients: Mutex, String>>>, } impl RegistryRlView { - pub fn new(registry: Arc) -> Self { - Self { registry } + pub fn new(registry: Arc, client_cache: Arc) -> Self { + Self { + registry, + client_cache, + control_clients: Mutex::new(HashMap::new()), + } + } + + /// The shared control client for `pool`: same TLS identity and roots as + /// every upstream client, HTTP/1.1 because the control apps are uvicorn. + fn control_client(&self, pool: &HttpPoolConfig) -> Result, String> { + let mut clients = self + .control_clients + .lock() + .unwrap_or_else(PoisonError::into_inner); + if let Some(client) = clients.get(pool) { + return client.clone(); + } + let client = self.client_cache.get(pool, false); + if let Err(e) = &client { + warn!( + error = %e, ?pool, + "no HTTP client for RL control endpoints with this pool config" + ); + } + clients.insert(pool.clone(), client.clone()); + client + } + + /// Drop the handles no registered gRPC or ZMQ worker needs anymore. + fn prune_control_clients(&self, workers: &[Arc]) { + let mut clients = self + .control_clients + .lock() + .unwrap_or_else(PoisonError::into_inner); + clients.retain(|pool, _| { + workers.iter().any(|w| { + *w.connection_mode() != ConnectionMode::Http && w.metadata().spec.http_pool == *pool + }) + }); + } + + /// Where control calls for `worker` go and which client carries them. + /// + /// HTTP workers are controlled through themselves, with the client the + /// gateway negotiated for them. Other transports need an advertised or + /// operator-supplied `rl.control_url`, resolved against the worker's own + /// host, and use the shared control client for their pool config. + fn control_endpoint( + &self, + worker: &Arc, + ) -> (Option, Result, String>) { + let spec = &worker.metadata().spec; + if *worker.connection_mode() == ConnectionMode::Http { + let client = worker + .http_client_handle_if_initialized() + .unwrap_or_else(|| Arc::new(worker.http_client().clone())); + return (Some(worker.base_url().to_string()), Ok(client)); + } + let url = spec + .labels + .get(CONTROL_URL_LABEL) + .map(String::as_str) + .filter(|v| !v.trim().is_empty()) + .map(|advertised| resolve_control_url(advertised, worker.base_url())); + (url, self.control_client(&spec.http_pool)) } fn info(&self, worker: &Arc) -> Option { let id = self.registry.get_id_by_url(worker.url())?; let spec = &worker.metadata().spec; - // Borrow the client the gateway negotiated for this worker, as the - // admin ops do; a worker without one is not spoken to over HTTP. - let http_client = (*worker.connection_mode() == ConnectionMode::Http).then(|| { - worker - .http_client_handle_if_initialized() - .unwrap_or_else(|| Arc::new(worker.http_client().clone())) - }); + let (control_url, control_client) = self.control_endpoint(worker); Some(RlWorkerInfo { id: id.as_str().to_string(), url: worker.url().to_string(), @@ -44,18 +117,18 @@ impl RegistryRlView { is_dp_aware: worker.is_dp_aware(), dp_size: worker.dp_size(), labels: spec.labels.clone(), - http_client, + control_url, + control_client, }) } } impl RlWorkerView for RegistryRlView { fn list(&self) -> Vec { - self.registry - .get_all() - .iter() - .filter_map(|w| self.info(w)) - .collect() + let workers = self.registry.get_all(); + let infos = workers.iter().filter_map(|w| self.info(w)).collect(); + self.prune_control_clients(&workers); + infos } fn get(&self, id: &str) -> Option { @@ -68,12 +141,16 @@ impl RlWorkerView for RegistryRlView { /// disabled path constructs nothing. pub fn build_rl_state( registry: &Arc, + client_cache: &Arc, config: &RouterConfig, ) -> Option> { if !config.rl.enabled { return None; } - let view = Arc::new(RegistryRlView::new(Arc::clone(registry))); + let view = Arc::new(RegistryRlView::new( + Arc::clone(registry), + Arc::clone(client_cache), + )); Some(Arc::new(RlState::new(view, config.rl.clone()))) } @@ -84,18 +161,35 @@ mod tests { use openai_protocol::worker::ConnectionMode; use super::*; - use crate::worker::BasicWorkerBuilder; + use crate::{config::RouterConfig, worker::BasicWorkerBuilder}; - fn registry_with(worker: Arc) -> Arc { + fn view_with(config: &RouterConfig, workers: Vec>) -> RegistryRlView { let registry = Arc::new(WorkerRegistry::new()); - registry.register(worker); - registry + for w in workers { + registry.register(w); + } + let cache = Arc::new(WorkerHttpClientCache::new(config)); + RegistryRlView::new(registry, cache) + } + + fn view_over(workers: Vec>) -> RegistryRlView { + view_with(&RouterConfig::default(), workers) } - /// Control calls must use the client the gateway negotiated for the - /// worker (HTTP version, TLS identity, pool tuning), not a second one. + /// A gRPC worker whose engine advertised a wildcard-bound control app. + fn grpc_worker(url: &str) -> Arc { + Arc::new( + BasicWorkerBuilder::new(url) + .connection_mode(ConnectionMode::Grpc) + .label("rl.control_url", "http://0.0.0.0:40100") + .build(), + ) + } + + /// Control calls to an HTTP worker use the client the gateway negotiated + /// for it (HTTP version, TLS identity, pool tuning), not a second one. #[test] - fn view_hands_out_the_worker_negotiated_http_client() { + fn http_worker_controls_itself_with_its_negotiated_client() { let client = Arc::new(reqwest::Client::new()); let worker: Arc = Arc::new( BasicWorkerBuilder::new("http://engine:30000") @@ -103,24 +197,98 @@ mod tests { .http_client(Arc::clone(&client)) .build(), ); - let view = RegistryRlView::new(registry_with(worker)); + let info = view_over(vec![worker]).list().pop().expect("one worker"); + assert_eq!(info.control_url.as_deref(), Some("http://engine:30000")); + assert!(Arc::ptr_eq(&info.control_client.expect("client"), &client)); + } + + #[test] + fn grpc_worker_with_a_label_gets_a_control_endpoint_and_a_client() { + let info = view_over(vec![grpc_worker("grpc://10.0.0.5:30000")]) + .list() + .pop() + .expect("one worker"); + assert_eq!( + info.control_url.as_deref(), + Some("http://10.0.0.5:40100"), + "wildcard bind host resolves to the worker host" + ); + assert!(info.control_client.is_ok()); + } - let info = view.list().pop().expect("one worker"); - let handed = info.http_client.expect("HTTP worker carries a client"); - assert!(Arc::ptr_eq(&handed, &client)); + #[test] + fn grpc_worker_without_a_label_has_no_control_endpoint() { + let worker: Arc = Arc::new( + BasicWorkerBuilder::new("grpc://engine:30000") + .connection_mode(ConnectionMode::Grpc) + .build(), + ); + let info = view_over(vec![worker]).list().pop().expect("one worker"); + assert!(info.control_url.is_none()); } - /// A worker the gateway does not speak HTTP to has no client to hand out. #[test] - fn view_gives_no_http_client_for_grpc_workers() { + fn grpc_worker_with_a_blank_label_has_no_control_endpoint() { let worker: Arc = Arc::new( BasicWorkerBuilder::new("grpc://engine:30000") .connection_mode(ConnectionMode::Grpc) + .label("rl.control_url", " ") .build(), ); - let view = RegistryRlView::new(registry_with(worker)); + let info = view_over(vec![worker]).list().pop().expect("one worker"); + assert!(info.control_url.is_none()); + } + + #[test] + fn control_clients_are_shared_across_workers_with_the_same_pool_config() { + let infos = view_over(vec![grpc_worker("grpc://a:1"), grpc_worker("grpc://b:1")]).list(); + let a = infos[0].control_client.as_ref().expect("a"); + let b = infos[1].control_client.as_ref().expect("b"); + assert!(Arc::ptr_eq(a, b)); + } + + /// The gateway cache holds only weak handles and a gRPC worker holds + /// none, so the view keeps the strong one: a client is built once, not + /// on every discovery or control call. + #[test] + fn the_view_holds_the_control_client_between_calls() { + let view = view_over(vec![grpc_worker("grpc://a:1")]); + let client = view + .list() + .pop() + .expect("one worker") + .control_client + .expect("client"); + assert_eq!( + Arc::strong_count(&client), + 2, + "the view holds the other handle" + ); + view.registry.remove_by_url("grpc://a:1"); + assert!(view.list().is_empty()); + assert_eq!( + Arc::strong_count(&client), + 1, + "no registered worker needs the handle anymore" + ); + } - let info = view.list().pop().expect("one worker"); - assert!(info.http_client.is_none()); + /// A gateway whose TLS settings cannot produce a client is the worker's + /// problem to report, not a silent `None`: the error rides on the info + /// and ends up in the 422. (An unparsable client identity fails the + /// build eagerly; CA bundles are only read when a connection is made.) + #[test] + fn a_failed_client_build_is_reported_on_the_worker() { + let config = RouterConfig { + client_identity: Some(b"not a pem".to_vec()), + ..RouterConfig::default() + }; + let info = view_with(&config, vec![grpc_worker("grpc://a:1")]) + .list() + .pop() + .expect("one worker"); + assert_eq!(info.control_url.as_deref(), Some("http://a:40100")); + let err = info.control_client.expect_err("no client"); + assert!(err.contains("client identity"), "{err}"); } } diff --git a/model_gateway/tests/common/mock_worker.rs b/model_gateway/tests/common/mock_worker.rs index 10b57f7a47..68c1a54847 100755 --- a/model_gateway/tests/common/mock_worker.rs +++ b/model_gateway/tests/common/mock_worker.rs @@ -13,7 +13,7 @@ use std::{ use axum::{ extract::{Json, Multipart, Path, State}, - http::{StatusCode, Version}, + http::{header::AUTHORIZATION, HeaderMap, StatusCode, Version}, response::{ sse::{Event, KeepAlive}, IntoResponse, Response, Sse, @@ -170,6 +170,7 @@ fn clear_scheduler_controls(port: u16) { pub struct RequestRecorder { bodies: Mutex>, versions: Mutex>, + authorizations: Mutex>>, } impl RequestRecorder { @@ -189,6 +190,21 @@ impl RequestRecorder { .clone() } + /// The `authorization` header of every RL control request received, + /// oldest first, bodyless requests and simulated failures included (so + /// not aligned with [`Self::bodies`], which only has the requests that + /// carried JSON). `None` is a request that arrived without the header. + #[expect( + clippy::expect_used, + reason = "test helper - panicking on failure is intentional" + )] + pub fn authorizations(&self) -> Vec> { + self.authorizations + .lock() + .expect("request recorder mutex poisoned") + .clone() + } + /// Every body received so far, oldest first. #[expect( clippy::expect_used, @@ -252,6 +268,21 @@ fn record_request(port: u16, version: Version, body: &serde_json::Value) { } } +/// Record the `authorization` header of one RL control request: every +/// control request, before the failure simulation and whether or not it +/// carried a body. +fn record_authorization(port: u16, value: Option) { + let recorder = request_recorders_table() + .lock() + .ok() + .and_then(|table| table.get(&port).cloned()); + if let Some(recorder) = recorder { + if let Ok(mut authorizations) = recorder.authorizations.lock() { + authorizations.push(value); + } + } +} + /// Remove a port's recorder on teardown so a later worker reusing the port /// does not append to it. Tolerant of a poisoned mutex since it runs from /// `Drop`. @@ -1518,9 +1549,17 @@ async fn responses_handler( async fn rl_control_handler( State(config): State>>, version: Version, + headers: HeaderMap, body: Option>, ) -> Response { let config = config.read().await; + record_authorization( + config.port, + headers + .get(AUTHORIZATION) + .and_then(|value| value.to_str().ok()) + .map(str::to_string), + ); if should_fail(&config) { return ( diff --git a/model_gateway/tests/rl_tokenspeed_control_endpoint_test.rs b/model_gateway/tests/rl_tokenspeed_control_endpoint_test.rs new file mode 100644 index 0000000000..c66737b92f --- /dev/null +++ b/model_gateway/tests/rl_tokenspeed_control_endpoint_test.rs @@ -0,0 +1,520 @@ +//! A TokenSpeed engine SMG speaks gRPC to is driven through `/v1/rl` via its +//! HTTP control endpoint. The workers register through `POST /workers`, so the +//! endpoint and the capabilities come out of the engine's own server info the +//! way they do in production: discovery reports them, the proxy reaches the +//! endpoint with the worker's bearer, a fan-out over a mixed HTTP+gRPC fleet +//! hits each worker once, a gRPC worker whose engine advertises nothing is +//! named in `failed[]`, and the gRPC `/generate` path reports the version the +//! engine stamped on the response. +//! +//! Each test builds its own fleet on ephemeral mock HTTP ports (the HTTP mock +//! picks a free port when given 0, and the request recorder is registered under +//! that port before the gateway sends anything), so the tests in one binary can +//! run concurrently without racing for a listener. + +#[path = "common/mod.rs"] +mod common; + +use std::{collections::BTreeMap, sync::Arc, time::Duration}; + +use axum::{ + body::Body, + http::{Request, StatusCode}, +}; +use common::{ + mock_worker::{ + set_request_recorder, HealthStatus, MockWorker, MockWorkerConfig, RequestRecorder, + WorkerType as MockWorkerType, + }, + test_app::create_test_app_with_context, +}; +use http_body_util::BodyExt; +use llm_tokenizer::{traits::Tokenizer, MockTokenizer, TokenizerRegistry}; +use openai_protocol::generate::GenerateRequest; +use serde_json::{json, Value}; +use smg::{ + app_context::AppContext, + config::{RouterConfig, RoutingMode}, + middleware::TenantRequestMeta, + routers::{RouterFactory, RouterTrait}, + tenant::TenantKey, +}; +use tokio::net::TcpListener; +use tower::ServiceExt; + +const MODEL: &str = "rl-ts-test-model"; +/// What the mock engine stamps on every generate response it produces. +const ENGINE_VERSION: &str = "v7"; +/// What the workers carry as their registration-time `weight_version` label, +/// so a response that reports `ENGINE_VERSION` can only have come from the +/// engine rather than from the label. +const REGISTERED_VERSION: &str = "registered"; +/// The bearer the engine's control app expects, registered as the worker's +/// `api_key`. +const CONTROL_KEY: &str = "ts-secret"; + +/// What a current TokenSpeed engine puts in its server info for the RL +/// control plane, with its control app at `control_url`. +fn advertisement(control_url: &str) -> BTreeMap { + [ + ("rl.control_url", control_url), + ("rl.pause_modes", "wait,abort,keep"), + ("rl.update_from", "distributed,mooncake"), + ("rl.abort", "true"), + ("rl.flush_cache", "true"), + ("rl.sleep_wake", "true"), + ("rl.reports_weight_version", "true"), + ] + .into_iter() + .map(|(k, v)| (k.to_string(), v.to_string())) + .collect() +} + +/// A canned mock TokenSpeed gRPC engine that advertises `server_args` and +/// stamps `ENGINE_VERSION` on every generate response. +#[expect( + clippy::expect_used, + clippy::disallowed_methods, + reason = "test helper - panicking on failure is intentional; the spawned \ + mock server task is fire-and-forget for the test process's lifetime" +)] +async fn start_mock_grpc_engine(server_args: BTreeMap) -> u16 { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("bind mock gRPC engine"); + let port = listener + .local_addr() + .expect("mock gRPC engine address") + .port(); + let cfg = Arc::new(mock_worker::config::Config { + host: "127.0.0.1".to_string(), + http_base_port: 0, + http_count: 0, + grpc_base_port: port, + grpc_count: 1, + zmq_handshake: None, + zmq_count: 0, + zmq_start_index: 0, + model_id: MODEL.to_string(), + tokenizer_path: MODEL.to_string(), + gen_delay: Duration::ZERO, + output_tokens: 3, + realistic: false, + engine: mock_worker::engine::EngineParams::default(), + server_args, + weight_version: Some(ENGINE_VERSION.to_string()), + ..mock_worker::config::Config::default() + }); + tokio::spawn(mock_worker::grpc::serve_with_listener(cfg, listener)); + port +} + +/// A gRPC-transport regular router context with the RL control plane mounted +/// and a tokenizer preloaded for `MODEL`: the gRPC router needs one, and the +/// mock engine carries no tokenizer artifacts of its own, so autoload stays +/// off. +#[expect( + clippy::unwrap_used, + reason = "test helper - panicking on failure is intentional" +)] +async fn grpc_rl_context() -> Arc { + // Load the tokenizer before building the config so no `RouterConfig` + // (a large struct) is held across an await point in this future. + let registry = Arc::new(TokenizerRegistry::new()); + let tokenizer = Arc::new(MockTokenizer::new()) as Arc; + registry + .load( + "tokenizer-id", + MODEL, + "test", + || async move { Ok(tokenizer) }, + ) + .await + .unwrap(); + + let mut config = RouterConfig::builder() + .mode(RoutingMode::Regular { + worker_urls: vec![], + }) + .grpc_connection() + .round_robin_policy() + .host("127.0.0.1") + .port(0) + .max_payload_size(1024 * 1024) + .build_unchecked(); + config.health_check.disable_health_check = true; + config.disable_tokenizer_autoload = true; + config.rl.enabled = true; + config.rl.control_timeout_secs = 5; + + common::create_test_context_with_tokenizer_registry(config, registry).await +} + +/// The port a started mock reports in its URL. +#[expect( + clippy::unwrap_used, + reason = "test helper - panicking on failure is intentional" +)] +fn port_of(url: &str) -> u16 { + url.rsplit(':').next().unwrap().parse().unwrap() +} + +fn mock_http(port: u16) -> MockWorkerConfig { + MockWorkerConfig { + port, + worker_type: MockWorkerType::Regular, + health_status: HealthStatus::Healthy, + response_delay_ms: 0, + fail_rate: 0.0, + } +} + +#[expect( + clippy::unwrap_used, + reason = "test helper - panicking on failure is intentional" +)] +async fn json_of(resp: axum::response::Response) -> Value { + serde_json::from_slice(&resp.into_body().collect().await.unwrap().to_bytes()) + .unwrap_or(Value::Null) +} + +/// Register `spec` through the worker API, as an operator or a launcher does. +#[expect( + clippy::unwrap_used, + reason = "test helper - panicking on failure is intentional" +)] +async fn register(app: &axum::Router, spec: Value) { + let resp = app + .clone() + .oneshot( + Request::post("/workers") + .header("content-type", "application/json") + .body(Body::from(spec.to_string())) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(resp.status(), StatusCode::ACCEPTED, "{spec}"); +} + +/// The discovery rows once `n` workers are registered and ready. Registration +/// runs in the background, so poll. +#[expect( + clippy::unwrap_used, + reason = "test helper - panicking on failure is intentional" +)] +async fn wait_for_workers(app: &axum::Router, n: usize) -> Vec { + let mut workers = Vec::new(); + for _ in 0..100 { + let resp = app + .clone() + .oneshot(Request::get("/v1/rl/workers").body(Body::empty()).unwrap()) + .await + .unwrap(); + let body = json_of(resp).await; + workers = body["workers"].as_array().cloned().unwrap_or_default(); + if workers.len() == n && workers.iter().all(|w| w["health"] == "ready") { + break; + } + tokio::time::sleep(Duration::from_millis(100)).await; + } + assert_eq!(workers.len(), n, "{workers:?}"); + workers +} + +fn tenant() -> TenantRequestMeta { + TenantRequestMeta::new(TenantKey::new("rl-ts-test-tenant")) +} + +#[expect( + clippy::unwrap_used, + reason = "test helper - panicking on failure is intentional" +)] +fn generate_request(stream: bool) -> GenerateRequest { + serde_json::from_value(json!({ + "model": MODEL, + "input_ids": [1, 2, 3], + "sampling_params": {"max_new_tokens": 3}, + "stream": stream, + })) + .unwrap() +} + +/// The gRPC `/generate` path answers with one element per sample; `n` is +/// unset here, so the single element is the whole answer. +fn first_generate_result(body: Value) -> Value { + match body { + Value::Array(items) => items.into_iter().next().unwrap_or(Value::Null), + other => other, + } +} + +/// Two TokenSpeed gRPC workers (one whose engine advertises a control app at +/// an HTTP mock that records what it receives, one whose engine advertises +/// nothing) plus one HTTP SGLang mock, all registered through `POST /workers`. +struct Fleet { + ctx: Arc, + app: axum::Router, + router: Arc, + /// The discovery rows, in the order `GET /v1/rl/workers` listed them. + workers: Vec, + control_recorder: Arc, + sglang_recorder: Arc, + control_url: String, + _control_app: MockWorker, + _sglang: MockWorker, +} + +#[expect( + clippy::unwrap_used, + clippy::expect_used, + reason = "test helper - panicking on failure is intentional" +)] +async fn fleet() -> Fleet { + let mut control_app = MockWorker::new(mock_http(0)); + let control_url = control_app.start().await.unwrap(); + let control_recorder = RequestRecorder::new(); + set_request_recorder(port_of(&control_url), control_recorder.clone()); + let mut sglang = MockWorker::new(mock_http(0)); + let sglang_url = sglang.start().await.unwrap(); + let sglang_recorder = RequestRecorder::new(); + set_request_recorder(port_of(&sglang_url), sglang_recorder.clone()); + + let ctx = grpc_rl_context().await; + let router: Arc = Arc::from( + RouterFactory::create_router(&ctx) + .await + .expect("gRPC RL router should build"), + ); + let app = create_test_app_with_context(Arc::clone(&router), Arc::clone(&ctx)); + + let advertising = start_mock_grpc_engine(advertisement(&control_url)).await; + let silent = start_mock_grpc_engine(BTreeMap::new()).await; + register( + &app, + json!({ + "url": format!("grpc://127.0.0.1:{advertising}"), + "runtime_type": "tokenspeed", + "api_key": CONTROL_KEY, + "labels": {"weight_version": REGISTERED_VERSION}, + }), + ) + .await; + register( + &app, + json!({ + "url": format!("grpc://127.0.0.1:{silent}"), + "runtime_type": "tokenspeed", + "labels": {"weight_version": REGISTERED_VERSION}, + }), + ) + .await; + register(&app, json!({"url": sglang_url})).await; + let workers = wait_for_workers(&app, 3).await; + + Fleet { + ctx, + app, + router, + workers, + control_recorder, + sglang_recorder, + control_url, + _control_app: control_app, + _sglang: sglang, + } +} + +impl Fleet { + /// The one discovery row for `engine` whose `control_url` is (or is not) + /// set. + #[expect( + clippy::expect_used, + reason = "test helper - panicking on failure is intentional" + )] + fn by_engine(&self, engine: &str, has_control: bool) -> &Value { + self.workers + .iter() + .find(|w| w["engine"] == engine && w["control_url"].is_string() == has_control) + .expect("worker present") + } + + #[expect( + clippy::expect_used, + reason = "test helper - panicking on failure is intentional" + )] + fn id(&self, engine: &str, has_control: bool) -> String { + self.by_engine(engine, has_control)["id"] + .as_str() + .expect("worker id") + .to_string() + } +} + +#[tokio::test] +async fn discovery_reports_the_advertised_endpoint_and_capabilities() { + let f = fleet().await; + + let ts = f.by_engine("tokenspeed", true); + assert_eq!(ts["connection_mode"], "grpc"); + assert_eq!(ts["control_url"], f.control_url); + assert_eq!(ts["capabilities"]["source"], "label"); + assert_eq!( + ts["capabilities"]["pause_modes"], + json!(["wait", "abort", "keep"]) + ); + assert_eq!( + ts["capabilities"]["update_from"], + json!(["distributed", "mooncake"]) + ); + assert_eq!( + ts["weight_version"], REGISTERED_VERSION, + "the registration label wins over the mock's model-info placeholder" + ); + + let silent = f.by_engine("tokenspeed", false); + assert_eq!(silent["control_url"], Value::Null); + assert_eq!( + silent["capabilities"]["source"], "static", + "an engine that advertises nothing gets the built-in TokenSpeed row" + ); + assert_eq!( + silent["capabilities"]["update_from"], + json!(["distributed"]) + ); + + let sglang = f.by_engine("sglang", true); + assert_eq!( + sglang["control_url"], sglang["base_url"], + "HTTP workers control themselves" + ); +} + +#[tokio::test] +async fn proxy_reaches_the_control_endpoint_with_the_worker_bearer() { + let f = fleet().await; + let id = f.id("tokenspeed", true); + + let resp = f + .app + .clone() + .oneshot( + Request::post(format!("/v1/rl/workers/{id}/engine/pause_generation")) + .header("content-type", "application/json") + .header("authorization", "Bearer caller-token") + .body(Body::from(r#"{"mode":"keep"}"#)) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(resp.status(), StatusCode::OK); + assert_eq!(f.control_recorder.only_body(), json!({"mode": "keep"})); + assert_eq!( + f.control_recorder.authorizations(), + vec![Some(format!("Bearer {CONTROL_KEY}"))], + "the worker's own key, not the caller's" + ); +} + +#[tokio::test] +async fn fanout_over_workers_with_endpoints_is_200_and_hits_each_once() { + let f = fleet().await; + let ts = f.id("tokenspeed", true); + let sglang = f.id("sglang", true); + + let resp = f + .app + .clone() + .oneshot( + Request::post(format!( + "/v1/rl/engine/pause_generation?selector=id%20in%20({ts},{sglang})" + )) + .header("content-type", "application/json") + .body(Body::from(r#"{"mode":"abort"}"#)) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(resp.status(), StatusCode::OK); + let body = json_of(resp).await; + assert_eq!(body["total"], 2); + assert_eq!(body["succeeded"], 2); + assert_eq!( + f.control_recorder.only_body(), + json!({"mode": "abort"}), + "the gRPC worker's control app, exactly once" + ); + assert_eq!( + f.sglang_recorder.only_body(), + json!({"mode": "abort"}), + "the HTTP worker itself, exactly once" + ); +} + +#[tokio::test] +async fn fanout_names_the_worker_without_an_endpoint() { + let f = fleet().await; + let resp = f + .app + .clone() + .oneshot( + Request::post("/v1/rl/engine/flush_cache?selector=worker_type%3Dregular") + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(resp.status(), StatusCode::MULTI_STATUS); + let body = json_of(resp).await; + assert_eq!(body["total"], 3); + assert_eq!( + body["succeeded"], 2, + "the HTTP SGLang mock and the TokenSpeed mock with an endpoint" + ); + let failed = body["failed"].as_array().unwrap(); + assert_eq!(failed.len(), 1); + assert_eq!(failed[0]["error"], "no_control_endpoint"); + assert_eq!(failed[0]["connection_mode"], "grpc"); + assert_eq!(failed[0]["worker_id"], f.id("tokenspeed", false)); + + // Control calls never touch breakers or load accounting. + for worker in f.ctx.worker_registry.get_all() { + assert!(worker.circuit_breaker_can_execute(), "{}", worker.url()); + assert_eq!(worker.load(), 0, "{}", worker.url()); + } +} + +#[tokio::test] +async fn grpc_generate_reports_the_engine_stamped_version() { + let f = fleet().await; + + let resp = f + .router + .route_generate(None, &tenant(), generate_request(false), MODEL) + .await; + assert_eq!(resp.status(), StatusCode::OK); + let first = first_generate_result(json_of(resp).await); + assert_eq!( + first["meta_info"]["weight_version"], ENGINE_VERSION, + "engine value beats the `{REGISTERED_VERSION}` label" + ); + + let resp = f + .router + .route_generate(None, &tenant(), generate_request(true), MODEL) + .await; + assert_eq!(resp.status(), StatusCode::OK); + let text = String::from_utf8( + resp.into_body() + .collect() + .await + .unwrap() + .to_bytes() + .to_vec(), + ) + .unwrap(); + assert!( + text.contains(&format!(r#""weight_version":"{ENGINE_VERSION}""#)), + "{text}" + ); +}