diff --git a/sgl-model-gateway/src/grpc_client/vllm_engine.rs b/sgl-model-gateway/src/grpc_client/vllm_engine.rs index 986682f8b..d9892e417 100644 --- a/sgl-model-gateway/src/grpc_client/vllm_engine.rs +++ b/sgl-model-gateway/src/grpc_client/vllm_engine.rs @@ -1,5 +1,4 @@ use std::{ - convert::TryFrom, pin::Pin, sync::{ atomic::{AtomicBool, Ordering}, @@ -198,25 +197,16 @@ impl VllmEngineClient { pub async fn abort_request( &self, request_id: String, - reason: String, + _reason: String, ) -> Result<(), Box> { - debug!( - "Sending abort request for {} (reason: {})", - request_id, reason - ); + debug!("Sending abort request for {}", request_id); let request = Request::new(proto::AbortRequest { - request_id: request_id.clone(), - reason, + request_ids: vec![request_id.clone()], }); let mut client = self.client.clone(); - let response = client.abort(request).await?; - debug!( - "Abort response for {}: success={}, message={}", - request_id, - response.get_ref().success, - response.get_ref().message - ); + let _response = client.abort(request).await?; + debug!("Abort response received for {}", request_id); Ok(()) } @@ -261,10 +251,12 @@ impl VllmEngineClient { let grpc_request = proto::GenerateRequest { request_id, - tokenized: Some(proto::TokenizedInput { - original_text: processed_text, - input_ids: token_ids, - }), + input: Some(proto::generate_request::Input::Tokenized( + proto::TokenizedInput { + original_text: processed_text, + input_ids: token_ids, + }, + )), sampling_params: Some(sampling_params), stream: body.stream, }; @@ -285,10 +277,12 @@ impl VllmEngineClient { let grpc_request = proto::GenerateRequest { request_id, - tokenized: Some(proto::TokenizedInput { - original_text: original_text.unwrap_or_default(), - input_ids: token_ids, - }), + input: Some(proto::generate_request::Input::Tokenized( + proto::TokenizedInput { + original_text: original_text.unwrap_or_default(), + input_ids: token_ids, + }, + )), sampling_params: Some(sampling_params), stream: body.stream, }; @@ -320,10 +314,12 @@ impl VllmEngineClient { let grpc_request = proto::GenerateRequest { request_id, - tokenized: Some(proto::TokenizedInput { - original_text: processed_text, - input_ids: token_ids, - }), + input: Some(proto::generate_request::Input::Tokenized( + proto::TokenizedInput { + original_text: processed_text, + input_ids: token_ids, + }, + )), sampling_params: Some(sampling_params), stream: body.stream.unwrap_or(false), }; @@ -339,7 +335,7 @@ impl VllmEngineClient { ) -> Result { let stop_sequences = self.extract_stop_strings(request); - let max_tokens = request.max_completion_tokens.map(|v| v as i32); + let max_tokens = request.max_completion_tokens; // Handle skip_special_tokens: set to false if tools are present and tool_choice is not "none" let skip_special_tokens = if request.tools.is_some() { @@ -353,9 +349,9 @@ impl VllmEngineClient { }; Ok(proto::SamplingParams { - temperature: request.temperature.unwrap_or(1.0), + temperature: request.temperature, top_p: request.top_p.unwrap_or(1.0), - top_k: request.top_k.unwrap_or(-1), + top_k: request.top_k.map(|v| v.max(0) as u32).unwrap_or(0), // 0 means disabled in vLLM min_p: request.min_p.unwrap_or(0.0), frequency_penalty: request.frequency_penalty.unwrap_or(0.0), presence_penalty: request.presence_penalty.unwrap_or(0.0), @@ -366,7 +362,7 @@ impl VllmEngineClient { skip_special_tokens, spaces_between_special_tokens: true, // Default from Python SamplingParams ignore_eos: request.ignore_eos, - n: request.n.unwrap_or(1) as i32, + n: request.n.unwrap_or(1), constraint: self.build_constraint_for_chat(request, tool_call_constraint)?, ..Default::default() }) @@ -450,15 +446,15 @@ impl VllmEngineClient { // Used by Harmony models only. Regular models use Chat API path. // Constraints come from Harmony preparation stage (structural_tag) or tool handling. - let max_tokens = request.max_output_tokens.map(|v| v as i32); + let max_tokens = request.max_output_tokens; Ok(proto::SamplingParams { - temperature: request.temperature.unwrap_or(1.0), + temperature: request.temperature, top_p: request.top_p.unwrap_or(1.0), - top_k: -1, // ResponsesRequest doesn't expose top_k - min_p: 0.0, // ResponsesRequest doesn't expose min_p - frequency_penalty: 0.0, // ResponsesRequest doesn't expose frequency_penalty - presence_penalty: 0.0, // ResponsesRequest doesn't expose presence_penalty + top_k: 0, // ResponsesRequest doesn't expose top_k (0 means disabled) + min_p: 0.0, // ResponsesRequest doesn't expose min_p + frequency_penalty: 0.0, // ResponsesRequest doesn't expose frequency_penalty + presence_penalty: 0.0, // ResponsesRequest doesn't expose presence_penalty repetition_penalty: 1.0, // ResponsesRequest doesn't expose repetition_penalty max_tokens, stop: vec![], // No stop sequences in Responses API @@ -525,9 +521,9 @@ impl VllmEngineClient { params: Option<&GenerateSamplingParams>, ) -> Result { let mut sampling = proto::SamplingParams { - temperature: 1.0, + temperature: Some(1.0), top_p: 1.0, - top_k: -1, + top_k: 0, // 0 means disabled in vLLM repetition_penalty: 1.0, n: 1, skip_special_tokens: true, @@ -539,24 +535,36 @@ impl VllmEngineClient { return Ok(sampling); }; - // Simple field mappings using a macro - macro_rules! map_field { - ($field:ident) => { - if let Some(val) = p.$field { - sampling.$field = val; - } - }; + // Handle temperature (now optional) + if let Some(val) = p.temperature { + sampling.temperature = Some(val); } - map_field!(temperature); - map_field!(top_p); - map_field!(top_k); - map_field!(frequency_penalty); - map_field!(presence_penalty); - map_field!(repetition_penalty); - map_field!(min_p); - map_field!(ignore_eos); - map_field!(skip_special_tokens); + // Simple field mappings + if let Some(val) = p.top_p { + sampling.top_p = val; + } + if let Some(val) = p.top_k { + sampling.top_k = val.max(0) as u32; // Clamp negative values to 0 (disabled) + } + if let Some(val) = p.frequency_penalty { + sampling.frequency_penalty = val; + } + if let Some(val) = p.presence_penalty { + sampling.presence_penalty = val; + } + if let Some(val) = p.repetition_penalty { + sampling.repetition_penalty = val; + } + if let Some(val) = p.min_p { + sampling.min_p = val; + } + if let Some(val) = p.ignore_eos { + sampling.ignore_eos = val; + } + if let Some(val) = p.skip_special_tokens { + sampling.skip_special_tokens = val; + } // Note: no_stop_trim not supported in vLLM // Handle stop sequences @@ -572,24 +580,19 @@ impl VllmEngineClient { sampling.stop_token_ids = stop_token_ids.clone(); } - // Handle max_tokens with conversion (read from internal max_new_tokens) + // Handle max_tokens (read from internal max_new_tokens) if let Some(max_new_tokens) = p.max_new_tokens { - sampling.max_tokens = Some( - i32::try_from(max_new_tokens) - .map_err(|_| "max_tokens must fit into a 32-bit signed integer".to_string())?, - ); + sampling.max_tokens = Some(max_new_tokens); } - // Handle min_tokens with conversion (read from internal min_new_tokens) + // Handle min_tokens (read from internal min_new_tokens) if let Some(min_new_tokens) = p.min_new_tokens { - sampling.min_tokens = i32::try_from(min_new_tokens) - .map_err(|_| "min_tokens must fit into a 32-bit signed integer".to_string())?; + sampling.min_tokens = min_new_tokens; } - // Handle n with conversion + // Handle n if let Some(n) = p.n { - sampling.n = i32::try_from(n) - .map_err(|_| "n must fit into a 32-bit signed integer".to_string())?; + sampling.n = n; } // Handle constraints (exactly one allowed) @@ -612,7 +615,7 @@ mod tests { #[test] fn test_generate_request_construction() { let sampling_params = proto::SamplingParams { - temperature: 0.7, + temperature: Some(0.7), max_tokens: Some(128), top_p: 0.9, top_k: 50, @@ -622,22 +625,24 @@ mod tests { let gen_req = proto::GenerateRequest { request_id: "test-req-123".to_string(), - tokenized: Some(proto::TokenizedInput { - original_text: "Hello world".to_string(), - input_ids: vec![9906, 1917], // Mock token IDs for "Hello world" - }), + input: Some(proto::generate_request::Input::Tokenized( + proto::TokenizedInput { + original_text: "Hello world".to_string(), + input_ids: vec![9906, 1917], // Mock token IDs for "Hello world" + }, + )), sampling_params: Some(sampling_params), stream: false, }; assert_eq!(gen_req.request_id, "test-req-123"); - if let Some(ref tokenized) = &gen_req.tokenized { + if let Some(proto::generate_request::Input::Tokenized(ref tokenized)) = gen_req.input { assert_eq!(tokenized.original_text, "Hello world"); } // vLLM: logprobs are in SamplingParams, not GenerateRequest let params = gen_req.sampling_params.unwrap(); - assert_eq!(params.temperature, 0.7); + assert_eq!(params.temperature, Some(0.7)); assert_eq!(params.max_tokens, Some(128)); assert_eq!(params.stop, vec![""]); } @@ -651,18 +656,17 @@ mod tests { #[test] fn test_abort_request_construction() { let abort_req = proto::AbortRequest { - request_id: "req-456".to_string(), - reason: "User canceled".to_string(), + request_ids: vec!["req-456".to_string(), "req-789".to_string()], }; - assert_eq!(abort_req.request_id, "req-456"); - assert_eq!(abort_req.reason, "User canceled"); + assert_eq!(abort_req.request_ids, vec!["req-456", "req-789"]); } #[test] fn test_sampling_params_defaults() { let params = proto::SamplingParams::default(); - // Numeric fields have proto defaults (0) - assert_eq!(params.temperature, 0.0); + // Optional float field defaults to None + assert_eq!(params.temperature, None); + // Non-optional numeric fields have proto defaults (0) assert_eq!(params.top_p, 0.0); assert_eq!(params.top_k, 0); assert_eq!(params.repetition_penalty, 0.0); @@ -672,7 +676,7 @@ mod tests { assert!(!params.spaces_between_special_tokens); assert!(!params.ignore_eos); assert!(!params.include_stop_str_in_output); - // Optional int fields should be None + // Optional fields should be None assert_eq!(params.max_tokens, None); assert_eq!(params.logprobs, None); // Other non-optional fields diff --git a/sgl-model-gateway/src/proto/vllm_engine.proto b/sgl-model-gateway/src/proto/vllm_engine.proto index f4e7934ba..bbb1b9b00 100644 --- a/sgl-model-gateway/src/proto/vllm_engine.proto +++ b/sgl-model-gateway/src/proto/vllm_engine.proto @@ -31,16 +31,16 @@ service VllmEngine { // Sampling parameters for text generation message SamplingParams { - float temperature = 1; + optional float temperature = 1; float top_p = 2; - int32 top_k = 3; + uint32 top_k = 3; float min_p = 4; float frequency_penalty = 5; float presence_penalty = 6; float repetition_penalty = 7; - optional int32 max_tokens = 8; - int32 min_tokens = 9; + optional uint32 max_tokens = 8; + uint32 min_tokens = 9; repeated string stop = 10; repeated uint32 stop_token_ids = 11; @@ -49,7 +49,7 @@ message SamplingParams { bool spaces_between_special_tokens = 13; bool ignore_eos = 14; - int32 n = 15; // Number of parallel samples + uint32 n = 15; // Number of parallel samples // Logprobs configuration optional int32 logprobs = 22; // Number of log probabilities per output token (-1 for all) @@ -90,14 +90,17 @@ message TokenizedInput { message GenerateRequest { string request_id = 1; - // Pre-tokenized input (required) - TokenizedInput tokenized = 2; + // Prompt input + oneof input { + TokenizedInput tokenized = 2; + string text = 3; + } // Generation parameters (includes logprobs config) - SamplingParams sampling_params = 3; + SamplingParams sampling_params = 4; // Streaming - bool stream = 4; + bool stream = 5; } // ===================== @@ -105,20 +108,17 @@ message GenerateRequest { // ===================== message GenerateResponse { - string request_id = 1; - oneof response { - GenerateStreamChunk chunk = 2; // For streaming - GenerateComplete complete = 3; // For final/non-streaming - GenerateError error = 4; // For errors + GenerateStreamChunk chunk = 1; // For streaming + GenerateComplete complete = 2; // For final/non-streaming } } message GenerateStreamChunk { repeated uint32 token_ids = 1; // Incremental tokens - int32 prompt_tokens = 2; - int32 completion_tokens = 3; - int32 cached_tokens = 4; + uint32 prompt_tokens = 2; + uint32 completion_tokens = 3; + uint32 cached_tokens = 4; // Logprobs support (TODO: implement in Phase 4) // OutputLogProbs output_logprobs = 5; @@ -128,21 +128,15 @@ message GenerateStreamChunk { message GenerateComplete { repeated uint32 output_ids = 1; // All output tokens string finish_reason = 2; // "stop", "length", "abort" - int32 prompt_tokens = 3; - int32 completion_tokens = 4; - int32 cached_tokens = 5; + uint32 prompt_tokens = 3; + uint32 completion_tokens = 4; + uint32 cached_tokens = 5; // Logprobs support (TODO: implement in Phase 4) // OutputLogProbs output_logprobs = 6; // InputLogProbs input_logprobs = 7; } -message GenerateError { - string message = 1; - string http_status_code = 2; - string details = 3; -} - // ===================== // Embedding Request // ===================== @@ -153,23 +147,9 @@ message EmbedRequest { } message EmbedResponse { - string request_id = 1; - - oneof response { - EmbedComplete complete = 2; - EmbedError error = 3; - } -} - -message EmbedComplete { repeated float embedding = 1; - int32 prompt_tokens = 2; - int32 embedding_dim = 3; -} - -message EmbedError { - string message = 1; - string code = 2; + uint32 prompt_tokens = 2; + uint32 embedding_dim = 3; } // ===================== @@ -184,13 +164,10 @@ message HealthCheckResponse { } message AbortRequest { - string request_id = 1; - string reason = 2; + repeated string request_ids = 1; } message AbortResponse { - bool success = 1; - string message = 2; } // ===================== @@ -202,15 +179,15 @@ message GetModelInfoRequest {} message GetModelInfoResponse { string model_path = 1; bool is_generation = 2; - int32 max_context_length = 3; - int32 vocab_size = 4; + uint32 max_context_length = 3; + uint32 vocab_size = 4; bool supports_vision = 5; } message GetServerInfoRequest {} message GetServerInfoResponse { - int32 active_requests = 1; + uint32 active_requests = 1; bool is_paused = 2; double last_receive_timestamp = 3; double uptime_seconds = 4; diff --git a/sgl-model-gateway/src/routers/grpc/proto_wrapper.rs b/sgl-model-gateway/src/routers/grpc/proto_wrapper.rs index 81dddaada..44e42b0b6 100644 --- a/sgl-model-gateway/src/routers/grpc/proto_wrapper.rs +++ b/sgl-model-gateway/src/routers/grpc/proto_wrapper.rs @@ -95,7 +95,7 @@ impl ProtoGenerateRequest { /// Unified GenerateResponse from stream pub enum ProtoGenerateResponse { - Sglang(sglang::GenerateResponse), + Sglang(Box), Vllm(vllm::GenerateResponse), } @@ -124,9 +124,7 @@ impl ProtoGenerateResponse { Some(vllm::generate_response::Response::Complete(complete)) => { ProtoResponseVariant::Complete(ProtoGenerateComplete::Vllm(complete)) } - Some(vllm::generate_response::Response::Error(error)) => { - ProtoResponseVariant::Error(ProtoGenerateError::Vllm(error)) - } + // Note: vLLM proto no longer has Error variant in GenerateResponse None => ProtoResponseVariant::None, }, } @@ -204,7 +202,7 @@ impl ProtoGenerateStreamChunk { pub fn prompt_tokens(&self) -> i32 { match self { Self::Sglang(c) => c.prompt_tokens, - Self::Vllm(c) => c.prompt_tokens, + Self::Vllm(c) => c.prompt_tokens as i32, } } @@ -212,7 +210,7 @@ impl ProtoGenerateStreamChunk { pub fn completion_tokens(&self) -> i32 { match self { Self::Sglang(c) => c.completion_tokens, - Self::Vllm(c) => c.completion_tokens, + Self::Vllm(c) => c.completion_tokens as i32, } } @@ -220,7 +218,7 @@ impl ProtoGenerateStreamChunk { pub fn cached_tokens(&self) -> i32 { match self { Self::Sglang(c) => c.cached_tokens, - Self::Vllm(c) => c.cached_tokens, + Self::Vllm(c) => c.cached_tokens as i32, } } } @@ -279,7 +277,7 @@ impl ProtoGenerateComplete { pub fn prompt_tokens(&self) -> i32 { match self { Self::Sglang(c) => c.prompt_tokens, - Self::Vllm(c) => c.prompt_tokens, + Self::Vllm(c) => c.prompt_tokens as i32, } } @@ -287,7 +285,7 @@ impl ProtoGenerateComplete { pub fn completion_tokens(&self) -> i32 { match self { Self::Sglang(c) => c.completion_tokens, - Self::Vllm(c) => c.completion_tokens, + Self::Vllm(c) => c.completion_tokens as i32, } } @@ -329,7 +327,7 @@ impl ProtoGenerateComplete { pub fn cached_tokens(&self) -> i32 { match self { Self::Sglang(c) => c.cached_tokens, - Self::Vllm(_) => 0, // vLLM doesn't have cached_tokens field + Self::Vllm(c) => c.cached_tokens as i32, } } @@ -351,10 +349,10 @@ impl ProtoGenerateComplete { } /// Unified GenerateError +/// Note: vLLM proto no longer has GenerateError - errors are returned via gRPC status #[derive(Clone)] pub enum ProtoGenerateError { Sglang(sglang::GenerateError), - Vllm(vllm::GenerateError), } impl ProtoGenerateError { @@ -362,7 +360,6 @@ impl ProtoGenerateError { pub fn message(&self) -> &str { match self { Self::Sglang(e) => &e.message, - Self::Vllm(e) => &e.message, } } } @@ -380,7 +377,7 @@ impl ProtoStream { Self::Sglang(stream) => stream .next() .await - .map(|result| result.map(ProtoGenerateResponse::Sglang)), + .map(|result| result.map(|r| ProtoGenerateResponse::Sglang(Box::new(r)))), Self::Vllm(stream) => stream .next() .await