Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
24 commits
Select commit Hold shift + click to select a range
74288b7
feat(rl): model a worker's control endpoint separately from its data …
key4ng Sep 22, 2026
db0db0b
feat(rl): proxy control calls to the worker's control endpoint regard…
key4ng Sep 22, 2026
58cc3b9
feat(rl): resolve control endpoints for gRPC and ZMQ workers from the…
key4ng Sep 22, 2026
0b625a6
feat(rl): report control_url in discovery, add the TokenSpeed capabil…
key4ng Sep 22, 2026
848bad9
test(mock_worker): advertise custom server_args and stamp a weight ve…
key4ng Sep 22, 2026
071f209
chore(mock_worker): record the prost-types dependency in Cargo.lock
key4ng Sep 22, 2026
a1c76a2
test(rl): drive a TokenSpeed gRPC worker through /v1/rl via its contr…
key4ng Sep 22, 2026
3bc1d46
feat(serve): wire the TokenSpeed control app for ZMQ workers and stam…
key4ng Sep 22, 2026
835121c
docs(rl): control endpoints, TokenSpeed drift rows, and the TokenSpee…
key4ng Sep 22, 2026
64a3ccb
fix(rl): treat a blank rl.control_url label as absent; document base_url
key4ng Sep 22, 2026
3adfb58
fix(serve): stop retrying 4xx when stamping rl.control_url; require zmq
key4ng Sep 22, 2026
54a9943
docs(rl): HTTP workers ignore rl.control_url; narrow recorder doc
key4ng Sep 22, 2026
5378dbf
style(rl): drop redundant serde_json qualifications in the discovery …
key4ng Sep 22, 2026
a74b03b
fix(rl): TokenSpeed refits are distributed-only; drop disk from the s…
key4ng Sep 22, 2026
229f15b
fix(rl): fall back to attn_tp_size when an engine reports no tp_size
key4ng Sep 22, 2026
8a64238
docs(rl): list the gateway-side files that carry RL data and name bot…
key4ng Sep 23, 2026
17e09d9
test(mock_worker): set the new mock fields in the context-length test…
key4ng Sep 23, 2026
be7ee23
style(rl): apply rustfmt to the control-endpoint resolver
key4ng Sep 23, 2026
54b96dc
fix(rl): hold control clients strongly and carry a client build failu…
key4ng Oct 6, 2026
cabc212
test(rl): register the TokenSpeed mock through POST /workers; a fan-o…
key4ng Oct 6, 2026
63a506d
fix(serve): wire no RL control endpoint for ZMQ TokenSpeed workers
key4ng Oct 6, 2026
f81491c
docs(rl): launch recipe that works against the merged engine
key4ng Oct 6, 2026
9ae9aa7
test(rl): make the failed-client-build case fail the build for real
key4ng Oct 6, 2026
5f12276
style(rl): build the test's RouterConfig with struct-update syntax
key4ng Oct 6, 2026
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions Cargo.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

8 changes: 5 additions & 3 deletions bindings/python/src/smg/rl.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)

Expand Down
3 changes: 1 addition & 2 deletions bindings/python/src/smg/serve.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
4 changes: 4 additions & 0 deletions bindings/python/tests/test_rl_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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"

Expand Down Expand Up @@ -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):
Expand Down
1 change: 1 addition & 0 deletions crates/mock_worker/Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,7 @@ serde_json.workspace = true
futures.workspace = true
tracing.workspace = true
tracing-subscriber.workspace = true
prost-types.workspace = true

[dev-dependencies]
tempfile = "3"
10 changes: 9 additions & 1 deletion crates/mock_worker/src/config.rs
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
//! Runtime configuration for the mock worker fleet, parsed from CLI flags.

use std::time::Duration;
use std::{collections::BTreeMap, time::Duration};

use crate::engine::EngineParams;

Expand Down Expand Up @@ -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<String, String>,
/// 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<String>,
}

impl Config {
Expand All @@ -58,6 +64,8 @@ impl Config {
output_tokens: 8,
realistic: false,
engine: EngineParams::default(),
server_args: BTreeMap::new(),
weight_version: None,
};

let mut args = std::env::args().skip(1);
Expand Down
46 changes: 39 additions & 7 deletions crates/mock_worker/src/grpc.rs
Original file line number Diff line number Diff line change
Expand Up @@ -92,6 +92,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(),
Expand All @@ -103,6 +104,7 @@ impl TokenSpeedScheduler for MockScheduler {
rx,
stream_chunks,
request_id,
weight_version,
)));
}

Expand All @@ -124,7 +126,7 @@ impl TokenSpeedScheduler for MockScheduler {
cached_tokens: 0,
output_logprobs: None,
index: 0,
weight_version: None,
weight_version: self.cfg.weight_version.clone(),
})),
}));
}
Expand All @@ -139,6 +141,7 @@ impl TokenSpeedScheduler for MockScheduler {
output_logprobs: None,
matched_stop: None,
index: 0,
weight_version: self.cfg.weight_version.clone(),
..Default::default()
})),
}));
Expand Down Expand Up @@ -193,8 +196,23 @@ impl TokenSpeedScheduler for MockScheduler {
&self,
_request: Request<ts::GetServerInfoRequest>,
) -> Result<Response<ts::GetServerInfoResponse>, 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,
Expand Down Expand Up @@ -295,11 +313,18 @@ fn generate_stream(
rx: mpsc::UnboundedReceiver<engine::GenEvent>,
stream_chunks: bool,
request_id: String,
weight_version: Option<String>,
) -> GenStream {
let init = (rx, Vec::<u32>::new(), stream_chunks, request_id);
let init = (
rx,
Vec::<u32>::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 {
Expand All @@ -318,10 +343,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.
}
Expand All @@ -342,10 +370,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,
}
Expand Down
4 changes: 4 additions & 0 deletions crates/mock_worker/src/zmq.rs
Original file line number Diff line number Diff line change
Expand Up @@ -284,6 +284,8 @@ fn map_finish(reason: &str) -> EngineCoreFinishReason {

#[cfg(test)]
mod tests {
use std::collections::BTreeMap;

use engine_zmq_client::{
connect_handshake, protocol::vllm::request::EngineCoreRequest, EngineCoreClient,
};
Expand All @@ -308,6 +310,8 @@ mod tests {
output_tokens: 4,
realistic: false,
engine: EngineParams::default(),
server_args: BTreeMap::new(),
weight_version: None,
}
}

Expand Down
11 changes: 8 additions & 3 deletions crates/protocols/src/rl.rs
Original file line number Diff line number Diff line change
Expand Up @@ -48,14 +48,18 @@ pub struct RlWorkerEntry {
pub id: String,
/// Registry URL; carries an `@<rank>` 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,
pub engine_version: Option<String>,
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<String>,
pub tp_size: Option<u64>,
pub dp_size: Option<u64>,
pub pp_size: Option<u64>,
Expand Down Expand Up @@ -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<u16>,
/// 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<String>,
}
Expand Down Expand Up @@ -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,
Expand Down
4 changes: 3 additions & 1 deletion crates/protocols/src/worker.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1101,7 +1101,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<usize>,
Expand Down
16 changes: 15 additions & 1 deletion crates/rl/COUPLING.md
Original file line number Diff line number Diff line change
Expand Up @@ -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<Arc<RlState>>` | `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()` |
Expand All @@ -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.
Expand Down
7 changes: 7 additions & 0 deletions crates/rl/NOTES.md
Original file line number Diff line number Diff line change
Expand Up @@ -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 |
26 changes: 22 additions & 4 deletions crates/rl/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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 <reachable> --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

Expand Down
Loading
Loading