Skip to content

Commit 2678123

Browse files
committed
feat(router): truncate routing tokens at configured media boundaries
Signed-off-by: Simo Lin <25425177+slin1237@users.noreply.github.com>
1 parent 08e4b34 commit 2678123

14 files changed

Lines changed: 580 additions & 15 deletions

File tree

‎bindings/python/src/lib.rs‎

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -516,6 +516,7 @@ struct Router {
516516
least_load_max_waiting_requests: u32,
517517
stream_request_bodies_over: u64,
518518
stream_body_stall_timeout_secs: u64,
519+
routing_token_boundaries: Vec<u32>,
519520
}
520521

521522
impl Router {
@@ -867,6 +868,7 @@ impl Router {
867868
.upstream_pool_idle_timeout_secs(self.upstream_pool_idle_timeout_secs)
868869
.stream_request_bodies_over(self.stream_request_bodies_over)
869870
.stream_body_stall_timeout_secs(self.stream_body_stall_timeout_secs)
871+
.routing_token_boundaries(self.routing_token_boundaries.clone())
870872
.multimodal_tensor_transport(multimodal_tensor_transport)
871873
.multimodal_shm_min_bytes(self.multimodal_shm_min_bytes)
872874
.routing_key_override(config::RoutingKeyOverrideConfig {
@@ -1030,6 +1032,7 @@ impl Router {
10301032
least_load_max_waiting_requests = 0,
10311033
stream_request_bodies_over = 0,
10321034
stream_body_stall_timeout_secs = 300,
1035+
routing_token_boundaries = vec![],
10331036
))]
10341037
#[expect(clippy::too_many_arguments)]
10351038
#[expect(
@@ -1170,6 +1173,7 @@ impl Router {
11701173
least_load_max_waiting_requests: u32,
11711174
stream_request_bodies_over: u64,
11721175
stream_body_stall_timeout_secs: u64,
1176+
routing_token_boundaries: Vec<u32>,
11731177
) -> PyResult<Self> {
11741178
let mut all_urls = worker_urls.clone();
11751179

@@ -1324,6 +1328,7 @@ impl Router {
13241328
least_load_max_waiting_requests,
13251329
stream_request_bodies_over,
13261330
stream_body_stall_timeout_secs,
1331+
routing_token_boundaries,
13271332
})
13281333
}
13291334

‎bindings/python/src/smg/router_args.py‎

Lines changed: 16 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -216,6 +216,7 @@ class RouterArgs:
216216
least_load_max_waiting_requests: int = 0
217217
stream_request_bodies_over: int = 0
218218
stream_body_stall_timeout_secs: int = 300
219+
routing_token_boundaries: list[int] = dataclasses.field(default_factory=list)
219220

220221
@staticmethod
221222
def add_cli_args(
@@ -613,6 +614,21 @@ def add_cli_args(
613614
" via --stream-request-bodies-over. 0 disables"
614615
),
615616
)
617+
routing_group.add_argument(
618+
f"--{prefix}routing-token-boundaries",
619+
type=int,
620+
nargs="*",
621+
action="extend",
622+
# None (not []) so an unset prefixed variant falls back to the
623+
# unprefixed backend value in from_cli_args.
624+
default=None,
625+
help=(
626+
"Token ids that end the routing-relevant prefix. Routing"
627+
" tokens are truncated at the first occurrence of any listed"
628+
" id before worker selection (e.g. multimodal placeholder"
629+
" ids). Empty disables truncation"
630+
),
631+
)
616632
routing_group.add_argument(
617633
f"--{prefix}dp-aware",
618634
action="store_true",

‎bindings/python/tests/test_arg_parser.py‎

Lines changed: 22 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -833,6 +833,26 @@ def test_prefixed_repeated_selector_accumulates(self):
833833

834834
assert router_args.selector == {"component": "engine", "env": "prod"}
835835

836+
def test_routing_token_boundaries_fall_back_to_backend_value(self):
837+
"""An unset prefixed list must not shadow the backend's value."""
838+
parser = argparse.ArgumentParser()
839+
RouterArgs.add_cli_args(parser, use_router_prefix=True)
840+
namespace = parser.parse_args([])
841+
namespace.routing_token_boundaries = [200091, 200038]
842+
843+
router_args = RouterArgs.from_cli_args(namespace, use_router_prefix=True)
844+
845+
assert router_args.routing_token_boundaries == [200091, 200038]
846+
847+
def test_routing_token_boundaries_prefixed_value_wins(self):
848+
parser = argparse.ArgumentParser()
849+
RouterArgs.add_cli_args(parser, use_router_prefix=True)
850+
namespace = parser.parse_args(["--router-routing-token-boundaries", "200091", "200038"])
851+
852+
router_args = RouterArgs.from_cli_args(namespace, use_router_prefix=True)
853+
854+
assert router_args.routing_token_boundaries == [200091, 200038]
855+
836856
def test_parse_repeated_model_alias_args(self):
837857
router_args = parse_router_args(
838858
[
@@ -1225,6 +1245,7 @@ class TestRouterArgsFieldOrder:
12251245
"least_load_max_waiting_requests",
12261246
"stream_request_bodies_over",
12271247
"stream_body_stall_timeout_secs",
1248+
"routing_token_boundaries",
12281249
]
12291250

12301251
def test_complete_field_sequence_is_frozen(self):
@@ -1249,6 +1270,7 @@ def test_new_fields_appended_after_positional_reserve(self):
12491270
"least_load_max_waiting_requests",
12501271
"stream_request_bodies_over",
12511272
"stream_body_stall_timeout_secs",
1273+
"routing_token_boundaries",
12521274
):
12531275
assert names.index(appended) > marker, (
12541276
f"{appended} must be appended after worker_startup_delay to "

‎model_gateway/src/config/builder.rs‎

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -215,6 +215,11 @@ impl RouterConfigBuilder {
215215
self
216216
}
217217

218+
pub fn routing_token_boundaries(mut self, ids: Vec<u32>) -> Self {
219+
self.config.routing_token_boundaries = ids;
220+
self
221+
}
222+
218223
pub fn upstream_pool_idle_timeout_secs(mut self, secs: u64) -> Self {
219224
self.config.upstream_pool_idle_timeout_secs = secs;
220225
self

‎model_gateway/src/config/types.rs‎

Lines changed: 28 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -37,6 +37,13 @@ pub struct RouterConfig {
3737
/// Per-request sticky-routing override (honors `X-SMG-Routing-Key`).
3838
#[serde(default)]
3939
pub routing_key_override: RoutingKeyOverrideConfig,
40+
/// Token ids that end the routing-relevant prefix. Routing tokens are
41+
/// truncated at the first occurrence of any listed id before worker
42+
/// selection: content past a media placeholder is never shareable across
43+
/// conversations, and match ratios over the full sequence shrink as
44+
/// conversations grow. Empty disables truncation.
45+
#[serde(default)]
46+
pub routing_token_boundaries: Vec<u32>,
4047
pub host: String,
4148
pub port: u16,
4249
/// Dedicated port for the isolated Kubernetes liveness/readiness/health
@@ -912,6 +919,7 @@ impl Default for RouterConfig {
912919
},
913920
policy: PolicyConfig::Random,
914921
routing_key_override: RoutingKeyOverrideConfig::default(),
922+
routing_token_boundaries: Vec::new(),
915923
host: "0.0.0.0".to_string(),
916924
port: 3001,
917925
health_check_port: None,
@@ -1206,6 +1214,26 @@ mod tests {
12061214
assert_eq!(with.stream_body_stall_timeout_secs, 0);
12071215
}
12081216

1217+
#[test]
1218+
fn test_routing_token_boundaries_serde_default_and_roundtrip() {
1219+
// Config files predating the field deserialize to no boundaries.
1220+
let mut json: serde_json::Value = serde_json::to_value(RouterConfig::default()).unwrap();
1221+
json.as_object_mut()
1222+
.unwrap()
1223+
.remove("routing_token_boundaries")
1224+
.unwrap();
1225+
let without: RouterConfig = serde_json::from_value(json).unwrap();
1226+
assert!(without.routing_token_boundaries.is_empty());
1227+
1228+
let config = RouterConfig::builder()
1229+
.regular_mode(vec![])
1230+
.routing_token_boundaries(vec![200091, 200038])
1231+
.build_unchecked();
1232+
let json = serde_json::to_string(&config).unwrap();
1233+
let with: RouterConfig = serde_json::from_str(&json).unwrap();
1234+
assert_eq!(with.routing_token_boundaries, vec![200091, 200038]);
1235+
}
1236+
12091237
#[test]
12101238
fn test_routing_mode_is_pd_mode() {
12111239
let regular = RoutingMode::Regular {

‎model_gateway/src/main.rs‎

Lines changed: 19 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -312,6 +312,12 @@ struct CliArgs {
312312
#[arg(long, default_value_t = false, help_heading = "Routing Policy")]
313313
routing_key_override: bool,
314314

315+
/// Token ids that end the routing-relevant prefix (e.g. multimodal
316+
/// placeholder ids); routing tokens are truncated at the first occurrence
317+
/// before worker selection. Empty disables
318+
#[arg(long, num_args = 0.., help_heading = "Routing Policy")]
319+
routing_token_boundaries: Vec<u32>,
320+
315321
/// Enable IGW (Inference Gateway) mode for multi-model support
316322
#[arg(long, default_value_t = false, help_heading = "Routing Policy")]
317323
enable_igw: bool,
@@ -1663,6 +1669,7 @@ impl CliArgs {
16631669
max_idle_secs: self.max_idle_secs,
16641670
assignment_mode: Self::parse_assignment_mode(&self.assignment_mode),
16651671
})
1672+
.routing_token_boundaries(self.routing_token_boundaries.clone())
16661673
.retries(!self.disable_retries)
16671674
.upstream_http2(self.upstream_http2)
16681675
.circuit_breaker(!self.disable_circuit_breaker)
@@ -1960,6 +1967,18 @@ mod tests {
19601967
assert_eq!(defaults.stream_body_stall_timeout_secs, 300);
19611968
}
19621969

1970+
/// Routing-token boundaries are a router-only setting and must flow into
1971+
/// `RouterConfig`.
1972+
#[test]
1973+
fn routing_token_boundaries_flag_flows_into_router_config() {
1974+
let cli = cli_args_from(&["--routing-token-boundaries", "200091", "200038"]);
1975+
let router_config = cli.to_router_config(vec![], vec![]).unwrap();
1976+
assert_eq!(router_config.routing_token_boundaries, vec![200091, 200038]);
1977+
1978+
let defaults = cli_args_from(&[]).to_router_config(vec![], vec![]).unwrap();
1979+
assert!(defaults.routing_token_boundaries.is_empty());
1980+
}
1981+
19631982
/// `--health-check-port` must flow into BOTH conversion paths
19641983
/// (`to_router_config` and `to_server_config`), mirroring the main
19651984
/// listener `--port` field exactly. This is the two-path config-plumbing

‎model_gateway/src/observability/metrics.rs‎

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1289,6 +1289,15 @@ impl Metrics {
12891289
.set(count as f64);
12901290
}
12911291

1292+
/// Record a routing-token truncation at a configured boundary id
1293+
pub fn record_routing_tokens_truncated(router_type: &'static str) {
1294+
counter!(
1295+
"smg_routing_tokens_truncated_total",
1296+
"router_type" => router_type
1297+
)
1298+
.increment(1);
1299+
}
1300+
12921301
/// Record health check result
12931302
pub fn record_worker_health_check(worker_type: &'static str, result: &'static str) {
12941303
counter!(

‎model_gateway/src/routers/common/mod.rs‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -28,5 +28,6 @@ pub mod openai_bridge;
2828
pub mod persistence_utils;
2929
pub mod realtime;
3030
pub mod retry;
31+
pub(crate) mod routing_tokens;
3132
pub mod sse;
3233
pub mod worker_selection;
Lines changed: 68 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,68 @@
1+
//! Routing-token boundary truncation shared by the HTTP, PD and gRPC
2+
//! selection paths.
3+
4+
use crate::observability::metrics::Metrics;
5+
6+
/// Slice `tokens` at the first configured boundary id. Content past a
7+
/// boundary is never shareable across conversations, and match ratios over
8+
/// the full sequence shrink as conversations grow.
9+
pub(crate) fn truncate_slice<'a>(
10+
tokens: &'a [u32],
11+
boundaries: &[u32],
12+
router_type: &'static str,
13+
) -> &'a [u32] {
14+
if boundaries.is_empty() {
15+
return tokens;
16+
}
17+
match tokens.iter().position(|id| boundaries.contains(id)) {
18+
Some(cut) => {
19+
Metrics::record_routing_tokens_truncated(router_type);
20+
&tokens[..cut]
21+
}
22+
None => tokens,
23+
}
24+
}
25+
26+
/// Owned variant: truncates in place, no reallocation.
27+
pub(crate) fn truncate_owned(
28+
mut tokens: Vec<u32>,
29+
boundaries: &[u32],
30+
router_type: &'static str,
31+
) -> Vec<u32> {
32+
let keep = truncate_slice(&tokens, boundaries, router_type).len();
33+
tokens.truncate(keep);
34+
tokens
35+
}
36+
37+
#[cfg(test)]
38+
mod tests {
39+
use super::*;
40+
41+
#[test]
42+
fn cuts_at_first_of_several_boundaries() {
43+
assert_eq!(
44+
truncate_slice(&[1, 2, 900, 3, 901], &[900, 901], "http"),
45+
&[1, 2]
46+
);
47+
assert_eq!(
48+
truncate_owned(vec![1, 2, 901, 3, 900], &[900, 901], "http"),
49+
vec![1, 2]
50+
);
51+
}
52+
53+
#[test]
54+
fn boundary_first_yields_empty() {
55+
assert_eq!(truncate_slice(&[900, 1], &[900], "http"), &[] as &[u32]);
56+
assert_eq!(
57+
truncate_owned(vec![900, 1], &[900], "http"),
58+
Vec::<u32>::new()
59+
);
60+
}
61+
62+
#[test]
63+
fn no_boundaries_or_no_match_is_identity() {
64+
assert_eq!(truncate_slice(&[1, 900, 2], &[], "http"), &[1, 900, 2]);
65+
assert_eq!(truncate_slice(&[1, 2, 3], &[900], "http"), &[1, 2, 3]);
66+
assert_eq!(truncate_owned(vec![1, 2], &[900], "http"), vec![1, 2]);
67+
}
68+
}

0 commit comments

Comments
 (0)