mirror of
https://github.com/tiennm99/serena.git
synced 2026-10-11 12:29:04 +00:00
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:
1 parent
eb53a4cee4
commit
91c2ec8dc2
10 files changed
+837
-516
No files matched your search
+26
-32
@@ -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
@@ -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
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
@@ -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)
|
||||
@@ -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:
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in new issue
Block a user