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 @@ -406,6 +406,7 @@ struct Router {
max_payload_size: usize,
dp_aware: bool,
dp_minimum_tokens_scheduler: bool,
upstream_http2: bool,
api_key: Option<String>,
log_dir: Option<String>,
log_level: Option<String>,
Expand Down Expand Up @@ -853,6 +854,7 @@ impl Router {
.maybe_storage_hook_wasm_path(self.storage_hook_wasm_path.as_deref())
.enable_wasm(self.enable_wasm)
.dp_aware(self.dp_aware)
.upstream_http2(self.upstream_http2)
.multimodal_tensor_transport(multimodal_tensor_transport)
.multimodal_shm_min_bytes(self.multimodal_shm_min_bytes)
.routing_key_override(config::RoutingKeyOverrideConfig {
Expand Down Expand Up @@ -1009,6 +1011,7 @@ impl Router {
prefix_token_count = 256,
prefix_hash_load_factor = 1.25,
prefix_hash_balance_abs_threshold = 10,
upstream_http2 = false,
))]
#[expect(clippy::too_many_arguments)]
#[expect(
Expand Down Expand Up @@ -1142,6 +1145,7 @@ impl Router {
prefix_token_count: usize,
prefix_hash_load_factor: f64,
prefix_hash_balance_abs_threshold: usize,
upstream_http2: bool,
) -> PyResult<Self> {
let mut all_urls = worker_urls.clone();

Expand Down Expand Up @@ -1192,6 +1196,7 @@ impl Router {
max_payload_size,
dp_aware,
dp_minimum_tokens_scheduler,
upstream_http2,
api_key,
log_dir,
log_level,
Expand Down
4 changes: 4 additions & 0 deletions bindings/python/src/smg/router.py
Original file line number Diff line number Diff line change
Expand Up @@ -191,6 +191,10 @@ class Router:
Default: 2^26
dp_aware: Enable data parallelism aware schedule. Default: False
dp_minimum_tokens_scheduler: Enable minimum tokens scheduler for data parallel group. Default: False
upstream_http2: Speak HTTP/2 to workers via prior knowledge (h2c on
cleartext), multiplexing every request to a worker over one
connection. Requires every HTTP worker to serve HTTP/2 without an
upgrade handshake. Default: False
enable_igw: Enable IGW (Inference-Gateway) mode for multi-model support. When
enabled, the router can manage multiple models simultaneously with per-model
load balancing policies. Default: False
Expand Down
11 changes: 11 additions & 0 deletions bindings/python/src/smg/router_args.py
Original file line number Diff line number Diff line change
Expand Up @@ -209,6 +209,7 @@ class RouterArgs:
prefix_token_count: int = 256
prefix_hash_load_factor: float = 1.25
prefix_hash_balance_abs_threshold: int = 10
upstream_http2: bool = False

@staticmethod
def add_cli_args(
Expand Down Expand Up @@ -329,6 +330,16 @@ def add_cli_args(
" (use brackets for IPv6, e.g., http://[::1]:8000 http://192.168.1.1:8000)"
),
)
worker_group.add_argument(
f"--{prefix}upstream-http2",
action="store_true",
help=(
"Speak HTTP/2 to workers via prior knowledge (h2c on cleartext),"
" multiplexing every request to a worker over one connection."
" Requires every HTTP worker to serve HTTP/2 without an upgrade"
" handshake."
),
)
worker_group.add_argument(
f"--{prefix}health-check-port",
type=int,
Expand Down
2 changes: 1 addition & 1 deletion model_gateway/Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -61,7 +61,7 @@ tracing-subscriber.workspace = true
# Workspace dependencies (with extra features)
axum = { workspace = true, features = ["macros", "multipart", "ws", "tracing"] }
bytemuck = { workspace = true, features = ["derive"] }
reqwest = { workspace = true, features = ["stream", "blocking", "json", "rustls", "multipart"] }
reqwest = { workspace = true, features = ["stream", "blocking", "json", "rustls", "multipart", "http2"] }
serde = { workspace = true, features = ["derive"] }
tokio = { workspace = true, features = ["full"] }
tokio-util.workspace = true
Expand Down
70 changes: 70 additions & 0 deletions model_gateway/src/app_context.rs
Original file line number Diff line number Diff line change
Expand Up @@ -464,6 +464,22 @@ impl AppContextBuilder {
.tcp_nodelay(true)
.tcp_keepalive(Some(Duration::from_secs(30)));

if config.upstream_http2 {
// Multiplex everything to a worker over one HTTP/2 connection.
// The default 64KB flow-control windows would let concurrent token
// streams throttle each other, so start large and let the adaptive
// window take over; h2 PING keepalives replace idle-connection
// churn and detect dead peers under long-lived streams.
client_builder = client_builder
.http2_prior_knowledge()
.http2_initial_stream_window_size(2 * 1024 * 1024)
.http2_initial_connection_window_size(16 * 1024 * 1024)
.http2_adaptive_window(true)
.http2_keep_alive_interval(Duration::from_secs(30))
.http2_keep_alive_timeout(Duration::from_secs(20))
.http2_keep_alive_while_idle(true);
}

// Force rustls backend when TLS is configured
if has_tls_config {
client_builder = client_builder.use_rustls_tls();
Expand Down Expand Up @@ -754,6 +770,60 @@ mod tests {
use super::*;
use crate::config::types::PolicyConfig;

/// Loopback echo server; axum::serve accepts HTTP/1.1 and prior-knowledge
/// h2c on the same listener, mirroring a dual-protocol engine.
async fn spawn_echo_server() -> String {
let app = axum::Router::new().route("/probe", axum::routing::get(|| async { "ok" }));
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
.await
.expect("bind echo server");
let addr = listener.local_addr().expect("echo server address");
#[expect(
clippy::disallowed_methods,
reason = "test server lives for the duration of the test process"
)]
tokio::spawn(async move {
axum::serve(listener, app).await.expect("echo serve");
});
format!("http://{addr}/probe")
}

fn built_client(upstream_http2: bool) -> Client {
let config = RouterConfig {
upstream_http2,
..RouterConfig::default()
};
AppContextBuilder::new()
.with_client(&config, 5)
.expect("client builds")
.client
.expect("client set")
}

#[tokio::test]
async fn upstream_http2_client_speaks_h2c_prior_knowledge() {
let url = spawn_echo_server().await;
let resp = built_client(true)
.get(&url)
.send()
.await
.expect("h2c request");
assert_eq!(resp.version(), http::Version::HTTP_2);
assert_eq!(resp.text().await.expect("body"), "ok");
}

#[tokio::test]
async fn default_client_stays_http1() {
let url = spawn_echo_server().await;
let resp = built_client(false)
.get(&url)
.send()
.await
.expect("h1 request");
assert_eq!(resp.version(), http::Version::HTTP_11);
assert_eq!(resp.text().await.expect("body"), "ok");
}

#[tokio::test]
async fn explicit_zero_rate_limit_disables_refill() {
let config = RouterConfig {
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 @@ -473,6 +473,11 @@ impl RouterConfigBuilder {
self
}

pub fn upstream_http2(mut self, enable: bool) -> Self {
self.config.upstream_http2 = enable;
self
}

// ==================== WASM ====================

pub fn enable_wasm(mut self, enable: bool) -> Self {
Expand Down
7 changes: 7 additions & 0 deletions model_gateway/src/config/types.rs
Original file line number Diff line number Diff line change
Expand Up @@ -195,6 +195,12 @@ pub struct RouterConfig {
/// PEM format, loaded from ca_cert_paths during config creation
#[serde(default)]
pub ca_certificates: Vec<Vec<u8>>,
/// Speak HTTP/2 to workers via prior knowledge (h2c on cleartext),
/// multiplexing every request to a worker over one connection instead of
/// one TCP connection per in-flight request. Requires every HTTP worker
/// to serve HTTP/2 without an upgrade handshake.
#[serde(default)]
pub upstream_http2: bool,
/// Loaded from mcp_config_path during config creation
#[serde(skip)]
pub mcp_config: Option<smg_mcp::McpConfig>,
Expand Down Expand Up @@ -907,6 +913,7 @@ impl Default for RouterConfig {
tokenizer_cache: TokenizerCacheConfig::default(),
client_identity: None,
ca_certificates: vec![],
upstream_http2: false,
mcp_config: None,
enable_wasm: false,
storage_hook_wasm_path: None,
Expand Down
7 changes: 7 additions & 0 deletions model_gateway/src/main.rs
Original file line number Diff line number Diff line change
Expand Up @@ -334,6 +334,12 @@ struct CliArgs {
#[arg(long, help_heading = "Worker Configuration")]
zmq_engine_count: Option<usize>,

/// Speak HTTP/2 to workers via prior knowledge (h2c on cleartext),
/// multiplexing every request to a worker over one connection. Requires
/// every HTTP worker to serve HTTP/2 without an upgrade handshake.
#[arg(long, default_value_t = false, help_heading = "Worker Configuration")]
upstream_http2: bool,

/// Interval in seconds between load monitor checks for PowerOfTwo routing
#[arg(long, default_value_t = 10, help_heading = "Load Monitoring")]
load_monitor_interval: u64,
Expand Down Expand Up @@ -1595,6 +1601,7 @@ impl CliArgs {
assignment_mode: Self::parse_assignment_mode(&self.assignment_mode),
})
.retries(!self.disable_retries)
.upstream_http2(self.upstream_http2)
.circuit_breaker(!self.disable_circuit_breaker)
.enable_wasm(self.enable_wasm)
.maybe_storage_hook_wasm_path(self.storage_hook_wasm_path.as_deref())
Expand Down
Loading