[model-gateway] Parallelize metrics requests (#14953)

This commit is contained in:
Praneth Paruchuri
2025-12-14 15:38:45 -08:00
committed by GitHub
parent 0e4108ba29
commit bab20a849e
+46 -32
View File
@@ -6,7 +6,10 @@
use std::{collections::HashMap, sync::Arc, time::Duration}; use std::{collections::HashMap, sync::Arc, time::Duration};
use axum::response::{IntoResponse, Response}; use axum::response::{IntoResponse, Response};
use futures::future; use futures::{
future,
stream::{self, StreamExt},
};
use http::{Method, StatusCode}; use http::{Method, StatusCode};
use serde_json::Value; use serde_json::Value;
use tokio::{ use tokio::{
@@ -21,6 +24,9 @@ use crate::{
protocols::worker_spec::{FlushCacheResult, WorkerLoadInfo, WorkerLoadsResult}, protocols::worker_spec::{FlushCacheResult, WorkerLoadInfo, WorkerLoadsResult},
}; };
const REQUEST_TIMEOUT: Duration = Duration::from_secs(5);
const MAX_CONCURRENT: usize = 32;
/// Unified worker management /// Unified worker management
pub struct WorkerManager; pub struct WorkerManager;
@@ -66,7 +72,7 @@ impl WorkerManager {
workers.len() workers.len()
); );
let mut tasks = Vec::new(); let mut tasks = Vec::with_capacity(http_workers.len());
for worker in &http_workers { for worker in &http_workers {
let url = worker.url().to_string(); let url = worker.url().to_string();
let flush_url = format!("{}/flush_cache", url); let flush_url = format!("{}/flush_cache", url);
@@ -131,6 +137,7 @@ impl WorkerManager {
message, message,
}) })
} }
pub async fn get_worker_load( pub async fn get_worker_load(
url: &str, url: &str,
api_key: Option<&str>, api_key: Option<&str>,
@@ -195,7 +202,7 @@ impl WorkerManager {
let total_workers = workers.len(); let total_workers = workers.len();
// Prepare tasks for parallel execution // Prepare tasks for parallel execution
let mut tasks = Vec::new(); let mut tasks = Vec::with_capacity(workers.len());
for worker in &workers { for worker in &workers {
let url = worker.url().to_string(); let url = worker.url().to_string();
let api_key = worker.api_key().clone(); let api_key = worker.api_key().clone();
@@ -272,50 +279,57 @@ impl WorkerManager {
method: Method, method: Method,
) -> Result<Vec<(String, String)>, Response> { ) -> Result<Vec<(String, String)>, Response> {
let workers = worker_registry.get_all(); let workers = worker_registry.get_all();
let worker_count = workers.len();
if workers.is_empty() { if workers.is_empty() {
return Err((StatusCode::SERVICE_UNAVAILABLE, "No available workers").into_response()); return Err((StatusCode::SERVICE_UNAVAILABLE, "No available workers").into_response());
} }
let mut responses = vec![]; let futures: Vec<_> = workers
// May do parallel requests later .into_iter()
for worker in workers { .map(|worker| {
let client = client.clone();
let worker_url = worker.url().to_string(); let worker_url = worker.url().to_string();
let url = format!("{worker_url}/{endpoint}");
let api_key = worker.api_key().clone();
let method = method.clone();
let url = format!("{}/{}", worker_url, endpoint); async move {
let mut request_builder = match method { let mut req = client.request(method, &url).timeout(REQUEST_TIMEOUT);
Method::GET => client.get(url),
Method::POST => client.post(url),
_ => {
return Err((
StatusCode::METHOD_NOT_ALLOWED,
"Unsupported method for simple routing",
)
.into_response())
}
};
if let Some(api_key) = worker.api_key() { if let Some(key) = api_key {
request_builder = req = req.bearer_auth(key);
request_builder.header("Authorization", format!("Bearer {}", api_key));
} }
match request_builder.send().await { match req.send().await {
Ok(res) if res.status().is_success() => match res.text().await {
Ok(body) => Some((worker_url, body)),
Err(e) => {
warn!("Failed reading response from {url}: {e}");
None
}
},
Ok(res) => { Ok(res) => {
let status = StatusCode::from_u16(res.status().as_u16()) warn!("Request to {url} failed: {}", res.status());
.unwrap_or(StatusCode::INTERNAL_SERVER_ERROR); None
match res.text().await {
Ok(body_text) => {
if status.is_success() {
responses.push((worker_url, body_text));
}
} }
Err(e) => { Err(e) => {
warn!("fan_out_simple_request failed when reading text: {}", e) warn!("Request to {url} failed: {e}");
None
} }
} }
} }
Err(e) => warn!("fan_out_simple_request failed when sending: {}", e), })
} .collect();
let responses: Vec<_> = stream::iter(futures)
.buffer_unordered(MAX_CONCURRENT)
.filter_map(|r| async { r })
.collect()
.await;
if responses.is_empty() && worker_count > 0 {
return Err((StatusCode::BAD_GATEWAY, "All backend requests failed").into_response());
} }
Ok(responses) Ok(responses)