[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:
co-authored by
Claude Opus 4.7
parent
e0474fdd9b
commit
d74d9bda49
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
];
|
||||
|
||||
|
||||
Reference in New Issue
Block a user