Add remaining LSP operations to the lsp facade; tools delegate to LspApi

LspApi now covers all language server-backed operations:
restart_language_server, get_symbols_overview, find_symbol,
find_referencing_symbols, find_implementations, find_declaration,
get_diagnostics_for_file, get_diagnostics_for_symbol,
replace_symbol_body, insert_after_symbol, insert_before_symbol,
rename_symbol and safe_delete_symbol.

Result objects carry their rendering policy:
* LspSymbolCollectionRenderer was generalised (symbol_dicts_,
  child_inclusion_predicate) and is reused by find_implementations
* LspSymbolsOverviewRenderer renders a file's overview with the
  depth-0/kind-count shortening ladder
* LspSymbol/LspSymbolRenderer represent a single symbol
  (find_declaration), preserving the dict output shape
* LspReferenceCollection with its renderer (context lines, per-file
  counts, total count)
* LspDiagnostics wrapping GroupedDiagnostics

The symbol tools are now thin adapters which delegate to the API via
the LspApiMixin (the JetBrains tools use JetBrainsApiMixin
analogously, replacing the intermediate tool base class). Editing
tools retain the DiagnosticsContext wrapper, which the API does not
use; DiagnosticsContext moved to serena.lsp.lsp_diagnostics and is
created via EditingToolWithDiagnostics.diagnostics_context.

SUCCESS_RESULT moved to serena.facades.facade (the API cannot import
from serena.tools without an import cycle); serena.tools re-exports it.
The project health check in the CLI uses LspApi directly instead of
tool internals. iter_subclasses now yields each class once.
This commit is contained in:
Dominik Jain authored and Dominik Jain committed 2026-09-15 12:50:43 +02:00
1 parent eb53a4cee4
commit 91c2ec8dc2
10 files changed
+837 -516

No files matched your search

+26 -32
View File
@@ -931,8 +931,8 @@ class ProjectCommands(AutoRegisteringGroup):
"""
# NOTE: completely written by Claude Code, only functionality was reviewed, not implementation
from serena.agent import SerenaAgent
from serena.facades.api.lsp import LspApi
from serena.project import Project
from serena.tools import FindReferencingSymbolsTool, FindSymbolTool, GetSymbolsOverviewTool
logging.configure(level=logging.INFO)
project_path = os.path.abspath(project)
@@ -977,61 +977,55 @@ class ProjectCommands(AutoRegisteringGroup):
if not target_file:
raise ProjectCommands._HealthCheckFailure("No analyzable files found")
# Get tools from agent
overview_tool = agent.get_tool(GetSymbolsOverviewTool)
find_symbol_tool = agent.get_tool(FindSymbolTool)
find_refs_tool = agent.get_tool(FindReferencingSymbolsTool)
api = LspApi(agent)
# Test 1: Get symbols overview
log.info("Testing GetSymbolsOverviewTool on file: %s", target_file)
overview_data = agent.execute_task(lambda: overview_tool.get_symbol_overview(target_file))
log.info(f"GetSymbolsOverviewTool returned: {overview_data}")
# Test 1: symbols overview
log.info("Testing get_symbols_overview on file: %s", target_file)
overview = agent.execute_task(lambda: api.get_symbols_overview(target_file))
log.info(f"get_symbols_overview returned: {overview.represent()}")
if not overview_data:
if len(overview) == 0:
raise ProjectCommands._HealthCheckFailure(f"No symbols found in target file {target_file}")
# Extract suitable symbol (prefer class or function over variables)
preferred_kinds = {SymbolKind.Class.name, SymbolKind.Function.name, SymbolKind.Method.name, SymbolKind.Constructor.name}
selected_symbol = None
for symbol in overview_data:
if symbol.get("kind") in preferred_kinds:
selected_symbol = symbol
break
preferred_kinds = {SymbolKind.Class, SymbolKind.Function, SymbolKind.Method, SymbolKind.Constructor}
selected_symbol = next((s for s in overview.symbols if s.symbol_kind in preferred_kinds), None)
# If no preferred symbol found, use first available
if not selected_symbol:
selected_symbol = overview_data[0]
if selected_symbol is None:
selected_symbol = overview.symbols[0]
log.info("No class or function found, using first available symbol")
symbol_name = selected_symbol["name"]
symbol_kind = selected_symbol["kind"]
log.info("Using symbol for testing: %s (kind: %s)", symbol_name, symbol_kind)
symbol_name = selected_symbol.name
log.info("Using symbol for testing: %s (kind: %s)", symbol_name, selected_symbol.symbol_kind_name)
# Test 2: FindSymbolTool
log.info("Testing FindSymbolTool for symbol: %s", symbol_name)
with find_symbol_tool.symbol_dict_grouper.disabled_context():
# Test 2: find_symbol
log.info("Testing find_symbol for symbol: %s", symbol_name)
with LspApi.find_symbol_dict_grouper_.disabled_context():
find_symbol_result = agent.execute_task(
lambda: find_symbol_tool.apply(symbol_name, relative_path=target_file, include_body=True)
lambda: api.find_symbol(symbol_name, relative_path=target_file, include_body=True).represent()
)
find_symbol_data = json.loads(find_symbol_result)
log.info("FindSymbolTool found %d matches for symbol %s", len(find_symbol_data), symbol_name)
log.info("find_symbol found %d matches for symbol %s", len(find_symbol_data), symbol_name)
if not find_symbol_data:
raise ProjectCommands._HealthCheckFailure("FindSymbolTool returned no results")
# Test 3: FindReferencingSymbolsTool
log.info("Testing FindReferencingSymbolsTool for symbol: %s", symbol_name)
# Test 3: find_referencing_symbols
log.info("Testing find_referencing_symbols for symbol: %s", symbol_name)
try:
with find_refs_tool.symbol_dict_grouper.disabled_context():
find_refs_result = agent.execute_task(lambda: find_refs_tool.apply(symbol_name, relative_path=target_file))
with LspApi.references_grouper_.disabled_context():
find_refs_result = agent.execute_task(
lambda: api.find_referencing_symbols(symbol_name, relative_path=target_file).represent()
)
find_refs_data = json.loads(find_refs_result)
log.info("FindReferencingSymbolsTool found %d references for symbol %s", len(find_refs_data), symbol_name)
log.info("find_referencing_symbols found %d references for symbol %s", len(find_refs_data), symbol_name)
except Exception as e:
# A symbol with no references at all is a legitimate result, so the number of
# references is not asserted - but a *failure* of the reference search means the
# language server is not functional, which is the single thing this command is
# asked to determine. Logging it as a warning let the command print
# "All tools working correctly" and exit 0 after the search had already failed.
raise ProjectCommands._HealthCheckFailure(f"FindReferencingSymbolsTool failed for symbol {symbol_name}: {e}") from e
raise ProjectCommands._HealthCheckFailure(f"find_referencing_symbols failed for symbol {symbol_name}: {e}") from e
log.info("Health check completed successfully")
+591 -45
View File
@@ -1,15 +1,27 @@
# SPDX-License-Identifier: GPL-3.0-or-later
"""
The implementation of language server (LSP)-backed operations.
"""
from collections import defaultdict
from collections.abc import Sequence
import os
from collections import Counter, defaultdict
from collections.abc import Callable, Sequence
from dataclasses import dataclass
from typing import TYPE_CHECKING
from typing import TYPE_CHECKING, Any
from serena.symbol import LanguageServerSymbol, LanguageServerSymbolDictGrouper, LanguageServerSymbolRetriever, SymbolDictGrouper
from serena.code_editor import LanguageServerCodeEditor
from serena.lsp.lsp_diagnostics import GroupedDiagnostics
from serena.symbol import (
LanguageServerSymbol,
LanguageServerSymbolDictGrouper,
LanguageServerSymbolRetriever,
ReferenceInLanguageServerSymbol,
SymbolDictGrouper,
)
from serena.util.text_utils import TextOutputUtils, find_text_coordinates
from solidlsp.lsp_protocol_handler.lsp_types import SymbolKind
from ...util.text_utils import TextOutputUtils
from ..facade import FacadeApi
from ..facade import SUCCESS_RESULT, FacadeApi
from ..representable import Renderer, RepresentableViaRenderer
if TYPE_CHECKING:
@@ -17,6 +29,12 @@ if TYPE_CHECKING:
class LspSymbolCollection(RepresentableViaRenderer):
"""
A collection of symbols retrieved via the language server.
Each symbol (`LanguageServerSymbol`) offers e.g. `get_name_path()`, `relative_path`, `symbol_kind_name`,
`body`, `get_body_line_numbers()`, `iter_children()`.
"""
def __init__(self, symbols: list[LanguageServerSymbol], renderer: "LspSymbolCollectionRenderer"):
"""
:param symbols: the list of symbols
@@ -25,7 +43,7 @@ class LspSymbolCollection(RepresentableViaRenderer):
super().__init__(renderer)
self.symbols = symbols
def __len__(self):
def __len__(self) -> int:
return len(self.symbols)
def relative_path_to_name_paths_(self) -> dict[str, list[str]]:
@@ -35,6 +53,20 @@ class LspSymbolCollection(RepresentableViaRenderer):
return result
class LspSymbol(RepresentableViaRenderer):
"""
A single symbol retrieved via the language server (see `LspSymbolCollection` for the symbol's interface).
"""
def __init__(self, symbol: LanguageServerSymbol, renderer: "LspSymbolRenderer"):
"""
:param symbol: the symbol
:param renderer: the renderer to use for representing the symbol
"""
super().__init__(renderer)
self.symbol = symbol
@dataclass(kw_only=True)
class SymbolOutputParams:
name_path: bool = True
@@ -49,9 +81,15 @@ class SymbolOutputParams:
relative_path: bool = False
include_body: bool = False
include_info: bool = False
child_inclusion_predicate: Callable[[LanguageServerSymbol], bool] | None = None
class LspSymbolCollectionRenderer(Renderer[LspSymbolCollection]):
"""
Renders a symbol collection as (optionally grouped) JSON according to the output parameters, falling back
to a mapping from files to name paths if the length limit is exceeded.
"""
def __init__(
self,
agent: "SerenaAgent",
@@ -65,26 +103,30 @@ class LspSymbolCollectionRenderer(Renderer[LspSymbolCollection]):
self._output_params = output_params
self._grouper = grouper
def set_grouper(self, grouper: SymbolDictGrouper) -> None:
self._grouper = grouper
def render(self, obj: LspSymbolCollection) -> str:
symbols = obj.symbols
def symbol_dicts_(self, symbols: list[LanguageServerSymbol]) -> list[LanguageServerSymbol.OutputDict]:
"""
:param symbols: the symbols to convert
:return: the dict representations of the symbols according to the output parameters (including info, if requested)
"""
p = self._output_params
symbol_dicts = [
s.to_dict(
kind=self._output_params.kind,
name_path=self._output_params.name_path,
name=self._output_params.name,
relative_path=self._output_params.relative_path,
body_location=self._output_params.body_location,
depth=self._output_params.depth,
body=self._output_params.include_body,
children_name=self._output_params.children_name,
children_name_path=self._output_params.children_name_path,
kind=p.kind,
name_path=p.name_path,
name=p.name,
location=p.location,
relative_path=p.relative_path,
body_location=p.body_location,
depth=p.depth,
body=p.include_body,
children_body=p.children_body,
children_name=p.children_name,
children_name_path=p.children_name_path,
child_inclusion_predicate=p.child_inclusion_predicate,
)
for s in symbols
]
if not self._output_params.include_body and self._output_params.include_info:
if not p.include_body and p.include_info:
info_by_symbol = self._symbol_retriever.request_info_for_symbol_batch(symbols)
for s, s_dict in zip(symbols, symbol_dicts, strict=True):
if symbol_info := info_by_symbol.get(s):
@@ -92,30 +134,239 @@ class LspSymbolCollectionRenderer(Renderer[LspSymbolCollection]):
# https://peps.python.org/pep-0728/
# If we ever upgrade to 3.15, we can remove the type: ignore[typeddict-unknown-key]
s_dict["info"] = symbol_info
return symbol_dicts
def _group(self, symbol_dicts: list[LanguageServerSymbol.OutputDict]) -> Any:
return self._grouper.group(symbol_dicts) if self._grouper is not None else symbol_dicts
def render(self, obj: LspSymbolCollection) -> str:
def create_short_result_relative_path_to_name_paths() -> str:
relative_path_to_name_paths = obj.relative_path_to_name_paths_()
return f"Shortened result:\n{TextOutputUtils.to_json(relative_path_to_name_paths)}"
return f"Shortened result:\n{TextOutputUtils.to_json(obj.relative_path_to_name_paths_())}"
if self._grouper is not None:
objects = self._grouper.group(symbol_dicts)
else:
objects = symbol_dicts
result = self._to_json(objects)
result = self._to_json(self._group(self.symbol_dicts_(obj.symbols)))
return self._limit_length(result, shortened_result_factories=[create_short_result_relative_path_to_name_paths])
class LspSymbolRenderer(Renderer[LspSymbol]):
"""
Renders a single symbol as JSON, using a collection renderer for the conversion.
"""
def __init__(self, agent: "SerenaAgent", max_answer_chars: int, collection_renderer: LspSymbolCollectionRenderer):
super().__init__(agent, max_answer_chars)
self._collection_renderer = collection_renderer
def render(self, obj: LspSymbol) -> str:
symbol_dict = self._collection_renderer.symbol_dicts_([obj.symbol])[0]
return self._limit_length(self._to_json(symbol_dict))
class LspSymbolsOverviewRenderer(LspSymbolCollectionRenderer):
"""
Renders a file's symbol overview, falling back to a depth-0 overview and finally symbol counts by kind
if the length limit is exceeded.
"""
def render(self, obj: LspSymbolCollection) -> str:
symbol_dicts = self.symbol_dicts_(obj.symbols)
result = self._to_json(self._group(symbol_dicts))
def make_kind_counts() -> str:
kind_names = [d.get("kind", "unknown") for d in symbol_dicts]
return f"Symbol counts by kind:\n{self._to_json(Counter(kind_names))}"
shortened_results: list[Callable[[], str]] = [make_kind_counts]
if self._output_params.depth > 0:
def make_depth_0_result() -> str:
depth_0_dicts = [d.copy() for d in symbol_dicts]
for d in depth_0_dicts:
d.pop("children", None)
return "Depth 0 overview:\n" + self._to_json(self._group(depth_0_dicts))
shortened_results.insert(0, make_depth_0_result)
return self._limit_length(result, shortened_result_factories=shortened_results)
class LspReferenceCollection(RepresentableViaRenderer):
"""
The references to a symbol, each a `ReferenceInLanguageServerSymbol` with the referencing `symbol`
(a `LanguageServerSymbol`) and the `line` of the reference.
"""
def __init__(self, references: list[ReferenceInLanguageServerSymbol], renderer: "LspReferenceCollectionRenderer"):
"""
:param references: the references
:param renderer: the renderer to use for representing the collection
"""
super().__init__(renderer)
self.references = references
def __len__(self) -> int:
return len(self.references)
class LspReferenceCollectionRenderer(Renderer[LspReferenceCollection]):
"""
Renders references as grouped JSON including the code around each reference, falling back to
references without code, per-file counts and finally the total count if the length limit is exceeded.
"""
def __init__(self, agent: "SerenaAgent", max_answer_chars: int, grouper: SymbolDictGrouper):
super().__init__(agent, max_answer_chars)
self._grouper = grouper
def render(self, obj: LspReferenceCollection) -> str:
project = self._agent.get_active_project_or_raise()
reference_dicts = []
ref_summaries = []
for ref in obj.references:
ref_dict = dict(ref.symbol.to_dict(kind=True, relative_path=True, depth=0, body=False, body_location=True))
ref_relative_path = ref.symbol.location.relative_path
assert ref_relative_path is not None, f"Referencing symbol {ref.symbol.name} has no relative path, this is likely a bug."
content_around_ref = project.retrieve_content_around_line(
relative_file_path=ref_relative_path, line=ref.line, context_lines_before=1, context_lines_after=1
)
ref_dict["content_around_reference"] = content_around_ref.to_display_string()
reference_dicts.append(ref_dict)
ref_summaries.append(
{
"name_path": ref_dict.get("name_path"),
"kind": ref_dict.get("kind"),
"relative_path": ref_dict.get("relative_path"),
"reference_line": ref.line,
}
)
result = self._to_json(self._grouper.group(reference_dicts))
# shortened result closures, from least to most aggressive shortening
def make_refs_without_context() -> str:
return f"References without surrounding lines:\n{self._to_json(self._grouper.group([dict(s) for s in ref_summaries]))}"
def make_per_file_counts() -> str:
counts = Counter(str(r["relative_path"]) for r in ref_summaries)
return f"Reference counts per file:\n{self._to_json(counts)}"
def make_summary() -> str:
return f"Found {len(ref_summaries)} references."
return self._limit_length(result, shortened_result_factories=[make_refs_without_context, make_per_file_counts, make_summary])
class LspDiagnostics(RepresentableViaRenderer):
"""
Diagnostics grouped as `relative_path -> severity -> name_path -> diagnostics`; see `grouped.get_dict()`.
"""
def __init__(self, grouped: GroupedDiagnostics, renderer: "LspDiagnosticsRenderer"):
"""
:param grouped: the grouped diagnostics
:param renderer: the renderer to use for representing the diagnostics
"""
super().__init__(renderer)
self.grouped = grouped
class LspDiagnosticsRenderer(Renderer[LspDiagnostics]):
def render(self, obj: LspDiagnostics) -> str:
return self._limit_length(self._to_json(obj.grouped.get_dict()))
class LspApi(FacadeApi):
def __init__(self, agent: "SerenaAgent") -> None:
super().__init__(agent, name="lsp", description="LSP-backed operations on the codebase (finding symbols, etc.)")
FILE_LEVEL_DIAGNOSTIC_BUCKET = "<file>"
"""the name path under which diagnostics that cannot be mapped to a symbol are grouped"""
def _create_symbol_retriever(self) -> LanguageServerSymbolRetriever:
assert self._agent.get_language_backend().is_lsp(), "Symbolic read operations require the language server backend"
return LanguageServerSymbolRetriever(self._get_project())
# group children by kind, keeping just the name (the parent's name_path makes it unambiguous);
# groupers for the various symbol collections; top-level symbols are grouped by the first key list,
# children by the second.
# For find_symbol, we group children by kind, keeping just the name (the parent's name_path makes it unambiguous);
# we don't group the top-level result list because many tests rely on it being a flat list of symbol dicts
find_symbol_dict_grouper_ = LanguageServerSymbolDictGrouper([], ["kind"], collapse_singleton=True)
references_grouper_ = LanguageServerSymbolDictGrouper(["relative_path", "kind"], ["kind"], collapse_singleton=True)
overview_grouper_ = LanguageServerSymbolDictGrouper(["kind"], ["kind"], collapse_singleton=True)
def __init__(self, agent: "SerenaAgent") -> None:
super().__init__(
agent,
name="lsp",
description="language server-backed operations on the codebase (finding symbols, references, implementations, "
"declarations and diagnostics; editing and renaming symbols)",
)
def _create_symbol_retriever(self) -> LanguageServerSymbolRetriever:
assert self._agent.get_language_backend().is_lsp(), "Language server operations require the language server backend"
return LanguageServerSymbolRetriever(self._get_project())
def _create_code_editor(self, symbol_retriever: LanguageServerSymbolRetriever | None = None) -> LanguageServerCodeEditor:
return LanguageServerCodeEditor(symbol_retriever or self._create_symbol_retriever())
@staticmethod
def _parse_kinds(kinds: Sequence[int]) -> Sequence[SymbolKind] | None:
return [SymbolKind(k) for k in kinds] if kinds else None
def _create_diagnostics(self, grouped: GroupedDiagnostics, max_answer_chars: int) -> LspDiagnostics:
return LspDiagnostics(grouped, LspDiagnosticsRenderer(self._agent, max_answer_chars))
# language server management
def restart_language_server(self) -> str:
"""
Restarts the language server(s). Use this only on explicit user request or after confirmation;
it may be necessary if a language server hangs.
:return: a success message
"""
self._agent.reset_language_server_manager()
return SUCCESS_RESULT
# read operations
def get_symbols_overview(self, relative_path: str, depth: int = -1, max_answer_chars: int = -1) -> LspSymbolCollection:
"""
Gets an overview of the top-level symbols defined in the given file (classes, methods, fields) — its
STRUCTURE, without their bodies. This is the cheap, structure-first way to learn what a file
contains: it costs far less context than reading the whole file.
:param relative_path: the relative path to the file to get the overview of
:param depth: depth up to which descendants shall be retrieved.
Default (-1) results in a language specific choice: 1 for java and kotlin and 0 for other languages
:param max_answer_chars: max result length; -1 for default. If exceeded, a shortened result is returned.
:return: the top-level symbols of the file
"""
# Note: file system sync not required (relevant file is opened in the language server explicitly)
if depth == -1:
depth = 1 if relative_path.endswith((".java", ".kt")) else 0
symbol_retriever = self._create_symbol_retriever()
# the symbol overview is capable of working with both files and directories, but we require a file
file_path = os.path.join(self._get_project().project_root, relative_path)
if not os.path.exists(file_path):
raise FileNotFoundError(f"File or directory {relative_path} does not exist in the project.")
if os.path.isdir(file_path):
raise ValueError(f"Expected a file path, but got a directory path: {relative_path}. ")
if not symbol_retriever.can_analyze_file(relative_path):
raise ValueError(
f"Cannot extract symbols from file {relative_path}. "
f"Active language servers: {[l.get_key() for l in self._agent.get_active_language_server_ids()]}"
)
symbols = symbol_retriever.get_symbol_overview(relative_path)[relative_path]
output_params = SymbolOutputParams(
name_path=False,
name=True,
depth=depth,
kind=True,
relative_path=False,
location=False,
child_inclusion_predicate=lambda s: not s.is_low_level(),
)
renderer = LspSymbolsOverviewRenderer(
self._agent, max_answer_chars, symbol_retriever, output_params, grouper=self.overview_grouper_
)
return LspSymbolCollection(symbols, renderer)
def find_symbol(
self,
@@ -154,31 +405,31 @@ class LspApi(FacadeApi):
:param relative_path: (optional) restrict search to this file or directory. If None, searches entire codebase.
If a directory is passed, the search will be restricted to the files in that directory.
If a file is passed, the search will be restricted to that file.
:param include_body: whether to include the symbol's source code. Use judiciously.
If you have some knowledge about the codebase, you should use this parameter, as it will significantly
speed up the search as well as reduce the number of results.
:param include_body: If True, include the symbol's source code. Use judiciously.
:param include_info: whether to include additional info (hover-like, typically including docstring and signature),
about the symbol (ignored if include_body is True). Info is never included for child symbols.
Note: Depending on the language, this can be slow (e.g., C/C++).
:param include_kinds: (optional) limits results to the given LSP symbol kinds (integers)
:param exclude_kinds: (optional) list of LSP symbol kinds (integers) to exclude.
:param substring_matching: If True, use substring matching for the last element of the pattern, such that
"Foo/get" would match "Foo/getValue" and "Foo/getData".
:param max_matches: maximum number of permitted matches. If exceeded, a shortened result is returned
:param substring_matching: If True, use substring matching for the last segment of `name_path_pattern`
(i.e. the name of the symbol, e.g. "foo" in "Class/foo" or "my_method" in "my_method").
:param max_matches: Maximum number of permitted matches. If exceeded, an error containing a shortened result is raised,
which allows refining the search. -1 (default) means no limit. Set to 1 to search for a unique symbol.
:param max_answer_chars: max result length; -1 for default
:return: collection of matching symbols
:return: the symbols (with locations) matching the name path pattern
"""
# Note: file system sync not required; the symbol finder opens all relevant source files explicitly in the case of changes
if include_body:
depth = 0 # ignore user-specified depth if include_body is True
assert max_matches != 0, "max_matches must be > 0 or equal to -1."
parsed_include_kinds: Sequence[SymbolKind] | None = [SymbolKind(k) for k in include_kinds] if include_kinds else None
parsed_exclude_kinds: Sequence[SymbolKind] | None = [SymbolKind(k) for k in exclude_kinds] if exclude_kinds else None
symbol_retriever = self._create_symbol_retriever()
symbols = symbol_retriever.find(
name_path_pattern,
include_kinds=parsed_include_kinds,
exclude_kinds=parsed_exclude_kinds,
include_kinds=self._parse_kinds(include_kinds),
exclude_kinds=self._parse_kinds(exclude_kinds),
substring_matching=substring_matching,
within_relative_path=relative_path,
)
@@ -208,3 +459,298 @@ class LspApi(FacadeApi):
)
return symbol_collection
def find_referencing_symbols(
self,
name_path: str,
relative_path: str,
include_kinds: Sequence[int] = (),
exclude_kinds: Sequence[int] = (),
max_answer_chars: int = -1,
) -> LspReferenceCollection:
"""
Finds references to the symbol at the given `name_path`. The result will contain metadata about the referencing symbols
as well as a short code snippet around the reference.
:param name_path: name path of the symbol
:param relative_path: the relative path to the file containing the symbol for which to find references.
:param include_kinds: (optional) limits results to the given LSP symbol kinds (integers)
:param exclude_kinds: optional list of LSP symbol kinds (integers) to exclude.
:param max_answer_chars: max result length; -1 for default
:return: the references to the symbol
"""
# file system sync needed for case where symbol finder does not perform a global search, updating everything
if relative_path:
self._get_project().ls_sync_file_system_changes()
symbol_retriever = self._create_symbol_retriever()
references = symbol_retriever.find_referencing_symbols(
name_path,
relative_file_path=relative_path,
include_body=False, # it is probably never a good idea to include the body of the referencing symbols
include_kinds=self._parse_kinds(include_kinds),
exclude_kinds=self._parse_kinds(exclude_kinds),
)
return LspReferenceCollection(references, LspReferenceCollectionRenderer(self._agent, max_answer_chars, self.references_grouper_))
def find_implementations(
self,
name_path: str,
relative_path: str,
include_info: bool = False,
include_kinds: Sequence[int] = (),
exclude_kinds: Sequence[int] = (),
max_answer_chars: int = -1,
) -> LspSymbolCollection:
"""
Finds implementations of the symbol at the given `name_path`.
:param name_path: the symbol's name path
:param relative_path: the relative path to the file containing the symbol for which to find implementations.
Note that here you can't pass a directory but must pass a file.
:param include_info: whether to include additional info (hover-like, typically including docstring and signature),
about the implementing symbols.
:param include_kinds: (optional) limits results to the given LSP symbol kinds (integers)
:param exclude_kinds: (optional) list of LSP symbol kinds (integers) to exclude.
:param max_answer_chars: max result length; -1 for default
:return: the symbols implementing the given symbol
"""
self._get_project().ls_sync_file_system_changes()
symbol_retriever = self._create_symbol_retriever()
symbols = symbol_retriever.find_implementing_symbols(
name_path,
relative_file_path=relative_path,
include_body=False,
include_kinds=self._parse_kinds(include_kinds),
exclude_kinds=self._parse_kinds(exclude_kinds),
)
output_params = SymbolOutputParams(kind=True, relative_path=True, body_location=True, include_info=include_info)
return LspSymbolCollection(symbols, LspSymbolCollectionRenderer(self._agent, max_answer_chars, symbol_retriever, output_params))
def find_declaration(
self,
relative_path: str,
regex: str,
containing_symbol_name_path: str | None = None,
include_body: bool = False,
include_info: bool = False,
) -> LspSymbol:
r"""
Finds the declaration of a symbol.
:param relative_path: the relative path to the source file containing the symbol for which to find the declaration.
:param regex: a regular expression with one group, where the group matches the symbol for which to perform the lookup.
For example, to find the declaration of the `process` method in a call like `obj.process()`,
pass an expression like "obj\.(process)\(process_input_arg=37\)".
Prefer regexes with sufficiently large context around the group to render the match unambiguous.
Uses Python syntax with MULTILINE and DOTALL flags enabled.
:param containing_symbol_name_path: optional name path of a containing symbol whose body shall be searched instead of the full file.
:param include_body: whether to include the symbol's body in the result. Default False.
:param include_info: whether to include additional info (hover-like). Default False.
:return: the declaring symbol
"""
self._get_project().ls_sync_file_system_changes()
symbol_retriever = self._create_symbol_retriever()
# find relevant location for lookup
editor = self._create_code_editor(symbol_retriever)
if not containing_symbol_name_path:
content = editor.read_file(relative_path)
coords = find_text_coordinates(content, regex, require_unique=True)
assert coords is not None
else:
symbol = symbol_retriever.find_unique(name_path_pattern=containing_symbol_name_path, within_relative_path=relative_path)
body_line_numbers = symbol.get_body_line_numbers_or_raise()
content = editor.read_file(relative_path, lines=body_line_numbers)
coords = find_text_coordinates(content, regex, require_unique=True)
assert coords is not None
coords.line += body_line_numbers[0]
# retrieve declaration
defining_symbol = symbol_retriever.find_declaration(
relative_file_path=relative_path, line=coords.line, column=coords.col, include_body=include_body
)
if defining_symbol is None:
raise ValueError(
f"No symbol declaration found at the location of the regex match. Location: {relative_path}:{coords.line}:{coords.col}."
)
output_params = SymbolOutputParams(
kind=True, relative_path=True, body_location=True, include_body=include_body, include_info=include_info
)
collection_renderer = LspSymbolCollectionRenderer(self._agent, -1, symbol_retriever, output_params)
return LspSymbol(defining_symbol, LspSymbolRenderer(self._agent, -1, collection_renderer))
def get_diagnostics_for_file(
self, relative_path: str, start_line: int = 0, end_line: int = -1, min_severity: int = 4, max_answer_chars: int = -1
) -> LspDiagnostics:
"""
Gets diagnostics for a file. Diagnostics are grouped as `relative_path -> severity -> name_path -> diagnostics_results`.
If a diagnostic cannot be mapped to a symbol, it is grouped under the special name path `<file>`.
:param relative_path: the relative path to the file to inspect.
:param start_line: the first 0-based line to include. Defaults to 0.
:param end_line: the last 0-based line to include. Defaults to -1, which means until the end of the file.
:param min_severity: minimum LSP severity to include, where 1=Error, 2=Warning, 3=Information, 4=Hint.
Diagnostics with lower-or-equal numeric severity are returned.
:param max_answer_chars: max result length; -1 for default
:return: the grouped diagnostics for the requested file.
"""
self._get_project().ls_sync_file_system_changes()
symbol_retriever = self._create_symbol_retriever()
diagnostics = symbol_retriever.get_file_diagnostics(
relative_file_path=relative_path, start_line=start_line, end_line=end_line, min_severity=min_severity
)
grouped_diagnostics = GroupedDiagnostics()
for diagnostic in diagnostics:
diag_start = diagnostic["range"]["start"]
owner_symbol = symbol_retriever.find_diagnostic_owner_symbol(
relative_file_path=relative_path, line=diag_start["line"], column=diag_start["character"]
)
name_path = owner_symbol.get_name_path() if owner_symbol is not None else self.FILE_LEVEL_DIAGNOSTIC_BUCKET
grouped_diagnostics.add(relative_path, name_path, diagnostic)
return self._create_diagnostics(grouped_diagnostics, max_answer_chars)
def get_diagnostics_for_symbol(
self,
name_path: str,
reference_file: str = "",
check_symbol_references: bool = False,
min_severity: int = 4,
max_answer_chars: int = -1,
) -> LspDiagnostics:
"""
Gets diagnostics for the specified symbol. When `check_symbol_references` is true, diagnostics for all
referencing symbols are also included. The result is grouped as
`relative_path -> severity -> name_path -> diagnostics_results`.
:param name_path: the name path of the symbol to inspect.
:param reference_file: optional file path used to disambiguate the symbol search.
:param check_symbol_references: whether to additionally collect diagnostics for symbols that reference the symbol.
:param min_severity: minimum LSP severity to include, where 1=Error, 2=Warning, 3=Information, 4=Hint.
Diagnostics with lower-or-equal numeric severity are returned.
:param max_answer_chars: max result length; -1 for default
:return: the grouped diagnostics for the requested symbol and, optionally, its referencing symbols.
"""
self._get_project().ls_sync_file_system_changes()
symbol_retriever = self._create_symbol_retriever()
diagnostics_by_symbol = symbol_retriever.get_symbol_diagnostics(
name_path=name_path,
reference_file=reference_file or None,
check_symbol_references=check_symbol_references,
min_severity=min_severity,
)
grouped_diagnostics = GroupedDiagnostics()
for symbol, diagnostics in diagnostics_by_symbol.items():
relative_path = symbol.relative_path
if relative_path is None:
continue
for diagnostic in diagnostics:
grouped_diagnostics.add(relative_path, symbol.get_name_path(), diagnostic)
return self._create_diagnostics(grouped_diagnostics, max_answer_chars)
# edit operations
def replace_symbol_body(self, name_path: str, relative_path: str, body: str) -> str:
"""
Replaces the body of the given symbol.
IMPORTANT: Only replace symbol bodies if you have previously made a retrieval with include_body=True and thus know what
constitutes the body!
:param name_path: name path of the symbol whose body to replace
:param relative_path: the relative path to the file containing the symbol
:param body: the new symbol body. The symbol body is the definition of a symbol
in the programming language, including e.g. the signature line for functions.
Depending on the language, it may or may not include a preceding docstring or other preceding annotations.
:return: a success message
"""
self._create_code_editor().replace_body(name_path, relative_file_path=relative_path, body=body)
return SUCCESS_RESULT
def insert_after_symbol(self, name_path: str, relative_path: str, body: str) -> str:
"""
Inserts code after a class/method/function definition.
Don't use this to insert after assignments (constants, fields).
:param name_path: name path of the symbol after which to insert content
:param relative_path: the relative path to the file containing the symbol
:param body: the body/content to be inserted. The inserted code shall begin with the next line after
the symbol.
:return: a success message
"""
self._create_code_editor().insert_after_symbol(name_path, relative_file_path=relative_path, body=body)
return SUCCESS_RESULT
def insert_before_symbol(self, name_path: str, relative_path: str, body: str) -> str:
"""
Inserts the given content before the beginning of the definition of the given symbol (via the symbol's location).
A typical use case is to insert a new class, function, method, field or variable assignment; or
a new import statement before the first symbol in the file.
:param name_path: name path of the symbol before which to insert content
:param relative_path: the relative path to the file containing the symbol
:param body: the body/content to be inserted before the line in which the referenced symbol is defined
:return: a success message
"""
self._create_code_editor().insert_before_symbol(name_path, relative_file_path=relative_path, body=body)
return SUCCESS_RESULT
def rename_symbol(self, name_path: str, relative_path: str, new_name: str) -> str:
"""
Renames the symbol with the given `name_path` to `new_name` throughout the entire codebase.
Note: for languages with method overloading, like Java, name_path may have to include a method's
signature to uniquely identify a method.
:param name_path: name path of the symbol to rename
:param relative_path: the relative path to the file containing the symbol to rename
:param new_name: the new name for the symbol
:return: a result summary indicating success or failure
"""
self._get_project().ls_sync_file_system_changes()
return self._create_code_editor().rename_symbol(name_path, relative_path=relative_path, new_name=new_name)
def safe_delete_symbol(self, name_path_pattern: str, relative_path: str) -> str:
"""
Deletes the symbol if it is safe to do so (i.e., if there are no references to it)
or returns a list of references to it.
:param name_path_pattern: name path of the symbol to delete
:param relative_path: the relative path to the file containing the symbol to delete
:return: a success message, or a message listing the references preventing deletion
"""
self._get_project().ls_sync_file_system_changes()
symbol_retriever = self._create_symbol_retriever()
symbol = symbol_retriever.find_unique(name_path_pattern, substring_matching=False, within_relative_path=relative_path)
symbol_rel_path = symbol.relative_path
assert symbol_rel_path is not None, f"Symbol {name_path_pattern} has no relative path, this is likely a bug."
assert symbol_rel_path == relative_path, f"Symbol {name_path_pattern} is not in the expected relative path {relative_path}."
symbol_name_path = symbol.get_name_path()
# check for references
symbol_line = symbol.line
symbol_col = symbol.column
assert symbol_line is not None and symbol_col is not None, (
f"Symbol {name_path_pattern} has no identifier position, this is likely a bug."
)
lang_server = symbol_retriever.get_language_server(symbol_rel_path)
references_locations = lang_server.request_references(symbol_rel_path, symbol_line, symbol_col)
file_to_lines: dict[str, list[int]] = defaultdict(list)
for ref_loc in references_locations or []:
ref_relative_path = ref_loc.get("relativePath")
if ref_relative_path is None:
continue
file_to_lines[ref_relative_path].append(ref_loc["range"]["start"]["line"])
if file_to_lines:
return f"Cannot delete, the symbol {symbol_name_path} is referenced in: {TextOutputUtils.to_json(file_to_lines)}"
self._create_code_editor(symbol_retriever).delete_symbol(symbol_name_path, relative_file_path=symbol_rel_path)
return SUCCESS_RESULT
+3
View File
@@ -14,6 +14,9 @@ from serena.project import Project
if TYPE_CHECKING:
from serena.agent import SerenaAgent
SUCCESS_RESULT = "OK"
"""the result returned by operations which have no result other than their success"""
class FacadeApi(ABC):
"""
@@ -3,12 +3,14 @@
import json
from collections.abc import Iterable
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any
from typing import TYPE_CHECKING, Any, Optional, Self
from serena.util.text_utils import TextOutputUtils
from solidlsp import ls_types
from solidlsp.lsp_protocol_handler.lsp_types import DiagnosticSeverity
if TYPE_CHECKING:
from serena.agent import SerenaAgent
from serena.symbol import LanguageServerSymbolRetriever
@@ -203,3 +205,55 @@ class DiagnosticsDiff:
def get_grouped_diagnostics(self) -> GroupedDiagnostics:
return self._grouped_diagnostics
class DiagnosticsContext:
ENABLE_DIAGNOSTICS_DEFAULT: bool = False
"""
Global flag to enable/disable diagnostics for LSP-based editing tools derived from this class.
The feature is currently disabled, because per-edit diagnostics are a questionable feature, since individual
edits often intentionally introduce diagnostics (e.g. function signature mismatches or even syntax errors) that
are then resolved in subsequent edits.
"""
DIAGNOSTICS_KEY = "diagnostics[warning-or-higher]"
def __init__(self, agent: "SerenaAgent", *edited_relative_paths: str, enable: bool = ENABLE_DIAGNOSTICS_DEFAULT) -> None:
self._is_diagnostics_enabled = enable and agent.is_using_language_server()
self._edited_files = [EditedFilePath(path, path) for path in edited_relative_paths]
self._before_edit_diagnostics_snapshot: PublishedDiagnosticsSnapshot | None = None
self._symbol_retriever: Optional["LanguageServerSymbolRetriever"] | None = None
if self._is_diagnostics_enabled:
from serena.symbol import LanguageServerSymbolRetriever # local import to avoid a circular dependency
self._symbol_retriever = LanguageServerSymbolRetriever(agent.get_active_project_or_raise())
self._before_edit_diagnostics_snapshot = PublishedDiagnosticsSnapshot(self._edited_files, self._symbol_retriever)
def __enter__(self) -> Self:
return self
def __exit__(self, exc_type, exc_val, exc_tb):
pass
def format_result(
self,
base_result: str,
) -> str:
if not self._is_diagnostics_enabled:
return base_result
if self._before_edit_diagnostics_snapshot is None:
return base_result
assert self._symbol_retriever is not None
diagnostics_diff = DiagnosticsDiff(self._before_edit_diagnostics_snapshot, self._edited_files, self._symbol_retriever)
grouped_diagnostics = diagnostics_diff.get_grouped_diagnostics().get_dict()
if not grouped_diagnostics:
return base_result
else:
result_dict = {
"result": base_result,
self.DIAGNOSTICS_KEY: grouped_diagnostics,
}
return TextOutputUtils.to_json(result_dict)
+6 -6
View File
@@ -69,7 +69,7 @@ class CreateTextFileTool(EditingToolWithDiagnostics):
:param content: the (appropriately encoded) content to write to the file
:return: a message indicating success or failure
"""
with self.DiagnosticsContext(self, relative_path) as diagnostics_context:
with self.diagnostics_context(relative_path) as diagnostics_context:
# validating the destination path
project_root = self.get_project_root()
abs_path = (Path(project_root) / relative_path).resolve()
@@ -206,7 +206,7 @@ class ReplaceContentTool(EditingToolWithDiagnostics):
:param allow_multiple_occurrences: whether to allow matching and replacing multiple occurrences.
If false and multiple occurrences are found, an error will be returned
"""
with self.DiagnosticsContext(self, relative_path) as diagnostics_context:
with self.diagnostics_context(relative_path) as diagnostics_context:
self.project.validate_relative_path(relative_path)
with EditedFileContext(relative_path, self.create_code_editor()) as context:
original_content = context.get_original_content()
@@ -427,7 +427,7 @@ class ReplaceInFilesTool(EditingToolWithDiagnostics):
occurrences_by_file: dict[str, list[ReplacementOccurrence]] = {}
for occ in occurrences:
occurrences_by_file.setdefault(occ.relative_path, []).append(occ)
with self.DiagnosticsContext(self, *occurrences_by_file.keys()) as diagnostics_context:
with self.diagnostics_context() as diagnostics_context:
code_editor = self.create_code_editor()
for path, file_occurrences in occurrences_by_file.items():
with EditedFileContext(path, code_editor) as context:
@@ -469,7 +469,7 @@ class DeleteLinesTool(EditingToolWithDiagnostics, ToolMarkerOptional):
:param start_line: the 0-based index of the first line to be deleted
:param end_line: the 0-based index of the last line to be deleted
"""
with self.DiagnosticsContext(self, relative_path) as diagnostics_context:
with self.diagnostics_context(relative_path) as diagnostics_context:
code_editor = self.create_code_editor()
code_editor.delete_lines(relative_path, start_line, end_line)
return diagnostics_context.format_result(SUCCESS_RESULT)
@@ -501,7 +501,7 @@ class ReplaceLinesTool(EditingToolWithDiagnostics, ToolMarkerOptional):
if not content.endswith("\n"):
content += "\n"
with self.DiagnosticsContext(self, relative_path) as diagnostics_context:
with self.diagnostics_context(relative_path) as diagnostics_context:
code_editor = self.create_code_editor()
code_editor.delete_lines(relative_path, start_line, end_line)
code_editor.insert_at_line(relative_path, start_line, content)
@@ -534,7 +534,7 @@ class InsertAtLineTool(EditingToolWithDiagnostics, ToolMarkerOptional):
if not content.endswith("\n"):
content += "\n"
with self.DiagnosticsContext(self, relative_path) as diagnostics_context:
with self.diagnostics_context(relative_path) as diagnostics_context:
code_editor = self.create_code_editor()
code_editor.insert_at_line(relative_path, line, content)
+21 -16
View File
@@ -1,24 +1,29 @@
# SPDX-License-Identifier: GPL-3.0-or-later
import logging
from typing import Literal
from typing import TYPE_CHECKING, Literal
from serena.facades.api.jb import JetBrainsApi
from serena.tools import Tool, ToolMarkerBeta, ToolMarkerOptional, ToolMarkerSymbolicEdit, ToolMarkerSymbolicRead
if TYPE_CHECKING:
from serena.agent import SerenaAgent
log = logging.getLogger(__name__)
class JetBrainsTool(Tool):
class JetBrainsApiMixin:
"""
Base class for tools which delegate to the JetBrains API
Mixin for tools which delegate to the JetBrains API
"""
agent: "SerenaAgent"
def _api(self) -> JetBrainsApi:
return JetBrainsApi(self.agent)
class JetBrainsFindSymbolTool(JetBrainsTool, ToolMarkerSymbolicRead, ToolMarkerOptional):
class JetBrainsFindSymbolTool(Tool, ToolMarkerSymbolicRead, ToolMarkerOptional, JetBrainsApiMixin):
"""
Performs a global (or local) search for symbols using the JetBrains backend
"""
@@ -106,7 +111,7 @@ class JetBrainsFindSymbolTool(JetBrainsTool, ToolMarkerSymbolicRead, ToolMarkerO
return {"name_path": "name_path_pattern"}
class JetBrainsMoveTool(JetBrainsTool, ToolMarkerSymbolicEdit, ToolMarkerOptional, ToolMarkerBeta):
class JetBrainsMoveTool(Tool, ToolMarkerSymbolicEdit, ToolMarkerOptional, ToolMarkerBeta, JetBrainsApiMixin):
"""
Moves a symbol, file or directory to a new location using the JetBrains backend, updating all references
"""
@@ -146,7 +151,7 @@ class JetBrainsMoveTool(JetBrainsTool, ToolMarkerSymbolicEdit, ToolMarkerOptiona
return self._api().move(relative_path, name_path, target_relative_path, target_parent_name_path).represent()
class JetBrainsSafeDeleteTool(JetBrainsTool, ToolMarkerSymbolicEdit, ToolMarkerOptional, ToolMarkerBeta):
class JetBrainsSafeDeleteTool(Tool, ToolMarkerSymbolicEdit, ToolMarkerOptional, ToolMarkerBeta, JetBrainsApiMixin):
"""
Safely deletes a symbol using the JetBrains backend, checking for remaining usages first
"""
@@ -178,7 +183,7 @@ class JetBrainsSafeDeleteTool(JetBrainsTool, ToolMarkerSymbolicEdit, ToolMarkerO
return self._api().safe_delete(relative_path, name_path, delete_even_if_used, propagate).represent()
class JetBrainsInlineSymbol(JetBrainsTool, ToolMarkerSymbolicEdit, ToolMarkerOptional, ToolMarkerBeta):
class JetBrainsInlineSymbol(Tool, ToolMarkerSymbolicEdit, ToolMarkerOptional, ToolMarkerBeta, JetBrainsApiMixin):
"""
Inlines a symbol using the JetBrains backend, replacing all call sites with the symbol's body
"""
@@ -205,7 +210,7 @@ class JetBrainsInlineSymbol(JetBrainsTool, ToolMarkerSymbolicEdit, ToolMarkerOpt
return self._api().inline_symbol(name_path, relative_path, keep_definition).represent()
class JetBrainsFindReferencingSymbolsTool(JetBrainsTool, ToolMarkerSymbolicRead, ToolMarkerOptional):
class JetBrainsFindReferencingSymbolsTool(Tool, ToolMarkerSymbolicRead, ToolMarkerOptional, JetBrainsApiMixin):
"""
Finds symbols that reference the given symbol using the JetBrains backend
"""
@@ -231,7 +236,7 @@ class JetBrainsFindReferencingSymbolsTool(JetBrainsTool, ToolMarkerSymbolicRead,
return self._api().find_referencing_symbols(name_path, relative_path, max_answer_chars).represent()
class JetBrainsGetSymbolsOverviewTool(JetBrainsTool, ToolMarkerSymbolicRead, ToolMarkerOptional):
class JetBrainsGetSymbolsOverviewTool(Tool, ToolMarkerSymbolicRead, ToolMarkerOptional, JetBrainsApiMixin):
"""
Retrieves an overview of the top-level symbols within a specified file using the JetBrains backend
"""
@@ -258,7 +263,7 @@ class JetBrainsGetSymbolsOverviewTool(JetBrainsTool, ToolMarkerSymbolicRead, Too
return self._api().get_symbols_overview(relative_path, depth, max_answer_chars, include_file_documentation).represent()
class JetBrainsTypeHierarchyTool(JetBrainsTool, ToolMarkerSymbolicRead, ToolMarkerOptional):
class JetBrainsTypeHierarchyTool(Tool, ToolMarkerSymbolicRead, ToolMarkerOptional, JetBrainsApiMixin):
"""
Retrieves the type hierarchy (supertypes and/or subtypes) of a symbol using the JetBrains backend
"""
@@ -287,7 +292,7 @@ class JetBrainsTypeHierarchyTool(JetBrainsTool, ToolMarkerSymbolicRead, ToolMark
return self._api().get_type_hierarchy(name_path, relative_path, hierarchy_type, depth, max_answer_chars).represent()
class JetBrainsFindDeclarationTool(JetBrainsTool, ToolMarkerSymbolicRead, ToolMarkerOptional):
class JetBrainsFindDeclarationTool(Tool, ToolMarkerSymbolicRead, ToolMarkerOptional, JetBrainsApiMixin):
"""
Finds the declaration of a symbol using the JetBrains backend
"""
@@ -309,7 +314,7 @@ class JetBrainsFindDeclarationTool(JetBrainsTool, ToolMarkerSymbolicRead, ToolMa
return self._api().find_declaration(relative_path, regex, include_body).represent()
class JetBrainsFindImplementationsTool(JetBrainsTool, ToolMarkerSymbolicRead, ToolMarkerOptional):
class JetBrainsFindImplementationsTool(Tool, ToolMarkerSymbolicRead, ToolMarkerOptional, JetBrainsApiMixin):
"""
Finds the implementations of a symbol using the JetBrains backend
"""
@@ -324,7 +329,7 @@ class JetBrainsFindImplementationsTool(JetBrainsTool, ToolMarkerSymbolicRead, To
return self._api().find_implementations(relative_path, name_path).represent()
class JetBrainsRenameTool(JetBrainsTool, ToolMarkerSymbolicEdit, ToolMarkerOptional):
class JetBrainsRenameTool(Tool, ToolMarkerSymbolicEdit, ToolMarkerOptional, JetBrainsApiMixin):
"""
Renames a symbol, file or directory throughout the codebase using the JetBrains backend.
"""
@@ -353,7 +358,7 @@ class JetBrainsRenameTool(JetBrainsTool, ToolMarkerSymbolicEdit, ToolMarkerOptio
return self._api().rename(relative_path, new_name, name_path, rename_in_comments, rename_in_text_occurrences).represent()
class JetBrainsDebugTool(JetBrainsTool, ToolMarkerOptional, ToolMarkerBeta):
class JetBrainsDebugTool(Tool, ToolMarkerOptional, ToolMarkerBeta, JetBrainsApiMixin):
"""
Provides debugging functionality (run configs, breakpoints, stepping, inspection, and evaluation)
via a persistent debug REPL connected to the JetBrains IDE.
@@ -378,7 +383,7 @@ class JetBrainsDebugTool(JetBrainsTool, ToolMarkerOptional, ToolMarkerBeta):
return self._api().debug_eval(expression, repl_key)
class JetBrainsRunInspectionsTool(JetBrainsTool, ToolMarkerSymbolicRead, ToolMarkerOptional):
class JetBrainsRunInspectionsTool(Tool, ToolMarkerSymbolicRead, ToolMarkerOptional, JetBrainsApiMixin):
"""
Runs JetBrains IDE inspections on a file and returns the results.
"""
@@ -413,7 +418,7 @@ class JetBrainsRunInspectionsTool(JetBrainsTool, ToolMarkerSymbolicRead, ToolMar
)
class JetBrainsListInspectionsTool(JetBrainsTool, ToolMarkerSymbolicRead, ToolMarkerOptional):
class JetBrainsListInspectionsTool(Tool, ToolMarkerSymbolicRead, ToolMarkerOptional, JetBrainsApiMixin):
"""
Lists available JetBrains IDE inspections, optionally filtered by language or group.
"""
+106 -364
View File
@@ -3,45 +3,46 @@ Language server-related tools
"""
# SPDX-License-Identifier: GPL-3.0-or-later
import copy
import os
from collections import Counter, defaultdict
from collections.abc import Callable, Sequence
from typing import Any
from typing import TYPE_CHECKING, cast
from serena.facades.api.lsp import LspApi
from serena.symbol import LanguageServerSymbol, LanguageServerSymbolDictGrouper
from serena.tools import (
SUCCESS_RESULT,
EditingToolWithDiagnostics,
Tool,
ToolMarkerSymbolicEdit,
ToolMarkerSymbolicRead,
)
from serena.tools.tools_base import ToolMarkerOptional
from serena.util.ls_diagnostics import GroupedDiagnostics
from serena.util.text_utils import find_text_coordinates
from solidlsp.ls_types import SymbolKind
if TYPE_CHECKING:
pass
class RestartLanguageServerTool(Tool, ToolMarkerOptional):
class LspApiMixin:
"""
Mixin for tools which delegate to the language server API
"""
def _api(self) -> LspApi:
tool = cast(Tool, cast(object, self))
return LspApi(tool.agent)
class RestartLanguageServerTool(Tool, ToolMarkerOptional, LspApiMixin):
"""Restarts the language server(s)."""
def apply(self) -> str:
"""Use this tool only on explicit user request or after confirmation.
It may be necessary to restart the language server if it hangs.
"""
self.agent.reset_language_server_manager()
return SUCCESS_RESULT
return self._api().restart_language_server()
class GetSymbolsOverviewTool(Tool, ToolMarkerSymbolicRead):
class GetSymbolsOverviewTool(Tool, ToolMarkerSymbolicRead, LspApiMixin):
"""
Gets an overview of the top-level symbols defined in a given file.
"""
symbol_dict_grouper = LanguageServerSymbolDictGrouper(["kind"], ["kind"], collapse_singleton=True)
def apply(self, relative_path: str, depth: int = -1, max_answer_chars: int = -1) -> str:
"""
Use this tool to get a high-level understanding of the code symbols in a file.
@@ -56,96 +57,14 @@ class GetSymbolsOverviewTool(Tool, ToolMarkerSymbolicRead):
Don't adjust unless there is really no other way to get the content required for the task.
:return: a JSON object containing symbols grouped by kind in a compact format.
"""
# Note: file system sync not required (relevant file is opened in the language server explicitly)
if depth == -1:
if relative_path.endswith((".java", ".kt")):
depth = 1
else:
depth = 0
result = self.get_symbol_overview(relative_path, depth=depth)
# capture kind names and depth-0 snapshots before grouping, which mutates the dicts
kind_names = [d.get("kind", "unknown") for d in result]
if depth > 0:
depth_0_result = [d.copy() for d in result]
for d in depth_0_result:
d.pop("children", None)
compact_result = self.symbol_dict_grouper.group(result)
result_json_str = self._to_json(compact_result)
# shortened result closures
def make_kind_counts() -> str:
return f"Symbol counts by kind:\n{self._to_json(Counter(kind_names))}"
shortened_results: list[Callable[[], str]]
if depth == 0:
shortened_results = [make_kind_counts]
else:
def make_depth_0_result() -> str:
compact_depth_0_result = self.symbol_dict_grouper.group(depth_0_result)
return "Depth 0 overview:\n" + self._to_json(compact_depth_0_result)
shortened_results = [make_depth_0_result, make_kind_counts]
return self._limit_length(result_json_str, max_answer_chars, shortened_result_factories=shortened_results)
def get_symbol_overview(self, relative_path: str, depth: int = 0) -> list[LanguageServerSymbol.OutputDict]:
"""
:param relative_path: relative path to a source file
:param depth: the depth up to which descendants shall be retrieved
:return: a list of symbol dictionaries representing the symbol overview of the file
"""
symbol_retriever = self.create_language_server_symbol_retriever()
# The symbol overview is capable of working with both files and directories,
# but we want to ensure that the user provides a file path.
file_path = os.path.join(self.project.project_root, relative_path)
if not os.path.exists(file_path):
raise FileNotFoundError(f"File or directory {relative_path} does not exist in the project.")
if os.path.isdir(file_path):
raise ValueError(f"Expected a file path, but got a directory path: {relative_path}. ")
if not symbol_retriever.can_analyze_file(relative_path):
raise ValueError(
f"Cannot extract symbols from file {relative_path}. Active language servers: {[l.value for l in self.agent.get_active_language_server_ids()]}"
)
symbols = symbol_retriever.get_symbol_overview(relative_path)[relative_path]
def child_inclusion_predicate(s: LanguageServerSymbol) -> bool:
return not s.is_low_level()
symbol_dicts = []
for symbol in symbols:
symbol_dicts.append(
symbol.to_dict(
name_path=False,
name=True,
depth=depth,
kind=True,
relative_path=False,
location=False,
child_inclusion_predicate=child_inclusion_predicate,
)
)
return symbol_dicts
return self._api().get_symbols_overview(relative_path, depth=depth, max_answer_chars=max_answer_chars).represent()
class FindSymbolTool(Tool, ToolMarkerSymbolicRead):
class FindSymbolTool(Tool, ToolMarkerSymbolicRead, LspApiMixin):
"""
Performs a global (or local) search using the language server backend.
"""
symbol_dict_grouper = LspApi.find_symbol_dict_grouper_
"""
Reference to the grouper that is indirectly used by this tool.
Made explicit such that grouping behaviour for this tool can be modified dynamically.
"""
# noinspection PyDefaultArgument
def apply(
self,
name_path_pattern: str,
@@ -183,46 +102,48 @@ class FindSymbolTool(Tool, ToolMarkerSymbolicRead):
:param relative_path: (optional) restrict search to this file or directory. If None, searches entire codebase.
If a directory is passed, the search will be restricted to the files in that directory.
If a file is passed, the search will be restricted to that file.
:param include_body: whether to include the symbol's source code. Use judiciously.
If you have some knowledge about the codebase, you should use this parameter, as it will significantly
speed up the search as well as reduce the number of results.
:param include_body: If True, include the symbol's source code. Use judiciously.
:param include_info: whether to include additional info (hover-like, typically including docstring and signature),
about the symbol (ignored if include_body is True). Info is never included for child symbols.
Note: Depending on the language, this can be slow (e.g., C/C++).
:param include_kinds: (optional) limits results to the given LSP symbol kinds (integers)
:param exclude_kinds: (optional) list of LSP symbol kinds (integers) to exclude.
:param substring_matching: If True, use substring matching for the last element of the pattern, such that
"Foo/get" would match "Foo/getValue" and "Foo/getData".
:param max_matches: maximum number of permitted matches. If exceeded, a shortened result is returned
:param substring_matching: If True, use substring matching for the last segment of `name_path_pattern`
(i.e. the name of the symbol, e.g. "foo" in "Class/foo" or "my_method" in "my_method").
:param max_matches: Maximum number of permitted matches. If exceeded, a shortened result is returned
which allows refining the search. -1 (default) means no limit. Set to 1 to search for a unique symbol.
:param max_answer_chars: max result length; -1 for default
:return: symbols (with locations) matching the name.
"""
collection = LspApi(self.agent).find_symbol(
name_path_pattern,
depth=depth,
relative_path=relative_path,
include_body=include_body,
include_info=include_info,
include_kinds=include_kinds,
exclude_kinds=exclude_kinds,
substring_matching=substring_matching,
max_matches=max_matches,
max_answer_chars=max_answer_chars,
return (
self._api()
.find_symbol(
name_path_pattern,
depth=depth,
relative_path=relative_path,
include_body=include_body,
include_info=include_info,
include_kinds=include_kinds,
exclude_kinds=exclude_kinds,
substring_matching=substring_matching,
max_matches=max_matches,
max_answer_chars=max_answer_chars,
)
.represent()
)
return collection.represent()
@classmethod
def get_param_aliases(cls) -> dict[str, str]:
return {"name_path": "name_path_pattern"}
class FindReferencingSymbolsTool(Tool, ToolMarkerSymbolicRead):
class FindReferencingSymbolsTool(Tool, ToolMarkerSymbolicRead, LspApiMixin):
"""
Finds symbols that reference the given symbol using the language server backend
Finds symbols that reference the given symbol
"""
symbol_dict_grouper = LanguageServerSymbolDictGrouper(["relative_path", "kind"], ["kind"], collapse_singleton=True)
# noinspection PyDefaultArgument
def apply(
self,
name_path: str,
@@ -242,75 +163,20 @@ class FindReferencingSymbolsTool(Tool, ToolMarkerSymbolicRead):
:param max_answer_chars: max result length; -1 for default
:return: a list of JSON objects with the symbols referencing the requested symbol
"""
# file system sync needed for case where symbol finder does not perform a global search, updating everything
if relative_path:
self.project.ls_sync_file_system_changes()
include_body = False # It is probably never a good idea to include the body of the referencing symbols
parsed_include_kinds: Sequence[SymbolKind] | None = [SymbolKind(k) for k in include_kinds] if include_kinds else None
parsed_exclude_kinds: Sequence[SymbolKind] | None = [SymbolKind(k) for k in exclude_kinds] if exclude_kinds else None
symbol_retriever = self.create_language_server_symbol_retriever()
references_in_symbols = symbol_retriever.find_referencing_symbols(
name_path,
relative_file_path=relative_path,
include_body=include_body,
include_kinds=parsed_include_kinds,
exclude_kinds=parsed_exclude_kinds,
return (
self._api()
.find_referencing_symbols(
name_path, relative_path, include_kinds=include_kinds, exclude_kinds=exclude_kinds, max_answer_chars=max_answer_chars
)
.represent()
)
reference_dicts = []
for ref in references_in_symbols:
ref_dict_orig = ref.symbol.to_dict(kind=True, relative_path=True, depth=0, body=include_body, body_location=True)
ref_dict = dict(ref_dict_orig)
if not include_body:
ref_relative_path = ref.symbol.location.relative_path
assert ref_relative_path is not None, f"Referencing symbol {ref.symbol.name} has no relative path, this is likely a bug."
content_around_ref = self.project.retrieve_content_around_line(
relative_file_path=ref_relative_path, line=ref.line, context_lines_before=1, context_lines_after=1
)
ref_dict["content_around_reference"] = content_around_ref.to_display_string()
reference_dicts.append(ref_dict)
# capture lightweight reference data before grouping
ref_summaries = []
for ref, d in zip(references_in_symbols, reference_dicts, strict=True):
ref_summaries.append(
{
"name_path": d.get("name_path"),
"kind": d.get("kind"),
"relative_path": d.get("relative_path"),
"reference_line": ref.line,
}
)
result = self.symbol_dict_grouper.group(reference_dicts)
# shortened result closures, from least to most aggressive shortening
def make_refs_without_context() -> str:
"""References with name_path and reference line, without surrounding code lines"""
grouped = self.symbol_dict_grouper.group(copy.deepcopy(ref_summaries))
return f"References without surrounding lines:\n{self._to_json(grouped)}"
def make_per_file_counts() -> str:
counts = Counter(str(r["relative_path"]) for r in ref_summaries)
return f"Reference counts per file:\n{self._to_json(counts)}"
def make_summary() -> str:
return f"Found {len(ref_summaries)} references."
shortened_results: list[Callable[[], str]] = [make_refs_without_context, make_per_file_counts, make_summary]
result_json = self._to_json(result)
return self._limit_length(result_json, max_answer_chars, shortened_result_factories=shortened_results)
class FindImplementationsTool(Tool, ToolMarkerSymbolicRead):
class FindImplementationsTool(Tool, ToolMarkerSymbolicRead, LspApiMixin):
"""
Finds symbols that implement the given symbol using the language server backend.
Finds the implementations of a symbol
"""
# noinspection PyDefaultArgument
def apply(
self,
name_path: str,
@@ -333,36 +199,21 @@ class FindImplementationsTool(Tool, ToolMarkerSymbolicRead):
:param max_answer_chars: max result length; -1 for default
:return: a list of JSON objects with the symbols implementing the requested symbol
"""
self.project.ls_sync_file_system_changes()
include_body = False
parsed_include_kinds: Sequence[SymbolKind] | None = [SymbolKind(k) for k in include_kinds] if include_kinds else None
parsed_exclude_kinds: Sequence[SymbolKind] | None = [SymbolKind(k) for k in exclude_kinds] if exclude_kinds else None
symbol_retriever = self.create_language_server_symbol_retriever()
implementing_symbols = symbol_retriever.find_implementing_symbols(
name_path,
relative_file_path=relative_path,
include_body=include_body,
include_kinds=parsed_include_kinds,
exclude_kinds=parsed_exclude_kinds,
return (
self._api()
.find_implementations(
name_path,
relative_path,
include_info=include_info,
include_kinds=include_kinds,
exclude_kinds=exclude_kinds,
max_answer_chars=max_answer_chars,
)
.represent()
)
symbol_dicts = [
dict(s.to_dict(kind=True, relative_path=True, depth=0, body=include_body, body_location=True)) for s in implementing_symbols
]
if include_info:
info_by_symbol = symbol_retriever.request_info_for_symbol_batch(implementing_symbols)
for s, s_dict in zip(implementing_symbols, symbol_dicts, strict=True):
if symbol_info := info_by_symbol.get(s):
s_dict["info"] = symbol_info
s_dict.pop("name", None) # name is included in the info
result = self._to_json(symbol_dicts)
return self._limit_length(result, max_answer_chars)
class FindDeclarationTool(Tool, ToolMarkerSymbolicRead):
class FindDeclarationTool(Tool, ToolMarkerSymbolicRead, LspApiMixin):
"""
Finds the declaration/definition of a symbol
"""
@@ -388,70 +239,26 @@ class FindDeclarationTool(Tool, ToolMarkerSymbolicRead):
:param include_body: whether to include the symbol's body in the result. Default False.
:param include_info: whether to include additional info (hover-like). Default False.
"""
self.project.ls_sync_file_system_changes()
symbol_retriever = self.create_language_server_symbol_retriever()
relative_path = self._sanitize_input_param(relative_path)
regex = self._sanitize_input_param(regex)
# find relevant location for lookup
editor = self.create_code_editor()
if not containing_symbol_name_path:
content = editor.read_file(relative_path)
coords = find_text_coordinates(content, regex, require_unique=True)
assert coords is not None
else:
symbol = symbol_retriever.find_unique(name_path_pattern=containing_symbol_name_path, within_relative_path=relative_path)
body_line_numers = symbol.get_body_line_numbers_or_raise()
content = editor.read_file(relative_path, lines=body_line_numers)
coords = find_text_coordinates(content, regex, require_unique=True)
assert coords is not None
coords.line += body_line_numers[0]
# retrieve declaration
defining_symbol = symbol_retriever.find_declaration(
relative_file_path=relative_path,
line=coords.line,
column=coords.col,
include_body=include_body,
)
if defining_symbol is None:
raise ValueError(
f"No symbol declaration found at the location of the regex match. Location: {relative_path}:{coords.line}:{coords.col}."
return (
self._api()
.find_declaration(
relative_path,
regex,
containing_symbol_name_path=containing_symbol_name_path,
include_body=include_body,
include_info=include_info,
)
# create output
symbol_dict = self._defining_symbol_to_result_dict(
symbol_retriever,
defining_symbol,
include_body,
include_info,
.represent()
)
result = self._to_json(symbol_dict)
return result
@staticmethod
def _defining_symbol_to_result_dict(
symbol_retriever: Any,
defining_symbol: LanguageServerSymbol,
include_body: bool,
include_info: bool,
) -> dict[str, Any]:
symbol_dict = dict(defining_symbol.to_dict(kind=True, relative_path=True, depth=0, body=include_body, body_location=True))
if not include_body and include_info:
if symbol_info := symbol_retriever.request_info_for_symbol(defining_symbol):
symbol_dict["info"] = symbol_info
symbol_dict.pop("name", None)
return symbol_dict
class GetDiagnosticsForFileTool(Tool, ToolMarkerSymbolicRead):
class GetDiagnosticsForFileTool(Tool, ToolMarkerSymbolicRead, LspApiMixin):
"""
Gets diagnostics for a file, optionally restricted to a line range, grouped by file, severity, and containing symbol.
Gets diagnostics for a file, grouped by symbol.
"""
FILE_LEVEL_DIAGNOSTIC_BUCKET = "<file>"
def apply(
self,
relative_path: str,
@@ -472,34 +279,16 @@ class GetDiagnosticsForFileTool(Tool, ToolMarkerSymbolicRead):
:param max_answer_chars: max result length; -1 for default
:return: grouped diagnostics for the requested file.
"""
self.project.ls_sync_file_system_changes()
symbol_retriever = self.create_language_server_symbol_retriever()
diagnostics = symbol_retriever.get_file_diagnostics(
relative_file_path=relative_path,
start_line=start_line,
end_line=end_line,
min_severity=min_severity,
return (
self._api()
.get_diagnostics_for_file(
relative_path, start_line=start_line, end_line=end_line, min_severity=min_severity, max_answer_chars=max_answer_chars
)
.represent()
)
grouped_diagnostics = GroupedDiagnostics()
for diagnostic in diagnostics:
diag_range = diagnostic["range"]["start"]
name_path = self.FILE_LEVEL_DIAGNOSTIC_BUCKET
owner_symbol = symbol_retriever.find_diagnostic_owner_symbol(
relative_file_path=relative_path,
line=diag_range["line"],
column=diag_range["character"],
)
if owner_symbol is not None:
name_path = owner_symbol.get_name_path()
grouped_diagnostics.add(relative_path, name_path, diagnostic)
result = self._to_json(grouped_diagnostics.get_dict())
return self._limit_length(result, max_answer_chars)
class GetDiagnosticsForSymbolTool(Tool, ToolMarkerSymbolicRead, ToolMarkerOptional):
class GetDiagnosticsForSymbolTool(Tool, ToolMarkerSymbolicRead, ToolMarkerOptional, LspApiMixin):
"""
Gets diagnostics for a symbol and, optionally, for symbols that reference it.
"""
@@ -525,30 +314,20 @@ class GetDiagnosticsForSymbolTool(Tool, ToolMarkerSymbolicRead, ToolMarkerOption
:param max_answer_chars: max result length; -1 for default
:return: grouped diagnostics for the requested symbol and, optionally, its referencing symbols.
"""
self.project.ls_sync_file_system_changes()
symbol_retriever = self.create_language_server_symbol_retriever()
diagnostics_by_symbol = symbol_retriever.get_symbol_diagnostics(
name_path=name_path,
reference_file=reference_file or None,
check_symbol_references=check_symbol_references,
min_severity=min_severity,
return (
self._api()
.get_diagnostics_for_symbol(
name_path,
reference_file=reference_file,
check_symbol_references=check_symbol_references,
min_severity=min_severity,
max_answer_chars=max_answer_chars,
)
.represent()
)
grouped_diagnostics = GroupedDiagnostics()
for symbol, diagnostics in diagnostics_by_symbol.items():
relative_path = symbol.relative_path
if relative_path is None:
continue
symbol_name_path = symbol.get_name_path()
for diagnostic in diagnostics:
grouped_diagnostics.add(relative_path, symbol_name_path, diagnostic)
result = self._to_json(grouped_diagnostics.get_dict())
return self._limit_length(result, max_answer_chars)
class ReplaceSymbolBodyTool(EditingToolWithDiagnostics):
class ReplaceSymbolBodyTool(EditingToolWithDiagnostics, LspApiMixin):
"""
Replaces the full definition of a symbol using the language server backend.
"""
@@ -571,17 +350,12 @@ class ReplaceSymbolBodyTool(EditingToolWithDiagnostics):
in the programming language, including e.g. the signature line for functions.
Depending on the language, it may or may not include a preceding docstring or other preceding annotations.
"""
with self.DiagnosticsContext(self, relative_path) as diagnostics_context:
code_editor = self.create_code_editor()
code_editor.replace_body(
name_path,
relative_file_path=relative_path,
body=body,
)
return diagnostics_context.format_result(SUCCESS_RESULT)
with self.diagnostics_context(relative_path) as diagnostics_context:
result = self._api().replace_symbol_body(name_path, relative_path, body)
return diagnostics_context.format_result(result)
class InsertAfterSymbolTool(EditingToolWithDiagnostics):
class InsertAfterSymbolTool(EditingToolWithDiagnostics, LspApiMixin):
"""
Inserts content after the end of the definition of a given symbol.
"""
@@ -601,13 +375,12 @@ class InsertAfterSymbolTool(EditingToolWithDiagnostics):
:param body: the body/content to be inserted. The inserted code shall begin with the next line after
the symbol.
"""
with self.DiagnosticsContext(self, relative_path) as diagnostics_context:
code_editor = self.create_code_editor()
code_editor.insert_after_symbol(name_path, relative_file_path=relative_path, body=body)
return diagnostics_context.format_result(SUCCESS_RESULT)
with self.diagnostics_context(relative_path) as diagnostics_context:
result = self._api().insert_after_symbol(name_path, relative_path, body)
return diagnostics_context.format_result(result)
class InsertBeforeSymbolTool(EditingToolWithDiagnostics):
class InsertBeforeSymbolTool(EditingToolWithDiagnostics, LspApiMixin):
"""
Inserts content before the beginning of the definition of a given symbol.
"""
@@ -627,13 +400,12 @@ class InsertBeforeSymbolTool(EditingToolWithDiagnostics):
:param relative_path: the relative path to the file containing the symbol
:param body: the body/content to be inserted before the line in which the referenced symbol is defined
"""
with self.DiagnosticsContext(self, relative_path) as diagnostics_context:
code_editor = self.create_code_editor()
code_editor.insert_before_symbol(name_path, relative_file_path=relative_path, body=body)
return diagnostics_context.format_result(SUCCESS_RESULT)
with self.diagnostics_context(relative_path) as diagnostics_context:
result = self._api().insert_before_symbol(name_path, relative_path, body)
return diagnostics_context.format_result(result)
class RenameSymbolTool(Tool, ToolMarkerSymbolicEdit):
class RenameSymbolTool(Tool, ToolMarkerSymbolicEdit, LspApiMixin):
"""
Renames a symbol throughout the codebase using language server refactoring capabilities.
For JB, we use a separate tool.
@@ -655,13 +427,10 @@ class RenameSymbolTool(Tool, ToolMarkerSymbolicEdit):
:param new_name: the new name for the symbol
:return: result summary indicating success or failure
"""
self.project.ls_sync_file_system_changes()
code_editor = self.create_ls_code_editor()
status_message = code_editor.rename_symbol(name_path, relative_path=relative_path, new_name=new_name)
return status_message
return self._api().rename_symbol(name_path, relative_path, new_name)
class SafeDeleteSymbol(Tool, ToolMarkerSymbolicEdit):
class SafeDeleteSymbol(Tool, ToolMarkerSymbolicEdit, LspApiMixin):
def apply(
self,
name_path_pattern: str,
@@ -674,31 +443,4 @@ class SafeDeleteSymbol(Tool, ToolMarkerSymbolicEdit):
:param name_path_pattern: name path of the symbol to delete
:param relative_path: the relative path to the file containing the symbol to delete
"""
self.project.ls_sync_file_system_changes()
ls_symbol_retriever = self.create_language_server_symbol_retriever()
symbol = ls_symbol_retriever.find_unique(name_path_pattern, substring_matching=False, within_relative_path=relative_path)
symbol_rel_path = symbol.relative_path
assert symbol_rel_path is not None, f"Symbol {name_path_pattern} has no relative path, this is likely a bug."
assert symbol_rel_path == relative_path, f"Symbol {name_path_pattern} is not in the expected relative path {relative_path}."
symbol_name_path = symbol.get_name_path()
symbol_line = symbol.line
symbol_col = symbol.column
assert symbol_line is not None and symbol_col is not None, (
f"Symbol {name_path_pattern} has no identifier position, this is likely a bug."
)
lang_server = ls_symbol_retriever.get_language_server(symbol_rel_path)
references_locations = lang_server.request_references(symbol_rel_path, symbol_line, symbol_col)
file_to_lines: dict[str, list[int]] = defaultdict(list)
if references_locations:
for ref_loc in references_locations:
ref_relative_path = ref_loc.get("relativePath")
if ref_relative_path is None:
continue
file_to_lines[ref_relative_path].append(ref_loc["range"]["start"]["line"])
if file_to_lines:
return f"Cannot delete, the symbol {symbol_name_path} is referenced in: {self._to_json(file_to_lines)}"
code_editor = self.create_ls_code_editor()
code_editor.delete_symbol(symbol_name_path, relative_file_path=symbol_rel_path)
return SUCCESS_RESULT
return self._api().safe_delete_symbol(name_path_pattern, relative_path)
+12 -43
View File
@@ -7,7 +7,7 @@ from collections.abc import Callable, Iterable
from dataclasses import dataclass
from functools import cached_property
from types import TracebackType
from typing import TYPE_CHECKING, Any, Optional, Protocol, Self, TypeVar, cast
from typing import TYPE_CHECKING, Any, Protocol, Self, TypeVar, cast
from mcp import Implementation
from mcp.server.fastmcp import Context
@@ -16,12 +16,13 @@ from sensai.util import logging
from sensai.util.string import dict_string
from serena.config.serena_config import LanguageBackend
from serena.facades.facade import SUCCESS_RESULT # noqa: F401 (re-exported for tools)
from serena.lsp.lsp_diagnostics import DiagnosticsContext
from serena.memories.memory_manager import MemoryManager
from serena.project import Project
from serena.prompt_factory import PromptFactory
from serena.util.class_decorators import singleton
from serena.util.inspection import iter_subclasses
from serena.util.ls_diagnostics import DiagnosticsDiff, EditedFilePath, PublishedDiagnosticsSnapshot
from serena.util.text_utils import TextOutputUtils
from solidlsp.ls_exceptions import SolidLSPException
@@ -32,7 +33,6 @@ if TYPE_CHECKING:
log = logging.getLogger(__name__)
T = TypeVar("T")
SUCCESS_RESULT = "OK"
class Component(ABC):
@@ -467,47 +467,16 @@ class EditingToolWithDiagnostics(Tool, ToolMarkerCanEdit):
are then resolved in subsequent edits.
"""
DIAGNOSTICS_KEY = "diagnostics[warning-or-higher]"
def diagnostics_context(self, *edited_relative_paths: str) -> DiagnosticsContext:
"""
Creates a context for use with the `with` statement, which captures the diagnostics before the edit,
such that changes can be reported
class DiagnosticsContext:
def __init__(self, tool: "EditingToolWithDiagnostics", *edited_relative_paths: str) -> None:
self._tool = tool
self._is_diagnostics_enabled = tool.ENABLE_DIAGNOSTICS and tool.agent.is_using_language_server()
self._edited_files = [EditedFilePath(path, path) for path in edited_relative_paths]
self._before_edit_diagnostics_snapshot: PublishedDiagnosticsSnapshot | None = None
self._symbol_retriever: Optional["LanguageServerSymbolRetriever"] | None = None
if self._is_diagnostics_enabled:
self._symbol_retriever = tool.create_language_server_symbol_retriever()
self._before_edit_diagnostics_snapshot = PublishedDiagnosticsSnapshot(self._edited_files, self._symbol_retriever)
def __enter__(self) -> Self:
return self
def __exit__(self, exc_type, exc_val, exc_tb):
pass
def format_result(
self,
base_result: str,
) -> str:
if not self._is_diagnostics_enabled:
return base_result
if self._before_edit_diagnostics_snapshot is None:
return base_result
assert self._symbol_retriever is not None
diagnostics_diff = DiagnosticsDiff(self._before_edit_diagnostics_snapshot, self._edited_files, self._symbol_retriever)
grouped_diagnostics = diagnostics_diff.get_grouped_diagnostics().get_dict()
if not grouped_diagnostics:
return base_result
else:
result_dict = {
"result": base_result,
EditingToolWithDiagnostics.DIAGNOSTICS_KEY: grouped_diagnostics,
}
return self._tool._to_json(result_dict)
:param edited_relative_paths: the relative paths of the files that are to be edited within the context
:return: a context which captures the diagnostics before the edit, such that changes can be reported
via `format_result`
"""
return DiagnosticsContext(self.agent, *edited_relative_paths, enable=self.ENABLE_DIAGNOSTICS)
class EditedFileContext:
+14 -6
View File
@@ -16,17 +16,25 @@ log = logging.getLogger(__name__)
def iter_subclasses(
cls: type[T], recursive: bool = True, inclusion_predicate: Callable[[type[T]], bool] = lambda t: True
) -> Iterator[type[T]]:
"""Iterate over all subclasses of a class.
"""Iterate over all subclasses of a class, yielding each subclass once (even if it is reachable via multiple base classes).
:param cls: The class whose subclasses to iterate over.
:param recursive: If True, also iterate over all subclasses of all subclasses.
:param inclusion_predicate: a predicate function to decide whether to include a subclass in the result
"""
for subclass in cls.__subclasses__():
if inclusion_predicate(subclass):
yield subclass
if recursive:
yield from iter_subclasses(subclass, recursive, inclusion_predicate)
seen: set[type] = set()
def iterate(c: type[T]) -> Iterator[type[T]]:
for subclass in c.__subclasses__():
if subclass in seen:
continue
seen.add(subclass)
if inclusion_predicate(subclass):
yield subclass
if recursive:
yield from iterate(subclass)
yield from iterate(cls)
def compute_language_server_support_composition(
+3 -3
View File
@@ -15,11 +15,11 @@ from _pytest.mark import Mark, MarkDecorator, ParameterSet
from serena.agent import SerenaAgent
from serena.config.context_mode import SerenaAgentContext
from serena.config.serena_config import ProjectConfig, RegisteredProject, SerenaConfig
from serena.lsp.lsp_diagnostics import DiagnosticsContext
from serena.project import Project
from serena.tools import (
SUCCESS_RESULT,
ActivateProjectTool,
EditingToolWithDiagnostics,
FindDeclarationTool,
FindImplementationsTool,
FindReferencingSymbolsTool,
@@ -824,9 +824,9 @@ def read_project_file(project: Project, relative_path: str) -> str:
def parse_edit_diagnostics_result(result: str) -> dict:
"""Utility function to parse the diagnostic payload returned by edit tools."""
assert EditingToolWithDiagnostics.DIAGNOSTICS_KEY in result
assert DiagnosticsContext.DIAGNOSTICS_KEY in result
d = json.loads(result)
return d[EditingToolWithDiagnostics.DIAGNOSTICS_KEY]
return d[DiagnosticsContext.DIAGNOSTICS_KEY]
@contextmanager