Extension of tools: filtering, substring matching, reading chunks, including bodies

This commit is contained in:
Michael Panchenko committed 2025-03-26 22:03:03 +01:00
1 parent 85cb44e425
commit eb3d2d3ec4
2 files changed
+239 -37

No files matched your search

+116 -24
View File
@@ -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)
+123 -13
View File
@@ -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)