Diagnostics, implementation and definition tools (WIP)

This commit is contained in:
Michael Panchenko authored and Dominik Jain committed 2026-04-29 23:52:26 +02:00
1 parent df0f476614
commit c55e0c900e
66 files changed
+3928 -38

No files matched your search

+170
View File
@@ -0,0 +1,170 @@
"""
Demonstrates diagnostics tools and edit-tool diagnostic reporting on the Serena repo itself.
The script creates a temporary Python file inside this repository, introduces one warning,
shows file and symbol diagnostics, then introduces another warning and verifies that the
second edit reports only the newly introduced warning.
"""
import json
import shutil
import tempfile
from pathlib import Path
from pprint import pprint
from serena.agent import SerenaAgent
from serena.config.serena_config import LanguageBackend, ProjectConfig, RegisteredProject, SerenaConfig
from serena.constants import REPO_ROOT
from serena.project import Project
from serena.tools import CreateTextFileTool, GetDiagnosticsForFileTool, GetDiagnosticsForSymbolTool, ReplaceContentTool, Tool
from solidlsp.ls_config import Language
SEPARATOR = "=" * 80
REPO_PATH = Path(REPO_ROOT)
EDIT_RESULT_PREFIX = "Edit introduced new warning-or-higher diagnostics: "
def make_agent() -> SerenaAgent:
"""Create an LSP-backed Serena agent for the Serena repository."""
serena_config = SerenaConfig.from_config_file()
serena_config.web_dashboard = False
serena_config.language_backend = LanguageBackend.LSP
project = Project(
project_root=str(REPO_PATH),
project_config=ProjectConfig(
project_name="demo_serena_repo",
languages=[Language.PYTHON],
ignored_paths=[],
excluded_tools=[],
read_only=False,
ignore_all_files_in_gitignore=True,
initial_prompt="",
encoding="utf-8",
),
serena_config=serena_config,
)
serena_config.projects = [RegisteredProject.from_project_instance(project)]
return SerenaAgent(project="demo_serena_repo", serena_config=serena_config)
def print_section(title: str) -> None:
"""Print a visibly separated section header."""
print(f"\n{SEPARATOR}")
print(title)
print(SEPARATOR)
def parse_json_result(result: str) -> object:
"""Parse and pretty-print JSON tool output."""
parsed = json.loads(result)
pprint(parsed, width=200)
return parsed
def parse_edit_diagnostics_result(result: str) -> dict:
"""Extract the grouped diagnostics payload from an edit-tool result."""
assert result.startswith(EDIT_RESULT_PREFIX), result
return json.loads(result[len(EDIT_RESULT_PREFIX) :])
if __name__ == "__main__":
Tool._ENABLE_DIAGNOSTICS = True
temp_dir = Path(tempfile.mkdtemp(prefix="serena_demo_", dir=REPO_PATH))
temp_file = temp_dir / "demo_temp_diagnostics.py"
relative_path = temp_file.relative_to(REPO_PATH).as_posix()
initial_content = """def demo_existing_issue() -> int:
value = 1
return value
"""
agent = make_agent()
try:
# letting the language server finish startup
agent.execute_task(lambda: None)
create_text_file_tool = agent.get_tool(CreateTextFileTool)
replace_content_tool = agent.get_tool(ReplaceContentTool)
get_diagnostics_for_file_tool = agent.get_tool(GetDiagnosticsForFileTool)
get_diagnostics_for_symbol_tool = agent.get_tool(GetDiagnosticsForSymbolTool)
# creating a clean temporary file
print_section("Create Temporary File")
create_result = agent.execute_task(lambda: create_text_file_tool.apply(relative_path=relative_path, content=initial_content))
print(create_result)
# showing file diagnostics before introducing any warning
print_section("Initial File Diagnostics")
initial_diagnostics_result = agent.execute_task(
lambda: get_diagnostics_for_file_tool.apply(relative_path=relative_path, min_severity=2)
)
initial_diagnostics = parse_json_result(initial_diagnostics_result)
assert initial_diagnostics == {}, initial_diagnostics
# introducing the first warning
print_section("First Edit Result")
first_edit_result = agent.execute_task(
lambda: replace_content_tool.apply(
relative_path=relative_path,
needle="value = 1",
repl="value = missing_one",
mode="literal",
)
)
print(first_edit_result)
first_edit_diagnostics = parse_edit_diagnostics_result(first_edit_result)
pprint(first_edit_diagnostics, width=200)
assert "missing_one" in json.dumps(first_edit_diagnostics), first_edit_diagnostics
# showing the file- and symbol-level diagnostics after the first warning
print_section("File Diagnostics After First Edit")
diagnostics_after_first_edit_result = agent.execute_task(
lambda: get_diagnostics_for_file_tool.apply(relative_path=relative_path, min_severity=2)
)
diagnostics_after_first_edit = parse_json_result(diagnostics_after_first_edit_result)
assert "missing_one" in json.dumps(diagnostics_after_first_edit), diagnostics_after_first_edit
print_section("Symbol Diagnostics After First Edit")
symbol_diagnostics_result = agent.execute_task(
lambda: get_diagnostics_for_symbol_tool.apply(
name_path="demo_existing_issue",
reference_file=relative_path,
min_severity=2,
)
)
symbol_diagnostics = parse_json_result(symbol_diagnostics_result)
assert "missing_one" in json.dumps(symbol_diagnostics), symbol_diagnostics
# introducing a second warning while keeping the first one unchanged
print_section("Second Edit Result")
second_edit_result = agent.execute_task(
lambda: replace_content_tool.apply(
relative_path=relative_path,
needle=" return value\n",
repl=" other = missing_two\n return value + other\n",
mode="literal",
)
)
print(second_edit_result)
second_edit_diagnostics = parse_edit_diagnostics_result(second_edit_result)
pprint(second_edit_diagnostics, width=200)
second_edit_json = json.dumps(second_edit_diagnostics)
assert "missing_two" in second_edit_json, second_edit_diagnostics
assert "missing_one" not in second_edit_json, second_edit_diagnostics
print("\nVerified: the second edit result reports only the newly introduced warning.")
# showing the complete file diagnostics after both warnings exist
print_section("File Diagnostics After Second Edit")
diagnostics_after_second_edit_result = agent.execute_task(
lambda: get_diagnostics_for_file_tool.apply(relative_path=relative_path, min_severity=2)
)
diagnostics_after_second_edit = parse_json_result(diagnostics_after_second_edit_result)
diagnostics_after_second_edit_json = json.dumps(diagnostics_after_second_edit)
assert "missing_one" in diagnostics_after_second_edit_json, diagnostics_after_second_edit
assert "missing_two" in diagnostics_after_second_edit_json, diagnostics_after_second_edit
finally:
agent.shutdown()
shutil.rmtree(temp_dir, ignore_errors=True)
+128
View File
@@ -0,0 +1,128 @@
"""
Demonstrates both defining-symbol tools on the Python test repository.
"""
import json
import re
from pathlib import Path
from pprint import pprint
from serena.agent import SerenaAgent
from serena.config.serena_config import LanguageBackend, ProjectConfig, RegisteredProject, SerenaConfig
from serena.constants import REPO_ROOT
from serena.project import Project
from serena.tools import FindDefiningSymbolAtLocationTool, FindDefiningSymbolTool
from solidlsp.ls_config import Language
SEPARATOR = "=" * 80
PYTHON_TEST_REPO = Path(REPO_ROOT) / "test" / "resources" / "repos" / "python" / "test_repo"
SERVICES_FILE = Path("test_repo") / "services.py"
def make_agent(project_root: Path, language: Language, project_name: str) -> SerenaAgent:
"""Create an LSP-backed Serena agent for a single explicit project."""
serena_config = SerenaConfig.from_config_file()
serena_config.web_dashboard = False
serena_config.language_backend = LanguageBackend.LSP
project = Project(
project_root=str(project_root),
project_config=ProjectConfig(
project_name=project_name,
languages=[language],
ignored_paths=[],
excluded_tools=[],
read_only=False,
ignore_all_files_in_gitignore=True,
initial_prompt="",
encoding="utf-8",
),
serena_config=serena_config,
)
serena_config.projects = [RegisteredProject.from_project_instance(project)]
return SerenaAgent(project=project_name, serena_config=serena_config)
def print_section(title: str) -> None:
"""Print a visibly separated section header."""
print(f"\n{SEPARATOR}")
print(title)
print(SEPARATOR)
def find_identifier_occurrence_position(file_path: Path, identifier: str, occurrence_index: int = 0) -> tuple[int, int]:
"""Find the 0-based position of an identifier occurrence in a file."""
pattern = re.compile(r"\b" + re.escape(identifier) + r"\b")
current_occurrence_index = 0
with file_path.open(encoding="utf-8") as f:
for line_index, line in enumerate(f):
for match in pattern.finditer(line):
if current_occurrence_index == occurrence_index:
return line_index, match.start()
current_occurrence_index += 1
raise ValueError(f"Could not find occurrence {occurrence_index} of {identifier!r} in {file_path}")
if __name__ == "__main__":
agent = make_agent(PYTHON_TEST_REPO, Language.PYTHON, "demo_python_test_repo")
try:
# letting the language server finish startup
agent.execute_task(lambda: None)
relative_path = SERVICES_FILE.as_posix()
services_abs_path = PYTHON_TEST_REPO / SERVICES_FILE
# resolving via exact location
line, column = find_identifier_occurrence_position(services_abs_path, "User", occurrence_index=1)
find_by_location_tool = agent.get_tool(FindDefiningSymbolAtLocationTool)
location_result = agent.execute_task(
lambda: find_by_location_tool.apply(
relative_path=relative_path,
line=line,
column=column,
include_info=True,
)
)
print_section("FindDefiningSymbolAtLocationTool")
location_symbol = json.loads(location_result)
pprint(location_symbol, width=200)
# resolving via regex over the full file
find_by_regex_tool = agent.get_tool(FindDefiningSymbolTool)
regex_result = agent.execute_task(
lambda: find_by_regex_tool.apply(
regex=r"from \.models import Item, (User)",
relative_path=relative_path,
include_info=True,
)
)
print_section("FindDefiningSymbolTool (File Regex)")
regex_symbol = json.loads(regex_result)
pprint(regex_symbol, width=200)
# resolving via regex restricted to one containing symbol body
contained_regex_result = agent.execute_task(
lambda: find_by_regex_tool.apply(
regex=r"=\s+(User)\(",
relative_path=relative_path,
containing_symbol_name_path="UserService/create_user",
include_info=True,
)
)
print_section("FindDefiningSymbolTool (Contained Regex)")
contained_regex_symbol = json.loads(contained_regex_result)
pprint(contained_regex_symbol, width=200)
# validating the demonstrated result
for symbol in [location_symbol, regex_symbol, contained_regex_symbol]:
assert symbol is not None, "Expected a defining symbol result"
assert symbol.get("relative_path") is not None
assert "models.py" in symbol["relative_path"], symbol
assert "User" in json.dumps(symbol), symbol
print("\nVerified definition target: User in models.py")
finally:
agent.shutdown()
+76
View File
@@ -0,0 +1,76 @@
"""
Demonstrates FindImplementationsTool on the Go test repository.
"""
import json
from pathlib import Path
from pprint import pprint
from serena.agent import SerenaAgent
from serena.config.serena_config import LanguageBackend, ProjectConfig, RegisteredProject, SerenaConfig
from serena.constants import REPO_ROOT
from serena.project import Project
from serena.tools import FindImplementationsTool
from solidlsp.ls_config import Language
SEPARATOR = "=" * 80
GO_TEST_REPO = Path(REPO_ROOT) / "test" / "resources" / "repos" / "go" / "test_repo"
def make_agent(project_root: Path, language: Language, project_name: str) -> SerenaAgent:
"""Create an LSP-backed Serena agent for a single explicit project."""
serena_config = SerenaConfig.from_config_file()
serena_config.web_dashboard = False
serena_config.language_backend = LanguageBackend.LSP
project = Project(
project_root=str(project_root),
project_config=ProjectConfig(
project_name=project_name,
languages=[language],
ignored_paths=[],
excluded_tools=[],
read_only=False,
ignore_all_files_in_gitignore=True,
initial_prompt="",
encoding="utf-8",
),
serena_config=serena_config,
)
serena_config.projects = [RegisteredProject.from_project_instance(project)]
return SerenaAgent(project=project_name, serena_config=serena_config)
def print_section(title: str) -> None:
"""Print a visibly separated section header."""
print(f"\n{SEPARATOR}")
print(title)
print(SEPARATOR)
if __name__ == "__main__":
agent = make_agent(GO_TEST_REPO, Language.GO, "demo_go_test_repo")
try:
# letting the language server finish startup
agent.execute_task(lambda: None)
# running the implementation lookup
find_implementations_tool = agent.get_tool(FindImplementationsTool)
result = agent.execute_task(
lambda: find_implementations_tool.apply(
name_path="Greeter/FormatGreeting",
relative_path="main.go",
include_info=True,
)
)
print_section("Find Implementations Result")
implementations = json.loads(result)
pprint(implementations, width=200)
# validating the demonstrated result
assert any(implementation["name_path"] == "(ConsoleGreeter).FormatGreeting" for implementation in implementations), result
print("\nVerified implementation target: (ConsoleGreeter).FormatGreeting")
finally:
agent.shutdown()
+53
View File
@@ -4,6 +4,7 @@ import os
from abc import ABC, abstractmethod
from collections.abc import Iterable, Iterator, Reversible
from contextlib import contextmanager
from dataclasses import dataclass
from typing import Generic, TypeVar, cast
from serena.jetbrains.jetbrains_plugin_client import JetBrainsPluginClient
@@ -18,11 +19,18 @@ log = logging.getLogger(__name__)
TSymbol = TypeVar("TSymbol", bound=Symbol)
@dataclass(frozen=True)
class EditedFilePath:
before_relative_path: str
after_relative_path: str
class CodeEditor(Generic[TSymbol], ABC):
def __init__(self, project: Project) -> None:
self.project_root = project.project_root
self.encoding = project.project_config.encoding
self.newline = project.line_ending.newline_str
self._last_edited_file_paths: list[EditedFilePath] = []
class EditedFile(ABC):
def __init__(self, relative_path: str) -> None:
@@ -83,6 +91,20 @@ class CodeEditor(Generic[TSymbol], ABC):
with open(abs_path, "w", encoding=self.encoding, newline=self.newline) as f:
f.write(new_contents)
def _set_last_edited_file_paths(self, edited_file_paths: Iterable[EditedFilePath]) -> None:
unique_paths: list[EditedFilePath] = []
seen_keys: set[tuple[str, str]] = set()
for edited_file_path in edited_file_paths:
key = (edited_file_path.before_relative_path, edited_file_path.after_relative_path)
if key in seen_keys:
continue
seen_keys.add(key)
unique_paths.append(edited_file_path)
self._last_edited_file_paths = unique_paths
def get_last_edited_file_paths(self) -> list[EditedFilePath]:
return list(self._last_edited_file_paths)
@abstractmethod
def _find_unique_symbol(self, name_path: str, relative_file_path: str) -> TSymbol:
"""
@@ -114,6 +136,8 @@ class CodeEditor(Generic[TSymbol], ABC):
edited_file.delete_text_between_positions(start_pos, end_pos)
edited_file.insert_text_at_position(start_pos, body)
self._set_last_edited_file_paths([EditedFilePath(relative_file_path, relative_file_path)])
@staticmethod
def _count_leading_newlines(text: Iterable) -> int:
cnt = 0
@@ -171,6 +195,8 @@ class CodeEditor(Generic[TSymbol], ABC):
with self.edited_file_context(relative_file_path) as edited_file:
edited_file.insert_text_at_position(PositionInFile(line, col), body)
self._set_last_edited_file_paths([EditedFilePath(relative_file_path, relative_file_path)])
def insert_before_symbol(self, name_path: str, relative_file_path: str, body: str) -> None:
"""
Inserts content before the symbol with the given name in the given file.
@@ -199,6 +225,8 @@ class CodeEditor(Generic[TSymbol], ABC):
with self.edited_file_context(relative_file_path) as edited_file:
edited_file.insert_text_at_position(PositionInFile(line=line, col=col), body)
self._set_last_edited_file_paths([EditedFilePath(relative_file_path, relative_file_path)])
def insert_at_line(self, relative_path: str, line: int, content: str) -> None:
"""
Inserts content at the given line in the given file.
@@ -210,6 +238,8 @@ class CodeEditor(Generic[TSymbol], ABC):
with self.edited_file_context(relative_path) as edited_file:
edited_file.insert_text_at_position(PositionInFile(line, 0), content)
self._set_last_edited_file_paths([EditedFilePath(relative_path, relative_path)])
def delete_lines(self, relative_path: str, start_line: int, end_line: int) -> None:
"""
Deletes lines in the given file.
@@ -226,6 +256,8 @@ class CodeEditor(Generic[TSymbol], ABC):
end_pos = PositionInFile(line=end_line_for_delete, col=end_col)
edited_file.delete_text_between_positions(start_pos, end_pos)
self._set_last_edited_file_paths([EditedFilePath(relative_path, relative_path)])
def delete_symbol(self, name_path: str, relative_file_path: str) -> None:
"""
Deletes the symbol with the given name in the given file.
@@ -236,6 +268,8 @@ class CodeEditor(Generic[TSymbol], ABC):
with self.edited_file_context(relative_file_path) as edited_file:
edited_file.delete_text_between_positions(start_pos, end_pos)
self._set_last_edited_file_paths([EditedFilePath(relative_file_path, relative_file_path)])
@abstractmethod
def rename_symbol(self, name_path: str, relative_path: str, new_name: str) -> str:
pass
@@ -292,6 +326,10 @@ class LanguageServerCodeEditor(CodeEditor[LanguageServerSymbol]):
def apply(self) -> None:
pass
@abstractmethod
def get_edited_file_paths(self) -> list[EditedFilePath]:
pass
class EditOperationFileTextEdits(EditOperation):
def __init__(self, code_editor: "LanguageServerCodeEditor", file_uri: str, text_edits: list[ls_types.TextEdit]):
self._code_editor = code_editor
@@ -303,6 +341,9 @@ class LanguageServerCodeEditor(CodeEditor[LanguageServerSymbol]):
edited_file = cast(LanguageServerCodeEditor.EditedFile, edited_file)
edited_file.apply_text_edits(self._text_edits)
def get_edited_file_paths(self) -> list[EditedFilePath]:
return [EditedFilePath(self._relative_path, self._relative_path)]
class EditOperationRenameFile(EditOperation):
def __init__(self, code_editor: "LanguageServerCodeEditor", old_uri: str, new_uri: str):
self._code_editor = code_editor
@@ -314,6 +355,9 @@ class LanguageServerCodeEditor(CodeEditor[LanguageServerSymbol]):
new_abs_path = os.path.join(self._code_editor.project_root, self._new_relative_path)
os.rename(old_abs_path, new_abs_path)
def get_edited_file_paths(self) -> list[EditedFilePath]:
return [EditedFilePath(self._old_relative_path, self._new_relative_path)]
def _workspace_edit_to_edit_operations(self, workspace_edit: ls_types.WorkspaceEdit) -> list["LanguageServerCodeEditor.EditOperation"]:
operations: list[LanguageServerCodeEditor.EditOperation] = []
@@ -343,6 +387,15 @@ class LanguageServerCodeEditor(CodeEditor[LanguageServerSymbol]):
:return: number of edit operations applied
"""
operations = self._workspace_edit_to_edit_operations(workspace_edit)
# recording the affected files
edited_file_paths: list[EditedFilePath] = []
for operation in operations:
edited_file_paths.extend(operation.get_edited_file_paths())
self._set_last_edited_file_paths(edited_file_paths)
# applying the edit operations
for operation in operations:
operation.apply()
return len(operations)
+260 -1
View File
@@ -3,6 +3,7 @@ import json
import logging
import os
from abc import ABC, abstractmethod
from collections import OrderedDict
from collections.abc import Callable, Iterable, Iterator, Sequence
from dataclasses import asdict, dataclass
from time import perf_counter
@@ -11,7 +12,7 @@ from typing import Any, Generic, Literal, NotRequired, Self, TypedDict, TypeVar
from sensai.util.string import ToStringMixin
import serena.jetbrains.jetbrains_types as jb
from solidlsp import SolidLanguageServer
from solidlsp import SolidLanguageServer, ls_types
from solidlsp.ls import LSPFileBuffer
from solidlsp.ls import ReferenceInSymbol as LSPReferenceInSymbol
from solidlsp.ls_types import Position, SymbolKind, UnifiedSymbolInformation
@@ -869,6 +870,264 @@ class LanguageServerSymbolRetriever:
return [ReferenceInLanguageServerSymbol.from_lsp_reference(r) for r in references]
def find_implementing_symbols(
self,
name_path: str,
relative_file_path: str,
include_body: bool = False,
include_kinds: Sequence[SymbolKind] | None = None,
exclude_kinds: Sequence[SymbolKind] | None = None,
) -> list[LanguageServerSymbol]:
"""
Find all symbols that implement the specified symbol, which is assumed to be unique.
:param name_path: the name path of the symbol to find implementations for. While this can be a matching pattern,
it should usually be the full path to ensure uniqueness.
:param relative_file_path: the relative path of the file in which the implemented symbol is defined.
:param include_body: whether to include the body of all symbols in the result.
:param include_kinds: which kinds of symbols to include in the result.
:param exclude_kinds: which kinds of symbols to exclude from the result.
"""
symbol = self.find_unique(name_path, substring_matching=False, within_relative_path=relative_file_path)
return self.find_implementing_symbols_by_location(
symbol.location,
include_body=include_body,
include_kinds=include_kinds,
exclude_kinds=exclude_kinds,
)
def find_implementing_symbols_by_location(
self,
symbol_location: LanguageServerSymbolLocation,
include_body: bool = False,
include_kinds: Sequence[SymbolKind] | None = None,
exclude_kinds: Sequence[SymbolKind] | None = None,
) -> list[LanguageServerSymbol]:
"""
Find all symbols that implement the symbol at the given location.
:param symbol_location: the location of the symbol for which to find implementations.
Does not need to include an end_line, as it is unused in the search.
:param include_body: whether to include the body of all symbols in the result.
:param include_kinds: an optional sequence of ints representing the LSP symbol kind.
If provided, only symbols of the given kinds will be included in the result.
:param exclude_kinds: If provided, symbols of the given kinds will be excluded from the result.
Takes precedence over include_kinds.
:return: a list of symbols that implement the given symbol
"""
if not symbol_location.has_position_in_file():
raise ValueError("Symbol location does not contain a valid position in a file")
assert symbol_location.relative_path is not None
assert symbol_location.line is not None
assert symbol_location.column is not None
lang_server = self.get_language_server(symbol_location.relative_path)
implementing_symbols = lang_server.request_implementing_symbols(
relative_file_path=symbol_location.relative_path,
line=symbol_location.line,
column=symbol_location.column,
include_body=include_body,
)
if include_kinds is not None:
implementing_symbols = [s for s in implementing_symbols if s["kind"] in include_kinds]
if exclude_kinds is not None:
implementing_symbols = [s for s in implementing_symbols if s["kind"] not in exclude_kinds]
return [LanguageServerSymbol(s) for s in implementing_symbols]
def find_defining_symbol(
self,
relative_file_path: str,
line: int,
column: int,
include_body: bool = False,
) -> LanguageServerSymbol | None:
"""
Find the symbol that defines the symbol at the given file position.
:param relative_file_path: the relative path to the file in which the symbol usage occurs.
:param line: the 0-based line number of the symbol usage.
:param column: the 0-based column number of the symbol usage.
:param include_body: whether to include the body of the defining symbol in the result.
:return: the defining symbol, or None if no definition could be resolved.
"""
lang_server = self.get_language_server(relative_file_path)
defining_symbol = lang_server.request_defining_symbol(
relative_file_path=relative_file_path,
line=line,
column=column,
include_body=include_body,
)
if defining_symbol is None:
return None
return LanguageServerSymbol(defining_symbol)
def get_file_diagnostics(
self,
relative_file_path: str,
start_line: int = 0,
end_line: int = -1,
min_severity: int = 4,
) -> list[ls_types.Diagnostic]:
"""
Get diagnostics for a file, optionally restricted to a line range and minimum severity.
:param relative_file_path: the relative path to the file.
:param start_line: the first 0-based line to include.
:param end_line: the last 0-based line to include. `-1` means until end of file.
:param min_severity: minimum LSP severity to include, where 1=Error, 2=Warning, 3=Information, 4=Hint.
:return: the diagnostics matching the requested constraints.
"""
lang_server = self.get_language_server(relative_file_path)
return lang_server.request_text_document_diagnostics(
relative_file_path=relative_file_path,
start_line=start_line,
end_line=end_line,
min_severity=min_severity,
)
@staticmethod
def _symbol_identity(symbol: LanguageServerSymbol) -> tuple[str | None, int | None, int | None, str]:
return (symbol.relative_path, symbol.line, symbol.column, symbol.get_name_path())
@staticmethod
def _normalize_symbol_for_diagnostics(symbol: LanguageServerSymbol) -> LanguageServerSymbol:
current_symbol = symbol
while current_symbol.is_low_level():
parent_symbol = current_symbol.get_parent()
if parent_symbol is None:
break
current_symbol = parent_symbol
return current_symbol
def find_diagnostic_owner_symbol(self, relative_file_path: str, line: int, column: int) -> LanguageServerSymbol | None:
"""
Find the symbol that should own a diagnostic at the given position.
This prefers the structural container of the diagnostic over low-level symbols such as
local variables, because diagnostics are typically more meaningful when grouped by the
surrounding function, method, class, or analogous construct.
:param relative_file_path: the relative path to the file containing the diagnostic.
:param line: the 0-based line of the diagnostic.
:param column: the 0-based column of the diagnostic.
:return: the owning symbol, or None if no symbol could be resolved.
"""
lang_server = self.get_language_server(relative_file_path)
symbol_dict = lang_server.request_symbol_at_location(
relative_file_path=relative_file_path,
line=line,
column=column,
)
if symbol_dict is None:
return None
return self._normalize_symbol_for_diagnostics(LanguageServerSymbol(symbol_dict))
def _get_diagnostics_for_symbol(self, symbol: LanguageServerSymbol, min_severity: int) -> list[ls_types.Diagnostic]:
relative_path = symbol.relative_path
if relative_path is None:
return []
start_line, end_line = symbol.get_body_line_numbers()
if start_line is None:
if symbol.line is None:
return []
start_line = symbol.line
if end_line is None:
end_line = start_line
return self.get_file_diagnostics(
relative_file_path=relative_path,
start_line=start_line,
end_line=end_line,
min_severity=min_severity,
)
def get_symbol_diagnostics(
self,
name_path: str,
reference_file: str | None = None,
check_symbol_references: bool = False,
min_severity: int = 4,
) -> dict[LanguageServerSymbol, list[ls_types.Diagnostic]]:
"""
Get diagnostics for the specified symbol and, optionally, for all symbols that reference it.
:param name_path: the name path of the symbol to find. It should usually be unique.
:param reference_file: optional file path used to disambiguate the symbol search.
:param check_symbol_references: whether to additionally collect diagnostics for referencing symbols.
:param min_severity: minimum LSP severity to include, where 1=Error, 2=Warning, 3=Information, 4=Hint.
:return: a mapping from symbols to the diagnostics that overlap their body ranges.
"""
symbol = self.find_unique(name_path, substring_matching=False, within_relative_path=reference_file or None)
return self.get_symbol_diagnostics_by_location(
symbol.location,
check_symbol_references=check_symbol_references,
min_severity=min_severity,
)
def get_symbol_diagnostics_by_location(
self,
symbol_location: LanguageServerSymbolLocation,
check_symbol_references: bool = False,
min_severity: int = 4,
) -> dict[LanguageServerSymbol, list[ls_types.Diagnostic]]:
"""
Get diagnostics for the symbol at the given location and, optionally, for all referencing symbols.
:param symbol_location: location of the symbol to inspect.
:param check_symbol_references: whether to additionally collect diagnostics for referencing symbols.
:param min_severity: minimum LSP severity to include, where 1=Error, 2=Warning, 3=Information, 4=Hint.
:return: an ordered mapping from symbols to the diagnostics that overlap their body ranges.
"""
if not symbol_location.has_position_in_file():
raise ValueError("Symbol location does not contain a valid position in a file")
symbol = self.find_by_location(symbol_location)
if symbol is None:
assert symbol_location.relative_path is not None
assert symbol_location.line is not None
assert symbol_location.column is not None
lang_server = self.get_language_server(symbol_location.relative_path)
symbol_dict = lang_server.request_symbol_at_location(
relative_file_path=symbol_location.relative_path,
line=symbol_location.line,
column=symbol_location.column,
)
if symbol_dict is None:
return {}
symbol = LanguageServerSymbol(symbol_dict)
symbols_to_check: "OrderedDict[tuple[str | None, int | None, int | None, str], LanguageServerSymbol]" = OrderedDict()
symbols_to_check[self._symbol_identity(symbol)] = symbol
if check_symbol_references:
reference_symbols = self.find_referencing_symbols_by_location(
symbol.location,
include_body=False,
exclude_kinds=[SymbolKind.File, SymbolKind.Module, SymbolKind.Package, SymbolKind.Namespace],
)
for reference in reference_symbols:
reference_relative_path = reference.get_relative_path()
normalized_reference_symbol = None
if reference_relative_path is not None:
normalized_reference_symbol = self.find_diagnostic_owner_symbol(
relative_file_path=reference_relative_path,
line=reference.line,
column=reference.character,
)
if normalized_reference_symbol is None:
normalized_reference_symbol = self._normalize_symbol_for_diagnostics(reference.symbol)
symbols_to_check.setdefault(self._symbol_identity(normalized_reference_symbol), normalized_reference_symbol)
result: dict[LanguageServerSymbol, list[ls_types.Diagnostic]] = {}
for current_symbol in symbols_to_check.values():
diagnostics = self._get_diagnostics_for_symbol(current_symbol, min_severity=min_severity)
if diagnostics:
result[current_symbol] = diagnostics
return result
def get_symbol_overview(self, relative_path: str) -> dict[str, list[LanguageServerSymbol]]:
"""
:param relative_path: the path of the file for which to get the symbol overview
+46 -9
View File
@@ -12,6 +12,7 @@ from fnmatch import fnmatch
from pathlib import Path
from typing import Literal
from serena.code_editor import EditedFilePath
from serena.tools import SUCCESS_RESULT, EditedFileContext, Tool, ToolMarkerCanEdit, ToolMarkerOptional
from serena.util.file_system import scan_directory
from serena.util.text_utils import ContentReplacer, search_files
@@ -61,6 +62,11 @@ class CreateTextFileTool(Tool, ToolMarkerCanEdit):
:param content: the (appropriately encoded) content to write to the file
:return: a message indicating success or failure
"""
# capturing diagnostics before the edit
edited_file_paths = [EditedFilePath(relative_path, relative_path)]
diagnostics_snapshot = self._capture_published_lsp_diagnostics_snapshot(edited_file_paths)
# validating the destination path
project_root = self.get_project_root()
abs_path = (Path(project_root) / relative_path).resolve()
will_overwrite_existing = abs_path.exists()
@@ -72,12 +78,14 @@ class CreateTextFileTool(Tool, ToolMarkerCanEdit):
f"Cannot create file outside of the project directory, got {relative_path=}"
)
# writing the file
abs_path.parent.mkdir(parents=True, exist_ok=True)
abs_path.write_text(content, encoding=self.project.project_config.encoding, newline=self.project.line_ending.newline_str)
answer = f"File created: {relative_path}."
if will_overwrite_existing:
answer += " Overwrote existing file."
return answer
return self._format_lsp_edit_result_with_new_diagnostics(answer, edited_file_paths, diagnostics_snapshot)
class ListDirTool(Tool):
@@ -161,6 +169,8 @@ class ReplaceContentTool(Tool, ToolMarkerCanEdit):
Replaces content in a file (optionally using regular expressions).
"""
_CAPTURE_DIAGNOSTICS: bool = False
def apply(
self,
relative_path: str,
@@ -212,13 +222,19 @@ class ReplaceContentTool(Tool, ToolMarkerCanEdit):
Performs the replacement, with additional options not exposed in the tool.
This function can be used internally by other tools.
"""
# capturing diagnostics before the edit
edited_file_paths = [EditedFilePath(relative_path, relative_path)]
diagnostics_snapshot = self._capture_published_lsp_diagnostics_snapshot(edited_file_paths)
# applying the content replacement
self.project.validate_relative_path(relative_path, require_not_ignored=require_not_ignored)
with EditedFileContext(relative_path, self.create_code_editor()) as context:
original_content = context.get_original_content()
replacer = ContentReplacer(mode=mode, allow_multiple_occurrences=allow_multiple_occurrences)
updated_content = replacer.replace(original_content, needle, repl)
context.set_updated_content(updated_content)
return SUCCESS_RESULT
return self._format_lsp_edit_result_with_new_diagnostics(SUCCESS_RESULT, edited_file_paths, diagnostics_snapshot)
class DeleteLinesTool(Tool, ToolMarkerCanEdit, ToolMarkerOptional):
@@ -241,9 +257,15 @@ class DeleteLinesTool(Tool, ToolMarkerCanEdit, 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
"""
# capturing diagnostics before the edit
edited_file_paths = [EditedFilePath(relative_path, relative_path)]
diagnostics_snapshot = self._capture_published_lsp_diagnostics_snapshot(edited_file_paths)
# applying the edit
code_editor = self.create_code_editor()
code_editor.delete_lines(relative_path, start_line, end_line)
return SUCCESS_RESULT
return self._format_lsp_edit_result_with_new_diagnostics(SUCCESS_RESULT, edited_file_paths, diagnostics_snapshot)
class ReplaceLinesTool(Tool, ToolMarkerCanEdit, ToolMarkerOptional):
@@ -268,13 +290,20 @@ class ReplaceLinesTool(Tool, ToolMarkerCanEdit, ToolMarkerOptional):
:param end_line: the 0-based index of the last line to be deleted
:param content: the content to insert
"""
# normalizing the replacement content
if not content.endswith("\n"):
content += "\n"
result = self.agent.get_tool(DeleteLinesTool).apply(relative_path, start_line, end_line)
if result != SUCCESS_RESULT:
return result
self.agent.get_tool(InsertAtLineTool).apply(relative_path, start_line, content)
return SUCCESS_RESULT
# capturing diagnostics before the edit
edited_file_paths = [EditedFilePath(relative_path, relative_path)]
diagnostics_snapshot = self._capture_published_lsp_diagnostics_snapshot(edited_file_paths)
# applying the replacement
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)
return self._format_lsp_edit_result_with_new_diagnostics(SUCCESS_RESULT, edited_file_paths, diagnostics_snapshot)
class InsertAtLineTool(Tool, ToolMarkerCanEdit, ToolMarkerOptional):
@@ -298,11 +327,19 @@ class InsertAtLineTool(Tool, ToolMarkerCanEdit, ToolMarkerOptional):
:param line: the 0-based index of the line to insert content at
:param content: the content to be inserted
"""
# normalizing the inserted content
if not content.endswith("\n"):
content += "\n"
# capturing diagnostics before the edit
edited_file_paths = [EditedFilePath(relative_path, relative_path)]
diagnostics_snapshot = self._capture_published_lsp_diagnostics_snapshot(edited_file_paths)
# applying the insertion
code_editor = self.create_code_editor()
code_editor.insert_at_line(relative_path, line, content)
return SUCCESS_RESULT
return self._format_lsp_edit_result_with_new_diagnostics(SUCCESS_RESULT, edited_file_paths, diagnostics_snapshot)
class SearchForPatternTool(Tool):
+468 -3
View File
@@ -4,9 +4,12 @@ Language server-related tools
import copy
import os
import re
from collections import Counter, defaultdict
from collections.abc import Sequence
from typing import Any
from serena.code_editor import EditedFilePath
from serena.symbol import LanguageServerSymbol, LanguageServerSymbolDictGrouper
from serena.tools import (
SUCCESS_RESULT,
@@ -15,7 +18,72 @@ from serena.tools import (
ToolMarkerSymbolicRead,
)
from serena.tools.tools_base import ToolMarkerOptional
from solidlsp import ls_types
from solidlsp.ls_types import SymbolKind
from solidlsp.lsp_protocol_handler.lsp_types import DiagnosticSeverity
FILE_LEVEL_DIAGNOSTIC_BUCKET = "<file>"
def _diagnostic_severity_name(severity: int | None) -> str:
if severity is None:
return "Unknown"
try:
return DiagnosticSeverity(severity).name
except ValueError:
return f"Severity_{severity}"
def _diagnostic_output_dict(diagnostic: ls_types.Diagnostic) -> dict[str, Any]:
result: dict[str, Any] = {
"message": diagnostic["message"],
"range": diagnostic["range"],
}
if "code" in diagnostic:
result["code"] = diagnostic["code"]
if "source" in diagnostic:
result["source"] = diagnostic["source"]
return result
def _add_grouped_diagnostic(
grouped_result: dict[str, dict[str, dict[str, list[dict[str, Any]]]]],
relative_path: str,
severity_name: str,
name_path: str,
diagnostic: ls_types.Diagnostic,
) -> None:
grouped_result.setdefault(relative_path, {}).setdefault(severity_name, {}).setdefault(name_path, []).append(
_diagnostic_output_dict(diagnostic)
)
def _offset_to_line_and_column(text: str, offset: int) -> tuple[int, int]:
if offset < 0 or offset > len(text):
raise ValueError(f"Offset out of range: {offset}")
prefix = text[:offset]
line = prefix.count("\n")
previous_newline_offset = prefix.rfind("\n")
column = offset if previous_newline_offset == -1 else offset - previous_newline_offset - 1
return line, column
def _line_and_column_to_offset(text: str, line: int, column: int) -> int:
if line < 0 or column < 0:
raise ValueError(f"Line and column must be non-negative, got {line=}, {column=}")
line_start_offsets = [0]
for match in re.finditer("\n", text):
line_start_offsets.append(match.end())
if line >= len(line_start_offsets):
raise ValueError(f"Line out of range: {line}")
line_start_offset = line_start_offsets[line]
line_end_offset = line_start_offsets[line + 1] - 1 if line + 1 < len(line_start_offsets) else len(text)
if column > line_end_offset - line_start_offset:
raise ValueError(f"Column out of range for line {line}: {column}")
return line_start_offset + column
class RestartLanguageServerTool(Tool, ToolMarkerOptional):
@@ -324,6 +392,385 @@ class FindReferencingSymbolsTool(Tool, ToolMarkerSymbolicRead):
return self._limit_length(result_json, max_answer_chars, shortened_result_factories=shortened_results)
class FindImplementationsTool(Tool, ToolMarkerSymbolicRead):
"""
Finds symbols that implement the given symbol using the language server backend.
"""
# noinspection PyDefaultArgument
def apply(
self,
name_path: str,
relative_path: str,
include_info: bool = False,
include_kinds: list[int] = [], # noqa: B006
exclude_kinds: list[int] = [], # noqa: B006
max_answer_chars: int = -1,
) -> str:
"""
Finds implementations of the symbol at the given `name_path`.
:param name_path: for finding the symbol to find implementations for, same logic as in the `find_symbol` tool.
: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: same as in the `find_symbol` tool.
:param exclude_kinds: same as in the `find_symbol` tool.
:param max_answer_chars: same as in the `find_symbol` tool.
:return: a list of JSON objects with the symbols implementing the requested symbol
"""
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,
)
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 FindDefiningSymbolAtLocationTool(Tool, ToolMarkerSymbolicRead, ToolMarkerOptional):
"""
Finds the symbol that defines the symbol at the given file position using the language server backend.
"""
@staticmethod
def _defining_symbol_to_result_dict(
symbol_retriever: Any,
defining_symbol: LanguageServerSymbol | None,
include_body: bool,
include_info: bool,
) -> dict[str, Any] | None:
if defining_symbol is None:
return None
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
# noinspection PyDefaultArgument
def apply(
self,
relative_path: str,
line: int,
column: int,
include_body: bool = False,
include_info: bool = False,
max_answer_chars: int = -1,
) -> str:
"""
Finds the defining symbol for the symbol at the given position.
:param relative_path: the relative path to the file containing the symbol usage.
:param line: the 0-based line number of the symbol usage.
:param column: the 0-based column number of the symbol usage.
:param include_body: whether to include the source code of the defining symbol.
:param include_info: whether to include additional info (hover-like, typically including docstring and signature)
about the defining symbol. Ignored if no defining symbol is found.
:param max_answer_chars: same as in the `find_symbol` tool.
:return: a JSON object representing the defining symbol, or `null` if no definition was found.
"""
symbol_retriever = self.create_language_server_symbol_retriever()
defining_symbol = symbol_retriever.find_defining_symbol(
relative_file_path=relative_path,
line=line,
column=column,
include_body=include_body,
)
if defining_symbol is None:
result = self._to_json(None)
return self._limit_length(result, max_answer_chars)
symbol_dict = self._defining_symbol_to_result_dict(symbol_retriever, defining_symbol, include_body, include_info)
assert symbol_dict is not None
result = self._to_json(symbol_dict)
return self._limit_length(result, max_answer_chars)
class FindDefiningSymbolTool(Tool, ToolMarkerSymbolicRead):
"""
Finds the symbol that defines a uniquely captured regex match in a file or containing symbol body.
"""
@staticmethod
def _describe_search_scope(relative_path: str, containing_symbol_name_path: str | None) -> str:
if containing_symbol_name_path:
return f"symbol body '{containing_symbol_name_path}' in file '{relative_path}'"
return f"file '{relative_path}'"
@classmethod
def _format_match_preview(cls, match: re.Match[str]) -> str:
matched_text = match.group(0).replace("\n", "\\n")
return matched_text[:120]
@classmethod
def _get_unique_captured_span(cls, match: re.Match[str], regex: str, search_scope_description: str) -> tuple[int, int] | str:
if match.re.groups == 0:
return (
f"Error: Regex '{regex}' must contain exactly one capturing group that identifies the symbol usage in "
f"{search_scope_description}."
)
matched_capture_spans = [span for span in match.regs[1:] if span != (-1, -1)]
if len(matched_capture_spans) != 1:
return (
f"Error: Regex '{regex}' must produce exactly one matched capture in {search_scope_description}, "
f"but produced {len(matched_capture_spans)} for match '{cls._format_match_preview(match)}'."
)
capture_start_offset, capture_end_offset = matched_capture_spans[0]
if capture_start_offset == capture_end_offset:
return (
f"Error: Regex '{regex}' produced an empty capture in {search_scope_description}; "
"the capture must select the referenced symbol text."
)
return capture_start_offset, capture_end_offset
def _find_unique_captured_location(
self,
relative_path: str,
regex: str,
containing_symbol_name_path: str | None,
) -> tuple[int, int] | str:
# retrieving the search region
file_content = self.project.read_file(relative_path)
symbol_retriever = self.create_language_server_symbol_retriever()
search_scope_description = self._describe_search_scope(relative_path, containing_symbol_name_path)
search_start_offset = 0
search_text = file_content
if containing_symbol_name_path:
try:
containing_symbol = symbol_retriever.find_unique(containing_symbol_name_path, within_relative_path=relative_path)
except Exception as e:
return f"Error: Could not resolve containing symbol '{containing_symbol_name_path}' in file '{relative_path}': {e}"
body_start_position = containing_symbol.get_body_start_position_or_raise()
body_end_position = containing_symbol.get_body_end_position_or_raise()
search_start_offset = _line_and_column_to_offset(file_content, body_start_position.line, body_start_position.col)
search_end_offset = _line_and_column_to_offset(file_content, body_end_position.line, body_end_position.col)
search_text = file_content[search_start_offset:search_end_offset]
# finding regex matches
try:
compiled_regex = re.compile(regex, re.MULTILINE)
except re.error as e:
return f"Error: Invalid regex '{regex}': {e}"
if compiled_regex.groups == 0:
return (
f"Error: Regex '{regex}' must contain exactly one capturing group that identifies the symbol usage in "
f"{search_scope_description}."
)
matches = list(compiled_regex.finditer(search_text))
if len(matches) != 1:
match_previews = [self._format_match_preview(match) for match in matches[:3]]
preview_suffix = f" Matches: {match_previews}" if match_previews else ""
return (
f"Error: Expected exactly one regex match for '{regex}' in {search_scope_description}, "
f"but found {len(matches)}.{preview_suffix}"
)
capture_span_or_error = self._get_unique_captured_span(matches[0], regex, search_scope_description)
if isinstance(capture_span_or_error, str):
return capture_span_or_error
capture_start_offset, _ = capture_span_or_error
absolute_capture_start_offset = search_start_offset + capture_start_offset
return _offset_to_line_and_column(file_content, absolute_capture_start_offset)
# noinspection PyDefaultArgument
def apply(
self,
regex: str,
relative_path: str,
containing_symbol_name_path: str = "",
include_body: bool = False,
include_info: bool = False,
max_answer_chars: int = -1,
) -> str:
r"""
Finds the defining symbol for a uniquely captured regex match in a file.
The regex must contain exactly one capturing group, and exactly one overall match must be found.
The capture identifies the symbol usage whose definition should be resolved.
The regex is compiled with ``re.MULTILINE`` enabled and ``re.DOTALL`` disabled. This keeps common
single-symbol matches predictable. If cross-line matching is needed, opt in explicitly with ``(?s)``
or ``[\\s\\S]*?``.
:param regex: a Python regular expression containing exactly one capturing group.
:param relative_path: the relative path to the file containing the symbol usage.
: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 source code of the defining symbol.
:param include_info: whether to include additional info (hover-like, typically including docstring and signature)
about the defining symbol. Ignored if no defining symbol is found.
:param max_answer_chars: same as in the `find_symbol` tool.
:return: a JSON object representing the defining symbol, or ``null`` if no definition was found.
If the regex does not identify a unique captured match, an informative error string is returned.
"""
captured_location_or_error = self._find_unique_captured_location(
relative_path=relative_path,
regex=regex,
containing_symbol_name_path=containing_symbol_name_path or None,
)
if isinstance(captured_location_or_error, str):
return captured_location_or_error
line, column = captured_location_or_error
symbol_retriever = self.create_language_server_symbol_retriever()
defining_symbol = symbol_retriever.find_defining_symbol(
relative_file_path=relative_path,
line=line,
column=column,
include_body=include_body,
)
symbol_dict = FindDefiningSymbolAtLocationTool._defining_symbol_to_result_dict(
symbol_retriever,
defining_symbol,
include_body,
include_info,
)
result = self._to_json(symbol_dict)
return self._limit_length(result, max_answer_chars)
class GetDiagnosticsForFileTool(Tool, ToolMarkerSymbolicRead):
"""
Gets diagnostics for a file, optionally restricted to a line range, grouped by file, severity, and containing symbol.
"""
_ENABLE_DIAGNOSTICS: bool = True
def apply(
self,
relative_path: str,
start_line: int = 0,
end_line: int = -1,
min_severity: int = 4,
max_answer_chars: int = -1,
) -> str:
"""
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: same as in the `find_symbol` tool.
:return: grouped diagnostics for the requested file.
"""
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,
)
grouped_result: dict[str, dict[str, dict[str, list[dict[str, Any]]]]] = {}
for diagnostic in diagnostics:
diag_range = diagnostic["range"]["start"]
name_path = 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()
_add_grouped_diagnostic(
grouped_result,
relative_path=relative_path,
severity_name=_diagnostic_severity_name(diagnostic.get("severity")),
name_path=name_path,
diagnostic=diagnostic,
)
result = self._to_json(grouped_result)
return self._limit_length(result, max_answer_chars)
class GetDiagnosticsForSymbolTool(Tool, ToolMarkerSymbolicRead, ToolMarkerOptional):
"""
Gets diagnostics for a symbol and, optionally, for symbols that reference it.
"""
_ENABLE_DIAGNOSTICS: bool = True
def apply(
self,
name_path: str,
reference_file: str = "",
check_symbol_references: bool = False,
min_severity: int = 4,
max_answer_chars: int = -1,
) -> str:
"""
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: same as in the `find_symbol` tool.
:return: grouped diagnostics for the requested symbol and, optionally, its referencing symbols.
"""
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,
)
grouped_result: dict[str, dict[str, dict[str, list[dict[str, Any]]]]] = {}
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:
_add_grouped_diagnostic(
grouped_result,
relative_path=relative_path,
severity_name=_diagnostic_severity_name(diagnostic.get("severity")),
name_path=symbol_name_path,
diagnostic=diagnostic,
)
result = self._to_json(grouped_result)
return self._limit_length(result, max_answer_chars)
class ReplaceSymbolBodyTool(Tool, ToolMarkerSymbolicEdit):
"""
Replaces the full definition of a symbol using the language server backend.
@@ -348,13 +795,19 @@ class ReplaceSymbolBodyTool(Tool, ToolMarkerSymbolicEdit):
in the programming language, including e.g. the signature line for functions.
IMPORTANT: The body does NOT include any preceding docstrings/comments or imports, in particular.
"""
# capturing diagnostics before the edit
edited_file_paths = [EditedFilePath(relative_path, relative_path)]
diagnostics_snapshot = self._capture_published_lsp_diagnostics_snapshot(edited_file_paths)
# applying the symbol replacement
code_editor = self.create_code_editor()
code_editor.replace_body(
name_path,
relative_file_path=relative_path,
body=body,
)
return SUCCESS_RESULT
return self._format_lsp_edit_result_with_new_diagnostics(SUCCESS_RESULT, edited_file_paths, diagnostics_snapshot)
class InsertAfterSymbolTool(Tool, ToolMarkerSymbolicEdit):
@@ -377,9 +830,15 @@ class InsertAfterSymbolTool(Tool, ToolMarkerSymbolicEdit):
:param body: the body/content to be inserted. The inserted code shall begin with the next line after
the symbol.
"""
# capturing diagnostics before the edit
edited_file_paths = [EditedFilePath(relative_path, relative_path)]
diagnostics_snapshot = self._capture_published_lsp_diagnostics_snapshot(edited_file_paths)
# applying the insertion
code_editor = self.create_code_editor()
code_editor.insert_after_symbol(name_path, relative_file_path=relative_path, body=body)
return SUCCESS_RESULT
return self._format_lsp_edit_result_with_new_diagnostics(SUCCESS_RESULT, edited_file_paths, diagnostics_snapshot)
class InsertBeforeSymbolTool(Tool, ToolMarkerSymbolicEdit):
@@ -402,9 +861,15 @@ class InsertBeforeSymbolTool(Tool, ToolMarkerSymbolicEdit):
: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
"""
# capturing diagnostics before the edit
edited_file_paths = [EditedFilePath(relative_path, relative_path)]
diagnostics_snapshot = self._capture_published_lsp_diagnostics_snapshot(edited_file_paths)
# applying the insertion
code_editor = self.create_code_editor()
code_editor.insert_before_symbol(name_path, relative_file_path=relative_path, body=body)
return SUCCESS_RESULT
return self._format_lsp_edit_result_with_new_diagnostics(SUCCESS_RESULT, edited_file_paths, diagnostics_snapshot)
class RenameSymbolTool(Tool, ToolMarkerSymbolicEdit):
+198 -2
View File
@@ -18,11 +18,13 @@ from serena.project import MemoriesManager, Project
from serena.prompt_factory import PromptFactory
from serena.util.class_decorators import singleton
from serena.util.inspection import iter_subclasses
from solidlsp import ls_types
from solidlsp.ls_exceptions import SolidLSPException
from solidlsp.lsp_protocol_handler.lsp_types import DiagnosticSeverity
if TYPE_CHECKING:
from serena.agent import SerenaAgent
from serena.code_editor import CodeEditor, LanguageServerCodeEditor
from serena.code_editor import CodeEditor, EditedFilePath, LanguageServerCodeEditor
from serena.symbol import LanguageServerSymbolRetriever
log = logging.getLogger(__name__)
@@ -30,6 +32,24 @@ T = TypeVar("T")
SUCCESS_RESULT = "OK"
@dataclass(frozen=True)
class DiagnosticIdentity:
message: str
start_line: int
start_character: int
end_line: int
end_character: int
severity: int | None
code_repr: str | None
source: str | None
@dataclass
class PublishedDiagnosticsSnapshot:
generation_by_after_path: dict[str, int]
warning_identities_by_before_path: dict[str, set[DiagnosticIdentity]]
class Component(ABC):
def __init__(self, agent: "SerenaAgent"):
self.agent = agent
@@ -145,6 +165,8 @@ class Tool(Component):
_last_tool_call_client_str: str | None = None
"""We can only get the client info from within a tool call. Each tool call will update this variable."""
_ENABLE_DIAGNOSTICS: bool = False
def __init__(self, agent: "SerenaAgent"):
super().__init__(agent)
@@ -304,6 +326,180 @@ class Tool(Component):
def is_symbolic(self) -> bool:
return issubclass(self.__class__, ToolMarkerSymbolicRead) or issubclass(self.__class__, ToolMarkerSymbolicEdit)
@staticmethod
def _diagnostic_code_repr(code: Any) -> str | None:
if code is None:
return None
try:
return json.dumps(code, sort_keys=True, ensure_ascii=False)
except TypeError:
return repr(code)
@classmethod
def _diagnostic_identity(cls, diagnostic: ls_types.Diagnostic) -> DiagnosticIdentity:
diagnostic_range = diagnostic["range"]
start = diagnostic_range["start"]
end = diagnostic_range["end"]
return DiagnosticIdentity(
message=diagnostic["message"],
start_line=start["line"],
start_character=start["character"],
end_line=end["line"],
end_character=end["character"],
severity=diagnostic.get("severity"),
code_repr=cls._diagnostic_code_repr(diagnostic.get("code")),
source=diagnostic.get("source"),
)
@staticmethod
def _diagnostic_severity_name(severity: int | None) -> str:
if severity is None:
return "Unknown"
try:
return DiagnosticSeverity(severity).name
except ValueError:
return f"Severity_{severity}"
@staticmethod
def _diagnostic_output_dict(diagnostic: ls_types.Diagnostic) -> dict[str, Any]:
result: dict[str, Any] = {
"message": diagnostic["message"],
"range": diagnostic["range"],
}
if "code" in diagnostic:
result["code"] = diagnostic["code"]
if "source" in diagnostic:
result["source"] = diagnostic["source"]
return result
@classmethod
def _add_grouped_diagnostic(
cls,
grouped_result: dict[str, dict[str, dict[str, list[dict[str, Any]]]]],
relative_path: str,
name_path: str,
diagnostic: ls_types.Diagnostic,
) -> None:
severity_name = cls._diagnostic_severity_name(diagnostic.get("severity"))
grouped_result.setdefault(relative_path, {}).setdefault(severity_name, {}).setdefault(name_path, []).append(
cls._diagnostic_output_dict(diagnostic)
)
def _capture_published_lsp_diagnostics_snapshot(
self,
edited_file_paths: Iterable["EditedFilePath"],
) -> PublishedDiagnosticsSnapshot | None:
if not self._ENABLE_DIAGNOSTICS:
return None
if self.agent.get_language_backend() != LanguageBackend.LSP:
return None
# collecting diagnostics state before the edit
symbol_retriever = self.create_language_server_symbol_retriever()
generation_by_after_path: dict[str, int] = {}
warning_identities_by_before_path: dict[str, set[DiagnosticIdentity]] = {}
for edited_file_path in edited_file_paths:
try:
language_server = symbol_retriever.get_language_server(edited_file_path.after_relative_path)
except Exception:
return None
generation_by_after_path[edited_file_path.after_relative_path] = language_server.get_published_diagnostics_generation(
edited_file_path.after_relative_path
)
cached_diagnostics = language_server.get_cached_published_text_document_diagnostics(
edited_file_path.before_relative_path,
min_severity=2,
)
if cached_diagnostics is None:
try:
cached_diagnostics = language_server.request_text_document_diagnostics(
edited_file_path.before_relative_path,
min_severity=2,
)
except Exception:
cached_diagnostics = []
warning_identities_by_before_path[edited_file_path.before_relative_path] = {
self._diagnostic_identity(diagnostic) for diagnostic in cached_diagnostics or []
}
return PublishedDiagnosticsSnapshot(
generation_by_after_path=generation_by_after_path,
warning_identities_by_before_path=warning_identities_by_before_path,
)
def _format_lsp_edit_result_with_new_diagnostics(
self,
default_result: str,
edited_file_paths: Iterable["EditedFilePath"],
before_edit_diagnostics_snapshot: PublishedDiagnosticsSnapshot | None,
) -> str:
if not self._ENABLE_DIAGNOSTICS:
return default_result
# TODO: this is weird, it works because before_edit_diagnostics_snapshot is only None when diagnostics
# are disabled, not when they are empty. A part of the general design flaws introduced by agent-generated code
if before_edit_diagnostics_snapshot is None or self.agent.get_language_backend() != LanguageBackend.LSP:
return default_result
# collecting diagnostics state after the edit
symbol_retriever = self.create_language_server_symbol_retriever()
grouped_result: dict[str, dict[str, dict[str, list[dict[str, Any]]]]] = {}
saw_diagnostics_result = False
for edited_file_path in edited_file_paths:
try:
language_server = symbol_retriever.get_language_server(edited_file_path.after_relative_path)
except Exception:
return default_result
published_diagnostics = language_server.request_published_text_document_diagnostics(
relative_file_path=edited_file_path.after_relative_path,
after_generation=before_edit_diagnostics_snapshot.generation_by_after_path.get(edited_file_path.after_relative_path, -1),
timeout=2.5,
min_severity=2,
allow_cached=True,
)
if not published_diagnostics:
try:
published_diagnostics = language_server.request_text_document_diagnostics(
edited_file_path.after_relative_path,
min_severity=2,
)
except Exception:
published_diagnostics = None
if published_diagnostics is None:
continue
saw_diagnostics_result = True
existing_warning_identities = before_edit_diagnostics_snapshot.warning_identities_by_before_path.get(
edited_file_path.before_relative_path, set()
)
new_warning_identities: set[DiagnosticIdentity] = set()
for diagnostic in published_diagnostics:
diagnostic_identity = self._diagnostic_identity(diagnostic)
if diagnostic_identity in existing_warning_identities or diagnostic_identity in new_warning_identities:
continue
new_warning_identities.add(diagnostic_identity)
diagnostic_start = diagnostic["range"]["start"]
owner_symbol = symbol_retriever.find_diagnostic_owner_symbol(
relative_file_path=edited_file_path.after_relative_path,
line=diagnostic_start["line"],
column=diagnostic_start["character"],
)
name_path = owner_symbol.get_name_path() if owner_symbol is not None else "<file>"
self._add_grouped_diagnostic(grouped_result, edited_file_path.after_relative_path, name_path, diagnostic)
# preserving the normal success result when no new diagnostics were introduced
if not saw_diagnostics_result or not grouped_result:
return default_result
return f"Edit introduced new warning-or-higher diagnostics: {self._to_json(grouped_result)}"
def apply_ex(self, log_call: bool = True, catch_exceptions: bool = True, mcp_ctx: Context | None = None, **kwargs) -> str: # type: ignore
"""
Applies the tool with logging and exception handling, using the given keyword arguments
@@ -319,7 +515,7 @@ class Tool(Component):
if client_str != self.get_last_tool_call_client_str():
log.debug(f"Updating client info: {client_info}")
self.set_last_tool_call_client_str(client_str)
except BaseException as e:
except Exception as e:
log.info(f"Failed to get client info: {e}.")
def task() -> str:
+274 -1
View File
@@ -14,7 +14,7 @@ from contextlib import contextmanager
from copy import copy
from pathlib import Path, PurePath
from time import perf_counter, sleep
from typing import Self, Union, cast
from typing import Any, Self, Union, cast
import pathspec
from sensai.util.pickle import getstate, load_pickle
@@ -511,6 +511,10 @@ class SolidLanguageServer(ABC):
self.language_id = language_id
self.open_file_buffers: dict[str, LSPFileBuffer] = {}
self.language = Language(language_id)
self._published_diagnostics: dict[str, list[ls_types.Diagnostic]] = {}
self._published_diagnostics_generation_by_uri: dict[str, int] = {}
self._published_diagnostics_generation = 0
self._published_diagnostics_condition = threading.Condition()
# initialise symbol caches
self.cache_dir = Path(self._solidlsp_settings.project_data_path) / self.CACHE_FOLDER_NAME / self.language_id
@@ -550,6 +554,7 @@ class SolidLanguageServer(ABC):
logger=logging_fn,
start_independent_lsp_process=config.start_independent_lsp_process,
)
self.server.on_any_notification(self._observe_server_notification)
# Set up the pathspec matcher for the ignored paths
# for all absolute paths in ignored_paths, convert them to relative paths
@@ -567,6 +572,274 @@ class SolidLanguageServer(ABC):
self._has_waited_for_cross_file_references = False
def _observe_server_notification(self, method: str, params: Any) -> None:
"""
Observe notifications sent by the language server.
This is used for generic cross-language bookkeeping that must work independently of
language-specific notification handlers.
"""
if method == "textDocument/publishDiagnostics":
self._store_published_diagnostics(params)
def _store_published_diagnostics(self, params: Any) -> None:
"""
Store diagnostics received through ``textDocument/publishDiagnostics``.
"""
if not isinstance(params, dict):
return
uri = params.get("uri")
diagnostics = params.get("diagnostics")
if not isinstance(uri, str) or not isinstance(diagnostics, list):
return
normalized_diagnostics: list[ls_types.Diagnostic] = []
for diagnostic in diagnostics:
if not isinstance(diagnostic, dict):
continue
if "message" not in diagnostic or "range" not in diagnostic:
continue
normalized_diagnostic: ls_types.Diagnostic = {
"uri": uri,
"message": diagnostic["message"],
"range": diagnostic["range"],
}
severity = diagnostic.get("severity")
if isinstance(severity, int):
normalized_diagnostic["severity"] = ls_types.DiagnosticSeverity(severity)
code = diagnostic.get("code")
if isinstance(code, int | str):
normalized_diagnostic["code"] = code
if "source" in diagnostic:
normalized_diagnostic["source"] = diagnostic["source"]
normalized_diagnostics.append(ls_types.Diagnostic(**normalized_diagnostic))
with self._published_diagnostics_condition:
self._published_diagnostics_generation += 1
self._published_diagnostics[uri] = normalized_diagnostics
self._published_diagnostics_generation_by_uri[uri] = self._published_diagnostics_generation
self._published_diagnostics_condition.notify_all()
def _get_published_diagnostics_generation(self, uri: str) -> int:
with self._published_diagnostics_condition:
return self._published_diagnostics_generation_by_uri.get(uri, -1)
def _wait_for_published_diagnostics(
self,
uri: str,
after_generation: int,
timeout: float,
) -> list[ls_types.Diagnostic] | None:
deadline = perf_counter() + timeout
with self._published_diagnostics_condition:
while True:
current_generation = self._published_diagnostics_generation_by_uri.get(uri, -1)
if current_generation > after_generation:
return list(self._published_diagnostics.get(uri, []))
remaining_timeout = deadline - perf_counter()
if remaining_timeout <= 0:
return None
self._published_diagnostics_condition.wait(timeout=remaining_timeout)
def _get_cached_published_diagnostics(self, uri: str) -> list[ls_types.Diagnostic] | None:
with self._published_diagnostics_condition:
diagnostics = self._published_diagnostics.get(uri)
if diagnostics is None:
return None
return list(diagnostics)
@staticmethod
def _diagnostic_matches_range(diagnostic: ls_types.Diagnostic, start_line: int, end_line: int) -> bool:
diagnostic_start_line = diagnostic["range"]["start"]["line"]
diagnostic_end_line = diagnostic["range"]["end"]["line"]
effective_end_line = end_line if end_line >= 0 else diagnostic_end_line
return diagnostic_start_line <= effective_end_line and diagnostic_end_line >= start_line
@staticmethod
def _diagnostic_matches_min_severity(diagnostic: ls_types.Diagnostic, min_severity: int) -> bool:
severity = diagnostic.get("severity")
if severity is None:
return True
return int(severity) <= min_severity
@classmethod
def _filter_diagnostics(
cls,
diagnostics: list[ls_types.Diagnostic],
start_line: int,
end_line: int,
min_severity: int,
) -> list[ls_types.Diagnostic]:
diagnostics = [d for d in diagnostics if cls._diagnostic_matches_range(d, start_line, end_line)]
diagnostics = [d for d in diagnostics if cls._diagnostic_matches_min_severity(d, min_severity)]
return diagnostics
def _validate_text_document_diagnostics_request(
self,
relative_file_path: str,
start_line: int,
end_line: int,
min_severity: int,
) -> str:
if not self.server_started:
log.error("request_text_document_diagnostics called before Language Server started")
raise SolidLSPException("Language Server not started")
if start_line < 0:
raise ValueError(f"start_line must be non-negative, got {start_line}")
if end_line != -1 and end_line < start_line:
raise ValueError(f"end_line must be -1 or >= start_line, got {end_line} < {start_line}")
if min_severity not in {1, 2, 3, 4}:
raise ValueError(f"min_severity must be one of 1, 2, 3, 4, got {min_severity}")
return pathlib.Path(str(PurePath(self.repository_root_path, relative_file_path))).as_uri()
def get_published_diagnostics_generation(self, relative_file_path: str) -> int:
"""
Get the generation number for the latest published diagnostics of a file.
:param relative_file_path: The relative path of the file.
:return: the generation number, or ``-1`` if none were published yet.
"""
uri = pathlib.Path(str(PurePath(self.repository_root_path, relative_file_path))).as_uri()
return self._get_published_diagnostics_generation(uri)
def get_cached_published_text_document_diagnostics(
self,
relative_file_path: str,
start_line: int = 0,
end_line: int = -1,
min_severity: int = 4,
) -> list[ls_types.Diagnostic] | None:
"""
Get cached diagnostics received through ``textDocument/publishDiagnostics``.
:param relative_file_path: The relative path of the file to retrieve diagnostics for.
:param start_line: the first 0-based line to include in the result.
:param end_line: the last 0-based line to include in the result. ``-1`` means no upper bound.
: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.
:return: the cached diagnostics, or ``None`` if no diagnostics were published yet.
"""
uri = self._validate_text_document_diagnostics_request(relative_file_path, start_line, end_line, min_severity)
diagnostics = self._get_cached_published_diagnostics(uri)
if diagnostics is None:
return None
return self._filter_diagnostics(diagnostics, start_line, end_line, min_severity)
def request_published_text_document_diagnostics(
self,
relative_file_path: str,
after_generation: int = -1,
timeout: float = 2.5,
start_line: int = 0,
end_line: int = -1,
min_severity: int = 4,
allow_cached: bool = True,
) -> list[ls_types.Diagnostic] | None:
"""
Wait for diagnostics received through ``textDocument/publishDiagnostics`` and return them.
:param relative_file_path: The relative path of the file to retrieve diagnostics for.
:param after_generation: only return diagnostics published after this generation. ``-1`` accepts the next publication.
:param timeout: the maximum time to wait for a newer publication.
:param start_line: the first 0-based line to include in the result.
:param end_line: the last 0-based line to include in the result. ``-1`` means no upper bound.
: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 allow_cached: whether to fall back to the current cached diagnostics if no newer publication arrives in time.
:return: the published diagnostics, or ``None`` if no diagnostics are available.
"""
uri = self._validate_text_document_diagnostics_request(relative_file_path, start_line, end_line, min_severity)
diagnostics: list[ls_types.Diagnostic] | None = None
# keeping the document open
with self.open_file(relative_file_path):
diagnostics = self._wait_for_published_diagnostics(uri=uri, after_generation=after_generation, timeout=timeout)
# falling back to cached diagnostics
if diagnostics is None and allow_cached:
diagnostics = self._get_cached_published_diagnostics(uri)
if diagnostics is None:
return None
return self._filter_diagnostics(diagnostics, start_line, end_line, min_severity)
def request_text_document_diagnostics(
self,
relative_file_path: str,
start_line: int = 0,
end_line: int = -1,
min_severity: int = 4,
) -> list[ls_types.Diagnostic]:
"""
Raise a [textDocument/diagnostic](https://microsoft.github.io/language-server-protocol/specifications/lsp/3.17/specification/#textDocument_diagnostic) request to the Language Server
to find diagnostics for the given file. Wait for the response and return the result.
:param relative_file_path: The relative path of the file to retrieve diagnostics for
:param start_line: the first 0-based line to include in the result.
:param end_line: the last 0-based line to include in the result. `-1` means no upper bound.
: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.
:return: A list of diagnostics for the file
"""
uri = self._validate_text_document_diagnostics_request(relative_file_path, start_line, end_line, min_severity)
diagnostics_before_request = self._get_published_diagnostics_generation(uri)
ret: list[ls_types.Diagnostic] | None = None
pull_diagnostics_failed = False
with self.open_file(relative_file_path):
try:
response = self.server.send.text_document_diagnostic(
{
LSPConstants.TEXT_DOCUMENT: { # type: ignore
LSPConstants.URI: uri,
}
}
)
except SolidLSPException as ex:
log.debug("Falling back to published diagnostics for %s due to pull-diagnostics error: %s", relative_file_path, ex)
response = None
pull_diagnostics_failed = True
if response is not None:
assert isinstance(response, dict), (
f"Unexpected response from Language Server (expected list, got {type(response)}): {response}"
)
ret = []
for item in response["items"]: # type: ignore
new_item: ls_types.Diagnostic = {
"uri": uri,
"severity": item["severity"],
"message": item["message"],
"range": item["range"],
"code": item["code"], # type: ignore
}
if "source" in item:
new_item["source"] = item["source"]
ret.append(ls_types.Diagnostic(**new_item))
if not ret:
published_diagnostics = self._wait_for_published_diagnostics(
uri=uri,
after_generation=diagnostics_before_request,
timeout=2.5 if pull_diagnostics_failed else 0.5,
)
if published_diagnostics is None:
published_diagnostics = self._get_cached_published_diagnostics(uri)
if published_diagnostics is not None:
ret = published_diagnostics
if ret is None:
return []
return self._filter_diagnostics(ret, start_line, end_line, min_severity)
def _create_dependency_provider(self) -> LanguageServerDependencyProvider:
"""
Creates the dependency provider for this language server.
+17
View File
@@ -150,6 +150,7 @@ class LanguageServerProcess:
self._pending_requests: dict[Any, Request] = {}
self.on_request_handlers: dict[str, Callable[[Any], Any]] = {}
self.on_notification_handlers: dict[str, Callable[[Any], None]] = {}
self._notification_observers: list[Callable[[str, Any], None]] = []
self._trace_log_fn = logger
self.tasks: dict[int, Any] = {}
self.task_counter = 0
@@ -510,6 +511,12 @@ class LanguageServerProcess:
"""
self.on_notification_handlers[method] = cb
def on_any_notification(self, cb: Callable[[str, Any], None]) -> None:
"""
Register an observer that is invoked for every notification received from the server.
"""
self._notification_observers.append(cb)
def _response_handler(self, response: StringDict) -> None:
"""
Handle the response received from the server for a request, using the id to determine the request
@@ -561,6 +568,16 @@ class LanguageServerProcess:
"""
method = response.get("method", "")
params = response.get("params")
for observer in self._notification_observers:
try:
observer(method, params)
except asyncio.CancelledError:
return
except Exception as ex:
if not self._is_shutting_down:
log.error("Error handling notification observer for method '%s': %s", method, ex, exc_info=ex)
handler = self.on_notification_handlers.get(method)
if not handler:
log.warning("Unhandled method '%s'", method)
+1 -1
View File
@@ -347,7 +347,7 @@ class Diagnostic(TypedDict):
""" The severity of the diagnostic. """
message: str
""" The diagnostic message. """
code: str
code: NotRequired[str | int]
""" The code of the diagnostic. """
source: NotRequired[str]
""" The source of the diagnostic, e.g. the name of the tool that produced it. """
+61
View File
@@ -1,6 +1,7 @@
import logging
import os
import platform
import re
import shutil as _sh
from collections.abc import Iterator
from contextlib import contextmanager
@@ -41,10 +42,13 @@ _LANGUAGE_REPO_ALIASES: dict[Language, Language] = {
Language.CPP_CCLS: Language.CPP,
Language.PHP_PHPACTOR: Language.PHP,
Language.PYTHON_JEDI: Language.PYTHON,
Language.PYTHON_TY: Language.PYTHON,
Language.RUBY_SOLARGRAPH: Language.RUBY,
Language.PYTHON_TY: Language.PYTHON,
}
PYTHON_LANGUAGE_BACKENDS = [Language.PYTHON, Language.PYTHON_TY]
def get_repo_path(language: Language) -> Path:
repo_language = _LANGUAGE_REPO_ALIASES.get(language, language)
@@ -329,3 +333,60 @@ def language_tests_enabled(language: Language) -> bool:
:return: True if tests for the language are enabled, False otherwise
"""
return language not in _disabled_languages
def language_supports_implementation(language: Language) -> bool:
return language.supports_implementation_request()
def languages_supporting_implementation(*languages: Language) -> list[Language]:
return [language for language in languages if language_supports_implementation(language)]
_VERIFIED_IMPLEMENTATION_LANGUAGES = {
Language.CSHARP,
Language.GO,
Language.JAVA,
Language.RUST,
Language.TYPESCRIPT,
}
def language_has_verified_implementation_support(language: Language) -> bool:
"""
True only for languages where the server advertises implementation support and
the repo fixtures contain a verified working go-to-implementation scenario.
"""
return language in _VERIFIED_IMPLEMENTATION_LANGUAGES and language_supports_implementation(language)
def find_identifier_position(file_path: Path, identifier: str) -> tuple[int, int] | None:
pattern = re.compile(r"\b" + re.escape(identifier) + r"\b")
with file_path.open(encoding="utf-8") as f:
for line_idx, line in enumerate(f):
match = pattern.search(line)
if match:
return line_idx, match.start()
return None
def find_identifier_occurrence_position(
file_path: Path,
identifier: str,
occurrence_index: int = 0,
column_offset: int = 0,
) -> tuple[int, int] | None:
if occurrence_index < 0:
raise ValueError("occurrence_index must be non-negative")
if column_offset < 0:
raise ValueError("column_offset must be non-negative")
pattern = re.compile(r"\b" + re.escape(identifier) + r"\b")
current_index = 0
with file_path.open(encoding="utf-8") as f:
for line_idx, line in enumerate(f):
for match in pattern.finditer(line):
if current_index == occurrence_index:
return line_idx, match.start() + column_offset
current_index += 1
return None
+230
View File
@@ -0,0 +1,230 @@
import os
from dataclasses import dataclass
from typing import cast
import pytest
from _pytest.mark import Mark, MarkDecorator
from solidlsp.ls_config import Language
from test.conftest import get_pytest_markers
@dataclass(frozen=True)
class DiagnosticCase:
language: Language
relative_path: str
primary_symbol_name_path: str
primary_symbol_identifier: str
reference_symbol_name_path: str
reference_symbol_identifier: str
primary_message_fragment: str
reference_message_fragment: str
def diagnostic_case_param(
case: DiagnosticCase,
*marks: MarkDecorator | Mark,
id: str,
):
return pytest.param(case.language, case, marks=[*get_pytest_markers(case.language), *marks], id=id)
DIAGNOSTIC_CASE_PARAMS = [
diagnostic_case_param(
DiagnosticCase(
language=Language.PYTHON,
relative_path=os.path.join("test_repo", "diagnostics_sample.py"),
primary_symbol_name_path="broken_factory",
primary_symbol_identifier="broken_factory",
reference_symbol_name_path="broken_consumer",
reference_symbol_identifier="broken_consumer",
primary_message_fragment="missing_user",
reference_message_fragment="undefined_name",
),
id="python_missing_user",
),
diagnostic_case_param(
DiagnosticCase(
language=Language.PYTHON_TY,
relative_path=os.path.join("test_repo", "diagnostics_sample.py"),
primary_symbol_name_path="broken_factory",
primary_symbol_identifier="broken_factory",
reference_symbol_name_path="broken_consumer",
reference_symbol_identifier="broken_consumer",
primary_message_fragment="missing_user",
reference_message_fragment="undefined_name",
),
id="python_ty_missing_user",
),
diagnostic_case_param(
DiagnosticCase(
language=Language.GO,
relative_path="diagnostics_sample.go",
primary_symbol_name_path="brokenFactory",
primary_symbol_identifier="brokenFactory",
reference_symbol_name_path="brokenConsumer",
reference_symbol_identifier="brokenConsumer",
primary_message_fragment="missingGreeting",
reference_message_fragment="missingConsumerValue",
),
id="go_missing_greeting",
),
diagnostic_case_param(
DiagnosticCase(
language=Language.JAVA,
relative_path=os.path.join("src", "main", "java", "test_repo", "DiagnosticsSample.java"),
primary_symbol_name_path="DiagnosticsSample/brokenFactory",
primary_symbol_identifier="brokenFactory",
reference_symbol_name_path="DiagnosticsSample/brokenConsumer",
reference_symbol_identifier="brokenConsumer",
primary_message_fragment="missingGreeting",
reference_message_fragment="missingConsumerValue",
),
id="java_missing_greeting",
),
diagnostic_case_param(
DiagnosticCase(
language=Language.KOTLIN,
relative_path=os.path.join("src", "main", "kotlin", "test_repo", "DiagnosticsSample.kt"),
primary_symbol_name_path="brokenFactory",
primary_symbol_identifier="brokenFactory",
reference_symbol_name_path="brokenConsumer",
reference_symbol_identifier="brokenConsumer",
primary_message_fragment="missingGreeting",
reference_message_fragment="missingConsumerValue",
),
id="kotlin_missing_greeting",
),
diagnostic_case_param(
DiagnosticCase(
language=Language.RUST,
relative_path=os.path.join("src", "diagnostics_sample.rs"),
primary_symbol_name_path="broken_factory",
primary_symbol_identifier="broken_factory",
reference_symbol_name_path="broken_consumer",
reference_symbol_identifier="broken_consumer",
primary_message_fragment="missing_greeting",
reference_message_fragment="missing_consumer_value",
),
pytest.mark.xfail(reason="rust-analyzer does not surface diagnostics for this fixture through Serena currently"),
id="rust_missing_greeting",
),
diagnostic_case_param(
DiagnosticCase(
language=Language.PHP,
relative_path="diagnostics_sample.php",
primary_symbol_name_path="brokenFactory",
primary_symbol_identifier="brokenFactory",
reference_symbol_name_path="brokenConsumer",
reference_symbol_identifier="brokenConsumer",
primary_message_fragment="missingGreeting",
reference_message_fragment="missingConsumerValue",
),
pytest.mark.xfail(reason="PHP LS integration does not expose document diagnostics in this environment"),
id="php_missing_greeting",
),
diagnostic_case_param(
DiagnosticCase(
language=Language.CLOJURE,
relative_path=os.path.join("src", "test_app", "diagnostics_sample.clj"),
primary_symbol_name_path="broken-factory",
primary_symbol_identifier="broken-factory",
reference_symbol_name_path="broken-consumer",
reference_symbol_identifier="broken-consumer",
primary_message_fragment="missing-greeting",
reference_message_fragment="missing-consumer-value",
),
id="clojure_missing_greeting",
),
diagnostic_case_param(
DiagnosticCase(
language=Language.CSHARP,
relative_path="DiagnosticsSample.cs",
primary_symbol_name_path="TestProject/DiagnosticsSample/BrokenFactory",
primary_symbol_identifier="BrokenFactory",
reference_symbol_name_path="TestProject/DiagnosticsSample/BrokenConsumer",
reference_symbol_identifier="BrokenConsumer",
primary_message_fragment="missingGreeting",
reference_message_fragment="missingConsumerValue",
),
id="csharp_missing_greeting",
),
diagnostic_case_param(
DiagnosticCase(
language=Language.POWERSHELL,
relative_path="diagnostics_sample.ps1",
primary_symbol_name_path="function Invoke-BrokenFactory ()",
primary_symbol_identifier="Invoke-BrokenFactory",
reference_symbol_name_path="function Invoke-BrokenConsumer ()",
reference_symbol_identifier="Invoke-BrokenConsumer",
primary_message_fragment="MissingGreeting",
reference_message_fragment="MissingConsumerValue",
),
pytest.mark.xfail(reason="PowerShell LS does not surface document diagnostics in this environment"),
id="powershell_missing_greeting",
),
diagnostic_case_param(
DiagnosticCase(
language=Language.CPP_CCLS,
relative_path="diagnostics_sample.cpp",
primary_symbol_name_path="brokenFactory",
primary_symbol_identifier="brokenFactory",
reference_symbol_name_path="brokenConsumer",
reference_symbol_identifier="brokenConsumer",
primary_message_fragment="missingGreeting",
reference_message_fragment="missingConsumerValue",
),
pytest.mark.xfail(reason="ccls does not expose document diagnostics through this integration"),
id="cpp_missing_greeting",
),
diagnostic_case_param(
DiagnosticCase(
language=Language.LEAN4,
relative_path="DiagnosticsSample.lean",
primary_symbol_name_path="brokenFactory",
primary_symbol_identifier="brokenFactory",
reference_symbol_name_path="brokenConsumer",
reference_symbol_identifier="brokenConsumer",
primary_message_fragment="missingGreeting",
reference_message_fragment="missingConsumerValue",
),
pytest.mark.xfail(reason="Lean4 LS does not reliably surface diagnostics in CI"),
id="lean_missing_greeting",
),
diagnostic_case_param(
DiagnosticCase(
language=Language.TYPESCRIPT,
relative_path="diagnostics_sample.ts",
primary_symbol_name_path="brokenFactory",
primary_symbol_identifier="brokenFactory",
reference_symbol_name_path="brokenConsumer",
reference_symbol_identifier="brokenConsumer",
primary_message_fragment="missingGreeting",
reference_message_fragment="missingConsumerValue",
),
pytest.mark.typescript,
pytest.mark.xfail(reason="TypeScript LS does not surface document diagnostics through this integration"),
id="typescript_missing_greeting",
),
diagnostic_case_param(
DiagnosticCase(
language=Language.FSHARP,
relative_path="DiagnosticsSample.fs",
primary_symbol_name_path="brokenFactory",
primary_symbol_identifier="brokenFactory",
reference_symbol_name_path="brokenConsumer",
reference_symbol_identifier="brokenConsumer",
primary_message_fragment="missingGreeting",
reference_message_fragment="missingConsumerValue",
),
pytest.mark.xfail(reason="F# LS does not expose document diagnostics through this integration"),
id="fsharp_missing_greeting",
),
]
WORKING_DIAGNOSTIC_TOOL_CASE_PARAMS = [
case_param
for case_param in DIAGNOSTIC_CASE_PARAMS
if cast(DiagnosticCase, case_param.values[1]).language
in {Language.PYTHON, Language.PYTHON_TY, Language.GO, Language.JAVA, Language.KOTLIN, Language.CSHARP}
]
@@ -0,0 +1,9 @@
(ns test-app.diagnostics-sample)
(defn broken-factory []
missing-greeting)
(defn broken-consumer []
(let [value (broken-factory)]
(println value)
missing-consumer-value))
@@ -8,5 +8,10 @@
"directory": ".",
"command": "g++ -std=c++17 -I . -c b.cpp",
"file": "b.cpp"
},
{
"directory": ".",
"command": "g++ -std=c++17 -I . -c diagnostics_sample.cpp",
"file": "diagnostics_sample.cpp"
}
]
]
@@ -0,0 +1,8 @@
int brokenFactory() {
return missingGreeting;
}
int brokenConsumer() {
int value = brokenFactory();
return value + missingConsumerValue;
}
@@ -0,0 +1,19 @@
using System;
namespace TestProject
{
public static class DiagnosticsSample
{
public static string BrokenFactory()
{
return missingGreeting;
}
public static void BrokenConsumer()
{
string value = BrokenFactory();
Console.WriteLine(value);
Console.WriteLine(missingConsumerValue);
}
}
}
@@ -1,4 +1,5 @@
using System;
using TestProject.Services;
namespace TestProject
{
@@ -7,10 +8,13 @@ namespace TestProject
static void Main(string[] args)
{
Console.WriteLine("Hello, World!");
var calculator = new Calculator();
int result = calculator.Add(5, 3);
Console.WriteLine($"5 + 3 = {result}");
IGreeter greeter = new ConsoleGreeter();
Console.WriteLine(greeter.FormatGreeting("World"));
}
}
@@ -40,4 +44,4 @@ namespace TestProject
return (double)a / b;
}
}
}
}
@@ -0,0 +1,10 @@
namespace TestProject.Services
{
public class ConsoleGreeter : IGreeter
{
public string FormatGreeting(string name)
{
return $"Hello, {name}!";
}
}
}
@@ -0,0 +1,7 @@
namespace TestProject.Services
{
public interface IGreeter
{
string FormatGreeting(string name);
}
}
@@ -0,0 +1,11 @@
module DiagnosticsSample
let brokenConsumerValue = 1
let brokenFactory () =
missingGreeting
let brokenConsumer () =
let value = brokenFactory ()
printfn "%A" value
missingConsumerValue
@@ -0,0 +1,8 @@
module Formatter
type IGreeter =
abstract member FormatGreeting: string -> string
type ConsoleGreeter() =
interface IGreeter with
member _.FormatGreeting(name: string) = sprintf "Hello, %s!" name
@@ -1,7 +1,9 @@
module Program
open Calculator
open Formatter
open Models
open Shapes
[<EntryPoint>]
let main argv =
@@ -24,6 +26,12 @@ let main argv =
let calc = CalculatorClass()
let classResult = calc.Add(20, 5)
printfn "Calculator class: 20 + 5 = %d" classResult
let greeter: IGreeter = ConsoleGreeter() :> IGreeter
printfn "%s" (greeter.FormatGreeting("World"))
let shape: Shape = Circle(2.0) :> Shape
printfn "Area: %.2f" (shape.Area())
// Test person module
let person = PersonModule.createPerson "Alice Smith" 25 (Some "alice@example.com")
@@ -34,4 +42,4 @@ let main argv =
let fact5 = factorial 5
printfn "5! = %d" fact5
0 // return success
0 // return success
@@ -0,0 +1,10 @@
module Shapes
[<AbstractClass>]
type Shape() =
abstract member Area: unit -> float
type Circle(radius: float) =
inherit Shape()
override _.Area() = System.Math.PI * radius * radius
@@ -7,8 +7,11 @@
<ItemGroup>
<Compile Include="Calculator.fs" />
<Compile Include="DiagnosticsSample.fs" />
<Compile Include="Formatter.fs" />
<Compile Include="Shapes.fs" />
<Compile Include="Models/Person.fs" />
<Compile Include="Program.fs" />
</ItemGroup>
</Project>
</Project>
@@ -0,0 +1,11 @@
package main
func brokenFactory() string {
return missingGreeting
}
func brokenConsumer() {
value := brokenFactory()
_ = value
_ = missingConsumerValue
}
+12
View File
@@ -5,6 +5,8 @@ import "fmt"
func main() {
fmt.Println("Hello, Go!")
Helper()
var greeter Greeter = ConsoleGreeter{}
fmt.Println(greeter.FormatGreeting("Go"))
}
func Helper() {
@@ -22,3 +24,13 @@ func (d *DemoStruct) Value() int {
func UsingHelper() {
Helper()
}
type Greeter interface {
FormatGreeting(name string) string
}
type ConsoleGreeter struct{}
func (ConsoleGreeter) FormatGreeting(name string) string {
return "Hello, " + name + "!"
}
@@ -0,0 +1,8 @@
package test_repo;
public class ConsoleGreeter implements Greeter {
@Override
public String formatGreeting(String name) {
return "Hello, " + name + "!";
}
}
@@ -0,0 +1,13 @@
package test_repo;
public class DiagnosticsSample {
public static String brokenFactory() {
return missingGreeting;
}
public static void brokenConsumer() {
String value = brokenFactory();
System.out.println(value);
System.out.println(missingConsumerValue);
}
}
@@ -0,0 +1,5 @@
package test_repo;
public interface Greeter {
String formatGreeting(String name);
}
@@ -6,6 +6,8 @@ public class Main {
Model model = new Model("Cascade");
System.out.println(model.getName());
acceptModel(model);
Greeter greeter = new ConsoleGreeter();
System.out.println(greeter.formatGreeting("Cascade"));
}
public static void acceptModel(Model m) {
// Do nothing, just for LSP reference
@@ -0,0 +1,7 @@
distributionBase=GRADLE_USER_HOME
distributionPath=wrapper/dists
distributionUrl=https\://services.gradle.org/distributions/gradle-9.0.0-bin.zip
networkTimeout=10000
validateDistributionUrl=true
zipStoreBase=GRADLE_USER_HOME
zipStorePath=wrapper/dists
+251
View File
@@ -0,0 +1,251 @@
#!/bin/sh
#
# Copyright © 2015 the original authors.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# https://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
#
# SPDX-License-Identifier: Apache-2.0
#
##############################################################################
#
# Gradle start up script for POSIX generated by Gradle.
#
# Important for running:
#
# (1) You need a POSIX-compliant shell to run this script. If your /bin/sh is
# noncompliant, but you have some other compliant shell such as ksh or
# bash, then to run this script, type that shell name before the whole
# command line, like:
#
# ksh Gradle
#
# Busybox and similar reduced shells will NOT work, because this script
# requires all of these POSIX shell features:
# * functions;
# * expansions «$var», «${var}», «${var:-default}», «${var+SET}»,
# «${var#prefix}», «${var%suffix}», and «$( cmd )»;
# * compound commands having a testable exit status, especially «case»;
# * various built-in commands including «command», «set», and «ulimit».
#
# Important for patching:
#
# (2) This script targets any POSIX shell, so it avoids extensions provided
# by Bash, Ksh, etc; in particular arrays are avoided.
#
# The "traditional" practice of packing multiple parameters into a
# space-separated string is a well documented source of bugs and security
# problems, so this is (mostly) avoided, by progressively accumulating
# options in "$@", and eventually passing that to Java.
#
# Where the inherited environment variables (DEFAULT_JVM_OPTS, JAVA_OPTS,
# and GRADLE_OPTS) rely on word-splitting, this is performed explicitly;
# see the in-line comments for details.
#
# There are tweaks for specific operating systems such as AIX, CygWin,
# Darwin, MinGW, and NonStop.
#
# (3) This script is generated from the Groovy template
# https://github.com/gradle/gradle/blob/HEAD/platforms/jvm/plugins-application/src/main/resources/org/gradle/api/internal/plugins/unixStartScript.txt
# within the Gradle project.
#
# You can find Gradle at https://github.com/gradle/gradle/.
#
##############################################################################
# Attempt to set APP_HOME
# Resolve links: $0 may be a link
app_path=$0
# Need this for daisy-chained symlinks.
while
APP_HOME=${app_path%"${app_path##*/}"} # leaves a trailing /; empty if no leading path
[ -h "$app_path" ]
do
ls=$( ls -ld "$app_path" )
link=${ls#*' -> '}
case $link in #(
/*) app_path=$link ;; #(
*) app_path=$APP_HOME$link ;;
esac
done
# This is normally unused
# shellcheck disable=SC2034
APP_BASE_NAME=${0##*/}
# Discard cd standard output in case $CDPATH is set (https://github.com/gradle/gradle/issues/25036)
APP_HOME=$( cd -P "${APP_HOME:-./}" > /dev/null && printf '%s\n' "$PWD" ) || exit
# Use the maximum available, or set MAX_FD != -1 to use that value.
MAX_FD=maximum
warn () {
echo "$*"
} >&2
die () {
echo
echo "$*"
echo
exit 1
} >&2
# OS specific support (must be 'true' or 'false').
cygwin=false
msys=false
darwin=false
nonstop=false
case "$( uname )" in #(
CYGWIN* ) cygwin=true ;; #(
Darwin* ) darwin=true ;; #(
MSYS* | MINGW* ) msys=true ;; #(
NONSTOP* ) nonstop=true ;;
esac
CLASSPATH="\\\"\\\""
# Determine the Java command to use to start the JVM.
if [ -n "$JAVA_HOME" ] ; then
if [ -x "$JAVA_HOME/jre/sh/java" ] ; then
# IBM's JDK on AIX uses strange locations for the executables
JAVACMD=$JAVA_HOME/jre/sh/java
else
JAVACMD=$JAVA_HOME/bin/java
fi
if [ ! -x "$JAVACMD" ] ; then
die "ERROR: JAVA_HOME is set to an invalid directory: $JAVA_HOME
Please set the JAVA_HOME variable in your environment to match the
location of your Java installation."
fi
else
JAVACMD=java
if ! command -v java >/dev/null 2>&1
then
die "ERROR: JAVA_HOME is not set and no 'java' command could be found in your PATH.
Please set the JAVA_HOME variable in your environment to match the
location of your Java installation."
fi
fi
# Increase the maximum file descriptors if we can.
if ! "$cygwin" && ! "$darwin" && ! "$nonstop" ; then
case $MAX_FD in #(
max*)
# In POSIX sh, ulimit -H is undefined. That's why the result is checked to see if it worked.
# shellcheck disable=SC2039,SC3045
MAX_FD=$( ulimit -H -n ) ||
warn "Could not query maximum file descriptor limit"
esac
case $MAX_FD in #(
'' | soft) :;; #(
*)
# In POSIX sh, ulimit -n is undefined. That's why the result is checked to see if it worked.
# shellcheck disable=SC2039,SC3045
ulimit -n "$MAX_FD" ||
warn "Could not set maximum file descriptor limit to $MAX_FD"
esac
fi
# Collect all arguments for the java command, stacking in reverse order:
# * args from the command line
# * the main class name
# * -classpath
# * -D...appname settings
# * --module-path (only if needed)
# * DEFAULT_JVM_OPTS, JAVA_OPTS, and GRADLE_OPTS environment variables.
# For Cygwin or MSYS, switch paths to Windows format before running java
if "$cygwin" || "$msys" ; then
APP_HOME=$( cygpath --path --mixed "$APP_HOME" )
CLASSPATH=$( cygpath --path --mixed "$CLASSPATH" )
JAVACMD=$( cygpath --unix "$JAVACMD" )
# Now convert the arguments - kludge to limit ourselves to /bin/sh
for arg do
if
case $arg in #(
-*) false ;; # don't mess with options #(
/?*) t=${arg#/} t=/${t%%/*} # looks like a POSIX filepath
[ -e "$t" ] ;; #(
*) false ;;
esac
then
arg=$( cygpath --path --ignore --mixed "$arg" )
fi
# Roll the args list around exactly as many times as the number of
# args, so each arg winds up back in the position where it started, but
# possibly modified.
#
# NB: a `for` loop captures its iteration list before it begins, so
# changing the positional parameters here affects neither the number of
# iterations, nor the values presented in `arg`.
shift # remove old arg
set -- "$@" "$arg" # push replacement arg
done
fi
# Add default JVM options here. You can also use JAVA_OPTS and GRADLE_OPTS to pass JVM options to this script.
DEFAULT_JVM_OPTS='"-Xmx64m" "-Xms64m"'
# Collect all arguments for the java command:
# * DEFAULT_JVM_OPTS, JAVA_OPTS, and optsEnvironmentVar are not allowed to contain shell fragments,
# and any embedded shellness will be escaped.
# * For example: A user cannot expect ${Hostname} to be expanded, as it is an environment variable and will be
# treated as '${Hostname}' itself on the command line.
set -- \
"-Dorg.gradle.appname=$APP_BASE_NAME" \
-classpath "$CLASSPATH" \
-jar "$APP_HOME/gradle/wrapper/gradle-wrapper.jar" \
"$@"
# Stop when "xargs" is not available.
if ! command -v xargs >/dev/null 2>&1
then
die "xargs is not available"
fi
# Use "xargs" to parse quoted args.
#
# With -n1 it outputs one arg per line, with the quotes and backslashes removed.
#
# In Bash we could simply go:
#
# readarray ARGS < <( xargs -n1 <<<"$var" ) &&
# set -- "${ARGS[@]}" "$@"
#
# but POSIX shell has neither arrays nor command substitution, so instead we
# post-process each arg (as a line of input to sed) to backslash-escape any
# character that might be a shell metacharacter, then use eval to reverse
# that process (while maintaining the separation between arguments), and wrap
# the whole thing up as a single "set" statement.
#
# This will of course break if any of these variables contains a newline or
# an unmatched quote.
#
eval "set -- $(
printf '%s\n' "$DEFAULT_JVM_OPTS $JAVA_OPTS $GRADLE_OPTS" |
xargs -n1 |
sed ' s~[^-[:alnum:]+,./:=@_]~\\&~g; ' |
tr '\n' ' '
)" '"$@"'
exec "$JAVACMD" "$@"
+94
View File
@@ -0,0 +1,94 @@
@rem
@rem Copyright 2015 the original author or authors.
@rem
@rem Licensed under the Apache License, Version 2.0 (the "License");
@rem you may not use this file except in compliance with the License.
@rem You may obtain a copy of the License at
@rem
@rem https://www.apache.org/licenses/LICENSE-2.0
@rem
@rem Unless required by applicable law or agreed to in writing, software
@rem distributed under the License is distributed on an "AS IS" BASIS,
@rem WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
@rem See the License for the specific language governing permissions and
@rem limitations under the License.
@rem
@rem SPDX-License-Identifier: Apache-2.0
@rem
@if "%DEBUG%"=="" @echo off
@rem ##########################################################################
@rem
@rem Gradle startup script for Windows
@rem
@rem ##########################################################################
@rem Set local scope for the variables with windows NT shell
if "%OS%"=="Windows_NT" setlocal
set DIRNAME=%~dp0
if "%DIRNAME%"=="" set DIRNAME=.
@rem This is normally unused
set APP_BASE_NAME=%~n0
set APP_HOME=%DIRNAME%
@rem Resolve any "." and ".." in APP_HOME to make it shorter.
for %%i in ("%APP_HOME%") do set APP_HOME=%%~fi
@rem Add default JVM options here. You can also use JAVA_OPTS and GRADLE_OPTS to pass JVM options to this script.
set DEFAULT_JVM_OPTS="-Xmx64m" "-Xms64m"
@rem Find java.exe
if defined JAVA_HOME goto findJavaFromJavaHome
set JAVA_EXE=java.exe
%JAVA_EXE% -version >NUL 2>&1
if %ERRORLEVEL% equ 0 goto execute
echo. 1>&2
echo ERROR: JAVA_HOME is not set and no 'java' command could be found in your PATH. 1>&2
echo. 1>&2
echo Please set the JAVA_HOME variable in your environment to match the 1>&2
echo location of your Java installation. 1>&2
goto fail
:findJavaFromJavaHome
set JAVA_HOME=%JAVA_HOME:"=%
set JAVA_EXE=%JAVA_HOME%/bin/java.exe
if exist "%JAVA_EXE%" goto execute
echo. 1>&2
echo ERROR: JAVA_HOME is set to an invalid directory: %JAVA_HOME% 1>&2
echo. 1>&2
echo Please set the JAVA_HOME variable in your environment to match the 1>&2
echo location of your Java installation. 1>&2
goto fail
:execute
@rem Setup the command line
set CLASSPATH=
@rem Execute Gradle
"%JAVA_EXE%" %DEFAULT_JVM_OPTS% %JAVA_OPTS% %GRADLE_OPTS% "-Dorg.gradle.appname=%APP_BASE_NAME%" -classpath "%CLASSPATH%" -jar "%APP_HOME%\gradle\wrapper\gradle-wrapper.jar" %*
:end
@rem End local scope for the variables with windows NT shell
if %ERRORLEVEL% equ 0 goto mainEnd
:fail
rem Set variable GRADLE_EXIT_CONSOLE if you need the _script_ return code instead of
rem the _cmd.exe /c_ return code!
set EXIT_CODE=%ERRORLEVEL%
if %EXIT_CODE% equ 0 set EXIT_CODE=1
if not ""=="%GRADLE_EXIT_CONSOLE%" exit %EXIT_CODE%
exit /b %EXIT_CODE%
:mainEnd
if "%OS%"=="Windows_NT" endlocal
:omega
@@ -0,0 +1,11 @@
package test_repo
fun brokenFactory(): String {
return missingGreeting
}
fun brokenConsumer() {
val value = brokenFactory()
println(value)
println(missingConsumerValue)
}
@@ -0,0 +1,6 @@
def brokenFactory : Nat :=
missingGreeting
def brokenConsumer : Nat :=
let value := brokenFactory
value + missingConsumerValue
+10 -1
View File
@@ -3,6 +3,7 @@
-- main.lua: Entry point for the test application
local calculator = require("src.calculator")
local animals = require("src.animals")
local utils = require("src.utils")
local function print_banner()
@@ -50,6 +51,12 @@ local function test_utils()
logger:warn("This is a warning")
end
local function test_animals()
print("\nTesting Animal Module:")
local dog = animals.Dog:new("Rex")
print(dog:speak())
end
local function interactive_calculator()
print("\nInteractive Calculator (type 'quit' to exit):")
while true do
@@ -99,11 +106,13 @@ local function main(args)
if #args == 0 then
test_calculator()
test_utils()
test_animals()
elseif args[1] == "interactive" then
interactive_calculator()
elseif args[1] == "test" then
test_calculator()
test_utils()
test_animals()
print("\nAll tests completed!")
else
print("Usage: lua main.lua [interactive|test]")
@@ -111,4 +120,4 @@ local function main(args)
end
-- Run main function
main(arg or {})
main(arg or {})
@@ -0,0 +1,37 @@
---@class Animal
local Animal = {}
Animal.__index = Animal
---@param name string
---@return Animal
function Animal:new(name)
local self = setmetatable({}, Animal)
self.name = name
return self
end
---@return string
function Animal:speak()
return self.name .. " makes a sound"
end
---@class Dog: Animal
local Dog = setmetatable({}, { __index = Animal })
Dog.__index = Dog
---@param name string
---@return Dog
function Dog:new(name)
local self = Animal.new(self, name)
return setmetatable(self, Dog)
end
---@return string
function Dog:speak()
return self.name .. " says woof"
end
return {
Animal = Animal,
Dog = Dog,
}
@@ -0,0 +1,11 @@
<?php
function brokenFactory(): string {
return $missingGreeting;
}
function brokenConsumer(): void {
$value = brokenFactory();
echo $value;
echo $missingConsumerValue;
}
@@ -0,0 +1,9 @@
function Invoke-BrokenFactory {
Invoke-MissingGreeting
}
function Invoke-BrokenConsumer {
$value = Invoke-BrokenFactory
Write-Output $value
Invoke-MissingConsumerValue
}
@@ -0,0 +1,11 @@
from .models import User
def broken_factory() -> User:
return missing_user
def broken_consumer() -> None:
created_user = broken_factory()
print(created_user)
print(undefined_name)
@@ -7,3 +7,15 @@ class Calculator
a - b
end
end
class Greeter
def format_greeting(name)
name
end
end
class ConsoleGreeter < Greeter
def format_greeting(name)
"Hello, #{name}!"
end
end
@@ -15,6 +15,8 @@ end
def helper_function(number = 42)
demo = DemoClass.new(number)
Calculator.new.add(demo.value, 10)
greeter = ConsoleGreeter.new
puts greeter.format_greeting("Ruby")
demo.print_value
end
@@ -0,0 +1,9 @@
pub fn broken_factory() -> String {
missing_greeting
}
pub fn broken_consumer() {
let value = broken_factory();
println!("{value}");
println!("{}", missing_consumer_value);
}
@@ -1,3 +1,5 @@
pub mod diagnostics_sample;
// This function returns the sum of 2 + 2
pub fn add() -> i32 {
let res = 2 + 2;
@@ -7,3 +9,15 @@ pub fn multiply() -> i32 {
2 * 3
}
pub trait Greeter {
fn format_greeting(&self, name: &str) -> String;
}
pub struct ConsoleGreeter;
impl Greeter for ConsoleGreeter {
fn format_greeting(&self, name: &str) -> String {
format!("Hello, {name}!")
}
}
@@ -1,8 +1,10 @@
use rsandbox::add;
use rsandbox::{add, ConsoleGreeter, Greeter};
fn main() {
println!("Hello, World!");
println!("Good morning!");
println!("add result: {}", add());
let greeter = ConsoleGreeter;
println!("{}", greeter.format_greeting("Rust"));
println!("inserted line");
}
@@ -1,7 +1,153 @@
# list of tool names to exclude.
# This extends the existing exclusions (e.g. from the global configuration)
#
# Below is the complete list of tools for convenience.
# To make sure you have the latest list of tools, and to view their descriptions,
# execute `uv run scripts/print_tool_overview.py`.
#
# * `activate_project`: Activates a project by name.
# * `check_onboarding_performed`: Checks whether project onboarding was already performed.
# * `create_text_file`: Creates/overwrites a file in the project directory.
# * `delete_lines`: Deletes a range of lines within a file.
# * `delete_memory`: Deletes a memory from Serena's project-specific memory store.
# * `execute_shell_command`: Executes a shell command.
# * `find_referencing_code_snippets`: Finds code snippets in which the symbol at the given location is referenced.
# * `find_referencing_symbols`: Finds symbols that reference the symbol at the given location (optionally filtered by type).
# * `find_symbol`: Performs a global (or local) search for symbols with/containing a given name/substring (optionally filtered by type).
# * `get_current_config`: Prints the current configuration of the agent, including the active and available projects, tools, contexts, and modes.
# * `get_symbols_overview`: Gets an overview of the top-level symbols defined in a given file.
# * `initial_instructions`: Gets the initial instructions for the current project.
# Should only be used in settings where the system prompt cannot be set,
# e.g. in clients you have no control over, like Claude Desktop.
# * `insert_after_symbol`: Inserts content after the end of the definition of a given symbol.
# * `insert_at_line`: Inserts content at a given line in a file.
# * `insert_before_symbol`: Inserts content before the beginning of the definition of a given symbol.
# * `list_dir`: Lists files and directories in the given directory (optionally with recursion).
# * `list_memories`: Lists memories in Serena's project-specific memory store.
# * `onboarding`: Performs onboarding (identifying the project structure and essential tasks, e.g. for testing or building).
# * `prepare_for_new_conversation`: Provides instructions for preparing for a new conversation (in order to continue with the necessary context).
# * `read_file`: Reads a file within the project directory.
# * `read_memory`: Reads the memory with the given name from Serena's project-specific memory store.
# * `remove_project`: Removes a project from the Serena configuration.
# * `replace_lines`: Replaces a range of lines within a file with new content.
# * `replace_symbol_body`: Replaces the full definition of a symbol.
# * `restart_language_server`: Restarts the language server, may be necessary when edits not through Serena happen.
# * `search_for_pattern`: Performs a search for a pattern in the project.
# * `summarize_changes`: Provides instructions for summarizing the changes made to the codebase.
# * `switch_modes`: Activates modes by providing a list of their names
# * `think_about_collected_information`: Thinking tool for pondering the completeness of collected information.
# * `think_about_task_adherence`: Thinking tool for determining whether the agent is still on track with the current task.
# * `think_about_whether_you_are_done`: Thinking tool for determining whether the task is truly completed.
# * `write_memory`: Writes a named memory (for future reference) to Serena's project-specific memory store.
excluded_tools: []
# whether to use project's .gitignore files to ignore files
ignore_all_files_in_gitignore: true
# list of additional paths to ignore in this project.
# Same syntax as gitignore, so you can use * and **.
# Note: global ignored_paths from serena_config.yml are also applied additively.
ignored_paths: []
# initial prompt for the project. It will always be given to the LLM upon activating the project
# (contrary to the memories, which are loaded on demand).
initial_prompt: ''
language: typescript
# the name by which the project can be referenced within Serena
project_name: test_repo
# whether the project is in read-only mode
# If set to true, all editing tools will be disabled and attempts to use them will result in an error
# Added on 2025-04-18
read_only: false
# list of tools to include that would otherwise be disabled (particularly optional tools that are disabled by default).
# This extends the existing inclusions (e.g. from the global configuration).
included_optional_tools: []
# fixed set of tools to use as the base tool set (if non-empty), replacing Serena's default set of tools.
# This cannot be combined with non-empty excluded_tools or included_optional_tools.
fixed_tools: []
# list of mode names to that are always to be included in the set of active modes
# The full set of modes to be activated is base_modes + default_modes.
# If the setting is undefined, the base_modes from the global configuration (serena_config.yml) apply.
# Otherwise, this setting overrides the global configuration.
# Set this to [] to disable base modes for this project.
# Set this to a list of mode names to always include the respective modes for this project.
base_modes:
# list of mode names that are to be activated by default.
# The full set of modes to be activated is base_modes + default_modes.
# If the setting is undefined, the default_modes from the global configuration (serena_config.yml) apply.
# Otherwise, this overrides the setting from the global configuration (serena_config.yml).
# This setting can, in turn, be overridden by CLI parameters (--mode).
default_modes:
# time budget (seconds) per tool call for the retrieval of additional symbol information
# such as docstrings or parameter information.
# This overrides the corresponding setting in the global configuration; see the documentation there.
# If null or missing, use the setting from the global configuration.
symbol_info_budget:
# The language backend to use for this project.
# If not set, the global setting from serena_config.yml is used.
# Valid values: LSP, JetBrains
# Note: the backend is fixed at startup. If a project with a different backend
# is activated post-init, an error will be returned.
language_backend:
# line ending convention to use when writing source files.
# Possible values: unset (use global setting), "lf", "crlf", or "native" (platform default)
# This does not affect Serena's own files (e.g. memories and configuration files), which always use native line endings.
line_ending:
# list of regex patterns which, when matched, mark a memory entry as read‑only.
# Extends the list from the global configuration, merging the two lists.
read_only_memory_patterns: []
# list of regex patterns for memories to completely ignore.
# Matching memories will not appear in list_memories or activate_project output
# and cannot be accessed via read_memory or write_memory.
# To access ignored memory files, use the read_file tool on the raw file path.
# Extends the list from the global configuration, merging the two lists.
# Example: ["_archive/.*", "_episodes/.*"]
ignored_memory_patterns: []
# advanced configuration option allowing to configure language server-specific options.
# Maps the language key to the options.
# Have a look at the docstring of the constructors of the LS implementations within solidlsp (e.g., for C# or PHP) to see which options are available.
# No documentation on options means no options are available.
ls_specific_settings: {}
# the encoding used by text files in the project
# For a list of possible encodings, see https://docs.python.org/3.11/library/codecs.html#standard-encodings
encoding: utf-8
# list of languages for which language servers are started; choose from:
# al bash clojure cpp csharp
# csharp_omnisharp dart elixir elm erlang
# fortran fsharp go groovy haskell
# java julia kotlin lua markdown
# matlab nix pascal perl php
# php_phpactor powershell python python_jedi python_ty
# r
# rego ruby ruby_solargraph rust scala
# swift terraform toml typescript typescript_vts
# vue yaml zig
# (This list may be outdated. For the current list, see values of Language enum here:
# https://github.com/oraios/serena/blob/main/src/solidlsp/ls_config.py
# For some languages, there are alternative language servers, e.g. csharp_omnisharp, ruby_solargraph.)
# Note:
# - For C, use cpp
# - For JavaScript, use typescript
# - For Free Pascal/Lazarus, use pascal
# Special requirements:
# Some languages require additional setup/installations.
# See here for details: https://oraios.github.io/serena/01-about/020_programming-languages.html#language-servers
# When using multiple languages, the first language server that supports a given file will be used for that file.
# The first language is the default language and the respective language server will be used as a fallback.
# Note that when using the JetBrains backend, language servers are not used and this list is correspondingly ignored.
languages:
- typescript
@@ -0,0 +1,9 @@
export function brokenFactory(): string {
return missingGreeting;
}
export function brokenConsumer(): void {
const value = brokenFactory();
console.log(value);
console.log(missingConsumerValue);
}
@@ -0,0 +1,9 @@
export interface Greeter {
formatGreeting(name: string): string;
}
export class ConsoleGreeter implements Greeter {
formatGreeting(name: string): string {
return `Hello, ${name}!`;
}
}
@@ -1,3 +1,5 @@
import { ConsoleGreeter, Greeter } from "./formatters";
export class DemoClass {
value: number;
constructor(value: number) {
@@ -11,6 +13,9 @@ export class DemoClass {
export function helperFunction() {
const demo = new DemoClass(42);
demo.printValue();
const greeter: Greeter = new ConsoleGreeter();
console.log(greeter.formatGreeting("World"));
}
helperFunction();
@@ -67,11 +67,16 @@
# ---
# name: test_delete_symbol[test_case1]
'''
import { ConsoleGreeter, Greeter } from "./formatters";
export function helperFunction() {
const demo = new DemoClass(42);
demo.printValue();
const greeter: Greeter = new ConsoleGreeter();
console.log(greeter.formatGreeting("World"));
}
helperFunction();
@@ -515,6 +520,8 @@
# ---
# name: test_insert_in_rel_to_symbol[test_case1-after]
'''
import { ConsoleGreeter, Greeter } from "./formatters";
export class DemoClass {
value: number;
constructor(value: number) {
@@ -532,6 +539,9 @@
export function helperFunction() {
const demo = new DemoClass(42);
demo.printValue();
const greeter: Greeter = new ConsoleGreeter();
console.log(greeter.formatGreeting("World"));
}
helperFunction();
@@ -544,6 +554,8 @@
# ---
# name: test_insert_in_rel_to_symbol[test_case1-before]
'''
import { ConsoleGreeter, Greeter } from "./formatters";
function newFunctionAfterClass(): void {
console.log("This function is after DemoClass.");
}
@@ -561,6 +573,9 @@
export function helperFunction() {
const demo = new DemoClass(42);
demo.printValue();
const greeter: Greeter = new ConsoleGreeter();
console.log(greeter.formatGreeting("World"));
}
helperFunction();
@@ -573,6 +588,8 @@
# ---
# name: test_insert_in_rel_to_symbol[test_case2-after]
'''
import { ConsoleGreeter, Greeter } from "./formatters";
export class DemoClass {
value: number;
constructor(value: number) {
@@ -586,6 +603,9 @@
export function helperFunction() {
const demo = new DemoClass(42);
demo.printValue();
const greeter: Greeter = new ConsoleGreeter();
console.log(greeter.formatGreeting("World"));
}
function newInsertedFunction(): void {
@@ -602,6 +622,8 @@
# ---
# name: test_insert_in_rel_to_symbol[test_case2-before]
'''
import { ConsoleGreeter, Greeter } from "./formatters";
export class DemoClass {
value: number;
constructor(value: number) {
@@ -619,6 +641,9 @@
export function helperFunction() {
const demo = new DemoClass(42);
demo.printValue();
const greeter: Greeter = new ConsoleGreeter();
console.log(greeter.formatGreeting("World"));
}
helperFunction();
@@ -1686,6 +1711,8 @@
# ---
# name: test_replace_body[test_case1]
'''
import { ConsoleGreeter, Greeter } from "./formatters";
export class DemoClass {
value: number;
constructor(value: number) {
@@ -1700,6 +1727,9 @@
export function helperFunction() {
const demo = new DemoClass(42);
demo.printValue();
const greeter: Greeter = new ConsoleGreeter();
console.log(greeter.formatGreeting("World"));
}
helperFunction();
+595 -1
View File
@@ -2,12 +2,15 @@ import json
import logging
import os
import re
import shutil
import time
from collections.abc import Iterator
from contextlib import contextmanager
from typing import Literal
import pytest
from test.diagnostics_cases import WORKING_DIAGNOSTIC_TOOL_CASE_PARAMS, DiagnosticCase
from _pytest.mark import Mark, MarkDecorator
from serena.agent import SerenaAgent
from serena.config.serena_config import ProjectConfig, RegisteredProject, SerenaConfig
@@ -15,8 +18,13 @@ from serena.project import Project
from serena.tools import (
SUCCESS_RESULT,
ActivateProjectTool,
FindDefiningSymbolAtLocationTool,
FindDefiningSymbolTool,
FindImplementationsTool,
FindReferencingSymbolsTool,
FindSymbolTool,
GetDiagnosticsForFileTool,
GetDiagnosticsForSymbolTool,
InitialInstructionsTool,
ReplaceContentTool,
ReplaceSymbolBodyTool,
@@ -25,9 +33,208 @@ from serena.tools import (
)
from solidlsp.ls_config import Language
from solidlsp.ls_types import SymbolKind
from test.conftest import get_repo_path, is_ci, language_tests_enabled
from test.conftest import (
find_identifier_occurrence_position,
get_repo_path,
is_ci,
language_has_verified_implementation_support,
language_tests_enabled,
)
from test.diagnostics_cases import WORKING_DIAGNOSTIC_TOOL_CASE_PARAMS, DiagnosticCase
from test.solidlsp import clojure as clj
DEFINING_SYMBOL_TOOL_TEST_CASES = [
pytest.param(
Language.PYTHON,
os.path.join("test_repo", "services.py"),
"User",
1,
1,
"User",
"models.py",
marks=pytest.mark.python,
),
pytest.param(
Language.PYTHON_TY,
os.path.join("test_repo", "services.py"),
"User",
1,
1,
"User",
"models.py",
marks=pytest.mark.python,
),
pytest.param(Language.GO, "main.go", "Helper", 0, 1, "Helper", "main.go", marks=pytest.mark.go),
pytest.param(
Language.JAVA,
os.path.join("src", "main", "java", "test_repo", "Main.java"),
"Model",
0,
1,
"Model",
"Model.java",
marks=pytest.mark.java,
),
pytest.param(
Language.KOTLIN,
os.path.join("src", "main", "kotlin", "test_repo", "Main.kt"),
"Model",
0,
1,
"Model",
"Model.kt",
marks=[pytest.mark.kotlin] + ([pytest.mark.skip(reason="Kotlin LSP JVM crashes on restart in CI")] if is_ci else []),
),
pytest.param(
Language.RUST,
os.path.join("src", "main.rs"),
"format_greeting",
0,
1,
"format_greeting",
"lib.rs",
marks=pytest.mark.rust,
),
pytest.param(Language.PHP, "index.php", "helperFunction", 0, 5, "helperFunction", "helper.php", marks=pytest.mark.php),
pytest.param(
Language.CLOJURE,
clj.UTILS_PATH,
"multiply",
0,
1,
"multiply",
clj.CORE_PATH,
marks=[
pytest.mark.clojure,
pytest.mark.skipif(not clj.is_clojure_cli_available(), reason="clojure CLI is not installed"),
],
),
pytest.param(Language.CSHARP, "Program.cs", "Add", 0, 1, "Add", "Program.cs", marks=pytest.mark.csharp),
pytest.param(
Language.POWERSHELL,
"main.ps1",
"Convert-ToUpperCase",
0,
1,
"function Convert-ToUpperCase ()",
"utils.ps1",
marks=pytest.mark.powershell,
),
pytest.param(Language.CPP_CCLS, "a.cpp", "add", 0, 1, "add", "b.cpp", marks=pytest.mark.cpp),
pytest.param(
Language.LEAN4,
"Main.lean",
"add",
0,
1,
"add",
"Helper.lean",
marks=[
pytest.mark.lean4,
pytest.mark.skipif(shutil.which("lean") is None, reason="Lean is not installed"),
],
),
pytest.param(Language.TYPESCRIPT, "index.ts", "helperFunction", 1, 1, "helperFunction", "index.ts", marks=pytest.mark.typescript),
pytest.param(
Language.FSHARP,
"Program.fs",
"add",
0,
1,
"add",
"Calculator.fs",
marks=[pytest.mark.fsharp, pytest.mark.xfail(reason="F# language server cannot reliably resolve defining symbols")],
),
]
REGEX_DEFINING_SYMBOL_TOOL_TEST_CASES = [
pytest.param(
Language.PYTHON,
os.path.join("test_repo", "services.py"),
r"from \.models import Item, (User)",
"",
"User",
"models.py",
marks=pytest.mark.python,
),
pytest.param(
Language.PYTHON,
os.path.join("test_repo", "services.py"),
r"=\s+(User)\(",
"UserService/create_user",
"User",
"models.py",
marks=pytest.mark.python,
),
pytest.param(
Language.PYTHON_TY,
os.path.join("test_repo", "services.py"),
r"=\s+(User)\(",
"UserService/create_user",
"User",
"models.py",
marks=pytest.mark.python,
),
pytest.param(
Language.GO,
"main.go",
r"var greeter (Greeter) =",
"main",
"Greeter",
"main.go",
marks=pytest.mark.go,
),
]
IMPLEMENTATION_TOOL_TEST_CASE_DATA = [
(
Language.CSHARP,
"IGreeter/FormatGreeting",
os.path.join("Services", "IGreeter.cs"),
os.path.join("Services", "ConsoleGreeter.cs"),
"FormatGreeting",
pytest.mark.csharp,
),
(
Language.GO,
"Greeter/FormatGreeting",
"main.go",
"main.go",
"(ConsoleGreeter).FormatGreeting",
pytest.mark.go,
),
(
Language.JAVA,
"Greeter/formatGreeting",
os.path.join("src", "main", "java", "test_repo", "Greeter.java"),
os.path.join("src", "main", "java", "test_repo", "ConsoleGreeter.java"),
"formatGreeting",
pytest.mark.java,
),
(
Language.RUST,
"Greeter/format_greeting",
os.path.join("src", "lib.rs"),
os.path.join("src", "lib.rs"),
"format_greeting",
pytest.mark.rust,
),
(
Language.TYPESCRIPT,
"Greeter/formatGreeting",
"formatters.ts",
"formatters.ts",
"formatGreeting",
pytest.mark.typescript,
),
]
IMPLEMENTATION_TOOL_TEST_CASES = [
pytest.param(language, symbol_name, def_file, impl_file, expected_symbol_name, marks=mark)
for language, symbol_name, def_file, impl_file, expected_symbol_name, mark in IMPLEMENTATION_TOOL_TEST_CASE_DATA
if language_has_verified_implementation_support(language)
]
@pytest.fixture
def serena_config():
@@ -37,6 +244,7 @@ def serena_config():
test_projects = []
for language in [
Language.PYTHON,
Language.PYTHON_TY,
Language.GO,
Language.JAVA,
Language.KOTLIN,
@@ -82,6 +290,13 @@ def read_project_file(project: Project, relative_path: str) -> str:
return f.read()
def parse_edit_diagnostics_result(result: str) -> dict:
"""Utility function to parse the diagnostic payload returned by edit tools."""
prefix = "Edit introduced new warning-or-higher diagnostics: "
assert result.startswith(prefix), result
return json.loads(result[len(prefix) :])
@contextmanager
def project_file_modification_context(serena_agent: SerenaAgent, relative_path: str) -> Iterator[None]:
"""Context manager to modify a project file and revert the changes after use."""
@@ -267,6 +482,174 @@ class TestSerenaAgent:
refs = json.loads(result)
assert contains_ref_with_relative_path(refs, ref_file), f"Expected to find reference to {symbol_name} in {ref_file}. refs={refs}"
def _assert_find_symbol_implementations(
self,
serena_agent: SerenaAgent,
symbol_name: str,
def_file: str,
impl_file: str,
expected_symbol_name: str,
) -> None:
agent = serena_agent
find_symbol_tool = agent.get_tool(FindSymbolTool)
result = find_symbol_tool.apply(name_path_pattern=symbol_name, relative_path=def_file)
symbols = json.loads(result)
assert symbols, f"Expected to find symbol {symbol_name} in {def_file}"
def_symbol = symbols[0]
find_impl_tool = agent.get_tool(FindImplementationsTool)
result = find_impl_tool.apply(name_path=def_symbol["name_path"], relative_path=def_symbol["relative_path"], include_info=True)
implementations = json.loads(result)
assert any(
impl_file in implementation["relative_path"]
and (
implementation.get("name") == expected_symbol_name
or implementation.get("name_path") == expected_symbol_name
or expected_symbol_name in implementation.get("info", "")
)
for implementation in implementations
), f"Expected to find implementation of {symbol_name} in {impl_file}. implementations={implementations}"
for implementation in implementations:
if implementation["kind"] in (SymbolKind.File.name, SymbolKind.Module.name):
continue
symbol_info = implementation.get("info")
assert symbol_info, f"Expected symbol info to be present for implementation: {implementation}"
def _assert_find_defining_symbol(
self,
serena_agent: SerenaAgent,
relative_path: str,
identifier: str,
occurrence_index: int,
column_offset: int,
expected_name: str,
expected_definition_file: str,
) -> None:
project_root = get_repo_path(serena_agent.get_active_lsp_languages()[0])
position = find_identifier_occurrence_position(project_root / relative_path, identifier, occurrence_index, column_offset)
assert position is not None, f"Could not find occurrence {occurrence_index} of {identifier!r} in {relative_path}"
find_defining_symbol_tool = serena_agent.get_tool(FindDefiningSymbolAtLocationTool)
result = find_defining_symbol_tool.apply(relative_path=relative_path, line=position[0], column=position[1], include_info=True)
defining_symbol = json.loads(result)
assert defining_symbol is not None, f"Expected defining symbol for {identifier!r} in {relative_path}"
assert defining_symbol.get("relative_path") is not None
assert expected_definition_file in defining_symbol["relative_path"], (
f"Expected defining symbol in {expected_definition_file!r}, got: {defining_symbol}"
)
assert (
defining_symbol.get("name") == expected_name
or defining_symbol.get("name_path") == expected_name
or expected_name in defining_symbol.get("info", "")
), f"Expected defining symbol name {expected_name!r}, got: {defining_symbol}"
if serena_agent.get_active_lsp_languages() == [Language.KOTLIN]:
return
if defining_symbol["kind"] not in (SymbolKind.File.name, SymbolKind.Module.name):
assert defining_symbol.get("info"), f"Expected defining symbol info to be present: {defining_symbol}"
def _assert_find_defining_symbol_by_regex(
self,
serena_agent: SerenaAgent,
relative_path: str,
regex: str,
containing_symbol_name_path: str,
expected_name: str,
expected_definition_file: str,
) -> None:
find_defining_symbol_tool = serena_agent.get_tool(FindDefiningSymbolTool)
result = find_defining_symbol_tool.apply(
regex=regex,
relative_path=relative_path,
containing_symbol_name_path=containing_symbol_name_path,
include_info=True,
)
defining_symbol = json.loads(result)
assert defining_symbol is not None, f"Expected defining symbol for regex {regex!r} in {relative_path}"
assert defining_symbol.get("relative_path") is not None
assert expected_definition_file in defining_symbol["relative_path"], (
f"Expected defining symbol in {expected_definition_file!r}, got: {defining_symbol}"
)
assert (
defining_symbol.get("name") == expected_name
or defining_symbol.get("name_path") == expected_name
or expected_name in defining_symbol.get("info", "")
), f"Expected defining symbol name {expected_name!r}, got: {defining_symbol}"
if serena_agent.get_active_lsp_languages() == [Language.KOTLIN]:
return
if defining_symbol["kind"] not in (SymbolKind.File.name, SymbolKind.Module.name):
assert defining_symbol.get("info"), f"Expected defining symbol info to be present: {defining_symbol}"
def _assert_diagnostics_for_file(
self,
serena_agent: SerenaAgent,
diagnostic_case: DiagnosticCase,
start_line: int = 0,
end_line: int = -1,
) -> None:
diagnostics_tool = serena_agent.get_tool(GetDiagnosticsForFileTool)
result = diagnostics_tool.apply(
relative_path=diagnostic_case.relative_path,
start_line=start_line,
end_line=end_line,
min_severity=1,
)
grouped_diagnostics = json.loads(result)
assert diagnostic_case.relative_path in grouped_diagnostics, grouped_diagnostics
severity_group = grouped_diagnostics[diagnostic_case.relative_path]
assert "Error" in severity_group, severity_group
name_path_group = severity_group["Error"]
for expected_name_path in [diagnostic_case.primary_symbol_name_path, diagnostic_case.reference_symbol_name_path]:
assert expected_name_path in name_path_group, name_path_group
diagnostic_messages = [
diagnostic["message"] for diagnostics_for_name_path in name_path_group.values() for diagnostic in diagnostics_for_name_path
]
for expected_fragment in [diagnostic_case.primary_message_fragment, diagnostic_case.reference_message_fragment]:
assert any(expected_fragment in message for message in diagnostic_messages), diagnostic_messages
def _assert_diagnostics_for_symbol(
self,
serena_agent: SerenaAgent,
diagnostic_case: DiagnosticCase,
check_symbol_references: bool,
) -> None:
diagnostics_tool = serena_agent.get_tool(GetDiagnosticsForSymbolTool)
result = diagnostics_tool.apply(
name_path=diagnostic_case.primary_symbol_name_path,
reference_file=diagnostic_case.relative_path,
check_symbol_references=check_symbol_references,
min_severity=1,
)
grouped_diagnostics = json.loads(result)
diagnostics_file = diagnostic_case.relative_path
assert diagnostics_file in grouped_diagnostics, grouped_diagnostics
severity_group = grouped_diagnostics[diagnostics_file]
assert "Error" in severity_group, severity_group
name_path_group = severity_group["Error"]
expected_name_paths = [diagnostic_case.primary_symbol_name_path]
expected_message_fragments = [diagnostic_case.primary_message_fragment]
if check_symbol_references:
expected_name_paths.append(diagnostic_case.reference_symbol_name_path)
expected_message_fragments.append(diagnostic_case.reference_message_fragment)
assert set(expected_name_paths).issubset(name_path_group.keys()), name_path_group
diagnostic_messages = [
diagnostic["message"] for diagnostics_for_name_path in name_path_group.values() for diagnostic in diagnostics_for_name_path
]
for expected_fragment in expected_message_fragments:
assert any(expected_fragment in message for message in diagnostic_messages), diagnostic_messages
@pytest.mark.parametrize(
"serena_agent,symbol_name,def_file,ref_file",
[
@@ -341,6 +724,161 @@ class TestSerenaAgent:
def test_find_symbol_references_fsharp(self, serena_agent: SerenaAgent, symbol_name: str, def_file: str, ref_file: str) -> None:
self._assert_find_symbol_references(serena_agent, symbol_name, def_file, ref_file)
@pytest.mark.parametrize(
"serena_agent,relative_path,identifier,occurrence_index,column_offset,expected_name,expected_definition_file",
DEFINING_SYMBOL_TOOL_TEST_CASES,
indirect=["serena_agent"],
)
def test_find_defining_symbol(
self,
serena_agent: SerenaAgent,
relative_path: str,
identifier: str,
occurrence_index: int,
column_offset: int,
expected_name: str,
expected_definition_file: str,
) -> None:
self._assert_find_defining_symbol(
serena_agent,
relative_path,
identifier,
occurrence_index,
column_offset,
expected_name,
expected_definition_file,
)
@pytest.mark.parametrize(
"serena_agent,relative_path,regex,containing_symbol_name_path,expected_name,expected_definition_file",
REGEX_DEFINING_SYMBOL_TOOL_TEST_CASES,
indirect=["serena_agent"],
)
def test_find_defining_symbol_by_regex(
self,
serena_agent: SerenaAgent,
relative_path: str,
regex: str,
containing_symbol_name_path: str,
expected_name: str,
expected_definition_file: str,
) -> None:
self._assert_find_defining_symbol_by_regex(
serena_agent,
relative_path,
regex,
containing_symbol_name_path,
expected_name,
expected_definition_file,
)
@pytest.mark.parametrize(
"serena_agent,relative_path,regex,containing_symbol_name_path,error_fragment",
[
pytest.param(
Language.PYTHON,
os.path.join("test_repo", "services.py"),
r"(User)",
"",
"Expected exactly one regex match",
marks=pytest.mark.python,
),
pytest.param(
Language.PYTHON,
os.path.join("test_repo", "services.py"),
r"User",
"UserService/create_user",
"must contain exactly one capturing group",
marks=pytest.mark.python,
),
],
indirect=["serena_agent"],
)
def test_find_defining_symbol_by_regex_error(
self,
serena_agent: SerenaAgent,
relative_path: str,
regex: str,
containing_symbol_name_path: str,
error_fragment: str,
) -> None:
find_defining_symbol_tool = serena_agent.get_tool(FindDefiningSymbolTool)
result = find_defining_symbol_tool.apply(
regex=regex,
relative_path=relative_path,
containing_symbol_name_path=containing_symbol_name_path,
)
assert result.startswith("Error: "), result
assert error_fragment in result, result
@pytest.mark.parametrize("serena_agent,diagnostic_case", WORKING_DIAGNOSTIC_TOOL_CASE_PARAMS, indirect=["serena_agent"])
def test_get_diagnostics_for_file(self, serena_agent: SerenaAgent, diagnostic_case: DiagnosticCase) -> None:
self._assert_diagnostics_for_file(
serena_agent,
diagnostic_case=diagnostic_case,
)
@pytest.mark.parametrize("serena_agent,diagnostic_case", WORKING_DIAGNOSTIC_TOOL_CASE_PARAMS, indirect=["serena_agent"])
def test_get_diagnostics_for_file_in_range(self, serena_agent: SerenaAgent, diagnostic_case: DiagnosticCase) -> None:
project_root = get_repo_path(diagnostic_case.language)
primary_position = find_identifier_occurrence_position(
project_root / diagnostic_case.relative_path, diagnostic_case.primary_symbol_identifier
)
reference_position = find_identifier_occurrence_position(
project_root / diagnostic_case.relative_path, diagnostic_case.reference_symbol_identifier
)
assert primary_position is not None
assert reference_position is not None
self._assert_diagnostics_for_file(
serena_agent,
diagnostic_case=DiagnosticCase(
language=diagnostic_case.language,
relative_path=diagnostic_case.relative_path,
primary_symbol_name_path=diagnostic_case.primary_symbol_name_path,
primary_symbol_identifier=diagnostic_case.primary_symbol_identifier,
reference_symbol_name_path=diagnostic_case.primary_symbol_name_path,
reference_symbol_identifier=diagnostic_case.reference_symbol_identifier,
primary_message_fragment=diagnostic_case.primary_message_fragment,
reference_message_fragment=diagnostic_case.primary_message_fragment,
),
start_line=primary_position[0],
end_line=reference_position[0] - 1,
)
@pytest.mark.parametrize("serena_agent,diagnostic_case", WORKING_DIAGNOSTIC_TOOL_CASE_PARAMS, indirect=["serena_agent"])
def test_get_diagnostics_for_symbol(self, serena_agent: SerenaAgent, diagnostic_case: DiagnosticCase) -> None:
self._assert_diagnostics_for_symbol(
serena_agent,
diagnostic_case=diagnostic_case,
check_symbol_references=False,
)
@pytest.mark.parametrize("serena_agent,diagnostic_case", WORKING_DIAGNOSTIC_TOOL_CASE_PARAMS, indirect=["serena_agent"])
def test_get_diagnostics_for_symbol_with_references(self, serena_agent: SerenaAgent, diagnostic_case: DiagnosticCase) -> None:
self._assert_diagnostics_for_symbol(
serena_agent,
diagnostic_case=diagnostic_case,
check_symbol_references=True,
)
if IMPLEMENTATION_TOOL_TEST_CASES:
@pytest.mark.parametrize(
"serena_agent,symbol_name,def_file,impl_file,expected_symbol_name",
IMPLEMENTATION_TOOL_TEST_CASES,
indirect=["serena_agent"],
)
def test_find_symbol_implementations(
self,
serena_agent: SerenaAgent,
symbol_name: str,
def_file: str,
impl_file: str,
expected_symbol_name: str,
) -> None:
self._assert_find_symbol_implementations(serena_agent, symbol_name, def_file, impl_file, expected_symbol_name)
@pytest.mark.parametrize(
"serena_agent,name_path,substring_matching,expected_symbol_name,expected_kind,expected_file",
[
@@ -588,6 +1126,62 @@ class TestSerenaAgent:
new_content = read_project_file(serena_agent.get_active_project(), relative_path)
assert repl in new_content
@pytest.mark.parametrize(
"serena_agent",
[
pytest.param(Language.PYTHON, marks=pytest.mark.python),
pytest.param(Language.PYTHON_TY, marks=pytest.mark.python),
],
indirect=["serena_agent"],
)
def test_replace_content_reports_new_diagnostics(self, serena_agent: SerenaAgent):
"""Tests that file-level edits report newly introduced diagnostics."""
relative_path = os.path.join("test_repo", "services.py")
replace_content_tool = serena_agent.get_tool(ReplaceContentTool)
with project_file_modification_context(serena_agent, relative_path):
result = replace_content_tool.apply(
relative_path=relative_path,
needle="return container",
repl="return missing_container",
mode="literal",
)
diagnostics = parse_edit_diagnostics_result(result)
relative_path_result = diagnostics[relative_path]
diagnostic_messages = json.dumps(relative_path_result)
assert "missing_container" in diagnostic_messages
assert "create_service_container" in diagnostic_messages
@pytest.mark.parametrize(
"serena_agent",
[
pytest.param(Language.PYTHON, marks=pytest.mark.python),
pytest.param(Language.PYTHON_TY, marks=pytest.mark.python),
],
indirect=["serena_agent"],
)
def test_replace_symbol_body_reports_new_diagnostics(self, serena_agent: SerenaAgent):
"""Tests that symbol-level edits report newly introduced diagnostics."""
relative_path = os.path.join("test_repo", "services.py")
replace_symbol_body_tool = serena_agent.get_tool(ReplaceSymbolBodyTool)
with project_file_modification_context(serena_agent, relative_path):
result = replace_symbol_body_tool.apply(
name_path="create_service_container",
relative_path=relative_path,
body="""
def create_service_container() -> dict[str, Any]:
return missing_container
""",
)
diagnostics = parse_edit_diagnostics_result(result)
relative_path_result = diagnostics[relative_path]
diagnostic_messages = json.dumps(relative_path_result)
assert "missing_container" in diagnostic_messages
assert "create_service_container" in diagnostic_messages
@pytest.mark.parametrize(
"serena_agent",
[
+2 -2
View File
@@ -13,8 +13,8 @@ def _test_clojure_cli() -> bool:
CLI_FAIL = _test_clojure_cli()
TEST_APP_PATH = Path("src") / "test_app"
CORE_PATH = str(TEST_APP_PATH / "core.clj")
UTILS_PATH = str(TEST_APP_PATH / "utils.clj")
CORE_PATH = str(TEST_APP_PATH / "core.clj").replace("\\", "/")
UTILS_PATH = str(TEST_APP_PATH / "utils.clj").replace("\\", "/")
def is_clojure_cli_available() -> bool:
+28
View File
@@ -18,6 +18,7 @@ from solidlsp.ls_config import Language, LanguageServerConfig
from solidlsp.ls_types import SymbolKind
from solidlsp.ls_utils import SymbolUtils
from solidlsp.settings import SolidLSPSettings
from test.conftest import find_identifier_position, get_repo_path, language_has_verified_implementation_support
from test.solidlsp.conftest import format_symbol_for_assert, has_malformed_name, request_all_symbols
@@ -176,6 +177,33 @@ class TestCSharpLanguageServer:
assert "bool" in method_hover_text, f"Hover should include 'bool' return type, got: {method_hover_text}"
assert "IsAdult" in method_hover_text, f"Hover should include 'IsAdult' method name, got: {method_hover_text}"
if language_has_verified_implementation_support(Language.CSHARP):
@pytest.mark.parametrize("language_server", [Language.CSHARP], indirect=True)
def test_find_implementations(self, language_server: SolidLanguageServer) -> None:
repo_path = get_repo_path(Language.CSHARP)
pos = find_identifier_position(repo_path / "Services" / "IGreeter.cs", "FormatGreeting")
assert pos is not None, "Could not find IGreeter.FormatGreeting in fixture"
implementations = language_server.request_implementation("Services/IGreeter.cs", *pos)
assert implementations, "Expected at least one implementation of IGreeter.FormatGreeting"
assert any("ConsoleGreeter.cs" in implementation.get("relativePath", "") for implementation in implementations), (
f"Expected ConsoleGreeter.FormatGreeting in implementations, got: {implementations}"
)
@pytest.mark.parametrize("language_server", [Language.CSHARP], indirect=True)
def test_request_implementing_symbols(self, language_server: SolidLanguageServer) -> None:
repo_path = get_repo_path(Language.CSHARP)
pos = find_identifier_position(repo_path / "Services" / "IGreeter.cs", "FormatGreeting")
assert pos is not None, "Could not find IGreeter.FormatGreeting in fixture"
implementing_symbols = language_server.request_implementing_symbols("Services/IGreeter.cs", *pos)
assert implementing_symbols, "Expected implementing symbols for IGreeter.FormatGreeting"
assert any(
symbol.get("name") == "FormatGreeting" and "ConsoleGreeter.cs" in symbol["location"].get("relativePath", "")
for symbol in implementing_symbols
), f"Expected ConsoleGreeter.FormatGreeting symbol, got: {implementing_symbols}"
@pytest.mark.csharp
class TestCSharpSolutionProjectOpening:
@@ -11,6 +11,7 @@ from solidlsp import SolidLanguageServer
from solidlsp.ls_config import Language
from solidlsp.ls_types import SymbolKind
from solidlsp.ls_utils import SymbolUtils
from test.conftest import find_identifier_position, get_repo_path, language_has_verified_implementation_support
from test.solidlsp.conftest import format_symbol_for_assert, has_malformed_name, request_all_symbols
# Mark all tests in this module as fortran tests
@@ -38,6 +39,33 @@ class TestFortranLanguageServer:
# Verify subroutine symbol
assert SymbolUtils.symbol_tree_contains_name(symbols, "print_result"), "print_result subroutine not found in symbol tree"
if language_has_verified_implementation_support(Language.FORTRAN):
@pytest.mark.parametrize("language_server", [Language.FORTRAN], indirect=True)
def test_find_implementations(self, language_server: SolidLanguageServer) -> None:
repo_path = get_repo_path(Language.FORTRAN)
pos = find_identifier_position(repo_path / "modules" / "geometry.f90", "distance")
assert pos is not None, "Could not find interface distance in geometry.f90"
implementations = language_server.request_implementation("modules/geometry.f90", *pos)
assert implementations, "Expected implementations for geometry_types.distance"
implementation_files = {implementation.get("relativePath", "") for implementation in implementations}
assert implementation_files == {"modules/geometry.f90"}, f"Unexpected implementation locations: {implementations}"
assert len(implementations) >= 2, f"Expected module procedure implementations, got: {implementations}"
@pytest.mark.parametrize("language_server", [Language.FORTRAN], indirect=True)
def test_request_implementing_symbols(self, language_server: SolidLanguageServer) -> None:
repo_path = get_repo_path(Language.FORTRAN)
pos = find_identifier_position(repo_path / "modules" / "geometry.f90", "distance")
assert pos is not None, "Could not find interface distance in geometry.f90"
implementing_symbols = language_server.request_implementing_symbols("modules/geometry.f90", *pos)
assert implementing_symbols, "Expected implementing symbols for geometry_types.distance"
implementing_symbol_names = {symbol.get("name") for symbol in implementing_symbols}
assert {"distance_2d", "distance_3d"}.issubset(implementing_symbol_names), (
f"Expected distance_2d and distance_3d, got: {implementing_symbols}"
)
@pytest.mark.parametrize("language_server", [Language.FORTRAN], indirect=True)
def test_request_document_symbols(self, language_server: SolidLanguageServer) -> None:
"""Test that document symbols can be retrieved from Fortran files."""
+28 -1
View File
@@ -7,7 +7,7 @@ import pytest
from solidlsp import SolidLanguageServer
from solidlsp.ls_config import Language
from solidlsp.ls_utils import SymbolUtils
from test.conftest import is_ci
from test.conftest import find_identifier_position, get_repo_path, is_ci, language_has_verified_implementation_support
from test.solidlsp.conftest import format_symbol_for_assert, has_malformed_name, request_all_symbols
@@ -39,6 +39,33 @@ class TestFSharpLanguageServer:
symbol_names = [s.get("name") for s in symbols]
assert "main" in symbol_names, "main function not found in Program.fs symbols"
if language_has_verified_implementation_support(Language.FSHARP):
@pytest.mark.parametrize("language_server", [Language.FSHARP], indirect=True)
def test_find_implementations(self, language_server: SolidLanguageServer) -> None:
repo_path = get_repo_path(Language.FSHARP)
pos = find_identifier_position(repo_path / "Formatter.fs", "FormatGreeting")
assert pos is not None, "Could not find IGreeter.FormatGreeting in fixture"
implementations = language_server.request_implementation("Formatter.fs", *pos)
assert implementations, "Expected at least one implementation of IGreeter.FormatGreeting"
assert any("Formatter.fs" in implementation.get("relativePath", "") for implementation in implementations), (
f"Expected ConsoleGreeter.FormatGreeting in implementations, got: {implementations}"
)
@pytest.mark.parametrize("language_server", [Language.FSHARP], indirect=True)
def test_request_implementing_symbols(self, language_server: SolidLanguageServer) -> None:
repo_path = get_repo_path(Language.FSHARP)
pos = find_identifier_position(repo_path / "Formatter.fs", "FormatGreeting")
assert pos is not None, "Could not find IGreeter.FormatGreeting in fixture"
implementing_symbols = language_server.request_implementing_symbols("Formatter.fs", *pos)
assert implementing_symbols, "Expected implementing symbols for IGreeter.FormatGreeting"
assert any(
symbol.get("name") == "FormatGreeting" and "Formatter.fs" in symbol["location"].get("relativePath", "")
for symbol in implementing_symbols
), f"Expected ConsoleGreeter.FormatGreeting symbol, got: {implementing_symbols}"
@pytest.mark.parametrize("language_server", [Language.FSHARP], indirect=True)
def test_get_document_symbols_calculator(self, language_server: SolidLanguageServer) -> None:
"""Test getting document symbols from Calculator.fs file."""
+28
View File
@@ -7,6 +7,7 @@ from serena.symbol import LanguageServerSymbol
from solidlsp import SolidLanguageServer
from solidlsp.ls_config import Language
from solidlsp.ls_utils import SymbolUtils
from test.conftest import find_identifier_position, get_repo_path, language_has_verified_implementation_support
from test.solidlsp.conftest import format_symbol_for_assert, has_malformed_name, request_all_symbols
@@ -47,6 +48,33 @@ class TestGoLanguageServer:
refs = language_server.request_references(file_path, sel_start["line"], sel_start["character"])
assert any("main.go" in ref.get("uri", "") for ref in refs), "Expected at least one reference result to point at main.go"
if language_has_verified_implementation_support(Language.GO):
@pytest.mark.parametrize("language_server", [Language.GO], indirect=True)
def test_find_implementations(self, language_server: SolidLanguageServer) -> None:
repo_path = get_repo_path(Language.GO)
pos = find_identifier_position(repo_path / "main.go", "FormatGreeting")
assert pos is not None, "Could not find Greeter.FormatGreeting in fixture"
implementations = language_server.request_implementation("main.go", *pos)
assert implementations, "Expected at least one implementation of Greeter.FormatGreeting"
assert any("main.go" in implementation.get("relativePath", "") for implementation in implementations), (
f"Expected ConsoleGreeter.FormatGreeting in implementations, got: {implementations}"
)
@pytest.mark.parametrize("language_server", [Language.GO], indirect=True)
def test_request_implementing_symbols(self, language_server: SolidLanguageServer) -> None:
repo_path = get_repo_path(Language.GO)
pos = find_identifier_position(repo_path / "main.go", "FormatGreeting")
assert pos is not None, "Could not find Greeter.FormatGreeting in fixture"
implementing_symbols = language_server.request_implementing_symbols("main.go", *pos)
assert implementing_symbols, "Expected implementing symbols for Greeter.FormatGreeting"
assert any(
symbol.get("name") == "(ConsoleGreeter).FormatGreeting" and "main.go" in symbol["location"].get("relativePath", "")
for symbol in implementing_symbols
), f"Expected ConsoleGreeter.FormatGreeting symbol, got: {implementing_symbols}"
def _filter_symbols_by_name_in_repo(symbols: list | None, target_name: str, repo_name: str = "test_repo") -> list:
"""Filter workspace symbols to exact name matches in the test repo."""
+33 -1
View File
@@ -5,7 +5,12 @@ import pytest
from solidlsp import SolidLanguageServer
from solidlsp.ls_config import Language
from solidlsp.ls_utils import SymbolUtils
from test.conftest import language_tests_enabled
from test.conftest import (
find_identifier_position,
get_repo_path,
language_has_verified_implementation_support,
language_tests_enabled,
)
from test.solidlsp.conftest import format_symbol_for_assert, has_malformed_name, request_all_symbols
pytestmark = [pytest.mark.java, pytest.mark.skipif(not language_tests_enabled(Language.JAVA), reason="Java tests disabled")]
@@ -52,6 +57,33 @@ class TestJavaLanguageServer:
assert SymbolUtils.symbol_tree_contains_name(symbols, "Utils"), "Utils missing from overview"
assert SymbolUtils.symbol_tree_contains_name(symbols, "Model"), "Model missing from overview"
if language_has_verified_implementation_support(Language.JAVA):
@pytest.mark.parametrize("language_server", [Language.JAVA], indirect=True)
def test_find_implementations(self, language_server: SolidLanguageServer) -> None:
repo_path = get_repo_path(Language.JAVA)
pos = find_identifier_position(repo_path / "src/main/java/test_repo/Greeter.java", "formatGreeting")
assert pos is not None, "Could not find Greeter.formatGreeting in fixture"
implementations = language_server.request_implementation("src/main/java/test_repo/Greeter.java", *pos)
assert implementations, "Expected at least one implementation of Greeter.formatGreeting"
assert any("ConsoleGreeter.java" in implementation.get("relativePath", "") for implementation in implementations), (
f"Expected ConsoleGreeter.formatGreeting in implementations, got: {implementations}"
)
@pytest.mark.parametrize("language_server", [Language.JAVA], indirect=True)
def test_request_implementing_symbols(self, language_server: SolidLanguageServer) -> None:
repo_path = get_repo_path(Language.JAVA)
pos = find_identifier_position(repo_path / "src/main/java/test_repo/Greeter.java", "formatGreeting")
assert pos is not None, "Could not find Greeter.formatGreeting in fixture"
implementing_symbols = language_server.request_implementing_symbols("src/main/java/test_repo/Greeter.java", *pos)
assert implementing_symbols, "Expected implementing symbols for Greeter.formatGreeting"
assert any(
symbol.get("name") == "formatGreeting" and "ConsoleGreeter.java" in symbol["location"].get("relativePath", "")
for symbol in implementing_symbols
), f"Expected ConsoleGreeter.formatGreeting symbol, got: {implementing_symbols}"
@pytest.mark.parametrize("language_server", [Language.JAVA], indirect=True)
def test_bare_symbol_names(self, language_server) -> None:
all_symbols = request_all_symbols(language_server)
+33 -5
View File
@@ -10,6 +10,7 @@ import pytest
from solidlsp import SolidLanguageServer
from solidlsp.ls_config import Language
from solidlsp.ls_types import SymbolKind
from test.conftest import find_identifier_position, get_repo_path, language_has_verified_implementation_support
from test.solidlsp.conftest import format_symbol_for_assert, has_malformed_name, request_all_symbols
@@ -82,6 +83,33 @@ class TestLuaLanguageServer:
# Check for Logger class/table
assert "Logger" in all_symbols or any("Logger" in s for s in all_symbols), "Logger not found in symbols"
if language_has_verified_implementation_support(Language.LUA):
@pytest.mark.parametrize("language_server", [Language.LUA], indirect=True)
def test_find_implementations(self, language_server: SolidLanguageServer) -> None:
repo_path = get_repo_path(Language.LUA)
pos = find_identifier_position(repo_path / "src" / "animals.lua", "speak")
assert pos is not None, "Could not find Animal:speak in fixture"
implementations = language_server.request_implementation("src/animals.lua", *pos)
assert implementations, "Expected at least one implementation of Animal:speak"
assert any("animals.lua" in implementation.get("relativePath", "") for implementation in implementations), (
f"Expected Dog:speak in implementations, got: {implementations}"
)
@pytest.mark.parametrize("language_server", [Language.LUA], indirect=True)
def test_request_implementing_symbols(self, language_server: SolidLanguageServer) -> None:
repo_path = get_repo_path(Language.LUA)
pos = find_identifier_position(repo_path / "src" / "animals.lua", "speak")
assert pos is not None, "Could not find Animal:speak in fixture"
implementing_symbols = language_server.request_implementing_symbols("src/animals.lua", *pos)
assert implementing_symbols, "Expected implementing symbols for Animal:speak"
assert any(
symbol.get("name") == "speak" and "animals.lua" in symbol["location"].get("relativePath", "")
for symbol in implementing_symbols
), f"Expected Dog:speak symbol, got: {implementing_symbols}"
@pytest.mark.parametrize("language_server", [Language.LUA], indirect=True)
def test_find_symbols_in_main(self, language_server: SolidLanguageServer) -> None:
"""Test finding functions in main.lua."""
@@ -133,7 +161,7 @@ class TestLuaLanguageServer:
assert refs is not None
assert isinstance(refs, list)
# add function appears in: main.lua (lines 16, 71), test_calculator.lua (lines 22, 23, 24)
# add function appears in: main.lua (lines 17, 78), test_calculator.lua (lines 22, 23, 24)
# Note: The declaration itself may or may not be included as a reference
assert len(refs) >= 5, f"Should find at least 5 references to calculator.add, found {len(refs)}"
@@ -153,7 +181,7 @@ class TestLuaLanguageServer:
# Check main.lua has usages
assert "main.lua" in ref_files, "Should find add usages in main.lua"
assert 15 in ref_files["main.lua"] or 70 in ref_files["main.lua"], (
assert 16 in ref_files["main.lua"] or 77 in ref_files["main.lua"], (
f"Should find add usage in main.lua, found at lines {ref_files.get('main.lua', [])}"
)
@@ -189,7 +217,7 @@ class TestLuaLanguageServer:
assert refs is not None
assert isinstance(refs, list)
# trim function appears in: usage (line 32 in main.lua)
# trim function appears in: usage (line 33 in main.lua)
# Note: The declaration itself may or may not be included as a reference
assert len(refs) >= 1, f"Should find at least 1 reference to utils.trim, found {len(refs)}"
@@ -209,8 +237,8 @@ class TestLuaLanguageServer:
# Check main.lua has usage
assert "main.lua" in ref_files, "Should find trim usage in main.lua"
assert 31 in ref_files["main.lua"], (
f"Should find trim usage at line 32 (0-indexed: 31) in main.lua, found at lines {ref_files.get('main.lua', [])}"
assert 32 in ref_files["main.lua"], (
f"Should find trim usage at line 33 (0-indexed: 32) in main.lua, found at lines {ref_files.get('main.lua', [])}"
)
# Check for cross-file references from main.lua
+21
View File
@@ -12,6 +12,7 @@ import pytest
from serena.project import Project
from serena.util.text_utils import LineType
from solidlsp import SolidLanguageServer
from test.conftest import PYTHON_LANGUAGE_BACKENDS
from test.solidlsp.conftest import PYTHON_BACKEND_LANGUAGES, format_symbol_for_assert, has_malformed_name, request_all_symbols
@@ -79,6 +80,26 @@ class TestPythonLanguageServerBasics:
references = language_server.request_references(file_path, sel_start["line"], sel_start["character"])
assert len(references) > 1, "Should get valid references for create_user (using selectionRange if present)"
@pytest.mark.parametrize("language_server", PYTHON_LANGUAGE_BACKENDS, indirect=True)
def test_request_text_document_diagnostics_with_filters(self, language_server: SolidLanguageServer) -> None:
file_path = os.path.join("test_repo", "diagnostics_sample.py")
diagnostics = language_server.request_text_document_diagnostics(file_path)
assert len(diagnostics) >= 2
diagnostic_messages = [diagnostic["message"] for diagnostic in diagnostics]
assert any("missing_user" in message for message in diagnostic_messages), diagnostic_messages
assert any("undefined_name" in message for message in diagnostic_messages), diagnostic_messages
factory_diagnostics = language_server.request_text_document_diagnostics(file_path, start_line=3, end_line=5, min_severity=1)
factory_messages = [diagnostic["message"] for diagnostic in factory_diagnostics]
assert factory_messages, "Expected diagnostics in broken_factory range"
assert all("missing_user" in message for message in factory_messages), factory_messages
consumer_diagnostics = language_server.request_text_document_diagnostics(file_path, start_line=7, end_line=10, min_severity=1)
consumer_messages = [diagnostic["message"] for diagnostic in consumer_diagnostics]
assert consumer_messages, "Expected diagnostics in broken_consumer range"
assert all("undefined_name" in message for message in consumer_messages), consumer_messages
class TestProjectBasics:
@pytest.mark.parametrize("project", PYTHON_BACKEND_LANGUAGES, indirect=True)
+28
View File
@@ -5,6 +5,7 @@ import pytest
from solidlsp import SolidLanguageServer
from solidlsp.ls_config import Language
from solidlsp.ls_utils import SymbolUtils
from test.conftest import find_identifier_position, get_repo_path, language_has_verified_implementation_support
from test.solidlsp.conftest import format_symbol_for_assert, has_malformed_name, request_all_symbols
@@ -57,6 +58,33 @@ class TestRustLanguageServer:
assert SymbolUtils.symbol_tree_contains_name(symbols, "main"), "main missing from overview"
assert SymbolUtils.symbol_tree_contains_name(symbols, "add"), "add missing from overview"
if language_has_verified_implementation_support(Language.RUST):
@pytest.mark.parametrize("language_server", [Language.RUST], indirect=True)
def test_find_implementations(self, language_server: SolidLanguageServer) -> None:
repo_path = get_repo_path(Language.RUST)
pos = find_identifier_position(repo_path / os.path.join("src", "lib.rs"), "format_greeting")
assert pos is not None, "Could not find Greeter.format_greeting in fixture"
implementations = language_server.request_implementation(os.path.join("src", "lib.rs"), *pos)
assert implementations, "Expected at least one implementation of Greeter.format_greeting"
assert any("src/lib.rs" in implementation.get("relativePath", "").replace("\\", "/") for implementation in implementations), (
f"Expected ConsoleGreeter.format_greeting in implementations, got: {implementations}"
)
@pytest.mark.parametrize("language_server", [Language.RUST], indirect=True)
def test_request_implementing_symbols(self, language_server: SolidLanguageServer) -> None:
repo_path = get_repo_path(Language.RUST)
pos = find_identifier_position(repo_path / os.path.join("src", "lib.rs"), "format_greeting")
assert pos is not None, "Could not find Greeter.format_greeting in fixture"
implementing_symbols = language_server.request_implementing_symbols(os.path.join("src", "lib.rs"), *pos)
assert implementing_symbols, "Expected implementing symbols for Greeter.format_greeting"
assert any(
symbol.get("name") == "format_greeting" and "src/lib.rs" in symbol["location"].get("relativePath", "").replace("\\", "/")
for symbol in implementing_symbols
), f"Expected ConsoleGreeter.format_greeting symbol, got: {implementing_symbols}"
@pytest.mark.parametrize("language_server", [Language.RUST], indirect=True)
def test_bare_symbol_names(self, language_server) -> None:
all_symbols = request_all_symbols(language_server)
@@ -13,6 +13,7 @@ import pytest
from solidlsp import SolidLanguageServer
from solidlsp.ls_config import Language
from test.conftest import language_has_verified_implementation_support
from test.solidlsp.conftest import format_symbol_for_assert, has_malformed_name, request_all_symbols
@@ -165,6 +166,33 @@ class TestSolidityLanguageServerBasics:
ref_files = {ref.get("uri", "") for ref in references}
assert any("Token.sol" in uri for uri in ref_files), "IERC20.transfer references should include Token.sol"
if language_has_verified_implementation_support(Language.SOLIDITY):
@pytest.mark.parametrize("language_server", [Language.SOLIDITY], indirect=True)
@pytest.mark.parametrize("repo_path", [Language.SOLIDITY], indirect=True)
def test_find_implementations(self, language_server: SolidLanguageServer, repo_path: Path) -> None:
pos = _find_identifier_position(repo_path / "contracts/interfaces/IERC20.sol", "transfer")
assert pos is not None, "Should find 'transfer' identifier in IERC20.sol"
implementations = language_server.request_implementation("contracts/interfaces/IERC20.sol", *pos)
assert implementations, "Expected Token.transfer to be returned as an implementation"
assert any("Token.sol" in implementation.get("relativePath", "") for implementation in implementations), (
f"Expected Token.transfer implementation, got: {implementations}"
)
@pytest.mark.parametrize("language_server", [Language.SOLIDITY], indirect=True)
@pytest.mark.parametrize("repo_path", [Language.SOLIDITY], indirect=True)
def test_request_implementing_symbols(self, language_server: SolidLanguageServer, repo_path: Path) -> None:
pos = _find_identifier_position(repo_path / "contracts/interfaces/IERC20.sol", "transfer")
assert pos is not None, "Should find 'transfer' identifier in IERC20.sol"
implementing_symbols = language_server.request_implementing_symbols("contracts/interfaces/IERC20.sol", *pos)
assert implementing_symbols, "Expected implementing symbols for IERC20.transfer"
assert any(
symbol.get("name") == "transfer" and "Token.sol" in symbol["location"].get("relativePath", "")
for symbol in implementing_symbols
), f"Expected Token.transfer symbol, got: {implementing_symbols}"
@pytest.mark.parametrize("language_server", [Language.SOLIDITY], indirect=True)
def test_bare_symbol_names(self, language_server) -> None:
all_symbols = request_all_symbols(language_server)
@@ -0,0 +1,139 @@
import os
import shutil
from pathlib import Path
import pytest
from solidlsp import SolidLanguageServer
from solidlsp.ls_config import Language
from test.conftest import find_identifier_occurrence_position, is_ci
from test.solidlsp import clojure as clj
@pytest.mark.parametrize(
"language_server,relative_path,identifier,occurrence_index,column_offset,expected_name,expected_definition_file",
[
pytest.param(
Language.PYTHON,
os.path.join("test_repo", "services.py"),
"User",
1,
1,
"User",
"models.py",
marks=pytest.mark.python,
),
pytest.param(
Language.PYTHON_TY,
os.path.join("test_repo", "services.py"),
"User",
1,
1,
"User",
"models.py",
marks=pytest.mark.python,
),
pytest.param(Language.GO, "main.go", "Helper", 0, 1, "Helper", "main.go", marks=pytest.mark.go),
pytest.param(
Language.JAVA,
os.path.join("src", "main", "java", "test_repo", "Main.java"),
"Model",
0,
1,
"Model",
"Model.java",
marks=pytest.mark.java,
),
pytest.param(
Language.KOTLIN,
os.path.join("src", "main", "kotlin", "test_repo", "Main.kt"),
"Model",
0,
1,
"Model",
"Model.kt",
marks=[pytest.mark.kotlin] + ([pytest.mark.skip(reason="Kotlin LSP JVM crashes on restart in CI")] if is_ci else []),
),
pytest.param(
Language.RUST,
os.path.join("src", "main.rs"),
"format_greeting",
0,
1,
"format_greeting",
"lib.rs",
marks=pytest.mark.rust,
),
pytest.param(Language.PHP, "index.php", "helperFunction", 0, 5, "helperFunction", "helper.php", marks=pytest.mark.php),
pytest.param(
Language.CLOJURE,
clj.UTILS_PATH,
"multiply",
0,
1,
"multiply",
clj.CORE_PATH,
marks=[
pytest.mark.clojure,
pytest.mark.skipif(not clj.is_clojure_cli_available(), reason="clojure CLI is not installed"),
],
),
pytest.param(Language.CSHARP, "Program.cs", "Add", 0, 1, "Add", "Program.cs", marks=pytest.mark.csharp),
pytest.param(
Language.POWERSHELL,
"main.ps1",
"Convert-ToUpperCase",
0,
1,
"function Convert-ToUpperCase ()",
"utils.ps1",
marks=pytest.mark.powershell,
),
pytest.param(Language.CPP_CCLS, "a.cpp", "add", 0, 1, "add", "b.cpp", marks=pytest.mark.cpp),
pytest.param(
Language.LEAN4,
"Main.lean",
"add",
0,
1,
"add",
"Helper.lean",
marks=[
pytest.mark.lean4,
pytest.mark.skipif(shutil.which("lean") is None, reason="Lean is not installed"),
pytest.mark.xfail(reason="Lean4 LS does not reliably resolve cross-file defining symbols in CI"),
],
),
pytest.param(Language.TYPESCRIPT, "index.ts", "helperFunction", 1, 1, "helperFunction", "index.ts", marks=pytest.mark.typescript),
pytest.param(
Language.FSHARP,
"Program.fs",
"add",
0,
1,
"add",
"Calculator.fs",
marks=[pytest.mark.fsharp, pytest.mark.xfail(reason="F# language server cannot reliably resolve defining symbols")],
),
],
indirect=["language_server"],
)
def test_request_defining_symbol_matrix(
language_server: SolidLanguageServer,
relative_path: str,
identifier: str,
occurrence_index: int,
column_offset: int,
expected_name: str,
expected_definition_file: str,
) -> None:
repo_root = Path(language_server.repository_root_path)
position = find_identifier_occurrence_position(repo_root / relative_path, identifier, occurrence_index, column_offset)
assert position is not None, f"Could not find occurrence {occurrence_index} of {identifier!r} in {relative_path}"
defining_symbol = language_server.request_defining_symbol(relative_path, *position)
assert defining_symbol is not None, f"Expected a defining symbol for {identifier!r} in {relative_path}"
assert defining_symbol.get("name") == expected_name, f"Expected defining symbol name {expected_name!r}, got: {defining_symbol}"
assert expected_definition_file in defining_symbol["location"].get("relativePath", ""), (
f"Expected defining symbol in {expected_definition_file!r}, got: {defining_symbol}"
)
+79
View File
@@ -0,0 +1,79 @@
from pathlib import Path
import pytest
from serena.symbol import LanguageServerSymbolRetriever
from solidlsp import SolidLanguageServer
from test.conftest import find_identifier_position
from test.diagnostics_cases import DIAGNOSTIC_CASE_PARAMS, DiagnosticCase
@pytest.mark.parametrize("language_server,diagnostic_case", DIAGNOSTIC_CASE_PARAMS, indirect=["language_server"])
def test_request_text_document_diagnostics_matrix(
language_server: SolidLanguageServer,
diagnostic_case: DiagnosticCase,
) -> None:
diagnostics = language_server.request_text_document_diagnostics(diagnostic_case.relative_path, min_severity=1)
assert diagnostics, f"Expected diagnostics for {diagnostic_case.language.value}:{diagnostic_case.relative_path}"
diagnostic_messages = [diagnostic["message"] for diagnostic in diagnostics]
assert any(diagnostic_case.primary_message_fragment in message for message in diagnostic_messages), diagnostic_messages
assert any(diagnostic_case.reference_message_fragment in message for message in diagnostic_messages), diagnostic_messages
repo_root = language_server.repository_root_path
primary_symbol_position = find_identifier_position(
Path(repo_root) / diagnostic_case.relative_path, diagnostic_case.primary_symbol_identifier
)
reference_symbol_position = find_identifier_position(
Path(repo_root) / diagnostic_case.relative_path, diagnostic_case.reference_symbol_identifier
)
assert primary_symbol_position is not None
assert reference_symbol_position is not None
primary_diagnostics = language_server.request_text_document_diagnostics(
diagnostic_case.relative_path,
start_line=primary_symbol_position[0],
end_line=reference_symbol_position[0] - 1,
min_severity=1,
)
primary_messages = [diagnostic["message"] for diagnostic in primary_diagnostics]
assert primary_messages, f"Expected range-filtered diagnostics for {diagnostic_case.primary_symbol_identifier}"
assert any(diagnostic_case.primary_message_fragment in message for message in primary_messages), primary_messages
assert all(diagnostic_case.reference_message_fragment not in message for message in primary_messages), primary_messages
@pytest.mark.parametrize("project_with_ls,diagnostic_case", DIAGNOSTIC_CASE_PARAMS, indirect=["project_with_ls"])
def test_get_symbol_diagnostics_matrix(project_with_ls, diagnostic_case: DiagnosticCase) -> None:
symbol_retriever = LanguageServerSymbolRetriever(project_with_ls)
diagnostics_by_symbol = symbol_retriever.get_symbol_diagnostics(
diagnostic_case.primary_symbol_name_path,
reference_file=diagnostic_case.relative_path,
min_severity=1,
)
diagnostic_messages_by_symbol = {
symbol.get_name_path(): [diagnostic["message"] for diagnostic in diagnostics]
for symbol, diagnostics in diagnostics_by_symbol.items()
}
assert diagnostic_case.primary_symbol_name_path in diagnostic_messages_by_symbol, diagnostic_messages_by_symbol
assert any(
diagnostic_case.primary_message_fragment in message
for message in diagnostic_messages_by_symbol[diagnostic_case.primary_symbol_name_path]
), diagnostic_messages_by_symbol
diagnostics_with_references = symbol_retriever.get_symbol_diagnostics(
diagnostic_case.primary_symbol_name_path,
reference_file=diagnostic_case.relative_path,
check_symbol_references=True,
min_severity=1,
)
diagnostic_messages_with_references = {
symbol.get_name_path(): [diagnostic["message"] for diagnostic in diagnostics]
for symbol, diagnostics in diagnostics_with_references.items()
}
assert diagnostic_case.primary_symbol_name_path in diagnostic_messages_with_references, diagnostic_messages_with_references
assert diagnostic_case.reference_symbol_name_path in diagnostic_messages_with_references, diagnostic_messages_with_references
assert any(
diagnostic_case.reference_message_fragment in message
for message in diagnostic_messages_with_references[diagnostic_case.reference_symbol_name_path]
), diagnostic_messages_with_references
+2 -2
View File
@@ -3,13 +3,13 @@ import os
import pytest
from solidlsp import SolidLanguageServer
from solidlsp.ls_config import Language
from test.conftest import PYTHON_LANGUAGE_BACKENDS
class TestLanguageServerCommonFunctionality:
"""Test common functionality of SolidLanguageServer base implementation (not language-specific behaviour)."""
@pytest.mark.parametrize("language_server", [Language.PYTHON], indirect=True)
@pytest.mark.parametrize("language_server", PYTHON_LANGUAGE_BACKENDS, indirect=True)
def test_open_file_cache_invalidate(self, language_server: SolidLanguageServer) -> None:
"""
Tests that the file buffer cache is invalidated when the file is changed on disk.
@@ -5,6 +5,7 @@ import pytest
from solidlsp import SolidLanguageServer
from solidlsp.ls_config import Language
from solidlsp.ls_utils import SymbolUtils
from test.conftest import find_identifier_position, get_repo_path, language_has_verified_implementation_support
from test.solidlsp.conftest import format_symbol_for_assert, has_malformed_name, request_all_symbols
@@ -33,6 +34,33 @@ class TestTypescriptLanguageServer:
"index.ts should reference helperFunction (tried all positions in selectionRange)"
)
if language_has_verified_implementation_support(Language.TYPESCRIPT):
@pytest.mark.parametrize("language_server", [Language.TYPESCRIPT], indirect=True)
def test_find_implementations(self, language_server: SolidLanguageServer) -> None:
repo_path = get_repo_path(Language.TYPESCRIPT)
pos = find_identifier_position(repo_path / "formatters.ts", "formatGreeting")
assert pos is not None, "Could not find Greeter.formatGreeting in fixture"
implementations = language_server.request_implementation("formatters.ts", *pos)
assert implementations, "Expected at least one implementation of Greeter.formatGreeting"
assert any("formatters.ts" in implementation.get("relativePath", "") for implementation in implementations), (
f"Expected ConsoleGreeter.formatGreeting in implementations, got: {implementations}"
)
@pytest.mark.parametrize("language_server", [Language.TYPESCRIPT], indirect=True)
def test_request_implementing_symbols(self, language_server: SolidLanguageServer) -> None:
repo_path = get_repo_path(Language.TYPESCRIPT)
pos = find_identifier_position(repo_path / "formatters.ts", "formatGreeting")
assert pos is not None, "Could not find Greeter.formatGreeting in fixture"
implementing_symbols = language_server.request_implementing_symbols("formatters.ts", *pos)
assert implementing_symbols, "Expected implementing symbols for Greeter.formatGreeting"
assert any(
symbol.get("name") == "formatGreeting" and "formatters.ts" in symbol["location"].get("relativePath", "")
for symbol in implementing_symbols
), f"Expected ConsoleGreeter.formatGreeting symbol, got: {implementing_symbols}"
@pytest.mark.parametrize("language_server", [Language.TYPESCRIPT], indirect=True)
def test_bare_symbol_names(self, language_server) -> None:
all_symbols = request_all_symbols(language_server)