mirror of
https://github.com/tiennm99/serena.git
synced 2026-10-11 03:13:51 +00:00
Extension of tools: filtering, substring matching, reading chunks, including bodies
This commit is contained in:
1 parent
85cb44e425
commit
eb3d2d3ec4
2 files changed
+239
-37
No files matched your search
+116
-24
@@ -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
@@ -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)
|
||||
Reference in new issue
Block a user