Support patching source code (#19561)
This commit is contained in:
66
test/registered/debug_utils/source_patcher/conftest.py
Normal file
66
test/registered/debug_utils/source_patcher/conftest.py
Normal file
@@ -0,0 +1,66 @@
|
||||
"""Shared fixtures for source_patcher tests.
|
||||
|
||||
The sample module is defined as an inline string and written to a temp file
|
||||
at test time, avoiding CI complaints about fixture files without test registration.
|
||||
"""
|
||||
|
||||
import importlib.util
|
||||
import sys
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
from types import ModuleType
|
||||
|
||||
import pytest
|
||||
|
||||
SAMPLE_MODULE_NAME = "_source_patcher_test_fixtures.sample_module"
|
||||
|
||||
SAMPLE_MODULE_SOURCE = '''\
|
||||
GLOBAL_VAR = "global_value"
|
||||
|
||||
|
||||
class HelperClass:
|
||||
"""Utility class referenced by SampleClass to test cross-class calls."""
|
||||
|
||||
@staticmethod
|
||||
def format_value(value: str) -> str:
|
||||
return f"[{value}]"
|
||||
|
||||
|
||||
class SampleClass:
|
||||
def greet(self, name: str) -> str:
|
||||
greeting = f"hello {name}"
|
||||
return greeting
|
||||
|
||||
def compute(self, x: int) -> int:
|
||||
result = x * 2 + 1
|
||||
return result
|
||||
|
||||
def uses_global(self) -> str:
|
||||
return f"value={GLOBAL_VAR}"
|
||||
|
||||
def uses_helper(self, value: str) -> str:
|
||||
return HelperClass.format_value(value)
|
||||
|
||||
|
||||
def standalone_function(a: int, b: int) -> int:
|
||||
return a + b
|
||||
'''
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def sample_module() -> ModuleType:
|
||||
"""Load the sample module from a temp file and register it in sys.modules."""
|
||||
if SAMPLE_MODULE_NAME in sys.modules:
|
||||
return sys.modules[SAMPLE_MODULE_NAME]
|
||||
|
||||
tmpdir = tempfile.mkdtemp(prefix="source_patcher_fixtures_")
|
||||
module_path = Path(tmpdir) / "sample_module.py"
|
||||
module_path.write_text(SAMPLE_MODULE_SOURCE)
|
||||
|
||||
spec = importlib.util.spec_from_file_location(SAMPLE_MODULE_NAME, module_path)
|
||||
assert spec is not None
|
||||
assert spec.loader is not None
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
sys.modules[SAMPLE_MODULE_NAME] = module
|
||||
spec.loader.exec_module(module)
|
||||
return module
|
||||
253
test/registered/debug_utils/source_patcher/test_code_patcher.py
Normal file
253
test/registered/debug_utils/source_patcher/test_code_patcher.py
Normal file
@@ -0,0 +1,253 @@
|
||||
from types import ModuleType
|
||||
|
||||
import pytest
|
||||
|
||||
from sglang.srt.debug_utils.source_patcher.code_patcher import (
|
||||
CodePatcher,
|
||||
_resolve_target,
|
||||
patch_function,
|
||||
)
|
||||
from sglang.srt.debug_utils.source_patcher.types import EditSpec, PatchSpec
|
||||
|
||||
SAMPLE_MODULE_NAME = "_source_patcher_test_fixtures.sample_module"
|
||||
|
||||
|
||||
class TestPatchFunction:
|
||||
def test_basic_patch_changes_behavior(self, sample_module: ModuleType) -> None:
|
||||
cls = sample_module.SampleClass
|
||||
obj = cls()
|
||||
assert obj.greet("world") == "hello world"
|
||||
|
||||
state = patch_function(
|
||||
target=cls.greet,
|
||||
edits=[
|
||||
EditSpec(
|
||||
match='greeting = f"hello {name}"',
|
||||
replacement='greeting = f"patched {name}"',
|
||||
)
|
||||
],
|
||||
)
|
||||
try:
|
||||
assert obj.greet("world") == "patched world"
|
||||
finally:
|
||||
state.restore()
|
||||
|
||||
assert obj.greet("world") == "hello world"
|
||||
|
||||
def test_globals_preserved_after_patch(self, sample_module: ModuleType) -> None:
|
||||
cls = sample_module.SampleClass
|
||||
obj = cls()
|
||||
assert obj.uses_global() == "value=global_value"
|
||||
|
||||
state = patch_function(
|
||||
target=cls.uses_global,
|
||||
edits=[
|
||||
EditSpec(
|
||||
match='return f"value={GLOBAL_VAR}"',
|
||||
replacement='return f"patched_value={GLOBAL_VAR}"',
|
||||
)
|
||||
],
|
||||
)
|
||||
try:
|
||||
assert obj.uses_global() == "patched_value=global_value"
|
||||
finally:
|
||||
state.restore()
|
||||
|
||||
def test_function_identity_preserved(self, sample_module: ModuleType) -> None:
|
||||
cls = sample_module.SampleClass
|
||||
fn_id_before = id(cls.greet)
|
||||
|
||||
state = patch_function(
|
||||
target=cls.greet,
|
||||
edits=[
|
||||
EditSpec(
|
||||
match='greeting = f"hello {name}"',
|
||||
replacement='greeting = f"patched {name}"',
|
||||
)
|
||||
],
|
||||
)
|
||||
try:
|
||||
assert id(cls.greet) == fn_id_before
|
||||
finally:
|
||||
state.restore()
|
||||
|
||||
def test_patch_standalone_function(self, sample_module: ModuleType) -> None:
|
||||
fn = sample_module.standalone_function
|
||||
assert fn(2, 3) == 5
|
||||
|
||||
state = patch_function(
|
||||
target=fn,
|
||||
edits=[
|
||||
EditSpec(
|
||||
match="return a + b",
|
||||
replacement="return a * b",
|
||||
)
|
||||
],
|
||||
)
|
||||
try:
|
||||
assert fn(2, 3) == 6
|
||||
finally:
|
||||
state.restore()
|
||||
|
||||
assert fn(2, 3) == 5
|
||||
|
||||
def test_patched_code_can_reference_global_variable(
|
||||
self, sample_module: ModuleType
|
||||
) -> None:
|
||||
"""Replacement code that references a module-level global should work."""
|
||||
cls = sample_module.SampleClass
|
||||
obj = cls()
|
||||
|
||||
state = patch_function(
|
||||
target=cls.greet,
|
||||
edits=[
|
||||
EditSpec(
|
||||
match='greeting = f"hello {name}"',
|
||||
replacement='greeting = f"{GLOBAL_VAR} {name}"',
|
||||
)
|
||||
],
|
||||
)
|
||||
try:
|
||||
assert obj.greet("world") == "global_value world"
|
||||
finally:
|
||||
state.restore()
|
||||
|
||||
def test_patched_code_can_call_another_class_method(
|
||||
self, sample_module: ModuleType
|
||||
) -> None:
|
||||
"""Replacement code that calls HelperClass.format_value should work."""
|
||||
cls = sample_module.SampleClass
|
||||
obj = cls()
|
||||
|
||||
state = patch_function(
|
||||
target=cls.greet,
|
||||
edits=[
|
||||
EditSpec(
|
||||
match='greeting = f"hello {name}"',
|
||||
replacement="greeting = HelperClass.format_value(name)",
|
||||
)
|
||||
],
|
||||
)
|
||||
try:
|
||||
assert obj.greet("world") == "[world]"
|
||||
finally:
|
||||
state.restore()
|
||||
|
||||
def test_patched_code_uses_helper_via_existing_method(
|
||||
self, sample_module: ModuleType
|
||||
) -> None:
|
||||
"""The uses_helper method already calls HelperClass; verify it survives patching."""
|
||||
cls = sample_module.SampleClass
|
||||
obj = cls()
|
||||
assert obj.uses_helper("test") == "[test]"
|
||||
|
||||
state = patch_function(
|
||||
target=cls.uses_helper,
|
||||
edits=[
|
||||
EditSpec(
|
||||
match="return HelperClass.format_value(value)",
|
||||
replacement='return HelperClass.format_value("patched_" + value)',
|
||||
)
|
||||
],
|
||||
)
|
||||
try:
|
||||
assert obj.uses_helper("test") == "[patched_test]"
|
||||
finally:
|
||||
state.restore()
|
||||
|
||||
assert obj.uses_helper("test") == "[test]"
|
||||
|
||||
|
||||
class TestResolveTarget:
|
||||
def test_resolve_class_method(self, sample_module: ModuleType) -> None:
|
||||
target = _resolve_target(f"{SAMPLE_MODULE_NAME}.SampleClass.greet")
|
||||
assert target is sample_module.SampleClass.greet
|
||||
|
||||
def test_resolve_standalone_function(self, sample_module: ModuleType) -> None:
|
||||
target = _resolve_target(f"{SAMPLE_MODULE_NAME}.standalone_function")
|
||||
assert target is sample_module.standalone_function
|
||||
|
||||
def test_resolve_nonexistent_raises(self, sample_module: ModuleType) -> None:
|
||||
with pytest.raises((ImportError, AttributeError)):
|
||||
_resolve_target(f"{SAMPLE_MODULE_NAME}.NonexistentClass.method")
|
||||
|
||||
|
||||
class TestCodePatcher:
|
||||
def test_context_manager_patches_and_restores(
|
||||
self, sample_module: ModuleType
|
||||
) -> None:
|
||||
cls = sample_module.SampleClass
|
||||
obj = cls()
|
||||
assert obj.greet("world") == "hello world"
|
||||
|
||||
patches = [
|
||||
PatchSpec(
|
||||
target=f"{SAMPLE_MODULE_NAME}.SampleClass.greet",
|
||||
edits=[
|
||||
EditSpec(
|
||||
match='greeting = f"hello {name}"',
|
||||
replacement='greeting = f"ctx_patched {name}"',
|
||||
)
|
||||
],
|
||||
)
|
||||
]
|
||||
|
||||
with CodePatcher(patches=patches):
|
||||
assert obj.greet("world") == "ctx_patched world"
|
||||
|
||||
assert obj.greet("world") == "hello world"
|
||||
|
||||
def test_context_manager_multiple_patches(self, sample_module: ModuleType) -> None:
|
||||
cls = sample_module.SampleClass
|
||||
obj = cls()
|
||||
|
||||
patches = [
|
||||
PatchSpec(
|
||||
target=f"{SAMPLE_MODULE_NAME}.SampleClass.greet",
|
||||
edits=[
|
||||
EditSpec(
|
||||
match='greeting = f"hello {name}"',
|
||||
replacement='greeting = f"p1 {name}"',
|
||||
)
|
||||
],
|
||||
),
|
||||
PatchSpec(
|
||||
target=f"{SAMPLE_MODULE_NAME}.SampleClass.compute",
|
||||
edits=[
|
||||
EditSpec(
|
||||
match="result = x * 2 + 1",
|
||||
replacement="result = x * 100",
|
||||
)
|
||||
],
|
||||
),
|
||||
]
|
||||
|
||||
with CodePatcher(patches=patches):
|
||||
assert obj.greet("world") == "p1 world"
|
||||
assert obj.compute(5) == 500
|
||||
|
||||
assert obj.greet("world") == "hello world"
|
||||
assert obj.compute(5) == 11
|
||||
|
||||
def test_restores_on_exception(self, sample_module: ModuleType) -> None:
|
||||
cls = sample_module.SampleClass
|
||||
obj = cls()
|
||||
|
||||
patches = [
|
||||
PatchSpec(
|
||||
target=f"{SAMPLE_MODULE_NAME}.SampleClass.greet",
|
||||
edits=[
|
||||
EditSpec(
|
||||
match='greeting = f"hello {name}"',
|
||||
replacement='greeting = f"err_patched {name}"',
|
||||
)
|
||||
],
|
||||
)
|
||||
]
|
||||
|
||||
with pytest.raises(RuntimeError):
|
||||
with CodePatcher(patches=patches):
|
||||
assert obj.greet("world") == "err_patched world"
|
||||
raise RuntimeError("test error")
|
||||
|
||||
assert obj.greet("world") == "hello world"
|
||||
289
test/registered/debug_utils/source_patcher/test_source_editor.py
Normal file
289
test/registered/debug_utils/source_patcher/test_source_editor.py
Normal file
@@ -0,0 +1,289 @@
|
||||
import pytest
|
||||
from pydantic import ValidationError
|
||||
|
||||
from sglang.srt.debug_utils.source_patcher.source_editor import apply_edits
|
||||
from sglang.srt.debug_utils.source_patcher.types import EditSpec, PatchApplicationError
|
||||
|
||||
|
||||
class TestApplyEdits:
|
||||
"""Tests for the apply_edits() source text transformation function."""
|
||||
|
||||
def test_single_line_match_to_multiline_replacement(self) -> None:
|
||||
source = "def foo():\n" " x = compute()\n" " return x\n"
|
||||
edits = [
|
||||
EditSpec(
|
||||
match="x = compute()",
|
||||
replacement="x = compute()\nprint(x)",
|
||||
)
|
||||
]
|
||||
result = apply_edits(source=source, edits=edits)
|
||||
assert result == (
|
||||
"def foo():\n" " x = compute()\n" " print(x)\n" " return x\n"
|
||||
)
|
||||
|
||||
def test_pure_insertion(self) -> None:
|
||||
source = "def foo():\n" " a = 1\n" " b = 2\n"
|
||||
edits = [
|
||||
EditSpec(
|
||||
match="a = 1",
|
||||
replacement="a = 1\nprint(a)",
|
||||
)
|
||||
]
|
||||
result = apply_edits(source=source, edits=edits)
|
||||
assert result == ("def foo():\n" " a = 1\n" " print(a)\n" " b = 2\n")
|
||||
|
||||
def test_pure_deletion_via_empty_replacement(self) -> None:
|
||||
source = "def foo():\n" " debug_log()\n" " return 42\n"
|
||||
edits = [
|
||||
EditSpec(
|
||||
match="debug_log()",
|
||||
replacement="",
|
||||
)
|
||||
]
|
||||
result = apply_edits(source=source, edits=edits)
|
||||
assert result == ("def foo():\n" " return 42\n")
|
||||
|
||||
def test_deletion_fewer_lines(self) -> None:
|
||||
source = "def foo():\n" " a = 1\n" " b = 2\n" " c = 3\n"
|
||||
edits = [
|
||||
EditSpec(
|
||||
match="a = 1\nb = 2",
|
||||
replacement="ab = 3",
|
||||
)
|
||||
]
|
||||
result = apply_edits(source=source, edits=edits)
|
||||
assert result == ("def foo():\n" " ab = 3\n" " c = 3\n")
|
||||
|
||||
def test_multiline_match_to_multiline_replacement(self) -> None:
|
||||
source = (
|
||||
"def foo():\n"
|
||||
" result = self.attn(\n"
|
||||
" q=q,\n"
|
||||
" k=k,\n"
|
||||
" )\n"
|
||||
" return result\n"
|
||||
)
|
||||
edits = [
|
||||
EditSpec(
|
||||
match="result = self.attn(\n q=q,\n k=k,\n)",
|
||||
replacement="result = self.attn(\n q=q,\n k=k,\n v=v,\n)",
|
||||
)
|
||||
]
|
||||
result = apply_edits(source=source, edits=edits)
|
||||
assert result == (
|
||||
"def foo():\n"
|
||||
" result = self.attn(\n"
|
||||
" q=q,\n"
|
||||
" k=k,\n"
|
||||
" v=v,\n"
|
||||
" )\n"
|
||||
" return result\n"
|
||||
)
|
||||
|
||||
def test_indent_alignment_deep_nesting(self) -> None:
|
||||
source = (
|
||||
"class Foo:\n"
|
||||
" class Bar:\n"
|
||||
" def method(self):\n"
|
||||
" x = compute()\n"
|
||||
" return x\n"
|
||||
)
|
||||
edits = [
|
||||
EditSpec(
|
||||
match="x = compute()",
|
||||
replacement="x = compute()\nprint(x)",
|
||||
)
|
||||
]
|
||||
result = apply_edits(source=source, edits=edits)
|
||||
assert result == (
|
||||
"class Foo:\n"
|
||||
" class Bar:\n"
|
||||
" def method(self):\n"
|
||||
" x = compute()\n"
|
||||
" print(x)\n"
|
||||
" return x\n"
|
||||
)
|
||||
|
||||
def test_match_not_found_raises(self) -> None:
|
||||
source = "def foo():\n return 1\n"
|
||||
edits = [EditSpec(match="nonexistent_call()", replacement="replaced()")]
|
||||
with pytest.raises(PatchApplicationError, match="not found"):
|
||||
apply_edits(source=source, edits=edits)
|
||||
|
||||
def test_match_found_multiple_times_raises(self) -> None:
|
||||
source = "def foo():\n" " print(1)\n" " print(1)\n"
|
||||
edits = [EditSpec(match="print(1)", replacement="print(2)")]
|
||||
with pytest.raises(PatchApplicationError, match="multiple"):
|
||||
apply_edits(source=source, edits=edits)
|
||||
|
||||
def test_multiple_edits_applied_sequentially(self) -> None:
|
||||
source = "def foo():\n" " a = 1\n" " b = 2\n" " return a + b\n"
|
||||
edits = [
|
||||
EditSpec(match="a = 1", replacement="a = 10"),
|
||||
EditSpec(match="b = 2", replacement="b = 20"),
|
||||
]
|
||||
result = apply_edits(source=source, edits=edits)
|
||||
assert result == (
|
||||
"def foo():\n" " a = 10\n" " b = 20\n" " return a + b\n"
|
||||
)
|
||||
|
||||
def test_strip_matching_ignores_leading_trailing_whitespace(self) -> None:
|
||||
source = "def foo():\n" " x = compute()\n" " return x\n"
|
||||
edits = [
|
||||
EditSpec(
|
||||
match=" x = compute() ",
|
||||
replacement="x = replaced()",
|
||||
)
|
||||
]
|
||||
result = apply_edits(source=source, edits=edits)
|
||||
assert result == ("def foo():\n" " x = replaced()\n" " return x\n")
|
||||
|
||||
def test_replacement_indented_text_realigned(self) -> None:
|
||||
"""replacement text with its own indentation gets realigned to match source."""
|
||||
source = "def foo():\n" " x = compute()\n" " return x\n"
|
||||
edits = [
|
||||
EditSpec(
|
||||
match="x = compute()",
|
||||
replacement="x = compute()\nprint(x)",
|
||||
)
|
||||
]
|
||||
result = apply_edits(source=source, edits=edits)
|
||||
assert result == (
|
||||
"def foo():\n"
|
||||
" x = compute()\n"
|
||||
" print(x)\n"
|
||||
" return x\n"
|
||||
)
|
||||
|
||||
def test_replacement_with_existing_indent_realigned(self) -> None:
|
||||
"""replacement text already has indentation that should be rebased."""
|
||||
source = "def foo():\n" " if True:\n" " x = 1\n" " return x\n"
|
||||
edits = [
|
||||
EditSpec(
|
||||
match="x = 1",
|
||||
replacement="x = 1\nif x > 0:\n print(x)",
|
||||
)
|
||||
]
|
||||
result = apply_edits(source=source, edits=edits)
|
||||
assert result == (
|
||||
"def foo():\n"
|
||||
" if True:\n"
|
||||
" x = 1\n"
|
||||
" if x > 0:\n"
|
||||
" print(x)\n"
|
||||
" return x\n"
|
||||
)
|
||||
|
||||
def test_append_keeps_match_and_adds_after(self) -> None:
|
||||
source = "def foo():\n" " x = compute()\n" " return x\n"
|
||||
edits = [EditSpec(match="x = compute()", append="print(x)")]
|
||||
result = apply_edits(source=source, edits=edits)
|
||||
assert result == (
|
||||
"def foo():\n" " x = compute()\n" " print(x)\n" " return x\n"
|
||||
)
|
||||
|
||||
def test_append_multiline_match(self) -> None:
|
||||
source = (
|
||||
"def foo():\n"
|
||||
" result = call(\n"
|
||||
" a=1,\n"
|
||||
" b=2,\n"
|
||||
" )\n"
|
||||
" return result\n"
|
||||
)
|
||||
edits = [
|
||||
EditSpec(
|
||||
match="result = call(\n a=1,\n b=2,\n)",
|
||||
append="dumper.dump('result', result)",
|
||||
)
|
||||
]
|
||||
result = apply_edits(source=source, edits=edits)
|
||||
assert result == (
|
||||
"def foo():\n"
|
||||
" result = call(\n"
|
||||
" a=1,\n"
|
||||
" b=2,\n"
|
||||
" )\n"
|
||||
" dumper.dump('result', result)\n"
|
||||
" return result\n"
|
||||
)
|
||||
|
||||
def test_prepend_adds_before_match(self) -> None:
|
||||
source = "def foo():\n" " x = compute()\n" " return x\n"
|
||||
edits = [EditSpec(match="x = compute()", prepend="print('before')")]
|
||||
result = apply_edits(source=source, edits=edits)
|
||||
assert result == (
|
||||
"def foo():\n"
|
||||
" print('before')\n"
|
||||
" x = compute()\n"
|
||||
" return x\n"
|
||||
)
|
||||
|
||||
def test_prepend_multiline(self) -> None:
|
||||
source = "def foo():\n" " return x\n"
|
||||
edits = [EditSpec(match="return x", prepend="a = 1\nb = 2")]
|
||||
result = apply_edits(source=source, edits=edits)
|
||||
assert result == ("def foo():\n" " a = 1\n" " b = 2\n" " return x\n")
|
||||
|
||||
def test_prepend_deep_indent(self) -> None:
|
||||
source = (
|
||||
"class Foo:\n"
|
||||
" class Bar:\n"
|
||||
" def method(self):\n"
|
||||
" return x\n"
|
||||
)
|
||||
edits = [EditSpec(match="return x", prepend="dumper.dump('x', x)")]
|
||||
result = apply_edits(source=source, edits=edits)
|
||||
assert result == (
|
||||
"class Foo:\n"
|
||||
" class Bar:\n"
|
||||
" def method(self):\n"
|
||||
" dumper.dump('x', x)\n"
|
||||
" return x\n"
|
||||
)
|
||||
|
||||
def test_prepend_multiline_match(self) -> None:
|
||||
source = (
|
||||
"def foo():\n"
|
||||
" result = call(\n"
|
||||
" a=1,\n"
|
||||
" )\n"
|
||||
" return result\n"
|
||||
)
|
||||
edits = [
|
||||
EditSpec(
|
||||
match="result = call(\n a=1,\n)",
|
||||
prepend="dumper.dump('before', x)",
|
||||
)
|
||||
]
|
||||
result = apply_edits(source=source, edits=edits)
|
||||
assert result == (
|
||||
"def foo():\n"
|
||||
" dumper.dump('before', x)\n"
|
||||
" result = call(\n"
|
||||
" a=1,\n"
|
||||
" )\n"
|
||||
" return result\n"
|
||||
)
|
||||
|
||||
def test_replacement_and_append_mutually_exclusive(self) -> None:
|
||||
with pytest.raises(ValidationError, match="only one of"):
|
||||
EditSpec(match="x = 1", replacement="x = 2", append="print(x)")
|
||||
|
||||
def test_replacement_and_prepend_mutually_exclusive(self) -> None:
|
||||
with pytest.raises(ValidationError, match="only one of"):
|
||||
EditSpec(match="x = 1", replacement="x = 2", prepend="print(x)")
|
||||
|
||||
def test_prepend_and_append_mutually_exclusive(self) -> None:
|
||||
with pytest.raises(ValidationError, match="only one of"):
|
||||
EditSpec(match="x = 1", prepend="a()", append="b()")
|
||||
|
||||
def test_second_edit_sees_result_of_first(self) -> None:
|
||||
"""Edits are applied sequentially; second edit matches modified source."""
|
||||
source = "def foo():\n" " x = 1\n" " return x\n"
|
||||
edits = [
|
||||
EditSpec(match="x = 1", replacement="x = 1\ny = 2"),
|
||||
EditSpec(match="y = 2", replacement="y = 20"),
|
||||
]
|
||||
result = apply_edits(source=source, edits=edits)
|
||||
assert result == ("def foo():\n" " x = 1\n" " y = 20\n" " return x\n")
|
||||
Reference in New Issue
Block a user