From eb3d2d3ec468e0f10394a8203ea6d5367a3a0246 Mon Sep 17 00:00:00 2001 From: Michael Panchenko Date: Wed, 26 Mar 2025 21:40:59 +0100 Subject: [PATCH] Extension of tools: filtering, substring matching, reading chunks, including bodies --- src/serena/mcp.py | 140 +++++++++++++++++++++++++++++++++++-------- src/serena/symbol.py | 136 +++++++++++++++++++++++++++++++++++++---- 2 files changed, 239 insertions(+), 37 deletions(-) diff --git a/src/serena/mcp.py b/src/serena/mcp.py index 4d1d811d..ed9a957f 100644 --- a/src/serena/mcp.py +++ b/src/serena/mcp.py @@ -8,7 +8,7 @@ import sys import traceback from abc import ABC, abstractmethod from collections import defaultdict -from collections.abc import AsyncIterator +from collections.abc import AsyncIterator, Sequence from contextlib import asynccontextmanager from dataclasses import dataclass from typing import Any, cast @@ -22,6 +22,7 @@ from sensai.util import logging from multilspy import SyncLanguageServer from multilspy.multilspy_config import Language, MultilspyConfig from multilspy.multilspy_logger import MultilspyLogger +from multilspy.multilspy_types import SymbolKind from serena.llm.prompt_factory import PromptFactory from serena.symbol import SymbolRetriever from serena.util.file_system import scan_directory @@ -61,6 +62,8 @@ async def server_lifespan(mcp_server: FastMCP) -> AsyncIterator[SerenaMCPRequest print(f"Project file not found: {project_file}", file=sys.stderr) sys.exit(1) + print(f"Starting serena server for project {project_file}") + # read project configuration with open(project_file, encoding="utf-8") as f: project_config = yaml.safe_load(f) @@ -93,10 +96,21 @@ class Component(ABC): self.prompt_factory = lifespan_context.prompt_factory +_DEFAULT_MAX_ANSWER_LENGTH = int(2e5) + + class Tool(Component): - def execute(self) -> str: + _on_too_long_answer_msg = """ +The answer is too long to display ({} characters). +Please try a more specific tool query or raise the max_answer_length parameter. +""" + + def execute(self, max_answer_chars: int = _DEFAULT_MAX_ANSWER_LENGTH) -> str: try: - return self._execute() + result = self._execute() + if (n_chars := len(result)) > max_answer_chars: + return self._on_too_long_answer_msg.format(n_chars) + return result except Exception as e: msg = f"Error executing tool: {e}\n{traceback.format_exc()}" return msg @@ -133,19 +147,35 @@ class SequentialPrompt(Component): @mcp.tool() -def read_file(ctx: Context, relative_path: str) -> str: +def read_file( + ctx: Context, relative_path: str, start_line: int = 0, end_line: int | None = None, max_answer_chars: int = _DEFAULT_MAX_ANSWER_LENGTH +) -> str: """ + Read the given file or a chunk of it (between start_line and end_line). + :param ctx: the context object, which will be created and provided automatically :param relative_path: the relative path to the file to read + :param start_line: the start line of the range to read + :param end_line: the end line of the range to read. If None, the entire file will be read. + :param max_answer_chars: if the file (chunk) is longer than this number of characters, + no content will be returned. Don't adjust unless there is really no other way to get the content + required for the task. :return: the full text of the file at the given relative path """ log.info(f"read_file: {relative_path=}") class ReadFileTool(Tool): def _execute(self) -> str: - return self.langsrv.retrieve_full_file_content(relative_path) + result = self.langsrv.retrieve_full_file_content(relative_path) + result_lines = result.splitlines() + if end_line is None: + result_lines = result_lines[start_line:] + else: + result_lines = result_lines[start_line:end_line] + result = "\n".join(result_lines) + return result - return ReadFileTool(ctx).execute() + return ReadFileTool(ctx).execute(max_answer_chars=max_answer_chars) @mcp.tool() @@ -169,11 +199,14 @@ def create_text_file(ctx: Context, relative_path: str, content: str) -> str: @mcp.tool() -def list_dir(ctx: Context, relative_path: str, recursive: bool) -> str: +def list_dir(ctx: Context, relative_path: str, recursive: bool, max_answer_chars: int = _DEFAULT_MAX_ANSWER_LENGTH) -> str: """ :param ctx: the context object, which will be created and provided automatically :param relative_path: the relative path to the directory to list; pass "." to scan the project root :param recursive: whether to scan subdirectories recursively + :param max_answer_chars: if the directory is longer than this number of characters, + no content will be returned. Don't adjust unless there is really no other way to get the content + required for the task. :return: a JSON object with the names of directories and files within the given directory """ log.info(f"list_dir: {relative_path=}") @@ -185,17 +218,21 @@ def list_dir(ctx: Context, relative_path: str, recursive: bool) -> str: ) return json.dumps({"dirs": dirs, "files": files}) - return ListDirTool(ctx).execute() + return ListDirTool(ctx).execute(max_answer_chars=max_answer_chars) @mcp.tool() -def get_dir_overview(ctx: Context, relative_path: str) -> str: +def get_dir_overview(ctx: Context, relative_path: str, max_answer_chars: int = _DEFAULT_MAX_ANSWER_LENGTH) -> str: """ Get an overview of the given directory. For each file in the directory, we list the top-level symbols in the file (name, kind, line). :param ctx: the context object, which will be created and provided automatically :param relative_path: the relative path to the directory to get the overview of + :param max_answer_chars: if the overview is longer than this number of characters, + no content will be returned. Don't adjust unless there is really no other way to get the content + required for the task. If the overview is too long, you should use a smaller directory instead, + (e.g. a subdirectory). :return: a JSON object mapping relative paths of all contained files to info about top-level symbols in the file (name, kind, line). """ log.info(f"get_dir_overview: {relative_path=}") @@ -204,47 +241,97 @@ def get_dir_overview(ctx: Context, relative_path: str) -> str: def _execute(self) -> str: return json.dumps(self.langsrv.request_dir_overview(relative_path)) - return GetDirOverviewTool(ctx).execute() + return GetDirOverviewTool(ctx).execute(max_answer_chars=max_answer_chars) @mcp.tool() -def find_symbol(ctx: Context, name: str, depth: int = 0) -> str: +def find_symbol( + ctx: Context, + name: str, + depth: int = 0, + include_body: bool = False, + include_kinds: Sequence[SymbolKind] | None = None, + exclude_kinds: Sequence[SymbolKind] | None = None, + substring_matching: bool = False, + dir_relative_path: str | None = None, + max_answer_chars: int = _DEFAULT_MAX_ANSWER_LENGTH, +) -> str: """ - Retrieves information on the symbol/code entity with the given name + Retrieves information on all symbols/code entities with the given name. :param ctx: the context object, which will be created and provided automatically - :param name: the name of the symbol - :param depth: specifies the depth up to which descendants of the symbol are to be retrieved - (e.g. depth 1 will retrieve methods for the case where the symbol refers to a class) - :return: a list of JSON objects with the result + :param name: the name of the symbols to find + :param dir_relative_path: pass a directory relative path to only consider symbols within this directory. + 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 + in the to_dict method. + :param include_kinds: an optional sequence of ints representing the LSP symbol kind. + If provided, only symbols of the given kinds will be included in the result. + :param exclude_kinds: If provided, symbols of the given kinds will be excluded from the result. + Takes precedence over include_kinds. + :param substring_matching: whether to use substring matching for the symbol name. + If True, the symbol name will be matched if it contains the given name as a substring. + :param max_answer_chars: if the output is longer than this number of characters, + no content will be returned. Don't adjust unless there is really no other way to get the content + required for the task. Instead, if the output is too long, you should + make a stricter query. + :return: a list of symbols that match the given name """ class FindSymbolTool(Tool): def _execute(self) -> str: - symbols = SymbolRetriever(self.langsrv).find(name) - symbol_dicts = [s.to_dict(kind=True, location=True, depth=depth) for s in symbols] + symbols = SymbolRetriever(self.langsrv).find( + name, + include_body=include_body, + include_kinds=include_kinds, + exclude_kinds=exclude_kinds, + substring_matching=substring_matching, + dir_relative_path=dir_relative_path, + ) + symbol_dicts = [s.to_dict(kind=True, location=True, depth=depth, include_body=include_body) for s in symbols] return json.dumps(symbol_dicts) - return FindSymbolTool(ctx).execute() + return FindSymbolTool(ctx).execute(max_answer_chars=max_answer_chars) @mcp.tool() -def find_referencing_symbols(ctx: Context, relative_path: str, line: int, column: int) -> str: +def find_referencing_symbols( + ctx: Context, + relative_path: str, + line: int, + column: int, + include_body: bool = False, + include_kinds: Sequence[SymbolKind] | None = None, + exclude_kinds: Sequence[SymbolKind] | None = None, + max_answer_chars: int = _DEFAULT_MAX_ANSWER_LENGTH, +) -> str: """ :param ctx: the context object, which will be created and provided automatically :param relative_path: the relative path to the file containing the symbol :param line: the line number :param column: the column + :param include_body: whether to include the body of the symbols in the result + :param include_kinds: an optional sequence of ints representing the LSP symbol kind. + If provided, only symbols of the given kinds will be included in the result. + :param exclude_kinds: If provided, symbols of the given kinds will be excluded from the result. + Takes precedence over include_kinds. + :param max_answer_chars: if the output is longer than this number of characters, + no content will be returned. Don't adjust unless there is really no other way to get the content + required for the task. Instead, if the output is too long, you should + make a stricter query. :return: a list of JSON objects with the symbols referencing the requested symbol """ class FindReferencingSymbolsTool(Tool): def _execute(self) -> str: - symbols = SymbolRetriever(self.langsrv).find_references(relative_path, line, column) - symbol_dicts = [s.to_dict(kind=True, location=True, depth=0) for s in symbols] + symbols = SymbolRetriever(self.langsrv).find_references( + relative_path, line, column, include_body=include_body, include_kinds=include_kinds, exclude_kinds=exclude_kinds + ) + symbol_dicts = [s.to_dict(kind=True, location=True, depth=0, include_body=include_body) for s in symbols] return json.dumps(symbol_dicts) - return FindReferencingSymbolsTool(ctx).execute() + return FindReferencingSymbolsTool(ctx).execute(max_answer_chars=max_answer_chars) @mcp.tool() @@ -270,6 +357,7 @@ def search_files_for_pattern( context_lines_after: int = 0, paths_include_glob: str | None = None, paths_exclude_glob: str | None = None, + max_answer_chars: int = _DEFAULT_MAX_ANSWER_LENGTH, ) -> str: """ Search for a pattern in all codefiles in the project. @@ -280,6 +368,10 @@ def search_files_for_pattern( :param context_lines_after: Number of lines of context to include after each match :param paths_include_glob: Glob pattern to filter which files to include in the search :param paths_exclude_glob: Glob pattern to filter which files to exclude from the search. Takes precedence over paths_include_glob. + :param max_answer_chars: if the output is longer than this number of characters, + no content will be returned. Don't adjust unless there is really no other way to get the content + required for the task. Instead, if the output is too long, you should + make a stricter query. :return: A JSON object mapping file paths to lists of matched consecutive lines (with context, if requested). """ @@ -300,4 +392,4 @@ def search_files_for_pattern( file_to_matches[match.source_file_path].append(match.to_display_string()) return json.dumps(file_to_matches) - return SearchInAllCodeTool(ctx).execute() + return SearchInAllCodeTool(ctx).execute(max_answer_chars=max_answer_chars) diff --git a/src/serena/symbol.py b/src/serena/symbol.py index e7551014..29762746 100644 --- a/src/serena/symbol.py +++ b/src/serena/symbol.py @@ -1,4 +1,5 @@ -from collections.abc import Iterator +import logging +from collections.abc import Iterator, Sequence from typing import Any, Self from sensai.util.string import ToStringMixin @@ -6,6 +7,8 @@ from sensai.util.string import ToStringMixin from multilspy import SyncLanguageServer from multilspy.multilspy_types import SymbolKind, UnifiedSymbolInformation +log = logging.getLogger(__name__) + class Symbol(ToStringMixin): def __init__(self, s: UnifiedSymbolInformation) -> None: @@ -23,7 +26,11 @@ class Symbol(ToStringMixin): @property def kind(self) -> str: - return SymbolKind(self.s["kind"]).name + return SymbolKind(self.symbol_kind).name + + @property + def symbol_kind(self) -> SymbolKind: + return self.s["kind"] @property def relative_path(self) -> str: @@ -37,15 +44,32 @@ class Symbol(ToStringMixin): def column(self) -> int: return self.s["selectionRange"]["start"]["character"] - def iter_children(self) -> Iterator[Self]: + @property + def body(self) -> str | None: + return self.s.get("body") + + def iter_children(self) -> Iterator["Symbol"]: for c in self.s["children"]: yield Symbol(c) - def find(self, name: str) -> list[Self]: + def find( + self, + name: str, + substring_matching: bool = False, + include_kinds: Sequence[SymbolKind] | None = None, + exclude_kinds: Sequence[SymbolKind] | None = None, + ) -> list["Symbol"]: result = [] - def traverse(s: Self) -> None: - if s.name == name: + def should_include(s: "Symbol") -> bool: + name_match = (substring_matching and name in s.name) or name == s.name + kind_include_match = include_kinds is None or s.symbol_kind in include_kinds + kind_exclude_match = exclude_kinds is None or s.symbol_kind not in exclude_kinds + + return name_match and kind_include_match and kind_exclude_match + + def traverse(s: "Symbol") -> None: + if should_include(s): result.append(s) for c in s.iter_children(): traverse(c) @@ -53,7 +77,22 @@ class Symbol(ToStringMixin): traverse(self) return result - def to_dict(self, kind: bool = False, location: bool = False, depth: int = 0) -> dict[str, Any]: + def to_dict( + self, kind: bool = False, location: bool = False, depth: int = 0, include_body: bool = False, include_children_body: bool = False + ) -> dict[str, Any]: + """ + Convert the symbol to a dictionary. + + :param kind: whether to include the kind of the symbol + :param location: whether to include the location of the symbol + :param depth: the depth of the symbol + :param include_body: whether to include the body of the top-level symbol. + :param include_children_body: whether to also include the body of the children. + Note that the body of the children is part of the body of the parent symbol, + so there is usually no need to set this to True unless you want process the output + and pass the children without passing the parent body to the LM. + :return: a dictionary representation of the symbol + """ result: dict[str, Any] = {"name": self.name} if kind: @@ -62,10 +101,23 @@ class Symbol(ToStringMixin): if location: result["location"] = {"relativePath": self.relative_path, "line": self.line, "column": self.column} + if include_body: + if self.body is None: + log.warning("Requested body for symbol, but it is not present. The symbol might have been loaded with include_body=False.") + result["body"] = self.body + def add_children(s: Self) -> list[dict[str, Any]]: children = [] for c in s.iter_children(): - children.append(c.to_dict(kind=kind, location=location, depth=depth - 1)) + children.append( + c.to_dict( + kind=kind, + location=location, + depth=depth - 1, + include_body=include_children_body, + include_children_body=include_children_body, + ) + ) return children if depth > 0: @@ -81,15 +133,73 @@ class SymbolRetriever: def _to_symbols(self, items: list[UnifiedSymbolInformation]) -> list[Symbol]: return [Symbol(s) for s in items] - def find(self, name: str) -> list[Symbol]: + def find( + self, + name: str, + dir_relative_path: str | None = None, + include_body: bool = False, + include_kinds: Sequence[SymbolKind] | None = None, + exclude_kinds: Sequence[SymbolKind] | None = None, + substring_matching: bool = False, + ) -> list[Symbol]: + """ + 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. + 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 + in the to_dict method. + :param include_kinds: an optional sequence of ints representing the LSP symbol kind. + If provided, only symbols of the given kinds will be included in the result. + :param exclude_kinds: If provided, symbols of the given kinds will be excluded from the result. + Takes precedence over include_kinds. + :param substring_matching: whether to use substring matching for the symbol name. + If True, the symbol name will be matched if it contains the given name as a substring. + :return: a list of symbols that match the given name + """ symbols: list[Symbol] = [] - symbol_roots = self.lang_server.request_full_symbol_tree() + symbol_roots = self.lang_server.request_full_symbol_tree(start_package_relative_path=dir_relative_path, include_body=include_body) for root in symbol_roots: - symbols.extend(Symbol(root).find(name)) + symbols.extend( + Symbol(root).find(name, include_kinds=include_kinds, exclude_kinds=exclude_kinds, substring_matching=substring_matching) + ) return symbols - def find_references(self, relative_path: str, line: int, column: int) -> list[Symbol]: + def find_references( + self, + relative_path: str, + line: int, + column: int, + include_body: bool = False, + include_kinds: Sequence[SymbolKind] | None = None, + exclude_kinds: Sequence[SymbolKind] | None = None, + ) -> list[Symbol]: + """ + Find all symbols that reference the given symbol. + + :param relative_path: the relative path to the file containing the symbol + :param line: the line number of the symbol (0-indexed). + :param column: the column number of the symbol. Note that this usually corresponds to the + column in `selectionRange` of the symbol (as opposed to the `range`). + :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 + in the to_dict method. + :param include_kinds: an optional sequence of ints representing the LSP symbol kind. + If provided, only symbols of the given kinds will be included in the result. + :param exclude_kinds: If provided, symbols of the given kinds will be excluded from the result. + Takes precedence over include_kinds. + :return: a list of symbols that reference the given symbol + """ symbol_dicts = self.lang_server.request_referencing_symbols( - relative_file_path=relative_path, line=line, column=column, include_imports=False, include_self=False, include_body=False + relative_file_path=relative_path, line=line, column=column, include_imports=False, include_self=False, include_body=include_body ) + + if include_kinds is not None: + symbol_dicts = [s for s in symbol_dicts if s["kind"] in include_kinds] + + if exclude_kinds is not None: + symbol_dicts = [s for s in symbol_dicts if s["kind"] not in exclude_kinds] + return self._to_symbols(symbol_dicts)