[gRPC] Native gRPC server: proto + Rust crate scaffold + server args (#22736)

This commit is contained in:
Alex Nails
2026-04-20 12:39:35 +08:00
committed by GitHub
parent c304d0d64d
commit 10e17cc55e
9 changed files with 483 additions and 1 deletions
+8
View File
@@ -639,6 +639,14 @@ jobs:
- name: Install uv
uses: astral-sh/setup-uv@v5
# Needed by setuptools-rust to build the bundled native gRPC extension
# (rust/sglang-grpc) when installing the main `sglang` wheel from source.
- name: Install protoc
run: sudo bash scripts/ci/utils/install_protoc.sh
- name: Install Rust toolchain
run: bash scripts/ci/utils/install_rustup.sh
# uv pip targets a venv by default; setup-python has no venv — install into that interpreter (see UV_SYSTEM_PYTHON in https://docs.astral.sh/uv/guides/integration/github/)
- name: Install dependencies
timeout-minutes: 20
+295
View File
@@ -0,0 +1,295 @@
syntax = "proto3";
package sglang.runtime.v1;
service SglangService {
// SGLang-native RPCs (typed proto)
rpc TextGenerate(TextGenerateRequest) returns (stream TextGenerateResponse);
rpc Generate(GenerateRequest) returns (stream GenerateResponse);
rpc TextEmbed(TextEmbedRequest) returns (TextEmbedResponse);
rpc Embed(EmbedRequest) returns (EmbedResponse);
rpc Classify(ClassifyRequest) returns (ClassifyResponse);
rpc Tokenize(TokenizeRequest) returns (TokenizeResponse);
rpc Detokenize(DetokenizeRequest) returns (DetokenizeResponse);
rpc HealthCheck(HealthCheckRequest) returns (HealthCheckResponse);
rpc GetModelInfo(GetModelInfoRequest) returns (GetModelInfoResponse);
rpc GetServerInfo(GetServerInfoRequest) returns (GetServerInfoResponse);
rpc ListModels(ListModelsRequest) returns (ListModelsResponse);
rpc GetLoad(GetLoadRequest) returns (GetLoadResponse);
rpc Abort(AbortRequest) returns (AbortResponse);
rpc FlushCache(FlushCacheRequest) returns (FlushCacheResponse);
rpc PauseGeneration(PauseGenerationRequest) returns (PauseGenerationResponse);
rpc ContinueGeneration(ContinueGenerationRequest) returns (ContinueGenerationResponse);
// OpenAI-compatible RPCs (JSON pass-through)
rpc ChatComplete(OpenAIRequest) returns (stream OpenAIStreamChunk);
rpc Complete(OpenAIRequest) returns (stream OpenAIStreamChunk);
rpc OpenAIEmbed(OpenAIRequest) returns (OpenAIResponse);
rpc OpenAIClassify(OpenAIRequest) returns (OpenAIResponse);
rpc Score(OpenAIRequest) returns (OpenAIResponse);
rpc Rerank(OpenAIRequest) returns (OpenAIResponse);
// Admin/Ops RPCs
rpc StartProfile(StartProfileRequest) returns (StartProfileResponse);
rpc StopProfile(StopProfileRequest) returns (StopProfileResponse);
rpc UpdateWeightsFromDisk(UpdateWeightsRequest) returns (UpdateWeightsResponse);
}
// Sampling parameters shared across text and tokenized RPCs.
message SamplingParams {
optional float temperature = 1;
optional float top_p = 2;
optional int32 top_k = 3;
optional float min_p = 4;
optional float frequency_penalty = 5;
optional float presence_penalty = 6;
optional float repetition_penalty = 7;
optional int32 max_new_tokens = 8;
optional int32 min_new_tokens = 9;
repeated string stop = 10;
repeated int32 stop_token_ids = 11;
optional bool ignore_eos = 12;
optional int32 n = 13;
optional string json_schema = 14;
optional string regex = 15;
}
// ---- Text-based generate (text in, text out) ----
message TextGenerateRequest {
string text = 1;
optional SamplingParams sampling_params = 2;
optional bool stream = 3;
optional bool return_logprob = 4;
optional int32 top_logprobs_num = 5;
optional int32 logprob_start_len = 6;
optional bool return_text_in_logprobs = 7;
optional string rid = 8;
optional string lora_path = 9;
optional string routing_key = 10;
optional int32 routed_dp_rank = 11;
map<string, string> trace_headers = 12;
}
message TextGenerateResponse {
string text = 1;
map<string, string> meta_info = 2;
bool finished = 3;
}
// ---- Tokenized generate (input_ids in, token_ids out) ----
message GenerateRequest {
repeated int32 input_ids = 1;
optional SamplingParams sampling_params = 2;
optional bool stream = 3;
optional bool return_logprob = 4;
optional int32 top_logprobs_num = 5;
optional int32 logprob_start_len = 6;
optional string rid = 7;
optional string lora_path = 8;
optional string routing_key = 9;
optional int32 routed_dp_rank = 10;
map<string, string> trace_headers = 11;
}
message GenerateResponse {
repeated int32 output_ids = 1;
map<string, string> meta_info = 2;
bool finished = 3;
}
// ---- Text-based embed (text in, embedding out) ----
message TextEmbedRequest {
string text = 1;
optional string rid = 2;
optional string routing_key = 3;
map<string, string> trace_headers = 4;
}
message TextEmbedResponse {
repeated float embedding = 1;
map<string, string> meta_info = 2;
}
// ---- Tokenized embed (input_ids in, embedding out) ----
message EmbedRequest {
repeated int32 input_ids = 1;
optional string rid = 2;
optional string routing_key = 3;
map<string, string> trace_headers = 4;
}
message EmbedResponse {
repeated float embedding = 1;
map<string, string> meta_info = 2;
}
// ---- Health check ----
message HealthCheckRequest {}
message HealthCheckResponse {
bool healthy = 1;
}
// ---- Model info ----
message GetModelInfoRequest {}
message GetModelInfoResponse {
string model_path = 1;
string json_info = 2;
}
// ---- Server info ----
message GetServerInfoRequest {}
message GetServerInfoResponse {
string json_info = 1;
}
// ---- Abort ----
message AbortRequest {
string rid = 1;
bool abort_all = 2;
}
message AbortResponse {
bool success = 1;
}
// ---- Classify (same internal path as embed, uses EmbeddingReqInput) ----
message ClassifyRequest {
string text = 1;
repeated int32 input_ids = 2;
optional string rid = 3;
optional string routing_key = 4;
map<string, string> trace_headers = 5;
}
message ClassifyResponse {
repeated float embedding = 1;
map<string, string> meta_info = 2;
}
// ---- Tokenize / Detokenize (local ops, no inference) ----
message TokenizeRequest {
string text = 1;
optional bool add_special_tokens = 2;
}
message TokenizeResponse {
repeated int32 tokens = 1;
int32 count = 2;
int32 max_model_len = 3;
string input_text = 4;
}
message DetokenizeRequest {
repeated int32 tokens = 1;
}
message DetokenizeResponse {
string text = 1;
}
// ---- List models ----
message ListModelsRequest {}
message ListModelsResponse {
repeated ModelCard models = 1;
}
message ModelCard {
string id = 1;
string root = 2;
optional string parent = 3;
optional int32 max_model_len = 4;
}
// ---- Get load ----
message GetLoadRequest {
optional int32 dp_rank = 1;
}
message GetLoadResponse {
string json_info = 1;
}
// ---- Flush cache ----
message FlushCacheRequest {}
message FlushCacheResponse {
bool success = 1;
string message = 2;
}
// ---- Pause / Continue generation ----
message PauseGenerationRequest {
string mode = 1;
}
message PauseGenerationResponse {
string message = 1;
}
message ContinueGenerationRequest {}
message ContinueGenerationResponse {
string message = 1;
}
// ---- OpenAI-compatible pass-through messages ----
message OpenAIRequest {
bytes json_body = 1;
map<string, string> trace_headers = 2;
}
message OpenAIStreamChunk {
bytes json_chunk = 1;
bool finished = 2;
}
message OpenAIResponse {
bytes json_body = 1;
int32 status_code = 2;
}
// ---- Admin: Profile ----
message StartProfileRequest {
optional string output_dir = 1;
}
message StartProfileResponse {
string message = 1;
}
message StopProfileRequest {}
message StopProfileResponse {
string message = 1;
}
// ---- Admin: Weight update ----
message UpdateWeightsRequest {
string model_path = 1;
optional string load_format = 2;
}
message UpdateWeightsResponse {
bool success = 1;
string message = 2;
}
+6 -1
View File
@@ -1,5 +1,5 @@
[build-system]
requires = ["setuptools>=61.0", "setuptools-scm>=8.0", "wheel"]
requires = ["setuptools>=61.0", "setuptools-scm>=8.0", "setuptools-rust>=1.10", "wheel"]
build-backend = "setuptools.build_meta"
[project]
@@ -217,5 +217,10 @@ git_describe_command = ["python3", "python/tools/get_version_tag.py"]
# Allow editable installs even when .git metadata is not available.
fallback_version = "0.0.0.dev0"
[[tool.setuptools-rust.ext-modules]]
target = "sglang.srt.grpc._core"
path = "../rust/sglang-grpc/Cargo.toml"
binding = "PyO3"
[tool.kernels.dependencies]
"kernels-community/sgl-flash-attn3" = 1
+3
View File
@@ -27,6 +27,9 @@ def run_server(server_args):
launch_server(server_args)
elif server_args.grpc_mode:
# TODO: Once the native Rust gRPC server starts alongside HTTP in the
# default path below (controlled by SGLANG_ENABLE_GRPC / SGLANG_GRPC_PORT),
# remove this legacy SMG path and the grpc_mode flag.
from sglang.srt.entrypoints.grpc_server import serve_grpc
asyncio.run(serve_grpc(server_args))
+4
View File
@@ -516,6 +516,10 @@ class Envs:
# Encoder receiver selection: http|grpc (used by EPD paths).
SGLANG_ENCODER_MM_RECEIVER_MODE = EnvStr("http")
# Native gRPC server (internal, not yet user-facing)
SGLANG_GRPC_PORT = EnvInt(None)
SGLANG_ENABLE_GRPC = EnvBool(False)
# External models
SGLANG_EXTERNAL_MODEL_PACKAGE = EnvStr("")
SGLANG_EXTERNAL_MM_MODEL_ARCH = EnvStr("")
+28
View File
@@ -999,6 +999,21 @@ class ServerArgs:
envs.SGLANG_SPEC_NAN_DETECTION.set(True)
envs.SGLANG_SPEC_OOB_DETECTION.set(True)
# Native gRPC flags — env-only for now, not exposed as CLI args.
# Set as instance attributes (not dataclass fields) to avoid
# argparse namespace lookup in from_cli_args.
self.enable_grpc = envs.SGLANG_ENABLE_GRPC.get()
grpc_port_env = envs.SGLANG_GRPC_PORT.get()
self.grpc_port = (
grpc_port_env if grpc_port_env is not None else self.port + 10000
)
if not (1 <= self.grpc_port <= 65535):
raise ValueError(
f"SGLANG_GRPC_PORT ({self.grpc_port}) must be between 1 and 65535"
)
def _handle_prefill_delayer_env_compat(self):
if envs.SGLANG_SCHEDULER_DECREASE_PREFILL_IDLE.get():
self.enable_prefill_delayer = True
@@ -6622,6 +6637,19 @@ class ServerArgs:
"When enabling two batch overlap, moe_a2a_backend cannot be 'none'."
)
if (
self.enable_grpc
and self.grpc_port is not None
and self.grpc_port == self.port
):
raise ValueError(
f"SGLANG_GRPC_PORT ({self.grpc_port}) must differ from --port ({self.port})"
)
# TODO: Also validate grpc_port != metrics_http_port and grpc_port != nccl_port
# to avoid opaque bind errors at runtime. Deferred because metrics_http_port
# and nccl_port have dynamic defaults that may not be resolved yet here.
if self.gc_threshold:
if not (1 <= len(self.gc_threshold) <= 3):
raise ValueError(
+39
View File
@@ -0,0 +1,39 @@
[package]
name = "sglang-grpc"
version = "0.1.0"
edition = "2024"
description = "In-process Rust gRPC server for SGLang"
license = "Apache-2.0"
[lib]
name = "_core"
crate-type = ["cdylib"]
[dependencies]
pyo3 = { version = "0.23", features = ["extension-module"] }
tokio = { version = "1", features = ["full"] }
tonic = { version = "0.12", features = ["gzip", "transport"] }
prost = "0.13"
crossbeam-channel = "0.5"
uuid = { version = "1", features = ["v4"] }
tracing = "0.1"
tracing-subscriber = { version = "0.3", features = ["env-filter"] }
serde_json = "1"
tokenizers = { version = "0.21", default-features = false, features = ["onig"] }
tokio-stream = "0.1"
async-stream = "0.3"
[build-dependencies]
tonic-build = "0.12"
[features]
default = ["pyo3/extension-module"]
[profile.release]
opt-level = 2
lto = "thin"
strip = true
[profile.dev]
opt-level = 0
debug = 1
+15
View File
@@ -0,0 +1,15 @@
fn main() -> Result<(), Box<dyn std::error::Error>> {
let proto_path = "../../proto/sglang/runtime/v1/sglang.proto";
tonic_build::configure()
.build_server(true)
.build_client(false)
.file_descriptor_set_path(
std::path::PathBuf::from(std::env::var("OUT_DIR").unwrap())
.join("sglang_descriptor.bin"),
)
.compile_protos(&[proto_path], &["../../proto"])?;
println!("cargo:rerun-if-changed={}", proto_path);
Ok(())
}
+85
View File
@@ -0,0 +1,85 @@
use pyo3::prelude::*;
use std::sync::Arc;
use tokio::sync::Notify;
pub mod proto {
tonic::include_proto!("sglang.runtime.v1");
}
/// Handle returned by `start_server` — used to shut down the gRPC server.
#[pyclass]
pub struct GrpcServerHandle {
shutdown: Arc<Notify>,
join_handle: Option<std::thread::JoinHandle<()>>,
}
#[pymethods]
impl GrpcServerHandle {
/// Signal the server to stop and wait for the background thread to exit.
fn shutdown(&mut self) {
self.shutdown.notify_one();
if let Some(handle) = self.join_handle.take() {
let _ = handle.join();
}
}
/// Returns `true` while the server thread is still running.
fn is_alive(&self) -> bool {
self.join_handle
.as_ref()
.map_or(false, |h| !h.is_finished())
}
}
/// Start the gRPC server in a background thread.
///
/// * `host` bind address (e.g. "0.0.0.0")
/// * `port` listen port
/// * `runtime_handle` Python `RuntimeHandle` object (from `grpc_bridge.py`)
///
/// Returns a `GrpcServerHandle` that can be used to shut the server down.
#[pyfunction]
fn start_server(host: String, port: u16, runtime_handle: PyObject) -> PyResult<GrpcServerHandle> {
let _ = &runtime_handle; // Will be used in Phase 1 PR 2
let shutdown = Arc::new(Notify::new());
let shutdown_clone = shutdown.clone();
let addr_str = format!("{}:{}", host, port);
let addr: std::net::SocketAddr = addr_str
.parse()
.map_err(|e| pyo3::exceptions::PyValueError::new_err(format!("Bad address: {e}")))?;
let join_handle = std::thread::Builder::new()
.name("grpc-server".into())
.spawn(move || {
let rt = tokio::runtime::Builder::new_multi_thread()
.worker_threads(4)
.enable_all()
.build()
.expect("Failed to build Tokio runtime");
rt.block_on(async move {
tracing::info!("gRPC server listening on {}", addr);
// Server implementation will be added in PR 2.
// For now, just wait for shutdown signal.
shutdown_clone.notified().await;
tracing::info!("gRPC server shutting down");
});
})
.map_err(|e| {
pyo3::exceptions::PyRuntimeError::new_err(format!("Failed to spawn thread: {e}"))
})?;
Ok(GrpcServerHandle {
shutdown,
join_handle: Some(join_handle),
})
}
/// Python module exported by the Rust extension.
#[pymodule]
fn _core(m: &Bound<'_, PyModule>) -> PyResult<()> {
m.add_function(wrap_pyfunction!(start_server, m)?)?;
m.add_class::<GrpcServerHandle>()?;
Ok(())
}