[Rust] Derive server address and accept signed env values (#37221)
This commit is contained in:
@@ -89,13 +89,11 @@ class RustServer:
|
|||||||
"ingress has no equivalent). Launch without SGLANG_RUST_SERVER, or "
|
"ingress has no equivalent). Launch without SGLANG_RUST_SERVER, or "
|
||||||
"drop --preferred-sampling-params and send those values per request."
|
"drop --preferred-sampling-params and send those values per request."
|
||||||
)
|
)
|
||||||
http_addr = f"{get_serving().host}:{get_serving().port}"
|
|
||||||
|
|
||||||
# Per-DP-rank HTTP port with client load balancing. `None` when DP is off,
|
# Per-DP-rank HTTP port with client load balancing. `None` when DP is off,
|
||||||
# so the rank is not conflated with rank 0 of a one-rank group.
|
# so the rank is not conflated with rank 0 of a one-rank group.
|
||||||
dp_rank = scheduler.ps.attn_dp_rank if scheduler.ps.dp_size > 1 else None
|
dp_rank = scheduler.ps.attn_dp_rank if scheduler.ps.dp_size > 1 else None
|
||||||
if dp_rank is not None:
|
listen_port = get_serving().port + (dp_rank or 0)
|
||||||
http_addr = f"{get_serving().host}:{get_serving().port + dp_rank}"
|
listen_addr = f"{get_serving().host}:{listen_port}"
|
||||||
|
|
||||||
launch_cores, server_cores = _partition_cores(
|
launch_cores, server_cores = _partition_cores(
|
||||||
mm_workers=(
|
mm_workers=(
|
||||||
@@ -109,7 +107,7 @@ class RustServer:
|
|||||||
_build_server_args(scheduler),
|
_build_server_args(scheduler),
|
||||||
# None -> run unpinned; the list carries the pinning decision.
|
# None -> run unpinned; the list carries the pinning decision.
|
||||||
cores=server_cores,
|
cores=server_cores,
|
||||||
http_addr=http_addr,
|
port_offset=dp_rank,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Multimodal models must have a Rust pipeline — there is no Python
|
# Multimodal models must have a Rust pipeline — there is no Python
|
||||||
@@ -161,7 +159,7 @@ class RustServer:
|
|||||||
)
|
)
|
||||||
logger.info(
|
logger.info(
|
||||||
"SGLANG_RUST_SERVER enabled, Rust server listen on %s%s",
|
"SGLANG_RUST_SERVER enabled, Rust server listen on %s%s",
|
||||||
http_addr,
|
listen_addr,
|
||||||
dp_note,
|
dp_note,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -338,10 +338,13 @@ fn router(state: Arc<Registry>) -> Router {
|
|||||||
|
|
||||||
/// Drop room entries
|
/// Drop room entries
|
||||||
async fn cleanup_sweeper(state: Arc<Registry>) {
|
async fn cleanup_sweeper(state: Arc<Registry>) {
|
||||||
let cleanup_interval = Duration::from_secs(environ::env_u64(
|
let cleanup_interval = Duration::from_secs(
|
||||||
ENTRY_CLEANUP_INTERVAL_ENV,
|
environ::env_i64(
|
||||||
ENTRY_CLEANUP_INTERVAL_DEFAULT_SECS,
|
ENTRY_CLEANUP_INTERVAL_ENV,
|
||||||
));
|
ENTRY_CLEANUP_INTERVAL_DEFAULT_SECS as i64,
|
||||||
|
)
|
||||||
|
.max(0) as u64,
|
||||||
|
);
|
||||||
loop {
|
loop {
|
||||||
tokio::time::sleep(cleanup_interval).await;
|
tokio::time::sleep(cleanup_interval).await;
|
||||||
state.rooms.sweep(cleanup_interval);
|
state.rooms.sweep(cleanup_interval);
|
||||||
|
|||||||
@@ -99,8 +99,9 @@ pub(super) fn native_error(code: StatusCode, message: &str, stream: bool) -> Res
|
|||||||
/// Python) decides whether `/health` shares it or is a plain 200 (routing the
|
/// Python) decides whether `/health` shares it or is a plain 200 (routing the
|
||||||
/// request already proves the frontend is up).
|
/// request already proves the frontend is up).
|
||||||
fn health_routes() -> Router<Arc<AppState>> {
|
fn health_routes() -> Router<Arc<AppState>> {
|
||||||
let timeout =
|
let timeout = std::time::Duration::from_secs(
|
||||||
std::time::Duration::from_secs(environ::env_u64("SGLANG_HEALTH_CHECK_TIMEOUT", 20));
|
environ::env_i64("SGLANG_HEALTH_CHECK_TIMEOUT", 20).max(0) as u64,
|
||||||
|
);
|
||||||
let probe = get(move |state: State<Arc<AppState>>| health_generate(state, timeout));
|
let probe = get(move |state: State<Arc<AppState>>| health_generate(state, timeout));
|
||||||
let health = if environ::env_bool("SGLANG_ENABLE_HEALTH_ENDPOINT_GENERATION", true) {
|
let health = if environ::env_bool("SGLANG_ENABLE_HEALTH_ENDPOINT_GENERATION", true) {
|
||||||
probe.clone()
|
probe.clone()
|
||||||
|
|||||||
@@ -17,8 +17,6 @@ mod multi_modality;
|
|||||||
mod tokenizer_manager;
|
mod tokenizer_manager;
|
||||||
mod utils;
|
mod utils;
|
||||||
|
|
||||||
use std::net::SocketAddr;
|
|
||||||
|
|
||||||
use pyo3::prelude::*;
|
use pyo3::prelude::*;
|
||||||
use pyo3::pybacked::PyBackedBytes;
|
use pyo3::pybacked::PyBackedBytes;
|
||||||
use pyo3::types::PyBytes;
|
use pyo3::types::PyBytes;
|
||||||
@@ -27,13 +25,9 @@ use crate::message::config::{
|
|||||||
DefaultSamplingParams, DisaggregationMode, MmFamily, MmResample, MmSpec, ModelConfig,
|
DefaultSamplingParams, DisaggregationMode, MmFamily, MmResample, MmSpec, ModelConfig,
|
||||||
RuntimeConfig, RustServerServerArgs, ServerArgs,
|
RuntimeConfig, RustServerServerArgs, ServerArgs,
|
||||||
};
|
};
|
||||||
|
use crate::utils::startup::{listen_addr, value_error};
|
||||||
use crate::utils::{logging, runtime};
|
use crate::utils::{logging, runtime};
|
||||||
|
|
||||||
/// A `ValueError` for a boot-time failure, as `"{context}: {err}"`.
|
|
||||||
fn value_error(context: &str, err: impl std::fmt::Display) -> PyErr {
|
|
||||||
pyo3::exceptions::PyValueError::new_err(format!("{context}: {err}"))
|
|
||||||
}
|
|
||||||
|
|
||||||
/// One drained MM result (see [`Server::take_mm_result`]), consumed by
|
/// One drained MM result (see [`Server::take_mm_result`]), consumed by
|
||||||
/// `RustMmProcessor.build_output` to build the scheduler's
|
/// `RustMmProcessor.build_output` to build the scheduler's
|
||||||
/// `MultimodalProcessorOutput`.
|
/// `MultimodalProcessorOutput`.
|
||||||
@@ -93,7 +87,7 @@ impl Server {
|
|||||||
#[new]
|
#[new]
|
||||||
#[pyo3(signature = (
|
#[pyo3(signature = (
|
||||||
server_args,
|
server_args,
|
||||||
http_addr = None,
|
port_offset = None,
|
||||||
to_scheduler_cap = 8192,
|
to_scheduler_cap = 8192,
|
||||||
from_scheduler_cap = 8192,
|
from_scheduler_cap = 8192,
|
||||||
stage_channel_cap = 8192,
|
stage_channel_cap = 8192,
|
||||||
@@ -104,7 +98,7 @@ impl Server {
|
|||||||
#[allow(clippy::too_many_arguments)]
|
#[allow(clippy::too_many_arguments)]
|
||||||
fn start(
|
fn start(
|
||||||
server_args: ServerArgs,
|
server_args: ServerArgs,
|
||||||
http_addr: Option<String>,
|
port_offset: Option<u16>, // DP rank; listen on server_args.port + offset
|
||||||
to_scheduler_cap: usize,
|
to_scheduler_cap: usize,
|
||||||
from_scheduler_cap: usize,
|
from_scheduler_cap: usize,
|
||||||
stage_channel_cap: usize,
|
stage_channel_cap: usize,
|
||||||
@@ -115,14 +109,10 @@ impl Server {
|
|||||||
server_args
|
server_args
|
||||||
.validate()
|
.validate()
|
||||||
.map_err(|e| value_error("server_args", e))?;
|
.map_err(|e| value_error("server_args", e))?;
|
||||||
// The HTTP listen address, tokenizer source/threads/shards all live in
|
// The host and base port come from `server_args`; DP ranks only supply
|
||||||
// `server_args`; resolve them from there so the scheduler doesn't re-pass
|
// their offset so this boundary has one source of truth for the address.
|
||||||
// them. The explicit params stay as optional overrides (per-DP-rank port,
|
let http_addr = listen_addr(&server_args, port_offset)
|
||||||
// pinning) and for standalone callers.
|
.map_err(|e| value_error("bad listen address", e))?;
|
||||||
let http_addr: SocketAddr = http_addr
|
|
||||||
.unwrap_or_else(|| server_args.bind())
|
|
||||||
.parse()
|
|
||||||
.map_err(|e| value_error("bad http_addr", e))?;
|
|
||||||
|
|
||||||
let cfg = RuntimeConfig {
|
let cfg = RuntimeConfig {
|
||||||
rust_server_args: RustServerServerArgs {
|
rust_server_args: RustServerServerArgs {
|
||||||
|
|||||||
@@ -14,7 +14,7 @@ use super::sampling::{SamplingParams, SamplingParamsInput};
|
|||||||
use super::types::{OneOrMany, OneOrManyItem, TokenIds};
|
use super::types::{OneOrMany, OneOrManyItem, TokenIds};
|
||||||
use crate::message::ids::Rid;
|
use crate::message::ids::Rid;
|
||||||
use crate::utils::fsm::RequestState;
|
use crate::utils::fsm::RequestState;
|
||||||
use crate::utils::{environ::env_u64, error::Error};
|
use crate::utils::{environ::env_i64, error::Error};
|
||||||
|
|
||||||
/// Hard cap on how many scheduler requests one `/generate` HTTP call may expand
|
/// Hard cap on how many scheduler requests one `/generate` HTTP call may expand
|
||||||
/// into. Every column below is allocated per item before anything is dispatched,
|
/// into. Every column below is allocated per item before anything is dispatched,
|
||||||
@@ -27,8 +27,12 @@ use crate::utils::{environ::env_u64, error::Error};
|
|||||||
/// `python/sglang/srt/environ.py`, which owns the default). Memoized because the
|
/// `python/sglang/srt/environ.py`, which owns the default). Memoized because the
|
||||||
/// value is process-static — Python sets it before launching this server — and a
|
/// value is process-static — Python sets it before launching this server — and a
|
||||||
/// per-request `env::var` would take a lock on the hot path for a constant.
|
/// per-request `env::var` would take a lock on the hot path for a constant.
|
||||||
static MAX_BATCH_REQS_PER_HTTP_REQ: LazyLock<usize> =
|
static MAX_BATCH_REQS_PER_HTTP_REQ: LazyLock<i64> =
|
||||||
LazyLock::new(|| env_u64("SGLANG_MAX_BATCH_REQS_PER_HTTP_REQ", 4096) as usize);
|
LazyLock::new(|| env_i64("SGLANG_MAX_BATCH_REQS_PER_HTTP_REQ", 4096));
|
||||||
|
|
||||||
|
fn batch_size_exceeds_limit(batch_size: usize, limit: i64) -> bool {
|
||||||
|
limit >= 0 && batch_size as u128 > limit as u128
|
||||||
|
}
|
||||||
|
|
||||||
/// Hard cap on the total bytes a broadcast value may clone into the batch (see
|
/// Hard cap on the total bytes a broadcast value may clone into the batch (see
|
||||||
/// the `One` arms of the fan-out).
|
/// the `One` arms of the fan-out).
|
||||||
@@ -170,7 +174,7 @@ impl GenerateBody {
|
|||||||
(None, Some(OneOrMany::Many(v))) => v.len(),
|
(None, Some(OneOrMany::Many(v))) => v.len(),
|
||||||
_ => 1,
|
_ => 1,
|
||||||
};
|
};
|
||||||
if declared_n > *MAX_BATCH_REQS_PER_HTTP_REQ {
|
if batch_size_exceeds_limit(declared_n, *MAX_BATCH_REQS_PER_HTTP_REQ) {
|
||||||
return Err(Error::Validation(format!(
|
return Err(Error::Validation(format!(
|
||||||
"batch size {declared_n} exceeds the maximum of {}",
|
"batch size {declared_n} exceeds the maximum of {}",
|
||||||
*MAX_BATCH_REQS_PER_HTTP_REQ
|
*MAX_BATCH_REQS_PER_HTTP_REQ
|
||||||
@@ -1115,19 +1119,16 @@ mod tests {
|
|||||||
/// capped before any column is built.
|
/// capped before any column is built.
|
||||||
#[test]
|
#[test]
|
||||||
fn oversized_batches_are_rejected_before_allocating() {
|
fn oversized_batches_are_rejected_before_allocating() {
|
||||||
let texts: Vec<String> = (0..*MAX_BATCH_REQS_PER_HTTP_REQ + 1)
|
let cap = usize::try_from(*MAX_BATCH_REQS_PER_HTTP_REQ).unwrap();
|
||||||
.map(|i| i.to_string())
|
let texts: Vec<String> = (0..cap + 1).map(|i| i.to_string()).collect();
|
||||||
.collect();
|
|
||||||
let body = serde_json::json!({ "text": texts }).to_string();
|
let body = serde_json::json!({ "text": texts }).to_string();
|
||||||
let err = requests(&body).unwrap_err().to_string();
|
let err = requests(&body).unwrap_err().to_string();
|
||||||
assert!(err.contains("exceeds the maximum"), "{err}");
|
assert!(err.contains("exceeds the maximum"), "{err}");
|
||||||
|
|
||||||
// At the cap it is accepted.
|
// At the cap it is accepted.
|
||||||
let texts: Vec<String> = (0..*MAX_BATCH_REQS_PER_HTTP_REQ)
|
let texts: Vec<String> = (0..cap).map(|i| i.to_string()).collect();
|
||||||
.map(|i| i.to_string())
|
|
||||||
.collect();
|
|
||||||
let (reqs, _) = requests(&serde_json::json!({ "text": texts }).to_string()).unwrap();
|
let (reqs, _) = requests(&serde_json::json!({ "text": texts }).to_string()).unwrap();
|
||||||
assert_eq!(reqs.len(), *MAX_BATCH_REQS_PER_HTTP_REQ);
|
assert_eq!(reqs.len(), cap);
|
||||||
|
|
||||||
// A small batch with a huge broadcast `custom_params` is the quadratic case:
|
// A small batch with a huge broadcast `custom_params` is the quadratic case:
|
||||||
// few items, but each clone carries the whole blob. The item count is a
|
// few items, but each clone carries the whole blob. The item count is a
|
||||||
@@ -1144,6 +1145,13 @@ mod tests {
|
|||||||
assert!(err.contains("would allocate more than"), "{err}");
|
assert!(err.contains("would allocate more than"), "{err}");
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn negative_batch_limit_disables_the_item_cap() {
|
||||||
|
assert!(!batch_size_exceeds_limit(usize::MAX, -1));
|
||||||
|
assert!(batch_size_exceeds_limit(11, 10));
|
||||||
|
assert!(!batch_size_exceeds_limit(10, 10));
|
||||||
|
}
|
||||||
|
|
||||||
/// `token_ids_logprob` mirrors Python `_normalize_batch`'s nested-structure
|
/// `token_ids_logprob` mirrors Python `_normalize_batch`'s nested-structure
|
||||||
/// branch: a flat list broadcasts to every prompt, a list of lists is
|
/// branch: a flat list broadcasts to every prompt, a list of lists is
|
||||||
/// per-prompt. Regression — the whole value used to be cloned to every item.
|
/// per-prompt. Regression — the whole value used to be cloned to every item.
|
||||||
|
|||||||
@@ -9,4 +9,5 @@ pub mod response;
|
|||||||
pub mod runtime;
|
pub mod runtime;
|
||||||
pub mod serialize;
|
pub mod serialize;
|
||||||
pub mod sock;
|
pub mod sock;
|
||||||
|
pub mod startup;
|
||||||
pub mod threads;
|
pub mod threads;
|
||||||
|
|||||||
@@ -13,14 +13,10 @@ pub fn env_bool(name: &str, default: bool) -> bool {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Deliberately restricted unsigned parser — NOT Python `int()` semantics.
|
/// Signed integer parser. Accepts the `i64::from_str` grammar, including
|
||||||
/// Accepts only what `u64::from_str` does: ASCII digits with an optional
|
/// negative values, while invalid or out-of-range values warn and use the
|
||||||
/// leading `+`, up to `u64::MAX`. Inputs Python's `EnvInt` would accept —
|
/// default.
|
||||||
/// surrounding whitespace (`" 45 "`), digit-group underscores (`"4_5"`),
|
pub fn env_i64(name: &str, default: i64) -> i64 {
|
||||||
/// negatives, values above `u64::MAX`, non-ASCII digits — warn and fall back
|
|
||||||
/// to the default, like any other invalid value. Callers are counts/sizes, so
|
|
||||||
/// strictness over parity is intentional here.
|
|
||||||
pub fn env_u64(name: &str, default: u64) -> u64 {
|
|
||||||
read(name, default, |raw| raw.parse().ok())
|
read(name, default, |raw| raw.parse().ok())
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -71,30 +67,32 @@ mod tests {
|
|||||||
assert!(!env_bool("SGLANG_TEST_ENV_BOOL_UNSET", false));
|
assert!(!env_bool("SGLANG_TEST_ENV_BOOL_UNSET", false));
|
||||||
}
|
}
|
||||||
|
|
||||||
/// `env_u64`: strict `u64::from_str` grammar; everything else — including
|
/// `env_i64`: strict `i64::from_str` grammar, including negative values;
|
||||||
/// int()-valid inputs the doc calls out as deliberately rejected — → default.
|
/// everything else falls back to the default.
|
||||||
#[test]
|
#[test]
|
||||||
fn env_u64_parses_or_defaults() {
|
fn env_i64_parses_or_defaults() {
|
||||||
for (i, (raw, want)) in [
|
for (i, (raw, want)) in [
|
||||||
("45", 45),
|
("45", 45),
|
||||||
("+45", 45), // u64::from_str allows a leading `+`
|
("+45", 45),
|
||||||
|
("-1", -1),
|
||||||
|
("-9223372036854775808", i64::MIN),
|
||||||
|
("9223372036854775807", i64::MAX),
|
||||||
// Invalid → default.
|
// Invalid → default.
|
||||||
("20s", 20),
|
("20s", 20),
|
||||||
("", 20),
|
("", 20),
|
||||||
// int()-valid but deliberately rejected → default.
|
|
||||||
(" 45 ", 20),
|
(" 45 ", 20),
|
||||||
("4_5", 20),
|
("4_5", 20),
|
||||||
("-1", 20),
|
("9223372036854775808", 20),
|
||||||
("18446744073709551616", 20), // u64::MAX + 1
|
("-9223372036854775809", 20),
|
||||||
("١٢", 20), // non-ASCII digits
|
("١٢", 20), // non-ASCII digits
|
||||||
]
|
]
|
||||||
.into_iter()
|
.into_iter()
|
||||||
.enumerate()
|
.enumerate()
|
||||||
{
|
{
|
||||||
let name = format!("SGLANG_TEST_ENV_U64_{i}");
|
let name = format!("SGLANG_TEST_ENV_I64_{i}");
|
||||||
unsafe { std::env::set_var(&name, raw) };
|
unsafe { std::env::set_var(&name, raw) };
|
||||||
assert_eq!(env_u64(&name, 20), want, "value {raw:?}");
|
assert_eq!(env_i64(&name, 20), want, "value {raw:?}");
|
||||||
}
|
}
|
||||||
assert_eq!(env_u64("SGLANG_TEST_ENV_U64_UNSET", 20), 20);
|
assert_eq!(env_i64("SGLANG_TEST_ENV_I64_UNSET", 20), 20);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,61 @@
|
|||||||
|
//! Helpers for the Python-facing server startup boundary.
|
||||||
|
|
||||||
|
use std::net::SocketAddr;
|
||||||
|
|
||||||
|
use pyo3::PyErr;
|
||||||
|
use pyo3::exceptions::PyValueError;
|
||||||
|
|
||||||
|
use crate::message::config::ServerArgs;
|
||||||
|
|
||||||
|
/// A `ValueError` for a boot-time failure, as `"{context}: {err}"`.
|
||||||
|
pub(crate) fn value_error(context: &str, err: impl std::fmt::Display) -> PyErr {
|
||||||
|
PyValueError::new_err(format!("{context}: {err}"))
|
||||||
|
}
|
||||||
|
|
||||||
|
pub(crate) fn listen_addr(
|
||||||
|
server_args: &ServerArgs,
|
||||||
|
port_offset: Option<u16>,
|
||||||
|
) -> Result<SocketAddr, String> {
|
||||||
|
let offset = port_offset.unwrap_or_default();
|
||||||
|
let port = server_args
|
||||||
|
.port
|
||||||
|
.checked_add(offset)
|
||||||
|
.ok_or_else(|| format!("port {} + offset {offset} exceeds 65535", server_args.port))?;
|
||||||
|
let mut addr: SocketAddr = server_args
|
||||||
|
.bind()
|
||||||
|
.parse()
|
||||||
|
.map_err(|err| format!("invalid host {:?}: {err}", server_args.host))?;
|
||||||
|
addr.set_port(port);
|
||||||
|
Ok(addr)
|
||||||
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod tests {
|
||||||
|
use super::*;
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn listen_addr_uses_server_host_and_port_offset() {
|
||||||
|
let args = ServerArgs {
|
||||||
|
host: "::".into(),
|
||||||
|
port: 30_000,
|
||||||
|
..Default::default()
|
||||||
|
};
|
||||||
|
assert_eq!(
|
||||||
|
listen_addr(&args, None).unwrap(),
|
||||||
|
"[::]:30000".parse().unwrap()
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
listen_addr(&args, Some(7)).unwrap(),
|
||||||
|
"[::]:30007".parse().unwrap()
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn listen_addr_rejects_port_overflow() {
|
||||||
|
let args = ServerArgs {
|
||||||
|
port: u16::MAX,
|
||||||
|
..Default::default()
|
||||||
|
};
|
||||||
|
assert!(listen_addr(&args, Some(1)).unwrap_err().contains("exceeds"));
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user