Support DP-aware PD router dispatch (#26245)
Co-authored-by: weireweire <20922698+weireweire@users.noreply.github.com>
This commit is contained in:
@@ -39,6 +39,36 @@ static WORKER_CLIENT: LazyLock<reqwest::Client> = LazyLock::new(|| {
|
||||
.expect("Failed to create worker HTTP client")
|
||||
});
|
||||
|
||||
pub(crate) fn parse_bootstrap_host_from_url(url: &str) -> String {
|
||||
let metadata_url = match url.rsplit_once('@') {
|
||||
Some((base_url, rank)) if rank.parse::<usize>().is_ok() => base_url,
|
||||
_ => url,
|
||||
};
|
||||
|
||||
match url::Url::parse(metadata_url) {
|
||||
Ok(parsed) => parsed.host_str().unwrap_or("localhost").to_string(),
|
||||
Err(_) if !metadata_url.contains("://") => {
|
||||
match url::Url::parse(&format!("http://{}", metadata_url)) {
|
||||
Ok(parsed) => parsed.host_str().unwrap_or("localhost").to_string(),
|
||||
Err(_) => {
|
||||
tracing::warn!(
|
||||
"Failed to parse URL '{}', defaulting to localhost",
|
||||
metadata_url
|
||||
);
|
||||
"localhost".to_string()
|
||||
}
|
||||
}
|
||||
}
|
||||
Err(_) => {
|
||||
tracing::warn!(
|
||||
"Failed to parse URL '{}', defaulting to localhost",
|
||||
metadata_url
|
||||
);
|
||||
"localhost".to_string()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub struct WorkerRoutingKeyLoad {
|
||||
url: String,
|
||||
active_routing_keys: dashmap::DashMap<String, usize>,
|
||||
@@ -963,6 +993,8 @@ pub struct DPAwareWorker {
|
||||
dp_size: usize,
|
||||
/// Base URL without DP suffix
|
||||
base_url: String,
|
||||
/// Bootstrap host parsed from the real base URL, not the virtual DP URL.
|
||||
bootstrap_host: String,
|
||||
}
|
||||
|
||||
impl DPAwareWorker {
|
||||
@@ -974,11 +1006,13 @@ impl DPAwareWorker {
|
||||
dp_rank: usize,
|
||||
dp_size: usize,
|
||||
) -> Self {
|
||||
let bootstrap_host = parse_bootstrap_host_from_url(&base_url);
|
||||
Self {
|
||||
base_worker,
|
||||
dp_rank,
|
||||
dp_size,
|
||||
base_url,
|
||||
bootstrap_host,
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -1001,6 +1035,10 @@ impl Worker for DPAwareWorker {
|
||||
self.base_worker.connection_mode()
|
||||
}
|
||||
|
||||
fn bootstrap_host(&self) -> &str {
|
||||
&self.bootstrap_host
|
||||
}
|
||||
|
||||
fn is_healthy(&self) -> bool {
|
||||
self.base_worker.is_healthy()
|
||||
}
|
||||
@@ -1274,6 +1312,22 @@ mod tests {
|
||||
DPAwareWorkerBuilder,
|
||||
};
|
||||
|
||||
#[test]
|
||||
fn test_parse_bootstrap_host_strips_dp_rank_suffix() {
|
||||
assert_eq!(
|
||||
parse_bootstrap_host_from_url("http://10.66.5.115:20664@3"),
|
||||
"10.66.5.115"
|
||||
);
|
||||
assert_eq!(
|
||||
parse_bootstrap_host_from_url("grpc://cluster.local@1"),
|
||||
"cluster.local"
|
||||
);
|
||||
assert_eq!(
|
||||
parse_bootstrap_host_from_url("localhost:8080@2"),
|
||||
"localhost"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_worker_type_display() {
|
||||
assert_eq!(WorkerType::Regular.to_string(), "Regular");
|
||||
@@ -1679,6 +1733,7 @@ mod tests {
|
||||
|
||||
assert_eq!(dp_worker.url(), "http://worker1:8080@2");
|
||||
assert_eq!(dp_worker.base_url(), "http://worker1:8080");
|
||||
assert_eq!(dp_worker.bootstrap_host(), "worker1");
|
||||
assert!(dp_worker.is_dp_aware());
|
||||
assert_eq!(dp_worker.dp_rank(), Some(2));
|
||||
assert_eq!(dp_worker.dp_size(), Some(4));
|
||||
@@ -1694,6 +1749,8 @@ mod tests {
|
||||
.build();
|
||||
|
||||
assert_eq!(dp_worker.url(), "http://worker1:8080@1");
|
||||
assert_eq!(dp_worker.bootstrap_host(), "worker1");
|
||||
assert_eq!(dp_worker.bootstrap_port(), Some(9090));
|
||||
assert!(dp_worker.is_dp_aware());
|
||||
assert_eq!(
|
||||
dp_worker.worker_type(),
|
||||
@@ -1714,6 +1771,23 @@ mod tests {
|
||||
assert_eq!(dp_worker.worker_type(), &WorkerType::Decode);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_dp_aware_worker_bootstrap_host_uses_base_url() {
|
||||
let dp_worker = DPAwareWorkerBuilder::new("http://10.66.5.240:21686", 1, 4)
|
||||
.worker_type(WorkerType::Prefill {
|
||||
bootstrap_port: None,
|
||||
})
|
||||
.build();
|
||||
|
||||
assert_eq!(dp_worker.url(), "http://10.66.5.240:21686@1");
|
||||
assert_eq!(
|
||||
dp_worker.endpoint_url("/generate"),
|
||||
"http://10.66.5.240:21686/generate"
|
||||
);
|
||||
assert_eq!(dp_worker.bootstrap_host(), "10.66.5.240");
|
||||
assert_eq!(dp_worker.bootstrap_port(), None);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_dp_aware_prepare_request() {
|
||||
let dp_worker = DPAwareWorkerBuilder::new("http://worker1:8080", 3, 8)
|
||||
|
||||
@@ -5,8 +5,8 @@ use super::{
|
||||
model_card::ModelCard,
|
||||
model_type::ModelType,
|
||||
worker::{
|
||||
BasicWorker, ConnectionMode, DPAwareWorker, HealthConfig, RuntimeType, WorkerMetadata,
|
||||
WorkerRoutingKeyLoad, WorkerType,
|
||||
parse_bootstrap_host_from_url, BasicWorker, ConnectionMode, DPAwareWorker, HealthConfig,
|
||||
RuntimeType, WorkerMetadata, WorkerRoutingKeyLoad, WorkerType,
|
||||
},
|
||||
};
|
||||
use crate::{observability::metrics::Metrics, routers::grpc::client::GrpcClient};
|
||||
@@ -133,28 +133,7 @@ impl BasicWorkerBuilder {
|
||||
|
||||
use tokio::sync::OnceCell;
|
||||
|
||||
let bootstrap_host = match url::Url::parse(&self.url) {
|
||||
Ok(parsed) => parsed.host_str().unwrap_or("localhost").to_string(),
|
||||
Err(_) if !self.url.contains("://") => {
|
||||
match url::Url::parse(&format!("http://{}", self.url)) {
|
||||
Ok(parsed) => parsed.host_str().unwrap_or("localhost").to_string(),
|
||||
Err(_) => {
|
||||
tracing::warn!(
|
||||
"Failed to parse URL '{}', defaulting to localhost",
|
||||
self.url
|
||||
);
|
||||
"localhost".to_string()
|
||||
}
|
||||
}
|
||||
}
|
||||
Err(_) => {
|
||||
tracing::warn!(
|
||||
"Failed to parse URL '{}', defaulting to localhost",
|
||||
self.url
|
||||
);
|
||||
"localhost".to_string()
|
||||
}
|
||||
};
|
||||
let bootstrap_host = parse_bootstrap_host_from_url(&self.url);
|
||||
|
||||
let bootstrap_port = match self.worker_type {
|
||||
WorkerType::Prefill { bootstrap_port } => bootstrap_port,
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
use std::{sync::Arc, time::Instant};
|
||||
use std::{borrow::Cow, sync::Arc, time::Instant};
|
||||
|
||||
use async_trait::async_trait;
|
||||
use axum::{
|
||||
@@ -56,6 +56,11 @@ pub struct PDRouter {
|
||||
pub enable_igw: bool,
|
||||
}
|
||||
|
||||
struct PreparedWorkerRequest<'a> {
|
||||
endpoint_url: String,
|
||||
body: Cow<'a, Value>,
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
struct PDRequestContext<'a> {
|
||||
route: &'static str,
|
||||
@@ -77,16 +82,20 @@ struct PDRequestContext<'a> {
|
||||
struct BreakerOutcomesRecorded;
|
||||
|
||||
impl PDRouter {
|
||||
fn worker_endpoint_url(worker: &dyn Worker, endpoint: &str) -> String {
|
||||
api_path(worker.base_url(), endpoint)
|
||||
}
|
||||
|
||||
async fn proxy_to_first_prefill_worker(
|
||||
&self,
|
||||
endpoint: &str,
|
||||
headers: Option<Vec<(String, String)>>,
|
||||
) -> Response {
|
||||
let workers = self.worker_registry.get_prefill_workers();
|
||||
let first_worker_url = workers.first().map(|w| w.url().to_string());
|
||||
|
||||
if let Some(worker_url) = first_worker_url {
|
||||
self.proxy_to_worker(worker_url, endpoint, headers).await
|
||||
if let Some(worker) = workers.first() {
|
||||
self.proxy_to_worker(worker.as_ref(), endpoint, headers)
|
||||
.await
|
||||
} else {
|
||||
error::service_unavailable("no_prefill_servers", "No prefill servers available")
|
||||
}
|
||||
@@ -94,11 +103,11 @@ impl PDRouter {
|
||||
|
||||
async fn proxy_to_worker(
|
||||
&self,
|
||||
worker_url: String,
|
||||
worker: &dyn Worker,
|
||||
endpoint: &str,
|
||||
headers: Option<Vec<(String, String)>>,
|
||||
) -> Response {
|
||||
let url = format!("{}/{}", worker_url, endpoint);
|
||||
let url = Self::worker_endpoint_url(worker, endpoint);
|
||||
let mut request_builder = self.client.get(&url);
|
||||
|
||||
if let Some(headers) = headers {
|
||||
@@ -224,6 +233,7 @@ impl PDRouter {
|
||||
const BOOTSTRAP_HOST_KEY: &'static str = "bootstrap_host";
|
||||
const BOOTSTRAP_PORT_KEY: &'static str = "bootstrap_port";
|
||||
const BOOTSTRAP_ROOM_KEY: &'static str = "bootstrap_room";
|
||||
const DISAGG_PREFILL_DP_RANK_KEY: &'static str = "disagg_prefill_dp_rank";
|
||||
|
||||
fn inject_bootstrap_into_value(
|
||||
mut original: Value,
|
||||
@@ -285,6 +295,73 @@ impl PDRouter {
|
||||
Ok(original)
|
||||
}
|
||||
|
||||
fn inject_prefill_dp_rank_for_decode<'a>(
|
||||
decode_request: Cow<'a, Value>,
|
||||
prefill_worker: &dyn Worker,
|
||||
) -> Result<Cow<'a, Value>, String> {
|
||||
let Some(prefill_dp_rank) = prefill_worker.dp_rank() else {
|
||||
return Ok(decode_request);
|
||||
};
|
||||
|
||||
let mut decode_request = decode_request.into_owned();
|
||||
let Some(obj) = decode_request.as_object_mut() else {
|
||||
return Err(
|
||||
"Failed to insert disagg_prefill_dp_rank because request body is not an object"
|
||||
.to_string(),
|
||||
);
|
||||
};
|
||||
|
||||
obj.insert(
|
||||
Self::DISAGG_PREFILL_DP_RANK_KEY.to_string(),
|
||||
Value::from(prefill_dp_rank as u64),
|
||||
);
|
||||
Ok(Cow::Owned(decode_request))
|
||||
}
|
||||
|
||||
async fn prepare_worker_request<'a>(
|
||||
route: &'static str,
|
||||
worker: &dyn Worker,
|
||||
json_request: Cow<'a, Value>,
|
||||
) -> Result<PreparedWorkerRequest<'a>, String> {
|
||||
let body = if worker.is_dp_aware() {
|
||||
Cow::Owned(
|
||||
worker
|
||||
.prepare_request(json_request.into_owned())
|
||||
.await
|
||||
.map_err(|err| {
|
||||
format!(
|
||||
"Failed to prepare request for worker {}: {}",
|
||||
worker.url(),
|
||||
err
|
||||
)
|
||||
})?,
|
||||
)
|
||||
} else {
|
||||
json_request
|
||||
};
|
||||
|
||||
Ok(PreparedWorkerRequest {
|
||||
endpoint_url: Self::worker_endpoint_url(worker, route),
|
||||
body,
|
||||
})
|
||||
}
|
||||
|
||||
async fn prepare_pd_worker_requests<'a>(
|
||||
route: &'static str,
|
||||
json_request: &'a Value,
|
||||
prefill: &dyn Worker,
|
||||
decode: &dyn Worker,
|
||||
) -> Result<(PreparedWorkerRequest<'a>, PreparedWorkerRequest<'a>), String> {
|
||||
let prefill_request =
|
||||
Self::prepare_worker_request(route, prefill, Cow::Borrowed(json_request)).await?;
|
||||
let decode_json_request =
|
||||
Self::inject_prefill_dp_rank_for_decode(Cow::Borrowed(json_request), prefill)?;
|
||||
let decode_request =
|
||||
Self::prepare_worker_request(route, decode, decode_json_request).await?;
|
||||
|
||||
Ok((prefill_request, decode_request))
|
||||
}
|
||||
|
||||
async fn execute_dual_dispatch<T: Serialize + Clone>(
|
||||
&self,
|
||||
headers: Option<&HeaderMap>,
|
||||
@@ -586,20 +663,33 @@ impl PDRouter {
|
||||
inject_trace_context_http(&mut headers_with_trace);
|
||||
let headers = Some(&headers_with_trace);
|
||||
|
||||
let (prepared_prefill, prepared_decode) = match Self::prepare_pd_worker_requests(
|
||||
context.route,
|
||||
&json_request,
|
||||
prefill.as_ref(),
|
||||
decode.as_ref(),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(requests) => requests,
|
||||
Err(e) => {
|
||||
error!("Failed to prepare PD worker requests: {}", e);
|
||||
return error::internal_error("pd_request_preparation_failed", e);
|
||||
}
|
||||
};
|
||||
|
||||
// Build both requests
|
||||
let prefill_request = self.build_post_with_headers(
|
||||
&self.client,
|
||||
prefill.url(),
|
||||
context.route,
|
||||
&json_request,
|
||||
&prepared_prefill.endpoint_url,
|
||||
&prepared_prefill.body,
|
||||
headers,
|
||||
false,
|
||||
);
|
||||
let decode_request = self.build_post_with_headers(
|
||||
&self.client,
|
||||
decode.url(),
|
||||
context.route,
|
||||
&json_request,
|
||||
&prepared_decode.endpoint_url,
|
||||
&prepared_decode.body,
|
||||
headers,
|
||||
false,
|
||||
);
|
||||
@@ -1201,13 +1291,12 @@ impl PDRouter {
|
||||
fn build_post_with_headers(
|
||||
&self,
|
||||
client: &Client,
|
||||
url: &str,
|
||||
route: &'static str,
|
||||
endpoint_url: &str,
|
||||
json_request: &Value,
|
||||
headers: Option<&HeaderMap>,
|
||||
connection_close: bool,
|
||||
) -> reqwest::RequestBuilder {
|
||||
let mut request = client.post(api_path(url, route)).json(json_request);
|
||||
let mut request = client.post(endpoint_url).json(json_request);
|
||||
if connection_close {
|
||||
request = request.header("Connection", "close");
|
||||
}
|
||||
@@ -1315,12 +1404,11 @@ impl RouterTrait for PDRouter {
|
||||
}
|
||||
};
|
||||
|
||||
let prefill_url = format!("{}/health_generate", prefill.url());
|
||||
let prefill_url = Self::worker_endpoint_url(prefill.as_ref(), "health_generate");
|
||||
let decode_url = Self::worker_endpoint_url(decode.as_ref(), "health_generate");
|
||||
let (prefill_result, decode_result) = tokio::join!(
|
||||
self.client.get(&prefill_url).send(),
|
||||
self.client
|
||||
.get(format!("{}/health_generate", decode.url()))
|
||||
.send()
|
||||
self.client.get(&decode_url).send()
|
||||
);
|
||||
|
||||
// Check results
|
||||
@@ -1561,7 +1649,7 @@ impl RouterTrait for PDRouter {
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::core::{BasicWorkerBuilder, WorkerType};
|
||||
use crate::core::{BasicWorkerBuilder, DPAwareWorkerBuilder, WorkerType};
|
||||
|
||||
fn create_test_pd_router() -> PDRouter {
|
||||
let worker_registry = Arc::new(WorkerRegistry::new());
|
||||
@@ -1679,6 +1767,101 @@ mod tests {
|
||||
assert!(result.unwrap_err().contains("No prefill workers available"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_worker_endpoint_url_uses_base_url_for_dp_aware_worker() {
|
||||
let worker = DPAwareWorkerBuilder::new("http://prefill:30000", 2, 4)
|
||||
.worker_type(WorkerType::Prefill {
|
||||
bootstrap_port: Some(8998),
|
||||
})
|
||||
.build();
|
||||
|
||||
assert_eq!(
|
||||
PDRouter::worker_endpoint_url(&worker, "health_generate"),
|
||||
"http://prefill:30000/health_generate"
|
||||
);
|
||||
assert_eq!(
|
||||
PDRouter::worker_endpoint_url(&worker, "/v1/models"),
|
||||
"http://prefill:30000/v1/models"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_prepare_pd_worker_requests_uses_dp_aware_rank() {
|
||||
let prefill = DPAwareWorkerBuilder::new("http://prefill:30000", 2, 4)
|
||||
.worker_type(WorkerType::Prefill {
|
||||
bootstrap_port: Some(8998),
|
||||
})
|
||||
.build();
|
||||
let decode = DPAwareWorkerBuilder::new("http://decode:30001", 1, 4)
|
||||
.worker_type(WorkerType::Decode)
|
||||
.build();
|
||||
let request = json!({
|
||||
"prompt": "shared prefix",
|
||||
"max_tokens": 8,
|
||||
"bootstrap_host": "prefill",
|
||||
"bootstrap_port": 8998,
|
||||
"bootstrap_room": 1234,
|
||||
});
|
||||
|
||||
let (prefill_request, decode_request) =
|
||||
PDRouter::prepare_pd_worker_requests("/v1/completions", &request, &prefill, &decode)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(
|
||||
prefill_request.endpoint_url,
|
||||
"http://prefill:30000/v1/completions"
|
||||
);
|
||||
assert_eq!(prefill_request.body["data_parallel_rank"], 2);
|
||||
assert!(prefill_request.body.get("disagg_prefill_dp_rank").is_none());
|
||||
|
||||
assert_eq!(
|
||||
decode_request.endpoint_url,
|
||||
"http://decode:30001/v1/completions"
|
||||
);
|
||||
assert_eq!(decode_request.body["data_parallel_rank"], 1);
|
||||
assert_eq!(decode_request.body["disagg_prefill_dp_rank"], 2);
|
||||
assert_eq!(decode_request.body["bootstrap_room"], 1234);
|
||||
assert!(matches!(prefill_request.body, Cow::Owned(_)));
|
||||
assert!(matches!(decode_request.body, Cow::Owned(_)));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_prepare_pd_worker_requests_preserves_non_dp_workers() {
|
||||
let prefill = BasicWorkerBuilder::new("http://prefill:30000")
|
||||
.worker_type(WorkerType::Prefill {
|
||||
bootstrap_port: Some(8998),
|
||||
})
|
||||
.build();
|
||||
let decode = BasicWorkerBuilder::new("http://decode:30001")
|
||||
.worker_type(WorkerType::Decode)
|
||||
.build();
|
||||
let request = json!({
|
||||
"prompt": "shared prefix",
|
||||
"max_tokens": 8,
|
||||
"bootstrap_room": 1234,
|
||||
});
|
||||
|
||||
let (prefill_request, decode_request) =
|
||||
PDRouter::prepare_pd_worker_requests("/v1/completions", &request, &prefill, &decode)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(
|
||||
prefill_request.endpoint_url,
|
||||
"http://prefill:30000/v1/completions"
|
||||
);
|
||||
assert_eq!(
|
||||
decode_request.endpoint_url,
|
||||
"http://decode:30001/v1/completions"
|
||||
);
|
||||
assert!(prefill_request.body.get("data_parallel_rank").is_none());
|
||||
assert!(decode_request.body.get("data_parallel_rank").is_none());
|
||||
assert!(decode_request.body.get("disagg_prefill_dp_rank").is_none());
|
||||
assert!(matches!(prefill_request.body, Cow::Borrowed(_)));
|
||||
assert!(matches!(decode_request.body, Cow::Borrowed(_)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_worker_load_metrics() {
|
||||
let prefill_worker: Arc<dyn Worker> = Arc::from(create_test_worker(
|
||||
|
||||
Reference in New Issue
Block a user