Add scripted-runtime unit, core integration, and chunked-prefill tests (#27413)

This commit is contained in:
fzyzcjy
2026-06-06 09:08:35 +08:00
committed by GitHub
parent 5a82db85f0
commit bf4f2ccc78
10 changed files with 1912 additions and 0 deletions
@@ -0,0 +1,180 @@
from __future__ import annotations
import asyncio
import threading
import unittest
from concurrent.futures import Future
from unittest.mock import MagicMock
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.scripted_runtime import background_http_poster as bg_poster
from sglang.test.scripted_runtime.background_http_poster import BackgroundHttpPoster
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=8, suite="base-a-test-cpu")
class _FakeResponse:
def __init__(self) -> None:
self.read_called = False
async def read(self) -> bytes:
self.read_called = True
return b"chunk-1chunk-2"
class _FakePostCM:
def __init__(self, response: _FakeResponse) -> None:
self._response = response
async def __aenter__(self) -> _FakeResponse:
return self._response
async def __aexit__(self, *exc_info: object) -> bool:
return False
class _FakeSession:
def __init__(self) -> None:
self.closed = False
self.calls: list[tuple[str, object]] = []
self.response = _FakeResponse()
def post(self, url: str, json: object) -> _FakePostCM:
self.calls.append((url, json))
return _FakePostCM(self.response)
async def close(self) -> None:
self.closed = True
class TestBackgroundHttpPosterLifecycle(CustomTestCase):
def test_init_starts_running_loop_on_daemon_thread(self):
poster = BackgroundHttpPoster()
self.addCleanup(poster.close)
self.assertIsNotNone(poster._loop)
self.assertTrue(poster._loop.is_running())
self.assertTrue(poster._thread.is_alive())
self.assertTrue(poster._thread.daemon)
def test_close_stops_loop_and_joins_thread(self):
poster = BackgroundHttpPoster()
poster.close()
self.assertFalse(poster._loop.is_running())
self.assertFalse(poster._thread.is_alive())
def test_close_is_safe_when_loop_never_started(self):
poster = BackgroundHttpPoster.__new__(BackgroundHttpPoster)
poster._loop = None
poster._thread = None
poster._session = None
poster.close()
class TestBackgroundHttpPosterSubmitCoro(CustomTestCase):
def test_submit_coro_runs_on_background_loop_thread(self):
poster = BackgroundHttpPoster()
self.addCleanup(poster.close)
done = threading.Event()
recorded: dict[str, str] = {}
async def record_thread() -> None:
recorded["thread_name"] = threading.current_thread().name
done.set()
poster.submit_coro(record_thread())
self.assertTrue(done.wait(timeout=5.0))
self.assertEqual(recorded["thread_name"], "scripted-runtime-async")
def test_log_coro_exception_logs_real_failure(self):
future: Future = Future()
future.set_exception(RuntimeError("boom"))
original = bg_poster.logger.exception
bg_poster.logger.exception = MagicMock()
try:
BackgroundHttpPoster._log_coro_exception(future)
bg_poster.logger.exception.assert_called_once()
finally:
bg_poster.logger.exception = original
def test_log_coro_exception_swallows_cancellation_silently(self):
future: Future = Future()
future.set_exception(asyncio.CancelledError())
original = bg_poster.logger.exception
bg_poster.logger.exception = MagicMock()
try:
BackgroundHttpPoster._log_coro_exception(future)
bg_poster.logger.exception.assert_not_called()
finally:
bg_poster.logger.exception = original
def test_log_coro_exception_quiet_on_success(self):
future: Future = Future()
future.set_result(None)
original = bg_poster.logger.exception
bg_poster.logger.exception = MagicMock()
try:
BackgroundHttpPoster._log_coro_exception(future)
bg_poster.logger.exception.assert_not_called()
finally:
bg_poster.logger.exception = original
class TestBackgroundHttpPosterEnsureSession(CustomTestCase):
def test_ensure_session_creates_reuses_then_recreates_when_closed(self):
poster = BackgroundHttpPoster()
self.addCleanup(poster.close)
sessions = [MagicMock(closed=False), MagicMock(closed=False)]
original = bg_poster.aiohttp.ClientSession
original_connector = bg_poster.aiohttp.TCPConnector
bg_poster.aiohttp.ClientSession = MagicMock(side_effect=sessions)
bg_poster.aiohttp.TCPConnector = MagicMock()
try:
first = poster._ensure_session()
self.assertIs(first, sessions[0])
reused = poster._ensure_session()
self.assertIs(reused, sessions[0])
sessions[0].closed = True
recreated = poster._ensure_session()
self.assertIs(recreated, sessions[1])
finally:
bg_poster.aiohttp.ClientSession = original
bg_poster.aiohttp.TCPConnector = original_connector
poster._session = None
class TestBackgroundHttpPosterPost(CustomTestCase):
def _run_on_loop(self, poster: BackgroundHttpPoster, coro) -> None:
asyncio.run_coroutine_threadsafe(coro, poster._loop).result(timeout=5.0)
def test_post_posts_json_and_reads_body(self):
poster = BackgroundHttpPoster()
self.addCleanup(poster.close)
session = _FakeSession()
poster._ensure_session = lambda: session
self._run_on_loop(poster, poster.post("http://h/flush", {"a": 1}))
self.assertEqual(session.calls, [("http://h/flush", {"a": 1})])
self.assertTrue(session.response.read_called)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,128 @@
from __future__ import annotations
import unittest
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.scripted_runtime.http_server import ScriptedHttpServer
from sglang.test.scripted_runtime.io_struct import (
HookReady,
RunScript,
ScriptFailed,
ScriptSucceeded,
)
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=15, suite="base-a-test-cpu")
def _sample_script(ctx, *args):
yield
_EXPECTED_FN_PATH = f"{_sample_script.__module__}:{_sample_script.__qualname__}"
class _FakePairSocket:
def __init__(self, *, poll_result: bool, reply: object = None) -> None:
self._poll_result = poll_result
self._reply = reply
self.sent: list = []
def send_pyobj(self, obj: object) -> None:
self.sent.append(obj)
def poll(self, timeout_ms: int) -> bool:
return self._poll_result
def recv_pyobj(self) -> object:
return self._reply
class _FakeProcess:
def __init__(self, *, alive: bool) -> None:
self._alive = alive
def is_alive(self) -> bool:
return self._alive
def _make_server(socket: _FakePairSocket, process: _FakeProcess) -> ScriptedHttpServer:
server = ScriptedHttpServer.__new__(ScriptedHttpServer)
server._socket = socket
server._server_process = process
server._dirty = None
return server
class TestExecuteScriptReplyMatching(CustomTestCase):
def test_returns_on_script_succeeded(self):
socket = _FakePairSocket(poll_result=True, reply=ScriptSucceeded())
server = _make_server(socket, _FakeProcess(alive=True))
server.execute_script(_sample_script)
self.assertEqual(socket.sent, [RunScript(fn_path=_EXPECTED_FN_PATH, args=())])
def test_forwards_args_in_run_script(self):
socket = _FakePairSocket(poll_result=True, reply=ScriptSucceeded())
server = _make_server(socket, _FakeProcess(alive=True))
server.execute_script(_sample_script, args=(1, "two"))
self.assertEqual(
socket.sent, [RunScript(fn_path=_EXPECTED_FN_PATH, args=(1, "two"))]
)
def test_script_failed_reply_raises_assertion_with_traceback(self):
socket = _FakePairSocket(
poll_result=True, reply=ScriptFailed(traceback="REMOTE-TB-MARKER")
)
server = _make_server(socket, _FakeProcess(alive=True))
with self.assertRaisesRegex(AssertionError, "REMOTE-TB-MARKER"):
server.execute_script(_sample_script)
def test_unexpected_reply_raises_runtime_error(self):
socket = _FakePairSocket(poll_result=True, reply=HookReady())
server = _make_server(socket, _FakeProcess(alive=True))
with self.assertRaisesRegex(RuntimeError, "unexpected message"):
server.execute_script(_sample_script)
class TestExecuteScriptNoReply(CustomTestCase):
def test_timeout_when_process_still_alive(self):
socket = _FakePairSocket(poll_result=False)
server = _make_server(socket, _FakeProcess(alive=True))
with self.assertRaisesRegex(TimeoutError, "timed out"):
server.execute_script(_sample_script, timeout_s=0.01)
self.assertIn("timed out", server._dirty)
def test_runtime_error_when_process_died(self):
socket = _FakePairSocket(poll_result=False)
server = _make_server(socket, _FakeProcess(alive=False))
with self.assertRaisesRegex(RuntimeError, "died before responding"):
server.execute_script(_sample_script, timeout_s=0.01)
self.assertIn("died before responding", server._dirty)
class TestExecuteScriptDirtyGuard(CustomTestCase):
def test_refuses_to_run_when_already_dirty(self):
socket = _FakePairSocket(poll_result=True, reply=ScriptSucceeded())
server = _make_server(socket, _FakeProcess(alive=True))
server._dirty = "prior timeout"
with self.assertRaisesRegex(RuntimeError, "dirty"):
server.execute_script(_sample_script)
self.assertEqual(socket.sent, [])
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,56 @@
from __future__ import annotations
import unittest
from unittest.mock import MagicMock
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.scripted_runtime import scheduler_hook
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=10, suite="base-a-test-cpu")
def _yielding_gen():
yield
yield
def _empty_gen():
return
yield # pragma: no cover — makes this a generator function
def _raising_gen():
raise ValueError("scripted-boom")
yield # pragma: no cover — makes this a generator function
class TestAdvanceGenerator(CustomTestCase):
def test_not_done_when_generator_yields(self):
done, exc_tb = scheduler_hook._advance_generator(_yielding_gen())
self.assertEqual((done, exc_tb), (False, None))
def test_done_without_traceback_on_stop_iteration(self):
done, exc_tb = scheduler_hook._advance_generator(_empty_gen())
self.assertEqual((done, exc_tb), (True, None))
def test_done_with_traceback_on_exception(self):
original = scheduler_hook.logger.exception
scheduler_hook.logger.exception = MagicMock()
try:
done, exc_tb = scheduler_hook._advance_generator(_raising_gen())
scheduler_hook.logger.exception.assert_called_once()
finally:
scheduler_hook.logger.exception = original
self.assertTrue(done)
self.assertIsNotNone(exc_tb)
self.assertIn("ValueError", exc_tb)
self.assertIn("scripted-boom", exc_tb)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,79 @@
from __future__ import annotations
import json
import os
import sys
import unittest
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.scripted_runtime.utils import ensure_script_importable, resolve_fn
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=2, suite="base-a-test-cpu")
class TestResolveFn(CustomTestCase):
def test_resolves_top_level_function(self):
self.assertIs(resolve_fn("json:dumps"), json.dumps)
def test_resolves_nested_attribute_path(self):
self.assertIs(resolve_fn("os:path.join"), os.path.join)
def test_rejects_missing_colon(self):
with self.assertRaisesRegex(ValueError, "module.path:function_name"):
resolve_fn("json.dumps")
def test_rejects_empty_module(self):
with self.assertRaisesRegex(ValueError, "module.path:function_name"):
resolve_fn(":dumps")
def test_rejects_empty_function(self):
with self.assertRaisesRegex(ValueError, "module.path:function_name"):
resolve_fn("json:")
def test_rejects_non_callable_target(self):
with self.assertRaisesRegex(TypeError, "not callable"):
resolve_fn("math:pi")
def test_propagates_missing_module_error(self):
with self.assertRaises(ModuleNotFoundError):
resolve_fn("sglang_no_such_module_zzz:foo")
def test_propagates_missing_attribute_error(self):
with self.assertRaises(AttributeError):
resolve_fn("json:no_such_attribute")
class TestEnsureScriptImportable(CustomTestCase):
_FAKE_ENTRY = "/tmp/__scripted_runtime_ut_fake_sys_path__"
def setUp(self):
self._orig_path = list(sys.path)
def tearDown(self):
sys.path[:] = self._orig_path
def test_inserts_new_entry_at_front(self):
self.assertNotIn(self._FAKE_ENTRY, sys.path)
ensure_script_importable(self._FAKE_ENTRY)
self.assertEqual(sys.path[0], self._FAKE_ENTRY)
def test_noop_when_entry_is_none(self):
ensure_script_importable(None)
self.assertEqual(sys.path, self._orig_path)
def test_noop_when_entry_already_present(self):
sys.path.insert(0, self._FAKE_ENTRY)
ensure_script_importable(self._FAKE_ENTRY)
self.assertEqual(sys.path.count(self._FAKE_ENTRY), 1)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,159 @@
from __future__ import annotations
from collections import deque
from dataclasses import dataclass
import zmq
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.scripted_runtime.tokenizer_recv_proxy import (
ScriptedTokenizerRecvProxy,
)
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
import unittest
@dataclass
class _ControlMsg:
tag: str = "flush"
@dataclass
class _StartReq:
rid: str
class _FakeUnderlyingSocket:
def __init__(self) -> None:
self._ready: deque = deque()
self._scheduled: list[list] = []
def feed(self, obj: object) -> None:
self._ready.append(obj)
def feed_after_drain_cycles(self, obj: object, *, cycles: int) -> None:
self._scheduled.append([cycles, obj])
def recv_pyobj(self, flags: int = 0) -> object:
if self._ready:
return self._ready.popleft()
for entry in self._scheduled:
entry[0] -= 1
ready_now = [obj for remaining, obj in self._scheduled if remaining <= 0]
self._scheduled = [entry for entry in self._scheduled if entry[0] > 0]
self._ready.extend(ready_now)
raise zmq.ZMQError(zmq.EAGAIN, "Resource temporarily unavailable")
def _is_control(obj: object) -> bool:
return isinstance(obj, _ControlMsg)
def _is_start_req(rid: str):
return lambda obj: isinstance(obj, _StartReq) and obj.rid == rid
class TestScriptedTokenizerRecvProxyRecv(CustomTestCase):
def test_recv_pyobj_drains_then_pops_fifo(self):
underlying = _FakeUnderlyingSocket()
proxy = ScriptedTokenizerRecvProxy(underlying=underlying)
first, second = _ControlMsg("a"), _ControlMsg("b")
underlying.feed(first)
underlying.feed(second)
self.assertIs(proxy.recv_pyobj(), first)
self.assertIs(proxy.recv_pyobj(), second)
def test_recv_pyobj_empty_noblock_raises_eagain(self):
proxy = ScriptedTokenizerRecvProxy(underlying=_FakeUnderlyingSocket())
with self.assertRaises(zmq.ZMQError) as ctx:
proxy.recv_pyobj(zmq.NOBLOCK)
self.assertEqual(ctx.exception.errno, zmq.EAGAIN)
def test_recv_pyobj_empty_blocking_raises_runtime_error(self):
proxy = ScriptedTokenizerRecvProxy(underlying=_FakeUnderlyingSocket())
with self.assertRaisesRegex(RuntimeError, "blocking recv is not supported"):
proxy.recv_pyobj()
class TestScriptedTokenizerRecvProxyWaitUntilArrived(CustomTestCase):
def _proxy_with_stale_control(self):
underlying = _FakeUnderlyingSocket()
proxy = ScriptedTokenizerRecvProxy(underlying=underlying)
stale = _ControlMsg("stale")
underlying.feed(stale)
proxy.wait_until_arrived(_is_control, timeout_s=1.0)
return proxy, underlying, stale
def test_wait_until_arrived_returns_on_first_match_when_buffer_empty(self):
underlying = _FakeUnderlyingSocket()
proxy = ScriptedTokenizerRecvProxy(underlying=underlying)
msg = _ControlMsg("first")
underlying.feed(msg)
proxy.wait_until_arrived(_is_control, timeout_s=1.0)
self.assertIs(proxy.recv_pyobj(), msg)
def test_wait_until_arrived_skips_stale_same_type_object(self):
proxy, _, _ = self._proxy_with_stale_control()
with self.assertRaises(TimeoutError):
proxy.wait_until_arrived(_is_control, timeout_s=0.05)
def test_wait_until_arrived_returns_on_new_object_after_stale(self):
proxy, underlying, stale = self._proxy_with_stale_control()
fresh = _ControlMsg("fresh")
underlying.feed_after_drain_cycles(fresh, cycles=1)
proxy.wait_until_arrived(_is_control, timeout_s=2.0)
self.assertIs(proxy.recv_pyobj(), stale)
self.assertIs(proxy.recv_pyobj(), fresh)
def test_wait_until_arrived_rid_predicate_ignores_stale_other_rid(self):
underlying = _FakeUnderlyingSocket()
proxy = ScriptedTokenizerRecvProxy(underlying=underlying)
old = _StartReq(rid="old")
underlying.feed(old)
proxy.wait_until_arrived(_is_start_req("old"), timeout_s=1.0)
new = _StartReq(rid="new")
underlying.feed(new)
proxy.wait_until_arrived(_is_start_req("new"), timeout_s=1.0)
self.assertIs(proxy.recv_pyobj(), old)
self.assertIs(proxy.recv_pyobj(), new)
def test_wait_until_arrived_rid_predicate_skips_stale_same_rid(self):
underlying = _FakeUnderlyingSocket()
proxy = ScriptedTokenizerRecvProxy(underlying=underlying)
underlying.feed(_StartReq(rid="reused"))
proxy.wait_until_arrived(_is_start_req("reused"), timeout_s=1.0)
with self.assertRaises(TimeoutError):
proxy.wait_until_arrived(_is_start_req("reused"), timeout_s=0.05)
def test_wait_until_arrived_timeout_message_names_description(self):
proxy = ScriptedTokenizerRecvProxy(underlying=_FakeUnderlyingSocket())
with self.assertRaisesRegex(TimeoutError, "FlushCacheReqInput"):
proxy.wait_until_arrived(
_is_control, timeout_s=0.02, description="FlushCacheReqInput"
)
if __name__ == "__main__":
unittest.main()