diff --git a/sgl-model-gateway/src/routers/http/pd_router.rs b/sgl-model-gateway/src/routers/http/pd_router.rs index a939b6427..45e801f36 100644 --- a/sgl-model-gateway/src/routers/http/pd_router.rs +++ b/sgl-model-gateway/src/routers/http/pd_router.rs @@ -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 } diff --git a/sgl-model-gateway/src/routers/http/router.rs b/sgl-model-gateway/src/routers/http/router.rs index b02f6638d..ffb036d91 100644 --- a/sgl-model-gateway/src/routers/http/router.rs +++ b/sgl-model-gateway/src/routers/http/router.rs @@ -732,7 +732,7 @@ impl RouterTrait for Router { } async fn get_model_info(&self, req: Request) -> Response { - self.proxy_get_request(req, "get_model_info").await + self.proxy_get_request(req, "model_info").await } async fn route_generate( diff --git a/sgl-model-gateway/src/server.rs b/sgl-model-gateway/src/server.rs index db23d0aad..99af905a3 100644 --- a/sgl-model-gateway/src/server.rs +++ b/sgl-model-gateway/src/server.rs @@ -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>, - headers: http::HeaderMap, - ValidatedJson(body): ValidatedJson, -) -> Response { - state - .router - .route_rerank(Some(&headers), &body, Some(&body.model)) - .await -} - async fn v1_rerank( State(state): State>, 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)) diff --git a/sgl-model-gateway/tests/api/api_endpoints_test.rs b/sgl-model-gateway/tests/api/api_endpoints_test.rs index 6e6ff125e..7fbb38c05 100644 --- a/sgl-model-gateway/tests/api/api_endpoints_test.rs +++ b/sgl-model-gateway/tests/api/api_endpoints_test.rs @@ -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; - } } diff --git a/sgl-model-gateway/tests/common/mock_worker.rs b/sgl-model-gateway/tests/common/mock_worker.rs index 166fc9d83..19e863ea1 100755 --- a/sgl-model-gateway/tests/common/mock_worker.rs +++ b/sgl-model-gateway/tests/common/mock_worker.rs @@ -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)) diff --git a/sgl-model-gateway/tests/routing/test_pd_routing.rs b/sgl-model-gateway/tests/routing/test_pd_routing.rs index b6c69576f..3587512b8 100644 --- a/sgl-model-gateway/tests/routing/test_pd_routing.rs +++ b/sgl-model-gateway/tests/routing/test_pd_routing.rs @@ -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 ];