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"
|
||||
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]]
|
||||
name = "arg_enum_proc_macro"
|
||||
version = "0.3.4"
|
||||
@@ -3688,6 +3697,7 @@ dependencies = [
|
||||
name = "sglang-server"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"arc-swap",
|
||||
"async-stream",
|
||||
"axum 0.8.9",
|
||||
"bytemuck",
|
||||
|
||||
@@ -32,7 +32,10 @@ tracing-subscriber = { workspace = true }
|
||||
tracing-appender = { workspace = true }
|
||||
uuid = { workspace = true }
|
||||
|
||||
arc-swap = "1"
|
||||
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"
|
||||
# the dynamo-tokenizers is deps on hf-hub, should bump version together.
|
||||
dynamo-tokenizers = "1.7.0"
|
||||
@@ -43,21 +46,19 @@ dynamo-parsers = "7.0.1"
|
||||
dynamo-protocols = "5.1.0"
|
||||
dynamo-renderer = "5.0.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 }
|
||||
itertools = "0.14"
|
||||
# POSIX shm for the MM feature fan-out (`mm::ShmSegment`).
|
||||
libc = "0.2"
|
||||
# Same major as the workspace pyo3: the zero-copy MM drain (`take_mm`) moves
|
||||
# Rust vectors into numpy arrays.
|
||||
numpy = "0.29.0"
|
||||
rmp-serde = "1"
|
||||
rmpv = { version = "1", features = ["with-serde"] }
|
||||
# Pinned EXACTLY: this crate's accepted grammar defines the
|
||||
# "anything Rust admits, Python can compile" invariant in `message::sampling`.
|
||||
# A minor bump can widen it and silently reopen a scheduler-killing hole.
|
||||
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
|
||||
# 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/.
|
||||
|
||||
@@ -4,12 +4,12 @@
|
||||
//! frames (`data: {json}` … `[DONE]`), byte-compatible with Python
|
||||
//! `http_server.generate_request`; `/server_info` reuses it for one control result.
|
||||
mod common;
|
||||
mod disaggregation;
|
||||
mod frame;
|
||||
mod guard;
|
||||
mod log;
|
||||
mod native_api;
|
||||
mod openai;
|
||||
mod pd_bootstrap;
|
||||
mod prefetch;
|
||||
mod submit;
|
||||
|
||||
@@ -20,6 +20,7 @@ use axum::Router;
|
||||
use crate::runtime::ServerArgs;
|
||||
use crate::tokenizer_manager::ActivityCounter;
|
||||
use crate::tokenizer_manager::Senders;
|
||||
use disaggregation::bootstrap as pd_bootstrap;
|
||||
|
||||
/// Shared handler state: submission handles, immutable server configuration,
|
||||
/// and the API-owned chat formatter.
|
||||
@@ -53,25 +54,32 @@ pub async fn serve(
|
||||
egress_activity,
|
||||
};
|
||||
// Each endpoint module registers its own routes and merges here.
|
||||
let mut app = Router::new()
|
||||
let router = Router::new()
|
||||
.merge(common::routes())
|
||||
.merge(native_api::routes())
|
||||
.merge(openai::routes())
|
||||
// TODO(auth): no API-key boundary yet. Python gates every route (except
|
||||
// /health*, /metrics*, OPTIONS) via `add_api_key_middleware`; until ported,
|
||||
// a configured `api_key` does NOT protect these routes.
|
||||
//
|
||||
// No body limit, matching the Python server.
|
||||
.merge(openai::routes());
|
||||
|
||||
// TODO(auth): no API-key boundary yet. Python gates every route (except
|
||||
// /health*, /metrics*, OPTIONS) via `add_api_key_middleware`; until ported,
|
||||
// a configured `api_key` does NOT protect these routes.
|
||||
//
|
||||
// No body limit, matching the Python server.
|
||||
let mut app = router
|
||||
.layer(axum::extract::DefaultBodyLimit::disable())
|
||||
.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() {
|
||||
// Merged after `with_state` (the registry carries its own state) and
|
||||
// before `log::apply`, so bootstrap traffic shows in the access log.
|
||||
let (bootstrap_routes, sweeper) = pd_bootstrap::router_and_sweeper();
|
||||
let (routes, sweeper) = pd_bootstrap::router_and_sweeper();
|
||||
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");
|
||||
}
|
||||
|
||||
// Apply logging and access log middleware.
|
||||
let app = log::apply(app, &server_args);
|
||||
|
||||
// The listener was already bound synchronously in `runtime::start` (so a port
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
pub(crate) mod bootstrap;
|
||||
+165
-165
@@ -1,48 +1,37 @@
|
||||
//! PD KV bootstrap registry — the rust port of Python
|
||||
//! `CommonKVBootstrapServer` (disaggregation/common/conn.py), which every
|
||||
//! transfer backend (mooncake / mori / nixl / ascend) subclasses without
|
||||
//! overrides, so this one implementation covers them all.
|
||||
//!
|
||||
//! 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.
|
||||
//! PD KV bootstrap registry — rust port of Python `CommonKVBootstrapServer` (shared by all
|
||||
//! transfer backends): prefill ranks PUT `/route`, decode ranks GET routes and the `-1`-sentinel
|
||||
//! topology, the PD router tracks per-room dp ranks; the wire format is Python-owned parity.
|
||||
//! Mounted on the prefill api listener (bootstrap port = api port) before `init_disaggregation`.
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::sync::{Arc, Mutex};
|
||||
use std::time::{Duration, Instant};
|
||||
|
||||
use arc_swap::ArcSwap;
|
||||
use axum::extract::{Query, State};
|
||||
use axum::http::StatusCode;
|
||||
use axum::response::{IntoResponse, Response};
|
||||
use axum::routing::{post, put};
|
||||
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_DEFAULT_SECS: u64 = 120;
|
||||
const ROOM_SHARD_COUNT: usize = 64;
|
||||
|
||||
/// One registered prefill rank, exactly the JSON the Python decode side
|
||||
/// consumes (`PrefillRankInfo`).
|
||||
#[derive(Clone, serde::Serialize)]
|
||||
struct RankInfo {
|
||||
/// Python's (`PrefillRankInfo`).
|
||||
#[derive(Clone, Serialize)]
|
||||
struct PrefillRankInfo {
|
||||
rank_ip: String,
|
||||
rank_port: i64,
|
||||
}
|
||||
|
||||
/// The all-`-1` sentinel response, exactly Python's
|
||||
/// `dataclasses.asdict(PrefillServerInfo)` shape.
|
||||
#[derive(serde::Serialize)]
|
||||
/// Python's (`PrefillServerInfo`).
|
||||
#[derive(Serialize)]
|
||||
struct PrefillServerInfo {
|
||||
attn_tp_size: i64,
|
||||
attn_cp_size: i64,
|
||||
@@ -60,12 +49,56 @@ struct RoomEntry {
|
||||
registered_at: Instant,
|
||||
}
|
||||
|
||||
/// Mirror of the Python server's mutable state. First PUT wins for the
|
||||
/// 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.
|
||||
/// Mirror of the Python server's mutable state, split by write pattern.
|
||||
#[derive(Default)]
|
||||
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_cp_size: Option<i64>,
|
||||
dp_size: Option<i64>,
|
||||
@@ -77,12 +110,11 @@ struct Registry {
|
||||
prefill_http_port: Option<i64>,
|
||||
/// Keyed `(dp_group, attn_cp_rank, attn_tp_rank, pp_rank)` — the flat form
|
||||
/// of Python's nested `prefill_port_table` dicts.
|
||||
prefill_ranks: HashMap<(i64, i64, i64, i64), RankInfo>,
|
||||
room_to_dp_rank: HashMap<i64, RoomEntry>,
|
||||
prefill_ranks: HashMap<(i64, i64, i64, i64), PrefillRankInfo>,
|
||||
registered_count: i64,
|
||||
}
|
||||
|
||||
impl Registry {
|
||||
impl Topology {
|
||||
/// `dp * cp * tp * pp` once every size is known (saturating: absurd sizes
|
||||
/// stay "never ready" instead of overflowing).
|
||||
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`).
|
||||
#[derive(serde::Deserialize)]
|
||||
struct RoutePut {
|
||||
#[derive(Deserialize)]
|
||||
struct Route {
|
||||
attn_tp_size: i64,
|
||||
attn_tp_rank: i64,
|
||||
attn_cp_size: i64,
|
||||
@@ -141,27 +146,21 @@ struct RoutePut {
|
||||
system_dp_size: i64,
|
||||
system_dp_rank: i64,
|
||||
rank_ip: String,
|
||||
rank_port: Int,
|
||||
page_size: Int,
|
||||
#[serde(deserialize_with = "parse_int")]
|
||||
rank_port: i64,
|
||||
#[serde(deserialize_with = "parse_int")]
|
||||
page_size: i64,
|
||||
#[serde(default)]
|
||||
kv_cache_dtype: Option<String>,
|
||||
#[serde(default)]
|
||||
prefill_http_port: Option<Int>,
|
||||
#[serde(default, deserialize_with = "parse_int_opt")]
|
||||
prefill_http_port: Option<i64>,
|
||||
#[serde(default)]
|
||||
load_balance_method: Option<String>,
|
||||
#[serde(default)]
|
||||
enable_dsa_cache_layer_split: Option<bool>,
|
||||
}
|
||||
|
||||
fn not_ready(registered_count: i64) -> 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 {
|
||||
async fn route_put(State(state): State<Arc<Registry>>, Json(body): Json<Route>) -> Response {
|
||||
// `system_dp_size == 1` → attention-dp topology; else system-dp topology.
|
||||
let dp_size = if body.system_dp_size == 1 {
|
||||
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
|
||||
};
|
||||
|
||||
let mut reg = state.lock().unwrap();
|
||||
reg.attn_tp_size.get_or_insert(body.attn_tp_size);
|
||||
reg.attn_cp_size.get_or_insert(body.attn_cp_size);
|
||||
reg.dp_size.get_or_insert(dp_size);
|
||||
reg.pp_size.get_or_insert(body.pp_size);
|
||||
reg.page_size.get_or_insert(body.page_size.0);
|
||||
if reg.kv_cache_dtype.is_none() {
|
||||
reg.kv_cache_dtype = body.kv_cache_dtype;
|
||||
}
|
||||
if reg.prefill_http_port.is_none() {
|
||||
reg.prefill_http_port = body.prefill_http_port.map(|p| p.0);
|
||||
}
|
||||
reg.follow_bootstrap_room.get_or_insert(
|
||||
body.load_balance_method
|
||||
.as_deref()
|
||||
.unwrap_or("follow_bootstrap_room")
|
||||
== "follow_bootstrap_room",
|
||||
);
|
||||
reg.enable_dsa_cache_layer_split
|
||||
.get_or_insert(body.enable_dsa_cache_layer_split.unwrap_or(false));
|
||||
|
||||
reg.prefill_ranks.insert(
|
||||
(dp_group, body.attn_cp_rank, body.attn_tp_rank, body.pp_rank),
|
||||
RankInfo {
|
||||
rank_ip: body.rank_ip.clone(),
|
||||
rank_port: body.rank_port.0,
|
||||
},
|
||||
);
|
||||
reg.registered_count += 1;
|
||||
// Copy-on-write update. `rcu` may re-run the closure under write
|
||||
// contention, so it only reads `body` and clones what it stores.
|
||||
state.topology.rcu(|current| {
|
||||
let mut topo = (**current).clone();
|
||||
topo.attn_tp_size.get_or_insert(body.attn_tp_size);
|
||||
topo.attn_cp_size.get_or_insert(body.attn_cp_size);
|
||||
topo.dp_size.get_or_insert(dp_size);
|
||||
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 topo.prefill_http_port.is_none() {
|
||||
topo.prefill_http_port = body.prefill_http_port;
|
||||
}
|
||||
topo.follow_bootstrap_room.get_or_insert(
|
||||
body.load_balance_method
|
||||
.as_deref()
|
||||
.unwrap_or("follow_bootstrap_room")
|
||||
== "follow_bootstrap_room",
|
||||
);
|
||||
topo.enable_dsa_cache_layer_split
|
||||
.get_or_insert(body.enable_dsa_cache_layer_split.unwrap_or(false));
|
||||
topo.prefill_ranks.insert(
|
||||
(dp_group, body.attn_cp_rank, body.attn_tp_rank, body.pp_rank),
|
||||
PrefillRankInfo {
|
||||
rank_ip: body.rank_ip.clone(),
|
||||
rank_port: body.rank_port,
|
||||
},
|
||||
);
|
||||
topo.registered_count += 1;
|
||||
topo
|
||||
});
|
||||
|
||||
let topo = state.topology.load();
|
||||
tracing::debug!(
|
||||
dp_group,
|
||||
cp = body.attn_cp_rank,
|
||||
tp = body.attn_tp_rank,
|
||||
pp = body.pp_rank,
|
||||
rank_ip = %body.rank_ip,
|
||||
rank_port = body.rank_port.0,
|
||||
registered = reg.registered_count,
|
||||
expected = reg.expected(),
|
||||
rank_port = body.rank_port,
|
||||
registered = topo.registered_count,
|
||||
expected = topo.expected(),
|
||||
"registered prefill bootstrap rank"
|
||||
);
|
||||
"OK".into_response()
|
||||
}
|
||||
|
||||
async fn route_get(
|
||||
State(state): State<Shared>,
|
||||
State(state): State<Arc<Registry>>,
|
||||
Query(query): Query<HashMap<String, String>>,
|
||||
) -> Response {
|
||||
// A missing, empty (Python truthiness), or non-integer param → 400.
|
||||
@@ -230,37 +234,41 @@ async fn route_get(
|
||||
rank("target_tp_rank"),
|
||||
rank("target_pp_rank"),
|
||||
) else {
|
||||
return (
|
||||
return json_error(
|
||||
StatusCode::BAD_REQUEST,
|
||||
"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.
|
||||
if !reg.is_ready() {
|
||||
return not_ready(reg.registered_count);
|
||||
if !topo.is_ready() {
|
||||
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) {
|
||||
// Aggregate-topology sentinel. The sizes are Some — `is_ready` above
|
||||
// requires all four.
|
||||
// Aggregate-topology sentinel.
|
||||
return Json(PrefillServerInfo {
|
||||
attn_tp_size: reg.attn_tp_size.unwrap(),
|
||||
attn_cp_size: reg.attn_cp_size.unwrap(),
|
||||
dp_size: reg.dp_size.unwrap(),
|
||||
pp_size: reg.pp_size.unwrap(),
|
||||
page_size: reg.page_size,
|
||||
kv_cache_dtype: reg.kv_cache_dtype.clone(),
|
||||
follow_bootstrap_room: reg.follow_bootstrap_room.unwrap_or(true),
|
||||
enable_dsa_cache_layer_split: reg.enable_dsa_cache_layer_split.unwrap_or(false),
|
||||
prefill_http_port: reg.prefill_http_port,
|
||||
attn_tp_size: topo.attn_tp_size.unwrap(),
|
||||
attn_cp_size: topo.attn_cp_size.unwrap(),
|
||||
dp_size: topo.dp_size.unwrap(),
|
||||
pp_size: topo.pp_size.unwrap(),
|
||||
page_size: topo.page_size,
|
||||
kv_cache_dtype: topo.kv_cache_dtype.clone(),
|
||||
follow_bootstrap_room: topo.follow_bootstrap_room.unwrap_or(true),
|
||||
enable_dsa_cache_layer_split: topo.enable_dsa_cache_layer_split.unwrap_or(false),
|
||||
prefill_http_port: topo.prefill_http_port,
|
||||
})
|
||||
.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(),
|
||||
None => (
|
||||
StatusCode::NOT_FOUND,
|
||||
@@ -273,49 +281,51 @@ async fn route_get(
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(serde::Deserialize)]
|
||||
#[derive(Deserialize)]
|
||||
struct RegisterDpRank {
|
||||
bootstrap_room: Int,
|
||||
dp_rank: Int,
|
||||
#[serde(deserialize_with = "parse_int")]
|
||||
bootstrap_room: i64,
|
||||
#[serde(deserialize_with = "parse_int")]
|
||||
dp_rank: i64,
|
||||
}
|
||||
|
||||
async fn register_dp_rank(
|
||||
State(state): State<Shared>,
|
||||
State(state): State<Arc<Registry>>,
|
||||
Json(body): Json<RegisterDpRank>,
|
||||
) -> Response {
|
||||
state.lock().unwrap().room_to_dp_rank.insert(
|
||||
body.bootstrap_room.0,
|
||||
state.rooms.insert(
|
||||
body.bootstrap_room,
|
||||
RoomEntry {
|
||||
dp_rank: body.dp_rank.0,
|
||||
dp_rank: body.dp_rank,
|
||||
registered_at: Instant::now(),
|
||||
},
|
||||
);
|
||||
"OK".into_response()
|
||||
}
|
||||
|
||||
#[derive(serde::Deserialize)]
|
||||
#[derive(Deserialize)]
|
||||
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
|
||||
/// 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 {
|
||||
let reg = state.lock().unwrap();
|
||||
async fn query_dp_ranks(
|
||||
State(state): State<Arc<Registry>>,
|
||||
Json(body): Json<QueryDpRanks>,
|
||||
) -> Response {
|
||||
let result: HashMap<String, i64> = body
|
||||
.bootstrap_rooms
|
||||
.iter()
|
||||
.filter_map(|room| {
|
||||
let entry = reg.room_to_dp_rank.get(&room.0)?;
|
||||
Some((room.0.to_string(), entry.dp_rank))
|
||||
})
|
||||
.filter_map(|room| Some((room.to_string(), state.rooms.dp_rank(*room)?)))
|
||||
.collect();
|
||||
Json(result).into_response()
|
||||
}
|
||||
|
||||
/// No `/health` here: the merged api router already serves it (same 200 "OK"
|
||||
/// the standalone Python bootstrap server answered, so probes are unchanged).
|
||||
fn router(state: Shared) -> Router {
|
||||
fn router(state: Arc<Registry>) -> Router {
|
||||
Router::new()
|
||||
// Unmatched methods on a routed path get axum's built-in 405, matching
|
||||
// Python's explicit method_not_allowed branch.
|
||||
@@ -325,32 +335,22 @@ fn router(state: Shared) -> Router {
|
||||
.with_state(state)
|
||||
}
|
||||
|
||||
/// Drop `room_to_dp_rank` entries older than `interval` — Python's
|
||||
/// `_cleanup_expired_entries` loop (interval doubles as both period and TTL).
|
||||
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();
|
||||
/// Drop room entries
|
||||
async fn cleanup_sweeper(state: Arc<Registry>) {
|
||||
let cleanup_interval = Duration::from_secs(crate::environ::env_u64(
|
||||
ENTRY_CLEANUP_INTERVAL_ENV,
|
||||
ENTRY_CLEANUP_INTERVAL_DEFAULT_SECS,
|
||||
));
|
||||
(
|
||||
router(state.clone()),
|
||||
cleanup_expired_entries(state, cleanup_interval),
|
||||
)
|
||||
loop {
|
||||
tokio::time::sleep(cleanup_interval).await;
|
||||
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)]
|
||||
@@ -2,4 +2,5 @@
|
||||
|
||||
pub mod regex;
|
||||
pub mod response;
|
||||
pub mod serialize;
|
||||
pub mod sock;
|
||||
|
||||
@@ -2,8 +2,8 @@
|
||||
//!
|
||||
//! Two mechanics live here; the WIRE SHAPES stay owned by their endpoints:
|
||||
//! the native `{"error": {"message", "code"}}` body (Python
|
||||
//! `http_server.generate_request` parity) built by [`error_value`], and the
|
||||
//! SSE variant [`sse_error_response`] used by any
|
||||
//! `http_server.generate_request` parity) built by [`error_value`] and formed
|
||||
//! by [`json_error`], and the SSE variant [`sse_error_response`] used by any
|
||||
//! endpoint family (native and OpenAI alike — the caller supplies its own
|
||||
//! body shape). The OpenAI error payload and the PD bootstrap registry's
|
||||
//! 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 } })
|
||||
}
|
||||
|
||||
/// 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
|
||||
/// the JSON `body`; streaming → 200 with one SSE error frame + `[DONE]` (the
|
||||
/// 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