From 72fcf3fa50e63afc67438ce3a0baf9a2fc94afa0 Mon Sep 17 00:00:00 2001 From: YanhuiDua Date: Fri, 10 Jul 2026 10:39:08 +0000 Subject: [PATCH] refine session server --- tests/rl/test_session_server.py | 337 +++++++ xtuner/v1/rl/rollout/session_server.py | 1264 +++++++++++++----------- 2 files changed, 1051 insertions(+), 550 deletions(-) create mode 100644 tests/rl/test_session_server.py diff --git a/tests/rl/test_session_server.py b/tests/rl/test_session_server.py new file mode 100644 index 000000000..33c75008a --- /dev/null +++ b/tests/rl/test_session_server.py @@ -0,0 +1,337 @@ +import json +import unittest +from types import SimpleNamespace + +from xtuner.v1.rl.rollout import session_server as session_server_mod + + +class _RemoteMethod: + def __init__(self, func): + self._func = func + + async def remote(self, *args, **kwargs): + return self._func(*args, **kwargs) + + +class _FakeTokenizer: + eos_token = "" + + def apply_chat_template(self, messages, tools=None, add_generation_prompt=False, tokenize=False): + del tools, tokenize + rendered = [] + for message in messages: + role = message["role"] + content = message.get("content") or "" + if role == "user": + rendered.append(f"user:{content}\n") + elif role == "assistant": + rendered.append(f"assistant:{content}") + else: + rendered.append(f"{role}:{content}\n") + text = "".join(rendered) + if add_generation_prompt: + text += "assistant:" + else: + text += self.eos_token + return text + + def encode(self, text, add_special_tokens=False): + del add_special_tokens + return [ord(ch) for ch in text] + + +class _FakeTraceStore: + def __init__(self, search_results): + self.search_results = list(search_results) + self.search_calls = [] + self.inserts = [] + self.search = _RemoteMethod(self._search) + self.insert = _RemoteMethod(self._insert) + + def _search(self, session_id, text, filter_none=False): + self.search_calls.append((session_id, text, filter_none)) + if self.search_results: + return self.search_results.pop(0) + return "", [] + + def _insert(self, *args, **kwargs): + self.inserts.append((args, kwargs)) + + +class _FakeWorkerResponse: + def __init__(self, payload, *, status=200): + self.status = status + self.headers = {"content-type": "application/json", "content-length": "999"} + self._payload = payload + + async def read(self): + return json.dumps(self._payload).encode("utf-8") + + +class TestSessionServerModes(unittest.IsolatedAsyncioTestCase): + def _rollout_request(self, data, *, stream=False, trace_enabled=False): + return session_server_mod._RolloutRequest( + original_request=None, + method="POST", + path="/v1/chat/completions", + query_string="", + headers={}, + body=json.dumps(data).encode("utf-8"), + data=data, + return_logprob=False, + return_token_ids=False, + return_routed_experts=True, + trace_enabled=trace_enabled, + session_id=data.get("session_id"), + messages=data.get("messages"), + tools=data.get("tools"), + stream=stream, + ) + + async def test_default_proxy_mode_preserves_default_proxy_options(self): + caller_data = { + "session_id": "session-1", + "messages": [{"role": "user", "content": "hello"}], + "logprobs": False, + "top_logprobs": 2, + "stream": False, + } + rollout_request = self._rollout_request(caller_data) + mode = session_server_mod._DefaultProxyMode( + worker_base_url="http://worker", + stop_word="", + ) + + worker_request = await mode.prepare_worker_request(rollout_request) + worker_payload = json.loads(worker_request.body) + + self.assertEqual(worker_request.method, "POST") + self.assertEqual(worker_request.target_url, "http://worker/v1/chat/completions") + self.assertNotIn("session_id", worker_payload) + self.assertNotIn("top_logprobs", worker_payload) + self.assertEqual(worker_payload["return_logprob"], False) + self.assertEqual(worker_payload["return_token_ids"], False) + self.assertEqual(worker_payload["return_routed_experts"], True) + self.assertIn("session_id", caller_data) + + def test_clean_caller_payload_uses_original_caller_options(self): + rollout_request = self._rollout_request({}) + rollout_request.return_routed_experts = False + payload = { + "output_ids": [1], + "output_token_logprobs": [[-0.1, 1]], + "routed_experts": "top-level-experts", + "choices": [ + { + "message": {"content": "hello"}, + "delta": {"content": "world"}, + "output_ids": [1], + "output_token_logprobs": [[-0.1, 1]], + "routed_experts": "choice-experts", + "logprobs": {"content": []}, + } + ], + } + mode = session_server_mod._DefaultProxyMode(worker_base_url="http://worker", stop_word="") + + cleaned = mode._clean_caller_payload(payload, rollout_request) + + self.assertNotIn("output_ids", cleaned) + self.assertNotIn("output_token_logprobs", cleaned) + self.assertNotIn("routed_experts", cleaned) + choice = cleaned["choices"][0] + self.assertNotIn("output_ids", choice) + self.assertNotIn("output_token_logprobs", choice) + self.assertNotIn("routed_experts", choice) + self.assertNotIn("logprobs", choice) + self.assertEqual(choice["message"]["content"], "hello") + self.assertEqual(choice["delta"]["content"], "world") + + def test_parse_sse_to_complete_response_builds_complete_message(self): + raw = b"".join( + [ + b'data: {"id":"cmpl-1","model":"m","choices":[{"delta":{"content":"he"},' + b'"output_ids":[10],"output_token_logprobs":[[-0.1,10]]}]}\n\n', + b'data: {"choices":[{"delta":{"content":"llo","reasoning_content":"r"},' + b'"output_ids":[11],"output_token_logprobs":[[-0.2,11]],' + b'"routed_experts":"expert-key","finish_reason":"stop"}],"usage":{"completion_tokens":2}}\n\n', + b"data: [DONE]\n\n", + ] + ) + response = session_server_mod._parse_sse_to_complete_response(raw) + + self.assertEqual(response["id"], "cmpl-1") + self.assertEqual(response["model"], "m") + choice = response["choices"][0] + self.assertEqual(choice["message"]["content"], "hello") + self.assertEqual(choice["message"]["reasoning_content"], "r") + self.assertEqual(choice["output_ids"], [10, 11]) + self.assertEqual(choice["output_token_logprobs"], [[-0.1, 10], [-0.2, 11]]) + self.assertEqual(choice["routed_experts"], "expert-key") + self.assertEqual(choice["finish_reason"], "stop") + self.assertEqual(response["usage"], {"completion_tokens": 2}) + + async def test_trace_store_request_prepares_incremental_input_ids_and_forces_trace_options(self): + tokenizer = _FakeTokenizer() + messages = [{"role": "user", "content": "hello"}] + prompt_text = tokenizer.apply_chat_template(messages, add_generation_prompt=True, tokenize=False) + prefix = "user:hello\n" + prefix_node = SimpleNamespace( + value=session_server_mod.TokenizedSegment(text=prefix, token_ids=[101, 102]) + ) + trace_store = _FakeTraceStore(search_results=[(prefix, [prefix_node])]) + mode = session_server_mod._TraceStoreMode( + tokenizer=tokenizer, + trace_store=trace_store, + worker_base_url="http://worker", + stop_word=tokenizer.eos_token, + ) + rollout_request = self._rollout_request( + { + "session_id": "session-1", + "messages": messages, + "logprobs": False, + "top_logprobs": 5, + "return_token_ids": True, + "temperature": 0.7, + }, + trace_enabled=True, + ) + + worker_request = await mode.prepare_worker_request(rollout_request) + worker_payload = json.loads(worker_request.body) + + delta = prompt_text[len(prefix) :] + self.assertEqual(worker_payload["input_ids"], [101, 102, *tokenizer.encode(delta)]) + self.assertEqual(worker_payload["messages"], []) + self.assertEqual(worker_payload["return_token_ids"], True) + self.assertEqual(worker_payload["return_routed_experts"], True) + self.assertEqual(worker_payload["return_logprob"], True) + self.assertEqual(worker_payload["include_stop_str_in_output"], True) + self.assertEqual(worker_payload["temperature"], 0.7) + self.assertNotIn("session_id", worker_payload) + self.assertNotIn("top_logprobs", worker_payload) + self.assertEqual(trace_store.search_calls, [("session-1", prompt_text, True)]) + insert_args, insert_kwargs = trace_store.inserts[0] + self.assertEqual(insert_args[:2], ("session-1", prompt_text)) + self.assertEqual(insert_kwargs, {}) + inserted_segment = insert_args[2] + self.assertEqual(inserted_segment.text, delta) + self.assertEqual(inserted_segment.token_ids, tokenizer.encode(delta)) + + async def test_trace_store_response_records_raw_worker_tokens_while_cleaning_caller_response(self): + tokenizer = _FakeTokenizer() + trace_store = _FakeTraceStore(search_results=[]) + mode = session_server_mod._TraceStoreMode( + tokenizer=tokenizer, + trace_store=trace_store, + worker_base_url="http://worker", + stop_word=tokenizer.eos_token, + ) + messages = [{"role": "user", "content": "hello"}] + rollout_request = self._rollout_request( + {"session_id": "session-1", "messages": messages}, + trace_enabled=True, + ) + worker_payload = { + "choices": [ + { + "message": {"role": "assistant", "content": "world"}, + "output_ids": [11, 12], + "output_token_logprobs": [[-0.1, 11], [-0.2, 12]], + "logprobs": {"content": []}, + } + ] + } + + response = await mode.handle_worker_response(rollout_request, _FakeWorkerResponse(worker_payload)) + + caller_payload = json.loads(response.body) + self.assertNotIn("output_ids", caller_payload["choices"][0]) + self.assertNotIn("output_token_logprobs", caller_payload["choices"][0]) + self.assertNotIn("logprobs", caller_payload["choices"][0]) + insert_args, insert_kwargs = trace_store.inserts[0] + self.assertEqual(insert_args, ("session-1",)) + self.assertEqual(insert_kwargs["key"], "user:hello\nassistant:world") + inserted_segment = insert_kwargs["value"] + self.assertEqual(inserted_segment.text, "world") + self.assertEqual(inserted_segment.token_ids, [11, 12]) + self.assertEqual(inserted_segment.labels, [11, 12]) + self.assertEqual(inserted_segment.logprobs, [-0.1, -0.2]) + + async def test_trace_store_response_returns_error_when_required_trace_fields_are_missing(self): + tokenizer = _FakeTokenizer() + trace_store = _FakeTraceStore(search_results=[]) + mode = session_server_mod._TraceStoreMode( + tokenizer=tokenizer, + trace_store=trace_store, + worker_base_url="http://worker", + stop_word=tokenizer.eos_token, + ) + rollout_request = self._rollout_request( + { + "session_id": "session-1", + "messages": [{"role": "user", "content": "hello"}], + }, + trace_enabled=True, + ) + + response = await mode.handle_worker_response( + rollout_request, + _FakeWorkerResponse({"choices": [{"message": {"role": "assistant", "content": "world"}}]}), + ) + + error_payload = json.loads(response.text) + self.assertEqual(response.status, 500) + self.assertEqual(error_payload["object"], "error") + self.assertIn("no output_ids", error_payload["message"]) + self.assertEqual(trace_store.inserts, []) + + def test_parse_sse_to_complete_response_merges_tool_call_deltas(self): + raw = b"".join( + [ + b'data: {"choices":[{"delta":{"tool_calls":[{"index":0,"id":"call-1",' + b'"type":"function","function":{"name":"get_"}}]}}]}\n\n', + b'data: {"choices":[{"delta":{"tool_calls":[{"index":0,' + b'"function":{"name":"weather","arguments":"{\\"city\\""}}]},' + b'"output_ids":[10],"output_token_logprobs":[[-0.1,10]]}]}\n\n', + b'data: {"choices":[{"delta":{"tool_calls":[{"index":0,' + b'"function":{"arguments":":\\"HZ\\"}"}}]},' + b'"output_ids":[11],"output_token_logprobs":[[-0.2,11]],' + b'"finish_reason":"tool_calls"}]}\n\n', + b"data: [DONE]\n\n", + ] + ) + + response = session_server_mod._parse_sse_to_complete_response(raw) + + choice = response["choices"][0] + tool_call = choice["message"]["tool_calls"][0] + self.assertEqual(tool_call["id"], "call-1") + self.assertEqual(tool_call["type"], "function") + self.assertEqual(tool_call["function"]["name"], "get_weather") + self.assertEqual(tool_call["function"]["arguments"], '{"city":"HZ"}') + self.assertEqual(choice["finish_reason"], "tool_calls") + self.assertEqual(choice["output_ids"], [10, 11]) + + def test_sse_traceability_rejects_incomplete_or_error_streams(self): + complete = ( + b'data: {"choices":[{"delta":{"content":"ok"},"output_ids":[1],' + b'"finish_reason":"stop"}]}\n\n' + b"data: [DONE]\n\n" + ) + missing_done = b'data: {"choices":[{"delta":{"content":"ok"},"output_ids":[1],"finish_reason":"stop"}]}\n\n' + error_choice = ( + b'data: {"choices":[{"delta":{"content":"bad"},"finish_reason":"error"}]}\n\n' + b"data: [DONE]\n\n" + ) + upstream_error = b'data: {"object":"error","message":"bad"}\n\n' + + self.assertTrue(session_server_mod._has_complete_traceable_sse_response(complete)) + self.assertFalse(session_server_mod._has_complete_traceable_sse_response(missing_done)) + self.assertFalse(session_server_mod._has_complete_traceable_sse_response(error_choice)) + self.assertFalse(session_server_mod._has_complete_traceable_sse_response(upstream_error)) + with self.assertRaisesRegex(RuntimeError, "without \\[DONE\\]"): + session_server_mod._parse_sse_to_complete_response(missing_done) + with self.assertRaisesRegex(RuntimeError, "finished with error"): + session_server_mod._parse_sse_to_complete_response(error_choice) diff --git a/xtuner/v1/rl/rollout/session_server.py b/xtuner/v1/rl/rollout/session_server.py index 92f65265a..abf522bc8 100644 --- a/xtuner/v1/rl/rollout/session_server.py +++ b/xtuner/v1/rl/rollout/session_server.py @@ -1,11 +1,15 @@ +from __future__ import annotations + import json +from contextlib import asynccontextmanager +from dataclasses import dataclass, field from functools import reduce from operator import add -from typing import Any, Optional +from typing import Any, AsyncIterator import numpy as np import ray -from aiohttp import ClientConnectionResetError, ClientSession, ClientTimeout, web +from aiohttp import ClientConnectionResetError, ClientResponse, ClientSession, ClientTimeout, web from transformers import AutoTokenizer from xtuner.v1.utils import get_logger @@ -14,101 +18,50 @@ from .trace_store import TokenizedSegment, get_store -def _is_error_payload(payload: dict) -> bool: - return payload.get("error") is not None or payload.get("type") == "error" or payload.get("object") == "error" - - -def _lmdeploy_error_payload(message: str, status: int = 500, error_type: str = "internal_server_error") -> dict: - return { - "message": message, - "type": error_type, - "code": status, - "object": "error", - } - - -def _stream_has_traceable_choices(raw: bytes) -> bool: - text = raw.decode("utf-8", errors="replace") - has_choices = False - saw_done = False - saw_terminal_finish = False - for line in text.split("\n"): - line = line.strip() - if line == "data: [DONE]": - saw_done = True - continue - if line.startswith("data: "): - try: - event = json.loads(line[6:]) - except json.JSONDecodeError: - return False - if _is_error_payload(event): - return False - if event.get("choices"): - if any(choice.get("finish_reason") == "error" for choice in event.get("choices", [])): - return False - if any(choice.get("finish_reason") for choice in event.get("choices", [])): - saw_terminal_finish = True - has_choices = True - return has_choices and saw_done and saw_terminal_finish - - -def _extract_output_logprobs(choice: dict, output_token_ids: list[int]) -> list[float]: - if not output_token_ids: - return [] - - output_token_logprobs = choice.get("output_token_logprobs") - if output_token_logprobs is None: - raise RuntimeError( - "SessionServer response choice has no output_token_logprobs; " - "the return_logprob protocol is required for training traces." - ) - - logprob_token_ids = [item[1] for item in output_token_logprobs] - if logprob_token_ids != output_token_ids: - raise RuntimeError( - "SessionServer response choice has mismatched output_token_logprobs: " - f"output_ids_len={len(output_token_ids)}, logprob_ids_len={len(logprob_token_ids)}" - ) - return [item[0] for item in output_token_logprobs] - - -_SESSION_SERVER_ONLY_KEYS = {"session_id"} - - -def _bool_request_value(value: Any, default: bool = False) -> bool: - if value is None: - return default - if isinstance(value, str): - return value.strip().lower() not in {"", "0", "false", "no", "off"} - return bool(value) - - -def _request_uses_trace_store(req_body: dict) -> bool: - if req_body.get("session_id") is None or "messages" not in req_body: - return False - return _bool_request_value(req_body.get("return_token_ids"), True) +@dataclass +class _RolloutRequest: + original_request: web.Request | None + method: str + path: str + query_string: str + headers: dict[str, str] + body: bytes + data: dict[str, Any] | None + return_logprob: bool + return_token_ids: bool + return_routed_experts: bool + trace_enabled: bool + session_id: str | None + messages: list[dict[str, Any]] | None + tools: Any | None + stream: bool + + +@dataclass +class _WorkerRequest: + method: str = "" + target_url: str = "" + headers: dict[str, str] = field(default_factory=dict) + body: bytes = b"" class SessionServer: - """SessionServer intercepts and records requests sent to a remote LLM API + """SessionServer is both a server for rollout callers and a client to the worker. - It acts as a reverse-proxy (or interceptor) in front of an already running - worker (like lmdeploy, sglang, or vllm). It binds to a specific (host, port) - and relays any received traffic to the actual worker URL. - - You can optionally provide before_request and after_response hooks to - perform extra logging, trace state mutations, or message cleanup before/after - routing the request to the worker backend. + It exposes an OpenAI-compatible HTTP server to rollout code, then forwards + each request as an HTTP client to an already running worker such as lmdeploy, + sglang, or vllm. _build_rollout_request parses caller requests, the selected + response handler writes caller responses, and _WorkerClient owns the worker + client transport. Args: - worker_base_url (str): The base URL of the real worker (e.g. "http://127.0.0.1:8000") + worker_base_url (str): The base URL of the real worker, for example "http://127.0.0.1:8000". tokenizer_path (str): The path to the tokenizer model. host (str): Host for this session server to listen on. port (int): Port for this session server to listen on. request_timeout (float): Total timeout in seconds for forwarding requests to the worker. - read_bufsize (int): Buffer limit for line reader in ClientSession. Default is 64MB (2**26). + read_bufsize (int): Buffer limit for line reader in ClientSession. """ def __init__( @@ -129,153 +82,40 @@ def __init__( self.store = get_store() self.stop_word = self.tokenizer.eos_token or "" - self._app: Optional[web.Application] = None - self._runner: Optional[web.AppRunner] = None - self._site: Optional[web.TCPSite] = None - self._lmdeploy_actor: Optional[ray.actor.ActorHandle] = None - - async def on_request(self, req_body: dict, *, trace_enabled: bool = True) -> dict: - """Hook for processing/modifying the request before forwarding.""" - - if not trace_enabled: - worker_req = {k: v for k, v in req_body.items() if k not in _SESSION_SERVER_ONLY_KEYS} - if "logprobs" in worker_req: - worker_req.setdefault("return_logprob", worker_req.pop("logprobs")) - if not _bool_request_value(worker_req.get("return_logprob"), False): - worker_req.pop("top_logprobs", None) - worker_req["return_logprob"] = False - worker_req["return_token_ids"] = False - worker_req.setdefault("return_routed_experts", True) - return worker_req - - session_id = req_body["session_id"] - # 1. chat_template render 出完整 prompt string,不 tokenize 全量 - prompt_text = self.tokenizer.apply_chat_template( - canonicalize_messages_for_chat_template(req_body["messages"]), - tools=req_body.get("tools", None), - add_generation_prompt=True, - tokenize=False, - ) - - # 2. Store 做 string prefix match。 - prefix, nodes = await self.store.search.remote(session_id, prompt_text, filter_none=True) - if prefix: - get_logger().debug(f"Hit prefix cache for session {session_id}") - delta, delta_ids = prompt_text[len(prefix) :], [] - if delta: - delta_ids = self.tokenizer.encode(delta, add_special_tokens=False) - await self.store.insert.remote(session_id, prompt_text, TokenizedSegment(text=delta, token_ids=delta_ids)) - input_ids = reduce(add, [node.value.token_ids for node in nodes] + [delta_ids]) - - # 3. 组装 OpenAI chat completions 请求。 - worker_req = { - **{ - k: v - for k, v in req_body.items() - if k not in _SESSION_SERVER_ONLY_KEYS | {"messages", "logprobs", "top_logprobs"} - }, - "messages": [], - "input_ids": input_ids, - "return_token_ids": True, - "return_routed_experts": True, - "return_logprob": True, - "include_stop_str_in_output": True, - } - return worker_req - - async def on_response(self, worker_resp: dict, *, trace_enabled: bool = True) -> dict: - """Hook for processing the parsed response received from the worker.""" - - if not trace_enabled: - return {k: v for k, v in worker_resp.items() if k not in {"messages", "tools"}} - - session_id = worker_resp["session_id"] - messages = worker_resp["messages"] - tools = worker_resp["tools"] - choice = worker_resp["choices"][0] - - output_token_ids = choice.get("output_ids") # len = N_out - if output_token_ids is None: - raise RuntimeError( - "SessionServer response choice has no output_ids; " - "cannot export a training trace for this assistant turn." - ) - output_logprobs = _extract_output_logprobs(choice, output_token_ids) - raw_routed_expert = choice.get("routed_experts") # 本次 call 的 raw routed_expert,可为 None - - # 2. Store 把 input_delta / assistant_output 两个节点补齐字段。 - old_prompt = self.tokenizer.apply_chat_template( - canonicalize_messages_for_chat_template(messages), tools=tools, add_generation_prompt=True, tokenize=False + self.default_proxy_mode = _DefaultProxyMode(worker_base_url=self.worker_base_url, stop_word=self.stop_word) + self.trace_store_mode = _TraceStoreMode( + tokenizer=self.tokenizer, + trace_store=self.store, + worker_base_url=self.worker_base_url, + stop_word=self.stop_word, ) - messages = [*messages, choice["message"]] - new_prompt = ( - self.tokenizer.apply_chat_template( - canonicalize_messages_for_chat_template(messages), - tools=tools, - add_generation_prompt=False, - tokenize=False, - ) - ).rstrip() - assert new_prompt.startswith(old_prompt) and new_prompt.endswith(self.stop_word) - - if raw_routed_expert is not None: - raw_routed_expert = await self._decode_routed_experts(raw_routed_expert) - if len(raw_routed_expert) > 0: - num_layers = raw_routed_expert.shape[1] - topk_experts = raw_routed_expert.shape[2] - dummy_expert = np.full((1, num_layers, topk_experts), 0, dtype=raw_routed_expert.dtype) - raw_routed_expert = np.concatenate([dummy_expert, raw_routed_expert], axis=0) - - _, nodes = await self.store.search.remote(session_id, old_prompt, filter_none=True) - - # last node in nodes corresponds to the delta inserted in on_request (if any) - if nodes: - delta_node_val: TokenizedSegment = nodes[-1].value - delta_len = len(delta_node_val.token_ids) - prefix_len = sum(len(n.value.token_ids) for n in nodes[:-1]) - assert prefix_len + delta_len + len(output_token_ids) == len(raw_routed_expert) - - # split raw_routed_expert - # raw_routed_expert target shape mapping: [prefix_len + delta_len + response_len, ...] - delta_expert = raw_routed_expert[prefix_len : prefix_len + delta_len] - response_expert = raw_routed_expert[prefix_len + delta_len :] - - if delta_len > 0: - delta_node_val.expert_key = ray.put(delta_expert) - # update delta node in store - await self.store.insert.remote(session_id, old_prompt, delta_node_val) - - raw_routed_expert = ray.put(response_expert) - else: - raw_routed_expert = ray.put(raw_routed_expert) - - await self.store.insert.remote( - session_id, - key=new_prompt, - value=TokenizedSegment( - text=new_prompt[len(old_prompt) :], - token_ids=output_token_ids, - logprobs=output_logprobs, - labels=output_token_ids, - expert_key=raw_routed_expert, - length=len(output_token_ids), - ), + self.client = _WorkerClient( + request_timeout=self.request_timeout, + read_bufsize=self.read_bufsize, ) - # 3. 返回标准 OpenAI response,session_id 由 SessionClient 层再剥 - resp = {k: v for k, v in worker_resp.items() if k != "messages"} - return resp + self._app: web.Application | None = None + self._runner: web.AppRunner | None = None + self._site: web.TCPSite | None = None @property def url(self) -> str: - """The bound URL for the SessionServer.""" + """The bound URL for the SessionServer. + + Returns: + str: The HTTP URL this SessionServer listens on. + """ + return f"http://{self.host}:{self.port}" - async def start(self): + async def start(self) -> None: """Start the SessionServer proxy application.""" + if self._site is not None: return + await self.client.start() + self._app = web.Application() self._app.router.add_route("*", "/{path:.*}", self._handle_request) @@ -285,344 +125,39 @@ async def start(self): await self._site.start() get_logger().info(f"SessionServer listening on {self.url} (Forwarding to {self.worker_base_url})") - async def stop(self): + async def stop(self) -> None: """Cleanly stop the SessionServer application.""" + if self._runner: await self._runner.cleanup() + await self.client.stop() self._site = None self._runner = None self._app = None get_logger().info("SessionServer stopped.") - async def _handle_request(self, request: web.Request) -> web.Response: - """Proxy handler for the worker API.""" - - # Read the request body - request_body = await request.read() - request_data = session_id = messages = None - trace_enabled = False - orig_return_logprob = orig_return_token_ids = orig_return_routed_experts = False - if request_body: - try: - request_data = json.loads(request_body) - - trace_enabled = _request_uses_trace_store(request_data) - orig_return_logprob = _bool_request_value( - request_data.get("return_logprob", request_data.get("logprobs")), False - ) - orig_return_token_ids = _bool_request_value(request_data.get("return_token_ids"), False) - orig_return_routed_experts = _bool_request_value(request_data.get("return_routed_experts"), True) - - session_id = request_data.get("session_id") - messages = request_data.get("messages") - tools = request_data.get("tools", None) - - # Apply purely abstract on_request processing - request_data = await self.on_request(request_data, trace_enabled=trace_enabled) - # Re-serialize the modified payload back to bytes - request_body = json.dumps(request_data).encode("utf-8") - except json.JSONDecodeError: - pass - except Exception as exc: - message = f"SessionServer request hook failed: {type(exc).__name__}: {exc}" - get_logger().error(message) - return web.json_response(_lmdeploy_error_payload(message), status=500) - - # Build forwarding headers, dropping original Host - forward_headers = dict(request.headers) - forward_headers.pop("Host", None) - forward_headers.pop("host", None) - forward_headers.pop("Content-Length", None) - forward_headers.pop("content-length", None) - - # Re-build Path - req_path = request.match_info["path"] - target_url = f"{self.worker_base_url}/{req_path.lstrip('/')}" - if request.query_string: - target_url += f"?{request.query_string}" - - is_stream = request_data.get("stream", False) if request_data else False - - def _clean_data(data: dict) -> bool: - modified = False - for key, drop in [ - ("output_token_logprobs", not orig_return_logprob), - ("output_ids", not orig_return_token_ids), - ("routed_experts", not orig_return_routed_experts), - ]: - if drop and key in data: - data.pop(key) - modified = True - if drop: - for c in data.get("choices", []): - if key in c: - c.pop(key) - modified = True - - for c in data.get("choices", []): - if "logprobs" in c: - c.pop("logprobs") - modified = True - - for c in data.get("choices", []): - if c.get("message") and isinstance(c["message"].get("content"), str): - if self.stop_word in c["message"]["content"]: - c["message"]["content"] = c["message"]["content"].replace(self.stop_word, "") - modified = True - if c.get("delta") and isinstance(c["delta"].get("content"), str): - if self.stop_word in c["delta"]["content"]: - c["delta"]["content"] = c["delta"]["content"].replace(self.stop_word, "") - modified = True - - return modified - - # Forward the request to the upstream worker - # read_bufsize controls StreamReader's line buffer limit; SSE events with large - # tool_calls/reasoning_content payloads can exceed the 64KB default and trigger - # "Chunk too big" from readuntil(b"\n"). - timeout = ClientTimeout(total=self.request_timeout, sock_connect=30) - async with ClientSession(read_bufsize=self.read_bufsize, timeout=timeout) as client: - async with client.request( - method=request.method, url=target_url, headers=forward_headers, data=request_body - ) as resp: - # Setup proper stream vs sync response objects - if is_stream: - response_chunks = [] - response = web.StreamResponse( - status=resp.status, - headers={ - k: v - for k, v in resp.headers.items() - if k.lower() not in ("transfer-encoding", "content-length", "content-encoding") - }, - ) - await response.prepare(request) - # If the downstream client closes the socket mid-stream - # (e.g. AsyncAPIClient bails out on a finish_reason=='error' - # chunk after the prompt overflowed the session window), - # keep draining the upstream so the trace is still recorded - # in full but stop attempting to write to the closed socket. - client_alive = True - async for line in resp.content: - # Keep unmodified line for trace store parsing - if trace_enabled: - response_chunks.append(line) - - # Dynamically prune added fields before writing to client - if request_data is not None and line.startswith(b"data: ") and line.strip() != b"data: [DONE]": - try: - text = line.decode("utf-8") - data = json.loads(text[6:]) - if _clean_data(data): - line = ("data: " + json.dumps(data) + "\n").encode("utf-8") - except Exception: - pass - - # Delay [DONE] only while a training trace still needs to be exported. - if client_alive and (not trace_enabled or line.strip() != b"data: [DONE]"): - try: - await response.write(line) - except (ConnectionError, ClientConnectionResetError): - client_alive = False - - raw_response = b"".join(response_chunks) if trace_enabled else b"" - else: - raw_response = await resp.read() - final_raw_response = raw_response - - if request_data is not None: - try: - clean_data = json.loads(raw_response) - if _clean_data(clean_data): - final_raw_response = json.dumps(clean_data).encode("utf-8") - except Exception: - pass - - response = web.Response( - status=resp.status, - headers={ - k: v - for k, v in resp.headers.items() - if k.lower() not in ("transfer-encoding", "content-length", "content-encoding") - }, - body=final_raw_response, # Modified raw response without our injected trace params - ) - - # Apply abstract on_response processing - response_data = None - skip_done = bool(is_stream and not trace_enabled) - session_error_msg = None - if request_data and trace_enabled: - if is_stream: - skip_done = not _stream_has_traceable_choices(raw_response) - if not skip_done: - try: - response_data = self._parse_stream_response(raw_response) - except Exception as exc: - session_error_msg = f"SessionServer stream trace failed: {type(exc).__name__}: {exc}" - else: - try: - response_data = json.loads(raw_response) - except json.JSONDecodeError: - pass - if isinstance(response_data, dict) and _is_error_payload(response_data): - response_data = None - - if response_data is not None: - try: - for c in response_data.get("choices", []): - if c.get("message") and isinstance(c["message"].get("content"), str): - c["message"]["content"] = c["message"]["content"].replace(self.stop_word, "") - - response_data["session_id"] = session_id - response_data["messages"] = messages - response_data["tools"] = tools - await self.on_response(response_data, trace_enabled=trace_enabled) - except Exception as exc: - session_error_msg = f"SessionServer response hook failed: {type(exc).__name__}: {exc}" - - if session_error_msg: - get_logger().error(session_error_msg) - - if is_stream: - try: - if session_error_msg: - error_payload = _lmdeploy_error_payload(session_error_msg) - await response.write( - ("data: " + json.dumps(error_payload, ensure_ascii=False) + "\n\n").encode("utf-8") - ) - skip_done = True - if not skip_done: - await response.write(b"data: [DONE]\n\n") - await response.write_eof() - except (ConnectionError, ClientConnectionResetError): - # Client already gone; trace was still recorded above. - pass - elif session_error_msg: - return web.json_response(_lmdeploy_error_payload(session_error_msg), status=500) - - return response + async def _handle_request(self, request: web.Request) -> web.StreamResponse | web.Response: + """Aiohttp entrypoint for one caller request through the worker + proxy.""" + + try: + rollout_request = await _build_rollout_request(request) + mode = self.trace_store_mode if rollout_request.trace_enabled else self.default_proxy_mode + worker_request = await mode.prepare_worker_request(rollout_request) + except _PrepareRequestError as exc: + get_logger().error(exc.message) + return web.json_response( + { + "message": exc.message, + "type": "internal_server_error", + "code": exc.status, + "object": "error", + }, + status=exc.status, + ) - async def _decode_routed_experts(self, routed_experts: Any) -> np.ndarray: - if isinstance(routed_experts, str): - if self._lmdeploy_actor is None: - self._lmdeploy_actor = ray.get_actor("shared_store", namespace="lmdeploy") - assert self._lmdeploy_actor is not None, "LMDeploy actor should be available in the shared store." - routed_experts_data = await self._lmdeploy_actor.get.remote(routed_experts) - return np.asarray(routed_experts_data) - return np.asarray(routed_experts) - - @staticmethod - def _parse_stream_response(raw: bytes) -> Optional[dict]: - """Parse SSE stream to reconstruct the complete final message state.""" - text = raw.decode("utf-8", errors="replace") - events = [] - saw_done = False - for line in text.split("\n"): - line = line.strip() - if line == "data: [DONE]": - saw_done = True - continue - if line.startswith("data: "): - event = json.loads(line[6:]) - if _is_error_payload(event): - raise RuntimeError(f"Upstream SSE stream returned error: {json.dumps(event, ensure_ascii=False)}") - events.append(event) - - if not events: - return None - if not any(event.get("choices") for event in events): - raise RuntimeError(f"Upstream SSE stream ended without choices: {json.dumps(events, ensure_ascii=False)}") - - # Reconstruct standard stream output (Assuming OpenAI format here) - message: dict[str, Any] = {"choices": [{"message": {"role": "assistant", "content": ""}}]} - content_parts: list[str] = [] - tool_calls_map: dict[int, dict[str, Any]] = {} - usage: dict[str, Any] = {} - - for event in events: - if event.get("id") and "id" not in message: - message["id"] = event["id"] - if event.get("model"): - message["model"] = event["model"] - - choices = event.get("choices", []) - for choice in choices: - if choice.get("finish_reason") == "error": - raise RuntimeError( - f"Upstream SSE choice finished with error: {json.dumps(event, ensure_ascii=False)}" - ) - delta = choice.get("delta", {}) - - # Check text content - if delta.get("content"): - content_parts.append(delta["content"]) - - # Check output ids - if choice.get("output_ids") is not None: - assistant_choice = message["choices"][0] - if "output_ids" not in assistant_choice: - assistant_choice["output_ids"] = [] - assistant_choice["output_ids"].extend(choice["output_ids"]) - - # Check routed experts. LMDeploy only emits this in the final - # chunk, often as a Ray shared-store key string. - if choice.get("routed_experts") is not None: - assistant_choice = message["choices"][0] - assistant_choice["routed_experts"] = choice["routed_experts"] - - # Check raw output logprobs from LMDeploy return_logprob protocol. - if choice.get("output_token_logprobs") is not None: - assistant_choice = message["choices"][0] - if "output_token_logprobs" not in assistant_choice: - assistant_choice["output_token_logprobs"] = [] - assistant_choice["output_token_logprobs"].extend(choice["output_token_logprobs"]) - - # Check reasoning content - if delta.get("reasoning_content"): - assistant_msg = message["choices"][0]["message"] - assistant_msg["reasoning_content"] = ( - assistant_msg.get("reasoning_content", "") + delta["reasoning_content"] - ) - - # Check tool calls - for tc_delta in delta.get("tool_calls") or []: - idx = tc_delta.get("index", 0) - if idx not in tool_calls_map: - tool_calls_map[idx] = { - "id": tc_delta.get("id", ""), - "type": tc_delta.get("type", "function"), - "function": {"name": "", "arguments": ""}, - } - tc = tool_calls_map[idx] - fn = tc_delta.get("function", {}) - if fn.get("name"): - tc["function"]["name"] += fn["name"] - if fn.get("arguments"): - tc["function"]["arguments"] += fn["arguments"] - - if choice.get("finish_reason"): - message["choices"][0]["finish_reason"] = choice["finish_reason"] - - if event.get("usage") is not None: - usage = event["usage"] - - msg = message["choices"][0]["message"] - msg["content"] = "".join(content_parts) - if tool_calls_map: - msg["tool_calls"] = [tool_calls_map[i] for i in sorted(tool_calls_map)] - if usage: - message["usage"] = usage - - assistant_choice = message["choices"][0] - if not saw_done: - raise RuntimeError("Upstream SSE stream ended without [DONE].") - if not assistant_choice.get("finish_reason"): - raise RuntimeError("Upstream SSE stream ended without terminal finish_reason.") - if assistant_choice.get("output_ids") is None: - raise RuntimeError("Upstream SSE stream ended without output_ids.") - - return message + async with self.client.request(worker_request) as worker_response: + return await mode.handle_worker_response(rollout_request, worker_response) class SessionServerActor: @@ -658,3 +193,632 @@ async def stop(self) -> None: if self.server is not None: await self.server.stop() self.server = None + + +class _WorkerClient: + """Worker-facing HTTP transport owned by SessionServer.""" + + def __init__( + self, + *, + request_timeout: float, + read_bufsize: int, + ): + self.request_timeout = request_timeout + self.read_bufsize = read_bufsize + self._session: ClientSession | None = None + + async def start(self) -> None: + if self._session is not None and not self._session.closed: + return + timeout = ClientTimeout(total=self.request_timeout, sock_connect=30) + self._session = ClientSession(read_bufsize=self.read_bufsize, timeout=timeout) + + async def stop(self) -> None: + if self._session is None: + return + await self._session.close() + self._session = None + + @asynccontextmanager + async def request(self, worker_request: _WorkerRequest) -> AsyncIterator[ClientResponse]: + if self._session is None or self._session.closed: + raise RuntimeError("Worker client must be started before forwarding requests.") + async with self._session.request( + method=worker_request.method, + url=worker_request.target_url, + headers=worker_request.headers, + data=worker_request.body, + ) as response: + yield response + + +async def _build_rollout_request(request: web.Request) -> _RolloutRequest: + request_body = await request.read() + request_data = None + if request_body: + try: + request_data = json.loads(request_body) + except json.JSONDecodeError: + request_data = None + + return_logprob = False + return_token_ids = False + return_routed_experts = True + trace_enabled = False + session_id = None + messages = None + tools = None + stream = False + if request_data is not None: + return_logprob = request_data.get("return_logprob", request_data.get("logprobs")) is True + return_token_ids = request_data.get("return_token_ids") is True + return_routed_experts_value = request_data.get("return_routed_experts") + return_routed_experts = True if return_routed_experts_value is None else return_routed_experts_value is True + trace_return_token_ids = request_data.get("return_token_ids") + trace_enabled = ( + request_data.get("session_id") is not None + and "messages" in request_data + and (True if trace_return_token_ids is None else trace_return_token_ids is True) + ) + session_id = request_data.get("session_id") + messages = request_data.get("messages") + tools = request_data.get("tools", None) + stream = request_data.get("stream") is True + + return _RolloutRequest( + original_request=request, + method=request.method, + path=request.match_info["path"], + query_string=request.query_string, + headers=dict(request.headers), + body=request_body, + data=request_data, + return_logprob=return_logprob, + return_token_ids=return_token_ids, + return_routed_experts=return_routed_experts, + trace_enabled=trace_enabled, + session_id=session_id, + messages=messages, + tools=tools, + stream=stream, + ) + + +class _BaseProxyMode: + def __init__(self, *, worker_base_url: str, stop_word: str = ""): + self.worker_base_url = worker_base_url.rstrip("/") + self.stop_word = stop_word + + async def prepare_worker_request(self, rollout_request: _RolloutRequest) -> _WorkerRequest: + worker_request = _WorkerRequest( + method=rollout_request.method, + target_url=f"{self.worker_base_url}/{rollout_request.path.lstrip('/')}", + headers=dict(rollout_request.headers), + ) + if rollout_request.query_string: + worker_request.target_url += f"?{rollout_request.query_string}" + for header in ("Host", "host", "Content-Length", "content-length"): + worker_request.headers.pop(header, None) + + if rollout_request.data is None: + worker_request.body = rollout_request.body + return worker_request + + try: + worker_payload = await self._build_worker_payload(rollout_request) + except Exception as exc: + message = f"SessionServer request hook failed: {type(exc).__name__}: {exc}" + raise _PrepareRequestError(message) from exc + + worker_request.body = json.dumps(worker_payload).encode("utf-8") + return worker_request + + async def _build_worker_payload(self, rollout_request: _RolloutRequest) -> dict[str, Any]: + raise NotImplementedError + + async def handle_worker_response( + self, + rollout_request: _RolloutRequest, + worker_response: ClientResponse, + ) -> web.StreamResponse | web.Response: + if rollout_request.stream: + response, _ = await self._relay_stream_response( + rollout_request, + worker_response, + capture_raw_response=False, + delay_done=False, + ) + try: + await response.write_eof() + except (ConnectionError, ClientConnectionResetError): + pass + return response + + response, _ = await self._read_non_stream_response(rollout_request, worker_response) + return response + + def _clean_caller_payload( + self, + payload: dict[str, Any], + rollout_request: _RolloutRequest, + ) -> dict[str, Any]: + for key, drop in [ + ("output_token_logprobs", not rollout_request.return_logprob), + ("output_ids", not rollout_request.return_token_ids), + ("routed_experts", not rollout_request.return_routed_experts), + ]: + if drop: + payload.pop(key, None) + for choice in payload.get("choices", []): + choice.pop(key, None) + + for choice in payload.get("choices", []): + choice.pop("logprobs", None) + self._remove_stop_word(choice.get("message")) + self._remove_stop_word(choice.get("delta")) + return payload + + async def _relay_stream_response( + self, + rollout_request: _RolloutRequest, + worker_response: ClientResponse, + *, + capture_raw_response: bool, + delay_done: bool, + ) -> tuple[web.StreamResponse, bytes]: + assert rollout_request.original_request is not None + raw_response_chunks: list[bytes] = [] + response = web.StreamResponse( + status=worker_response.status, + headers=_filter_response_headers(worker_response.headers), + ) + await response.prepare(rollout_request.original_request) + + client_alive = True + async for line in worker_response.content: + if capture_raw_response: + raw_response_chunks.append(line) + + if rollout_request.data is not None and line.startswith(b"data: ") and line.strip() != b"data: [DONE]": + try: + payload = json.loads(line.decode("utf-8")[6:]) + payload = self._clean_caller_payload(payload, rollout_request) + line = ("data: " + json.dumps(payload) + "\n").encode("utf-8") + except Exception: + pass + + if client_alive and (not delay_done or line.strip() != b"data: [DONE]"): + try: + await response.write(line) + except (ConnectionError, ClientConnectionResetError): + client_alive = False + + return response, b"".join(raw_response_chunks) + + async def _read_non_stream_response( + self, + rollout_request: _RolloutRequest, + worker_response: ClientResponse, + ) -> tuple[web.Response, bytes]: + raw_response = await worker_response.read() + final_raw_response = raw_response + if rollout_request.data is not None: + try: + payload = json.loads(raw_response) + payload = self._clean_caller_payload(payload, rollout_request) + final_raw_response = json.dumps(payload).encode("utf-8") + except Exception: + pass + + return ( + web.Response( + status=worker_response.status, + headers=_filter_response_headers(worker_response.headers), + body=final_raw_response, + ), + raw_response, + ) + + def _remove_stop_word(self, message: Any) -> None: + if not self.stop_word or not isinstance(message, dict): + return + content = message.get("content") + if isinstance(content, str) and self.stop_word in content: + message["content"] = content.replace(self.stop_word, "") + + +class _DefaultProxyMode(_BaseProxyMode): + async def _build_worker_payload(self, rollout_request: _RolloutRequest) -> dict[str, Any]: + assert rollout_request.data is not None + worker_req = {k: v for k, v in rollout_request.data.items() if k not in {"session_id"}} + if "logprobs" in worker_req: + worker_req.setdefault("return_logprob", worker_req.pop("logprobs")) + if worker_req.get("return_logprob") is not True: + worker_req.pop("top_logprobs", None) + worker_req["return_logprob"] = False + worker_req["return_token_ids"] = False + worker_req.setdefault("return_routed_experts", True) + return worker_req + + +class _TraceStoreMode(_BaseProxyMode): + def __init__(self, *, tokenizer: Any, trace_store: Any, worker_base_url: str, stop_word: str = ""): + super().__init__(worker_base_url=worker_base_url, stop_word=stop_word) + self.tokenizer = tokenizer + self.trace_store = trace_store + self._lmdeploy_actor: ray.actor.ActorHandle | None = None + + async def _build_worker_payload(self, rollout_request: _RolloutRequest) -> dict[str, Any]: + assert rollout_request.data is not None + input_ids = await self._prepare_input_ids(rollout_request) + return { + **{ + k: v + for k, v in rollout_request.data.items() + if k not in {"session_id", "messages", "logprobs", "top_logprobs"} + }, + "messages": [], + "input_ids": input_ids, + "return_token_ids": True, + "return_routed_experts": True, + "return_logprob": True, + "include_stop_str_in_output": True, + } + + async def _prepare_input_ids(self, rollout_request: _RolloutRequest) -> list[int]: + if rollout_request.session_id is None or rollout_request.messages is None: + raise RuntimeError("Trace-store requests require session_id and messages.") + + prompt_text = self.tokenizer.apply_chat_template( + canonicalize_messages_for_chat_template(rollout_request.messages), + tools=rollout_request.tools, + add_generation_prompt=True, + tokenize=False, + ) + + prefix, nodes = await self.trace_store.search.remote(rollout_request.session_id, prompt_text, filter_none=True) + if prefix: + get_logger().debug(f"Hit prefix cache for session {rollout_request.session_id}") + delta = prompt_text[len(prefix) :] + delta_ids = [] + if delta: + delta_ids = self.tokenizer.encode(delta, add_special_tokens=False) + await self.trace_store.insert.remote( + rollout_request.session_id, + prompt_text, + TokenizedSegment(text=delta, token_ids=delta_ids), + ) + return reduce(add, [node.value.token_ids for node in nodes] + [delta_ids]) + + async def handle_worker_response( + self, + rollout_request: _RolloutRequest, + worker_response: ClientResponse, + ) -> web.StreamResponse | web.Response: + if rollout_request.stream: + return await self._handle_stream_worker_response(rollout_request, worker_response) + + response, raw_response = await self._read_non_stream_response(rollout_request, worker_response) + try: + await self._record_raw_response(rollout_request, raw_response=raw_response, stream=False) + except Exception as exc: + return web.json_response( + { + "message": f"SessionServer response failed: {type(exc).__name__}: {exc}", + "type": "internal_server_error", + "code": 500, + "object": "error", + }, + status=500, + ) + return response + + async def _handle_stream_worker_response( + self, + rollout_request: _RolloutRequest, + worker_response: ClientResponse, + ) -> web.StreamResponse: + response, raw_response = await self._relay_stream_response( + rollout_request, + worker_response, + capture_raw_response=True, + delay_done=True, + ) + should_write_done = _has_complete_traceable_sse_response(raw_response) + error = None + if should_write_done: + try: + await self._record_raw_response(rollout_request, raw_response=raw_response, stream=True) + except Exception as exc: + error = exc + + try: + if error is not None: + error_payload = { + "message": f"SessionServer response failed: {type(error).__name__}: {error}", + "type": "internal_server_error", + "code": 500, + "object": "error", + } + await response.write( + ("data: " + json.dumps(error_payload, ensure_ascii=False) + "\n\n").encode("utf-8") + ) + elif should_write_done: + await response.write(b"data: [DONE]\n\n") + await response.write_eof() + except (ConnectionError, ClientConnectionResetError): + pass + return response + + async def _record_raw_response( + self, + rollout_request: _RolloutRequest, + *, + raw_response: bytes, + stream: bool, + ) -> None: + worker_response = self._parse_trace_response(raw=raw_response, stream=stream) + if worker_response is not None: + await self._record_response(rollout_request, worker_response) + + def _parse_trace_response(self, *, raw: bytes, stream: bool) -> dict[str, Any] | None: + if stream: + response = _parse_sse_to_complete_response(raw) + else: + response = json.loads(raw) + + if _is_error_payload(response): + return None + return self._validate_and_normalize_trace_response(response) + + def _validate_and_normalize_trace_response(self, response: dict[str, Any]) -> dict[str, Any]: + choices = response.get("choices") or [] + if not choices: + raise RuntimeError("SessionServer response has no choices; cannot export a training trace.") + choice = choices[0] + output_ids = choice.get("output_ids") + if output_ids is None: + raise RuntimeError( + "SessionServer response choice has no output_ids; " + "cannot export a training trace for this assistant turn." + ) + self._extract_output_logprobs(choice, output_ids) + return choice + + async def _record_response(self, rollout_request: _RolloutRequest, worker_response: dict[str, Any]) -> None: + if rollout_request.session_id is None or rollout_request.messages is None: + raise RuntimeError("Trace-store responses require session_id and messages.") + + output_ids = worker_response["output_ids"] + output_logprobs = self._extract_output_logprobs(worker_response, output_ids) + old_prompt = self.tokenizer.apply_chat_template( + canonicalize_messages_for_chat_template(rollout_request.messages), + tools=rollout_request.tools, + add_generation_prompt=True, + tokenize=False, + ) + messages = [*rollout_request.messages, worker_response["message"]] + new_prompt = ( + self.tokenizer.apply_chat_template( + canonicalize_messages_for_chat_template(messages), + tools=rollout_request.tools, + add_generation_prompt=False, + tokenize=False, + ) + ).rstrip() + assert new_prompt.startswith(old_prompt) and new_prompt.endswith(self.stop_word) + + routed_experts = None + if worker_response.get("routed_experts") is not None: + routed_experts = await self._decode_and_split_routed_experts( + session_id=rollout_request.session_id, + old_prompt=old_prompt, + output_ids=output_ids, + routed_experts=worker_response["routed_experts"], + ) + + await self.trace_store.insert.remote( + rollout_request.session_id, + key=new_prompt, + value=TokenizedSegment( + text=new_prompt[len(old_prompt) :], + token_ids=output_ids, + logprobs=output_logprobs, + labels=output_ids, + expert_key=routed_experts, + length=len(output_ids), + ), + ) + + async def _decode_and_split_routed_experts( + self, + *, + session_id: str, + old_prompt: str, + output_ids: list[int], + routed_experts: Any, + ) -> Any: + if isinstance(routed_experts, str): + if self._lmdeploy_actor is None: + self._lmdeploy_actor = ray.get_actor("shared_store", namespace="lmdeploy") + assert self._lmdeploy_actor is not None, "LMDeploy actor should be available in the shared store." + routed_experts = await self._lmdeploy_actor.get.remote(routed_experts) + raw_routed_expert = np.asarray(routed_experts) + if len(raw_routed_expert) > 0: + num_layers = raw_routed_expert.shape[1] + topk_experts = raw_routed_expert.shape[2] + dummy_expert = np.full((1, num_layers, topk_experts), 0, dtype=raw_routed_expert.dtype) + raw_routed_expert = np.concatenate([dummy_expert, raw_routed_expert], axis=0) + + _, nodes = await self.trace_store.search.remote(session_id, old_prompt, filter_none=True) + if not nodes: + return ray.put(raw_routed_expert) + + delta_node_val: TokenizedSegment = nodes[-1].value + delta_len = len(delta_node_val.token_ids) + prefix_len = sum(len(n.value.token_ids) for n in nodes[:-1]) + assert prefix_len + delta_len + len(output_ids) == len(raw_routed_expert) + + delta_expert = raw_routed_expert[prefix_len : prefix_len + delta_len] + response_expert = raw_routed_expert[prefix_len + delta_len :] + if delta_len > 0: + delta_node_val.expert_key = ray.put(delta_expert) + await self.trace_store.insert.remote(session_id, old_prompt, delta_node_val) + return ray.put(response_expert) + + def _extract_output_logprobs(self, choice: dict, output_token_ids: list[int]) -> list[float]: + if not output_token_ids: + return [] + + output_token_logprobs = choice.get("output_token_logprobs") + if output_token_logprobs is None: + raise RuntimeError( + "SessionServer response choice has no output_token_logprobs; " + "the return_logprob protocol is required for training traces." + ) + + logprob_token_ids = [item[1] for item in output_token_logprobs] + if logprob_token_ids != output_token_ids: + raise RuntimeError( + "SessionServer response choice has mismatched output_token_logprobs: " + f"output_ids_len={len(output_token_ids)}, logprob_ids_len={len(logprob_token_ids)}" + ) + return [item[0] for item in output_token_logprobs] + + +def _parse_sse_to_complete_response(raw: bytes) -> dict[str, Any]: + text = raw.decode("utf-8", errors="replace") + events = [] + saw_done = False + for line in text.split("\n"): + line = line.strip() + if line == "data: [DONE]": + saw_done = True + continue + if line.startswith("data: "): + event = json.loads(line[6:]) + if _is_error_payload(event): + raise RuntimeError(f"Upstream SSE stream returned error: {json.dumps(event, ensure_ascii=False)}") + events.append(event) + + if not events: + return {} + if not any(event.get("choices") for event in events): + raise RuntimeError(f"Upstream SSE stream ended without choices: {json.dumps(events, ensure_ascii=False)}") + + message: dict[str, Any] = {"choices": [{"message": {"role": "assistant", "content": ""}}]} + content_parts: list[str] = [] + tool_calls_map: dict[int, dict[str, Any]] = {} + usage: dict[str, Any] = {} + + for event in events: + if event.get("id") and "id" not in message: + message["id"] = event["id"] + if event.get("model"): + message["model"] = event["model"] + + for choice in event.get("choices", []): + if choice.get("finish_reason") == "error": + raise RuntimeError(f"Upstream SSE choice finished with error: {json.dumps(event, ensure_ascii=False)}") + delta = choice.get("delta", {}) + + if delta.get("content"): + content_parts.append(delta["content"]) + if choice.get("output_ids") is not None: + assistant_choice = message["choices"][0] + assistant_choice.setdefault("output_ids", []).extend(choice["output_ids"]) + if choice.get("routed_experts") is not None: + message["choices"][0]["routed_experts"] = choice["routed_experts"] + if choice.get("output_token_logprobs") is not None: + assistant_choice = message["choices"][0] + assistant_choice.setdefault("output_token_logprobs", []).extend(choice["output_token_logprobs"]) + if delta.get("reasoning_content"): + assistant_msg = message["choices"][0]["message"] + assistant_msg["reasoning_content"] = ( + assistant_msg.get("reasoning_content", "") + delta["reasoning_content"] + ) + + for tool_call_delta in delta.get("tool_calls") or []: + idx = tool_call_delta.get("index", 0) + tool_call = tool_calls_map.setdefault( + idx, + { + "id": tool_call_delta.get("id", ""), + "type": tool_call_delta.get("type", "function"), + "function": {"name": "", "arguments": ""}, + }, + ) + function_delta = tool_call_delta.get("function", {}) + if function_delta.get("name"): + tool_call["function"]["name"] += function_delta["name"] + if function_delta.get("arguments"): + tool_call["function"]["arguments"] += function_delta["arguments"] + + if choice.get("finish_reason"): + message["choices"][0]["finish_reason"] = choice["finish_reason"] + + if event.get("usage") is not None: + usage = event["usage"] + + assistant_message = message["choices"][0]["message"] + assistant_message["content"] = "".join(content_parts) + if tool_calls_map: + assistant_message["tool_calls"] = [tool_calls_map[i] for i in sorted(tool_calls_map)] + if usage: + message["usage"] = usage + + assistant_choice = message["choices"][0] + if not saw_done: + raise RuntimeError("Upstream SSE stream ended without [DONE].") + if not assistant_choice.get("finish_reason"): + raise RuntimeError("Upstream SSE stream ended without terminal finish_reason.") + if assistant_choice.get("output_ids") is None: + raise RuntimeError("Upstream SSE stream ended without output_ids.") + + return message + + +def _has_complete_traceable_sse_response(raw: bytes) -> bool: + text = raw.decode("utf-8", errors="replace") + has_choices = False + saw_done = False + saw_terminal_finish = False + for line in text.split("\n"): + line = line.strip() + if line == "data: [DONE]": + saw_done = True + continue + if line.startswith("data: "): + try: + event = json.loads(line[6:]) + except json.JSONDecodeError: + return False + if _is_error_payload(event): + return False + if event.get("choices"): + if any(choice.get("finish_reason") == "error" for choice in event.get("choices", [])): + return False + if any(choice.get("finish_reason") for choice in event.get("choices", [])): + saw_terminal_finish = True + has_choices = True + return has_choices and saw_done and saw_terminal_finish + + +def _is_error_payload(payload: dict[str, Any]) -> bool: + return payload.get("error") is not None or payload.get("type") == "error" or payload.get("object") == "error" + + +class _PrepareRequestError(RuntimeError): + def __init__(self, message: str, *, status: int = 500): + self.message = message + self.status = status + super().__init__(message) + + +def _filter_response_headers(headers: Any) -> dict[str, str]: + return { + k: v + for k, v in headers.items() + if k.lower() not in ("transfer-encoding", "content-length", "content-encoding") + }