Add test case abstraction for editing tests

Add implementation for "delete symbol", removing unnecessary dry-run
flag and overloads from the implementation
This commit is contained in:
Dominik Jain
2025-05-29 18:15:07 +02:00
parent 27b9dd0f9b
commit bed72f417a
3 changed files with 90 additions and 55 deletions
+4 -28
View File
@@ -949,43 +949,19 @@ class SymbolManager:
self.lang_server.delete_text_between_positions(relative_path, start_pos, end_pos)
return None
@overload
def delete_symbol_at_location(self, location: SymbolLocation, *, dry_run: Literal[False] = False) -> None: ...
@overload
def delete_symbol_at_location(self, location: SymbolLocation, *, dry_run: Literal[True]) -> CodeDiff: ...
def delete_symbol_at_location(self, location: SymbolLocation, *, dry_run: bool = False) -> CodeDiff | None:
def delete_symbol_at_location(self, location: SymbolLocation) -> None:
"""
Deletes the symbol at the given location.
:param dry_run: if True, return a CodeDiff instead of modifying the file
"""
with self._edited_symbol_location(location) as symbol:
assert location.relative_path is not None
assert symbol.body_start_position is not None
assert symbol.body_end_position is not None
if dry_run:
original_content = self._get_code_file_content(location.relative_path)
modified_content, _ = TextUtils.delete_text_between_positions(
original_content,
symbol.body_start_position["line"],
symbol.body_start_position["character"],
symbol.body_end_position["line"],
symbol.body_end_position["character"],
)
return CodeDiff(relative_path=location.relative_path, original_content=original_content, modified_content=modified_content)
else:
self.lang_server.delete_text_between_positions(location.relative_path, symbol.body_start_position, symbol.body_end_position)
return None
self.lang_server.delete_text_between_positions(location.relative_path, symbol.body_start_position, symbol.body_end_position)
@overload
def delete_symbol(self, name_path: str, relative_file_path: str, *, dry_run: Literal[False] = False) -> None: ...
@overload
def delete_symbol(self, name_path: str, relative_file_path: str, *, dry_run: Literal[True]) -> CodeDiff: ...
def delete_symbol(self, name_path: str, relative_file_path: str, *, dry_run: bool = False) -> CodeDiff | None:
def delete_symbol(self, name_path: str, relative_file_path: str) -> None:
"""
Deletes the symbol with the given name in the given file.
:param dry_run: if True, return a CodeDiff instead of modifying the file
"""
symbol_candidates = self.find_by_name(name_path, within_relative_path=relative_file_path)
if len(symbol_candidates) == 0:
@@ -997,4 +973,4 @@ class SymbolManager:
"Their locations are: \n " + json.dumps([s.location.to_dict() for s in symbol_candidates], indent=2)
)
symbol = symbol_candidates[0]
return self.delete_symbol_at_location(symbol.location, dry_run=dry_run) # type: ignore[call-overload]
self.delete_symbol_at_location(symbol.location)
+6 -2
View File
@@ -22,13 +22,17 @@ def get_repo_path(language: Language) -> Path:
return Path(__file__).parent / "resources" / "repos" / language / "test_repo"
def create_default_ls(language: Language) -> SyncLanguageServer:
def create_ls(language: Language, repo_path: str):
config = MultilspyConfig(code_language=language)
repo_path = str(get_repo_path(language))
logger = MultilspyLogger()
return SyncLanguageServer.create(config, logger, repo_path)
def create_default_ls(language: Language) -> SyncLanguageServer:
repo_path = str(get_repo_path(language))
return create_ls(language, repo_path)
@pytest.fixture(scope="session")
def repo_path(request: LanguageParamRequest) -> Path:
"""Get the repository path for a specific language.
+80 -25
View File
@@ -1,10 +1,66 @@
import os
import shutil
import tempfile
from abc import abstractmethod
from collections.abc import Iterator
from contextlib import contextmanager
from pathlib import Path
import pytest
from multilspy import SyncLanguageServer
from multilspy.multilspy_config import Language
from serena.symbol import CodeDiff
from src.serena.symbol import SymbolManager
from test.conftest import create_ls, get_repo_path
class EditingTest:
def __init__(self, language: Language, rel_path: str):
"""
:param language: the language
:param rel_path: the relative path of the edited file
"""
self.rel_path = rel_path
self.language = language
self.original_repo_path = get_repo_path(language)
self.repo_path: Path | None = None
@contextmanager
def _setup(self) -> Iterator[SymbolManager]:
"""Context manager for setup/teardown with a temporary directory, providing the symbol manager."""
temp_dir = Path(tempfile.mkdtemp())
self.repo_path = temp_dir / self.original_repo_path.name
try:
shutil.copytree(self.original_repo_path, self.repo_path)
language_server = create_ls(self.language, str(self.repo_path))
language_server.start()
yield SymbolManager(lang_server=language_server)
finally:
shutil.rmtree(temp_dir)
def _read_file(self, rel_path: str) -> str:
"""Read the content of a file in the test repository."""
file_path = self.repo_path / rel_path
with open(file_path, encoding="utf-8") as f:
return f.read()
def run_test(self) -> None:
with self._setup() as symbol_manager:
content_before = self._read_file(self.rel_path)
self._apply_edit(symbol_manager)
content_after = self._read_file(self.rel_path)
code_diff = CodeDiff(self.rel_path, original_content=content_before, modified_content=content_after)
self._test_diff(code_diff)
@abstractmethod
def _apply_edit(self, symbol_manager: SymbolManager) -> None:
pass
@abstractmethod
def _test_diff(self, code_diff: CodeDiff) -> None:
pass
# Python test file path
PYTHON_TEST_REL_FILE_PATH = os.path.join("test_repo", "variables.py")
@@ -63,43 +119,42 @@ EXPECTED_DELETED_DEMOCLASS_TYPESCRIPT = """export class DemoClass {
}"""
class DeleteSymbolTest(EditingTest):
def __init__(self, language: Language, rel_path: str, deleted_symbol: str, expected_deleted_lines: str):
super().__init__(language, rel_path)
self.expected_deleted_lines = expected_deleted_lines
self.deleted_symbol = deleted_symbol
self.rel_path = rel_path
def _apply_edit(self, symbol_manager: SymbolManager) -> None:
symbol_manager.delete_symbol(self.deleted_symbol, self.rel_path)
def _test_diff(self, code_diff: CodeDiff) -> None:
assert code_diff.original_content != code_diff.modified_content
actual_deleted_lines = [line.strip() for _, line in code_diff.deleted_lines if line.strip()]
normalized_expected_deleted_lines = [line.strip() for line in self.expected_deleted_lines.splitlines() if line.strip()]
assert actual_deleted_lines == normalized_expected_deleted_lines
@pytest.mark.parametrize(
"language_server, relative_file_path, symbol_name, expected_deleted_lines",
"test_case",
[
(
DeleteSymbolTest(
Language.PYTHON,
PYTHON_TEST_REL_FILE_PATH,
"VariableContainer",
EXPECTED_DELETED_VARIABLE_CONTAINER_PYTHON.strip().splitlines(),
EXPECTED_DELETED_VARIABLE_CONTAINER_PYTHON,
),
(
DeleteSymbolTest(
Language.TYPESCRIPT,
TYPESCRIPT_TEST_FILE,
"DemoClass",
EXPECTED_DELETED_DEMOCLASS_TYPESCRIPT.strip().splitlines(),
EXPECTED_DELETED_DEMOCLASS_TYPESCRIPT,
),
],
indirect=["language_server"],
)
def test_delete_symbol_dry_run(
language_server: SyncLanguageServer,
relative_file_path: str,
symbol_name: str,
expected_deleted_lines: list[str],
):
symbol_manager = SymbolManager(lang_server=language_server)
code_diff = symbol_manager.delete_symbol(symbol_name, relative_file_path, dry_run=True)
assert code_diff is not None
assert code_diff.relative_path == relative_file_path
assert len(code_diff.added_lines) == 0
actual_deleted_lines = [line.strip() for _, line in code_diff.deleted_lines if line.strip()]
# Normalize expected lines by stripping whitespace and removing empty lines
normalized_expected_deleted_lines = [line.strip() for line in expected_deleted_lines if line.strip()]
assert actual_deleted_lines == normalized_expected_deleted_lines
assert code_diff.original_content != code_diff.modified_content
def test_delete_symbol(test_case):
test_case.run_test()
NEW_PYTHON_FUNCTION = """def new_inserted_function():