Tiny support http headers in bench serving (#16606)

This commit is contained in:
fzyzcjy
2026-01-07 10:15:17 +08:00
committed by GitHub
parent 9a21d89c5b
commit d874c8bba4
3 changed files with 109 additions and 5 deletions
+22 -4
View File
@@ -132,6 +132,17 @@ def get_auth_headers() -> Dict[str, str]:
return {} return {}
def parse_custom_headers(header_list: List[str]) -> Dict[str, str]:
return {k: v for h in header_list for k, _, v in [h.partition("=")] if k and v}
def get_request_headers() -> Dict[str, str]:
headers = get_auth_headers()
if h := getattr(args, "header", None):
headers.update(parse_custom_headers(h))
return headers
# trt llm does not support ignore_eos # trt llm does not support ignore_eos
# https://github.com/triton-inference-server/tensorrtllm_backend/issues/505 # https://github.com/triton-inference-server/tensorrtllm_backend/issues/505
async def async_request_trt_llm( async def async_request_trt_llm(
@@ -236,7 +247,7 @@ async def async_request_openai_completions(
if request_func_input.image_data: if request_func_input.image_data:
payload.update({"image_data": request_func_input.image_data}) payload.update({"image_data": request_func_input.image_data})
headers = get_auth_headers() headers = get_request_headers()
if request_func_input.routing_key: if request_func_input.routing_key:
headers[_ROUTING_KEY_HEADER] = request_func_input.routing_key headers[_ROUTING_KEY_HEADER] = request_func_input.routing_key
@@ -377,7 +388,7 @@ async def async_request_openai_chat_completions(
payload["model"] = request_func_input.lora_name payload["model"] = request_func_input.lora_name
payload["lora_path"] = request_func_input.lora_name payload["lora_path"] = request_func_input.lora_name
headers = get_auth_headers() headers = get_request_headers()
if request_func_input.routing_key: if request_func_input.routing_key:
headers[_ROUTING_KEY_HEADER] = request_func_input.routing_key headers[_ROUTING_KEY_HEADER] = request_func_input.routing_key
@@ -495,7 +506,7 @@ async def async_request_truss(
"ignore_eos": not args.disable_ignore_eos, "ignore_eos": not args.disable_ignore_eos,
**request_func_input.extra_request_body, **request_func_input.extra_request_body,
} }
headers = get_auth_headers() headers = get_request_headers()
output = RequestFuncOutput.init_new(request_func_input) output = RequestFuncOutput.init_new(request_func_input)
@@ -583,7 +594,7 @@ async def async_request_sglang_generate(
if request_func_input.image_data: if request_func_input.image_data:
payload["image_data"] = request_func_input.image_data payload["image_data"] = request_func_input.image_data
headers = get_auth_headers() headers = get_request_headers()
if request_func_input.routing_key: if request_func_input.routing_key:
headers[_ROUTING_KEY_HEADER] = request_func_input.routing_key headers[_ROUTING_KEY_HEADER] = request_func_input.routing_key
@@ -3228,5 +3239,12 @@ if __name__ == "__main__":
parser.add_argument( parser.add_argument(
"--tag", type=str, default=None, help="The tag to be dumped to output." "--tag", type=str, default=None, help="The tag to be dumped to output."
) )
parser.add_argument(
"--header",
type=str,
nargs="+",
default=None,
help="Custom HTTP headers in Key=Value format. Example: --header MyHeader=MY_VALUE MyAnotherHeader=myanothervalue",
)
args = parser.parse_args() args = parser.parse_args()
run_benchmark(args) run_benchmark(args)
+2
View File
@@ -778,6 +778,7 @@ def get_benchmark_args(
gsp_question_len=32, gsp_question_len=32,
gsp_output_len=32, gsp_output_len=32,
gsp_num_turns=1, gsp_num_turns=1,
header=None,
): ):
return SimpleNamespace( return SimpleNamespace(
backend=backend, backend=backend,
@@ -818,6 +819,7 @@ def get_benchmark_args(
gsp_question_len=gsp_question_len, gsp_question_len=gsp_question_len,
gsp_output_len=gsp_output_len, gsp_output_len=gsp_output_len,
gsp_num_turns=gsp_num_turns, gsp_num_turns=gsp_num_turns,
header=header,
) )
@@ -1,10 +1,12 @@
import json import json
import tempfile import tempfile
import threading
import time import time
import unittest import unittest
from http.server import BaseHTTPRequestHandler, HTTPServer
from pathlib import Path from pathlib import Path
from sglang.bench_serving import run_benchmark from sglang.bench_serving import parse_custom_headers, run_benchmark
from sglang.srt.utils import kill_process_tree from sglang.srt.utils import kill_process_tree
from sglang.test.ci.ci_register import register_cuda_ci from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.test_utils import ( from sglang.test.test_utils import (
@@ -96,5 +98,87 @@ class TestBenchServingFunctionality(CustomTestCase):
) )
class TestBenchServingCustomHeaders(CustomTestCase):
def test_parse_custom_headers(self):
headers = parse_custom_headers(["MyHeader=MY_VALUE", "Another=value=hello"])
self.assertEqual(headers, {"MyHeader": "MY_VALUE", "Another": "value=hello"})
headers = parse_custom_headers(["InvalidNoEquals"])
self.assertEqual(headers, {})
headers = parse_custom_headers(["=NoKey"])
self.assertEqual(headers, {})
# TODO: Using well-implemented mock server, e.g. the on in sgl-router
def test_custom_headers_sent_to_server(self):
import queue
received_requests = queue.Queue()
class HeaderEchoHandler(BaseHTTPRequestHandler):
def _handle(self):
received_requests.put(
{
"method": self.command,
"path": self.path,
"headers": dict(self.headers),
}
)
self.send_response(200)
self.send_header("Content-Type", "application/json")
self.end_headers()
if self.path == "/v1/models":
self.wfile.write(json.dumps({"data": [{"id": "gpt2"}]}).encode())
elif self.path == "/generate":
self.wfile.write(
json.dumps(
{"text": "ok", "meta_info": {"completion_tokens": 1}}
).encode()
)
else:
self.wfile.write(json.dumps({}).encode())
do_GET = do_POST = _handle
server = HTTPServer(("127.0.0.1", 0), HeaderEchoHandler)
port = server.server_address[1]
server_thread = threading.Thread(target=server.serve_forever)
server_thread.daemon = True
server_thread.start()
try:
args = get_benchmark_args(
base_url=f"http://127.0.0.1:{port}",
backend="sglang",
dataset_name="random",
tokenizer="gpt2",
num_prompts=1,
random_input_len=8,
random_output_len=8,
header=["X-Custom-Test=TestValue123", "X-Another=AnotherVal"],
)
args.warmup_requests = 0
args.disable_tqdm = True
run_benchmark(args)
except Exception:
pass
finally:
server.shutdown()
all_reqs = []
while not received_requests.empty():
all_reqs.append(received_requests.get_nowait())
generate_reqs = [r for r in all_reqs if r["path"] == "/generate"]
self.assertGreater(
len(generate_reqs),
0,
f"No /generate request. All: {[r['path'] for r in all_reqs]}",
)
headers = generate_reqs[0]["headers"]
self.assertEqual(headers.get("X-Custom-Test"), "TestValue123")
self.assertEqual(headers.get("X-Another"), "AnotherVal")
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()