From d33b67eec2cb61c01795b5ccf546a733bb263d75 Mon Sep 17 00:00:00 2001 From: Dominik Jain Date: Thu, 27 Mar 2025 23:06:52 +0100 Subject: [PATCH] Revamp SymbolRetriever to SymbolManager, adding editing function to replace a symbol's body --- src/serena/mcp.py | 41 +++++++++++++++++-- src/serena/symbol.py | 96 ++++++++++++++++++++++++++++++++++++++------ 2 files changed, 120 insertions(+), 17 deletions(-) diff --git a/src/serena/mcp.py b/src/serena/mcp.py index 639ce72..fdfec2d 100644 --- a/src/serena/mcp.py +++ b/src/serena/mcp.py @@ -24,7 +24,7 @@ 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.symbol import SymbolLocation, SymbolManager from serena.util.file_system import scan_directory log = logging.getLogger(__name__) @@ -306,7 +306,7 @@ def find_symbol( class FindSymbolTool(Tool): def _execute(self) -> str: - symbols = SymbolRetriever(self.langsrv).find( + symbols = SymbolManager(self.langsrv).find_by_name( name, include_body=include_body, include_kinds=include_kinds, @@ -354,8 +354,11 @@ def find_referencing_symbols( class FindReferencingSymbolsTool(Tool): def _execute(self) -> str: - symbols = SymbolRetriever(self.langsrv).find_references( - relative_path, line, column, include_body=include_body, include_kinds=include_kinds, exclude_kinds=exclude_kinds + symbols = SymbolManager(self.langsrv).find_referencing_symbols( + SymbolLocation(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) @@ -363,6 +366,36 @@ def find_referencing_symbols( return FindReferencingSymbolsTool(ctx).execute(max_answer_chars=max_answer_chars) +@mcp.tool() +def replace_symbol_body( + ctx: Context, + relative_path: str, + line: int, + column: int, + body: str, +) -> str: + """ + Replaces the body of the symbol at the given location + + :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 body: the new symbol body. Important: Provide the correct level of indentation + (as the original body) + """ + + class ReplaceSymbolBodyTool(Tool): + def _execute(self) -> str: + SymbolManager(self.langsrv).replace_body( + SymbolLocation(relative_path, line, column), + body=body, + ) + return "OK" + + return ReplaceSymbolBodyTool(ctx).execute() + + @mcp.tool() def onboarding(ctx: Context) -> str: """ diff --git a/src/serena/symbol.py b/src/serena/symbol.py index 2976274..81288aa 100644 --- a/src/serena/symbol.py +++ b/src/serena/symbol.py @@ -1,15 +1,32 @@ import logging +import os from collections.abc import Iterator, Sequence +from contextlib import contextmanager +from dataclasses import asdict, dataclass from typing import Any, Self from sensai.util.string import ToStringMixin from multilspy import SyncLanguageServer -from multilspy.multilspy_types import SymbolKind, UnifiedSymbolInformation +from multilspy.multilspy_types import Position, SymbolKind, UnifiedSymbolInformation log = logging.getLogger(__name__) +@dataclass +class SymbolLocation: + """ + Represents the (start) location of a symbol identifier + """ + + relative_path: str + line: int + column: int + + def to_dict(self) -> dict[str, Any]: + return asdict(self) + + class Symbol(ToStringMixin): def __init__(self, s: UnifiedSymbolInformation) -> None: self.s = s @@ -36,6 +53,21 @@ class Symbol(ToStringMixin): def relative_path(self) -> str: return self.s["location"]["relativePath"] + @property + def location(self) -> SymbolLocation: + """ + :return: the start location of the actual symbol identifier + """ + return SymbolLocation(relative_path=self.relative_path, line=self.line, column=self.column) + + @property + def body_start_position(self) -> Position: + return self.s["location"]["range"]["start"] + + @property + def body_end_position(self) -> Position: + return self.s["location"]["range"]["end"] + @property def line(self) -> int: return self.s["selectionRange"]["start"]["line"] @@ -99,7 +131,7 @@ class Symbol(ToStringMixin): result["kind"] = self.kind if location: - result["location"] = {"relativePath": self.relative_path, "line": self.line, "column": self.column} + result["location"] = self.location.to_dict() if include_body: if self.body is None: @@ -126,14 +158,14 @@ class Symbol(ToStringMixin): return result -class SymbolRetriever: +class SymbolManager: def __init__(self, lang_server: SyncLanguageServer) -> None: self.lang_server = lang_server def _to_symbols(self, items: list[UnifiedSymbolInformation]) -> list[Symbol]: return [Symbol(s) for s in items] - def find( + def find_by_name( self, name: str, dir_relative_path: str | None = None, @@ -167,11 +199,22 @@ class SymbolRetriever: ) return symbols - def find_references( + def get_document_symbols(self, relative_path: str) -> list[Symbol]: + symbol_dicts, roots = self.lang_server.request_document_symbols(relative_path, include_body=False) + symbols = [Symbol(s) for s in symbol_dicts] + return symbols + + def find_by_location(self, location: SymbolLocation) -> Symbol | None: + symbol_dicts, roots = self.lang_server.request_document_symbols(location.relative_path, include_body=False) + for symbol_dict in symbol_dicts: + symbol = Symbol(symbol_dict) + if symbol.location == location: + return symbol + return None + + def find_referencing_symbols( self, - relative_path: str, - line: int, - column: int, + symbol_location: SymbolLocation, include_body: bool = False, include_kinds: Sequence[SymbolKind] | None = None, exclude_kinds: Sequence[SymbolKind] | None = None, @@ -179,10 +222,7 @@ class SymbolRetriever: """ 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 symbol_location: the location of the symbol for which to find references :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. @@ -193,7 +233,12 @@ class SymbolRetriever: :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=include_body + relative_file_path=symbol_location.relative_path, + line=symbol_location.line, + column=symbol_location.column, + include_imports=False, + include_self=False, + include_body=include_body, ) if include_kinds is not None: @@ -203,3 +248,28 @@ class SymbolRetriever: symbol_dicts = [s for s in symbol_dicts if s["kind"] not in exclude_kinds] return self._to_symbols(symbol_dicts) + + @contextmanager + def _edited_file(self, relative_path: str) -> Iterator[None]: + with self.lang_server.open_file(relative_path) as file_buffer: + yield + root_path = self.lang_server.language_server.repository_root_path + abs_path = os.path.join(root_path, relative_path) + with open(abs_path, "w") as f: + f.write(file_buffer.contents) + + def replace_body(self, location: SymbolLocation, body: str) -> None: + """ + Replace the body of the symbol at the given location with the given body + + :param location: the location of the symbol to replace + :param body: the new body + """ + symbol = self.find_by_location(location) + if symbol is None: + raise ValueError("Symbol not found") + with self._edited_file(location.relative_path): + self.lang_server.delete_text_between_positions(location.relative_path, symbol.body_start_position, symbol.body_end_position) + self.lang_server.insert_text_at_position( + location.relative_path, symbol.body_start_position["line"], symbol.body_start_position["character"], body + )