From ba03ec431fb11682fb5743e6a95143b606f0378f Mon Sep 17 00:00:00 2001 From: Michael Panchenko Date: Sun, 6 Apr 2025 22:23:06 +0200 Subject: [PATCH] FindSymbolTool: allow passing a file for restricting search, not just a directory --- CHANGELOG.md | 1 + src/multilspy/language_server.py | 38 ++++++++++++------------- src/serena/agent.py | 8 ++++-- src/serena/symbol.py | 8 ++++-- test/multilspy/test_symbol_retrieval.py | 2 +- 5 files changed, 30 insertions(+), 27 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 599b9d96..9bd34d7f 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -3,5 +3,6 @@ ## 06.04.2025 - New tool: FindReferencingCodeSnippets - Adjusted prompt in CreateTextFileTool to prevent writing partial content (see [here](https://www.reddit.com/r/ClaudeAI/comments/1jpavtm/comment/mloek1x/?utm_source=share&utm_medium=web3x&utm_name=web3xcss&utm_term=1&utm_content=share_button)). +- FindSymbolTool: allow passing a file for restricting search, not just a directory (Gemini was too dumb to pass directories) ## 01.04.2025: Initial Release \ No newline at end of file diff --git a/src/multilspy/language_server.py b/src/multilspy/language_server.py index 41fc43c9..ae5fa594 100644 --- a/src/multilspy/language_server.py +++ b/src/multilspy/language_server.py @@ -746,7 +746,7 @@ class LanguageServer: self._cache_has_changed = True return result - async def request_full_symbol_tree(self, start_dir_relative_path: str | None = None, include_body: bool = False) -> List[multilspy_types.UnifiedSymbolInformation]: + async def request_full_symbol_tree(self, within_relative_path: str | None = None, include_body: bool = False) -> List[multilspy_types.UnifiedSymbolInformation]: """ Will go through all files in the project and build a tree of symbols. Note: this may be slow the first time it is called. @@ -755,19 +755,16 @@ class LanguageServer: that are within the repository. Will ignore all directories that start with a dot (.) and __pycache__ directories. - Args: - start_dir_relative_path: if passed, only the symbols within this directory will be considered. - include_body: whether to include the body of the symbols in the result. + :param within_relative_path: pass a relative path to only consider symbols within this path. + If a file is passed, only the symbols within this file will be considered. + If a directory is passed, all files within this directory will be considered. + :param include_body: whether to include the body of the symbols in the result. - Returns: - A list of root symbols representing the top-level packages/modules in the project. + :return: A list of root symbols representing the top-level packages/modules in the project. """ - if not self.server_started: - self.logger.log( - "request_full_symbol_tree called before Language Server started", - logging.ERROR, - ) - raise MultilspyException("Language Server not started") + if within_relative_path is not None and os.path.isfile(within_relative_path): + _, root_nodes = await self.request_document_symbols(within_relative_path, include_body=include_body) + return root_nodes # Helper function to check if a path should be ignored def should_ignore_dir(path: str) -> bool: @@ -854,7 +851,7 @@ class LanguageServer: return result # Start from the root or the specified directory - start_path = start_dir_relative_path or "." + start_path = within_relative_path or "." return await process_directory(start_path) @staticmethod @@ -1579,7 +1576,7 @@ class SyncLanguageServer: ).result() return result - def request_full_symbol_tree(self, start_package_relative_path: str | None = None, include_body: bool = False) -> List[multilspy_types.UnifiedSymbolInformation]: + def request_full_symbol_tree(self, within_relative_path: str | None = None, include_body: bool = False) -> List[multilspy_types.UnifiedSymbolInformation]: """ Will go through all files in the project and build a tree of symbols. Note: this may be slow the first time it is called. @@ -1588,15 +1585,16 @@ class SyncLanguageServer: that are within the repository. Will ignore all directories that start with a dot (.) and __pycache__ directories. - Args: - start_package_relative_path: if passed, only the symbols within this directory will be considered. - include_body: whether to include the body of the symbols in the result. + :param within_relative_path: pass a relative path to only consider symbols within this path. + If a file is passed, only the symbols within this file will be considered. + If a directory is passed, all files within this directory will be considered. + If None, the entire codebase will be considered. + :param include_body: whether to include the body of the symbols in the result. - Returns: - A list of root symbols representing the top-level packages/modules in the project. + :return: A list of root symbols representing the top-level packages/modules in the project. """ result = asyncio.run_coroutine_threadsafe( - self.language_server.request_full_symbol_tree(start_package_relative_path, include_body), self.loop + self.language_server.request_full_symbol_tree(within_relative_path, include_body), self.loop ).result() return result diff --git a/src/serena/agent.py b/src/serena/agent.py index 4b949b3a..9ea4c03a 100644 --- a/src/serena/agent.py +++ b/src/serena/agent.py @@ -430,11 +430,11 @@ class FindSymbolTool(Tool): self, name: str, depth: int = 0, + within_relative_path: str | None = None, include_body: bool = False, include_kinds: list[int] | None = None, exclude_kinds: list[int] | None = None, substring_matching: bool = False, - dir_relative_path: str | None = None, max_answer_chars: int = _DEFAULT_MAX_ANSWER_LENGTH, ) -> str: """ @@ -450,7 +450,9 @@ class FindSymbolTool(Tool): (e.g. depth 1 will retrieve methods and attributes for the case where the symbol refers to a class). Provide a non-zero depth if you intend to subsequently query symbols that are contained in the retrieved symbol. - :param dir_relative_path: pass a directory relative path to only consider symbols within this directory. + :param within_relative_path: pass a relative path to only consider symbols within this path. + If a file is passed, only the symbols within this file will be considered. + If a directory is passed, all files within this directory will be considered. If None, the entire codebase will be considered. :param include_body: whether to include the body of all symbols in the result. You should only use this if you actually need the body of the symbol for the task at hand (for example, for a deep analysis @@ -479,7 +481,7 @@ class FindSymbolTool(Tool): include_kinds=include_kinds, exclude_kinds=exclude_kinds, substring_matching=substring_matching, - dir_relative_path=dir_relative_path, + within_relative_path=within_relative_path, ) symbol_dicts = [s.to_dict(kind=True, location=True, depth=depth, include_body=include_body) for s in symbols] result = json.dumps(symbol_dicts) diff --git a/src/serena/symbol.py b/src/serena/symbol.py index 9c698f43..1819b1a2 100644 --- a/src/serena/symbol.py +++ b/src/serena/symbol.py @@ -200,7 +200,7 @@ class SymbolManager: def find_by_name( self, name: str, - dir_relative_path: str | None = None, + within_relative_path: str | None = None, include_body: bool = False, include_kinds: Sequence[SymbolKind] | None = None, exclude_kinds: Sequence[SymbolKind] | None = None, @@ -210,7 +210,9 @@ class SymbolManager: Find all symbols that match the given name. :param name: the name of the symbol to find - :param dir_relative_path: pass a directory relative path to only consider symbols within this directory. + :param within_relative_path: pass a relative path to only consider symbols within this path. + If a file is passed, only the symbols within this file will be considered. + If a directory is passed, all files within this directory will be considered. If None, the entire codebase will be considered. :param include_body: whether to include the body of all symbols in the result. Note: you can filter out the bodies of the children if you set include_children_body=False @@ -224,7 +226,7 @@ class SymbolManager: :return: a list of symbols that match the given name """ symbols: list[Symbol] = [] - symbol_roots = self.lang_server.request_full_symbol_tree(start_package_relative_path=dir_relative_path, include_body=include_body) + symbol_roots = self.lang_server.request_full_symbol_tree(within_relative_path=within_relative_path, include_body=include_body) for root in symbol_roots: symbols.extend( Symbol(root).find(name, include_kinds=include_kinds, exclude_kinds=exclude_kinds, substring_matching=substring_matching) diff --git a/test/multilspy/test_symbol_retrieval.py b/test/multilspy/test_symbol_retrieval.py index dd630b5e..14fe1fc7 100644 --- a/test/multilspy/test_symbol_retrieval.py +++ b/test/multilspy/test_symbol_retrieval.py @@ -366,7 +366,7 @@ class TestLanguageServerSymbols: def test_symbol_tree_structure_subdir(self, language_server: SyncLanguageServer): """Test that the symbol tree structure is correctly built.""" # Get all symbols in the test file - examples_package_roots = language_server.request_full_symbol_tree(start_package_relative_path="examples") + examples_package_roots = language_server.request_full_symbol_tree(within_relative_path="examples") assert len(examples_package_roots) == 1 examples_package = examples_package_roots[0] assert examples_package["name"] == "examples"