[model-gateway] Optimize HTTP Router Fan-out: Replace Serial Execution with Concurrent Streams (#16042)

This commit is contained in:
Praneth Paruchuri
2026-01-04 23:38:45 -08:00
committed by GitHub
parent 012dc5866d
commit b12258bfaa
2 changed files with 48 additions and 30 deletions
-1
View File
@@ -149,7 +149,6 @@ tonic-v12 = { version = "0.12.3", package = "tonic" }
serial_test = "3.0" serial_test = "3.0"
rsa = { version = "0.9", features = ["sha2"] } rsa = { version = "0.9", features = ["sha2"] }
[[bench]] [[bench]]
name = "request_processing" name = "request_processing"
harness = false harness = false
+48 -29
View File
@@ -7,7 +7,7 @@ use axum::{
response::{IntoResponse, Response}, response::{IntoResponse, Response},
Json, Json,
}; };
use futures_util::StreamExt; use futures_util::{stream, StreamExt};
use reqwest::Client; use reqwest::Client;
use tokio_stream::wrappers::UnboundedReceiverStream; use tokio_stream::wrappers::UnboundedReceiverStream;
use tracing::{debug, error}; use tracing::{debug, error};
@@ -362,46 +362,65 @@ impl Router {
}) })
.unwrap_or_default(); .unwrap_or_default();
let mut last_response: Option<Response> = None; let futures: Vec<_> = workers
for worker in workers { .into_iter()
let worker_url = worker.url(); .map(|worker| {
let base = self.worker_base_url(worker_url); let worker_url = worker.url();
let base = self.worker_base_url(worker_url);
let url = format!("{}/{}", base, endpoint);
let client = self.client.clone();
let method = method.clone();
let url = format!("{}/{}", base, endpoint); let headers = filtered_headers.clone();
let mut request_builder = match method {
Method::GET => self.client.get(url), let api_key = worker.api_key().clone();
Method::POST => self.client.post(url),
_ => { async move {
return error::method_not_allowed( let mut request_builder = match method {
"unsupported_method", Method::GET => client.get(url),
"Unsupported method for simple routing", Method::POST => client.post(url),
) _ => {
return Err(error::method_not_allowed(
"unsupported_method",
"Unsupported method for simple routing",
))
}
};
if let Some(key) = api_key {
let mut auth_header = String::with_capacity(7 + key.len());
auth_header.push_str("Bearer ");
auth_header.push_str(&key);
request_builder = request_builder.header("Authorization", auth_header);
}
for (name, value) in headers {
request_builder = request_builder.header(name.clone(), value.clone());
}
request_builder.send().await.map_err(convert_reqwest_error)
} }
}; })
.collect();
if let Some(api_key) = worker.api_key() { // Now execute the collected futures concurrently
// Pre-allocate string with capacity to avoid reallocation let mut stream = stream::iter(futures).buffer_unordered(32);
let mut auth_header = String::with_capacity(7 + api_key.len()); let mut last_response: Option<Response> = None;
auth_header.push_str("Bearer ");
auth_header.push_str(api_key);
request_builder = request_builder.header("Authorization", auth_header);
}
// Apply pre-filtered headers while let Some(result) = stream.next().await {
for (name, value) in &filtered_headers { match result {
request_builder = request_builder.header(*name, *value);
}
match request_builder.send().await {
Ok(res) => { Ok(res) => {
let status = StatusCode::from_u16(res.status().as_u16()) let status = StatusCode::from_u16(res.status().as_u16())
.unwrap_or(StatusCode::INTERNAL_SERVER_ERROR); .unwrap_or(StatusCode::INTERNAL_SERVER_ERROR);
let response_headers = header_utils::preserve_response_headers(res.headers()); let response_headers = header_utils::preserve_response_headers(res.headers());
match res.bytes().await { match res.bytes().await {
Ok(body) => { Ok(body) => {
let mut response = Response::new(Body::from(body)); let mut response = Response::new(Body::from(body));
*response.status_mut() = status; *response.status_mut() = status;
*response.headers_mut() = response_headers; *response.headers_mut() = response_headers;
if status.is_success() { if status.is_success() {
return response; return response;
} }
@@ -416,7 +435,7 @@ impl Router {
} }
} }
Err(e) => { Err(e) => {
last_response = Some(convert_reqwest_error(e)); last_response = Some(e);
} }
} }
} }