Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 5 additions & 0 deletions bindings/python/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -528,6 +528,7 @@ struct Router {
kv_connector_annotation: String,
kv_engine_id_annotation: String,
mm_per_request_image_limit: Option<usize>,
pd_admission_wait_secs: u64,
}

impl Router {
Expand Down Expand Up @@ -848,6 +849,7 @@ impl Router {
.worker_overload_protection(self.worker_overload_protection)
.disable_load_monitoring(self.disable_load_monitoring)
.load_monitor_interval_secs(self.load_monitor_interval)
.pd_admission_wait_secs(self.pd_admission_wait_secs)
.max_concurrent_requests(self.max_concurrent_requests)
.queue_size(self.queue_size)
.queue_timeout_secs(self.queue_timeout_secs)
Expand Down Expand Up @@ -1098,6 +1100,7 @@ impl Router {
kv_connector_annotation = String::from("smg.ai/kv-connector"),
kv_engine_id_annotation = String::from("smg.ai/kv-engine-id"),
mm_per_request_image_limit = None,
pd_admission_wait_secs = 30,
))]
#[expect(clippy::too_many_arguments)]
#[expect(
Expand Down Expand Up @@ -1251,6 +1254,7 @@ impl Router {
kv_connector_annotation: String,
kv_engine_id_annotation: String,
mm_per_request_image_limit: Option<usize>,
pd_admission_wait_secs: u64,
) -> PyResult<Self> {
let mut all_urls = worker_urls.clone();

Expand Down Expand Up @@ -1418,6 +1422,7 @@ impl Router {
kv_connector_annotation,
kv_engine_id_annotation,
mm_per_request_image_limit,
pd_admission_wait_secs,
})
}

Expand Down
15 changes: 15 additions & 0 deletions bindings/python/src/smg/router_args.py
Original file line number Diff line number Diff line change
Expand Up @@ -251,6 +251,9 @@ class RouterArgs:
kv_engine_id_annotation: str = "smg.ai/kv-engine-id"
# Per-request image-count limit replacing model spec limits; None keeps spec limits
mm_per_request_image_limit: int | None = None
# Seconds a PD dispatch waits for a slot in the decode engine's running
# window before shedding; 0 sheds immediately
pd_admission_wait_secs: int = 30

@staticmethod
def add_cli_args(
Expand Down Expand Up @@ -728,6 +731,18 @@ def add_cli_args(
" retries"
),
)
routing_group.add_argument(
f"--{prefix}pd-admission-wait-secs",
type=int,
default=RouterArgs.pd_admission_wait_secs,
help=(
"Seconds a prefill/decode dispatch waits for a free slot in"
" the decode engine's running window before shedding with 503"
" worker_overload_protection_shed. Keep it well under the"
" engine's bootstrap deadline. 0 sheds immediately; engines"
" that report no running window are never gated"
),
)
routing_group.add_argument(
f"--{prefix}stream-body-stall-timeout-secs",
type=int,
Expand Down
1 change: 1 addition & 0 deletions bindings/python/tests/test_arg_parser.py
Original file line number Diff line number Diff line change
Expand Up @@ -1440,6 +1440,7 @@ class TestRouterArgsFieldOrder:
"kv_connector_annotation",
"kv_engine_id_annotation",
"mm_per_request_image_limit",
"pd_admission_wait_secs",
]

def test_complete_field_sequence_is_frozen(self):
Expand Down
57 changes: 57 additions & 0 deletions grpc_servicer/smg_grpc_servicer/tokenspeed/loads.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,57 @@
"""Conversion of TokenSpeed scheduler load replies into the gateway's protobuf.

Kept free of engine imports so the field mapping can be unit-tested without
TokenSpeed installed (see grpc_servicer/tests/test_tokenspeed_loads.py),
the same split the SGLang side uses.
"""

from __future__ import annotations

from typing import Any

from smg_grpc_proto import tokenspeed_scheduler_pb2


def running_window(server_args: Any) -> int:
"""The scheduler's configured admission window, or 0 when unknown.

``GetLoadReqOutput`` carries no admission bound, so the window has to come
from the server args. TokenSpeed spells it ``max_num_seqs``; the protobuf
field is the one SGLang fills with the same number, so a frontend reads a
single name across both engines. It matters most on a disaggregated decode
worker: that window is exactly the bound its prefill peer's bootstrap
deadline is racing, and a frontend that cannot see it cannot pace its
dispatch.
"""
return int(
getattr(server_args, "max_num_seqs", 0)
or getattr(server_args, "max_running_requests", 0)
or 0
)


def convert_load_to_protobuf(
load_output: Any,
*,
page_size: int,
max_total_num_tokens: int,
max_running_requests: int,
) -> tokenspeed_scheduler_pb2.SchedulerLoad:
"""Convert one rank's ``GetLoadReqOutput`` to a protobuf ``SchedulerLoad``.

The reply counts ``num_reqs`` as running + waiting and reports KV usage in
pages, so running requests and used tokens are derived here.
"""
num_waiting_reqs = int(load_output.num_waiting_reqs)
num_total_reqs = int(load_output.num_reqs)
num_used_tokens = int(load_output.num_pages) * page_size
return tokenspeed_scheduler_pb2.SchedulerLoad(
dp_rank=int(load_output.dp_rank),
num_running_reqs=max(0, num_total_reqs - num_waiting_reqs),
num_waiting_reqs=num_waiting_reqs,
num_total_reqs=num_total_reqs,
num_used_tokens=num_used_tokens,
max_total_num_tokens=max_total_num_tokens,
token_usage=(num_used_tokens / max_total_num_tokens if max_total_num_tokens > 0 else 0.0),
max_running_requests=max_running_requests,
)
38 changes: 14 additions & 24 deletions grpc_servicer/smg_grpc_servicer/tokenspeed/servicer.py
Original file line number Diff line number Diff line change
Expand Up @@ -47,6 +47,7 @@
from smg_grpc_servicer.tokenizer_bundle import CHUNK_SIZE, build_tokenizer_zip
from smg_grpc_servicer.tokenspeed.health_servicer import TokenSpeedHealthServicer
from smg_grpc_servicer.tokenspeed.kv_events import resolve_kv_events_config
from smg_grpc_servicer.tokenspeed.loads import convert_load_to_protobuf, running_window

if TYPE_CHECKING:
# Type-only — keeps these out of the cold-path graph when the servicer is
Expand Down Expand Up @@ -674,31 +675,20 @@ async def GetLoads(
or getattr(self.async_llm.server_args, "max_total_num_tokens", 0)
or 0
)

scheduler_loads: list[tokenspeed_scheduler_pb2.SchedulerLoad] = []
total_running = 0
total_waiting = 0
token_usages: list[float] = []
for lo in load_outputs:
num_running = max(0, int(lo.num_reqs) - int(lo.num_waiting_reqs))
num_used_tokens = int(lo.num_pages) * page_size
token_usage = (
num_used_tokens / max_total_num_tokens if max_total_num_tokens > 0 else 0.0
)
scheduler_loads.append(
tokenspeed_scheduler_pb2.SchedulerLoad(
dp_rank=int(lo.dp_rank),
num_running_reqs=num_running,
num_waiting_reqs=int(lo.num_waiting_reqs),
num_total_reqs=int(lo.num_reqs),
num_used_tokens=num_used_tokens,
max_total_num_tokens=max_total_num_tokens,
token_usage=token_usage,
)
max_running_requests = running_window(self.async_llm.server_args)

scheduler_loads = [
convert_load_to_protobuf(
lo,
page_size=page_size,
max_total_num_tokens=max_total_num_tokens,
max_running_requests=max_running_requests,
)
total_running += num_running
total_waiting += int(lo.num_waiting_reqs)
token_usages.append(token_usage)
for lo in load_outputs
]
token_usages = [load.token_usage for load in scheduler_loads]
total_running = sum(load.num_running_reqs for load in scheduler_loads)
total_waiting = sum(load.num_waiting_reqs for load in scheduler_loads)

aggregate = tokenspeed_scheduler_pb2.AggregateMetrics(
total_running_reqs=total_running,
Expand Down
75 changes: 75 additions & 0 deletions grpc_servicer/tests/test_tokenspeed_loads.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,75 @@
"""GetLoads must report the scheduler's configured running window.

``GetLoadReqOutput`` carries no admission bound, so the servicer left
``max_running_requests`` at zero on every TokenSpeed load report. A frontend
could not tell how many requests the worker will run at once — and on a
disaggregated decode worker that window is exactly the bound its prefill
peer's bootstrap deadline is racing.

Run with: pytest grpc_servicer/tests/test_tokenspeed_loads.py
"""

import importlib.util
from pathlib import Path
from types import SimpleNamespace

import pytest

pytest.importorskip("smg_grpc_proto")


@pytest.fixture(scope="module")
def loads_mod():
"""Load loads.py by path: the tokenspeed package __init__ imports the engine."""
path = Path(__file__).resolve().parent.parent / "smg_grpc_servicer" / "tokenspeed" / "loads.py"
spec = importlib.util.spec_from_file_location("tokenspeed_loads_under_test", path)
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
return module


class TestRunningWindow:
def test_max_num_seqs_is_the_window(self, loads_mod):
assert loads_mod.running_window(SimpleNamespace(max_num_seqs=16)) == 16

def test_max_running_requests_spelling_is_accepted(self, loads_mod):
assert loads_mod.running_window(SimpleNamespace(max_running_requests=64)) == 64

def test_absent_window_reports_zero_rather_than_guessing(self, loads_mod):
assert loads_mod.running_window(SimpleNamespace()) == 0
assert loads_mod.running_window(SimpleNamespace(max_num_seqs=None)) == 0


class TestLoadConversion:
def test_window_and_derived_counters_reach_the_protobuf(self, loads_mod):
load_output = SimpleNamespace(dp_rank=1, num_reqs=5, num_waiting_reqs=2, num_pages=8)

load = loads_mod.convert_load_to_protobuf(
load_output,
page_size=16,
max_total_num_tokens=1024,
max_running_requests=16,
)

assert load.max_running_requests == 16
assert load.dp_rank == 1
# num_reqs counts running + waiting; used tokens come from pages.
assert load.num_running_reqs == 3
assert load.num_waiting_reqs == 2
assert load.num_total_reqs == 5
assert load.num_used_tokens == 128
assert load.max_total_num_tokens == 1024
assert load.token_usage == pytest.approx(0.125)

def test_unknown_capacity_reports_zero_usage_instead_of_dividing(self, loads_mod):
load_output = SimpleNamespace(dp_rank=0, num_reqs=1, num_waiting_reqs=0, num_pages=4)

load = loads_mod.convert_load_to_protobuf(
load_output,
page_size=1,
max_total_num_tokens=0,
max_running_requests=0,
)

assert load.token_usage == 0.0
assert load.max_running_requests == 0
7 changes: 6 additions & 1 deletion model_gateway/src/app_context.rs
Original file line number Diff line number Diff line change
Expand Up @@ -21,7 +21,9 @@ use crate::{
policies::PolicyRegistry,
rate_limit::RateLimitManager,
routers::{
common::{openai_bridge::FormatRegistry, overload, realtime::RealtimeRegistry},
common::{
openai_bridge::FormatRegistry, overload, pd_admission, realtime::RealtimeRegistry,
},
gateway::Gateway,
grpc::multimodal::MultimodalConfigRegistry,
},
Expand Down Expand Up @@ -631,6 +633,9 @@ impl AppContextBuilder {
// The overload shed advertises the poll interval as Retry-After — the
// veto cannot clear between polls.
overload::set_shed_retry_after_secs(config.load_monitor_interval_secs);
// PD dispatch waits here, not in the decode engine's queue, when the
// pair's running window is full.
pd_admission::set_pd_admission_wait_secs(config.pd_admission_wait_secs);
// Wire the backend load-snapshot feed into every policy that consumes
// it; the monitor polls every group by default, conditionally under
// `--disable-load-monitoring`.
Expand Down
5 changes: 5 additions & 0 deletions model_gateway/src/config/builder.rs
Original file line number Diff line number Diff line change
Expand Up @@ -263,6 +263,11 @@ impl RouterConfigBuilder {
self
}

pub fn pd_admission_wait_secs(mut self, secs: u64) -> Self {
self.config.pd_admission_wait_secs = secs;
self
}

pub fn disable_load_monitoring(mut self, disabled: bool) -> Self {
self.config.disable_load_monitoring = disabled;
self
Expand Down
35 changes: 35 additions & 0 deletions model_gateway/src/config/types.rs
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@ pub use smg_data_connector::{

use super::{validation::ConfigValidator, ConfigResult};
use crate::{
routers::common::pd_admission::DEFAULT_PD_ADMISSION_WAIT_SECS,
tenant::DEFAULT_TENANT_HEADER_NAME,
worker::{ConnectionMode, RuntimeType},
};
Expand Down Expand Up @@ -91,6 +92,14 @@ pub struct RouterConfig {
pub job_queue_concurrency: usize,
#[serde(default = "default_load_monitor_interval_secs")]
pub load_monitor_interval_secs: u64,
/// How long a disaggregated (PD) dispatch waits for a slot in the decode
/// engine's running window before shedding. Must stay well under the
/// engine's bootstrap deadline (120s on TokenSpeed): a request that waits
/// out this budget and then dispatches still has the whole deadline ahead
/// of it. `0` sheds immediately instead of waiting. Ignored for engines
/// that report no running window.
#[serde(default = "default_pd_admission_wait_secs")]
pub pd_admission_wait_secs: u64,
/// Restore the conditional load-monitor poll gate: only poll worker groups
/// when a load-aware routing policy, `engine_metrics`, or overload
/// protection needs the data. Default `false` — the monitor polls every
Expand Down Expand Up @@ -333,6 +342,10 @@ fn default_load_monitor_interval_secs() -> u64 {
10
}

fn default_pd_admission_wait_secs() -> u64 {
DEFAULT_PD_ADMISSION_WAIT_SECS
}

fn default_job_queue_capacity() -> usize {
1000
}
Expand Down Expand Up @@ -1064,6 +1077,7 @@ impl Default for RouterConfig {
job_queue_capacity: default_job_queue_capacity(),
job_queue_concurrency: default_job_queue_concurrency(),
load_monitor_interval_secs: 10,
pd_admission_wait_secs: default_pd_admission_wait_secs(),
disable_load_monitoring: false,
worker_overload_protection: false,
worker_overload_waiting_requests: None,
Expand Down Expand Up @@ -1387,6 +1401,27 @@ mod tests {
assert_eq!(with.stream_body_stall_timeout_secs, 0);
}

#[test]
fn test_pd_admission_wait_serde_default_and_roundtrip() {
// Config files predating the field deserialize to the 30s default.
let mut json: serde_json::Value = serde_json::to_value(RouterConfig::default()).unwrap();
json.as_object_mut()
.unwrap()
.remove("pd_admission_wait_secs")
.unwrap();
let without: RouterConfig = serde_json::from_value(json).unwrap();
assert_eq!(without.pd_admission_wait_secs, 30);

// The shed-immediately zero round-trips instead of reverting.
let config = RouterConfig::builder()
.regular_mode(vec![])
.pd_admission_wait_secs(0)
.build_unchecked();
let json = serde_json::to_string(&config).unwrap();
let with: RouterConfig = serde_json::from_str(&json).unwrap();
assert_eq!(with.pd_admission_wait_secs, 0);
}

#[test]
fn alias_field_spellings_deserialize_and_serialize_canonically() {
// Config files may use the intent-revealing spellings; alias in,
Expand Down
Loading
Loading