[gateway] Align /v1/loads and /model_info with sglang server; drop dead /rerank (#24167)

Co-authored-by: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
Kangyan-Zhou
2026-05-02 11:43:30 -07:00
committed by GitHub
co-authored by Claude Opus 4.7
parent e0474fdd9b
commit d74d9bda49
6 changed files with 10 additions and 275 deletions
@@ -1241,7 +1241,7 @@ impl RouterTrait for PDRouter {
let headers = header_utils::copy_request_headers(&req);
// Proxy to first prefill worker
self.proxy_to_first_prefill_worker("get_model_info", Some(headers))
self.proxy_to_first_prefill_worker("model_info", Some(headers))
.await
}
+1 -1
View File
@@ -732,7 +732,7 @@ impl RouterTrait for Router {
}
async fn get_model_info(&self, req: Request<Body>) -> Response {
self.proxy_get_request(req, "get_model_info").await
self.proxy_get_request(req, "model_info").await
}
async fn route_generate(
+5 -13
View File
@@ -46,7 +46,7 @@ use crate::{
embedding::EmbeddingRequest,
generate::GenerateRequest,
parser::{ParseFunctionCallRequest, SeparateReasoningRequest},
rerank::{RerankRequest, V1RerankReqInput},
rerank::V1RerankReqInput,
responses::{ResponsesGetParams, ResponsesRequest},
tokenize::{AddTokenizerRequest, DetokenizeRequest, TokenizeRequest},
validated::ValidatedJson,
@@ -203,17 +203,6 @@ async fn v1_completions(
.await
}
async fn rerank(
State(state): State<Arc<AppState>>,
headers: http::HeaderMap,
ValidatedJson(body): ValidatedJson<RerankRequest>,
) -> Response {
state
.router
.route_rerank(Some(&headers), &body, Some(&body.model))
.await
}
async fn v1_rerank(
State(state): State<Arc<AppState>>,
headers: http::HeaderMap,
@@ -556,7 +545,6 @@ pub fn build_app(
.route("/generate", post(generate))
.route("/v1/chat/completions", post(v1_chat_completions))
.route("/v1/completions", post(v1_completions))
.route("/rerank", post(rerank))
.route("/v1/rerank", post(v1_rerank))
.route("/v1/responses", post(v1_responses))
.route("/v1/embeddings", post(v1_embeddings))
@@ -609,6 +597,8 @@ pub fn build_app(
.route("/health_generate", get(health_generate))
.route("/engine_metrics", get(engine_metrics))
.route("/v1/models", get(v1_models))
.route("/model_info", get(get_model_info))
// TODO: Remove `/get_model_info` alias after one release-cycle deprecation window.
.route("/get_model_info", get(get_model_info))
.route("/server_info", get(get_server_info))
// TODO: Remove `/get_server_info` alias after one release-cycle deprecation window.
@@ -617,6 +607,8 @@ pub fn build_app(
// Build admin routes with control plane auth if configured, otherwise use simple API key auth
let admin_routes = Router::new()
.route("/flush_cache", post(flush_cache))
.route("/v1/loads", get(get_loads))
// TODO: Remove `/get_loads` alias after one release-cycle deprecation window.
.route("/get_loads", get(get_loads))
.route("/parse/function_call", post(parse_function_call))
.route("/parse/reasoning", post(parse_reasoning))
@@ -1570,182 +1570,6 @@ mod rerank_tests {
use super::*;
// Note: RerankRequest and RerankResult are available for future use
#[tokio::test]
async fn test_rerank_success() {
let ctx = AppTestContext::new(vec![MockWorkerConfig {
port: 18105,
worker_type: WorkerType::Regular,
health_status: HealthStatus::Healthy,
response_delay_ms: 0,
fail_rate: 0.0,
}])
.await;
let app = ctx.create_app().await;
let payload = json!({
"query": "machine learning algorithms",
"documents": [
"Introduction to machine learning concepts",
"Deep learning neural networks tutorial"
],
"model": "test-rerank-model",
"top_k": 2,
"return_documents": true,
"rid": "test-request-123"
});
let req = Request::builder()
.method("POST")
.uri("/rerank")
.header(CONTENT_TYPE, "application/json")
.body(Body::from(serde_json::to_string(&payload).unwrap()))
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
let body = axum::body::to_bytes(resp.into_body(), usize::MAX)
.await
.unwrap();
let body_json: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert!(body_json.get("results").is_some());
assert!(body_json.get("model").is_some());
assert_eq!(body_json["model"], "test-rerank-model");
let results = body_json["results"].as_array().unwrap();
assert_eq!(results.len(), 2);
assert!(results[0]["score"].as_f64().unwrap() >= results[1]["score"].as_f64().unwrap());
ctx.shutdown().await;
}
#[tokio::test]
async fn test_rerank_with_top_k() {
let ctx = AppTestContext::new(vec![MockWorkerConfig {
port: 18106,
worker_type: WorkerType::Regular,
health_status: HealthStatus::Healthy,
response_delay_ms: 0,
fail_rate: 0.0,
}])
.await;
let app = ctx.create_app().await;
let payload = json!({
"query": "test query",
"documents": [
"Document 1",
"Document 2",
"Document 3"
],
"model": "test-model",
"top_k": 1,
"return_documents": true
});
let req = Request::builder()
.method("POST")
.uri("/rerank")
.header(CONTENT_TYPE, "application/json")
.body(Body::from(serde_json::to_string(&payload).unwrap()))
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
let body = axum::body::to_bytes(resp.into_body(), usize::MAX)
.await
.unwrap();
let body_json: serde_json::Value = serde_json::from_slice(&body).unwrap();
// Should only return top_k results
let results = body_json["results"].as_array().unwrap();
assert_eq!(results.len(), 1);
ctx.shutdown().await;
}
#[tokio::test]
async fn test_rerank_without_documents() {
let ctx = AppTestContext::new(vec![MockWorkerConfig {
port: 18107,
worker_type: WorkerType::Regular,
health_status: HealthStatus::Healthy,
response_delay_ms: 0,
fail_rate: 0.0,
}])
.await;
let app = ctx.create_app().await;
let payload = json!({
"query": "test query",
"documents": ["Document 1", "Document 2"],
"model": "test-model",
"return_documents": false
});
let req = Request::builder()
.method("POST")
.uri("/rerank")
.header(CONTENT_TYPE, "application/json")
.body(Body::from(serde_json::to_string(&payload).unwrap()))
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
let body = axum::body::to_bytes(resp.into_body(), usize::MAX)
.await
.unwrap();
let body_json: serde_json::Value = serde_json::from_slice(&body).unwrap();
// Documents should be null when return_documents is false
let results = body_json["results"].as_array().unwrap();
for result in results {
assert!(result.get("document").is_none());
}
ctx.shutdown().await;
}
#[tokio::test]
async fn test_rerank_worker_failure() {
let ctx = AppTestContext::new(vec![MockWorkerConfig {
port: 18108,
worker_type: WorkerType::Regular,
health_status: HealthStatus::Healthy,
response_delay_ms: 0,
fail_rate: 1.0, // Always fail
}])
.await;
let app = ctx.create_app().await;
let payload = json!({
"query": "test query",
"documents": ["Document 1"],
"model": "test-model"
});
let req = Request::builder()
.method("POST")
.uri("/rerank")
.header(CONTENT_TYPE, "application/json")
.body(Body::from(serde_json::to_string(&payload).unwrap()))
.unwrap();
let resp = app.oneshot(req).await.unwrap();
// Should return the worker's error response
assert_eq!(resp.status(), StatusCode::INTERNAL_SERVER_ERROR);
ctx.shutdown().await;
}
#[tokio::test]
async fn test_v1_rerank_compatibility() {
let ctx = AppTestContext::new(vec![MockWorkerConfig {
@@ -1802,85 +1626,4 @@ mod rerank_tests {
ctx.shutdown().await;
}
#[tokio::test]
async fn test_rerank_invalid_request() {
let ctx = AppTestContext::new(vec![MockWorkerConfig {
port: 18111,
worker_type: WorkerType::Regular,
health_status: HealthStatus::Healthy,
response_delay_ms: 0,
fail_rate: 0.0,
}])
.await;
let app = ctx.create_app().await;
let payload = json!({
"query": "",
"documents": ["Document 1", "Document 2"],
"model": "test-model"
});
let req = Request::builder()
.method("POST")
.uri("/rerank")
.header(CONTENT_TYPE, "application/json")
.body(Body::from(serde_json::to_string(&payload).unwrap()))
.unwrap();
let resp = app.clone().oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::BAD_REQUEST);
let payload = json!({
"query": " ",
"documents": ["Document 1", "Document 2"],
"model": "test-model"
});
let req = Request::builder()
.method("POST")
.uri("/rerank")
.header(CONTENT_TYPE, "application/json")
.body(Body::from(serde_json::to_string(&payload).unwrap()))
.unwrap();
let resp = app.clone().oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::BAD_REQUEST);
let payload = json!({
"query": "test query",
"documents": [],
"model": "test-model"
});
let req = Request::builder()
.method("POST")
.uri("/rerank")
.header(CONTENT_TYPE, "application/json")
.body(Body::from(serde_json::to_string(&payload).unwrap()))
.unwrap();
let resp = app.clone().oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::BAD_REQUEST);
let payload = json!({
"query": "test query",
"documents": ["Document 1", "Document 2"],
"model": "test-model",
"top_k": 0
});
let req = Request::builder()
.method("POST")
.uri("/rerank")
.header(CONTENT_TYPE, "application/json")
.body(Body::from(serde_json::to_string(&payload).unwrap()))
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::BAD_REQUEST);
ctx.shutdown().await;
}
}
@@ -83,7 +83,7 @@ impl MockWorker {
.route("/health", get(health_handler))
.route("/health_generate", get(health_generate_handler))
.route("/server_info", get(server_info_handler))
.route("/get_model_info", get(model_info_handler))
.route("/model_info", get(model_info_handler))
.route("/generate", post(generate_handler))
.route("/v1/chat/completions", post(chat_completions_handler))
.route("/v1/completions", post(completions_handler))
@@ -767,12 +767,12 @@ mod pd_routing_unit_tests {
("/health_generate", "GET", true), // Note: Python uses POST, we use GET
("/server_info", "GET", true),
("/v1/models", "GET", true),
("/get_model_info", "GET", true),
("/model_info", "GET", true),
("/generate", "POST", true),
("/v1/chat/completions", "POST", true),
("/v1/completions", "POST", true),
("/flush_cache", "POST", true),
("/get_loads", "GET", true),
("/v1/loads", "GET", true),
("/register", "POST", false), // NOT IMPLEMENTED - needs dynamic worker management
];