mirror of
https://github.com/tiennm99/serena.git
synced 2026-10-11 12:29:04 +00:00
Diagnostics, implementation and definition tools (WIP)
This commit is contained in:
1 parent
df0f476614
commit
c55e0c900e
66 files changed
+3928
-38
No files matched your search
@@ -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)
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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
@@ -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.
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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. """
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
|
||||
+7
@@ -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
@@ -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" "$@"
|
||||
@@ -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
|
||||
@@ -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();
|
||||
|
||||
@@ -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",
|
||||
[
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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}"
|
||||
)
|
||||
@@ -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
|
||||
@@ -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)
|
||||
|
||||
Reference in new issue
Block a user