diff --git a/sgl-model-gateway/src/routers/error.rs b/sgl-model-gateway/src/routers/error.rs index 6f1bb9b8e..035d8086a 100644 --- a/sgl-model-gateway/src/routers/error.rs +++ b/sgl-model-gateway/src/routers/error.rs @@ -29,6 +29,14 @@ pub fn not_implemented(code: impl Into, message: impl Into) -> R create_error(StatusCode::NOT_IMPLEMENTED, code, message) } +pub fn bad_gateway(code: impl Into, message: impl Into) -> Response { + create_error(StatusCode::BAD_GATEWAY, code, message) +} + +pub fn method_not_allowed(code: impl Into, message: impl Into) -> Response { + create_error(StatusCode::METHOD_NOT_ALLOWED, code, message) +} + fn create_error( status: StatusCode, code: impl Into, diff --git a/sgl-model-gateway/src/routers/http/pd_router.rs b/sgl-model-gateway/src/routers/http/pd_router.rs index 534755918..4c9e24225 100644 --- a/sgl-model-gateway/src/routers/http/pd_router.rs +++ b/sgl-model-gateway/src/routers/http/pd_router.rs @@ -33,7 +33,7 @@ use crate::{ generate::GenerateRequest, rerank::RerankRequest, }, - routers::{header_utils, RouterTrait}, + routers::{error, header_utils, RouterTrait}, }; #[derive(Debug)] @@ -68,11 +68,7 @@ impl PDRouter { if let Some(worker_url) = first_worker_url { self.proxy_to_worker(worker_url, endpoint, headers).await } else { - ( - StatusCode::SERVICE_UNAVAILABLE, - "No prefill servers available".to_string(), - ) - .into_response() + error::service_unavailable("no_prefill_servers", "No prefill servers available") } } @@ -104,26 +100,50 @@ impl PDRouter { } Err(e) => { error!("Failed to read response body: {}", e); - ( - StatusCode::INTERNAL_SERVER_ERROR, + error::internal_error( + "read_response_body_failed", format!("Failed to read response body: {}", e), ) - .into_response() } } } Ok(res) => { let status = StatusCode::from_u16(res.status().as_u16()) .unwrap_or(StatusCode::INTERNAL_SERVER_ERROR); - (status, format!("{} server returned status: ", res.status())).into_response() + // Use the status code to determine which error function to use + match status { + StatusCode::BAD_REQUEST => error::bad_request( + "server_bad_request", + format!("Server returned status: {}", res.status()), + ), + StatusCode::NOT_FOUND => error::not_found( + "server_not_found", + format!("Server returned status: {}", res.status()), + ), + StatusCode::INTERNAL_SERVER_ERROR => error::internal_error( + "server_internal_error", + format!("Server returned status: {}", res.status()), + ), + StatusCode::SERVICE_UNAVAILABLE => error::service_unavailable( + "server_unavailable", + format!("Server returned status: {}", res.status()), + ), + StatusCode::BAD_GATEWAY => error::bad_gateway( + "server_bad_gateway", + format!("Server returned status: {}", res.status()), + ), + _ => error::internal_error( + "server_error", + format!("Server returned status: {}", res.status()), + ), + } } Err(e) => { error!("Failed to proxy request server: {}", e); - ( - StatusCode::INTERNAL_SERVER_ERROR, + error::internal_error( + "proxy_request_failed", format!("Failed to proxy request: {}", e), ) - .into_response() } } } @@ -142,20 +162,15 @@ impl PDRouter { fn handle_server_selection_error(error: String) -> Response { error!("Failed to select PD pair error={}", error); RouterMetrics::record_pd_error("server_selection"); - ( - StatusCode::SERVICE_UNAVAILABLE, + error::service_unavailable( + "server_selection_failed", format!("No available servers: {}", error), ) - .into_response() } fn handle_serialization_error(error: impl std::fmt::Display) -> Response { error!("Failed to serialize request error={}", error); - ( - StatusCode::INTERNAL_SERVER_ERROR, - "Failed to serialize request", - ) - .into_response() + error::internal_error("serialization_failed", "Failed to serialize request") } fn get_generate_batch_size(req: &GenerateRequest) -> Option { @@ -378,8 +393,71 @@ impl PDRouter { } else { // Handle non-streaming error response match res.bytes().await { - Ok(error_body) => (status, error_body).into_response(), - Err(e) => (status, format!("Decode server error: {}", e)).into_response(), + Ok(error_body) => { + // Try to parse error message from body, fallback to status-based error + let error_message = if let Ok(error_json) = + serde_json::from_slice::(&error_body) + { + if let Some(msg) = error_json + .get("error") + .and_then(|e| e.get("message")) + .and_then(|m| m.as_str()) + { + msg.to_string() + } else if let Some(msg) = error_json.get("message").and_then(|m| m.as_str()) + { + msg.to_string() + } else { + String::from_utf8_lossy(&error_body).to_string() + } + } else { + String::from_utf8_lossy(&error_body).to_string() + }; + + let status_code = StatusCode::from_u16(status.as_u16()) + .unwrap_or(StatusCode::INTERNAL_SERVER_ERROR); + match status_code { + StatusCode::BAD_REQUEST => { + error::bad_request("decode_bad_request", error_message) + } + StatusCode::NOT_FOUND => { + error::not_found("decode_not_found", error_message) + } + StatusCode::INTERNAL_SERVER_ERROR => { + error::internal_error("decode_internal_error", error_message) + } + StatusCode::SERVICE_UNAVAILABLE => { + error::service_unavailable("decode_unavailable", error_message) + } + StatusCode::BAD_GATEWAY => { + error::bad_gateway("decode_bad_gateway", error_message) + } + _ => error::internal_error("decode_error", error_message), + } + } + Err(e) => { + let error_message = format!("Decode server error: {}", e); + let status_code = StatusCode::from_u16(status.as_u16()) + .unwrap_or(StatusCode::INTERNAL_SERVER_ERROR); + match status_code { + StatusCode::BAD_REQUEST => { + error::bad_request("decode_read_failed", error_message) + } + StatusCode::NOT_FOUND => { + error::not_found("decode_read_failed", error_message) + } + StatusCode::INTERNAL_SERVER_ERROR => { + error::internal_error("decode_read_failed", error_message) + } + StatusCode::SERVICE_UNAVAILABLE => { + error::service_unavailable("decode_read_failed", error_message) + } + StatusCode::BAD_GATEWAY => { + error::bad_gateway("decode_read_failed", error_message) + } + _ => error::internal_error("decode_read_failed", error_message), + } + } } } } @@ -535,8 +613,10 @@ impl PDRouter { } Err(e) => { error!("Failed to read decode response: {}", e); - (StatusCode::INTERNAL_SERVER_ERROR, "Failed to read response") - .into_response() + error::internal_error( + "read_response_failed", + "Failed to read response", + ) } } } @@ -549,11 +629,7 @@ impl PDRouter { "Decode request failed" ); RouterMetrics::record_pd_decode_error(decode.url()); - ( - StatusCode::BAD_GATEWAY, - format!("Decode server error: {}", e), - ) - .into_response() + error::bad_gateway("decode_server_error", format!("Decode server error: {}", e)) } } } @@ -759,8 +835,7 @@ impl PDRouter { Ok(decode_body) => decode_body, Err(e) => { error!("Failed to read decode response: {}", e); - return (StatusCode::INTERNAL_SERVER_ERROR, "Failed to read response") - .into_response(); + return error::internal_error("read_response_failed", "Failed to read response"); } }; @@ -812,14 +887,13 @@ impl PDRouter { ); // Return error immediately - don't wait for decode to timeout - return Err(( - StatusCode::BAD_GATEWAY, + return Err(error::bad_gateway( + "prefill_server_error", format!( "Prefill server error: {}. This will cause decode timeout.", e ), - ) - .into_response()); + )); } }; @@ -841,11 +915,34 @@ impl PDRouter { prefill_url, prefill_status, error_msg ); - return Err(( - prefill_status, - format!("Prefill server error ({}): {}", prefill_status, error_msg), - ) - .into_response()); + // Map prefill_status to appropriate error function + let error_response = match prefill_status { + StatusCode::BAD_REQUEST => error::bad_request( + "prefill_bad_request", + format!("Prefill server error ({}): {}", prefill_status, error_msg), + ), + StatusCode::NOT_FOUND => error::not_found( + "prefill_not_found", + format!("Prefill server error ({}): {}", prefill_status, error_msg), + ), + StatusCode::INTERNAL_SERVER_ERROR => error::internal_error( + "prefill_internal_error", + format!("Prefill server error ({}): {}", prefill_status, error_msg), + ), + StatusCode::SERVICE_UNAVAILABLE => error::service_unavailable( + "prefill_unavailable", + format!("Prefill server error ({}): {}", prefill_status, error_msg), + ), + StatusCode::BAD_GATEWAY => error::bad_gateway( + "prefill_bad_gateway", + format!("Prefill server error ({}): {}", prefill_status, error_msg), + ), + _ => error::internal_error( + "prefill_error", + format!("Prefill server error ({}): {}", prefill_status, error_msg), + ), + }; + return Err(error_response); } // Read prefill body if needed for logprob merging @@ -990,11 +1087,10 @@ impl RouterTrait for PDRouter { let (prefill, decode) = match self.select_pd_pair(None, None).await { Ok(pair) => pair, Err(e) => { - return ( - StatusCode::SERVICE_UNAVAILABLE, + return error::service_unavailable( + "no_healthy_worker_pair", format!("No healthy worker pair available: {}", e), - ) - .into_response(); + ); } }; @@ -1055,11 +1151,10 @@ impl RouterTrait for PDRouter { ) .into_response() } else { - ( - StatusCode::SERVICE_UNAVAILABLE, + error::service_unavailable( + "health_generate_failed", format!("Health generate failed: {:?}", errors), ) - .into_response() } } diff --git a/sgl-model-gateway/src/routers/http/router.rs b/sgl-model-gateway/src/routers/http/router.rs index 9e44e4ab6..0cb028865 100644 --- a/sgl-model-gateway/src/routers/http/router.rs +++ b/sgl-model-gateway/src/routers/http/router.rs @@ -37,7 +37,7 @@ use crate::{ rerank::{RerankRequest, RerankResponse, RerankResult}, responses::{ResponsesGetParams, ResponsesRequest}, }, - routers::{header_utils, RouterTrait}, + routers::{error, header_utils, RouterTrait}, }; /// Regular router that uses injected load balancing policies @@ -106,21 +106,18 @@ impl Router { *response.headers_mut() = response_headers; response } - Err(e) => ( - StatusCode::INTERNAL_SERVER_ERROR, + Err(e) => error::internal_error( + "read_response_failed", format!("Failed to read response: {}", e), - ) - .into_response(), + ), } } - Err(e) => ( - StatusCode::INTERNAL_SERVER_ERROR, - format!("Request failed: {}", e), - ) - .into_response(), + Err(e) => { + error::internal_error("request_failed", format!("Request failed: {}", e)) + } } } - Err(e) => (StatusCode::SERVICE_UNAVAILABLE, e).into_response(), + Err(e) => error::service_unavailable("no_workers", e), } } @@ -214,11 +211,10 @@ impl Router { Some(w) => w, None => { RouterMetrics::record_request_error(route, "no_available_workers"); - return ( - StatusCode::SERVICE_UNAVAILABLE, + return error::service_unavailable( + "no_available_workers", "No available workers (all circuits open or unhealthy)", - ) - .into_response(); + ); } }; @@ -298,7 +294,7 @@ impl Router { // Eventually, we need to have router to manage the chat history with a proper database, will update this implementation accordingly. let workers = self.worker_registry.get_all(); if workers.is_empty() { - return (StatusCode::SERVICE_UNAVAILABLE, "No available workers").into_response(); + return error::service_unavailable("no_workers", "No available workers"); } // Pre-filter headers once before the loop to avoid repeated lowercasing @@ -323,11 +319,10 @@ impl Router { Method::GET => self.client.get(url), Method::POST => self.client.post(url), _ => { - return ( - StatusCode::METHOD_NOT_ALLOWED, + return error::method_not_allowed( + "unsupported_method", "Unsupported method for simple routing", ) - .into_response() } }; @@ -360,30 +355,24 @@ impl Router { last_response = Some(response); } Err(e) => { - last_response = Some( - ( - StatusCode::INTERNAL_SERVER_ERROR, - format!("Failed to read response: {}", e), - ) - .into_response(), - ); + last_response = Some(error::internal_error( + "read_response_failed", + format!("Failed to read response: {}", e), + )); } } } Err(e) => { - last_response = Some( - ( - StatusCode::INTERNAL_SERVER_ERROR, - format!("Request failed: {}", e), - ) - .into_response(), - ); + last_response = Some(error::internal_error( + "request_failed", + format!("Request failed: {}", e), + )); } } } last_response - .unwrap_or_else(|| (StatusCode::BAD_GATEWAY, "No worker response").into_response()) + .unwrap_or_else(|| error::bad_gateway("no_worker_response", "No worker response")) } // Route a GET request with provided headers to a specific endpoint @@ -441,22 +430,20 @@ impl Router { Ok(tup) => tup, Err(e) => { error!("Failed to extract dp_rank: {}", e); - return ( - StatusCode::INTERNAL_SERVER_ERROR, + return error::internal_error( + "dp_rank_extraction_failed", format!("Failed to extract dp_rank: {}", e), - ) - .into_response(); + ); } }; let mut json_val = match serde_json::to_value(typed_req) { Ok(j) => j, Err(e) => { - return ( - StatusCode::BAD_REQUEST, + return error::bad_request( + "serialization_failed", format!("Convert into serde_json::Value failed: {}", e), - ) - .into_response(); + ); } }; @@ -471,11 +458,10 @@ impl Router { ); } } else { - return ( - StatusCode::BAD_REQUEST, + return error::bad_request( + "dp_rank_insertion_failed", "Failed to insert the data_parallel_rank field into the request body", - ) - .into_response(); + ); } self.client @@ -520,11 +506,7 @@ impl Router { } } - return ( - StatusCode::INTERNAL_SERVER_ERROR, - format!("Request failed: {}", e), - ) - .into_response(); + return error::internal_error("request_failed", format!("Request failed: {}", e)); } }; @@ -553,7 +535,7 @@ impl Router { } let error_msg = format!("Failed to get response body: {}", e); - (StatusCode::INTERNAL_SERVER_ERROR, error_msg).into_response() + error::internal_error("read_response_body_failed", error_msg) } }; @@ -825,11 +807,10 @@ impl RouterTrait for Router { Ok(rerank_response) => rerank_response, Err(e) => { error!("Failed to build rerank response: {}", e); - return ( - StatusCode::INTERNAL_SERVER_ERROR, - "Failed to build rerank response".to_string(), - ) - .into_response(); + return error::internal_error( + "rerank_response_build_failed", + "Failed to build rerank response", + ); } } } else {