From 5bdeb49f9f5773c77738bf0206c4f02926198750 Mon Sep 17 00:00:00 2001 From: key4ng Date: Tue, 22 Sep 2026 01:32:53 -0700 Subject: [PATCH 01/24] feat(rl): model a worker's control endpoint separately from its data transport Signed-off-by: key4ng --- crates/rl/src/control.rs | 97 ++++++++++++++++++++++++++++++++++++++++ crates/rl/src/lib.rs | 2 + crates/rl/src/state.rs | 10 ++--- crates/rl/src/testing.rs | 6 ++- crates/rl/src/view.rs | 42 ++++++++++++----- 5 files changed, 140 insertions(+), 17 deletions(-) create mode 100644 crates/rl/src/control.rs diff --git a/crates/rl/src/control.rs b/crates/rl/src/control.rs new file mode 100644 index 0000000000..bfbd951d9a --- /dev/null +++ b/crates/rl/src/control.rs @@ -0,0 +1,97 @@ +//! 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 URL. +/// +/// An engine that bound `0.0.0.0` (or `::`) advertises what it bound; the +/// only host the gateway knows is reachable is the worker's own, so use it. +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(""); + let worker_host = worker_host.split('@').next().unwrap_or(worker_host); + 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@2"), + "http://10.0.0.5:40100", + "a DP-rank suffix on the worker URL is not part of the host" + ); + } + + #[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/lib.rs b/crates/rl/src/lib.rs index a4b1895e53..11558c9335 100644 --- a/crates/rl/src/lib.rs +++ b/crates/rl/src/lib.rs @@ -4,6 +4,7 @@ pub mod capability; pub mod config; +pub 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/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..cf420d308c 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: Some(test_client()), } } diff --git a/crates/rl/src/view.rs b/crates/rl/src/view.rs index 1b08fe3add..17c01a33fd 100644 --- a/crates/rl/src/view.rs +++ b/crates/rl/src/view.rs @@ -25,10 +25,16 @@ 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 used 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 cached + /// client with the same TLS settings on HTTP/1.1. `None` when the worker + /// has no control endpoint. + pub control_client: Option>, } impl fmt::Debug for RlWorkerInfo { @@ -46,7 +52,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 +73,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 +87,24 @@ mod tests { is_dp_aware: false, dp_size: None, labels: HashMap::new(), - http_client: None, - }; + control_url: None, + control_client: None, + } + } + + #[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\")")); + } } From f4f8482a901cc4aca1f7a7ff68ff5126cc2b9241 Mon Sep 17 00:00:00 2001 From: key4ng Date: Tue, 22 Sep 2026 01:38:04 -0700 Subject: [PATCH 02/24] feat(rl): proxy control calls to the worker's control endpoint regardless of data transport Signed-off-by: key4ng --- crates/protocols/src/rl.rs | 4 +- crates/rl/src/error.rs | 31 +++++++++--- crates/rl/src/fanout.rs | 9 +++- crates/rl/src/proxy.rs | 101 ++++++++++++++++++++++++++++++------- 4 files changed, 115 insertions(+), 30 deletions(-) diff --git a/crates/protocols/src/rl.rs b/crates/protocols/src/rl.rs index e5275fdfd9..fbc3c4b5b6 100644 --- a/crates/protocols/src/rl.rs +++ b/crates/protocols/src/rl.rs @@ -110,13 +110,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, } 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..1bf47e0812 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,8 @@ 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; + grpc.control_client = None; let app = crate::router::<()>(state( vec![ worker("w1", &good.url, RuntimeType::Sglang), @@ -261,7 +265,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/proxy.rs b/crates/rl/src/proxy.rs index 898e118237..7d013cbc47 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; register the worker with one or upgrade the engine to advertise it", + )); + }; + let Some(client) = worker.control_client.as_ref() else { + return Err(no_control_endpoint( + worker, + "the gateway could not build an HTTP client for the control endpoint (see the gateway log)", + )); }; - let url = req.url_for(worker); + 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,8 @@ 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; + grpc.control_client = None; let app = crate::router::<()>(state( vec![worker("w1", &engine.url, RuntimeType::Sglang), grpc], 5, @@ -430,7 +446,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 +472,62 @@ 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 = None; + 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"); + assert!(body["message"].as_str().unwrap().contains("HTTP client")); + } + /// 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 = Some(Arc::new(reqwest::Client::new())); let app = crate::router::<()>(state(vec![w], 1)); let r = app From dfa72c9bcb3631337b6802fc6a4bbf91bd753af9 Mon Sep 17 00:00:00 2001 From: key4ng Date: Tue, 22 Sep 2026 01:42:09 -0700 Subject: [PATCH 03/24] feat(rl): resolve control endpoints for gRPC and ZMQ workers from the rl.control_url label Signed-off-by: key4ng --- crates/rl/COUPLING.md | 2 +- model_gateway/src/app_context.rs | 6 +- model_gateway/src/rl_adapter.rs | 131 ++++++++++++++++++++++++------- 3 files changed, 108 insertions(+), 31 deletions(-) diff --git a/crates/rl/COUPLING.md b/crates/rl/COUPLING.md index 5d8c7d0d5b..3c97793146 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()` | diff --git a/model_gateway/src/app_context.rs b/model_gateway/src/app_context.rs index bb9445b4eb..282be4e9d9 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..15f1e3b1e8 100644 --- a/model_gateway/src/rl_adapter.rs +++ b/model_gateway/src/rl_adapter.rs @@ -4,33 +4,69 @@ use std::sync::Arc; use openai_protocol::worker::ConnectionMode; -use smg_rl::{RlState, RlWorkerInfo, RlWorkerView}; +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, } impl RegistryRlView { - pub fn new(registry: Arc) -> Self { - Self { registry } + pub fn new(registry: Arc, client_cache: Arc) -> Self { + Self { + registry, + client_cache, + } + } + + /// 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`; the client comes from the shared + /// cache (same TLS identity and roots as every upstream client, HTTP/1.1 + /// because the control apps are uvicorn). + fn control_endpoint( + &self, + worker: &Arc, + ) -> (Option, Option>) { + 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()), Some(client)); + } + let Some(advertised) = spec.labels.get(CONTROL_URL_LABEL) else { + return (None, None); + }; + let url = resolve_control_url(advertised, worker.url()); + match self.client_cache.get(&spec.http_pool, false) { + Ok(client) => (Some(url), Some(client)), + Err(e) => { + warn!( + worker = %worker.url(), error = %e, + "no HTTP client for the RL control endpoint" + ); + (Some(url), None) + } + } } 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,7 +80,8 @@ impl RegistryRlView { is_dp_aware: worker.is_dp_aware(), dp_size: worker.dp_size(), labels: spec.labels.clone(), - http_client, + control_url, + control_client, }) } } @@ -68,12 +105,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 +125,21 @@ 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_over(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(&RouterConfig::default())); + RegistryRlView::new(registry, cache) } - /// Control calls must use the client the gateway negotiated for the - /// worker (HTTP version, TLS identity, pool tuning), not a second one. + /// 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 +147,53 @@ 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)); + } - 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_with_a_label_gets_a_control_endpoint_and_a_cached_client() { + let worker: Arc = Arc::new( + BasicWorkerBuilder::new("grpc://10.0.0.5:30000") + .connection_mode(ConnectionMode::Grpc) + .label("rl.control_url", "http://0.0.0.0:40100") + .build(), + ); + let info = view_over(vec![worker]).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_some()); } - /// 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_without_a_label_has_no_control_endpoint() { let worker: Arc = Arc::new( BasicWorkerBuilder::new("grpc://engine:30000") .connection_mode(ConnectionMode::Grpc) .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()); + assert!(info.control_client.is_none()); + } - let info = view.list().pop().expect("one worker"); - assert!(info.http_client.is_none()); + #[test] + fn cached_control_clients_are_shared_across_workers_with_the_same_pool_config() { + let mk = |url: &str| -> Arc { + Arc::new( + BasicWorkerBuilder::new(url) + .connection_mode(ConnectionMode::Grpc) + .label("rl.control_url", "http://0.0.0.0:40100") + .build(), + ) + }; + let infos = view_over(vec![mk("grpc://a:1"), mk("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)); } } From fd9c279a80d5d61d2df6807c98cfe177ebf64a00 Mon Sep 17 00:00:00 2001 From: key4ng Date: Tue, 22 Sep 2026 02:10:00 -0700 Subject: [PATCH 04/24] feat(rl): report control_url in discovery, add the TokenSpeed capability row, expose control_url in the Python client Signed-off-by: key4ng --- bindings/python/src/smg/rl.py | 8 ++++-- bindings/python/tests/test_rl_client.py | 4 +++ crates/protocols/src/rl.rs | 5 ++++ crates/rl/src/capability.rs | 35 ++++++++++++++++++++++- crates/rl/src/discovery.rs | 37 ++++++++++++++++++++++++- crates/rl/src/view.rs | 2 +- examples/rl/refit_from_disk.py | 4 +-- 7 files changed, 87 insertions(+), 8 deletions(-) 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/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/protocols/src/rl.rs b/crates/protocols/src/rl.rs index fbc3c4b5b6..a7e198ba3c 100644 --- a/crates/protocols/src/rl.rs +++ b/crates/protocols/src/rl.rs @@ -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, @@ -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/rl/src/capability.rs b/crates/rl/src/capability.rs index 93c2893404..8fde331679 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(&["disk", "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,29 @@ 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. + #[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, ["disk", "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", "disk,distributed"), + ]), + ); + assert_eq!(t.source, RlCapabilitySource::Label); + assert_eq!(t.pause_modes, ["wait", "abort", "keep"]); + } } diff --git a/crates/rl/src/discovery.rs b/crates/rl/src/discovery.rs index 63fa66b611..f82cc19c79 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![ @@ -308,4 +314,33 @@ 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; + n.control_client = 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"], serde_json::Value::Null); + assert_eq!(by_id("g1")["engine"], "tokenspeed"); + } } diff --git a/crates/rl/src/view.rs b/crates/rl/src/view.rs index 17c01a33fd..ef2387988c 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, 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 From edd4fdeee774ef57533e01eac9640b7e4621eff2 Mon Sep 17 00:00:00 2001 From: key4ng Date: Tue, 22 Sep 2026 02:29:41 -0700 Subject: [PATCH 05/24] test(mock_worker): advertise custom server_args and stamp a weight version on gRPC generates Signed-off-by: key4ng --- crates/mock_worker/Cargo.toml | 1 + crates/mock_worker/src/config.rs | 20 +++++++- crates/mock_worker/src/grpc.rs | 46 ++++++++++++++++--- crates/mock_worker/src/zmq.rs | 4 ++ .../grpc/regular/streaming/eof_tests.rs | 4 +- model_gateway/tests/grpc_pd_fanout_test.rs | 8 +++- .../tests/tenant_rate_limiting_grpc_test.rs | 4 +- model_gateway/tests/zmq_backend_test.rs | 4 +- 8 files changed, 79 insertions(+), 12 deletions(-) diff --git a/crates/mock_worker/Cargo.toml b/crates/mock_worker/Cargo.toml index 32027848a3..d6db8c7c9a 100644 --- a/crates/mock_worker/Cargo.toml +++ b/crates/mock_worker/Cargo.toml @@ -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" diff --git a/crates/mock_worker/src/config.rs b/crates/mock_worker/src/config.rs index 8140db7ba4..4443ea534a 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::EngineParams; @@ -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, /// Settings for replay testing (gRPC workers only). pub replay: ReplayConfig, } @@ -76,6 +82,8 @@ impl Config { output_tokens: 8, realistic: false, engine: EngineParams::default(), + server_args: BTreeMap::new(), + weight_version: None, replay: ReplayConfig::default(), }; @@ -126,6 +134,14 @@ impl Config { "--prefix-cache" => { cfg.engine.prefix_cache = parse(value(&mut args, &flag)?, &flag)? } + "--server-arg" => { + let raw = value(&mut args, &flag)?; + let (k, v) = raw + .split_once('=') + .ok_or_else(|| format!("--server-arg expects key=value, got {raw}"))?; + cfg.server_args.insert(k.to_string(), v.to_string()); + } + "--weight-version" => cfg.weight_version = Some(value(&mut args, &flag)?), "-h" | "--help" => return Err(usage()), other => return Err(format!("unknown flag: {other}\n\n{}", usage())), } @@ -178,6 +194,8 @@ fn usage() -> String { --tokenizer tokenizer path for gRPC autoload (default = model)\n\ --gen-ms canned per-request latency (default 0)\n\ --output-tokens output tokens per request when unspecified (default 8)\n\ + --server-arg extra GetServerInfo server_args entry (repeatable)\n\ + --weight-version weight version stamped on generate responses (default unset)\n\ --capture append each gRPC Generate request to as a JSON line\n\ \n\ Realistic engine simulator (continuous batching; opt-in):\n\ diff --git a/crates/mock_worker/src/grpc.rs b/crates/mock_worker/src/grpc.rs index fc754691bf..ce4901852c 100644 --- a/crates/mock_worker/src/grpc.rs +++ b/crates/mock_worker/src/grpc.rs @@ -121,6 +121,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(), @@ -132,6 +133,7 @@ impl TokenSpeedScheduler for MockScheduler { rx, stream_chunks, request_id, + weight_version, ))); } @@ -153,7 +155,7 @@ impl TokenSpeedScheduler for MockScheduler { cached_tokens: 0, output_logprobs: None, index: 0, - weight_version: None, + weight_version: self.cfg.weight_version.clone(), })), })); } @@ -168,6 +170,7 @@ impl TokenSpeedScheduler for MockScheduler { output_logprobs: None, matched_stop: None, index: 0, + weight_version: self.cfg.weight_version.clone(), ..Default::default() })), })); @@ -222,8 +225,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, @@ -324,11 +342,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 { @@ -347,10 +372,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. } @@ -371,10 +399,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/mock_worker/src/zmq.rs b/crates/mock_worker/src/zmq.rs index c58d9079c7..b07b30bb36 100644 --- a/crates/mock_worker/src/zmq.rs +++ b/crates/mock_worker/src/zmq.rs @@ -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, }; @@ -308,6 +310,8 @@ mod tests { output_tokens: 4, realistic: false, engine: EngineParams::default(), + server_args: BTreeMap::new(), + weight_version: None, replay: Default::default(), } } diff --git a/model_gateway/src/routers/grpc/regular/streaming/eof_tests.rs b/model_gateway/src/routers/grpc/regular/streaming/eof_tests.rs index 984e6821f9..bff869c4f0 100644 --- a/model_gateway/src/routers/grpc/regular/streaming/eof_tests.rs +++ b/model_gateway/src/routers/grpc/regular/streaming/eof_tests.rs @@ -1,6 +1,6 @@ //! Test gateway parsers and SSE events with exact gRPC frames. -use std::time::Duration; +use std::{collections::BTreeMap, time::Duration}; use axum::http::{HeaderMap, HeaderValue, StatusCode}; use bytes::BufMut; @@ -132,6 +132,8 @@ async fn scripted_stream( output_tokens: 0, realistic: false, engine: mock_worker::engine::EngineParams::default(), + server_args: BTreeMap::new(), + weight_version: None, replay: Default::default(), }); let server = tokio::spawn(mock_worker::grpc::serve_with_listener(config, listener)); diff --git a/model_gateway/tests/grpc_pd_fanout_test.rs b/model_gateway/tests/grpc_pd_fanout_test.rs index 5a5954678e..ae5248cc68 100644 --- a/model_gateway/tests/grpc_pd_fanout_test.rs +++ b/model_gateway/tests/grpc_pd_fanout_test.rs @@ -11,7 +11,11 @@ #[path = "common/mod.rs"] mod common; -use std::{collections::BTreeSet, sync::Arc, time::Duration}; +use std::{ + collections::{BTreeMap, BTreeSet}, + sync::Arc, + time::Duration, +}; use llm_tokenizer::{ chat_template::ChatTemplateParams, @@ -129,6 +133,8 @@ async fn start_mock_grpc_worker() -> u16 { output_tokens: OUTPUT_TOKENS, realistic: false, engine: mock_worker::engine::EngineParams::default(), + server_args: BTreeMap::new(), + weight_version: None, replay: Default::default(), }); tokio::spawn(mock_worker::grpc::serve_with_listener(cfg, listener)); diff --git a/model_gateway/tests/tenant_rate_limiting_grpc_test.rs b/model_gateway/tests/tenant_rate_limiting_grpc_test.rs index d5775eca98..35d45db4a1 100644 --- a/model_gateway/tests/tenant_rate_limiting_grpc_test.rs +++ b/model_gateway/tests/tenant_rate_limiting_grpc_test.rs @@ -22,7 +22,7 @@ #[path = "common/mod.rs"] mod common; -use std::{sync::Arc, time::Duration}; +use std::{collections::BTreeMap, sync::Arc, time::Duration}; use llm_tokenizer::{traits::Tokenizer, MockTokenizer, TokenizerRegistry}; use openai_protocol::{ @@ -78,6 +78,8 @@ async fn start_mock_grpc_worker(output_tokens: u32) -> u16 { output_tokens, realistic: false, engine: mock_worker::engine::EngineParams::default(), + server_args: BTreeMap::new(), + weight_version: None, replay: Default::default(), }); tokio::spawn(mock_worker::grpc::serve_with_listener(cfg, listener)); diff --git a/model_gateway/tests/zmq_backend_test.rs b/model_gateway/tests/zmq_backend_test.rs index 24679d19c2..83b69daffa 100644 --- a/model_gateway/tests/zmq_backend_test.rs +++ b/model_gateway/tests/zmq_backend_test.rs @@ -18,7 +18,7 @@ #[path = "common/mod.rs"] mod common; -use std::{sync::Arc, time::Duration}; +use std::{collections::BTreeMap, sync::Arc, time::Duration}; use llm_tokenizer::{traits::Tokenizer, MockTokenizer, TokenizerRegistry}; use openai_protocol::{ @@ -87,6 +87,8 @@ fn start_mock_zmq_engines(handshake: &str, count: u16) { output_tokens: OUTPUT_TOKENS, realistic: false, engine: mock_worker::engine::EngineParams::default(), + server_args: BTreeMap::new(), + weight_version: None, replay: Default::default(), }); for rank in 0..u32::from(count) { From 93fe65e4e882e4876bed0580f0d41997bc408fbd Mon Sep 17 00:00:00 2001 From: key4ng Date: Tue, 22 Sep 2026 02:42:39 -0700 Subject: [PATCH 06/24] chore(mock_worker): record the prost-types dependency in Cargo.lock Signed-off-by: key4ng --- Cargo.lock | 1 + 1 file changed, 1 insertion(+) diff --git a/Cargo.lock b/Cargo.lock index 610eaf5e32..4834ba0858 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -4368,6 +4368,7 @@ dependencies = [ "axum", "engine-zmq-client", "futures", + "prost-types", "serde_json", "smg-grpc-client", "tempfile", From 42c2ba039d5542ec97b4ce7bd370789ce36532aa Mon Sep 17 00:00:00 2001 From: key4ng Date: Tue, 22 Sep 2026 03:01:43 -0700 Subject: [PATCH 07/24] test(rl): drive a TokenSpeed gRPC worker through /v1/rl via its control endpoint Signed-off-by: key4ng --- model_gateway/tests/common/mock_worker.rs | 39 +- .../rl_tokenspeed_control_endpoint_test.rs | 440 ++++++++++++++++++ 2 files changed, 478 insertions(+), 1 deletion(-) create mode 100644 model_gateway/tests/rl_tokenspeed_control_endpoint_test.rs diff --git a/model_gateway/tests/common/mock_worker.rs b/model_gateway/tests/common/mock_worker.rs index 4fb3cf958a..1273ee9d33 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,20 @@ impl RequestRecorder { .clone() } + /// The `authorization` header each recorded RL control request carried, + /// oldest first and index-aligned with [`Self::bodies`]. `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 +267,20 @@ fn record_request(port: u16, version: Version, body: &serde_json::Value) { } } +/// Record the `authorization` header of one RL control request. Called +/// alongside [`record_request`] so the two vectors stay index-aligned. +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`. @@ -1517,6 +1546,7 @@ async fn responses_handler( async fn rl_control_handler( State(config): State>>, version: Version, + headers: HeaderMap, body: Option>, ) -> Response { let config = config.read().await; @@ -1531,6 +1561,13 @@ async fn rl_control_handler( if let Some(Json(body)) = body { record_request(config.port, version, &body); + record_authorization( + config.port, + headers + .get(AUTHORIZATION) + .and_then(|value| value.to_str().ok()) + .map(str::to_string), + ); } Json(json!({"success": true, "message": "ok"})).into_response() 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..b817b140a8 --- /dev/null +++ b/model_gateway/tests/rl_tokenspeed_control_endpoint_test.rs @@ -0,0 +1,440 @@ +//! A TokenSpeed engine SMG speaks gRPC to is driven through `/v1/rl` via its +//! HTTP control endpoint: discovery reports it, the proxy reaches it with the +//! worker's bearer, a mixed HTTP+gRPC fleet fans out, a gRPC worker without an +//! endpoint 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 its own pair of mock HTTP ports: the HTTP +//! mock binds the port it is configured with, and the tests in one binary run +//! concurrently, so a shared pair would race for the same 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, model_card::ModelCard, worker::HealthCheckConfig, +}; +use serde_json::{json, Value}; +use smg::{ + app_context::AppContext, + config::{RouterConfig, RoutingMode}, + middleware::TenantRequestMeta, + routers::{RouterFactory, RouterTrait}, + tenant::TenantKey, + worker::{BasicWorkerBuilder, ConnectionMode, RuntimeType, WorkerType}, +}; +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"; + +/// 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()), + }); + 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). +#[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.rl.enabled = true; + config.rl.control_timeout_secs = 5; + + common::create_test_context_with_tokenizer_registry(config, registry).await +} + +fn health_off() -> HealthCheckConfig { + HealthCheckConfig { + disable_health_check: true, + ..Default::default() + } +} + +/// Register a TokenSpeed gRPC worker whose RL capabilities come from labels. +/// `control_url` sets the `rl.control_url` label an engine would advertise; +/// `api_key` is the bearer the gateway presents to that control app. +#[expect( + clippy::expect_used, + reason = "test helper - panicking on failure is intentional" +)] +fn register_tokenspeed( + ctx: &Arc, + grpc_port: u16, + control_url: Option<&str>, + api_key: Option<&str>, +) { + let mut builder = BasicWorkerBuilder::new(format!("grpc://127.0.0.1:{grpc_port}")) + .worker_type(WorkerType::Regular) + .connection_mode(ConnectionMode::Grpc) + .runtime_type(RuntimeType::TokenSpeed) + .model(ModelCard::new(MODEL)) + .health_config(health_off()) + .label("weight_version", REGISTERED_VERSION) + .label("rl.pause_modes", "wait,abort,keep") + .label("rl.update_from", "disk,distributed") + .label("rl.abort", "true") + .label("rl.flush_cache", "true") + .label("rl.sleep_wake", "true") + .label("rl.reports_weight_version", "true"); + if let Some(url) = control_url { + builder = builder.label("rl.control_url", url); + } + if let Some(key) = api_key { + builder = builder.api_key(key); + } + ctx.worker_registry + .register(Arc::new(builder.build())) + .expect("TokenSpeed worker registered"); +} + +/// Register an HTTP SGLang worker, which controls itself over its own URL. +#[expect( + clippy::expect_used, + reason = "test helper - panicking on failure is intentional" +)] +fn register_http_sglang(ctx: &Arc, url: &str) { + let worker = BasicWorkerBuilder::new(url) + .worker_type(WorkerType::Regular) + .connection_mode(ConnectionMode::Http) + .runtime_type(RuntimeType::Sglang) + .model(ModelCard::new(MODEL)) + .health_config(health_off()) + .build(); + ctx.worker_registry + .register(Arc::new(worker)) + .expect("SGLang worker registered"); +} + +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) +} + +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 with a control endpoint at an HTTP mock +/// that records what it receives, one without) plus one HTTP SGLang mock. +struct Fleet { + ctx: Arc, + app: axum::Router, + router: Arc, + control_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(control_port: u16, sglang_port: u16) -> Fleet { + let control_recorder = RequestRecorder::new(); + set_request_recorder(control_port, control_recorder.clone()); + let mut control_app = MockWorker::new(mock_http(control_port)); + let control_url = control_app.start().await.unwrap(); + let mut sglang = MockWorker::new(mock_http(sglang_port)); + let sglang_url = sglang.start().await.unwrap(); + + let ctx = grpc_rl_context().await; + let advertised = BTreeMap::from([("rl.control_url".to_string(), control_url.clone())]); + let with_endpoint = start_mock_grpc_engine(advertised).await; + let without_endpoint = start_mock_grpc_engine(BTreeMap::new()).await; + register_tokenspeed( + &ctx, + with_endpoint, + Some(control_url.as_str()), + Some("ts-secret"), + ); + register_tokenspeed(&ctx, without_endpoint, None, None); + register_http_sglang(&ctx, &sglang_url); + + 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)); + Fleet { + ctx, + app, + router, + control_recorder, + control_url, + _control_app: control_app, + _sglang: sglang, + } +} + +/// 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<'a>(workers: &'a [Value], engine: &str, has_control: bool) -> &'a Value { + workers + .iter() + .find(|w| w["engine"] == engine && w["control_url"].is_string() == has_control) + .expect("worker present") +} + +#[tokio::test] +async fn discovery_reports_control_endpoint_and_advertised_capabilities() { + let f = fleet(18921, 18922).await; + let resp = f + .app + .clone() + .oneshot(Request::get("/v1/rl/workers").body(Body::empty()).unwrap()) + .await + .unwrap(); + assert_eq!(resp.status(), StatusCode::OK); + let body = json_of(resp).await; + assert_eq!(body["total"], 3); + let workers = body["workers"].as_array().unwrap(); + + let ts = by_engine(workers, "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"]) + ); + + let bare = by_engine(workers, "tokenspeed", false); + assert_eq!(bare["control_url"], Value::Null); + assert_eq!(bare["capabilities"]["source"], "label"); + + let sglang = by_engine(workers, "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(18923, 18924).await; + let workers = json_of( + f.app + .clone() + .oneshot(Request::get("/v1/rl/workers").body(Body::empty()).unwrap()) + .await + .unwrap(), + ) + .await; + let id = by_engine(workers["workers"].as_array().unwrap(), "tokenspeed", true)["id"] + .as_str() + .unwrap() + .to_string(); + + 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("Bearer ts-secret".to_string())], + "the worker's own key, not the caller's" + ); +} + +#[tokio::test] +async fn fanout_spans_transports_and_names_the_worker_without_an_endpoint() { + let f = fleet(18925, 18926).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"); + + // 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(18927, 18928).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}" + ); +} From 0d263490d8c1645491afbb098c4e8a4d41f11685 Mon Sep 17 00:00:00 2001 From: key4ng Date: Tue, 22 Sep 2026 03:32:47 -0700 Subject: [PATCH 08/24] feat(serve): wire the TokenSpeed control app for ZMQ workers and stamp rl.control_url when RL is on Signed-off-by: key4ng --- bindings/python/src/smg/serve.py | 91 ++++++++++++++++++++++++++++- bindings/python/tests/test_serve.py | 66 +++++++++++++++++++++ 2 files changed, 155 insertions(+), 2 deletions(-) diff --git a/bindings/python/src/smg/serve.py b/bindings/python/src/smg/serve.py index 6a20667b6d..090fcbb592 100644 --- a/bindings/python/src/smg/serve.py +++ b/bindings/python/src/smg/serve.py @@ -8,6 +8,7 @@ import argparse import atexit +import json import logging import os import random @@ -15,7 +16,9 @@ import socket import subprocess import sys +import threading import time +import urllib.request from abc import ABC, abstractmethod from smg.launch_router import launch_router @@ -83,6 +86,19 @@ def _zmq_handshake_port(ipc_url: str) -> int: return _ZMQ_HANDSHAKE_PORT_BASE + (h % _ZMQ_HANDSHAKE_PORT_SPAN) +def _rl_control_port(port: int) -> int: + """Port for a TokenSpeed worker's in-engine RL control app. + + Offset from the worker port so co-located workers (consecutive ports) and + the engine's own +233 distributed store never collide, reflected below the + u16 ceiling and hopped past SMG's ZMQ handshake band like ``dist_port``. + """ + p = port + 400 if port + 400 <= 65535 else port - 400 + if _ZMQ_HANDSHAKE_PORT_BASE <= p < _ZMQ_HANDSHAKE_PORT_BASE + _ZMQ_HANDSHAKE_PORT_SPAN: + p += _ZMQ_HANDSHAKE_PORT_SPAN + return p + + def _reject_handshake_port_collisions(ports: list[int]) -> None: """Fail before launch if two workers derive the same ZMQ handshake port. @@ -422,6 +438,10 @@ class TokenspeedWorkerLauncher(WorkerLauncher): def _get_tp_size(self, args: argparse.Namespace) -> int: return getattr(args, "tensor_parallel_size", 1) or 1 + def control_url(self, port: int) -> str: + """URL of the in-engine RL control app this launcher started for ``port``.""" + return f"http://127.0.0.1:{_rl_control_port(port)}" + def build_command( self, args: argparse.Namespace, backend_args: list[str], host: str, port: int ) -> list[str]: @@ -488,6 +508,10 @@ def _build_zmq_command( str(rpc_port), "--zmq-engine-index", "0", + "--rl-control-host", + "127.0.0.1", + "--rl-control-port", + str(_rl_control_port(port)), ] cmd.extend( self._backend_arg_defaults( @@ -517,6 +541,8 @@ def _build_zmq_command( "--data-parallel-address", "--data-parallel-rpc-port", "--zmq-engine-index", + "--rl-control-host", + "--rl-control-port", ], ) ) @@ -615,8 +641,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 @@ -976,6 +1000,52 @@ def parse_serve_args( _WORKER_SHUTDOWN_TIMEOUT = 30 +def _stamp_rl_control_labels( + gateway_url: str, + api_key: str | None, + targets: list[tuple[str, str]], + deadline_s: float, +) -> None: + """Label each ZMQ worker with its control endpoint once the gateway lists it. + + ZMQ discovery yields no labels, so the launcher, which owns both ends, + stamps ``rl.control_url`` through the worker update route (labels merge). + """ + headers = {"Content-Type": "application/json"} + if api_key: + headers["Authorization"] = f"Bearer {api_key}" + pending = dict(targets) + stop_at = time.monotonic() + deadline_s + while pending and time.monotonic() < stop_at: + try: + req = urllib.request.Request(f"{gateway_url}/workers", headers=headers, method="GET") + with urllib.request.urlopen(req, timeout=5) as resp: + workers = json.loads(resp.read()).get("workers", []) + except Exception as e: # noqa: BLE001 — the gateway may not be up yet + logger.debug("rl.control_url stamping: gateway not ready: %s", e) + time.sleep(1) + continue + for w in workers: + url = str(w.get("url", "")) + if url not in pending: + continue + body = json.dumps({"labels": {"rl.control_url": pending[url]}}).encode() + req = urllib.request.Request( + f"{gateway_url}/workers/{w['id']}", data=body, headers=headers, method="PATCH" + ) + try: + with urllib.request.urlopen(req, timeout=5): + pass + logger.info("stamped rl.control_url=%s on %s", pending[url], url) + del pending[url] + except Exception as e: # noqa: BLE001 + logger.warning("rl.control_url stamping failed for %s: %s", url, e) + if pending: + time.sleep(1) + for url in pending: + logger.warning("rl.control_url never stamped on %s (gateway did not list it)", url) + + class ServeOrchestrator: """Coordinate worker launch, health checking, router startup, and shutdown.""" @@ -998,6 +1068,23 @@ def run(self) -> None: self._launch_workers() self._wait_healthy() router_args = self._build_router_args() + if getattr(router_args, "enable_rl", False) and self.backend == "tokenspeed": + control = getattr(self.launcher, "control_url", None) + if callable(control): + targets = [ + ( + self.launcher.worker_url(self.args, self.args.worker_host, port), + control(port), + ) + for _, port in self.workers + ] + gateway_url = f"http://127.0.0.1:{router_args.port}" + threading.Thread( + target=_stamp_rl_control_labels, + args=(gateway_url, getattr(router_args, "api_key", None), targets, 300.0), + name="smg-rl-control-labels", + daemon=True, + ).start() launch_router(router_args) finally: self._cleanup_workers() diff --git a/bindings/python/tests/test_serve.py b/bindings/python/tests/test_serve.py index 0aa5a6f939..7287284808 100644 --- a/bindings/python/tests/test_serve.py +++ b/bindings/python/tests/test_serve.py @@ -1675,3 +1675,69 @@ def capture_launch(a, b, host, port, env): assert launched_envs[1]["CUDA_VISIBLE_DEVICES"] == "2,3" assert launched_envs[0]["PYTHONUNBUFFERED"] == "1" assert launched_envs[1]["PYTHONUNBUFFERED"] == "1" + + +def test_rl_control_port_avoids_neighbors_and_the_handshake_band(): + from smg.serve import _ZMQ_HANDSHAKE_PORT_BASE, _ZMQ_HANDSHAKE_PORT_SPAN, _rl_control_port + + assert _rl_control_port(30000) == 30400 + assert _rl_control_port(65500) == 65100 + inside = _ZMQ_HANDSHAKE_PORT_BASE - 400 + 10 + assert not ( + _ZMQ_HANDSHAKE_PORT_BASE + <= _rl_control_port(inside) + < _ZMQ_HANDSHAKE_PORT_BASE + _ZMQ_HANDSHAKE_PORT_SPAN + ) + + +def test_tokenspeed_zmq_command_wires_the_control_app(monkeypatch): + from types import SimpleNamespace + + from smg.serve import TokenspeedWorkerLauncher + + args = SimpleNamespace(model="/models/q", connection_mode="zmq", tensor_parallel_size=1) + cmd = TokenspeedWorkerLauncher().build_command(args, [], "127.0.0.1", 30000) + assert cmd[cmd.index("--rl-control-host") + 1] == "127.0.0.1" + assert cmd[cmd.index("--rl-control-port") + 1] == "30400" + assert TokenspeedWorkerLauncher().control_url(30000) == "http://127.0.0.1:30400" + + +def test_stamp_rl_control_labels_patches_each_worker(monkeypatch): + import json + + from smg import serve + + calls = [] + + class _Resp: + def __init__(self, payload): + self._payload = json.dumps(payload).encode() + self.status = 200 + + def read(self): + return self._payload + + def __enter__(self): + return self + + def __exit__(self, *a): + return False + + def fake_urlopen(req, timeout=0): + calls.append((req.get_method(), req.full_url, req.data, req.get_header("Authorization"))) + if req.get_method() == "GET": + return _Resp({"workers": [{"id": "w1", "url": "ipc:///tmp/engine-30000"}]}) + return _Resp({}) + + monkeypatch.setattr(serve.urllib.request, "urlopen", fake_urlopen) + serve._stamp_rl_control_labels( + "http://127.0.0.1:8000", + "adm", + [("ipc:///tmp/engine-30000", "http://127.0.0.1:30400")], + deadline_s=1.0, + ) + patch = [c for c in calls if c[0] == "PATCH"] + assert len(patch) == 1 + assert patch[0][1] == "http://127.0.0.1:8000/workers/w1" + assert json.loads(patch[0][2]) == {"labels": {"rl.control_url": "http://127.0.0.1:30400"}} + assert patch[0][3] == "Bearer adm" From 984fbc538478884ce7150abab6226516d186d4c8 Mon Sep 17 00:00:00 2001 From: key4ng Date: Tue, 22 Sep 2026 03:47:24 -0700 Subject: [PATCH 09/24] docs(rl): control endpoints, TokenSpeed drift rows, and the TokenSpeed launch guide Signed-off-by: key4ng --- crates/rl/NOTES.md | 4 +++ crates/rl/README.md | 22 ++++++++++++++--- docs/guides/rl-tokenspeed.md | 48 ++++++++++++++++++++++++++++++++++++ 3 files changed, 70 insertions(+), 4 deletions(-) create mode 100644 docs/guides/rl-tokenspeed.md diff --git a/crates/rl/NOTES.md b/crates/rl/NOTES.md index 62caeb4c02..f9b2baeeb8 100644 --- a/crates/rl/NOTES.md +++ b/crates/rl/NOTES.md @@ -7,3 +7,7 @@ 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=disk,distributed` | diff --git a/crates/rl/README.md b/crates/rl/README.md index 9a7641275a..27b9783b54 100644 --- a/crates/rl/README.md +++ b/crates/rl/README.md @@ -14,13 +14,27 @@ 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. 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. A wildcard bind host in +the advertised URL (`0.0.0.0`, `::`) 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`; see `docs/guides/rl-tokenspeed.md`. ## Python client diff --git a/docs/guides/rl-tokenspeed.md b/docs/guides/rl-tokenspeed.md new file mode 100644 index 0000000000..0f5b92103c --- /dev/null +++ b/docs/guides/rl-tokenspeed.md @@ -0,0 +1,48 @@ +# RL rollouts on TokenSpeed behind SMG + +SMG's RL control plane (`--enable-rl`, `/v1/rl/*`) drives TokenSpeed engines +through their in-engine SGLang-compatible control app 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 0.0.0.0 --rl-control-port 30400 --rl-control-api-key "$RL_KEY" \ + --enable-output-logprobs + +The gateway: + + 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 + +Register the engines with their control key so the proxy authenticates: + + curl -X POST http://smg:30000/workers -H 'content-type: application/json' \ + -d '{"url":"grpc://rollout-1:30000","api_key":"'"$RL_KEY"'"}' + +`GET /v1/rl/workers` then shows `engine: tokenspeed`, `connection_mode: grpc`, +`control_url: http://rollout-1:30400` and the engine's advertised capabilities. + +## Refit + + python3 examples/rl/refit_from_disk.py --smg http://smg:30000 \ + --model-path /ckpt/step-42 --weight-version 42 --selector engine=tokenspeed + +The next `/generate` through SMG reports `meta_info.weight_version: "42"`, +stamped by the engine on the gRPC response. + +## Security + +The control app accepts weight updates from anyone who can reach it. Bind it +on a routable host only with `--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`, `disk` and `distributed`) and need the label supplied at +registration: `{"url":"grpc://…","labels":{"rl.control_url":"http://…:30400"}}`. +See `crates/rl/NOTES.md` for their route-level drift. From bb6243c31f6d52c03c32f44e4d2b01dafeeb5894 Mon Sep 17 00:00:00 2001 From: key4ng Date: Tue, 22 Sep 2026 04:48:31 -0700 Subject: [PATCH 10/24] fix(rl): treat a blank rl.control_url label as absent; document base_url An operator PATCH can set rl.control_url to an empty or whitespace-only string, which previously produced control_url: Some("") in discovery and a 502 upstream_unreachable instead of the 422 no_control_endpoint the situation deserves. Filter blank labels the same way capability.rs and ProtoGenerateComplete::weight_version already treat empty as unset. Also fix the base_url doc comment in the wire type, which still claimed to be the address control calls are sent to; that is now control_url's job, and this field's twin in crates/rl/src/view.rs was already fixed. Signed-off-by: key4ng --- crates/protocols/src/rl.rs | 2 +- model_gateway/src/rl_adapter.rs | 20 +++++++++++++++++++- 2 files changed, 20 insertions(+), 2 deletions(-) diff --git a/crates/protocols/src/rl.rs b/crates/protocols/src/rl.rs index a7e198ba3c..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, diff --git a/model_gateway/src/rl_adapter.rs b/model_gateway/src/rl_adapter.rs index 15f1e3b1e8..e3a5b3cd45 100644 --- a/model_gateway/src/rl_adapter.rs +++ b/model_gateway/src/rl_adapter.rs @@ -47,7 +47,12 @@ impl RegistryRlView { .unwrap_or_else(|| Arc::new(worker.http_client().clone())); return (Some(worker.base_url().to_string()), Some(client)); } - let Some(advertised) = spec.labels.get(CONTROL_URL_LABEL) else { + let Some(advertised) = spec + .labels + .get(CONTROL_URL_LABEL) + .map(String::as_str) + .filter(|v| !v.trim().is_empty()) + else { return (None, None); }; let url = resolve_control_url(advertised, worker.url()); @@ -181,6 +186,19 @@ mod tests { assert!(info.control_client.is_none()); } + #[test] + 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 info = view_over(vec![worker]).list().pop().expect("one worker"); + assert!(info.control_url.is_none()); + assert!(info.control_client.is_none()); + } + #[test] fn cached_control_clients_are_shared_across_workers_with_the_same_pool_config() { let mk = |url: &str| -> Arc { From 0216a2c744cde859f6aa0f6e6c2ff7b79ce14070 Mon Sep 17 00:00:00 2001 From: key4ng Date: Tue, 22 Sep 2026 04:49:04 -0700 Subject: [PATCH 11/24] fix(serve): stop retrying 4xx when stamping rl.control_url; require zmq _stamp_rl_control_labels retried every PATCH failure, including a permanent 401/403/404, once a second for the full 300s deadline, logging a warning each time. Give up on the target immediately for a 4xx status and keep the existing retry behavior for everything else. Also gate the stamping thread on connection_mode == "zmq" explicitly. This was previously safe only because TokenspeedWorkerLauncher.build_command raises for any other mode, so a gRPC launch died before the orchestrator reached the stamping code; make the coupling explicit instead of load-bearing-by-accident. The new tests exercise smg.serve, which needs the native smg.smg_rs extension; they were verified by careful reading and could not be run in this environment (ModuleNotFoundError: No module named 'smg.smg_rs'). Signed-off-by: key4ng --- bindings/python/src/smg/serve.py | 15 +++++++++- bindings/python/tests/test_serve.py | 44 +++++++++++++++++++++++++++++ 2 files changed, 58 insertions(+), 1 deletion(-) diff --git a/bindings/python/src/smg/serve.py b/bindings/python/src/smg/serve.py index 090fcbb592..f279ac6c2b 100644 --- a/bindings/python/src/smg/serve.py +++ b/bindings/python/src/smg/serve.py @@ -18,6 +18,7 @@ import sys import threading import time +import urllib.error import urllib.request from abc import ABC, abstractmethod @@ -1038,6 +1039,14 @@ def _stamp_rl_control_labels( pass logger.info("stamped rl.control_url=%s on %s", pending[url], url) del pending[url] + except urllib.error.HTTPError as e: + if 400 <= e.code < 500: + logger.warning( + "rl.control_url stamping got HTTP %s for %s; not retrying", e.code, url + ) + del pending[url] + else: + logger.warning("rl.control_url stamping failed for %s: %s", url, e) except Exception as e: # noqa: BLE001 logger.warning("rl.control_url stamping failed for %s: %s", url, e) if pending: @@ -1068,7 +1077,11 @@ def run(self) -> None: self._launch_workers() self._wait_healthy() router_args = self._build_router_args() - if getattr(router_args, "enable_rl", False) and self.backend == "tokenspeed": + if ( + getattr(router_args, "enable_rl", False) + and self.backend == "tokenspeed" + and getattr(self.args, "connection_mode", "grpc") == "zmq" + ): control = getattr(self.launcher, "control_url", None) if callable(control): targets = [ diff --git a/bindings/python/tests/test_serve.py b/bindings/python/tests/test_serve.py index 7287284808..3db7e3f1c4 100644 --- a/bindings/python/tests/test_serve.py +++ b/bindings/python/tests/test_serve.py @@ -1741,3 +1741,47 @@ def fake_urlopen(req, timeout=0): assert patch[0][1] == "http://127.0.0.1:8000/workers/w1" assert json.loads(patch[0][2]) == {"labels": {"rl.control_url": "http://127.0.0.1:30400"}} assert patch[0][3] == "Bearer adm" + + +def test_stamp_rl_control_labels_stops_retrying_after_a_4xx(monkeypatch): + import json + import time + import urllib.error + + from smg import serve + + calls = [] + + class _Resp: + def __init__(self, payload): + self._payload = json.dumps(payload).encode() + self.status = 200 + + def read(self): + return self._payload + + def __enter__(self): + return self + + def __exit__(self, *a): + return False + + def fake_urlopen(req, timeout=0): + calls.append((req.get_method(), req.full_url, req.data, req.get_header("Authorization"))) + if req.get_method() == "GET": + return _Resp({"workers": [{"id": "w1", "url": "ipc:///tmp/engine-30000"}]}) + raise urllib.error.HTTPError(req.full_url, 401, "unauthorized", {}, None) + + monkeypatch.setattr(serve.urllib.request, "urlopen", fake_urlopen) + start = time.monotonic() + serve._stamp_rl_control_labels( + "http://127.0.0.1:8000", + "adm", + [("ipc:///tmp/engine-30000", "http://127.0.0.1:30400")], + deadline_s=5.0, + ) + elapsed = time.monotonic() - start + + patch = [c for c in calls if c[0] == "PATCH"] + assert len(patch) == 1, "a 4xx must not be retried" + assert elapsed < 3.0, "giving up on the 4xx must not wait out the deadline" From 33ea55cce27580f2945b1a924cc9ad4c1b5b88f0 Mon Sep 17 00:00:00 2001 From: key4ng Date: Tue, 22 Sep 2026 04:49:22 -0700 Subject: [PATCH 12/24] docs(rl): HTTP workers ignore rl.control_url; narrow recorder doc crates/rl/README.md never stated that an rl.control_url label on an HTTP worker is ignored (it always controls through itself); one clause closes that. MockWorker::authorizations() claimed to be index-aligned with bodies(), but record_authorization is only called from the RL control handler while record_request has five call sites, so the alignment does not hold. Narrow the doc instead of changing the recorder's behavior. Signed-off-by: key4ng --- crates/rl/README.md | 5 +++-- model_gateway/tests/common/mock_worker.rs | 7 ++++--- 2 files changed, 7 insertions(+), 5 deletions(-) diff --git a/crates/rl/README.md b/crates/rl/README.md index 27b9783b54..3256050da0 100644 --- a/crates/rl/README.md +++ b/crates/rl/README.md @@ -17,8 +17,9 @@ reports `protocol_version` (currently 1), bumped only for incompatible changes. ## Control endpoints Control calls go to a worker's **control endpoint**, not necessarily its data -transport. An HTTP worker is controlled through itself. A gRPC or ZMQ worker -needs the `rl.control_url` label: TokenSpeed engines advertise it in server +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. A wildcard bind host in the advertised URL (`0.0.0.0`, `::`) is replaced by the worker's own host. The diff --git a/model_gateway/tests/common/mock_worker.rs b/model_gateway/tests/common/mock_worker.rs index 1273ee9d33..debb868041 100755 --- a/model_gateway/tests/common/mock_worker.rs +++ b/model_gateway/tests/common/mock_worker.rs @@ -190,9 +190,10 @@ impl RequestRecorder { .clone() } - /// The `authorization` header each recorded RL control request carried, - /// oldest first and index-aligned with [`Self::bodies`]. `None` is a - /// request that arrived without the header. + /// The `authorization` header of each RL control request received, + /// oldest first (RL control routes only; not aligned with + /// [`Self::bodies`]). `None` is a request that arrived without the + /// header. #[expect( clippy::expect_used, reason = "test helper - panicking on failure is intentional" From 2952df9e2d805205d693faa05934857391efbc96 Mon Sep 17 00:00:00 2001 From: key4ng Date: Tue, 22 Sep 2026 12:48:36 -0700 Subject: [PATCH 13/24] style(rl): drop redundant serde_json qualifications in the discovery tests Signed-off-by: key4ng --- crates/rl/src/discovery.rs | 16 ++++++---------- 1 file changed, 6 insertions(+), 10 deletions(-) diff --git a/crates/rl/src/discovery.rs b/crates/rl/src/discovery.rs index f82cc19c79..83c7150cd2 100644 --- a/crates/rl/src/discovery.rs +++ b/crates/rl/src/discovery.rs @@ -227,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); @@ -236,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); } @@ -264,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); } @@ -282,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"); @@ -340,7 +336,7 @@ mod tests { .unwrap() }; assert_eq!(by_id("g1")["control_url"], "http://a:40100"); - assert_eq!(by_id("n1")["control_url"], serde_json::Value::Null); + assert_eq!(by_id("n1")["control_url"], Value::Null); assert_eq!(by_id("g1")["engine"], "tokenspeed"); } } From bf6a9a2ebf825b34a702949c57a6acb7ca317fe4 Mon Sep 17 00:00:00 2001 From: key4ng Date: Tue, 22 Sep 2026 13:01:11 -0700 Subject: [PATCH 14/24] fix(rl): TokenSpeed refits are distributed-only; drop disk from the static row and docs The companion TokenSpeed branch's scheduler has no receive path for a disk or tensor weight update: update_weights_from_disk and update_weights_from_tensor both answer HTTP 501 and advertise rl.update_from=distributed. Only the trainer-driven NCCL broadcast (update_weights_from_distributed) works. Drop 'disk' from TokenSpeed's static capability row so pre-advertisement builds report only what they can actually do, and update the NOTES.md drift log, crates/rl/README.md, and docs/guides/rl-tokenspeed.md to point callers at the distributed refit path instead of refit_from_disk.py. Signed-off-by: key4ng --- crates/rl/NOTES.md | 1 + crates/rl/README.md | 2 +- crates/rl/src/capability.rs | 11 +++++++---- docs/guides/rl-tokenspeed.md | 27 +++++++++++++++++++++------ 4 files changed, 30 insertions(+), 11 deletions(-) diff --git a/crates/rl/NOTES.md b/crates/rl/NOTES.md index f9b2baeeb8..deae294812 100644 --- a/crates/rl/NOTES.md +++ b/crates/rl/NOTES.md @@ -11,3 +11,4 @@ planning docs, with date, engine version, and what was done. | 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=disk,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) | diff --git a/crates/rl/README.md b/crates/rl/README.md index 3256050da0..9b761b5c2b 100644 --- a/crates/rl/README.md +++ b/crates/rl/README.md @@ -35,7 +35,7 @@ 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`; see `docs/guides/rl-tokenspeed.md`. +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` answers 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 8fde331679..383e6e8df7 100644 --- a/crates/rl/src/capability.rs +++ b/crates/rl/src/capability.rs @@ -34,7 +34,7 @@ fn static_for(runtime: RuntimeType) -> RlCapabilities { RuntimeType::TokenSpeed => RlCapabilities { source: RlCapabilitySource::Static, pause_modes: strings(&["wait", "abort"]), - update_from: strings(&["disk", "distributed"]), + update_from: strings(&["distributed"]), abort: true, flush_cache: true, sleep_wake: true, @@ -178,13 +178,15 @@ mod tests { /// 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. + /// 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, ["disk", "distributed"]); + assert_eq!(t.update_from, ["distributed"]); assert!(t.abort && t.flush_cache && t.sleep_wake && t.reports_weight_version); } @@ -194,10 +196,11 @@ mod tests { RuntimeType::TokenSpeed, &labels(&[ ("rl.pause_modes", "wait,abort,keep"), - ("rl.update_from", "disk,distributed"), + ("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/docs/guides/rl-tokenspeed.md b/docs/guides/rl-tokenspeed.md index 0f5b92103c..0cc7222358 100644 --- a/docs/guides/rl-tokenspeed.md +++ b/docs/guides/rl-tokenspeed.md @@ -28,11 +28,26 @@ Register the engines with their control key so the proxy authenticates: ## Refit - python3 examples/rl/refit_from_disk.py --smg http://smg:30000 \ - --model-path /ckpt/step-42 --weight-version 42 --selector engine=tokenspeed - -The next `/generate` through SMG reports `meta_info.weight_version: "42"`, -stamped by the engine on the gRPC response. +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`. ## Security @@ -43,6 +58,6 @@ as the worker's `api_key`. ## Older engines Engines that predate advertisement get SMG's static capability row (`wait` -and `abort`, `disk` and `distributed`) and need the label supplied at +and `abort`, `distributed` only) and need the label supplied at registration: `{"url":"grpc://…","labels":{"rl.control_url":"http://…:30400"}}`. See `crates/rl/NOTES.md` for their route-level drift. From dbd178e9678ceee42138504403cb122347622add Mon Sep 17 00:00:00 2001 From: key4ng Date: Tue, 22 Sep 2026 14:49:00 -0700 Subject: [PATCH 15/24] fix(rl): fall back to attn_tp_size when an engine reports no tp_size TokenSpeed spells its tensor-parallel width attn_tp_size, which TOKENSPEED_GRPC_KEYS already lifts into the worker's labels, but discovery only read tp_size -- so a TokenSpeed engine reported tp_size null and a trainer laying out NCCL ranks had nothing to go on. Read attn_tp_size when tp_size is absent or unparseable; an explicit tp_size still wins. Signed-off-by: key4ng --- crates/rl/src/discovery.rs | 33 ++++++++++++++++++++++++++++++++- docs/guides/rl-tokenspeed.md | 5 +++++ 2 files changed, 37 insertions(+), 1 deletion(-) diff --git a/crates/rl/src/discovery.rs b/crates/rl/src/discovery.rs index 83c7150cd2..2234c21624 100644 --- a/crates/rl/src/discovery.rs +++ b/crates/rl/src/discovery.rs @@ -109,7 +109,12 @@ pub fn entry(info: &RlWorkerInfo, dp_ranks: usize) -> RlWorkerEntry { 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"), + // TokenSpeed spells its tensor-parallel width `attn_tp_size` and leaves + // it unset unless the operator asks for one, so a plain launch reports + // no `tp_size` at all. A trainer laying out NCCL ranks needs the real + // width, and a wrong one deadlocks the weight-update group. + tp_size: int_label(&info.labels, "tp_size") + .or_else(|| int_label(&info.labels, "attn_tp_size")), dp_size: int_label(&info.labels, "dp_size"), pp_size: int_label(&info.labels, "pp_size"), dp_ranks, @@ -257,6 +262,32 @@ mod tests { assert_eq!(m["tp_size"], "1"); } + #[test] + fn tp_size_falls_back_to_tokenspeeds_attn_tp_size() { + // A TokenSpeed engine launched without an explicit parallelism flag + // reports `attn_tp_size` and no `tp_size`; a trainer needs the width to + // lay out its NCCL ranks. + let mut w = worker("w1", "http://a:1", RuntimeType::TokenSpeed); + w.labels.remove("tp_size"); + w.labels.insert("attn_tp_size".to_string(), "2".to_string()); + assert_eq!(entry(&w, 1).tp_size, Some(2)); + } + + #[test] + fn an_explicit_tp_size_wins_over_attn_tp_size() { + let mut w = worker("w1", "http://a:1", RuntimeType::TokenSpeed); + w.labels.insert("tp_size".to_string(), "4".to_string()); + w.labels.insert("attn_tp_size".to_string(), "2".to_string()); + assert_eq!(entry(&w, 1).tp_size, Some(4)); + } + + #[test] + fn tp_size_is_null_when_neither_label_is_present() { + let mut w = worker("w1", "http://a:1", RuntimeType::TokenSpeed); + w.labels.remove("tp_size"); + assert_eq!(entry(&w, 1).tp_size, None); + } + #[tokio::test] async fn list_reports_the_protocol_version() { let app = crate::router::<()>(state(vec![worker("w1", "http://a:1", RuntimeType::Sglang)])); diff --git a/docs/guides/rl-tokenspeed.md b/docs/guides/rl-tokenspeed.md index 0cc7222358..1c9300d293 100644 --- a/docs/guides/rl-tokenspeed.md +++ b/docs/guides/rl-tokenspeed.md @@ -48,6 +48,11 @@ with `pause_generation` / `continue_generation` fanned out around the 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`: discovery reads it from the engine's +server args and falls back to TokenSpeed's own spelling, `attn_tp_size`, so an +engine launched with either reports a width. An engine launched with neither +(TokenSpeed leaves `attn_tp_size` unset unless asked) reports `tp_size: null`, +and a trainer must then assume 1 or be told. ## Security From be35c03d47282a91126750ff0cd40a1ab914a626 Mon Sep 17 00:00:00 2001 From: key4ng Date: Wed, 23 Sep 2026 15:34:43 -0700 Subject: [PATCH 16/24] docs(rl): list the gateway-side files that carry RL data and name both refused refit routes Signed-off-by: key4ng --- crates/rl/COUPLING.md | 16 ++++++++++++++++ crates/rl/README.md | 2 +- 2 files changed, 17 insertions(+), 1 deletion(-) diff --git a/crates/rl/COUPLING.md b/crates/rl/COUPLING.md index 3c97793146..597f63571b 100644 --- a/crates/rl/COUPLING.md +++ b/crates/rl/COUPLING.md @@ -22,6 +22,22 @@ 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`; `src/routers/grpc/router.rs` and +`src/routers/grpc/pipeline.rs` give slime's model-less single-prompt +`/generate` SGLang's shape; `src/routers/http/router.rs` (with +`crates/protocols/src/generate.rs`) stops forwarding the wildcard `model` +placeholder. `model_gateway/tests/rl_tokenspeed_control_endpoint_test.rs` +drives a TokenSpeed gRPC worker through `/v1/rl` via its control endpoint and +checks the generate shape. + 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/README.md b/crates/rl/README.md index 9b761b5c2b..a8330b7919 100644 --- a/crates/rl/README.md +++ b/crates/rl/README.md @@ -35,7 +35,7 @@ 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` answers HTTP 501 on TokenSpeed. See `docs/guides/rl-tokenspeed.md`. +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 From 4b943f0b85c9c39c50f9b339711964c24706c541 Mon Sep 17 00:00:00 2001 From: key4ng Date: Wed, 23 Sep 2026 15:35:25 -0700 Subject: [PATCH 17/24] test(mock_worker): set the new mock fields in the context-length test that landed on main Signed-off-by: key4ng --- model_gateway/tests/grpc_context_length_test.rs | 2 ++ 1 file changed, 2 insertions(+) diff --git a/model_gateway/tests/grpc_context_length_test.rs b/model_gateway/tests/grpc_context_length_test.rs index 6fad2b926e..d84b377139 100644 --- a/model_gateway/tests/grpc_context_length_test.rs +++ b/model_gateway/tests/grpc_context_length_test.rs @@ -68,6 +68,8 @@ async fn start_mock_grpc_worker() -> u16 { output_tokens: 2, realistic: false, engine: mock_worker::engine::EngineParams::default(), + server_args: std::collections::BTreeMap::new(), + weight_version: None, replay: Default::default(), }); tokio::spawn(mock_worker::grpc::serve_with_listener(cfg, listener)); From a2104bfbedb2b51ed48018647fef75b28adf297d Mon Sep 17 00:00:00 2001 From: key4ng Date: Wed, 23 Sep 2026 15:35:25 -0700 Subject: [PATCH 18/24] style(rl): apply rustfmt to the control-endpoint resolver Signed-off-by: key4ng --- crates/rl/src/control.rs | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/crates/rl/src/control.rs b/crates/rl/src/control.rs index bfbd951d9a..0e8a1eb51d 100644 --- a/crates/rl/src/control.rs +++ b/crates/rl/src/control.rs @@ -43,7 +43,10 @@ fn is_wildcard_host(host: &str) -> bool { /// `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(']')) { + 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); From 69ac411360ae3ac722985209e5eac6098fe7f462 Mon Sep 17 00:00:00 2001 From: key4ng Date: Tue, 6 Oct 2026 11:52:16 -0700 Subject: [PATCH 19/24] fix(rl): hold control clients strongly and carry a client build failure into the 422 The gateway's worker client cache keeps weak handles, and a gRPC or ZMQ worker holds no HTTP client of its own, so every discovery and every control call on such a fleet rebuilt a reqwest client (TLS roots parsed again each time). RegistryRlView now keeps one strong handle per pool config, shared by the workers that use it and pruned when no registered non-HTTP worker needs it; HttpPoolConfig derives Eq and Hash for the key. RlWorkerInfo.control_client is a Result: a failed build is logged once and its reason rides on the worker, so the 422 names it instead of pointing at the gateway log. The no-label hint now says what the usual cause is (a wildcard --rl-control-host, which the engine does not advertise) rather than suggesting an engine upgrade. resolve_control_url takes the worker's base URL and leaves the label unchanged when the worker has no host (an ipc:// ZMQ worker), instead of producing http://:port; the control module is private, its one function re-exported. The attn_tp_size fallback in discovery is gone: normalize_grpc_keys already folds it into tp_size before labels reach the crate, so the code and its three tests exercised nothing. Signed-off-by: key4ng --- crates/protocols/src/worker.rs | 4 +- crates/rl/src/control.rs | 26 +++-- crates/rl/src/discovery.rs | 34 +------ crates/rl/src/fanout.rs | 1 - crates/rl/src/lib.rs | 2 +- crates/rl/src/proxy.rs | 20 ++-- crates/rl/src/testing.rs | 2 +- crates/rl/src/view.rs | 15 +-- model_gateway/src/rl_adapter.rs | 174 +++++++++++++++++++++++--------- 9 files changed, 167 insertions(+), 111 deletions(-) diff --git a/crates/protocols/src/worker.rs b/crates/protocols/src/worker.rs index fcc327cd03..67aa8aceab 100644 --- a/crates/protocols/src/worker.rs +++ b/crates/protocols/src/worker.rs @@ -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, diff --git a/crates/rl/src/control.rs b/crates/rl/src/control.rs index 0e8a1eb51d..4aaba2977d 100644 --- a/crates/rl/src/control.rs +++ b/crates/rl/src/control.rs @@ -3,10 +3,12 @@ //! (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 URL. +/// Resolve an advertised control URL against the worker's own base URL. /// -/// An engine that bound `0.0.0.0` (or `::`) advertises what it bound; the -/// only host the gateway knows is reachable is the worker's own, so use it. +/// 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 { @@ -24,7 +26,9 @@ pub fn resolve_control_url(advertised: &str, worker_url: &str) -> String { .next() .map(|authority| split_host_port(authority).0) .unwrap_or(""); - let worker_host = worker_host.split('@').next().unwrap_or(worker_host); + if worker_host.is_empty() { + return advertised.to_string(); + } let mut out = format!("{scheme}://{worker_host}"); if let Some(port) = port { out.push(':'); @@ -84,9 +88,17 @@ mod tests { "http://[fd00::5]:40100" ); assert_eq!( - resolve_control_url("http://:40100", "grpc://10.0.0.5:30000@2"), - "http://10.0.0.5:40100", - "a DP-rank suffix on the worker URL is not part of the host" + 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" ); } diff --git a/crates/rl/src/discovery.rs b/crates/rl/src/discovery.rs index 2234c21624..b912fa6572 100644 --- a/crates/rl/src/discovery.rs +++ b/crates/rl/src/discovery.rs @@ -109,12 +109,7 @@ pub fn entry(info: &RlWorkerInfo, dp_ranks: usize) -> RlWorkerEntry { worker_type: enum_str(&info.worker_type), connection_mode: enum_str(&info.connection_mode), control_url: info.control_url.clone(), - // TokenSpeed spells its tensor-parallel width `attn_tp_size` and leaves - // it unset unless the operator asks for one, so a plain launch reports - // no `tp_size` at all. A trainer laying out NCCL ranks needs the real - // width, and a wrong one deadlocks the weight-update group. - tp_size: int_label(&info.labels, "tp_size") - .or_else(|| int_label(&info.labels, "attn_tp_size")), + tp_size: int_label(&info.labels, "tp_size"), dp_size: int_label(&info.labels, "dp_size"), pp_size: int_label(&info.labels, "pp_size"), dp_ranks, @@ -262,32 +257,6 @@ mod tests { assert_eq!(m["tp_size"], "1"); } - #[test] - fn tp_size_falls_back_to_tokenspeeds_attn_tp_size() { - // A TokenSpeed engine launched without an explicit parallelism flag - // reports `attn_tp_size` and no `tp_size`; a trainer needs the width to - // lay out its NCCL ranks. - let mut w = worker("w1", "http://a:1", RuntimeType::TokenSpeed); - w.labels.remove("tp_size"); - w.labels.insert("attn_tp_size".to_string(), "2".to_string()); - assert_eq!(entry(&w, 1).tp_size, Some(2)); - } - - #[test] - fn an_explicit_tp_size_wins_over_attn_tp_size() { - let mut w = worker("w1", "http://a:1", RuntimeType::TokenSpeed); - w.labels.insert("tp_size".to_string(), "4".to_string()); - w.labels.insert("attn_tp_size".to_string(), "2".to_string()); - assert_eq!(entry(&w, 1).tp_size, Some(4)); - } - - #[test] - fn tp_size_is_null_when_neither_label_is_present() { - let mut w = worker("w1", "http://a:1", RuntimeType::TokenSpeed); - w.labels.remove("tp_size"); - assert_eq!(entry(&w, 1).tp_size, None); - } - #[tokio::test] async fn list_reports_the_protocol_version() { let app = crate::router::<()>(state(vec![worker("w1", "http://a:1", RuntimeType::Sglang)])); @@ -350,7 +319,6 @@ mod tests { let mut n = worker("n1", "grpc://b:1", RuntimeType::TokenSpeed); n.connection_mode = ConnectionMode::Grpc; n.control_url = None; - n.control_client = None; let app = crate::router::<()>(state(vec![g, n])); let resp = app .oneshot(Request::get("/workers").body(Body::empty()).unwrap()) diff --git a/crates/rl/src/fanout.rs b/crates/rl/src/fanout.rs index 1bf47e0812..c836571c3d 100644 --- a/crates/rl/src/fanout.rs +++ b/crates/rl/src/fanout.rs @@ -228,7 +228,6 @@ mod tests { let mut grpc = worker("g1", &format!("{}/grpc", good.url), RuntimeType::Sglang); grpc.connection_mode = ConnectionMode::Grpc; grpc.control_url = None; - grpc.control_client = None; let app = crate::router::<()>(state( vec![ worker("w1", &good.url, RuntimeType::Sglang), diff --git a/crates/rl/src/lib.rs b/crates/rl/src/lib.rs index 11558c9335..12f8a00e1b 100644 --- a/crates/rl/src/lib.rs +++ b/crates/rl/src/lib.rs @@ -4,7 +4,7 @@ pub mod capability; pub mod config; -pub mod control; +mod control; pub mod discovery; pub mod error; pub mod fanout; diff --git a/crates/rl/src/proxy.rs b/crates/rl/src/proxy.rs index 7d013cbc47..fafdbfd334 100644 --- a/crates/rl/src/proxy.rs +++ b/crates/rl/src/proxy.rs @@ -139,15 +139,15 @@ pub async fn call_worker( let Some(base) = worker.control_url.as_deref() else { return Err(no_control_endpoint( worker, - "no `rl.control_url` label; register the worker with one or upgrade the engine to advertise it", + "no `rl.control_url` label: launch the engine with a routable --rl-control-host (a wildcard bind is not advertised), or set the label at registration", )); }; - let Some(client) = worker.control_client.as_ref() else { - return Err(no_control_endpoint( + let client = worker.control_client.as_ref().map_err(|e| { + no_control_endpoint( worker, - "the gateway could not build an HTTP client for the control endpoint (see the gateway log)", - )); - }; + &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 @@ -418,7 +418,6 @@ mod tests { let mut grpc = worker("g1", &engine.url, RuntimeType::Sglang); grpc.connection_mode = ConnectionMode::Grpc; grpc.control_url = None; - grpc.control_client = None; let app = crate::router::<()>(state( vec![worker("w1", &engine.url, RuntimeType::Sglang), grpc], 5, @@ -505,7 +504,7 @@ mod tests { 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 = None; + w.control_client = Err("bad CA bundle".to_string()); let app = crate::router::<()>(state(vec![w], 5)); let r = app .oneshot( @@ -518,7 +517,8 @@ mod tests { assert_eq!(r.status(), StatusCode::UNPROCESSABLE_ENTITY); let body = json_body(r).await; assert_eq!(body["error"], "no_control_endpoint"); - assert!(body["message"].as_str().unwrap().contains("HTTP client")); + 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 @@ -527,7 +527,7 @@ mod tests { 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.control_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/testing.rs b/crates/rl/src/testing.rs index cf420d308c..ac6f038828 100644 --- a/crates/rl/src/testing.rs +++ b/crates/rl/src/testing.rs @@ -52,7 +52,7 @@ pub fn worker(id: &str, url: &str, runtime: RuntimeType) -> RlWorkerInfo { dp_size: None, labels, control_url: Some(base_url), - control_client: Some(test_client()), + control_client: Ok(test_client()), } } diff --git a/crates/rl/src/view.rs b/crates/rl/src/view.rs index ef2387988c..fae825f745 100644 --- a/crates/rl/src/view.rs +++ b/crates/rl/src/view.rs @@ -29,12 +29,13 @@ pub struct RlWorkerInfo { /// 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 used 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 cached - /// client with the same TLS settings on HTTP/1.1. `None` when the worker - /// has no control endpoint. - pub control_client: 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 { @@ -88,7 +89,7 @@ mod tests { dp_size: None, labels: HashMap::new(), control_url: None, - control_client: None, + control_client: Err("no client".to_string()), } } diff --git a/model_gateway/src/rl_adapter.rs b/model_gateway/src/rl_adapter.rs index e3a5b3cd45..8b8ab0f7bd 100644 --- a/model_gateway/src/rl_adapter.rs +++ b/model_gateway/src/rl_adapter.rs @@ -1,9 +1,12 @@ //! 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 openai_protocol::worker::{ConnectionMode, HttpPoolConfig}; use smg_rl::{resolve_control_url, RlState, RlWorkerInfo, RlWorkerView}; use tracing::warn; @@ -19,6 +22,13 @@ pub const CONTROL_URL_LABEL: &str = "rl.control_url"; 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 { @@ -26,46 +36,68 @@ impl RegistryRlView { 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`; the client comes from the shared - /// cache (same TLS identity and roots as every upstream client, HTTP/1.1 - /// because the control apps are uvicorn). + /// 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, Option>) { + ) -> (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()), Some(client)); + return (Some(worker.base_url().to_string()), Ok(client)); } - let Some(advertised) = spec + let url = spec .labels .get(CONTROL_URL_LABEL) .map(String::as_str) .filter(|v| !v.trim().is_empty()) - else { - return (None, None); - }; - let url = resolve_control_url(advertised, worker.url()); - match self.client_cache.get(&spec.http_pool, false) { - Ok(client) => (Some(url), Some(client)), - Err(e) => { - warn!( - worker = %worker.url(), error = %e, - "no HTTP client for the RL control endpoint" - ); - (Some(url), None) - } - } + .map(|advertised| resolve_control_url(advertised, worker.base_url())); + (url, self.control_client(&spec.http_pool)) } fn info(&self, worker: &Arc) -> Option { @@ -93,11 +125,10 @@ impl RegistryRlView { 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 { @@ -132,15 +163,29 @@ mod tests { use super::*; use crate::{config::RouterConfig, worker::BasicWorkerBuilder}; - fn view_over(workers: Vec>) -> RegistryRlView { + fn view_with(config: &RouterConfig, workers: Vec>) -> RegistryRlView { let registry = Arc::new(WorkerRegistry::new()); for w in workers { registry.register(w); } - let cache = Arc::new(WorkerHttpClientCache::new(&RouterConfig::default())); + let cache = Arc::new(WorkerHttpClientCache::new(config)); RegistryRlView::new(registry, cache) } + fn view_over(workers: Vec>) -> RegistryRlView { + view_with(&RouterConfig::default(), workers) + } + + /// 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] @@ -158,20 +203,17 @@ mod tests { } #[test] - fn grpc_worker_with_a_label_gets_a_control_endpoint_and_a_cached_client() { - let worker: Arc = Arc::new( - BasicWorkerBuilder::new("grpc://10.0.0.5:30000") - .connection_mode(ConnectionMode::Grpc) - .label("rl.control_url", "http://0.0.0.0:40100") - .build(), - ); - let info = view_over(vec![worker]).list().pop().expect("one worker"); + 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_some()); + assert!(info.control_client.is_ok()); } #[test] @@ -183,7 +225,6 @@ mod tests { ); let info = view_over(vec![worker]).list().pop().expect("one worker"); assert!(info.control_url.is_none()); - assert!(info.control_client.is_none()); } #[test] @@ -196,22 +237,55 @@ mod tests { ); let info = view_over(vec![worker]).list().pop().expect("one worker"); assert!(info.control_url.is_none()); - assert!(info.control_client.is_none()); } #[test] - fn cached_control_clients_are_shared_across_workers_with_the_same_pool_config() { - let mk = |url: &str| -> Arc { - Arc::new( - BasicWorkerBuilder::new(url) - .connection_mode(ConnectionMode::Grpc) - .label("rl.control_url", "http://0.0.0.0:40100") - .build(), - ) - }; - let infos = view_over(vec![mk("grpc://a:1"), mk("grpc://b:1")]).list(); + 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" + ); + } + + /// A pool config the gateway cannot build a client for is the worker's + /// problem to report, not a silent `None`: the error rides on the info + /// and ends up in the 422. + #[test] + fn a_failed_client_build_is_reported_on_the_worker() { + let mut config = RouterConfig::default(); + config.ca_certificates = vec![b"not a certificate".to_vec()]; + 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("CA certificate"), "{err}"); + } } From cc8640ee3a9974cea70c5b2c6063e71569b00db8 Mon Sep 17 00:00:00 2001 From: key4ng Date: Tue, 6 Oct 2026 11:52:16 -0700 Subject: [PATCH 20/24] test(rl): register the TokenSpeed mock through POST /workers; a fan-out over endpoints hits each worker once The integration test registered its workers with BasicWorkerBuilder and hand-copied labels, so the mock engine's server_args and the label lifting in discovery fed no assertion. The fleet now registers through POST /workers the way an operator does: the control URL and the rl.* capabilities come out of the mock's own server info, the api_key is the registration's, and the worker whose engine advertises nothing is shown to fall back to the static TokenSpeed row. The mock advertises what a current engine does (update_from distributed,mooncake). Added the spec's mixed-fleet case: a fan-out whose selector covers only workers with endpoints answers 200 and each control app records exactly one request. The mock worker's --server-arg and --weight-version CLI flags are dropped (tests set the fields on Config directly; nothing passed the flags), and the recorder comment no longer claims the authorization and body vectors are index-aligned. Signed-off-by: key4ng --- crates/mock_worker/src/config.rs | 10 - model_gateway/tests/common/mock_worker.rs | 5 +- .../rl_tokenspeed_control_endpoint_test.rs | 323 +++++++++++------- 3 files changed, 199 insertions(+), 139 deletions(-) diff --git a/crates/mock_worker/src/config.rs b/crates/mock_worker/src/config.rs index 4443ea534a..3f6286030c 100644 --- a/crates/mock_worker/src/config.rs +++ b/crates/mock_worker/src/config.rs @@ -134,14 +134,6 @@ impl Config { "--prefix-cache" => { cfg.engine.prefix_cache = parse(value(&mut args, &flag)?, &flag)? } - "--server-arg" => { - let raw = value(&mut args, &flag)?; - let (k, v) = raw - .split_once('=') - .ok_or_else(|| format!("--server-arg expects key=value, got {raw}"))?; - cfg.server_args.insert(k.to_string(), v.to_string()); - } - "--weight-version" => cfg.weight_version = Some(value(&mut args, &flag)?), "-h" | "--help" => return Err(usage()), other => return Err(format!("unknown flag: {other}\n\n{}", usage())), } @@ -194,8 +186,6 @@ fn usage() -> String { --tokenizer tokenizer path for gRPC autoload (default = model)\n\ --gen-ms canned per-request latency (default 0)\n\ --output-tokens output tokens per request when unspecified (default 8)\n\ - --server-arg extra GetServerInfo server_args entry (repeatable)\n\ - --weight-version weight version stamped on generate responses (default unset)\n\ --capture append each gRPC Generate request to as a JSON line\n\ \n\ Realistic engine simulator (continuous batching; opt-in):\n\ diff --git a/model_gateway/tests/common/mock_worker.rs b/model_gateway/tests/common/mock_worker.rs index debb868041..dea62f604d 100755 --- a/model_gateway/tests/common/mock_worker.rs +++ b/model_gateway/tests/common/mock_worker.rs @@ -268,8 +268,9 @@ fn record_request(port: u16, version: Version, body: &serde_json::Value) { } } -/// Record the `authorization` header of one RL control request. Called -/// alongside [`record_request`] so the two vectors stay index-aligned. +/// Record the `authorization` header of one RL control request. Only the +/// RL control routes call it, so this vector is shorter than `bodies` when +/// the recorder's port also served other routes. fn record_authorization(port: u16, value: Option) { let recorder = request_recorders_table() .lock() diff --git a/model_gateway/tests/rl_tokenspeed_control_endpoint_test.rs b/model_gateway/tests/rl_tokenspeed_control_endpoint_test.rs index b817b140a8..70f942ba96 100644 --- a/model_gateway/tests/rl_tokenspeed_control_endpoint_test.rs +++ b/model_gateway/tests/rl_tokenspeed_control_endpoint_test.rs @@ -1,8 +1,11 @@ //! A TokenSpeed engine SMG speaks gRPC to is driven through `/v1/rl` via its -//! HTTP control endpoint: discovery reports it, the proxy reaches it with the -//! worker's bearer, a mixed HTTP+gRPC fleet fans out, a gRPC worker without an -//! endpoint is named in `failed[]`, and the gRPC `/generate` path reports the -//! version the engine stamped on the response. +//! 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 its own pair of mock HTTP ports: the HTTP //! mock binds the port it is configured with, and the tests in one binary run @@ -26,9 +29,7 @@ use common::{ }; use http_body_util::BodyExt; use llm_tokenizer::{traits::Tokenizer, MockTokenizer, TokenizerRegistry}; -use openai_protocol::{ - generate::GenerateRequest, model_card::ModelCard, worker::HealthCheckConfig, -}; +use openai_protocol::generate::GenerateRequest; use serde_json::{json, Value}; use smg::{ app_context::AppContext, @@ -36,7 +37,6 @@ use smg::{ middleware::TenantRequestMeta, routers::{RouterFactory, RouterTrait}, tenant::TenantKey, - worker::{BasicWorkerBuilder, ConnectionMode, RuntimeType, WorkerType}, }; use tokio::net::TcpListener; use tower::ServiceExt; @@ -48,6 +48,26 @@ const ENGINE_VERSION: &str = "v7"; /// 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. @@ -88,8 +108,9 @@ async fn start_mock_grpc_engine(server_args: BTreeMap) -> u16 { } /// 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). +/// 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" @@ -120,74 +141,13 @@ async fn grpc_rl_context() -> Arc { .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 } -fn health_off() -> HealthCheckConfig { - HealthCheckConfig { - disable_health_check: true, - ..Default::default() - } -} - -/// Register a TokenSpeed gRPC worker whose RL capabilities come from labels. -/// `control_url` sets the `rl.control_url` label an engine would advertise; -/// `api_key` is the bearer the gateway presents to that control app. -#[expect( - clippy::expect_used, - reason = "test helper - panicking on failure is intentional" -)] -fn register_tokenspeed( - ctx: &Arc, - grpc_port: u16, - control_url: Option<&str>, - api_key: Option<&str>, -) { - let mut builder = BasicWorkerBuilder::new(format!("grpc://127.0.0.1:{grpc_port}")) - .worker_type(WorkerType::Regular) - .connection_mode(ConnectionMode::Grpc) - .runtime_type(RuntimeType::TokenSpeed) - .model(ModelCard::new(MODEL)) - .health_config(health_off()) - .label("weight_version", REGISTERED_VERSION) - .label("rl.pause_modes", "wait,abort,keep") - .label("rl.update_from", "disk,distributed") - .label("rl.abort", "true") - .label("rl.flush_cache", "true") - .label("rl.sleep_wake", "true") - .label("rl.reports_weight_version", "true"); - if let Some(url) = control_url { - builder = builder.label("rl.control_url", url); - } - if let Some(key) = api_key { - builder = builder.api_key(key); - } - ctx.worker_registry - .register(Arc::new(builder.build())) - .expect("TokenSpeed worker registered"); -} - -/// Register an HTTP SGLang worker, which controls itself over its own URL. -#[expect( - clippy::expect_used, - reason = "test helper - panicking on failure is intentional" -)] -fn register_http_sglang(ctx: &Arc, url: &str) { - let worker = BasicWorkerBuilder::new(url) - .worker_type(WorkerType::Regular) - .connection_mode(ConnectionMode::Http) - .runtime_type(RuntimeType::Sglang) - .model(ModelCard::new(MODEL)) - .health_config(health_off()) - .build(); - ctx.worker_registry - .register(Arc::new(worker)) - .expect("SGLang worker registered"); -} - fn mock_http(port: u16) -> MockWorkerConfig { MockWorkerConfig { port, @@ -207,6 +167,50 @@ async fn json_of(resp: axum::response::Response) -> Value { .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")) } @@ -234,13 +238,17 @@ fn first_generate_result(body: Value) -> Value { } } -/// Two TokenSpeed gRPC workers (one with a control endpoint at an HTTP mock -/// that records what it receives, one without) plus one HTTP SGLang mock. +/// 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, @@ -256,66 +264,87 @@ async fn fleet(control_port: u16, sglang_port: u16) -> Fleet { set_request_recorder(control_port, control_recorder.clone()); let mut control_app = MockWorker::new(mock_http(control_port)); let control_url = control_app.start().await.unwrap(); + let sglang_recorder = RequestRecorder::new(); + set_request_recorder(sglang_port, sglang_recorder.clone()); let mut sglang = MockWorker::new(mock_http(sglang_port)); let sglang_url = sglang.start().await.unwrap(); let ctx = grpc_rl_context().await; - let advertised = BTreeMap::from([("rl.control_url".to_string(), control_url.clone())]); - let with_endpoint = start_mock_grpc_engine(advertised).await; - let without_endpoint = start_mock_grpc_engine(BTreeMap::new()).await; - register_tokenspeed( - &ctx, - with_endpoint, - Some(control_url.as_str()), - Some("ts-secret"), - ); - register_tokenspeed(&ctx, without_endpoint, None, None); - register_http_sglang(&ctx, &sglang_url); - 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, } } -/// 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<'a>(workers: &'a [Value], engine: &str, has_control: bool) -> &'a Value { - workers - .iter() - .find(|w| w["engine"] == engine && w["control_url"].is_string() == has_control) - .expect("worker present") +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_control_endpoint_and_advertised_capabilities() { +async fn discovery_reports_the_advertised_endpoint_and_capabilities() { let f = fleet(18921, 18922).await; - let resp = f - .app - .clone() - .oneshot(Request::get("/v1/rl/workers").body(Body::empty()).unwrap()) - .await - .unwrap(); - assert_eq!(resp.status(), StatusCode::OK); - let body = json_of(resp).await; - assert_eq!(body["total"], 3); - let workers = body["workers"].as_array().unwrap(); - let ts = by_engine(workers, "tokenspeed", true); + 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"); @@ -323,12 +352,27 @@ async fn discovery_reports_control_endpoint_and_advertised_capabilities() { 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 bare = by_engine(workers, "tokenspeed", false); - assert_eq!(bare["control_url"], Value::Null); - assert_eq!(bare["capabilities"]["source"], "label"); + 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 = by_engine(workers, "sglang", true); + let sglang = f.by_engine("sglang", true); assert_eq!( sglang["control_url"], sglang["base_url"], "HTTP workers control themselves" @@ -338,18 +382,7 @@ async fn discovery_reports_control_endpoint_and_advertised_capabilities() { #[tokio::test] async fn proxy_reaches_the_control_endpoint_with_the_worker_bearer() { let f = fleet(18923, 18924).await; - let workers = json_of( - f.app - .clone() - .oneshot(Request::get("/v1/rl/workers").body(Body::empty()).unwrap()) - .await - .unwrap(), - ) - .await; - let id = by_engine(workers["workers"].as_array().unwrap(), "tokenspeed", true)["id"] - .as_str() - .unwrap() - .to_string(); + let id = f.id("tokenspeed", true); let resp = f .app @@ -367,14 +400,49 @@ async fn proxy_reaches_the_control_endpoint_with_the_worker_bearer() { assert_eq!(f.control_recorder.only_body(), json!({"mode": "keep"})); assert_eq!( f.control_recorder.authorizations(), - vec![Some("Bearer ts-secret".to_string())], + vec![Some(format!("Bearer {CONTROL_KEY}"))], "the worker's own key, not the caller's" ); } #[tokio::test] -async fn fanout_spans_transports_and_names_the_worker_without_an_endpoint() { +async fn fanout_over_workers_with_endpoints_is_200_and_hits_each_once() { let f = fleet(18925, 18926).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(18927, 18928).await; let resp = f .app .clone() @@ -396,6 +464,7 @@ async fn fanout_spans_transports_and_names_the_worker_without_an_endpoint() { 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() { @@ -406,7 +475,7 @@ async fn fanout_spans_transports_and_names_the_worker_without_an_endpoint() { #[tokio::test] async fn grpc_generate_reports_the_engine_stamped_version() { - let f = fleet(18927, 18928).await; + let f = fleet(18929, 18930).await; let resp = f .router From ac568ed797e951b98b472b72a164dc7601b7ed2a Mon Sep 17 00:00:00 2001 From: key4ng Date: Tue, 6 Oct 2026 11:52:16 -0700 Subject: [PATCH 21/24] fix(serve): wire no RL control endpoint for ZMQ TokenSpeed workers smg serve passed --rl-control-host/--rl-control-port to the headless TokenSpeed engine and stamped rl.control_url on the worker. A headless engine (launch_scheduler_headless) builds no AsyncLLM, and only AsyncLLM starts the control app, so the label pointed at a port nothing served: discovery showed a control URL and every control call failed 502 upstream_unreachable, where the worker without the label gets the honest 422. The launcher no longer passes the flags, the stamping thread and its tests are gone, and NOTES.md records the drift. Signed-off-by: key4ng --- bindings/python/src/smg/serve.py | 101 ------------------------- bindings/python/tests/test_serve.py | 110 ---------------------------- crates/rl/NOTES.md | 4 +- 3 files changed, 3 insertions(+), 212 deletions(-) diff --git a/bindings/python/src/smg/serve.py b/bindings/python/src/smg/serve.py index f279ac6c2b..a012bf3db0 100644 --- a/bindings/python/src/smg/serve.py +++ b/bindings/python/src/smg/serve.py @@ -8,7 +8,6 @@ import argparse import atexit -import json import logging import os import random @@ -16,9 +15,7 @@ import socket import subprocess import sys -import threading import time -import urllib.error import urllib.request from abc import ABC, abstractmethod @@ -87,19 +84,6 @@ def _zmq_handshake_port(ipc_url: str) -> int: return _ZMQ_HANDSHAKE_PORT_BASE + (h % _ZMQ_HANDSHAKE_PORT_SPAN) -def _rl_control_port(port: int) -> int: - """Port for a TokenSpeed worker's in-engine RL control app. - - Offset from the worker port so co-located workers (consecutive ports) and - the engine's own +233 distributed store never collide, reflected below the - u16 ceiling and hopped past SMG's ZMQ handshake band like ``dist_port``. - """ - p = port + 400 if port + 400 <= 65535 else port - 400 - if _ZMQ_HANDSHAKE_PORT_BASE <= p < _ZMQ_HANDSHAKE_PORT_BASE + _ZMQ_HANDSHAKE_PORT_SPAN: - p += _ZMQ_HANDSHAKE_PORT_SPAN - return p - - def _reject_handshake_port_collisions(ports: list[int]) -> None: """Fail before launch if two workers derive the same ZMQ handshake port. @@ -439,10 +423,6 @@ class TokenspeedWorkerLauncher(WorkerLauncher): def _get_tp_size(self, args: argparse.Namespace) -> int: return getattr(args, "tensor_parallel_size", 1) or 1 - def control_url(self, port: int) -> str: - """URL of the in-engine RL control app this launcher started for ``port``.""" - return f"http://127.0.0.1:{_rl_control_port(port)}" - def build_command( self, args: argparse.Namespace, backend_args: list[str], host: str, port: int ) -> list[str]: @@ -509,10 +489,6 @@ def _build_zmq_command( str(rpc_port), "--zmq-engine-index", "0", - "--rl-control-host", - "127.0.0.1", - "--rl-control-port", - str(_rl_control_port(port)), ] cmd.extend( self._backend_arg_defaults( @@ -542,8 +518,6 @@ def _build_zmq_command( "--data-parallel-address", "--data-parallel-rpc-port", "--zmq-engine-index", - "--rl-control-host", - "--rl-control-port", ], ) ) @@ -1001,60 +975,6 @@ def parse_serve_args( _WORKER_SHUTDOWN_TIMEOUT = 30 -def _stamp_rl_control_labels( - gateway_url: str, - api_key: str | None, - targets: list[tuple[str, str]], - deadline_s: float, -) -> None: - """Label each ZMQ worker with its control endpoint once the gateway lists it. - - ZMQ discovery yields no labels, so the launcher, which owns both ends, - stamps ``rl.control_url`` through the worker update route (labels merge). - """ - headers = {"Content-Type": "application/json"} - if api_key: - headers["Authorization"] = f"Bearer {api_key}" - pending = dict(targets) - stop_at = time.monotonic() + deadline_s - while pending and time.monotonic() < stop_at: - try: - req = urllib.request.Request(f"{gateway_url}/workers", headers=headers, method="GET") - with urllib.request.urlopen(req, timeout=5) as resp: - workers = json.loads(resp.read()).get("workers", []) - except Exception as e: # noqa: BLE001 — the gateway may not be up yet - logger.debug("rl.control_url stamping: gateway not ready: %s", e) - time.sleep(1) - continue - for w in workers: - url = str(w.get("url", "")) - if url not in pending: - continue - body = json.dumps({"labels": {"rl.control_url": pending[url]}}).encode() - req = urllib.request.Request( - f"{gateway_url}/workers/{w['id']}", data=body, headers=headers, method="PATCH" - ) - try: - with urllib.request.urlopen(req, timeout=5): - pass - logger.info("stamped rl.control_url=%s on %s", pending[url], url) - del pending[url] - except urllib.error.HTTPError as e: - if 400 <= e.code < 500: - logger.warning( - "rl.control_url stamping got HTTP %s for %s; not retrying", e.code, url - ) - del pending[url] - else: - logger.warning("rl.control_url stamping failed for %s: %s", url, e) - except Exception as e: # noqa: BLE001 - logger.warning("rl.control_url stamping failed for %s: %s", url, e) - if pending: - time.sleep(1) - for url in pending: - logger.warning("rl.control_url never stamped on %s (gateway did not list it)", url) - - class ServeOrchestrator: """Coordinate worker launch, health checking, router startup, and shutdown.""" @@ -1077,27 +997,6 @@ def run(self) -> None: self._launch_workers() self._wait_healthy() router_args = self._build_router_args() - if ( - getattr(router_args, "enable_rl", False) - and self.backend == "tokenspeed" - and getattr(self.args, "connection_mode", "grpc") == "zmq" - ): - control = getattr(self.launcher, "control_url", None) - if callable(control): - targets = [ - ( - self.launcher.worker_url(self.args, self.args.worker_host, port), - control(port), - ) - for _, port in self.workers - ] - gateway_url = f"http://127.0.0.1:{router_args.port}" - threading.Thread( - target=_stamp_rl_control_labels, - args=(gateway_url, getattr(router_args, "api_key", None), targets, 300.0), - name="smg-rl-control-labels", - daemon=True, - ).start() launch_router(router_args) finally: self._cleanup_workers() diff --git a/bindings/python/tests/test_serve.py b/bindings/python/tests/test_serve.py index 3db7e3f1c4..0aa5a6f939 100644 --- a/bindings/python/tests/test_serve.py +++ b/bindings/python/tests/test_serve.py @@ -1675,113 +1675,3 @@ def capture_launch(a, b, host, port, env): assert launched_envs[1]["CUDA_VISIBLE_DEVICES"] == "2,3" assert launched_envs[0]["PYTHONUNBUFFERED"] == "1" assert launched_envs[1]["PYTHONUNBUFFERED"] == "1" - - -def test_rl_control_port_avoids_neighbors_and_the_handshake_band(): - from smg.serve import _ZMQ_HANDSHAKE_PORT_BASE, _ZMQ_HANDSHAKE_PORT_SPAN, _rl_control_port - - assert _rl_control_port(30000) == 30400 - assert _rl_control_port(65500) == 65100 - inside = _ZMQ_HANDSHAKE_PORT_BASE - 400 + 10 - assert not ( - _ZMQ_HANDSHAKE_PORT_BASE - <= _rl_control_port(inside) - < _ZMQ_HANDSHAKE_PORT_BASE + _ZMQ_HANDSHAKE_PORT_SPAN - ) - - -def test_tokenspeed_zmq_command_wires_the_control_app(monkeypatch): - from types import SimpleNamespace - - from smg.serve import TokenspeedWorkerLauncher - - args = SimpleNamespace(model="/models/q", connection_mode="zmq", tensor_parallel_size=1) - cmd = TokenspeedWorkerLauncher().build_command(args, [], "127.0.0.1", 30000) - assert cmd[cmd.index("--rl-control-host") + 1] == "127.0.0.1" - assert cmd[cmd.index("--rl-control-port") + 1] == "30400" - assert TokenspeedWorkerLauncher().control_url(30000) == "http://127.0.0.1:30400" - - -def test_stamp_rl_control_labels_patches_each_worker(monkeypatch): - import json - - from smg import serve - - calls = [] - - class _Resp: - def __init__(self, payload): - self._payload = json.dumps(payload).encode() - self.status = 200 - - def read(self): - return self._payload - - def __enter__(self): - return self - - def __exit__(self, *a): - return False - - def fake_urlopen(req, timeout=0): - calls.append((req.get_method(), req.full_url, req.data, req.get_header("Authorization"))) - if req.get_method() == "GET": - return _Resp({"workers": [{"id": "w1", "url": "ipc:///tmp/engine-30000"}]}) - return _Resp({}) - - monkeypatch.setattr(serve.urllib.request, "urlopen", fake_urlopen) - serve._stamp_rl_control_labels( - "http://127.0.0.1:8000", - "adm", - [("ipc:///tmp/engine-30000", "http://127.0.0.1:30400")], - deadline_s=1.0, - ) - patch = [c for c in calls if c[0] == "PATCH"] - assert len(patch) == 1 - assert patch[0][1] == "http://127.0.0.1:8000/workers/w1" - assert json.loads(patch[0][2]) == {"labels": {"rl.control_url": "http://127.0.0.1:30400"}} - assert patch[0][3] == "Bearer adm" - - -def test_stamp_rl_control_labels_stops_retrying_after_a_4xx(monkeypatch): - import json - import time - import urllib.error - - from smg import serve - - calls = [] - - class _Resp: - def __init__(self, payload): - self._payload = json.dumps(payload).encode() - self.status = 200 - - def read(self): - return self._payload - - def __enter__(self): - return self - - def __exit__(self, *a): - return False - - def fake_urlopen(req, timeout=0): - calls.append((req.get_method(), req.full_url, req.data, req.get_header("Authorization"))) - if req.get_method() == "GET": - return _Resp({"workers": [{"id": "w1", "url": "ipc:///tmp/engine-30000"}]}) - raise urllib.error.HTTPError(req.full_url, 401, "unauthorized", {}, None) - - monkeypatch.setattr(serve.urllib.request, "urlopen", fake_urlopen) - start = time.monotonic() - serve._stamp_rl_control_labels( - "http://127.0.0.1:8000", - "adm", - [("ipc:///tmp/engine-30000", "http://127.0.0.1:30400")], - deadline_s=5.0, - ) - elapsed = time.monotonic() - start - - patch = [c for c in calls if c[0] == "PATCH"] - assert len(patch) == 1, "a 4xx must not be retried" - assert elapsed < 3.0, "giving up on the 4xx must not wait out the deadline" diff --git a/crates/rl/NOTES.md b/crates/rl/NOTES.md index deae294812..74804e719f 100644 --- a/crates/rl/NOTES.md +++ b/crates/rl/NOTES.md @@ -10,5 +10,7 @@ planning docs, with date, engine version, and what was done. | 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=disk,distributed` | +| 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 | From 26d072f3a3bee3b056fbd3afd9c9eeba0a333585 Mon Sep 17 00:00:00 2001 From: key4ng Date: Tue, 6 Oct 2026 11:52:49 -0700 Subject: [PATCH 22/24] docs(rl): launch recipe that works against the merged engine The guide started each engine with --rl-control-host 0.0.0.0 and promised a control URL; the merged engine advertises no rl.control_url for a wildcard bind, so discovery showed null and every call was 422. It also re-registered a --worker-urls worker through POST /workers, a 409. The recipe now uses a routable --rl-control-host, says that wildcard hosts are not advertised and that loopback is right only on the gateway's own machine, and registers the engines with their key through POST /workers only. The README says the same about advertisement. The tp_size sentence describes what discovery does today (attn_tp_size folded into tp_size; the newest engines nest parallelism under mapping.*, which it does not read yet). COUPLING.md describes this PR's files only, and NOTES.md no longer claims the engine advertises a disk refit source. Signed-off-by: key4ng --- crates/rl/COUPLING.md | 12 +++++------ crates/rl/README.md | 7 ++++-- docs/guides/rl-tokenspeed.md | 42 ++++++++++++++++++++++-------------- 3 files changed, 36 insertions(+), 25 deletions(-) diff --git a/crates/rl/COUPLING.md b/crates/rl/COUPLING.md index 597f63571b..fc6d24b050 100644 --- a/crates/rl/COUPLING.md +++ b/crates/rl/COUPLING.md @@ -30,13 +30,11 @@ 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`; `src/routers/grpc/router.rs` and -`src/routers/grpc/pipeline.rs` give slime's model-less single-prompt -`/generate` SGLang's shape; `src/routers/http/router.rs` (with -`crates/protocols/src/generate.rs`) stops forwarding the wildcard `model` -placeholder. `model_gateway/tests/rl_tokenspeed_control_endpoint_test.rs` -drives a TokenSpeed gRPC worker through `/v1/rl` via its control endpoint and -checks the generate shape. +`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 diff --git a/crates/rl/README.md b/crates/rl/README.md index a8330b7919..f15942e5e5 100644 --- a/crates/rl/README.md +++ b/crates/rl/README.md @@ -21,8 +21,11 @@ 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. A wildcard bind host in -the advertised URL (`0.0.0.0`, `::`) is replaced by the worker's own host. The +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 diff --git a/docs/guides/rl-tokenspeed.md b/docs/guides/rl-tokenspeed.md index 1c9300d293..9de6ab75d5 100644 --- a/docs/guides/rl-tokenspeed.md +++ b/docs/guides/rl-tokenspeed.md @@ -1,30 +1,40 @@ # RL rollouts on TokenSpeed behind SMG SMG's RL control plane (`--enable-rl`, `/v1/rl/*`) drives TokenSpeed engines -through their in-engine SGLang-compatible control app while the data plane -stays on gRPC. +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 0.0.0.0 --rl-control-port 30400 --rl-control-api-key "$RL_KEY" \ + --rl-control-host 10.0.0.11 --rl-control-port 30400 --rl-control-api-key "$RL_KEY" \ --enable-output-logprobs -The gateway: +`--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. - smg launch --worker-urls grpc://rollout-1:30000 grpc://rollout-2:30000 \ - --policy cache_aware --enable-rl --disable-health-check --disable-circuit-breaker \ +The gateway, with no startup workers: + + smg launch --policy cache_aware --enable-rl --disable-health-check --disable-circuit-breaker \ --request-timeout-secs 14400 -Register the engines with their control key so the proxy authenticates: +Register each engine with its control key, which the proxy sends as the +bearer to the control app: curl -X POST http://smg:30000/workers -H 'content-type: application/json' \ -d '{"url":"grpc://rollout-1:30000","api_key":"'"$RL_KEY"'"}' +A worker named in `--worker-urls` registers without a key, and registering +it again this way is a 409, so list no startup workers and register every +engine as above. + `GET /v1/rl/workers` then shows `engine: tokenspeed`, `connection_mode: grpc`, -`control_url: http://rollout-1:30400` and the engine's advertised capabilities. +`control_url: http://10.0.0.11:30400` and the engine's advertised capabilities. ## Refit @@ -48,17 +58,17 @@ with `pause_generation` / `continue_generation` fanned out around the 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`: discovery reads it from the engine's -server args and falls back to TokenSpeed's own spelling, `attn_tp_size`, so an -engine launched with either reports a width. An engine launched with neither -(TokenSpeed leaves `attn_tp_size` unset unless asked) reports `tp_size: null`, -and a trainer must then assume 1 or be told. +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. Bind it -on a routable host only with `--rl-control-api-key`, and give SMG the same key -as the worker's `api_key`. +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 From c8ebeb370c9acc919e4c3da4dc83907302e8db32 Mon Sep 17 00:00:00 2001 From: key4ng Date: Tue, 6 Oct 2026 12:05:57 -0700 Subject: [PATCH 23/24] test(rl): make the failed-client-build case fail the build for real reqwest keeps a CA bundle as bytes and only parses it when a connection is made, so a bogus bundle still produced a client and the test's expect_err panicked. An unparsable client identity is rejected when the client is built, which is the failure the test means to observe. Signed-off-by: key4ng --- model_gateway/src/rl_adapter.rs | 9 +++++---- 1 file changed, 5 insertions(+), 4 deletions(-) diff --git a/model_gateway/src/rl_adapter.rs b/model_gateway/src/rl_adapter.rs index 8b8ab0f7bd..3505480a21 100644 --- a/model_gateway/src/rl_adapter.rs +++ b/model_gateway/src/rl_adapter.rs @@ -273,19 +273,20 @@ mod tests { ); } - /// A pool config the gateway cannot build a client for is the worker's + /// 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. + /// 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 mut config = RouterConfig::default(); - config.ca_certificates = vec![b"not a certificate".to_vec()]; + config.client_identity = Some(b"not a pem".to_vec()); 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("CA certificate"), "{err}"); + assert!(err.contains("client identity"), "{err}"); } } From 4d1d7b339ec7eb9cf9bf52730a20449ee73bbac1 Mon Sep 17 00:00:00 2001 From: key4ng Date: Tue, 6 Oct 2026 12:32:30 -0700 Subject: [PATCH 24/24] style(rl): build the test's RouterConfig with struct-update syntax clippy (field_reassign_with_default) rejects assigning a field on a value just created with Default::default(); CI runs with -D warnings. Signed-off-by: key4ng --- model_gateway/src/rl_adapter.rs | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/model_gateway/src/rl_adapter.rs b/model_gateway/src/rl_adapter.rs index 3505480a21..32886585f5 100644 --- a/model_gateway/src/rl_adapter.rs +++ b/model_gateway/src/rl_adapter.rs @@ -279,8 +279,10 @@ mod tests { /// build eagerly; CA bundles are only read when a connection is made.) #[test] fn a_failed_client_build_is_reported_on_the_worker() { - let mut config = RouterConfig::default(); - config.client_identity = Some(b"not a pem".to_vec()); + 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()