Support patching source code (#19561)

This commit is contained in:
fzyzcjy
2026-02-28 18:05:45 +08:00
committed by GitHub
parent b73aa53d7e
commit 4097eb5ce9
7 changed files with 979 additions and 0 deletions

View 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

View 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"

View 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")