diff --git a/benchmarks/backend_request_func.py b/benchmarks/backend_request_func.py index 066d18d794d..c42af6c7147 100644 --- a/benchmarks/backend_request_func.py +++ b/benchmarks/backend_request_func.py @@ -28,6 +28,7 @@ import uuid from dataclasses import dataclass, field from typing import Optional +from urllib.parse import urlsplit, urlunsplit import aiohttp from tqdm.asyncio import tqdm @@ -35,6 +36,20 @@ AIOHTTP_TIMEOUT = aiohttp.ClientTimeout(total=int(os.environ.get("AIOHTTP_TIMEOUT", 6 * 60 * 60))) +def _get_control_url(api_url: str, endpoint: str) -> str: + parsed = urlsplit(api_url) + return urlunsplit((parsed.scheme, parsed.netloc, endpoint, "", "")) + + +async def _close_session(session: aiohttp.ClientSession, api_url: str, session_id: str) -> None: + url = _get_control_url(api_url, "/close_session") + async with session.post(url=url, json={"session_id": session_id}) as response: + if response.status != 200: + raise RuntimeError( + f"Failed to close session {session_id}: {response.status} {await response.text()}" + ) + + @dataclass class RequestFuncInput: """Input for requesting LLMs via API""" @@ -64,6 +79,8 @@ class RequestFuncInput: stream: bool = True session_id: Optional[str] = None turn_idx: Optional[int] = None + enable_session_control: bool = False + session_control_url: Optional[str] = None @dataclass @@ -411,6 +428,9 @@ async def async_request_eb_openai_chat_completions( if request_func_input.ignore_eos: payload["ignore_eos"] = request_func_input.ignore_eos + if request_func_input.session_id: + payload["session_id"] = request_func_input.session_id + headers = { "Content-Type": "application/json", "Authorization": f"Bearer {os.environ.get('OPENAI_API_KEY')}", @@ -727,7 +747,7 @@ async def simple_tool_call(model_output, tool_url: str, timeout=60): return str(e), False, tool_name, tool_id -async def async_request_eb_openai_chat_completions_multi_turn( +async def _async_request_eb_openai_chat_completions_multi_turn( request_func_input: RequestFuncInput, pbar: Optional[tqdm] = None, ): @@ -1024,6 +1044,35 @@ async def async_request_eb_openai_chat_completions_multi_turn( return outputs, metrics +async def async_request_eb_openai_chat_completions_multi_turn( + request_func_input: RequestFuncInput, + pbar: Optional[tqdm] = None, +): + if not request_func_input.enable_session_control: + return await _async_request_eb_openai_chat_completions_multi_turn(request_func_input, pbar) + + request_func_input.session_id = request_func_input.session_id or uuid.uuid4().hex + result = None + try: + result = await _async_request_eb_openai_chat_completions_multi_turn(request_func_input, pbar) + finally: + try: + async with aiohttp.ClientSession(trust_env=True, timeout=AIOHTTP_TIMEOUT) as session: + await _close_session( + session, + request_func_input.session_control_url or request_func_input.api_url, + request_func_input.session_id, + ) + except Exception: + close_error = traceback.format_exc() + logging.error("Failed to close session %s: %s", request_func_input.session_id, close_error) + if result and result[0]: + result[0][-1].success = False + result[0][-1].error += close_error + + return result + + async def async_request_eb_openai_completions( request_func_input: RequestFuncInput, pbar: Optional[tqdm] = None, diff --git a/benchmarks/benchmark_serving.py b/benchmarks/benchmark_serving.py index 95875592a6f..a3e94e931dc 100644 --- a/benchmarks/benchmark_serving.py +++ b/benchmarks/benchmark_serving.py @@ -408,6 +408,8 @@ async def benchmark( lora_modules: Optional[Iterable[str]], extra_body: Optional[dict], ip_list: Optional[list[str]] = None, + enable_session_control: bool = False, + session_control_url: Optional[str] = None, ): """Benchmarks an API endpoint using a given set of sample inputs and returns""" if backend in ASYNC_REQUEST_FUNCS: @@ -467,6 +469,8 @@ async def benchmark( stream=args.stream, session_id=f"warmup-{uuid.uuid4().hex}", turn_idx=0, + enable_session_control=enable_session_control, + session_control_url=session_control_url, ) if args.warmup: @@ -581,6 +585,8 @@ async def limited_request_func(request_func_input, pbar): tokenizer_model=args.tokenizer_model, tokenizer_path=args.tokenizer_path, stream=args.stream, + enable_session_control=enable_session_control, + session_control_url=session_control_url, ) tasks.append(asyncio.create_task(limited_request_func(request_func_input=request_func_input, pbar=pbar))) @@ -670,6 +676,8 @@ async def limited_request_func_per_ip(req_input, semaphore, pbar): tokenizer_model=args.tokenizer_model, tokenizer_path=args.tokenizer_path, stream=args.stream, + enable_session_control=enable_session_control, + session_control_url=session_control_url, ) tasks.append(asyncio.create_task(limited_request_func_per_ip(req_input, semaphore, pbar))) @@ -1249,6 +1257,8 @@ def main(args: argparse.Namespace): np.random.seed(args.seed) backend = args.backend + if args.enable_session_control and not args.multi_turn: + raise ValueError("--enable-session-control requires --multi-turn.") # 支持多轮对话方式请求,仅支持chat接口 if args.multi_turn: backend = "openai-chat-multi-turn" @@ -1362,6 +1372,8 @@ def main(args: argparse.Namespace): lora_modules=args.lora_modules, extra_body=sampling_params, ip_list=ip_list, + enable_session_control=args.enable_session_control, + session_control_url=args.session_control_url, ) ) @@ -1569,6 +1581,17 @@ def main(args: argparse.Namespace): action="store_true", help="按多轮对话方式请求", ) + parser.add_argument( + "--enable-session-control", + action="store_true", + help="为 SGLang 多轮请求携带 session id,并在 session 结束后调用 close_session 释放 radix cache", + ) + parser.add_argument( + "--session-control-url", + type=str, + default=None, + help="SGLang session 控制接口的 base URL;通过不代理 close_session 的 router 发压时需要指定", + ) parser.add_argument( "--no-warmup", action="store_false",