move the PD bootstrap registry under api_server::disaggregation (#33895)
This commit is contained in:
Generated
+10
@@ -120,6 +120,15 @@ version = "1.4.2"
|
|||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
checksum = "c3d036a3c4ab069c7b410a2ce876bd74808d2d0888a82667669f8e783a898bf1"
|
checksum = "c3d036a3c4ab069c7b410a2ce876bd74808d2d0888a82667669f8e783a898bf1"
|
||||||
|
|
||||||
|
[[package]]
|
||||||
|
name = "arc-swap"
|
||||||
|
version = "1.9.2"
|
||||||
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
|
checksum = "c049c0be4daef0b145cb3555416b3b8ef5b7888a38aea1a3a155801fe7b0810b"
|
||||||
|
dependencies = [
|
||||||
|
"rustversion",
|
||||||
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "arg_enum_proc_macro"
|
name = "arg_enum_proc_macro"
|
||||||
version = "0.3.4"
|
version = "0.3.4"
|
||||||
@@ -3688,6 +3697,7 @@ dependencies = [
|
|||||||
name = "sglang-server"
|
name = "sglang-server"
|
||||||
version = "0.1.0"
|
version = "0.1.0"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
|
"arc-swap",
|
||||||
"async-stream",
|
"async-stream",
|
||||||
"axum 0.8.9",
|
"axum 0.8.9",
|
||||||
"bytemuck",
|
"bytemuck",
|
||||||
|
|||||||
@@ -32,7 +32,10 @@ tracing-subscriber = { workspace = true }
|
|||||||
tracing-appender = { workspace = true }
|
tracing-appender = { workspace = true }
|
||||||
uuid = { workspace = true }
|
uuid = { workspace = true }
|
||||||
|
|
||||||
|
arc-swap = "1"
|
||||||
axum = { version = "0.8.9", features = ["json", "tokio"] }
|
axum = { version = "0.8.9", features = ["json", "tokio"] }
|
||||||
|
# Safe POD slice casts (feature buffers viewed as bytes for the shm copy).
|
||||||
|
bytemuck = "1"
|
||||||
core_affinity = "0.8"
|
core_affinity = "0.8"
|
||||||
# the dynamo-tokenizers is deps on hf-hub, should bump version together.
|
# the dynamo-tokenizers is deps on hf-hub, should bump version together.
|
||||||
dynamo-tokenizers = "1.7.0"
|
dynamo-tokenizers = "1.7.0"
|
||||||
@@ -43,21 +46,19 @@ dynamo-parsers = "7.0.1"
|
|||||||
dynamo-protocols = "5.1.0"
|
dynamo-protocols = "5.1.0"
|
||||||
dynamo-renderer = "5.0.0"
|
dynamo-renderer = "5.0.0"
|
||||||
flume = "0.12.0"
|
flume = "0.12.0"
|
||||||
# Safe POD slice casts (feature buffers viewed as bytes for the shm copy).
|
|
||||||
bytemuck = "1"
|
|
||||||
itertools = "0.14"
|
|
||||||
hf-hub = { version = "0.4", default-features = false }
|
hf-hub = { version = "0.4", default-features = false }
|
||||||
|
itertools = "0.14"
|
||||||
# POSIX shm for the MM feature fan-out (`mm::ShmSegment`).
|
# POSIX shm for the MM feature fan-out (`mm::ShmSegment`).
|
||||||
libc = "0.2"
|
libc = "0.2"
|
||||||
# Same major as the workspace pyo3: the zero-copy MM drain (`take_mm`) moves
|
# Same major as the workspace pyo3: the zero-copy MM drain (`take_mm`) moves
|
||||||
# Rust vectors into numpy arrays.
|
# Rust vectors into numpy arrays.
|
||||||
numpy = "0.29.0"
|
numpy = "0.29.0"
|
||||||
rmp-serde = "1"
|
|
||||||
rmpv = { version = "1", features = ["with-serde"] }
|
|
||||||
# Pinned EXACTLY: this crate's accepted grammar defines the
|
# Pinned EXACTLY: this crate's accepted grammar defines the
|
||||||
# "anything Rust admits, Python can compile" invariant in `message::sampling`.
|
# "anything Rust admits, Python can compile" invariant in `message::sampling`.
|
||||||
# A minor bump can widen it and silently reopen a scheduler-killing hole.
|
# A minor bump can widen it and silently reopen a scheduler-killing hole.
|
||||||
regex-syntax = "=0.8.11"
|
regex-syntax = "=0.8.11"
|
||||||
|
rmp-serde = "1"
|
||||||
|
rmpv = { version = "1", features = ["with-serde"] }
|
||||||
# The pure-Rust core of the MM pipeline. `default-features = false` drops the
|
# The pure-Rust core of the MM pipeline. `default-features = false` drops the
|
||||||
# pyo3 bindings so it links as a plain rlib; renamed so `use sglang_mm::…` reads
|
# pyo3 bindings so it links as a plain rlib; renamed so `use sglang_mm::…` reads
|
||||||
# naturally while the crate keeps its own artifact name in the shared target/.
|
# naturally while the crate keeps its own artifact name in the shared target/.
|
||||||
|
|||||||
@@ -4,12 +4,12 @@
|
|||||||
//! frames (`data: {json}` … `[DONE]`), byte-compatible with Python
|
//! frames (`data: {json}` … `[DONE]`), byte-compatible with Python
|
||||||
//! `http_server.generate_request`; `/server_info` reuses it for one control result.
|
//! `http_server.generate_request`; `/server_info` reuses it for one control result.
|
||||||
mod common;
|
mod common;
|
||||||
|
mod disaggregation;
|
||||||
mod frame;
|
mod frame;
|
||||||
mod guard;
|
mod guard;
|
||||||
mod log;
|
mod log;
|
||||||
mod native_api;
|
mod native_api;
|
||||||
mod openai;
|
mod openai;
|
||||||
mod pd_bootstrap;
|
|
||||||
mod prefetch;
|
mod prefetch;
|
||||||
mod submit;
|
mod submit;
|
||||||
|
|
||||||
@@ -20,6 +20,7 @@ use axum::Router;
|
|||||||
use crate::runtime::ServerArgs;
|
use crate::runtime::ServerArgs;
|
||||||
use crate::tokenizer_manager::ActivityCounter;
|
use crate::tokenizer_manager::ActivityCounter;
|
||||||
use crate::tokenizer_manager::Senders;
|
use crate::tokenizer_manager::Senders;
|
||||||
|
use disaggregation::bootstrap as pd_bootstrap;
|
||||||
|
|
||||||
/// Shared handler state: submission handles, immutable server configuration,
|
/// Shared handler state: submission handles, immutable server configuration,
|
||||||
/// and the API-owned chat formatter.
|
/// and the API-owned chat formatter.
|
||||||
@@ -53,25 +54,32 @@ pub async fn serve(
|
|||||||
egress_activity,
|
egress_activity,
|
||||||
};
|
};
|
||||||
// Each endpoint module registers its own routes and merges here.
|
// Each endpoint module registers its own routes and merges here.
|
||||||
let mut app = Router::new()
|
let router = Router::new()
|
||||||
.merge(common::routes())
|
.merge(common::routes())
|
||||||
.merge(native_api::routes())
|
.merge(native_api::routes())
|
||||||
.merge(openai::routes())
|
.merge(openai::routes());
|
||||||
|
|
||||||
// TODO(auth): no API-key boundary yet. Python gates every route (except
|
// TODO(auth): no API-key boundary yet. Python gates every route (except
|
||||||
// /health*, /metrics*, OPTIONS) via `add_api_key_middleware`; until ported,
|
// /health*, /metrics*, OPTIONS) via `add_api_key_middleware`; until ported,
|
||||||
// a configured `api_key` does NOT protect these routes.
|
// a configured `api_key` does NOT protect these routes.
|
||||||
//
|
//
|
||||||
// No body limit, matching the Python server.
|
// No body limit, matching the Python server.
|
||||||
|
let mut app = router
|
||||||
.layer(axum::extract::DefaultBodyLimit::disable())
|
.layer(axum::extract::DefaultBodyLimit::disable())
|
||||||
.with_state(state);
|
.with_state(state);
|
||||||
|
|
||||||
|
// Prefill-only KV bootstrap registry. Merged AFTER `with_state` — its
|
||||||
|
// router carries its own Arc<Registry> state, so it cannot merge into the
|
||||||
|
// Router<AppState> above — and before `log::apply`, so bootstrap traffic
|
||||||
|
// shows in the access log.
|
||||||
if server_args.enable_pd_bootstrap() {
|
if server_args.enable_pd_bootstrap() {
|
||||||
// Merged after `with_state` (the registry carries its own state) and
|
let (routes, sweeper) = pd_bootstrap::router_and_sweeper();
|
||||||
// before `log::apply`, so bootstrap traffic shows in the access log.
|
|
||||||
let (bootstrap_routes, sweeper) = pd_bootstrap::router_and_sweeper();
|
|
||||||
tokio::spawn(sweeper); // cancelled with the runtime on shutdown
|
tokio::spawn(sweeper); // cancelled with the runtime on shutdown
|
||||||
app = app.merge(bootstrap_routes);
|
app = app.merge(routes);
|
||||||
tracing::info!("PD KV bootstrap registry mounted on the api listener");
|
tracing::info!("PD KV bootstrap registry mounted on the api listener");
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Apply logging and access log middleware.
|
||||||
let app = log::apply(app, &server_args);
|
let app = log::apply(app, &server_args);
|
||||||
|
|
||||||
// The listener was already bound synchronously in `runtime::start` (so a port
|
// The listener was already bound synchronously in `runtime::start` (so a port
|
||||||
|
|||||||
@@ -0,0 +1 @@
|
|||||||
|
pub(crate) mod bootstrap;
|
||||||
+153
-153
@@ -1,48 +1,37 @@
|
|||||||
//! PD KV bootstrap registry — the rust port of Python
|
//! PD KV bootstrap registry — rust port of Python `CommonKVBootstrapServer` (shared by all
|
||||||
//! `CommonKVBootstrapServer` (disaggregation/common/conn.py), which every
|
//! transfer backends): prefill ranks PUT `/route`, decode ranks GET routes and the `-1`-sentinel
|
||||||
//! transfer backend (mooncake / mori / nixl / ascend) subclasses without
|
//! topology, the PD router tracks per-room dp ranks; the wire format is Python-owned parity.
|
||||||
//! overrides, so this one implementation covers them all.
|
//! Mounted on the prefill api listener (bootstrap port = api port) before `init_disaggregation`.
|
||||||
//!
|
|
||||||
//! Prefill ranks PUT their `{rank_ip, rank_port}` to `/route`; decode ranks
|
|
||||||
//! GET per-rank routes and the aggregate topology (the all `-1` sentinel
|
|
||||||
//! query); the PD router registers/queries per-room dp ranks. The wire
|
|
||||||
//! protocol is Python-owned — field names, status codes, and the `-1`
|
|
||||||
//! sentinel below are parity pins, not this crate's design.
|
|
||||||
//!
|
|
||||||
//! Served on the api listener itself: `api_server::serve` merges
|
|
||||||
//! [`router_and_sweeper`]'s routes on every prefill server
|
|
||||||
//! (`ServerArgs::enable_pd_bootstrap()`). In rust-server mode the resolved
|
|
||||||
//! `disaggregation_bootstrap_port` is aliased to the api port, so KV managers
|
|
||||||
//! register here without knowing about the merge. The scheduler starts the
|
|
||||||
//! rust server BEFORE `init_disaggregation` — the KV managers register
|
|
||||||
//! synchronously there, with only a few bounded retries, so the routes must
|
|
||||||
//! already be accepting.
|
|
||||||
|
|
||||||
use std::collections::HashMap;
|
use std::collections::HashMap;
|
||||||
use std::sync::{Arc, Mutex};
|
use std::sync::{Arc, Mutex};
|
||||||
use std::time::{Duration, Instant};
|
use std::time::{Duration, Instant};
|
||||||
|
|
||||||
|
use arc_swap::ArcSwap;
|
||||||
use axum::extract::{Query, State};
|
use axum::extract::{Query, State};
|
||||||
use axum::http::StatusCode;
|
use axum::http::StatusCode;
|
||||||
use axum::response::{IntoResponse, Response};
|
use axum::response::{IntoResponse, Response};
|
||||||
use axum::routing::{post, put};
|
use axum::routing::{post, put};
|
||||||
use axum::{Json, Router};
|
use axum::{Json, Router};
|
||||||
|
use serde::{Deserialize, Serialize};
|
||||||
|
|
||||||
/// Python default: `SGLANG_DISAGGREGATION_BOOTSTRAP_ENTRY_CLEANUP_INTERVAL = EnvInt(120)`.
|
use crate::utils::response::json_error;
|
||||||
|
use crate::utils::serialize::{parse_int, parse_int_opt, parse_int_vec};
|
||||||
|
|
||||||
|
/// Python default: `SGLANG_DISAGGREGATION_BOOTSTRAP_ENTRY_CLEANUP_INTERVAL`.
|
||||||
const ENTRY_CLEANUP_INTERVAL_ENV: &str = "SGLANG_DISAGGREGATION_BOOTSTRAP_ENTRY_CLEANUP_INTERVAL";
|
const ENTRY_CLEANUP_INTERVAL_ENV: &str = "SGLANG_DISAGGREGATION_BOOTSTRAP_ENTRY_CLEANUP_INTERVAL";
|
||||||
const ENTRY_CLEANUP_INTERVAL_DEFAULT_SECS: u64 = 120;
|
const ENTRY_CLEANUP_INTERVAL_DEFAULT_SECS: u64 = 120;
|
||||||
|
const ROOM_SHARD_COUNT: usize = 64;
|
||||||
|
|
||||||
/// One registered prefill rank, exactly the JSON the Python decode side
|
/// Python's (`PrefillRankInfo`).
|
||||||
/// consumes (`PrefillRankInfo`).
|
#[derive(Clone, Serialize)]
|
||||||
#[derive(Clone, serde::Serialize)]
|
struct PrefillRankInfo {
|
||||||
struct RankInfo {
|
|
||||||
rank_ip: String,
|
rank_ip: String,
|
||||||
rank_port: i64,
|
rank_port: i64,
|
||||||
}
|
}
|
||||||
|
|
||||||
/// The all-`-1` sentinel response, exactly Python's
|
/// Python's (`PrefillServerInfo`).
|
||||||
/// `dataclasses.asdict(PrefillServerInfo)` shape.
|
#[derive(Serialize)]
|
||||||
#[derive(serde::Serialize)]
|
|
||||||
struct PrefillServerInfo {
|
struct PrefillServerInfo {
|
||||||
attn_tp_size: i64,
|
attn_tp_size: i64,
|
||||||
attn_cp_size: i64,
|
attn_cp_size: i64,
|
||||||
@@ -60,12 +49,56 @@ struct RoomEntry {
|
|||||||
registered_at: Instant,
|
registered_at: Instant,
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Mirror of the Python server's mutable state. First PUT wins for the
|
/// Mirror of the Python server's mutable state, split by write pattern.
|
||||||
/// topology scalars (Python's `if self.x is None` pattern); `registered_count`
|
|
||||||
/// counts raw PUTs, and readiness is `count >= dp*cp*tp*pp` — both verbatim
|
|
||||||
/// from Python, re-registrations included.
|
|
||||||
#[derive(Default)]
|
#[derive(Default)]
|
||||||
struct Registry {
|
struct Registry {
|
||||||
|
/// Copy-on-write topology.
|
||||||
|
topology: ArcSwap<Topology>,
|
||||||
|
/// Per shard locking map.
|
||||||
|
rooms: RoomShards,
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Room→dp-rank entries, sharded `room % `[`ROOM_SHARD_COUNT`].
|
||||||
|
struct RoomShards([Mutex<HashMap<i64, RoomEntry>>; ROOM_SHARD_COUNT]);
|
||||||
|
|
||||||
|
// Manual: `Default` is only derivable for arrays up to 32 elements.
|
||||||
|
impl Default for RoomShards {
|
||||||
|
fn default() -> Self {
|
||||||
|
Self(std::array::from_fn(|_| Mutex::new(HashMap::new())))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
impl RoomShards {
|
||||||
|
fn shard(&self, room: i64) -> &Mutex<HashMap<i64, RoomEntry>> {
|
||||||
|
&self.0[(room as u64 % ROOM_SHARD_COUNT as u64) as usize]
|
||||||
|
}
|
||||||
|
|
||||||
|
fn insert(&self, room: i64, entry: RoomEntry) {
|
||||||
|
self.shard(room).lock().unwrap().insert(room, entry);
|
||||||
|
}
|
||||||
|
|
||||||
|
fn dp_rank(&self, room: i64) -> Option<i64> {
|
||||||
|
self.shard(room)
|
||||||
|
.lock()
|
||||||
|
.unwrap()
|
||||||
|
.get(&room)
|
||||||
|
.map(|entry| entry.dp_rank)
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Drop entries older than `ttl`, one shard at a time.
|
||||||
|
fn sweep(&self, ttl: Duration) {
|
||||||
|
for shard in &self.0 {
|
||||||
|
shard
|
||||||
|
.lock()
|
||||||
|
.unwrap()
|
||||||
|
.retain(|_, entry| entry.registered_at.elapsed() <= ttl);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// The registration topology.
|
||||||
|
#[derive(Clone, Default)]
|
||||||
|
struct Topology {
|
||||||
attn_tp_size: Option<i64>,
|
attn_tp_size: Option<i64>,
|
||||||
attn_cp_size: Option<i64>,
|
attn_cp_size: Option<i64>,
|
||||||
dp_size: Option<i64>,
|
dp_size: Option<i64>,
|
||||||
@@ -77,12 +110,11 @@ struct Registry {
|
|||||||
prefill_http_port: Option<i64>,
|
prefill_http_port: Option<i64>,
|
||||||
/// Keyed `(dp_group, attn_cp_rank, attn_tp_rank, pp_rank)` — the flat form
|
/// Keyed `(dp_group, attn_cp_rank, attn_tp_rank, pp_rank)` — the flat form
|
||||||
/// of Python's nested `prefill_port_table` dicts.
|
/// of Python's nested `prefill_port_table` dicts.
|
||||||
prefill_ranks: HashMap<(i64, i64, i64, i64), RankInfo>,
|
prefill_ranks: HashMap<(i64, i64, i64, i64), PrefillRankInfo>,
|
||||||
room_to_dp_rank: HashMap<i64, RoomEntry>,
|
|
||||||
registered_count: i64,
|
registered_count: i64,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl Registry {
|
impl Topology {
|
||||||
/// `dp * cp * tp * pp` once every size is known (saturating: absurd sizes
|
/// `dp * cp * tp * pp` once every size is known (saturating: absurd sizes
|
||||||
/// stay "never ready" instead of overflowing).
|
/// stay "never ready" instead of overflowing).
|
||||||
fn expected(&self) -> Option<i64> {
|
fn expected(&self) -> Option<i64> {
|
||||||
@@ -100,36 +132,9 @@ impl Registry {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
type Shared = Arc<Mutex<Registry>>;
|
|
||||||
|
|
||||||
/// i64 that also accepts a numeric string — used exactly where Python coerces
|
|
||||||
/// with `int(data[...])`, so the wire stays as tolerant as the original.
|
|
||||||
#[derive(Clone, Copy)]
|
|
||||||
struct Int(i64);
|
|
||||||
|
|
||||||
impl<'de> serde::Deserialize<'de> for Int {
|
|
||||||
fn deserialize<D: serde::Deserializer<'de>>(d: D) -> Result<Self, D::Error> {
|
|
||||||
#[derive(serde::Deserialize)]
|
|
||||||
#[serde(untagged)]
|
|
||||||
enum Raw {
|
|
||||||
Int(i64),
|
|
||||||
Str(String),
|
|
||||||
}
|
|
||||||
match Raw::deserialize(d)? {
|
|
||||||
Raw::Int(v) => Ok(Int(v)),
|
|
||||||
// `int(...)` tolerates surrounding whitespace.
|
|
||||||
Raw::Str(s) => s
|
|
||||||
.trim()
|
|
||||||
.parse()
|
|
||||||
.map(Int)
|
|
||||||
.map_err(|_| serde::de::Error::custom(format!("invalid int: {s:?}"))),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
/// PUT /route payload (`CommonKVManager.register_to_bootstrap`).
|
/// PUT /route payload (`CommonKVManager.register_to_bootstrap`).
|
||||||
#[derive(serde::Deserialize)]
|
#[derive(Deserialize)]
|
||||||
struct RoutePut {
|
struct Route {
|
||||||
attn_tp_size: i64,
|
attn_tp_size: i64,
|
||||||
attn_tp_rank: i64,
|
attn_tp_rank: i64,
|
||||||
attn_cp_size: i64,
|
attn_cp_size: i64,
|
||||||
@@ -141,27 +146,21 @@ struct RoutePut {
|
|||||||
system_dp_size: i64,
|
system_dp_size: i64,
|
||||||
system_dp_rank: i64,
|
system_dp_rank: i64,
|
||||||
rank_ip: String,
|
rank_ip: String,
|
||||||
rank_port: Int,
|
#[serde(deserialize_with = "parse_int")]
|
||||||
page_size: Int,
|
rank_port: i64,
|
||||||
|
#[serde(deserialize_with = "parse_int")]
|
||||||
|
page_size: i64,
|
||||||
#[serde(default)]
|
#[serde(default)]
|
||||||
kv_cache_dtype: Option<String>,
|
kv_cache_dtype: Option<String>,
|
||||||
#[serde(default)]
|
#[serde(default, deserialize_with = "parse_int_opt")]
|
||||||
prefill_http_port: Option<Int>,
|
prefill_http_port: Option<i64>,
|
||||||
#[serde(default)]
|
#[serde(default)]
|
||||||
load_balance_method: Option<String>,
|
load_balance_method: Option<String>,
|
||||||
#[serde(default)]
|
#[serde(default)]
|
||||||
enable_dsa_cache_layer_split: Option<bool>,
|
enable_dsa_cache_layer_split: Option<bool>,
|
||||||
}
|
}
|
||||||
|
|
||||||
fn not_ready(registered_count: i64) -> Response {
|
async fn route_put(State(state): State<Arc<Registry>>, Json(body): Json<Route>) -> Response {
|
||||||
(
|
|
||||||
StatusCode::SERVICE_UNAVAILABLE,
|
|
||||||
format!("Prefill server not fully registered yet ({registered_count} workers registered)."),
|
|
||||||
)
|
|
||||||
.into_response()
|
|
||||||
}
|
|
||||||
|
|
||||||
async fn route_put(State(state): State<Shared>, Json(body): Json<RoutePut>) -> Response {
|
|
||||||
// `system_dp_size == 1` → attention-dp topology; else system-dp topology.
|
// `system_dp_size == 1` → attention-dp topology; else system-dp topology.
|
||||||
let dp_size = if body.system_dp_size == 1 {
|
let dp_size = if body.system_dp_size == 1 {
|
||||||
body.attn_dp_size
|
body.attn_dp_size
|
||||||
@@ -174,52 +173,57 @@ async fn route_put(State(state): State<Shared>, Json(body): Json<RoutePut>) -> R
|
|||||||
body.system_dp_rank
|
body.system_dp_rank
|
||||||
};
|
};
|
||||||
|
|
||||||
let mut reg = state.lock().unwrap();
|
// Copy-on-write update. `rcu` may re-run the closure under write
|
||||||
reg.attn_tp_size.get_or_insert(body.attn_tp_size);
|
// contention, so it only reads `body` and clones what it stores.
|
||||||
reg.attn_cp_size.get_or_insert(body.attn_cp_size);
|
state.topology.rcu(|current| {
|
||||||
reg.dp_size.get_or_insert(dp_size);
|
let mut topo = (**current).clone();
|
||||||
reg.pp_size.get_or_insert(body.pp_size);
|
topo.attn_tp_size.get_or_insert(body.attn_tp_size);
|
||||||
reg.page_size.get_or_insert(body.page_size.0);
|
topo.attn_cp_size.get_or_insert(body.attn_cp_size);
|
||||||
if reg.kv_cache_dtype.is_none() {
|
topo.dp_size.get_or_insert(dp_size);
|
||||||
reg.kv_cache_dtype = body.kv_cache_dtype;
|
topo.pp_size.get_or_insert(body.pp_size);
|
||||||
|
topo.page_size.get_or_insert(body.page_size);
|
||||||
|
if topo.kv_cache_dtype.is_none() {
|
||||||
|
topo.kv_cache_dtype = body.kv_cache_dtype.clone();
|
||||||
}
|
}
|
||||||
if reg.prefill_http_port.is_none() {
|
if topo.prefill_http_port.is_none() {
|
||||||
reg.prefill_http_port = body.prefill_http_port.map(|p| p.0);
|
topo.prefill_http_port = body.prefill_http_port;
|
||||||
}
|
}
|
||||||
reg.follow_bootstrap_room.get_or_insert(
|
topo.follow_bootstrap_room.get_or_insert(
|
||||||
body.load_balance_method
|
body.load_balance_method
|
||||||
.as_deref()
|
.as_deref()
|
||||||
.unwrap_or("follow_bootstrap_room")
|
.unwrap_or("follow_bootstrap_room")
|
||||||
== "follow_bootstrap_room",
|
== "follow_bootstrap_room",
|
||||||
);
|
);
|
||||||
reg.enable_dsa_cache_layer_split
|
topo.enable_dsa_cache_layer_split
|
||||||
.get_or_insert(body.enable_dsa_cache_layer_split.unwrap_or(false));
|
.get_or_insert(body.enable_dsa_cache_layer_split.unwrap_or(false));
|
||||||
|
topo.prefill_ranks.insert(
|
||||||
reg.prefill_ranks.insert(
|
|
||||||
(dp_group, body.attn_cp_rank, body.attn_tp_rank, body.pp_rank),
|
(dp_group, body.attn_cp_rank, body.attn_tp_rank, body.pp_rank),
|
||||||
RankInfo {
|
PrefillRankInfo {
|
||||||
rank_ip: body.rank_ip.clone(),
|
rank_ip: body.rank_ip.clone(),
|
||||||
rank_port: body.rank_port.0,
|
rank_port: body.rank_port,
|
||||||
},
|
},
|
||||||
);
|
);
|
||||||
reg.registered_count += 1;
|
topo.registered_count += 1;
|
||||||
|
topo
|
||||||
|
});
|
||||||
|
|
||||||
|
let topo = state.topology.load();
|
||||||
tracing::debug!(
|
tracing::debug!(
|
||||||
dp_group,
|
dp_group,
|
||||||
cp = body.attn_cp_rank,
|
cp = body.attn_cp_rank,
|
||||||
tp = body.attn_tp_rank,
|
tp = body.attn_tp_rank,
|
||||||
pp = body.pp_rank,
|
pp = body.pp_rank,
|
||||||
rank_ip = %body.rank_ip,
|
rank_ip = %body.rank_ip,
|
||||||
rank_port = body.rank_port.0,
|
rank_port = body.rank_port,
|
||||||
registered = reg.registered_count,
|
registered = topo.registered_count,
|
||||||
expected = reg.expected(),
|
expected = topo.expected(),
|
||||||
"registered prefill bootstrap rank"
|
"registered prefill bootstrap rank"
|
||||||
);
|
);
|
||||||
"OK".into_response()
|
"OK".into_response()
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn route_get(
|
async fn route_get(
|
||||||
State(state): State<Shared>,
|
State(state): State<Arc<Registry>>,
|
||||||
Query(query): Query<HashMap<String, String>>,
|
Query(query): Query<HashMap<String, String>>,
|
||||||
) -> Response {
|
) -> Response {
|
||||||
// A missing, empty (Python truthiness), or non-integer param → 400.
|
// A missing, empty (Python truthiness), or non-integer param → 400.
|
||||||
@@ -230,37 +234,41 @@ async fn route_get(
|
|||||||
rank("target_tp_rank"),
|
rank("target_tp_rank"),
|
||||||
rank("target_pp_rank"),
|
rank("target_pp_rank"),
|
||||||
) else {
|
) else {
|
||||||
return (
|
return json_error(
|
||||||
StatusCode::BAD_REQUEST,
|
StatusCode::BAD_REQUEST,
|
||||||
"Missing inputs for bootstrap server.",
|
"Missing inputs for bootstrap server.",
|
||||||
)
|
);
|
||||||
.into_response();
|
|
||||||
};
|
};
|
||||||
|
|
||||||
let reg = state.lock().unwrap();
|
let topo = state.topology.load();
|
||||||
// Python checks readiness in both branches; hoisted, same behavior.
|
// Python checks readiness in both branches; hoisted, same behavior.
|
||||||
if !reg.is_ready() {
|
if !topo.is_ready() {
|
||||||
return not_ready(reg.registered_count);
|
let registered_count = topo.registered_count;
|
||||||
|
return json_error(
|
||||||
|
StatusCode::SERVICE_UNAVAILABLE,
|
||||||
|
&format!(
|
||||||
|
"Prefill server not fully registered yet ({registered_count} workers registered)."
|
||||||
|
),
|
||||||
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
if (dp, cp, tp, pp) == (-1, -1, -1, -1) {
|
if (dp, cp, tp, pp) == (-1, -1, -1, -1) {
|
||||||
// Aggregate-topology sentinel. The sizes are Some — `is_ready` above
|
// Aggregate-topology sentinel.
|
||||||
// requires all four.
|
|
||||||
return Json(PrefillServerInfo {
|
return Json(PrefillServerInfo {
|
||||||
attn_tp_size: reg.attn_tp_size.unwrap(),
|
attn_tp_size: topo.attn_tp_size.unwrap(),
|
||||||
attn_cp_size: reg.attn_cp_size.unwrap(),
|
attn_cp_size: topo.attn_cp_size.unwrap(),
|
||||||
dp_size: reg.dp_size.unwrap(),
|
dp_size: topo.dp_size.unwrap(),
|
||||||
pp_size: reg.pp_size.unwrap(),
|
pp_size: topo.pp_size.unwrap(),
|
||||||
page_size: reg.page_size,
|
page_size: topo.page_size,
|
||||||
kv_cache_dtype: reg.kv_cache_dtype.clone(),
|
kv_cache_dtype: topo.kv_cache_dtype.clone(),
|
||||||
follow_bootstrap_room: reg.follow_bootstrap_room.unwrap_or(true),
|
follow_bootstrap_room: topo.follow_bootstrap_room.unwrap_or(true),
|
||||||
enable_dsa_cache_layer_split: reg.enable_dsa_cache_layer_split.unwrap_or(false),
|
enable_dsa_cache_layer_split: topo.enable_dsa_cache_layer_split.unwrap_or(false),
|
||||||
prefill_http_port: reg.prefill_http_port,
|
prefill_http_port: topo.prefill_http_port,
|
||||||
})
|
})
|
||||||
.into_response();
|
.into_response();
|
||||||
}
|
}
|
||||||
|
|
||||||
match reg.prefill_ranks.get(&(dp, cp, tp, pp)) {
|
match topo.prefill_ranks.get(&(dp, cp, tp, pp)) {
|
||||||
Some(info) => Json(info.clone()).into_response(),
|
Some(info) => Json(info.clone()).into_response(),
|
||||||
None => (
|
None => (
|
||||||
StatusCode::NOT_FOUND,
|
StatusCode::NOT_FOUND,
|
||||||
@@ -273,49 +281,51 @@ async fn route_get(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
#[derive(serde::Deserialize)]
|
#[derive(Deserialize)]
|
||||||
struct RegisterDpRank {
|
struct RegisterDpRank {
|
||||||
bootstrap_room: Int,
|
#[serde(deserialize_with = "parse_int")]
|
||||||
dp_rank: Int,
|
bootstrap_room: i64,
|
||||||
|
#[serde(deserialize_with = "parse_int")]
|
||||||
|
dp_rank: i64,
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn register_dp_rank(
|
async fn register_dp_rank(
|
||||||
State(state): State<Shared>,
|
State(state): State<Arc<Registry>>,
|
||||||
Json(body): Json<RegisterDpRank>,
|
Json(body): Json<RegisterDpRank>,
|
||||||
) -> Response {
|
) -> Response {
|
||||||
state.lock().unwrap().room_to_dp_rank.insert(
|
state.rooms.insert(
|
||||||
body.bootstrap_room.0,
|
body.bootstrap_room,
|
||||||
RoomEntry {
|
RoomEntry {
|
||||||
dp_rank: body.dp_rank.0,
|
dp_rank: body.dp_rank,
|
||||||
registered_at: Instant::now(),
|
registered_at: Instant::now(),
|
||||||
},
|
},
|
||||||
);
|
);
|
||||||
"OK".into_response()
|
"OK".into_response()
|
||||||
}
|
}
|
||||||
|
|
||||||
#[derive(serde::Deserialize)]
|
#[derive(Deserialize)]
|
||||||
struct QueryDpRanks {
|
struct QueryDpRanks {
|
||||||
bootstrap_rooms: Vec<Int>,
|
#[serde(deserialize_with = "parse_int_vec")]
|
||||||
|
bootstrap_rooms: Vec<i64>,
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Unknown rooms are silently omitted from the response, not an error. JSON
|
/// Unknown rooms are silently omitted from the response, not an error. JSON
|
||||||
/// object keys are strings — Python's `str(room_int)` for free.
|
/// object keys are strings — Python's `str(room_int)` for free.
|
||||||
async fn query_dp_ranks(State(state): State<Shared>, Json(body): Json<QueryDpRanks>) -> Response {
|
async fn query_dp_ranks(
|
||||||
let reg = state.lock().unwrap();
|
State(state): State<Arc<Registry>>,
|
||||||
|
Json(body): Json<QueryDpRanks>,
|
||||||
|
) -> Response {
|
||||||
let result: HashMap<String, i64> = body
|
let result: HashMap<String, i64> = body
|
||||||
.bootstrap_rooms
|
.bootstrap_rooms
|
||||||
.iter()
|
.iter()
|
||||||
.filter_map(|room| {
|
.filter_map(|room| Some((room.to_string(), state.rooms.dp_rank(*room)?)))
|
||||||
let entry = reg.room_to_dp_rank.get(&room.0)?;
|
|
||||||
Some((room.0.to_string(), entry.dp_rank))
|
|
||||||
})
|
|
||||||
.collect();
|
.collect();
|
||||||
Json(result).into_response()
|
Json(result).into_response()
|
||||||
}
|
}
|
||||||
|
|
||||||
/// No `/health` here: the merged api router already serves it (same 200 "OK"
|
/// No `/health` here: the merged api router already serves it (same 200 "OK"
|
||||||
/// the standalone Python bootstrap server answered, so probes are unchanged).
|
/// the standalone Python bootstrap server answered, so probes are unchanged).
|
||||||
fn router(state: Shared) -> Router {
|
fn router(state: Arc<Registry>) -> Router {
|
||||||
Router::new()
|
Router::new()
|
||||||
// Unmatched methods on a routed path get axum's built-in 405, matching
|
// Unmatched methods on a routed path get axum's built-in 405, matching
|
||||||
// Python's explicit method_not_allowed branch.
|
// Python's explicit method_not_allowed branch.
|
||||||
@@ -325,32 +335,22 @@ fn router(state: Shared) -> Router {
|
|||||||
.with_state(state)
|
.with_state(state)
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Drop `room_to_dp_rank` entries older than `interval` — Python's
|
/// Drop room entries
|
||||||
/// `_cleanup_expired_entries` loop (interval doubles as both period and TTL).
|
async fn cleanup_sweeper(state: Arc<Registry>) {
|
||||||
async fn cleanup_expired_entries(state: Shared, interval: Duration) {
|
|
||||||
loop {
|
|
||||||
tokio::time::sleep(interval).await;
|
|
||||||
state
|
|
||||||
.lock()
|
|
||||||
.unwrap()
|
|
||||||
.room_to_dp_rank
|
|
||||||
.retain(|_, entry| entry.registered_at.elapsed() <= interval);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
/// The registry routes plus their expiry sweeper, ready to mount on the api
|
|
||||||
/// router (`merge` the router, `tokio::spawn` the sweeper on the api runtime —
|
|
||||||
/// the runtime drop on shutdown cancels it along with the handlers).
|
|
||||||
pub(crate) fn router_and_sweeper() -> (Router, impl std::future::Future<Output = ()>) {
|
|
||||||
let state = Shared::default();
|
|
||||||
let cleanup_interval = Duration::from_secs(crate::environ::env_u64(
|
let cleanup_interval = Duration::from_secs(crate::environ::env_u64(
|
||||||
ENTRY_CLEANUP_INTERVAL_ENV,
|
ENTRY_CLEANUP_INTERVAL_ENV,
|
||||||
ENTRY_CLEANUP_INTERVAL_DEFAULT_SECS,
|
ENTRY_CLEANUP_INTERVAL_DEFAULT_SECS,
|
||||||
));
|
));
|
||||||
(
|
loop {
|
||||||
router(state.clone()),
|
tokio::time::sleep(cleanup_interval).await;
|
||||||
cleanup_expired_entries(state, cleanup_interval),
|
state.rooms.sweep(cleanup_interval);
|
||||||
)
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
pub(crate) fn router_and_sweeper() -> (Router, impl std::future::Future<Output = ()>) {
|
||||||
|
let state = Arc::new(Registry::default());
|
||||||
|
let sweeper = cleanup_sweeper(state.clone());
|
||||||
|
(router(state), sweeper)
|
||||||
}
|
}
|
||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
@@ -2,4 +2,5 @@
|
|||||||
|
|
||||||
pub mod regex;
|
pub mod regex;
|
||||||
pub mod response;
|
pub mod response;
|
||||||
|
pub mod serialize;
|
||||||
pub mod sock;
|
pub mod sock;
|
||||||
|
|||||||
@@ -2,8 +2,8 @@
|
|||||||
//!
|
//!
|
||||||
//! Two mechanics live here; the WIRE SHAPES stay owned by their endpoints:
|
//! Two mechanics live here; the WIRE SHAPES stay owned by their endpoints:
|
||||||
//! the native `{"error": {"message", "code"}}` body (Python
|
//! the native `{"error": {"message", "code"}}` body (Python
|
||||||
//! `http_server.generate_request` parity) built by [`error_value`], and the
|
//! `http_server.generate_request` parity) built by [`error_value`] and formed
|
||||||
//! SSE variant [`sse_error_response`] used by any
|
//! by [`json_error`], and the SSE variant [`sse_error_response`] used by any
|
||||||
//! endpoint family (native and OpenAI alike — the caller supplies its own
|
//! endpoint family (native and OpenAI alike — the caller supplies its own
|
||||||
//! body shape). The OpenAI error payload and the PD bootstrap registry's
|
//! body shape). The OpenAI error payload and the PD bootstrap registry's
|
||||||
//! plain-text bodies are protocol-owned and deliberately not unified here.
|
//! plain-text bodies are protocol-owned and deliberately not unified here.
|
||||||
@@ -25,6 +25,11 @@ pub fn error_value(code: u16, message: &str) -> serde_json::Value {
|
|||||||
serde_json::json!({ "error": { "message": message, "code": code } })
|
serde_json::json!({ "error": { "message": message, "code": code } })
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Unary native-shape error response: `code` + [`error_value`] body.
|
||||||
|
pub fn json_error(code: StatusCode, message: &str) -> Response {
|
||||||
|
error_response(code, error_value(code.as_u16(), message), false)
|
||||||
|
}
|
||||||
|
|
||||||
/// Form an error in the shape the client committed to: unary → `code` plus
|
/// Form an error in the shape the client committed to: unary → `code` plus
|
||||||
/// the JSON `body`; streaming → 200 with one SSE error frame + `[DONE]` (the
|
/// the JSON `body`; streaming → 200 with one SSE error frame + `[DONE]` (the
|
||||||
/// client is already reading a stream — Python answers in-stream too, from
|
/// client is already reading a stream — Python answers in-stream too, from
|
||||||
|
|||||||
@@ -0,0 +1,118 @@
|
|||||||
|
//! Python-`int(...)`-tolerant integer deserialization: accepts a JSON number
|
||||||
|
//! or a numeric string (surrounding whitespace ok), so wire fields the Python
|
||||||
|
//! side coerces with `int(data[...])` stay as tolerant here as the original.
|
||||||
|
//! Field types remain plain integers — apply per field with
|
||||||
|
//! `#[serde(deserialize_with = "parse_int")]`; generic over any
|
||||||
|
//! `FromStr + Deserialize` integer width. The `_opt` / `_vec` variants exist
|
||||||
|
//! because serde's `deserialize_with` does not compose through containers.
|
||||||
|
|
||||||
|
use serde::{Deserialize, Deserializer};
|
||||||
|
|
||||||
|
/// Wire form: a number or a numeric string.
|
||||||
|
#[derive(Deserialize)]
|
||||||
|
#[serde(untagged)]
|
||||||
|
enum RawInt<T> {
|
||||||
|
Num(T),
|
||||||
|
Str(String),
|
||||||
|
}
|
||||||
|
|
||||||
|
impl<T: std::str::FromStr> RawInt<T> {
|
||||||
|
fn resolve<E: serde::de::Error>(self) -> Result<T, E> {
|
||||||
|
match self {
|
||||||
|
RawInt::Num(v) => Ok(v),
|
||||||
|
// `int(...)` tolerates surrounding whitespace.
|
||||||
|
RawInt::Str(s) => s
|
||||||
|
.trim()
|
||||||
|
.parse()
|
||||||
|
.map_err(|_| E::custom(format!("invalid int: {s:?}"))),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn parse_int<'de, D, T>(deserializer: D) -> Result<T, D::Error>
|
||||||
|
where
|
||||||
|
D: Deserializer<'de>,
|
||||||
|
T: Deserialize<'de> + std::str::FromStr,
|
||||||
|
{
|
||||||
|
RawInt::deserialize(deserializer)?.resolve()
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn parse_int_opt<'de, D, T>(deserializer: D) -> Result<Option<T>, D::Error>
|
||||||
|
where
|
||||||
|
D: Deserializer<'de>,
|
||||||
|
T: Deserialize<'de> + std::str::FromStr,
|
||||||
|
{
|
||||||
|
Option::<RawInt<T>>::deserialize(deserializer)?
|
||||||
|
.map(RawInt::resolve)
|
||||||
|
.transpose()
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn parse_int_vec<'de, D, T>(deserializer: D) -> Result<Vec<T>, D::Error>
|
||||||
|
where
|
||||||
|
D: Deserializer<'de>,
|
||||||
|
T: Deserialize<'de> + std::str::FromStr,
|
||||||
|
{
|
||||||
|
Vec::<RawInt<T>>::deserialize(deserializer)?
|
||||||
|
.into_iter()
|
||||||
|
.map(RawInt::resolve)
|
||||||
|
.collect()
|
||||||
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod tests {
|
||||||
|
use super::*;
|
||||||
|
|
||||||
|
/// One struct exercising all three container shapes and two widths.
|
||||||
|
#[derive(Deserialize)]
|
||||||
|
struct Probe {
|
||||||
|
#[serde(deserialize_with = "parse_int")]
|
||||||
|
scalar: i64,
|
||||||
|
#[serde(deserialize_with = "parse_int")]
|
||||||
|
narrow: u16,
|
||||||
|
#[serde(default, deserialize_with = "parse_int_opt")]
|
||||||
|
opt: Option<i64>,
|
||||||
|
#[serde(deserialize_with = "parse_int_vec")]
|
||||||
|
vec: Vec<i64>,
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Python `int(...)` parity across scalar/Option/Vec and integer widths:
|
||||||
|
/// numbers and (whitespace-padded) numeric strings both parse; a missing
|
||||||
|
/// optional defaults. Guards the wire tolerance for every consumer, not
|
||||||
|
/// just the one field the pd_bootstrap HTTP contract test pins.
|
||||||
|
#[test]
|
||||||
|
fn accepts_numbers_and_numeric_strings() {
|
||||||
|
let p: Probe = serde_json::from_value(serde_json::json!({
|
||||||
|
"scalar": " 17000 ",
|
||||||
|
"narrow": "8998",
|
||||||
|
"opt": 3,
|
||||||
|
"vec": [1, "2", " 3 "],
|
||||||
|
}))
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(
|
||||||
|
(p.scalar, p.narrow, p.opt, p.vec),
|
||||||
|
(17000, 8998, Some(3), vec![1, 2, 3])
|
||||||
|
);
|
||||||
|
|
||||||
|
let p: Probe =
|
||||||
|
serde_json::from_value(serde_json::json!({"scalar": 1, "narrow": 2, "vec": []}))
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(p.opt, None, "missing optional defaults to None");
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Non-numeric strings and out-of-range values are errors, not silent
|
||||||
|
/// defaults — the tolerance is exactly `int(...)`-wide, no wider.
|
||||||
|
#[test]
|
||||||
|
fn rejects_non_numeric_and_out_of_range() {
|
||||||
|
for body in [
|
||||||
|
serde_json::json!({"scalar": "abc", "narrow": 1, "vec": []}),
|
||||||
|
serde_json::json!({"scalar": 1, "narrow": "70000", "vec": []}), // > u16::MAX
|
||||||
|
serde_json::json!({"scalar": 1, "narrow": 1, "vec": ["4x"]}),
|
||||||
|
serde_json::json!({"scalar": 1, "narrow": 1, "opt": "no", "vec": []}),
|
||||||
|
] {
|
||||||
|
assert!(
|
||||||
|
serde_json::from_value::<Probe>(body.clone()).is_err(),
|
||||||
|
"{body}"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user