From 7235a7fbe9a7a92d4f1ee3b95c8ae97326ad50de Mon Sep 17 00:00:00 2001 From: Simo Lin Date: Fri, 5 Dec 2025 06:51:00 -0800 Subject: [PATCH] [misc] add model arch and type to server info and use it for harmony (#14456) --- python/sglang/srt/entrypoints/grpc_server.py | 12 ++- python/sglang/srt/entrypoints/http_server.py | 7 +- python/sglang/srt/grpc/sglang_scheduler.proto | 1 + .../sglang/srt/grpc/sglang_scheduler_pb2.py | 16 ++-- .../sglang/srt/grpc/sglang_scheduler_pb2.pyi | 6 +- sgl-router/src/core/model_card.rs | 23 +++++ .../workflow/steps/worker_registration.rs | 95 ++++++++++++++++--- sgl-router/src/proto/sglang_scheduler.proto | 1 + sgl-router/src/routers/grpc/client.rs | 8 +- .../src/routers/grpc/harmony/detector.rs | 60 ++++++++++++ sgl-router/src/routers/grpc/router.rs | 10 +- 11 files changed, 205 insertions(+), 34 deletions(-) diff --git a/python/sglang/srt/entrypoints/grpc_server.py b/python/sglang/srt/entrypoints/grpc_server.py index e3beff72e..090f65007 100644 --- a/python/sglang/srt/entrypoints/grpc_server.py +++ b/python/sglang/srt/entrypoints/grpc_server.py @@ -22,6 +22,7 @@ from grpc_health.v1 import health_pb2_grpc from grpc_reflection.v1alpha import reflection import sglang +from sglang.srt.configs.model_config import ModelConfig from sglang.srt.disaggregation.utils import FAKE_BOOTSTRAP_HOST, DisaggregationMode from sglang.srt.grpc import sglang_scheduler_pb2, sglang_scheduler_pb2_grpc from sglang.srt.grpc.grpc_request_manager import GrpcRequestManager @@ -321,7 +322,8 @@ class SGLangSchedulerServicer(sglang_scheduler_pb2_grpc.SglangSchedulerServicer) max_context_length=self.model_info["max_context_length"], vocab_size=self.model_info["vocab_size"], supports_vision=self.model_info["supports_vision"], - model_type=self.model_info["model_type"], + model_type=self.model_info.get("model_type") or "", + architectures=self.model_info.get("architectures") or [], eos_token_ids=self.model_info["eos_token_ids"], pad_token_id=self.model_info["pad_token_id"], bos_token_id=self.model_info["bos_token_id"], @@ -718,7 +720,10 @@ async def serve_grpc( server_args=server_args, ) - # Update model info from scheduler info + # Load model config to get HF config info (same as TokenizerManager does) + model_config = ModelConfig.from_server_args(server_args) + + # Update model info from scheduler info and model config if model_info is None: model_info = { "model_name": server_args.model_path, @@ -727,7 +732,8 @@ async def serve_grpc( ), "vocab_size": scheduler_info.get("vocab_size", 128256), "supports_vision": scheduler_info.get("supports_vision", False), - "model_type": scheduler_info.get("model_type", "transformer"), + "model_type": getattr(model_config.hf_config, "model_type", None), + "architectures": getattr(model_config.hf_config, "architectures", None), "max_req_input_len": scheduler_info.get("max_req_input_len", 8192), "eos_token_ids": scheduler_info.get("eos_token_ids", []), "pad_token_id": scheduler_info.get("pad_token_id", 0), diff --git a/python/sglang/srt/entrypoints/http_server.py b/python/sglang/srt/entrypoints/http_server.py index cf905378d..cf0a3784f 100644 --- a/python/sglang/srt/entrypoints/http_server.py +++ b/python/sglang/srt/entrypoints/http_server.py @@ -497,14 +497,17 @@ async def get_model_info(): @app.get("/model_info") async def model_info(): """Get the model information.""" + model_config = _global_state.tokenizer_manager.model_config result = { "model_path": _global_state.tokenizer_manager.model_path, "tokenizer_path": _global_state.tokenizer_manager.server_args.tokenizer_path, "is_generation": _global_state.tokenizer_manager.is_generation, "preferred_sampling_params": _global_state.tokenizer_manager.server_args.preferred_sampling_params, "weight_version": _global_state.tokenizer_manager.server_args.weight_version, - "has_image_understanding": _global_state.tokenizer_manager.model_config.is_image_understandable_model, - "has_audio_understanding": _global_state.tokenizer_manager.model_config.is_audio_understandable_model, + "has_image_understanding": model_config.is_image_understandable_model, + "has_audio_understanding": model_config.is_audio_understandable_model, + "model_type": getattr(model_config.hf_config, "model_type", None), + "architectures": getattr(model_config.hf_config, "architectures", None), } return result diff --git a/python/sglang/srt/grpc/sglang_scheduler.proto b/python/sglang/srt/grpc/sglang_scheduler.proto index f1d330d43..2c919214b 100644 --- a/python/sglang/srt/grpc/sglang_scheduler.proto +++ b/python/sglang/srt/grpc/sglang_scheduler.proto @@ -427,6 +427,7 @@ message GetModelInfoResponse { int32 pad_token_id = 12; int32 bos_token_id = 13; int32 max_req_input_len = 14; + repeated string architectures = 15; } // Get server information diff --git a/python/sglang/srt/grpc/sglang_scheduler_pb2.py b/python/sglang/srt/grpc/sglang_scheduler_pb2.py index ae57a9710..dc7cbcd75 100644 --- a/python/sglang/srt/grpc/sglang_scheduler_pb2.py +++ b/python/sglang/srt/grpc/sglang_scheduler_pb2.py @@ -29,7 +29,7 @@ from google.protobuf import timestamp_pb2 as google_dot_protobuf_dot_timestamp__ from google.protobuf import struct_pb2 as google_dot_protobuf_dot_struct__pb2 -DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile(b'\n\x16sglang_scheduler.proto\x12\x15sglang.grpc.scheduler\x1a\x1fgoogle/protobuf/timestamp.proto\x1a\x1cgoogle/protobuf/struct.proto\"\xd0\x05\n\x0eSamplingParams\x12\x13\n\x0btemperature\x18\x01 \x01(\x02\x12\r\n\x05top_p\x18\x02 \x01(\x02\x12\r\n\x05top_k\x18\x03 \x01(\x05\x12\r\n\x05min_p\x18\x04 \x01(\x02\x12\x19\n\x11\x66requency_penalty\x18\x05 \x01(\x02\x12\x18\n\x10presence_penalty\x18\x06 \x01(\x02\x12\x1a\n\x12repetition_penalty\x18\x07 \x01(\x02\x12\x1b\n\x0emax_new_tokens\x18\x08 \x01(\x05H\x01\x88\x01\x01\x12\x0c\n\x04stop\x18\t \x03(\t\x12\x16\n\x0estop_token_ids\x18\n \x03(\r\x12\x1b\n\x13skip_special_tokens\x18\x0b \x01(\x08\x12%\n\x1dspaces_between_special_tokens\x18\x0c \x01(\x08\x12\x0f\n\x05regex\x18\r \x01(\tH\x00\x12\x15\n\x0bjson_schema\x18\x0e \x01(\tH\x00\x12\x16\n\x0c\x65\x62nf_grammar\x18\x0f \x01(\tH\x00\x12\x18\n\x0estructural_tag\x18\x10 \x01(\tH\x00\x12\t\n\x01n\x18\x11 \x01(\x05\x12\x16\n\x0emin_new_tokens\x18\x12 \x01(\x05\x12\x12\n\nignore_eos\x18\x13 \x01(\x08\x12\x14\n\x0cno_stop_trim\x18\x14 \x01(\x08\x12\x1c\n\x0fstream_interval\x18\x15 \x01(\x05H\x02\x88\x01\x01\x12H\n\nlogit_bias\x18\x16 \x03(\x0b\x32\x34.sglang.grpc.scheduler.SamplingParams.LogitBiasEntry\x12.\n\rcustom_params\x18\x17 \x01(\x0b\x32\x17.google.protobuf.Struct\x1a\x30\n\x0eLogitBiasEntry\x12\x0b\n\x03key\x18\x01 \x01(\t\x12\r\n\x05value\x18\x02 \x01(\x02:\x02\x38\x01\x42\x0c\n\nconstraintB\x11\n\x0f_max_new_tokensB\x12\n\x10_stream_interval\"]\n\x13\x44isaggregatedParams\x12\x16\n\x0e\x62ootstrap_host\x18\x01 \x01(\t\x12\x16\n\x0e\x62ootstrap_port\x18\x02 \x01(\x05\x12\x16\n\x0e\x62ootstrap_room\x18\x03 \x01(\x05\"\xe2\x04\n\x0fGenerateRequest\x12\x12\n\nrequest_id\x18\x01 \x01(\t\x12\x38\n\ttokenized\x18\x02 \x01(\x0b\x32%.sglang.grpc.scheduler.TokenizedInput\x12:\n\tmm_inputs\x18\x03 \x01(\x0b\x32\'.sglang.grpc.scheduler.MultimodalInputs\x12>\n\x0fsampling_params\x18\x04 \x01(\x0b\x32%.sglang.grpc.scheduler.SamplingParams\x12\x16\n\x0ereturn_logprob\x18\x05 \x01(\x08\x12\x19\n\x11logprob_start_len\x18\x06 \x01(\x05\x12\x18\n\x10top_logprobs_num\x18\x07 \x01(\x05\x12\x19\n\x11token_ids_logprob\x18\x08 \x03(\r\x12\x1c\n\x14return_hidden_states\x18\t \x01(\x08\x12H\n\x14\x64isaggregated_params\x18\n \x01(\x0b\x32*.sglang.grpc.scheduler.DisaggregatedParams\x12\x1e\n\x16\x63ustom_logit_processor\x18\x0b \x01(\t\x12-\n\ttimestamp\x18\x0c \x01(\x0b\x32\x1a.google.protobuf.Timestamp\x12\x13\n\x0blog_metrics\x18\r \x01(\x08\x12\x14\n\x0cinput_embeds\x18\x0e \x03(\x02\x12\x0f\n\x07lora_id\x18\x0f \x01(\t\x12\x1a\n\x12\x64\x61ta_parallel_rank\x18\x10 \x01(\x05\x12\x0e\n\x06stream\x18\x11 \x01(\x08\":\n\x0eTokenizedInput\x12\x15\n\roriginal_text\x18\x01 \x01(\t\x12\x11\n\tinput_ids\x18\x02 \x03(\r\"\xd3\x01\n\x10MultimodalInputs\x12\x12\n\nimage_urls\x18\x01 \x03(\t\x12\x12\n\nvideo_urls\x18\x02 \x03(\t\x12\x12\n\naudio_urls\x18\x03 \x03(\t\x12\x33\n\x12processed_features\x18\x04 \x01(\x0b\x32\x17.google.protobuf.Struct\x12\x12\n\nimage_data\x18\x05 \x03(\x0c\x12\x12\n\nvideo_data\x18\x06 \x03(\x0c\x12\x12\n\naudio_data\x18\x07 \x03(\x0c\x12\x12\n\nmodalities\x18\x08 \x03(\t\"\xe3\x01\n\x10GenerateResponse\x12\x12\n\nrequest_id\x18\x01 \x01(\t\x12;\n\x05\x63hunk\x18\x02 \x01(\x0b\x32*.sglang.grpc.scheduler.GenerateStreamChunkH\x00\x12;\n\x08\x63omplete\x18\x03 \x01(\x0b\x32\'.sglang.grpc.scheduler.GenerateCompleteH\x00\x12\x35\n\x05\x65rror\x18\x04 \x01(\x0b\x32$.sglang.grpc.scheduler.GenerateErrorH\x00\x42\n\n\x08response\"\x95\x02\n\x13GenerateStreamChunk\x12\x11\n\ttoken_ids\x18\x01 \x03(\r\x12\x15\n\rprompt_tokens\x18\x02 \x01(\x05\x12\x19\n\x11\x63ompletion_tokens\x18\x03 \x01(\x05\x12\x15\n\rcached_tokens\x18\x04 \x01(\x05\x12>\n\x0foutput_logprobs\x18\x05 \x01(\x0b\x32%.sglang.grpc.scheduler.OutputLogProbs\x12\x15\n\rhidden_states\x18\x06 \x03(\x02\x12<\n\x0einput_logprobs\x18\x07 \x01(\x0b\x32$.sglang.grpc.scheduler.InputLogProbs\x12\r\n\x05index\x18\x08 \x01(\r\"\x9b\x03\n\x10GenerateComplete\x12\x12\n\noutput_ids\x18\x01 \x03(\r\x12\x15\n\rfinish_reason\x18\x02 \x01(\t\x12\x15\n\rprompt_tokens\x18\x03 \x01(\x05\x12\x19\n\x11\x63ompletion_tokens\x18\x04 \x01(\x05\x12\x15\n\rcached_tokens\x18\x05 \x01(\x05\x12>\n\x0foutput_logprobs\x18\x06 \x01(\x0b\x32%.sglang.grpc.scheduler.OutputLogProbs\x12>\n\x11\x61ll_hidden_states\x18\x07 \x03(\x0b\x32#.sglang.grpc.scheduler.HiddenStates\x12\x1a\n\x10matched_token_id\x18\x08 \x01(\rH\x00\x12\x1a\n\x10matched_stop_str\x18\t \x01(\tH\x00\x12<\n\x0einput_logprobs\x18\n \x01(\x0b\x32$.sglang.grpc.scheduler.InputLogProbs\x12\r\n\x05index\x18\x0b \x01(\rB\x0e\n\x0cmatched_stop\"K\n\rGenerateError\x12\x0f\n\x07message\x18\x01 \x01(\t\x12\x18\n\x10http_status_code\x18\x02 \x01(\t\x12\x0f\n\x07\x64\x65tails\x18\x03 \x01(\t\"u\n\x0eOutputLogProbs\x12\x16\n\x0etoken_logprobs\x18\x01 \x03(\x02\x12\x11\n\ttoken_ids\x18\x02 \x03(\x05\x12\x38\n\x0ctop_logprobs\x18\x03 \x03(\x0b\x32\".sglang.grpc.scheduler.TopLogProbs\"\x9e\x01\n\rInputLogProbs\x12@\n\x0etoken_logprobs\x18\x01 \x03(\x0b\x32(.sglang.grpc.scheduler.InputTokenLogProb\x12\x11\n\ttoken_ids\x18\x02 \x03(\x05\x12\x38\n\x0ctop_logprobs\x18\x03 \x03(\x0b\x32\".sglang.grpc.scheduler.TopLogProbs\"1\n\x11InputTokenLogProb\x12\x12\n\x05value\x18\x01 \x01(\x02H\x00\x88\x01\x01\x42\x08\n\x06_value\"0\n\x0bTopLogProbs\x12\x0e\n\x06values\x18\x01 \x03(\x02\x12\x11\n\ttoken_ids\x18\x02 \x03(\x05\"?\n\x0cHiddenStates\x12\x0e\n\x06values\x18\x01 \x03(\x02\x12\r\n\x05layer\x18\x02 \x01(\x05\x12\x10\n\x08position\x18\x03 \x01(\x05\"\xca\x02\n\x0c\x45mbedRequest\x12\x12\n\nrequest_id\x18\x01 \x01(\t\x12\x38\n\ttokenized\x18\x02 \x01(\x0b\x32%.sglang.grpc.scheduler.TokenizedInput\x12:\n\tmm_inputs\x18\x04 \x01(\x0b\x32\'.sglang.grpc.scheduler.MultimodalInputs\x12>\n\x0fsampling_params\x18\x05 \x01(\x0b\x32%.sglang.grpc.scheduler.SamplingParams\x12\x13\n\x0blog_metrics\x18\x06 \x01(\x08\x12\x16\n\x0etoken_type_ids\x18\x07 \x03(\x05\x12\x1a\n\x12\x64\x61ta_parallel_rank\x18\x08 \x01(\x05\x12\x18\n\x10is_cross_encoder\x18\t \x01(\x08\x12\r\n\x05texts\x18\n \x03(\t\"\x9d\x01\n\rEmbedResponse\x12\x12\n\nrequest_id\x18\x01 \x01(\t\x12\x38\n\x08\x63omplete\x18\x02 \x01(\x0b\x32$.sglang.grpc.scheduler.EmbedCompleteH\x00\x12\x32\n\x05\x65rror\x18\x03 \x01(\x0b\x32!.sglang.grpc.scheduler.EmbedErrorH\x00\x42\n\n\x08response\"\xa3\x01\n\rEmbedComplete\x12\x11\n\tembedding\x18\x01 \x03(\x02\x12\x15\n\rprompt_tokens\x18\x02 \x01(\x05\x12\x15\n\rcached_tokens\x18\x03 \x01(\x05\x12\x15\n\rembedding_dim\x18\x04 \x01(\x05\x12:\n\x10\x62\x61tch_embeddings\x18\x05 \x03(\x0b\x32 .sglang.grpc.scheduler.Embedding\"*\n\tEmbedding\x12\x0e\n\x06values\x18\x01 \x03(\x02\x12\r\n\x05index\x18\x02 \x01(\x05\"<\n\nEmbedError\x12\x0f\n\x07message\x18\x01 \x01(\t\x12\x0c\n\x04\x63ode\x18\x02 \x01(\t\x12\x0f\n\x07\x64\x65tails\x18\x03 \x01(\t\"\x14\n\x12HealthCheckRequest\"7\n\x13HealthCheckResponse\x12\x0f\n\x07healthy\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\"2\n\x0c\x41\x62ortRequest\x12\x12\n\nrequest_id\x18\x01 \x01(\t\x12\x0e\n\x06reason\x18\x02 \x01(\t\"1\n\rAbortResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\"I\n\x0fLoadLoRARequest\x12\x12\n\nadapter_id\x18\x01 \x01(\t\x12\x14\n\x0c\x61\x64\x61pter_path\x18\x02 \x01(\t\x12\x0c\n\x04rank\x18\x03 \x01(\x05\"H\n\x10LoadLoRAResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x12\n\nadapter_id\x18\x02 \x01(\t\x12\x0f\n\x07message\x18\x03 \x01(\t\"\'\n\x11UnloadLoRARequest\x12\x12\n\nadapter_id\x18\x01 \x01(\t\"6\n\x12UnloadLoRAResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\"w\n\x14UpdateWeightsRequest\x12\x13\n\tdisk_path\x18\x01 \x01(\tH\x00\x12\x15\n\x0btensor_data\x18\x02 \x01(\x0cH\x00\x12\x14\n\nremote_url\x18\x03 \x01(\tH\x00\x12\x13\n\x0bweight_name\x18\x04 \x01(\tB\x08\n\x06source\"9\n\x15UpdateWeightsResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\"-\n\x17GetInternalStateRequest\x12\x12\n\nstate_keys\x18\x01 \x03(\t\"B\n\x18GetInternalStateResponse\x12&\n\x05state\x18\x01 \x01(\x0b\x32\x17.google.protobuf.Struct\"A\n\x17SetInternalStateRequest\x12&\n\x05state\x18\x01 \x01(\x0b\x32\x17.google.protobuf.Struct\"<\n\x18SetInternalStateResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\"\x15\n\x13GetModelInfoRequest\"\xea\x02\n\x14GetModelInfoResponse\x12\x12\n\nmodel_path\x18\x01 \x01(\t\x12\x16\n\x0etokenizer_path\x18\x02 \x01(\t\x12\x15\n\ris_generation\x18\x03 \x01(\x08\x12!\n\x19preferred_sampling_params\x18\x04 \x01(\t\x12\x16\n\x0eweight_version\x18\x05 \x01(\t\x12\x19\n\x11served_model_name\x18\x06 \x01(\t\x12\x1a\n\x12max_context_length\x18\x07 \x01(\x05\x12\x12\n\nvocab_size\x18\x08 \x01(\x05\x12\x17\n\x0fsupports_vision\x18\t \x01(\x08\x12\x12\n\nmodel_type\x18\n \x01(\t\x12\x15\n\reos_token_ids\x18\x0b \x03(\x05\x12\x14\n\x0cpad_token_id\x18\x0c \x01(\x05\x12\x14\n\x0c\x62os_token_id\x18\r \x01(\x05\x12\x19\n\x11max_req_input_len\x18\x0e \x01(\x05\"\x16\n\x14GetServerInfoRequest\"\xb7\x02\n\x15GetServerInfoResponse\x12,\n\x0bserver_args\x18\x01 \x01(\x0b\x32\x17.google.protobuf.Struct\x12/\n\x0escheduler_info\x18\x02 \x01(\x0b\x32\x17.google.protobuf.Struct\x12\x17\n\x0f\x61\x63tive_requests\x18\x03 \x01(\x05\x12\x11\n\tis_paused\x18\x04 \x01(\x08\x12\x1e\n\x16last_receive_timestamp\x18\x05 \x01(\x01\x12\x16\n\x0euptime_seconds\x18\x06 \x01(\x01\x12\x16\n\x0esglang_version\x18\x07 \x01(\t\x12\x13\n\x0bserver_type\x18\x08 \x01(\t\x12.\n\nstart_time\x18\t \x01(\x0b\x32\x1a.google.protobuf.Timestamp2\xd3\x04\n\x0fSglangScheduler\x12]\n\x08Generate\x12&.sglang.grpc.scheduler.GenerateRequest\x1a\'.sglang.grpc.scheduler.GenerateResponse0\x01\x12R\n\x05\x45mbed\x12#.sglang.grpc.scheduler.EmbedRequest\x1a$.sglang.grpc.scheduler.EmbedResponse\x12\x64\n\x0bHealthCheck\x12).sglang.grpc.scheduler.HealthCheckRequest\x1a*.sglang.grpc.scheduler.HealthCheckResponse\x12R\n\x05\x41\x62ort\x12#.sglang.grpc.scheduler.AbortRequest\x1a$.sglang.grpc.scheduler.AbortResponse\x12g\n\x0cGetModelInfo\x12*.sglang.grpc.scheduler.GetModelInfoRequest\x1a+.sglang.grpc.scheduler.GetModelInfoResponse\x12j\n\rGetServerInfo\x12+.sglang.grpc.scheduler.GetServerInfoRequest\x1a,.sglang.grpc.scheduler.GetServerInfoResponseb\x06proto3') +DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile(b'\n\x16sglang_scheduler.proto\x12\x15sglang.grpc.scheduler\x1a\x1fgoogle/protobuf/timestamp.proto\x1a\x1cgoogle/protobuf/struct.proto\"\xd0\x05\n\x0eSamplingParams\x12\x13\n\x0btemperature\x18\x01 \x01(\x02\x12\r\n\x05top_p\x18\x02 \x01(\x02\x12\r\n\x05top_k\x18\x03 \x01(\x05\x12\r\n\x05min_p\x18\x04 \x01(\x02\x12\x19\n\x11\x66requency_penalty\x18\x05 \x01(\x02\x12\x18\n\x10presence_penalty\x18\x06 \x01(\x02\x12\x1a\n\x12repetition_penalty\x18\x07 \x01(\x02\x12\x1b\n\x0emax_new_tokens\x18\x08 \x01(\x05H\x01\x88\x01\x01\x12\x0c\n\x04stop\x18\t \x03(\t\x12\x16\n\x0estop_token_ids\x18\n \x03(\r\x12\x1b\n\x13skip_special_tokens\x18\x0b \x01(\x08\x12%\n\x1dspaces_between_special_tokens\x18\x0c \x01(\x08\x12\x0f\n\x05regex\x18\r \x01(\tH\x00\x12\x15\n\x0bjson_schema\x18\x0e \x01(\tH\x00\x12\x16\n\x0c\x65\x62nf_grammar\x18\x0f \x01(\tH\x00\x12\x18\n\x0estructural_tag\x18\x10 \x01(\tH\x00\x12\t\n\x01n\x18\x11 \x01(\x05\x12\x16\n\x0emin_new_tokens\x18\x12 \x01(\x05\x12\x12\n\nignore_eos\x18\x13 \x01(\x08\x12\x14\n\x0cno_stop_trim\x18\x14 \x01(\x08\x12\x1c\n\x0fstream_interval\x18\x15 \x01(\x05H\x02\x88\x01\x01\x12H\n\nlogit_bias\x18\x16 \x03(\x0b\x32\x34.sglang.grpc.scheduler.SamplingParams.LogitBiasEntry\x12.\n\rcustom_params\x18\x17 \x01(\x0b\x32\x17.google.protobuf.Struct\x1a\x30\n\x0eLogitBiasEntry\x12\x0b\n\x03key\x18\x01 \x01(\t\x12\r\n\x05value\x18\x02 \x01(\x02:\x02\x38\x01\x42\x0c\n\nconstraintB\x11\n\x0f_max_new_tokensB\x12\n\x10_stream_interval\"]\n\x13\x44isaggregatedParams\x12\x16\n\x0e\x62ootstrap_host\x18\x01 \x01(\t\x12\x16\n\x0e\x62ootstrap_port\x18\x02 \x01(\x05\x12\x16\n\x0e\x62ootstrap_room\x18\x03 \x01(\x05\"\xe2\x04\n\x0fGenerateRequest\x12\x12\n\nrequest_id\x18\x01 \x01(\t\x12\x38\n\ttokenized\x18\x02 \x01(\x0b\x32%.sglang.grpc.scheduler.TokenizedInput\x12:\n\tmm_inputs\x18\x03 \x01(\x0b\x32\'.sglang.grpc.scheduler.MultimodalInputs\x12>\n\x0fsampling_params\x18\x04 \x01(\x0b\x32%.sglang.grpc.scheduler.SamplingParams\x12\x16\n\x0ereturn_logprob\x18\x05 \x01(\x08\x12\x19\n\x11logprob_start_len\x18\x06 \x01(\x05\x12\x18\n\x10top_logprobs_num\x18\x07 \x01(\x05\x12\x19\n\x11token_ids_logprob\x18\x08 \x03(\r\x12\x1c\n\x14return_hidden_states\x18\t \x01(\x08\x12H\n\x14\x64isaggregated_params\x18\n \x01(\x0b\x32*.sglang.grpc.scheduler.DisaggregatedParams\x12\x1e\n\x16\x63ustom_logit_processor\x18\x0b \x01(\t\x12-\n\ttimestamp\x18\x0c \x01(\x0b\x32\x1a.google.protobuf.Timestamp\x12\x13\n\x0blog_metrics\x18\r \x01(\x08\x12\x14\n\x0cinput_embeds\x18\x0e \x03(\x02\x12\x0f\n\x07lora_id\x18\x0f \x01(\t\x12\x1a\n\x12\x64\x61ta_parallel_rank\x18\x10 \x01(\x05\x12\x0e\n\x06stream\x18\x11 \x01(\x08\":\n\x0eTokenizedInput\x12\x15\n\roriginal_text\x18\x01 \x01(\t\x12\x11\n\tinput_ids\x18\x02 \x03(\r\"\xd3\x01\n\x10MultimodalInputs\x12\x12\n\nimage_urls\x18\x01 \x03(\t\x12\x12\n\nvideo_urls\x18\x02 \x03(\t\x12\x12\n\naudio_urls\x18\x03 \x03(\t\x12\x33\n\x12processed_features\x18\x04 \x01(\x0b\x32\x17.google.protobuf.Struct\x12\x12\n\nimage_data\x18\x05 \x03(\x0c\x12\x12\n\nvideo_data\x18\x06 \x03(\x0c\x12\x12\n\naudio_data\x18\x07 \x03(\x0c\x12\x12\n\nmodalities\x18\x08 \x03(\t\"\xe3\x01\n\x10GenerateResponse\x12\x12\n\nrequest_id\x18\x01 \x01(\t\x12;\n\x05\x63hunk\x18\x02 \x01(\x0b\x32*.sglang.grpc.scheduler.GenerateStreamChunkH\x00\x12;\n\x08\x63omplete\x18\x03 \x01(\x0b\x32\'.sglang.grpc.scheduler.GenerateCompleteH\x00\x12\x35\n\x05\x65rror\x18\x04 \x01(\x0b\x32$.sglang.grpc.scheduler.GenerateErrorH\x00\x42\n\n\x08response\"\x95\x02\n\x13GenerateStreamChunk\x12\x11\n\ttoken_ids\x18\x01 \x03(\r\x12\x15\n\rprompt_tokens\x18\x02 \x01(\x05\x12\x19\n\x11\x63ompletion_tokens\x18\x03 \x01(\x05\x12\x15\n\rcached_tokens\x18\x04 \x01(\x05\x12>\n\x0foutput_logprobs\x18\x05 \x01(\x0b\x32%.sglang.grpc.scheduler.OutputLogProbs\x12\x15\n\rhidden_states\x18\x06 \x03(\x02\x12<\n\x0einput_logprobs\x18\x07 \x01(\x0b\x32$.sglang.grpc.scheduler.InputLogProbs\x12\r\n\x05index\x18\x08 \x01(\r\"\x9b\x03\n\x10GenerateComplete\x12\x12\n\noutput_ids\x18\x01 \x03(\r\x12\x15\n\rfinish_reason\x18\x02 \x01(\t\x12\x15\n\rprompt_tokens\x18\x03 \x01(\x05\x12\x19\n\x11\x63ompletion_tokens\x18\x04 \x01(\x05\x12\x15\n\rcached_tokens\x18\x05 \x01(\x05\x12>\n\x0foutput_logprobs\x18\x06 \x01(\x0b\x32%.sglang.grpc.scheduler.OutputLogProbs\x12>\n\x11\x61ll_hidden_states\x18\x07 \x03(\x0b\x32#.sglang.grpc.scheduler.HiddenStates\x12\x1a\n\x10matched_token_id\x18\x08 \x01(\rH\x00\x12\x1a\n\x10matched_stop_str\x18\t \x01(\tH\x00\x12<\n\x0einput_logprobs\x18\n \x01(\x0b\x32$.sglang.grpc.scheduler.InputLogProbs\x12\r\n\x05index\x18\x0b \x01(\rB\x0e\n\x0cmatched_stop\"K\n\rGenerateError\x12\x0f\n\x07message\x18\x01 \x01(\t\x12\x18\n\x10http_status_code\x18\x02 \x01(\t\x12\x0f\n\x07\x64\x65tails\x18\x03 \x01(\t\"u\n\x0eOutputLogProbs\x12\x16\n\x0etoken_logprobs\x18\x01 \x03(\x02\x12\x11\n\ttoken_ids\x18\x02 \x03(\x05\x12\x38\n\x0ctop_logprobs\x18\x03 \x03(\x0b\x32\".sglang.grpc.scheduler.TopLogProbs\"\x9e\x01\n\rInputLogProbs\x12@\n\x0etoken_logprobs\x18\x01 \x03(\x0b\x32(.sglang.grpc.scheduler.InputTokenLogProb\x12\x11\n\ttoken_ids\x18\x02 \x03(\x05\x12\x38\n\x0ctop_logprobs\x18\x03 \x03(\x0b\x32\".sglang.grpc.scheduler.TopLogProbs\"1\n\x11InputTokenLogProb\x12\x12\n\x05value\x18\x01 \x01(\x02H\x00\x88\x01\x01\x42\x08\n\x06_value\"0\n\x0bTopLogProbs\x12\x0e\n\x06values\x18\x01 \x03(\x02\x12\x11\n\ttoken_ids\x18\x02 \x03(\x05\"?\n\x0cHiddenStates\x12\x0e\n\x06values\x18\x01 \x03(\x02\x12\r\n\x05layer\x18\x02 \x01(\x05\x12\x10\n\x08position\x18\x03 \x01(\x05\"\xca\x02\n\x0c\x45mbedRequest\x12\x12\n\nrequest_id\x18\x01 \x01(\t\x12\x38\n\ttokenized\x18\x02 \x01(\x0b\x32%.sglang.grpc.scheduler.TokenizedInput\x12:\n\tmm_inputs\x18\x04 \x01(\x0b\x32\'.sglang.grpc.scheduler.MultimodalInputs\x12>\n\x0fsampling_params\x18\x05 \x01(\x0b\x32%.sglang.grpc.scheduler.SamplingParams\x12\x13\n\x0blog_metrics\x18\x06 \x01(\x08\x12\x16\n\x0etoken_type_ids\x18\x07 \x03(\x05\x12\x1a\n\x12\x64\x61ta_parallel_rank\x18\x08 \x01(\x05\x12\x18\n\x10is_cross_encoder\x18\t \x01(\x08\x12\r\n\x05texts\x18\n \x03(\t\"\x9d\x01\n\rEmbedResponse\x12\x12\n\nrequest_id\x18\x01 \x01(\t\x12\x38\n\x08\x63omplete\x18\x02 \x01(\x0b\x32$.sglang.grpc.scheduler.EmbedCompleteH\x00\x12\x32\n\x05\x65rror\x18\x03 \x01(\x0b\x32!.sglang.grpc.scheduler.EmbedErrorH\x00\x42\n\n\x08response\"\xa3\x01\n\rEmbedComplete\x12\x11\n\tembedding\x18\x01 \x03(\x02\x12\x15\n\rprompt_tokens\x18\x02 \x01(\x05\x12\x15\n\rcached_tokens\x18\x03 \x01(\x05\x12\x15\n\rembedding_dim\x18\x04 \x01(\x05\x12:\n\x10\x62\x61tch_embeddings\x18\x05 \x03(\x0b\x32 .sglang.grpc.scheduler.Embedding\"*\n\tEmbedding\x12\x0e\n\x06values\x18\x01 \x03(\x02\x12\r\n\x05index\x18\x02 \x01(\x05\"<\n\nEmbedError\x12\x0f\n\x07message\x18\x01 \x01(\t\x12\x0c\n\x04\x63ode\x18\x02 \x01(\t\x12\x0f\n\x07\x64\x65tails\x18\x03 \x01(\t\"\x14\n\x12HealthCheckRequest\"7\n\x13HealthCheckResponse\x12\x0f\n\x07healthy\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\"2\n\x0c\x41\x62ortRequest\x12\x12\n\nrequest_id\x18\x01 \x01(\t\x12\x0e\n\x06reason\x18\x02 \x01(\t\"1\n\rAbortResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\"I\n\x0fLoadLoRARequest\x12\x12\n\nadapter_id\x18\x01 \x01(\t\x12\x14\n\x0c\x61\x64\x61pter_path\x18\x02 \x01(\t\x12\x0c\n\x04rank\x18\x03 \x01(\x05\"H\n\x10LoadLoRAResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x12\n\nadapter_id\x18\x02 \x01(\t\x12\x0f\n\x07message\x18\x03 \x01(\t\"\'\n\x11UnloadLoRARequest\x12\x12\n\nadapter_id\x18\x01 \x01(\t\"6\n\x12UnloadLoRAResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\"w\n\x14UpdateWeightsRequest\x12\x13\n\tdisk_path\x18\x01 \x01(\tH\x00\x12\x15\n\x0btensor_data\x18\x02 \x01(\x0cH\x00\x12\x14\n\nremote_url\x18\x03 \x01(\tH\x00\x12\x13\n\x0bweight_name\x18\x04 \x01(\tB\x08\n\x06source\"9\n\x15UpdateWeightsResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\"-\n\x17GetInternalStateRequest\x12\x12\n\nstate_keys\x18\x01 \x03(\t\"B\n\x18GetInternalStateResponse\x12&\n\x05state\x18\x01 \x01(\x0b\x32\x17.google.protobuf.Struct\"A\n\x17SetInternalStateRequest\x12&\n\x05state\x18\x01 \x01(\x0b\x32\x17.google.protobuf.Struct\"<\n\x18SetInternalStateResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x0f\n\x07message\x18\x02 \x01(\t\"\x15\n\x13GetModelInfoRequest\"\x81\x03\n\x14GetModelInfoResponse\x12\x12\n\nmodel_path\x18\x01 \x01(\t\x12\x16\n\x0etokenizer_path\x18\x02 \x01(\t\x12\x15\n\ris_generation\x18\x03 \x01(\x08\x12!\n\x19preferred_sampling_params\x18\x04 \x01(\t\x12\x16\n\x0eweight_version\x18\x05 \x01(\t\x12\x19\n\x11served_model_name\x18\x06 \x01(\t\x12\x1a\n\x12max_context_length\x18\x07 \x01(\x05\x12\x12\n\nvocab_size\x18\x08 \x01(\x05\x12\x17\n\x0fsupports_vision\x18\t \x01(\x08\x12\x12\n\nmodel_type\x18\n \x01(\t\x12\x15\n\reos_token_ids\x18\x0b \x03(\x05\x12\x14\n\x0cpad_token_id\x18\x0c \x01(\x05\x12\x14\n\x0c\x62os_token_id\x18\r \x01(\x05\x12\x19\n\x11max_req_input_len\x18\x0e \x01(\x05\x12\x15\n\rarchitectures\x18\x0f \x03(\t\"\x16\n\x14GetServerInfoRequest\"\xb7\x02\n\x15GetServerInfoResponse\x12,\n\x0bserver_args\x18\x01 \x01(\x0b\x32\x17.google.protobuf.Struct\x12/\n\x0escheduler_info\x18\x02 \x01(\x0b\x32\x17.google.protobuf.Struct\x12\x17\n\x0f\x61\x63tive_requests\x18\x03 \x01(\x05\x12\x11\n\tis_paused\x18\x04 \x01(\x08\x12\x1e\n\x16last_receive_timestamp\x18\x05 \x01(\x01\x12\x16\n\x0euptime_seconds\x18\x06 \x01(\x01\x12\x16\n\x0esglang_version\x18\x07 \x01(\t\x12\x13\n\x0bserver_type\x18\x08 \x01(\t\x12.\n\nstart_time\x18\t \x01(\x0b\x32\x1a.google.protobuf.Timestamp2\xd3\x04\n\x0fSglangScheduler\x12]\n\x08Generate\x12&.sglang.grpc.scheduler.GenerateRequest\x1a\'.sglang.grpc.scheduler.GenerateResponse0\x01\x12R\n\x05\x45mbed\x12#.sglang.grpc.scheduler.EmbedRequest\x1a$.sglang.grpc.scheduler.EmbedResponse\x12\x64\n\x0bHealthCheck\x12).sglang.grpc.scheduler.HealthCheckRequest\x1a*.sglang.grpc.scheduler.HealthCheckResponse\x12R\n\x05\x41\x62ort\x12#.sglang.grpc.scheduler.AbortRequest\x1a$.sglang.grpc.scheduler.AbortResponse\x12g\n\x0cGetModelInfo\x12*.sglang.grpc.scheduler.GetModelInfoRequest\x1a+.sglang.grpc.scheduler.GetModelInfoResponse\x12j\n\rGetServerInfo\x12+.sglang.grpc.scheduler.GetServerInfoRequest\x1a,.sglang.grpc.scheduler.GetServerInfoResponseb\x06proto3') _globals = globals() _builder.BuildMessageAndEnumDescriptors(DESCRIPTOR, _globals) @@ -109,11 +109,11 @@ if not _descriptor._USE_C_DESCRIPTORS: _globals['_GETMODELINFOREQUEST']._serialized_start=4881 _globals['_GETMODELINFOREQUEST']._serialized_end=4902 _globals['_GETMODELINFORESPONSE']._serialized_start=4905 - _globals['_GETMODELINFORESPONSE']._serialized_end=5267 - _globals['_GETSERVERINFOREQUEST']._serialized_start=5269 - _globals['_GETSERVERINFOREQUEST']._serialized_end=5291 - _globals['_GETSERVERINFORESPONSE']._serialized_start=5294 - _globals['_GETSERVERINFORESPONSE']._serialized_end=5605 - _globals['_SGLANGSCHEDULER']._serialized_start=5608 - _globals['_SGLANGSCHEDULER']._serialized_end=6203 + _globals['_GETMODELINFORESPONSE']._serialized_end=5290 + _globals['_GETSERVERINFOREQUEST']._serialized_start=5292 + _globals['_GETSERVERINFOREQUEST']._serialized_end=5314 + _globals['_GETSERVERINFORESPONSE']._serialized_start=5317 + _globals['_GETSERVERINFORESPONSE']._serialized_end=5628 + _globals['_SGLANGSCHEDULER']._serialized_start=5631 + _globals['_SGLANGSCHEDULER']._serialized_end=6226 # @@protoc_insertion_point(module_scope) diff --git a/python/sglang/srt/grpc/sglang_scheduler_pb2.pyi b/python/sglang/srt/grpc/sglang_scheduler_pb2.pyi index 034dd653c..958a06269 100644 --- a/python/sglang/srt/grpc/sglang_scheduler_pb2.pyi +++ b/python/sglang/srt/grpc/sglang_scheduler_pb2.pyi @@ -432,7 +432,7 @@ class GetModelInfoRequest(_message.Message): def __init__(self) -> None: ... class GetModelInfoResponse(_message.Message): - __slots__ = ("model_path", "tokenizer_path", "is_generation", "preferred_sampling_params", "weight_version", "served_model_name", "max_context_length", "vocab_size", "supports_vision", "model_type", "eos_token_ids", "pad_token_id", "bos_token_id", "max_req_input_len") + __slots__ = ("model_path", "tokenizer_path", "is_generation", "preferred_sampling_params", "weight_version", "served_model_name", "max_context_length", "vocab_size", "supports_vision", "model_type", "eos_token_ids", "pad_token_id", "bos_token_id", "max_req_input_len", "architectures") MODEL_PATH_FIELD_NUMBER: _ClassVar[int] TOKENIZER_PATH_FIELD_NUMBER: _ClassVar[int] IS_GENERATION_FIELD_NUMBER: _ClassVar[int] @@ -447,6 +447,7 @@ class GetModelInfoResponse(_message.Message): PAD_TOKEN_ID_FIELD_NUMBER: _ClassVar[int] BOS_TOKEN_ID_FIELD_NUMBER: _ClassVar[int] MAX_REQ_INPUT_LEN_FIELD_NUMBER: _ClassVar[int] + ARCHITECTURES_FIELD_NUMBER: _ClassVar[int] model_path: str tokenizer_path: str is_generation: bool @@ -461,7 +462,8 @@ class GetModelInfoResponse(_message.Message): pad_token_id: int bos_token_id: int max_req_input_len: int - def __init__(self, model_path: _Optional[str] = ..., tokenizer_path: _Optional[str] = ..., is_generation: bool = ..., preferred_sampling_params: _Optional[str] = ..., weight_version: _Optional[str] = ..., served_model_name: _Optional[str] = ..., max_context_length: _Optional[int] = ..., vocab_size: _Optional[int] = ..., supports_vision: bool = ..., model_type: _Optional[str] = ..., eos_token_ids: _Optional[_Iterable[int]] = ..., pad_token_id: _Optional[int] = ..., bos_token_id: _Optional[int] = ..., max_req_input_len: _Optional[int] = ...) -> None: ... + architectures: _containers.RepeatedScalarFieldContainer[str] + def __init__(self, model_path: _Optional[str] = ..., tokenizer_path: _Optional[str] = ..., is_generation: bool = ..., preferred_sampling_params: _Optional[str] = ..., weight_version: _Optional[str] = ..., served_model_name: _Optional[str] = ..., max_context_length: _Optional[int] = ..., vocab_size: _Optional[int] = ..., supports_vision: bool = ..., model_type: _Optional[str] = ..., eos_token_ids: _Optional[_Iterable[int]] = ..., pad_token_id: _Optional[int] = ..., bos_token_id: _Optional[int] = ..., max_req_input_len: _Optional[int] = ..., architectures: _Optional[_Iterable[str]] = ...) -> None: ... class GetServerInfoRequest(_message.Message): __slots__ = () diff --git a/sgl-router/src/core/model_card.rs b/sgl-router/src/core/model_card.rs index f1d31bcb7..eee56a54c 100644 --- a/sgl-router/src/core/model_card.rs +++ b/sgl-router/src/core/model_card.rs @@ -123,6 +123,15 @@ pub struct ModelCard { #[serde(default = "default_model_type")] pub model_type: ModelType, + /// HuggingFace model type string (e.g., "llama", "qwen2", "gpt-oss") + /// This is different from `model_type` which is capability bitflags. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub hf_model_type: Option, + + /// Model architectures from HuggingFace config (e.g., ["LlamaForCausalLM"]) + #[serde(default, skip_serializing_if = "Vec::is_empty")] + pub architectures: Vec, + /// Provider hint for API transformations. /// `None` means native/passthrough (no transformation needed). #[serde(default, skip_serializing_if = "Option::is_none")] @@ -172,6 +181,8 @@ impl ModelCard { display_name: None, aliases: Vec::new(), model_type: ModelType::LLM, + hf_model_type: None, + architectures: Vec::new(), provider: None, context_length: None, tokenizer_path: None, @@ -208,6 +219,18 @@ impl ModelCard { self } + /// Set the HuggingFace model type string + pub fn with_hf_model_type(mut self, hf_model_type: impl Into) -> Self { + self.hf_model_type = Some(hf_model_type.into()); + self + } + + /// Set the model architectures + pub fn with_architectures(mut self, architectures: Vec) -> Self { + self.architectures = architectures; + self + } + /// Set the provider type (for external API transformations) pub fn with_provider(mut self, provider: ProviderType) -> Self { self.provider = Some(provider); diff --git a/sgl-router/src/core/workflow/steps/worker_registration.rs b/sgl-router/src/core/workflow/steps/worker_registration.rs index c8b4a09ca..e59204eb2 100644 --- a/sgl-router/src/core/workflow/steps/worker_registration.rs +++ b/sgl-router/src/core/workflow/steps/worker_registration.rs @@ -55,6 +55,18 @@ struct ServerInfo { max_num_reqs: Option, } +/// Model information returned from /model_info endpoint +#[derive(Debug, Clone, Deserialize, Serialize)] +struct ModelInfo { + model_path: Option, + tokenizer_path: Option, + is_generation: Option, + /// HuggingFace model type string (e.g., "llama", "qwen2", "gpt_oss") + model_type: Option, + /// Model architectures from HuggingFace config (e.g., ["LlamaForCausalLM"]) + architectures: Option>, +} + #[derive(Debug, Clone)] pub struct DpInfo { pub dp_size: usize, @@ -69,7 +81,7 @@ fn parse_server_info(json: Value) -> Result { /// Get server info from /get_server_info endpoint async fn get_server_info(url: &str, api_key: Option<&str>) -> Result { let base_url = url.trim_end_matches('/'); - let server_info_url = format!("{}/get_server_info", base_url); + let server_info_url = format!("{}/server_info", base_url); let mut req = HTTP_CLIENT.get(&server_info_url); if let Some(key) = api_key { @@ -97,6 +109,35 @@ async fn get_server_info(url: &str, api_key: Option<&str>) -> Result) -> Result { + let base_url = url.trim_end_matches('/'); + let model_info_url = format!("{}/model_info", base_url); + + let mut req = HTTP_CLIENT.get(&model_info_url); + if let Some(key) = api_key { + req = req.bearer_auth(key); + } + + let response = req + .send() + .await + .map_err(|e| format!("Failed to connect to {}: {}", model_info_url, e))?; + + if !response.status().is_success() { + return Err(format!( + "Server returned status {} from {}", + response.status(), + model_info_url + )); + } + + response + .json::() + .await + .map_err(|e| format!("Failed to parse response from {}: {}", model_info_url, e)) +} + /// Get DP info for a worker URL async fn get_dp_info(url: &str, api_key: Option<&str>) -> Result { let info = get_server_info(url, api_key).await?; @@ -319,21 +360,37 @@ impl StepExecutor for DiscoverMetadataStep { let (discovered_labels, detected_runtime) = match connection_mode.as_ref() { ConnectionMode::Http => { - match get_server_info(&config.url, config.api_key.as_deref()).await { - Ok(server_info) => { - let mut labels = HashMap::new(); - if let Some(model_path) = server_info.model_path.filter(|s| !s.is_empty()) { - labels.insert("model_path".to_string(), model_path); - } - if let Some(served_model_name) = - server_info.served_model_name.filter(|s| !s.is_empty()) - { - labels.insert("served_model_name".to_string(), served_model_name); - } - Ok((labels, None)) + let mut labels = HashMap::new(); + + // Fetch from /get_server_info for server-related metadata + if let Ok(server_info) = + get_server_info(&config.url, config.api_key.as_deref()).await + { + if let Some(model_path) = server_info.model_path.filter(|s| !s.is_empty()) { + labels.insert("model_path".to_string(), model_path); + } + if let Some(served_model_name) = + server_info.served_model_name.filter(|s| !s.is_empty()) + { + labels.insert("served_model_name".to_string(), served_model_name); } - Err(e) => Err(e), } + + // Fetch from /model_info for model-related metadata (model_type, architectures) + if let Ok(model_info) = get_model_info(&config.url, config.api_key.as_deref()).await + { + if let Some(model_type) = model_info.model_type.filter(|s| !s.is_empty()) { + labels.insert("model_type".to_string(), model_type); + } + if let Some(architectures) = model_info.architectures.filter(|a| !a.is_empty()) + { + if let Ok(json_str) = serde_json::to_string(&architectures) { + labels.insert("architectures".to_string(), json_str); + } + } + } + + Ok((labels, None)) } ConnectionMode::Grpc { .. } => { let runtime_type = config.runtime.as_deref(); @@ -481,6 +538,16 @@ impl StepExecutor for CreateWorkerStep { if let Some(ref chat_template) = config.chat_template { card = card.with_chat_template(chat_template.clone()); } + // Set HuggingFace model type from discovered labels + if let Some(model_type_str) = final_labels.get("model_type") { + card = card.with_hf_model_type(model_type_str.clone()); + } + // Set architectures from discovered labels (JSON array string) + if let Some(architectures_json) = final_labels.get("architectures") { + if let Ok(architectures) = serde_json::from_str::>(architectures_json) { + card = card.with_architectures(architectures); + } + } card }; diff --git a/sgl-router/src/proto/sglang_scheduler.proto b/sgl-router/src/proto/sglang_scheduler.proto index f1d330d43..2c919214b 100644 --- a/sgl-router/src/proto/sglang_scheduler.proto +++ b/sgl-router/src/proto/sglang_scheduler.proto @@ -427,6 +427,7 @@ message GetModelInfoResponse { int32 pad_token_id = 12; int32 bos_token_id = 13; int32 max_req_input_len = 14; + repeated string architectures = 15; } // Get server information diff --git a/sgl-router/src/routers/grpc/client.rs b/sgl-router/src/routers/grpc/client.rs index 79f1e084f..4d6b76a75 100644 --- a/sgl-router/src/routers/grpc/client.rs +++ b/sgl-router/src/routers/grpc/client.rs @@ -166,7 +166,13 @@ impl ModelInfo { serde_json::Value::Bool(true) => { labels.insert(key, "true".to_string()); } - // Skip empty strings, zeros, false, nulls, arrays, objects + // Insert non-empty arrays as JSON strings (for architectures, etc.) + serde_json::Value::Array(arr) if !arr.is_empty() => { + if let Ok(json_str) = serde_json::to_string(&arr) { + labels.insert(key, json_str); + } + } + // Skip empty strings, zeros, false, nulls, empty arrays, objects _ => {} } } diff --git a/sgl-router/src/routers/grpc/harmony/detector.rs b/sgl-router/src/routers/grpc/harmony/detector.rs index 88ef6a6a5..18fcf4b6e 100644 --- a/sgl-router/src/routers/grpc/harmony/detector.rs +++ b/sgl-router/src/routers/grpc/harmony/detector.rs @@ -1,11 +1,48 @@ //! Harmony model detection +use crate::core::{Worker, WorkerRegistry}; + /// Harmony model detector /// /// Detects if a model name indicates support for Harmony encoding/parsing. pub struct HarmonyDetector; impl HarmonyDetector { + /// Check if a worker is a Harmony/GPT-OSS model. + /// + /// Detection priority: + /// 1. Check if any model card has architectures containing "GptOssForCausalLM" + /// 2. Check if any model card has hf_model_type equal to "gpt_oss" + /// 3. Check if model_id contains "gpt-oss" substring (case-insensitive) + pub fn is_harmony_worker(worker: &dyn Worker) -> bool { + for model_card in worker.models() { + // 1. Check architectures for GptOssForCausalLM + if model_card + .architectures + .iter() + .any(|arch| arch == "GptOssForCausalLM") + { + return true; + } + + // 2. Check hf_model_type for gpt_oss + if let Some(ref model_type) = model_card.hf_model_type { + if model_type == "gpt_oss" { + return true; + } + } + + // 3. Check model id for gpt-oss substring + if Self::is_harmony_model(&model_card.id) { + return true; + } + } + + // Fallback: check worker's model_id directly + Self::is_harmony_model(worker.model_id()) + } + + /// Check if a model name contains "gpt-oss" (case-insensitive). pub fn is_harmony_model(model_name: &str) -> bool { // Case-insensitive substring search without heap allocation // More efficient than to_lowercase() which allocates a new String @@ -14,4 +51,27 @@ impl HarmonyDetector { .windows(7) // "gpt-oss".len() .any(|window| window.eq_ignore_ascii_case(b"gpt-oss")) } + + /// Check if any worker for the given model is a Harmony/GPT-OSS worker. + /// + /// This method looks up workers from the registry by model name and checks + /// if any of them are Harmony workers based on their metadata (architectures, + /// hf_model_type). + /// + /// Falls back to string-based detection if no workers are registered for + /// the model (e.g., during startup before workers are discovered). + pub fn is_harmony_model_in_registry(registry: &WorkerRegistry, model_name: &str) -> bool { + // Get workers for this model + let workers = registry.get_by_model_fast(model_name); + + if workers.is_empty() { + // No workers found - fall back to string-based detection + return Self::is_harmony_model(model_name); + } + + // Check if any worker is a Harmony worker + workers + .iter() + .any(|worker| Self::is_harmony_worker(worker.as_ref())) + } } diff --git a/sgl-router/src/routers/grpc/router.rs b/sgl-router/src/routers/grpc/router.rs index da3b89ae2..ad0bcef0b 100644 --- a/sgl-router/src/routers/grpc/router.rs +++ b/sgl-router/src/routers/grpc/router.rs @@ -142,8 +142,9 @@ impl GrpcRouter { body: &ChatCompletionRequest, model_id: Option<&str>, ) -> Response { - // Choose Harmony pipeline if model indicates Harmony - let is_harmony = HarmonyDetector::is_harmony_model(&body.model); + // Choose Harmony pipeline if workers indicate Harmony (checks architectures, hf_model_type) + let is_harmony = + HarmonyDetector::is_harmony_model_in_registry(&self.worker_registry, &body.model); debug!( "Processing chat completion request for model: {:?}, using_harmony={}", @@ -203,8 +204,9 @@ impl GrpcRouter { return error_response; } - // Choose implementation based on Harmony model detection - let is_harmony = HarmonyDetector::is_harmony_model(&body.model); + // Choose implementation based on Harmony model detection (checks worker metadata) + let is_harmony = + HarmonyDetector::is_harmony_model_in_registry(&self.worker_registry, &body.model); if is_harmony { debug!(