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
41 changes: 41 additions & 0 deletions bindings/python/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -415,6 +415,7 @@ struct Router {
shutdown_grace_period_secs: u64,
request_id_headers: Option<Vec<String>>,
trust_tenant_header: bool,
prefer_trusted_tenant_header: bool,
tenant_header_name: String,
storage_context_headers: HashMap<String, String>,
pd_disaggregation: bool,
Expand Down Expand Up @@ -505,6 +506,11 @@ struct Router {
adaptive_admission_max_segments: usize,
adaptive_admission_min_load_coverage: f64,
adaptive_admission_cold_start_output_tokens: u32,
adaptive_admission_strategy: String,
adaptive_admission_feedback_probe_requests_per_healthy_replica: u32,
adaptive_admission_feedback_max_waiting_requests_per_healthy_replica: u32,
adaptive_admission_feedback_max_token_usage: f64,
adaptive_admission_feedback_throughput_improvement_ratio: f64,
}

impl Router {
Expand Down Expand Up @@ -714,6 +720,14 @@ impl Router {
reason,
}
})?;
let adaptive_admission_strategy =
self.adaptive_admission_strategy.parse().map_err(|reason| {
config::ConfigError::InvalidValue {
field: "adaptive_admission_strategy".to_string(),
value: self.adaptive_admission_strategy.clone(),
reason,
}
})?;

let history_backend = match self.history_backend {
HistoryBackendType::Memory => config::HistoryBackend::Memory,
Expand Down Expand Up @@ -793,12 +807,20 @@ impl Router {
.priority_scheduler_tenant_metric_top_n(self.priority_scheduler_tenant_metric_top_n)
.adaptive_admission(config::AdaptiveAdmissionConfig {
mode: adaptive_admission_mode,
strategy: adaptive_admission_strategy,
work_horizon_secs: self.adaptive_admission_work_horizon_secs,
estimator_half_life_secs: self.adaptive_admission_estimator_half_life_secs,
prior_observations: self.adaptive_admission_prior_observations,
max_segments: self.adaptive_admission_max_segments,
min_load_coverage: self.adaptive_admission_min_load_coverage,
cold_start_output_tokens: self.adaptive_admission_cold_start_output_tokens,
feedback_probe_requests_per_healthy_replica: self
.adaptive_admission_feedback_probe_requests_per_healthy_replica,
feedback_max_waiting_requests_per_healthy_replica: self
.adaptive_admission_feedback_max_waiting_requests_per_healthy_replica,
feedback_max_token_usage: self.adaptive_admission_feedback_max_token_usage,
feedback_throughput_improvement_ratio: self
.adaptive_admission_feedback_throughput_improvement_ratio,
})
.cors_allowed_origins(self.cors_allowed_origins.clone())
.retry_config(config::RetryConfig {
Expand Down Expand Up @@ -840,6 +862,7 @@ impl Router {
.maybe_log_level(self.log_level.as_ref())
.maybe_request_id_headers(self.request_id_headers.clone())
.trust_tenant_header(self.trust_tenant_header)
.prefer_trusted_tenant_header(self.prefer_trusted_tenant_header)
.tenant_header_name(&self.tenant_header_name)
.maybe_storage_context_headers(
(!self.storage_context_headers.is_empty())
Expand Down Expand Up @@ -935,6 +958,7 @@ impl Router {
shutdown_grace_period_secs = 180,
request_id_headers = None,
trust_tenant_header = false,
prefer_trusted_tenant_header = false,
tenant_header_name = String::from("x-smg-tenant-id"),
storage_context_headers = HashMap::new(),
pd_disaggregation = false,
Expand Down Expand Up @@ -1026,6 +1050,11 @@ impl Router {
adaptive_admission_max_segments = 50000,
adaptive_admission_min_load_coverage = 0.8,
adaptive_admission_cold_start_output_tokens = 4096,
adaptive_admission_strategy = String::from("predicted_work"),
adaptive_admission_feedback_probe_requests_per_healthy_replica = 2,
adaptive_admission_feedback_max_waiting_requests_per_healthy_replica = 2,
adaptive_admission_feedback_max_token_usage = 0.9,
adaptive_admission_feedback_throughput_improvement_ratio = 0.02,
))]
#[expect(clippy::too_many_arguments)]
#[expect(
Expand Down Expand Up @@ -1080,6 +1109,7 @@ impl Router {
shutdown_grace_period_secs: u64,
request_id_headers: Option<Vec<String>>,
trust_tenant_header: bool,
prefer_trusted_tenant_header: bool,
tenant_header_name: String,
storage_context_headers: HashMap<String, String>,
pd_disaggregation: bool,
Expand Down Expand Up @@ -1170,6 +1200,11 @@ impl Router {
adaptive_admission_max_segments: usize,
adaptive_admission_min_load_coverage: f64,
adaptive_admission_cold_start_output_tokens: u32,
adaptive_admission_strategy: String,
adaptive_admission_feedback_probe_requests_per_healthy_replica: u32,
adaptive_admission_feedback_max_waiting_requests_per_healthy_replica: u32,
adaptive_admission_feedback_max_token_usage: f64,
adaptive_admission_feedback_throughput_improvement_ratio: f64,
) -> PyResult<Self> {
let mut all_urls = worker_urls.clone();

Expand Down Expand Up @@ -1241,6 +1276,7 @@ impl Router {
shutdown_grace_period_secs,
request_id_headers,
trust_tenant_header,
prefer_trusted_tenant_header,
tenant_header_name,
storage_context_headers,
pd_disaggregation,
Expand Down Expand Up @@ -1328,6 +1364,11 @@ impl Router {
adaptive_admission_max_segments,
adaptive_admission_min_load_coverage,
adaptive_admission_cold_start_output_tokens,
adaptive_admission_strategy,
adaptive_admission_feedback_probe_requests_per_healthy_replica,
adaptive_admission_feedback_max_waiting_requests_per_healthy_replica,
adaptive_admission_feedback_max_token_usage,
adaptive_admission_feedback_throughput_improvement_ratio,
})
}

Expand Down
67 changes: 61 additions & 6 deletions bindings/python/src/smg/router_args.py
Original file line number Diff line number Diff line change
Expand Up @@ -104,6 +104,9 @@ class RouterArgs:
prometheus_duration_buckets: list[float] | None = None
# Request ID headers configuration
request_id_headers: list[str] | None = None
trust_tenant_header: bool = False
prefer_trusted_tenant_header: bool = False
tenant_header_name: str = "x-smg-tenant-id"
# HTTP header to storage hook context mapping
storage_context_headers: dict[str, str] = dataclasses.field(default_factory=dict)
# Request timeout in seconds
Expand All @@ -124,12 +127,17 @@ class RouterArgs:
# Engine telemetry and predictive token-work admission.
engine_metrics: bool = False
adaptive_admission_mode: str = "off"
adaptive_admission_strategy: str = "predicted_work"
adaptive_admission_work_horizon_secs: float = 30.0
adaptive_admission_estimator_half_life_secs: float = 900.0
adaptive_admission_prior_observations: float = 20.0
adaptive_admission_max_segments: int = 50_000
adaptive_admission_min_load_coverage: float = 0.8
adaptive_admission_cold_start_output_tokens: int = 4096
adaptive_admission_feedback_probe_requests_per_healthy_replica: int = 2
adaptive_admission_feedback_max_waiting_requests_per_healthy_replica: int = 2
adaptive_admission_feedback_max_token_usage: float = 0.9
adaptive_admission_feedback_throughput_improvement_ratio: float = 0.02
# Token bucket refill rate (tokens per second). If not set, defaults to max_concurrent_requests
rate_limit_tokens_per_second: int | None = None
# Cluster-wide requests-per-second ceiling. Requires mesh and the same value on every gateway.
Expand Down Expand Up @@ -782,6 +790,25 @@ def add_cli_args(
" If not specified, uses common defaults."
),
)
request_group.add_argument(
f"--{prefix}trust-tenant-header",
action="store_true",
help="Trust the configured upstream tenant identity header",
)
request_group.add_argument(
f"--{prefix}prefer-trusted-tenant-header",
action="store_true",
help=(
"Prefer the trusted tenant header over authenticated proxy identity; "
"requires --trust-tenant-header"
),
)
request_group.add_argument(
f"--{prefix}tenant-header-name",
type=str,
default=RouterArgs.tenant_header_name,
help="Trusted tenant identity header name",
)
request_group.add_argument(
f"--{prefix}storage-context-headers",
type=str,
Expand Down Expand Up @@ -891,6 +918,12 @@ def add_cli_args(
default=RouterArgs.adaptive_admission_mode,
help="Predictive token-work admission mode",
)
adaptive_admission_group.add_argument(
f"--{prefix}adaptive-admission-strategy",
choices=["predicted_work", "engine_feedback"],
default=RouterArgs.adaptive_admission_strategy,
help="Admission signal: output-work prediction or direct engine feedback",
)
adaptive_admission_group.add_argument(
f"--{prefix}adaptive-admission-work-horizon-secs",
type=float,
Expand Down Expand Up @@ -927,6 +960,32 @@ def add_cli_args(
default=RouterArgs.adaptive_admission_cold_start_output_tokens,
help="Cold-start output-token prediction before observations",
)
adaptive_admission_group.add_argument(
f"--{prefix}adaptive-admission-feedback-probe-requests-per-healthy-replica",
type=int,
default=RouterArgs.adaptive_admission_feedback_probe_requests_per_healthy_replica,
help="Per-replica exploration margin above the learned throughput knee",
)
adaptive_admission_group.add_argument(
f"--{prefix}adaptive-admission-feedback-max-waiting-requests-per-healthy-replica",
type=int,
default=(
RouterArgs.adaptive_admission_feedback_max_waiting_requests_per_healthy_replica
),
help="Per-replica engine waiting queue that closes feedback admission",
)
adaptive_admission_group.add_argument(
f"--{prefix}adaptive-admission-feedback-max-token-usage",
type=float,
default=RouterArgs.adaptive_admission_feedback_max_token_usage,
help="Engine token/KV usage ratio that closes feedback admission",
)
adaptive_admission_group.add_argument(
f"--{prefix}adaptive-admission-feedback-throughput-improvement-ratio",
type=float,
default=RouterArgs.adaptive_admission_feedback_throughput_improvement_ratio,
help="Relative throughput gain required to raise the learned concurrency knee",
)

# Retry configuration
retry_group.add_argument(
Expand Down Expand Up @@ -1517,13 +1576,9 @@ def _parse_model_policies(values: list[str] | None) -> dict[str, str]:
def _validate_router_args(self):
if self.global_rate_limit_requests_per_second is not None:
if self.global_rate_limit_requests_per_second <= 0:
raise ValueError(
"global_rate_limit_requests_per_second must be greater than zero"
)
raise ValueError("global_rate_limit_requests_per_second must be greater than zero")
if not self.enable_mesh:
raise ValueError(
"global_rate_limit_requests_per_second requires enable_mesh=True"
)
raise ValueError("global_rate_limit_requests_per_second requires enable_mesh=True")

# Validate configuration based on mode
if self.epd_disaggregation:
Expand Down
42 changes: 38 additions & 4 deletions bindings/python/tests/test_arg_parser.py
Original file line number Diff line number Diff line change
Expand Up @@ -51,14 +51,22 @@ def test_default_values(self):
assert args.priority_scheduler_default_max_class == "default"
assert args.priority_scheduler_config is None
assert args.priority_scheduler_tenant_metric_top_n == 32
assert args.trust_tenant_header is False
assert args.prefer_trusted_tenant_header is False
assert args.tenant_header_name == "x-smg-tenant-id"
assert args.engine_metrics is False
assert args.adaptive_admission_mode == "off"
assert args.adaptive_admission_strategy == "predicted_work"
assert args.adaptive_admission_work_horizon_secs == 30.0
assert args.adaptive_admission_estimator_half_life_secs == 900.0
assert args.adaptive_admission_prior_observations == 20.0
assert args.adaptive_admission_max_segments == 50_000
assert args.adaptive_admission_min_load_coverage == 0.8
assert args.adaptive_admission_cold_start_output_tokens == 4096
assert args.adaptive_admission_feedback_probe_requests_per_healthy_replica == 2
assert args.adaptive_admission_feedback_max_waiting_requests_per_healthy_replica == 2
assert args.adaptive_admission_feedback_max_token_usage == 0.9
assert args.adaptive_admission_feedback_throughput_improvement_ratio == 0.02

def test_parse_priority_scheduler_options(self):
args = parse_router_args(
Expand All @@ -78,12 +86,28 @@ def test_parse_priority_scheduler_options(self):
assert args.priority_scheduler_config == "/tmp/priority.yaml"
assert args.priority_scheduler_tenant_metric_top_n == 16

def test_parse_preferred_trusted_tenant_header_options(self):
args = parse_router_args(
[
"--trust-tenant-header",
"--prefer-trusted-tenant-header",
"--tenant-header-name",
"x-comet-user",
]
)

assert args.trust_tenant_header is True
assert args.prefer_trusted_tenant_header is True
assert args.tenant_header_name == "x-comet-user"

def test_parse_adaptive_admission_options(self):
args = parse_router_args(
[
"--engine-metrics",
"--adaptive-admission-mode",
"shadow",
"--adaptive-admission-strategy",
"engine_feedback",
"--adaptive-admission-work-horizon-secs",
"45",
"--adaptive-admission-estimator-half-life-secs",
Expand All @@ -96,17 +120,30 @@ def test_parse_adaptive_admission_options(self):
"0.75",
"--adaptive-admission-cold-start-output-tokens",
"2048",
"--adaptive-admission-feedback-probe-requests-per-healthy-replica",
"3",
"--adaptive-admission-feedback-max-waiting-requests-per-healthy-replica",
"4",
"--adaptive-admission-feedback-max-token-usage",
"0.85",
"--adaptive-admission-feedback-throughput-improvement-ratio",
"0.03",
]
)

assert args.engine_metrics is True
assert args.adaptive_admission_mode == "shadow"
assert args.adaptive_admission_strategy == "engine_feedback"
assert args.adaptive_admission_work_horizon_secs == 45.0
assert args.adaptive_admission_estimator_half_life_secs == 600.0
assert args.adaptive_admission_prior_observations == 12.0
assert args.adaptive_admission_max_segments == 12_345
assert args.adaptive_admission_min_load_coverage == 0.75
assert args.adaptive_admission_cold_start_output_tokens == 2048
assert args.adaptive_admission_feedback_probe_requests_per_healthy_replica == 3
assert args.adaptive_admission_feedback_max_waiting_requests_per_healthy_replica == 4
assert args.adaptive_admission_feedback_max_token_usage == 0.85
assert args.adaptive_admission_feedback_throughput_improvement_ratio == 0.03

def test_parse_selector_valid(self):
"""Test parsing valid selector arguments."""
Expand Down Expand Up @@ -540,10 +577,7 @@ def test_valid_policies(self):
assert policy_from_str("round_robin") == PolicyType.RoundRobin
assert policy_from_str("cache_aware") == PolicyType.CacheAware
assert policy_from_str("power_of_two") == PolicyType.PowerOfTwo
assert (
policy_from_str("size_aware_power_of_two")
== PolicyType.SizeAwarePowerOfTwo
)
assert policy_from_str("size_aware_power_of_two") == PolicyType.SizeAwarePowerOfTwo
assert policy_from_str("consistent_hashing") == PolicyType.ConsistentHashing
assert policy_from_str("prefix_hash") == PolicyType.PrefixHash

Expand Down
Loading
Loading