diff --git a/src/serena/symbol.py b/src/serena/symbol.py index b2c764a..4821bc2 100644 --- a/src/serena/symbol.py +++ b/src/serena/symbol.py @@ -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) diff --git a/test/conftest.py b/test/conftest.py index 4e403df..a5ca841 100644 --- a/test/conftest.py +++ b/test/conftest.py @@ -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. diff --git a/test/serena/test_symbol_editing.py b/test/serena/test_symbol_editing.py index 9f2b5cf..0ea19bc 100644 --- a/test/serena/test_symbol_editing.py +++ b/test/serena/test_symbol_editing.py @@ -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():