Merge pull request #284 from oraios/jetbrains

JetBrains plugin support
This commit is contained in:
Dominik Jain
2025-07-11 00:03:46 +02:00
committed by GitHub
29 changed files with 1654 additions and 1026 deletions
+1
View File
@@ -225,6 +225,7 @@ ignore = [
"B904", # forces use of raise from other_exception
"RUF012", # forbids mutable attributes as ClassVar
"SIM117", # forbids nested with statements
"C400", # wants to unnecessarily force use of list comprehension
]
unfixable = ["F841", "F601", "F602", "B018"]
extend-fixable = ["F401", "B905", "W291"]
Binary file not shown.
+5 -2
View File
@@ -3,10 +3,13 @@ This script demonstrates how to use Serena's tools locally, useful
for testing or development. Here the tools will be operation the serena repo itself.
"""
import json
from pprint import pprint
from serena.agent import *
from serena.agent import SerenaAgent
from serena.config.serena_config import SerenaConfig
from serena.constants import REPO_ROOT
from serena.tools import FindFileTool, FindReferencingSymbolsTool, SearchForPatternTool
if __name__ == "__main__":
agent = SerenaAgent(project=REPO_ROOT, serena_config=SerenaConfig(gui_log_window_enabled=False, web_dashboard=False))
@@ -18,7 +21,7 @@ if __name__ == "__main__":
result = agent.execute_task(
lambda: search_pattern_tool.apply(
r"def request_parsed_files.*?\).*?\)",
r"def request_full_.*?\).*?\)",
restrict_search_to_code_files=False,
relative_path="src/solidlsp",
paths_include_glob="**/ls.py",
+1 -1
View File
@@ -1,4 +1,4 @@
from serena.agent import ToolRegistry
if __name__ == "__main__":
ToolRegistry.print_tool_overview()
ToolRegistry().print_tool_overview()
+60 -182
View File
@@ -8,7 +8,6 @@ import platform
import sys
import threading
import webbrowser
from abc import ABC, abstractmethod
from collections import defaultdict
from collections.abc import Callable
from concurrent.futures import Future, ThreadPoolExecutor
@@ -17,24 +16,21 @@ from pathlib import Path
from typing import TYPE_CHECKING, Any, TypeVar, Union
import click
from pathspec import PathSpec
from sensai.util import logging
from sensai.util.logging import LogTime
from tqdm import tqdm
from serena import serena_version
from serena.config.context_mode import SerenaAgentContext, SerenaAgentMode
from serena.config.serena_config import Project, SerenaConfig, get_serena_managed_dir
from serena.config.serena_config import SerenaConfig, ToolSet, get_serena_managed_dir
from serena.constants import (
SERENA_LOG_FORMAT,
)
from serena.dashboard import MemoryLogHandler, SerenaDashboardAPI
from serena.project import Project
from serena.prompt_factory import SerenaPromptFactory
from serena.symbol import SymbolManager
from serena.tools import Tool, ToolRegistry
from serena.util.file_system import GitignoreParser, match_path
from solidlsp import SolidLanguageServer
from solidlsp.ls_config import LanguageServerConfig
from solidlsp.ls_logger import LanguageServerLogger
if TYPE_CHECKING:
from serena.gui_log_viewer import GuiLogViewerHandler
@@ -43,7 +39,6 @@ log = logging.getLogger(__name__)
TTool = TypeVar("TTool", bound="Tool")
T = TypeVar("T")
SUCCESS_RESULT = "OK"
DEFAULT_TOOL_TIMEOUT: float = 240
class ProjectNotFoundError(Exception):
@@ -66,25 +61,7 @@ class LinesRead:
del self.files[relative_path]
class MemoriesManager(ABC):
@abstractmethod
def load_memory(self, name: str) -> str:
pass
@abstractmethod
def save_memory(self, name: str, content: str) -> str:
pass
@abstractmethod
def list_memories(self) -> list[str]:
pass
@abstractmethod
def delete_memory(self, name: str) -> str:
pass
class MemoriesManagerMDFilesInProject(MemoriesManager):
class MemoriesManager:
def __init__(self, project_root: str):
self._memory_dir = Path(get_serena_managed_dir(project_root)) / "memories"
self._memory_dir.mkdir(parents=True, exist_ok=True)
@@ -117,56 +94,6 @@ class MemoriesManagerMDFilesInProject(MemoriesManager):
return f"Memory {name} deleted."
def create_ls_for_project(
project: str | Project,
log_level: int = logging.INFO,
ls_timeout: float | None = DEFAULT_TOOL_TIMEOUT - 5,
trace_lsp_communication: bool = False,
) -> SolidLanguageServer:
"""
Create a language server for a project. Note that you will have to start it
before performing any LS operations.
:param project: either a path to the project root or a ProjectConfig instance.
If no project.yml is found, the default project configuration will be used.
:param log_level: the log level for the language server
:param ls_timeout: the timeout for the language server
:param trace_lsp_communication: whether to trace LSP communication
:return: the language server
"""
if isinstance(project, str):
project_instance = Project.load(project, autogenerate=True)
else:
project_instance = project
project_config = project_instance.project_config
ignored_paths = project_config.ignored_paths
if len(ignored_paths) > 0:
log.info(f"Using {len(ignored_paths)} ignored paths from the explicit project configuration.")
log.debug(f"Ignored paths: {ignored_paths}")
if project_config.ignore_all_files_in_gitignore:
log.info(f"Parsing all gitignore files in {project_instance.project_root}")
gitignore_parser = GitignoreParser(project_instance.project_root)
log.info(f"Found {len(gitignore_parser.get_ignore_specs())} gitignore files.")
for spec in gitignore_parser.get_ignore_specs():
log.debug(f"Adding {len(spec.patterns)} patterns from {spec.file_path} to the ignored paths.")
ignored_paths.extend(spec.patterns)
log.debug(f"Using {len(ignored_paths)} ignored paths in total.")
multilspy_config = LanguageServerConfig(
code_language=project_instance.language,
ignored_paths=ignored_paths,
trace_lsp_communication=trace_lsp_communication,
)
ls_logger = LanguageServerLogger(log_level=log_level)
log.info(f"Creating language server instance for {project_instance.project_root}.")
return SolidLanguageServer.create(
multilspy_config,
ls_logger,
project_instance.project_root,
timeout=ls_timeout,
)
@click.command()
@click.argument("project", type=click.Path(exists=True), required=False, default=os.getcwd())
@click.option("--log-level", type=click.Choice(["DEBUG", "INFO", "WARNING", "ERROR", "CRITICAL"]), default="WARNING")
@@ -177,11 +104,22 @@ def index_project(project: str, log_level: str = "INFO") -> None:
:param project: the project to index. By default, the current working directory is used.
"""
log_level_int = logging.getLevelNamesMapping()[log_level.upper()]
project = os.path.abspath(project)
project_instance = Project.load(os.path.abspath(project))
print(f"Indexing symbols in project {project}")
ls = create_ls_for_project(project, log_level=log_level_int)
ls = project_instance.create_language_server(log_level=log_level_int)
save_after_n_files = 10
with ls.start_server():
ls.index_repository()
parsed_files = project_instance.gather_source_files()
files_processed = 0
pbar = tqdm(parsed_files, disable=False)
for relative_file_path in pbar:
pbar.set_description(f"Indexing ({os.path.basename(relative_file_path)})")
ls.request_document_symbols(relative_file_path, include_body=False)
ls.request_document_symbols(relative_file_path, include_body=True)
files_processed += 1
if files_processed % save_after_n_files == 0:
ls.save_cache()
ls.save_cache()
print(f"Symbols saved to {ls.cache_path}")
@@ -234,14 +172,9 @@ class SerenaAgent:
self._context = context
# instantiate all tool classes
self._all_tools: dict[type[Tool], Tool] = {tool_class: tool_class(self) for tool_class in ToolRegistry.get_all_tool_classes()}
self._all_tools: dict[type[Tool], Tool] = {tool_class: tool_class(self) for tool_class in ToolRegistry().get_all_tool_classes()}
tool_names = [tool.get_name_from_cls() for tool in self._all_tools.values()]
# determine the set exposed tools (which e.g. the MCP shall see), limited by the context
# (which is fixed for the session)
excluded_tool_classes = set(self._context.get_excluded_tool_classes())
self._exposed_tools = {tc: t for tc, t in self._all_tools.items() if tc not in excluded_tool_classes}
# If GUI log window is enabled, set the tool names for highlighting
if self._gui_log_handler is not None:
self._gui_log_handler.log_viewer.set_tool_names(tool_names)
@@ -261,6 +194,16 @@ class SerenaAgent:
log.info(f"Starting Serena server (version={serena_version()}, process id={os.getpid()}, parent process id={os.getppid()})")
log.info("Configuration file: %s", self.serena_config.config_file_path)
log.info("Available projects: {}".format(", ".join(self.serena_config.project_names)))
log.info(f"Loaded tools ({len(self._all_tools)}): {', '.join([tool.get_name_from_cls() for tool in self._all_tools.values()])}")
# determine the base toolset defining the set of exposed tools (which e.g. the MCP shall see),
# limited by the Serena config, the context (which is fixed for the session) and JetBrains mode
tool_inclusion_definitions = [self.serena_config, self._context]
if self.serena_config.jetbrains:
tool_inclusion_definitions.append(SerenaAgentMode.from_name_internal("jetbrains"))
self._base_tool_set = ToolSet.default().apply(*tool_inclusion_definitions)
self._exposed_tools = {tc: t for tc, t in self._all_tools.items() if self._base_tool_set.includes_name(t.get_name())}
log.info(f"Number of exposed tools: {len(self._exposed_tools)}")
# create executor for starting the language server and running tools in another thread
# This executor is used to achieve linear task execution, so it is important to use a single-threaded executor.
@@ -276,21 +219,14 @@ class SerenaAgent:
self._active_project: Project | None = None
self._active_project_root: str | None = None
self.language_server: SolidLanguageServer | None = None
self.symbol_manager: SymbolManager | None = None
self.memories_manager: MemoriesManager | None = None
self.lines_read: LinesRead | None = None
self.ignore_spec: PathSpec # not set to None to avoid assert statements
"""Ignore spec, extracted from the project's gitignore files and the explicitly configured ignored paths."""
# set the active modes
if modes is None:
modes = SerenaAgentMode.load_default_modes()
self._modes = modes
# log tool information
log.info(f"Loaded tools ({len(self._all_tools)}): {', '.join([tool.get_name_from_cls() for tool in self._all_tools.values()])}")
log.info(f"Number of exposed tools given {self._context}: {len(self._exposed_tools)}")
self._active_tools: dict[type[Tool], Tool] = {}
self._update_active_tools()
@@ -325,46 +261,6 @@ class SerenaAgent:
raise ValueError("Cannot get project root if no project is active.")
return project.project_root
def path_is_inside_project(self, path: str | Path) -> bool:
"""
Checks if the given (absolute or relative) path is inside the project directory.
Note that even relative paths may be outside if the contain ".." or point to symlinks.
"""
path = Path(path)
_proj_root = Path(self.get_project_root())
if not path.is_absolute():
path = _proj_root / path
path = path.resolve()
return path.is_relative_to(_proj_root)
def path_is_gitignored(self, path: str | Path) -> bool:
"""
Checks if the given path is ignored by git. Non absolute paths are assumed to be relative to the project root.
"""
path = Path(path)
if path.is_absolute():
relative_path = path.relative_to(self.get_project_root())
else:
relative_path = path
# always ignore paths inside .git
if len(relative_path.parts) > 0 and relative_path.parts[0] == ".git":
return True
return match_path(str(relative_path), self.ignore_spec, root_path=self.get_project_root())
def validate_relative_path(self, relative_path: str) -> None:
"""
Validates that the given relative path is safe to read or edit,
meaning it's inside the project directory and is not ignored by git.
"""
if not self.path_is_inside_project(relative_path):
raise ValueError(f"{relative_path=} points to path outside of the repository root, can't use it for safety reasons")
if self.path_is_gitignored(relative_path):
raise ValueError(f"File {relative_path} is gitignored, can't read or edit it for safety reasons")
def get_exposed_tool_instances(self) -> list["Tool"]:
"""
:return: the tool instances which are exposed (e.g. to the MCP client).
@@ -382,6 +278,15 @@ class SerenaAgent:
"""
return self._active_project
def get_active_project_or_raise(self) -> Project:
"""
:return: the active project or raises an exception if no project is active
"""
project = self.get_active_project()
if project is None:
raise ValueError("No active project. Please activate a project first.")
return project
def set_modes(self, modes: list[SerenaAgentMode]) -> None:
"""
Set the current mode configurations.
@@ -407,43 +312,23 @@ class SerenaAgent:
def _update_active_tools(self) -> None:
"""
Update the active tools based on context, modes, and project configuration.
All tool exclusions are merged together.
Update the active tools based on enabled modes and the active project.
The base tool set already takes the Serena configuration and the context into account
(as well as any internal modes that are not handled dynamically, such as JetBrains mode).
"""
excluded_tool_classes: set[type[Tool]] = set()
# modes
for mode in self._modes:
mode_excluded_tool_classes = mode.get_excluded_tool_classes()
if len(mode_excluded_tool_classes) > 0:
log.info(
f"Mode {mode.name} excluded {len(mode_excluded_tool_classes)} tools: {', '.join([tool.get_name_from_cls() for tool in mode_excluded_tool_classes])}"
)
excluded_tool_classes.update(mode_excluded_tool_classes)
# context
context_excluded_tool_classes = self._context.get_excluded_tool_classes()
if len(context_excluded_tool_classes) > 0:
log.info(
f"Context {self._context.name} excluded {len(context_excluded_tool_classes)} tools: {', '.join([tool.get_name_from_cls() for tool in context_excluded_tool_classes])}"
)
excluded_tool_classes.update(context_excluded_tool_classes)
# project config
tool_set = self._base_tool_set.apply(*self._modes)
if self._active_project is not None:
project_excluded_tool_classes = self._active_project.project_config.get_excluded_tool_classes()
if len(project_excluded_tool_classes) > 0:
log.info(
f"Project {self._active_project.project_name} excluded {len(project_excluded_tool_classes)} tools: {', '.join([tool.get_name_from_cls() for tool in project_excluded_tool_classes])}"
)
excluded_tool_classes.update(project_excluded_tool_classes)
tool_set = tool_set.apply(self._active_project.project_config)
if self._active_project.project_config.read_only:
for tool_class in self._all_tools:
if tool_class.can_edit():
excluded_tool_classes.add(tool_class)
tool_set = tool_set.without_editing_tools()
self._active_tools = {
tool_class: tool_instance for tool_class, tool_instance in self._all_tools.items() if tool_class not in excluded_tool_classes
tool_class: tool_instance
for tool_class, tool_instance in self._all_tools.items()
if tool_set.includes_name(tool_instance.get_name())
}
log.info(f"Active tools after all exclusions ({len(self._active_tools)}): {', '.join(self.get_active_tool_names())}")
log.info(f"Active tools ({len(self._active_tools)}): {', '.join(self.get_active_tool_names())}")
def issue_task(self, task: Callable[[], Any], name: str | None = None) -> Future:
"""
@@ -476,31 +361,30 @@ class SerenaAgent:
future = self.issue_task(task)
return future.result()
def is_using_language_server(self) -> bool:
"""
:return: whether this agent uses language server-based code analysis
"""
return not self.serena_config.jetbrains
def _activate_project(self, project: Project) -> None:
log.info(f"Activating {project.project_name} at {project.project_root}")
self._active_project = project
self._update_active_tools()
# initialize project-specific instances which do not depend on the language server
self.memories_manager = MemoriesManagerMDFilesInProject(project.project_root)
self.memories_manager = MemoriesManager(project.project_root)
self.lines_read = LinesRead()
# reset project-specific instances that depend on the language server
self.symbol_manager = None
def init_language_server() -> None:
# start the language server
with LogTime("Language server initialization", logger=log):
self.reset_language_server()
assert self.language_server is not None
self.ignore_spec = self.language_server.get_ignore_spec()
# initialize project-specific instances which depend on the language server
log.debug(f"Initializing symbol and memories manager for {project.project_name} at {project.project_root}")
self.symbol_manager = SymbolManager(self.language_server, self)
# initialize the language server in the background
self.issue_task(init_language_server)
# initialize the language server in the background (if in language server mode)
if self.is_using_language_server():
self.issue_task(init_language_server)
if self._project_activation_callback is not None:
self._project_activation_callback()
@@ -626,8 +510,7 @@ class SerenaAgent:
# instantiate and start the language server
assert self._active_project is not None
self.language_server = create_ls_for_project(
self._active_project,
self.language_server = self._active_project.create_language_server(
log_level=self.serena_config.log_level,
ls_timeout=ls_timeout,
trace_lsp_communication=self.serena_config.trace_lsp_communication,
@@ -638,17 +521,12 @@ class SerenaAgent:
raise RuntimeError(
f"Failed to start the language server for {self._active_project.project_name} at {self._active_project.project_root}"
)
if self.symbol_manager is not None:
log.debug("Setting the language server in the agent's symbol manager")
self.symbol_manager.set_language_server(self.language_server)
else:
log.debug("No symbol manager available yet, skipping setting the language server")
def get_tool(self, tool_class: type[TTool]) -> TTool:
return self._all_tools[tool_class] # type: ignore
def print_tool_overview(self) -> None:
ToolRegistry.print_tool_overview(self._active_tools.values())
ToolRegistry().print_tool_overview(self._active_tools.values())
def mark_file_modified(self, relativ_path: str) -> None:
assert self.lines_read is not None
+297
View File
@@ -0,0 +1,297 @@
import json
import logging
import os
from abc import ABC, abstractmethod
from collections.abc import Iterable, Iterator, Reversible
from contextlib import contextmanager
from typing import TYPE_CHECKING, Generic, Optional, TypeVar
from serena.symbol import JetBrainsSymbol, LanguageServerSymbol, LanguageServerSymbolRetriever, PositionInFile, Symbol
from solidlsp import SolidLanguageServer
from solidlsp.ls import LSPFileBuffer
from solidlsp.ls_utils import TextUtils
from .project import Project
from .tools.jetbrains_plugin_client import JetBrainsPluginClient
if TYPE_CHECKING:
from .agent import SerenaAgent
log = logging.getLogger(__name__)
TSymbol = TypeVar("TSymbol", bound=Symbol)
class CodeEditor(Generic[TSymbol], ABC):
def __init__(self, project_root: str, agent: Optional["SerenaAgent"] = None) -> None:
self.project_root = project_root
self.agent = agent
class EditedFile(ABC):
@abstractmethod
def get_contents(self) -> str:
"""
:return: the contents of the file.
"""
@abstractmethod
def delete_text_between_positions(self, start_pos: PositionInFile, end_pos: PositionInFile) -> None:
pass
@abstractmethod
def insert_text_at_position(self, pos: PositionInFile, text: str) -> None:
pass
@contextmanager
def _open_file_context(self, relative_path: str) -> Iterator["CodeEditor.EditedFile"]:
"""
Context manager for opening a file
"""
raise NotImplementedError("This method must be overridden for each subclass")
@contextmanager
def _edited_file_context(self, relative_path: str) -> Iterator["CodeEditor.EditedFile"]:
"""
Context manager for editing a file.
"""
with self._open_file_context(relative_path) as edited_file:
yield edited_file
# save the file
abs_path = os.path.join(self.project_root, relative_path)
with open(abs_path, "w", encoding="utf-8") as f:
f.write(edited_file.get_contents())
# notify agent (if provided)
if self.agent is not None:
self.agent.mark_file_modified(relative_path)
@abstractmethod
def _find_unique_symbol(self, name_path: str, relative_file_path: str) -> TSymbol:
"""
Finds the unique symbol with the given name in the given file.
If no such symbol exists, raises a ValueError.
:param name_path: the name path
:param relative_file_path: the relative path of the file in which to search for the symbol.
:return: the unique symbol
"""
def replace_body(self, name_path: str, relative_file_path: str, body: str) -> None:
"""
Replaces the body of the symbol with the given name_path in the given file.
:param name_path: the name path of the symbol to replace.
:param relative_file_path: the relative path of the file in which the symbol is defined.
:param body: the new body
"""
symbol = self._find_unique_symbol(name_path, relative_file_path)
start_pos = symbol.get_body_start_position_or_raise()
end_pos = symbol.get_body_end_position_or_raise()
with self._edited_file_context(relative_file_path) as edited_file:
# make sure the replacement adds no additional newlines (before or after) - all newlines
# and whitespace before/after should remain the same, so we strip it entirely
body = body.strip()
edited_file.delete_text_between_positions(start_pos, end_pos)
edited_file.insert_text_at_position(start_pos, body)
@staticmethod
def _count_leading_newlines(text: Iterable) -> int:
cnt = 0
for c in text:
if c == "\n":
cnt += 1
elif c == "\r":
continue
else:
break
return cnt
@classmethod
def _count_trailing_newlines(cls, text: Reversible) -> int:
return cls._count_leading_newlines(reversed(text))
def insert_after_symbol(self, name_path: str, relative_file_path: str, body: str) -> None:
"""
Inserts content after the symbol with the given name in the given file.
"""
symbol = self._find_unique_symbol(name_path, relative_file_path)
# make sure body always ends with at least one newline
if not body.endswith("\n"):
body += "\n"
pos = symbol.get_body_end_position_or_raise()
# start at the beginning of the next line
col = 0
line = pos.line + 1
# make sure a suitable number of leading empty lines is used (at least 0/1 depending on the symbol type,
# otherwise as many as the caller wanted to insert)
original_leading_newlines = self._count_leading_newlines(body)
body = body.lstrip("\r\n")
min_empty_lines = 0
if symbol.is_neighbouring_definition_separated_by_empty_line():
min_empty_lines = 1
num_leading_empty_lines = max(min_empty_lines, original_leading_newlines)
if num_leading_empty_lines:
body = ("\n" * num_leading_empty_lines) + body
# make sure the one line break succeeding the original symbol, which we repurposed as prefix via
# `line += 1`, is replaced
body = body.rstrip("\r\n") + "\n"
with self._edited_file_context(relative_file_path) as edited_file:
edited_file.insert_text_at_position(PositionInFile(line, col), body)
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.
"""
symbol = self._find_unique_symbol(name_path, relative_file_path)
symbol_start_pos = symbol.get_body_start_position_or_raise()
# insert position is the start of line where the symbol is defined
line = symbol_start_pos.line
col = 0
original_trailing_empty_lines = self._count_trailing_newlines(body) - 1
# ensure eol is present at end
body = body.rstrip() + "\n"
# add suitable number of trailing empty lines after the body (at least 0/1 depending on the symbol type,
# otherwise as many as the caller wanted to insert)
min_trailing_empty_lines = 0
if symbol.is_neighbouring_definition_separated_by_empty_line():
min_trailing_empty_lines = 1
num_trailing_newlines = max(min_trailing_empty_lines, original_trailing_empty_lines)
body += "\n" * num_trailing_newlines
# apply edit
with self._edited_file_context(relative_file_path) as edited_file:
edited_file.insert_text_at_position(PositionInFile(line=line, col=col), body)
def insert_at_line(self, relative_path: str, line: int, content: str) -> None:
"""
Inserts content at the given line in the given file.
:param relative_path: the relative path of the file in which to insert content
:param line: the 0-based index of the line to insert content at
:param content: the content to insert
"""
with self._edited_file_context(relative_path) as edited_file:
edited_file.insert_text_at_position(PositionInFile(line, 0), content)
def delete_lines(self, relative_path: str, start_line: int, end_line: int) -> None:
"""
Deletes lines in the given file.
:param relative_path: the relative path of the file in which to delete lines
:param start_line: the 0-based index of the first line to delete (inclusive)
:param end_line: the 0-based index of the last line to delete (inclusive)
"""
start_col = 0
end_line_for_delete = end_line + 1
end_col = 0
with self._edited_file_context(relative_path) as edited_file:
start_pos = PositionInFile(line=start_line, col=start_col)
end_pos = PositionInFile(line=end_line_for_delete, col=end_col)
edited_file.delete_text_between_positions(start_pos, end_pos)
def delete_symbol(self, name_path: str, relative_file_path: str) -> None:
"""
Deletes the symbol with the given name in the given file.
"""
symbol = self._find_unique_symbol(name_path, relative_file_path)
start_pos = symbol.get_body_start_position_or_raise()
end_pos = symbol.get_body_end_position_or_raise()
with self._edited_file_context(relative_file_path) as edited_file:
edited_file.delete_text_between_positions(start_pos, end_pos)
class LanguageServerCodeEditor(CodeEditor[LanguageServerSymbol]):
def __init__(self, symbol_retriever: LanguageServerSymbolRetriever, agent: Optional["SerenaAgent"] = None):
super().__init__(project_root=symbol_retriever.get_language_server().repository_root_path, agent=agent)
self._symbol_retriever = symbol_retriever
@property
def _lang_server(self) -> SolidLanguageServer:
return self._symbol_retriever.get_language_server()
class EditedFile(CodeEditor.EditedFile):
def __init__(self, lang_server: SolidLanguageServer, relative_path: str, file_buffer: LSPFileBuffer):
self._lang_server = lang_server
self._relative_path = relative_path
self._file_buffer = file_buffer
def get_contents(self) -> str:
return self._file_buffer.contents
def delete_text_between_positions(self, start_pos: PositionInFile, end_pos: PositionInFile) -> None:
self._lang_server.delete_text_between_positions(self._relative_path, start_pos.to_lsp_position(), end_pos.to_lsp_position())
def insert_text_at_position(self, pos: PositionInFile, text: str) -> None:
self._lang_server.insert_text_at_position(self._relative_path, pos.line, pos.col, text)
@contextmanager
def _open_file_context(self, relative_path: str) -> Iterator["CodeEditor.EditedFile"]:
with self._lang_server.open_file(relative_path) as file_buffer:
yield self.EditedFile(self._lang_server, relative_path, file_buffer)
def _get_code_file_content(self, relative_path: str) -> str:
"""Get the content of a file using the language server."""
return self._lang_server.language_server.retrieve_full_file_content(relative_path)
def _find_unique_symbol(self, name_path: str, relative_file_path: str) -> LanguageServerSymbol:
symbol_candidates = self._symbol_retriever.find_by_name(name_path, within_relative_path=relative_file_path)
if len(symbol_candidates) == 0:
raise ValueError(f"No symbol with name {name_path} found in file {relative_file_path}")
if len(symbol_candidates) > 1:
raise ValueError(
f"Found multiple {len(symbol_candidates)} symbols with name {name_path} in file {relative_file_path}. "
"Their locations are: \n " + json.dumps([s.location.to_dict() for s in symbol_candidates], indent=2)
)
return symbol_candidates[0]
class JetBrainsCodeEditor(CodeEditor[JetBrainsSymbol]):
def __init__(self, project: Project, agent: Optional["SerenaAgent"] = None) -> None:
self._project = project
super().__init__(project_root=project.project_root, agent=agent)
class EditedFile(CodeEditor.EditedFile):
def __init__(self, relative_path: str, project: Project):
path = os.path.join(project.project_root, relative_path)
log.info("Editing file: %s", path)
with open(path, encoding=project.project_config.encoding) as f:
self._content = f.read()
def get_contents(self) -> str:
return self._content
def delete_text_between_positions(self, start_pos: PositionInFile, end_pos: PositionInFile) -> None:
self._content, _ = TextUtils.delete_text_between_positions(
self._content, start_pos.line, start_pos.col, end_pos.line, end_pos.col
)
def insert_text_at_position(self, pos: PositionInFile, text: str) -> None:
self._content, _, _ = TextUtils.insert_text_at_position(self._content, pos.line, pos.col, text)
@contextmanager
def _open_file_context(self, relative_path: str) -> Iterator["CodeEditor.EditedFile"]:
yield self.EditedFile(relative_path, self._project)
def _find_unique_symbol(self, name_path: str, relative_file_path: str) -> JetBrainsSymbol:
with JetBrainsPluginClient() as client:
result = client.find_symbol(name_path, relative_path=relative_file_path, include_body=False, depth=0, include_location=True)
symbols = result["symbols"]
if not symbols:
raise ValueError(f"No symbol with name {name_path} found in file {relative_file_path}")
if len(symbols) > 1:
raise ValueError(
f"Found multiple {len(symbols)} symbols with name {name_path} in file {relative_file_path}. "
"Their locations are: \n " + json.dumps([s["location"] for s in symbols], indent=2)
)
return JetBrainsSymbol(symbols[0], self._project)
+19 -21
View File
@@ -4,7 +4,7 @@ Context and Mode configuration loader
import os
from copy import copy
from dataclasses import asdict, dataclass, field
from dataclasses import asdict, dataclass
from enum import Enum
from pathlib import Path
from typing import TYPE_CHECKING, Self
@@ -13,16 +13,17 @@ import yaml
from sensai.util import logging
from sensai.util.string import ToStringMixin
from serena.constants import CONTEXT_YAMLS_DIR, DEFAULT_CONTEXT, DEFAULT_MODES, MODE_YAMLS_DIR
from serena.config.serena_config import ToolInclusionDefinition
from serena.constants import CONTEXT_YAMLS_DIR, DEFAULT_CONTEXT, DEFAULT_MODES, INTERNAL_MODE_YAMLS_DIR, MODE_YAMLS_DIR
if TYPE_CHECKING:
from serena.agent import Tool
pass
log = logging.getLogger(__name__)
@dataclass
class SerenaAgentMode:
@dataclass(kw_only=True)
class SerenaAgentMode(ToolInclusionDefinition, ToStringMixin):
"""Represents a mode of operation for the agent, typically read off a YAML file.
An agent can be in multiple modes simultaneously as long as they are not mutually exclusive.
The modes can be adjusted after the agent is running, for example for switching from planning to editing.
@@ -31,7 +32,9 @@ class SerenaAgentMode:
name: str
prompt: str
description: str = ""
excluded_tools: set[str] = field(default_factory=set)
def _tostring_includes(self) -> list[str]:
return ["name"]
def to_json_dict(self) -> dict[str, str | list[str]]:
result = asdict(self)
@@ -50,12 +53,6 @@ class SerenaAgentMode:
if self.excluded_tools:
print(" excluded tools:\n " + ", ".join(sorted(self.excluded_tools)))
def get_excluded_tool_classes(self) -> list[type["Tool"]]:
"""Get the list of tool classes that are excluded from the mode."""
from serena.agent import ToolRegistry
return [ToolRegistry.get_tool_class_by_name(tool_name) for tool_name in self.excluded_tools]
@classmethod
def from_yaml(cls, yaml_path: str | Path) -> Self:
"""Load a mode from a YAML file."""
@@ -75,6 +72,14 @@ class SerenaAgentMode:
)
return cls.from_yaml(yaml_path)
@classmethod
def from_name_internal(cls, name: str) -> Self:
"""Loads an internal Serena mode"""
yaml_path = os.path.join(INTERNAL_MODE_YAMLS_DIR, f"{name}.yml")
if not os.path.exists(yaml_path):
raise FileNotFoundError(f"Internal mode '{name}' not found in {INTERNAL_MODE_YAMLS_DIR}")
return cls.from_yaml(yaml_path)
@classmethod
def list_registered_mode_names(cls) -> list[str]:
"""Names of all registered modes (from the corresponding YAML files in the serena repo)."""
@@ -93,8 +98,8 @@ class SerenaAgentMode:
return cls.from_yaml(name_or_path)
@dataclass
class SerenaAgentContext(ToStringMixin):
@dataclass(kw_only=True)
class SerenaAgentContext(ToolInclusionDefinition, ToStringMixin):
"""Represents a context where the agent is operating (an IDE, a chat, etc.), typically read off a YAML file.
An agent can only be in a single context at a time.
The contexts cannot be changed after the agent is running.
@@ -103,7 +108,6 @@ class SerenaAgentContext(ToStringMixin):
name: str
prompt: str
description: str = ""
excluded_tools: set[str] = field(default_factory=set)
def _tostring_includes(self) -> list[str]:
return ["name"]
@@ -119,12 +123,6 @@ class SerenaAgentContext(ToStringMixin):
data["excluded_tools"] = set(data["excluded_tools"])
return cls(**data)
def get_excluded_tool_classes(self) -> list[type["Tool"]]:
"""Get the list of tool classes that are excluded from the context."""
from serena.agent import ToolRegistry
return [ToolRegistry.get_tool_class_by_name(tool_name) for tool_name in self.excluded_tools]
@classmethod
def from_yaml(cls, yaml_path: str | Path) -> Self:
"""Load a context from a YAML file."""
+100 -39
View File
@@ -4,11 +4,12 @@ The Serena Model Context Protocol (MCP) Server
import os
import shutil
from collections.abc import Iterable
from copy import deepcopy
from dataclasses import dataclass, field
from functools import cached_property
from pathlib import Path
from typing import TYPE_CHECKING, Any, Self, TypeVar
from typing import TYPE_CHECKING, Any, Optional, Self, TypeVar
import yaml
from ruamel.yaml.comments import CommentedMap
@@ -22,19 +23,90 @@ from serena.constants import (
SELENA_CONFIG_TEMPLATE_FILE,
SERENA_MANAGED_DIR_NAME,
)
from serena.tools import Tool, ToolRegistry
from serena.util.general import load_yaml, save_yaml
from serena.util.inspection import determine_programming_language_composition
from solidlsp.ls_config import Language
if TYPE_CHECKING:
pass
from ..project import Project
log = logging.getLogger(__name__)
T = TypeVar("T")
DEFAULT_TOOL_TIMEOUT: float = 240
class ToolSet:
def __init__(self, tool_names: set[str]) -> None:
self._tool_names = tool_names
@classmethod
def default(cls) -> "ToolSet":
"""
:return: the default tool set, which contains all tools that are enabled by default
"""
from serena.tools import ToolRegistry
return cls(set(ToolRegistry().get_tool_names_default_enabled()))
def apply(self, *tool_inclusion_definitions: "ToolInclusionDefinition") -> "ToolSet":
"""
:param tool_inclusion_definitions: the definitions to apply
:return: a new tool set with the definitions applied
"""
from serena.tools import ToolRegistry
registry = ToolRegistry()
tool_names = set(self._tool_names)
for definition in tool_inclusion_definitions:
included_tools = []
excluded_tools = []
for included_tool in definition.included_optional_tools:
if not registry.is_valid_tool_name(included_tool):
raise ValueError(f"Invalid tool name '{included_tool}' provided for inclusion")
if included_tool not in tool_names:
tool_names.add(included_tool)
included_tools.append(included_tool)
for excluded_tool in definition.excluded_tools:
if not registry.is_valid_tool_name(excluded_tool):
raise ValueError(f"Invalid tool name '{excluded_tool}' provided for exclusion")
if excluded_tool in self._tool_names:
tool_names.remove(excluded_tool)
excluded_tools.append(excluded_tool)
if included_tools:
log.info(f"{definition} included {len(included_tools)} tools: {', '.join(included_tools)}")
if excluded_tools:
log.info(f"{definition} excluded {len(excluded_tools)} tools: {', '.join(excluded_tools)}")
return ToolSet(tool_names)
def without_editing_tools(self) -> "ToolSet":
"""
:return: a new tool set that excludes all tools that can edit
"""
from serena.tools import ToolRegistry
registry = ToolRegistry()
tool_names = set(self._tool_names)
for tool_name in self._tool_names:
if registry.get_tool_class_by_name(tool_name).can_edit():
tool_names.remove(tool_name)
return ToolSet(tool_names)
def get_tool_names(self) -> set[str]:
"""
Returns the names of the tools that are currently included in the tool set.
"""
return self._tool_names
def includes_name(self, tool_name: str) -> bool:
return tool_name in self._tool_names
@dataclass
class ToolInclusionDefinition:
excluded_tools: Iterable[str] = ()
included_optional_tools: Iterable[str] = ()
class SerenaConfigError(Exception):
pass
@@ -56,12 +128,11 @@ def is_running_in_docker() -> bool:
return False
@dataclass
class ProjectConfig(ToStringMixin):
@dataclass(kw_only=True)
class ProjectConfig(ToolInclusionDefinition, ToStringMixin):
project_name: str
language: Language
ignored_paths: list[str] = field(default_factory=list)
excluded_tools: set[str] = field(default_factory=set)
read_only: bool = False
ignore_all_files_in_gitignore: bool = True
initial_prompt: str = ""
@@ -69,6 +140,9 @@ class ProjectConfig(ToStringMixin):
SERENA_DEFAULT_PROJECT_FILE = "project.yml"
def _tostring_includes(self) -> list[str]:
return ["project_name"]
@classmethod
def autogenerate(cls, project_root: str | Path, project_name: str | None = None, save_to_disk: bool = True) -> Self:
"""
@@ -128,7 +202,8 @@ class ProjectConfig(ToStringMixin):
project_name=project_name,
language=language,
ignored_paths=data.get("ignored_paths", []),
excluded_tools=set(data.get("excluded_tools", [])),
excluded_tools=data.get("excluded_tools", []),
included_optional_tools=data.get("included_optional_tools", []),
read_only=data.get("read_only", False),
ignore_all_files_in_gitignore=data.get("ignore_all_files_in_gitignore", True),
initial_prompt=data.get("initial_prompt", ""),
@@ -153,44 +228,16 @@ class ProjectConfig(ToStringMixin):
yaml_data["project_name"] = project_root.name
return cls._from_dict(yaml_data)
def get_excluded_tool_classes(self) -> set[type["Tool"]]:
return set(ToolRegistry.get_tool_class_by_name(tool_name) for tool_name in self.excluded_tools)
@dataclass
class Project:
project_root: str
project_config: ProjectConfig
@property
def project_name(self) -> str:
return self.project_config.project_name
@property
def language(self) -> Language:
return self.project_config.language
@classmethod
def load(cls, project_root: str | Path, autogenerate: bool = True) -> Self:
project_root = Path(project_root).resolve()
if not project_root.exists():
raise FileNotFoundError(f"Project root not found: {project_root}")
project_config = ProjectConfig.load(project_root, autogenerate=autogenerate)
return cls(project_root=str(project_root), project_config=project_config)
def path_to_project_yml(self) -> str:
return os.path.join(self.project_root, self.project_config.rel_path_to_project_yml())
@dataclass(kw_only=True)
class SerenaConfig:
class SerenaConfig(ToolInclusionDefinition, ToStringMixin):
"""
Holds the Serena agent configuration, which is typically loaded from a YAML configuration file
(when instantiated via :method:`from_config_file`), which is updated when projects are added or removed.
For testing purposes, it can also be instantiated directly with the desired parameters.
"""
projects: list[Project] = field(default_factory=list)
projects: list["Project"] = field(default_factory=list)
gui_log_window_enabled: bool = False
log_level: int = logging.INFO
trace_lsp_communication: bool = False
@@ -203,10 +250,17 @@ class SerenaConfig:
the path to the configuration file to which updates of the configuration shall be saved;
if None, the configuration is not saved to disk
"""
jetbrains: bool = False
"""
whether to apply JetBrains mode
"""
CONFIG_FILE = "serena_config.yml"
CONFIG_FILE_DOCKER = "serena_config.docker.yml" # Docker-specific config file; auto-generated if missing, mounted via docker-compose for user customization
def _tostring_includes(self) -> list[str]:
return ["config_file_path"]
@classmethod
def _generate_config_file(cls, config_file_path: str) -> None:
"""
@@ -243,6 +297,8 @@ class SerenaConfig:
"""
Static constructor to create SerenaConfig from the configuration file
"""
from ..project import Project
config_file_path = cls._determine_config_file_path()
# create the configuration file from the template if necessary
@@ -292,6 +348,9 @@ class SerenaConfig:
instance.web_dashboard_open_on_launch = loaded_commented_yaml.get("web_dashboard_open_on_launch", True)
instance.tool_timeout = loaded_commented_yaml.get("tool_timeout", DEFAULT_TOOL_TIMEOUT)
instance.trace_lsp_communication = loaded_commented_yaml.get("trace_lsp_communication", False)
instance.excluded_tools = loaded_commented_yaml.get("excluded_tools", [])
instance.included_optional_tools = loaded_commented_yaml.get("included_optional_tools", [])
instance.jetbrains = loaded_commented_yaml.get("jetbrains", False)
# re-save the configuration file if any migrations were performed
if num_project_migrations > 0:
@@ -334,7 +393,7 @@ class SerenaConfig:
def project_names(self) -> list[str]:
return sorted(project.project_config.project_name for project in self.projects)
def get_project(self, project_root_or_name: str) -> Project | None:
def get_project(self, project_root_or_name: str) -> Optional["Project"]:
for project in self.projects:
if project.project_config.project_name == project_root_or_name:
return project
@@ -345,7 +404,7 @@ class SerenaConfig:
return project
return None
def add_project_from_path(self, project_root: Path | str, project_name: str | None = None) -> tuple[Project, bool]:
def add_project_from_path(self, project_root: Path | str, project_name: str | None = None) -> tuple["Project", bool]:
"""
Add a project to the Serena configuration from a given path. Will raise a FileExistsError if the
name or path is already registered.
@@ -357,6 +416,8 @@ class SerenaConfig:
saved to disk. It may be that no new project configuration was generated if the project configuration already
exists on disk but the project itself was not added yet to the Serena configuration.
"""
from ..project import Project
project_root = Path(project_root).resolve()
if not project_root.exists():
raise FileNotFoundError(f"Error: Path does not exist: {project_root}")
+1
View File
@@ -7,6 +7,7 @@ REPO_ROOT = str(_repo_root_path)
PROMPT_TEMPLATES_DIR = str(_serena_pkg_path / "resources" / "config" / "prompt_templates")
CONTEXT_YAMLS_DIR = str(_serena_pkg_path / "resources" / "config" / "contexts")
MODE_YAMLS_DIR = str(_serena_pkg_path / "resources" / "config" / "modes")
INTERNAL_MODE_YAMLS_DIR = str(_serena_pkg_path / "resources" / "config" / "internal_modes")
SERENA_DASHBOARD_DIR = str(_serena_pkg_path / "resources" / "dashboard")
SERENA_ICON_DIR = str(_serena_pkg_path / "resources" / "icons")
+4 -4
View File
@@ -25,7 +25,7 @@ from serena.agent import (
)
from serena.config.context_mode import SerenaAgentContext, SerenaAgentMode
from serena.constants import DEFAULT_CONTEXT, DEFAULT_MODES
from serena.tools import ToolInterface
from serena.tools import Tool
from serena.util.exception import show_fatal_exception_safe
log = logging.getLogger(__name__)
@@ -62,7 +62,7 @@ class SerenaMCPFactory:
self.project = project
@staticmethod
def make_mcp_tool(tool: ToolInterface) -> MCPTool:
def make_mcp_tool(tool: Tool) -> MCPTool:
func_name = tool.get_name()
func_doc = tool.get_apply_docstring() or ""
func_arg_metadata = tool.get_apply_fn_metadata()
@@ -106,7 +106,7 @@ class SerenaMCPFactory:
)
@abstractmethod
def _iter_tools(self) -> Iterator[ToolInterface]:
def _iter_tools(self) -> Iterator[Tool]:
pass
# noinspection PyProtectedMember
@@ -204,7 +204,7 @@ class SerenaMCPFactorySingleProcess(SerenaMCPFactory):
def _instantiate_agent(self, serena_config: SerenaConfig, modes: list[SerenaAgentMode]) -> None:
self.agent = SerenaAgent(project=self.project, serena_config=serena_config, context=self.context, modes=modes)
def _iter_tools(self) -> Iterator[ToolInterface]:
def _iter_tools(self) -> Iterator[Tool]:
assert self.agent is not None
yield from self.agent.get_exposed_tool_instances()
+274
View File
@@ -0,0 +1,274 @@
import logging
import os
from pathlib import Path
from typing import Self
import pathspec
from serena.config.serena_config import DEFAULT_TOOL_TIMEOUT, ProjectConfig
from serena.text_utils import MatchedConsecutiveLines, search_files
from serena.util.file_system import GitignoreParser, match_path
from solidlsp import SolidLanguageServer
from solidlsp.ls_config import Language, LanguageServerConfig
from solidlsp.ls_logger import LanguageServerLogger
log = logging.getLogger(__name__)
class Project:
def __init__(self, project_root: str, project_config: ProjectConfig):
self.project_root = project_root
self.project_config = project_config
# gather ignored paths from the project configuration and gitignore files
ignored_patterns = project_config.ignored_paths
if len(ignored_patterns) > 0:
log.info(f"Using {len(ignored_patterns)} ignored paths from the explicit project configuration.")
log.debug(f"Ignored paths: {ignored_patterns}")
if project_config.ignore_all_files_in_gitignore:
log.info(f"Parsing all gitignore files in {self.project_root}")
gitignore_parser = GitignoreParser(self.project_root)
log.info(f"Found {len(gitignore_parser.get_ignore_specs())} gitignore files.")
for spec in gitignore_parser.get_ignore_specs():
log.debug(f"Adding {len(spec.patterns)} patterns from {spec.file_path} to the ignored paths.")
ignored_patterns.extend(spec.patterns)
self._ignored_patterns = ignored_patterns
# Set up the pathspec matcher for the ignored paths
# for all absolute paths in ignored_paths, convert them to relative paths
processed_patterns = []
for pattern in set(ignored_patterns):
# Normalize separators (pathspec expects forward slashes)
pattern = pattern.replace(os.path.sep, "/")
processed_patterns.append(pattern)
log.debug(f"Processing {len(processed_patterns)} ignored paths")
self._ignore_spec = pathspec.PathSpec.from_lines(pathspec.patterns.GitWildMatchPattern, processed_patterns)
@property
def project_name(self) -> str:
return self.project_config.project_name
@property
def language(self) -> Language:
return self.project_config.language
@classmethod
def load(cls, project_root: str | Path, autogenerate: bool = True) -> Self:
project_root = Path(project_root).resolve()
if not project_root.exists():
raise FileNotFoundError(f"Project root not found: {project_root}")
project_config = ProjectConfig.load(project_root, autogenerate=autogenerate)
return cls(project_root=str(project_root), project_config=project_config)
def path_to_project_yml(self) -> str:
return os.path.join(self.project_root, self.project_config.rel_path_to_project_yml())
def read_file(self, relative_path: str) -> str:
"""
Reads a file relative to the project root.
:param relative_path: the path to the file relative to the project root
:return: the content of the file
"""
abs_path = Path(self.project_root) / relative_path
if not abs_path.exists():
raise FileNotFoundError(f"File not found: {abs_path}")
return abs_path.read_text(encoding=self.project_config.encoding)
def get_ignore_spec(self) -> pathspec.PathSpec:
"""
:return: the pathspec matcher for the paths that were configured to be ignored,
either explicitly or implicitly through .gitignore files.
"""
return self._ignore_spec
def _is_ignored_dirname(self, dirname: str) -> bool:
return dirname.startswith(".")
def _is_ignored_relative_path(self, relative_path: str, ignore_non_source_files: bool = True) -> bool:
"""
Determine whether a path should be ignored based on file type and ignore patterns.
:param relative_path: Relative path to check
:param ignore_non_source_files: whether files that are not source files (according to the file masks
determined by the project's programming language) shall be ignored
:return: whether the path should be ignored
"""
abs_path = os.path.join(self.project_root, relative_path)
if not os.path.exists(abs_path):
raise FileNotFoundError(f"File {abs_path} not found, the ignore check cannot be performed")
# Check file extension if it's a file
is_file = os.path.isfile(abs_path)
if is_file and ignore_non_source_files:
fn_matcher = self.language.get_source_fn_matcher()
if not fn_matcher.is_relevant_filename(abs_path):
return True
# Create normalized path for consistent handling
rel_path = Path(relative_path)
# Check each part of the path against always fulfilled ignore conditions
dir_parts = rel_path.parts
if is_file:
dir_parts = dir_parts[:-1]
for part in dir_parts:
if not part: # Skip empty parts (e.g., from leading '/')
continue
if self._is_ignored_dirname(part):
return True
return match_path(relative_path, self.get_ignore_spec(), root_path=self.project_root)
def is_ignored_path(self, path: str | Path) -> bool:
"""
Checks whether the given path is ignored
:param path: the path to check, can be absolute or relative
"""
path = Path(path)
if path.is_absolute():
relative_path = path.relative_to(self.project_root)
else:
relative_path = path
# always ignore paths inside .git
if len(relative_path.parts) > 0 and relative_path.parts[0] == ".git":
return True
return match_path(str(relative_path), self.get_ignore_spec(), root_path=self.project_root)
def is_path_in_project(self, path: str | Path) -> bool:
"""
Checks if the given (absolute or relative) path is inside the project directory.
Note that even relative paths may be outside if they contain ".." or point to symlinks.
"""
path = Path(path)
_proj_root = Path(self.project_root)
if not path.is_absolute():
path = _proj_root / path
path = path.resolve()
return path.is_relative_to(_proj_root)
def validate_relative_path(self, relative_path: str) -> None:
"""
Validates that the given relative path is safe to read or edit,
meaning it's inside the project directory and is not ignored by git.
"""
if not self.is_path_in_project(relative_path):
raise ValueError(f"{relative_path=} points to path outside of the repository root; cannot access for safety reasons")
if self.is_ignored_path(relative_path):
raise ValueError(f"Path {relative_path} is ignored; cannot access for safety reasons")
def gather_source_files(self, relative_path: str = "") -> list[str]:
"""Retrieves relative paths of all source files, optionally limited to the given path
:param relative_path: if provided, restrict search to this path
"""
rel_file_paths = []
start_path = os.path.join(self.project_root, relative_path)
if not os.path.exists(start_path):
raise FileNotFoundError(f"Relative path {start_path} not found.")
if os.path.isfile(start_path):
return [relative_path]
else:
for root, dirs, files in os.walk(start_path, followlinks=True):
dirs[:] = [d for d in dirs if not self._is_ignored_relative_path(os.path.join(root, d))]
for file in files:
rel_file_path = os.path.relpath(os.path.join(root, file), start=self.project_root)
try:
if not self._is_ignored_relative_path(rel_file_path):
rel_file_paths.append(rel_file_path)
except FileNotFoundError:
log.warning(
f"File {rel_file_path} not found (possibly due it being a symlink), skipping it in request_parsed_files",
)
return rel_file_paths
def search_source_files_for_pattern(
self,
pattern: str,
relative_path: str = "",
context_lines_before: int = 0,
context_lines_after: int = 0,
paths_include_glob: str | None = None,
paths_exclude_glob: str | None = None,
) -> list[MatchedConsecutiveLines]:
"""
Search for a pattern across all (non-ignored) source files
:param pattern: Regular expression pattern to search for, either as a compiled Pattern or string
:param relative_path:
:param context_lines_before: Number of lines of context to include before each match
:param context_lines_after: Number of lines of context to include after each match
:param paths_include_glob: Glob pattern to filter which files to include in the search
:param paths_exclude_glob: Glob pattern to filter which files to exclude from the search. Takes precedence over paths_include_glob.
:return: List of matched consecutive lines with context
"""
relative_file_paths = self.gather_source_files(relative_path=relative_path)
return search_files(
relative_file_paths,
pattern,
root_path=self.project_root,
context_lines_before=context_lines_before,
context_lines_after=context_lines_after,
paths_include_glob=paths_include_glob,
paths_exclude_glob=paths_exclude_glob,
)
def retrieve_content_around_line(
self, relative_file_path: str, line: int, context_lines_before: int = 0, context_lines_after: int = 0
) -> MatchedConsecutiveLines:
"""
Retrieve the content of the given file around the given line.
:param relative_file_path: The relative path of the file to retrieve the content from
:param line: The line number to retrieve the content around
:param context_lines_before: The number of lines to retrieve before the given line
:param context_lines_after: The number of lines to retrieve after the given line
:return MatchedConsecutiveLines: A container with the desired lines.
"""
file_contents = self.read_file(relative_file_path)
return MatchedConsecutiveLines.from_file_contents(
file_contents,
line=line,
context_lines_before=context_lines_before,
context_lines_after=context_lines_after,
source_file_path=relative_file_path,
)
def create_language_server(
self,
log_level: int = logging.INFO,
ls_timeout: float | None = DEFAULT_TOOL_TIMEOUT - 5,
trace_lsp_communication: bool = False,
) -> SolidLanguageServer:
"""
Create a language server for a project. Note that you will have to start it
before performing any LS operations.
:param project: either a path to the project root or a ProjectConfig instance.
If no project.yml is found, the default project configuration will be used.
:param log_level: the log level for the language server
:param ls_timeout: the timeout for the language server
:param trace_lsp_communication: whether to trace LSP communication
:return: the language server
"""
ls_config = LanguageServerConfig(
code_language=self.language,
ignored_paths=self._ignored_patterns,
trace_lsp_communication=trace_lsp_communication,
)
ls_logger = LanguageServerLogger(log_level=log_level)
log.info(f"Creating language server instance for {self.project_root}.")
return SolidLanguageServer.create(
ls_config,
ls_logger,
self.project_root,
timeout=ls_timeout,
)
@@ -0,0 +1,15 @@
description: JetBrains tools replace language server-based tools
prompt: |
You have access to the very powerful JetBrains tools for symbolic operations:
* `jet_brains_find_symbol` replaces `find_symbol`
* `jet_brains_find_referencing_symbols` replaces `find_referencing_symbols`
* `jet_brains_get_symbols_overview` replaces `get_symbols_overview`
excluded_tools:
- find_symbol
- find_referencing_symbols
- get_symbols_overview
- restart_language_server
included_optional_tools:
- jet_brains_find_symbol
- jet_brains_find_referencing_symbols
- jet_brains_get_symbols_overview
@@ -23,7 +23,7 @@ web_dashboard: True
web_dashboard_open_on_launch: True
# whether to open a browser window with the web dashboard when Serena starts (provided that web_dashboard
# is enabled). If set to False, you can still open the dashboard manually by navigating to
# http://localhost:24282/dashboard/ in your web browser.
# http://localhost:24282/dashboard/ in your web browser (24282 = 0x5EDA, SErena DAshboard).
# If you have multiple instances running, a higher port will be used; try port 24283, 24284, etc.
log_level: 20
@@ -36,6 +36,18 @@ trace_lsp_communication: False
tool_timeout: 240
# timeout, in seconds, after which tool executions are terminated
excluded_tools: []
# list of tools to be globally excluded
included_optional_tools: []
# list of optional tools (which are disabled by default) to be included
jetbrains: False
# whether to enable JetBrains mode and use tools based on the Serena JetBrains IDE plugin
# instead of language server-based tools
# NOTE: The plugin is yet unreleased. This is for Serena developers only.
# MANAGED BY SERENA, KEEP AT THE BOTTOM OF THE YAML AND DON'T EDIT WITHOUT NEED
# The list of registered projects.
# To add a project, within a chat, simply ask Serena to "activate the project /path/to/project" or,
+169 -452
View File
@@ -1,11 +1,10 @@
import json
import logging
import os
from collections.abc import Iterable, Iterator, Reversible, Sequence
from contextlib import contextmanager
from dataclasses import asdict, dataclass, field
from difflib import SequenceMatcher
from typing import TYPE_CHECKING, Any, Literal, NamedTuple, Self, Union
from abc import ABC, abstractmethod
from collections.abc import Iterator, Sequence
from dataclasses import asdict, dataclass
from typing import TYPE_CHECKING, Any, Self, Union
from sensai.util.string import ToStringMixin
@@ -13,159 +12,16 @@ from solidlsp import SolidLanguageServer
from solidlsp.ls import ReferenceInSymbol as LSPReferenceInSymbol
from solidlsp.ls_types import Position, SymbolKind, UnifiedSymbolInformation
from .project import Project
if TYPE_CHECKING:
from .agent import SerenaAgent
log = logging.getLogger(__name__)
class LineChange(NamedTuple):
"""Represents a change to a specific line or range of lines."""
operation: Literal["insert", "delete", "replace"]
original_start: int
original_end: int
modified_start: int
modified_end: int
original_lines: list[str]
modified_lines: list[str]
@dataclass
class CodeDiff:
"""
Represents the difference between original and modified code.
Provides object-oriented access to diff information including line numbers.
"""
relative_path: str
original_content: str
modified_content: str
_line_changes: list[LineChange] = field(init=False)
def __post_init__(self) -> None:
"""Compute the diff using difflib's SequenceMatcher."""
original_lines = self.original_content.splitlines(keepends=True)
modified_lines = self.modified_content.splitlines(keepends=True)
matcher = SequenceMatcher(None, original_lines, modified_lines)
self._line_changes = []
for tag, orig_start, orig_end, mod_start, mod_end in matcher.get_opcodes():
if tag == "equal":
continue
if tag == "insert":
self._line_changes.append(
LineChange(
operation="insert",
original_start=orig_start,
original_end=orig_start,
modified_start=mod_start,
modified_end=mod_end,
original_lines=[],
modified_lines=modified_lines[mod_start:mod_end],
)
)
elif tag == "delete":
self._line_changes.append(
LineChange(
operation="delete",
original_start=orig_start,
original_end=orig_end,
modified_start=mod_start,
modified_end=mod_start,
original_lines=original_lines[orig_start:orig_end],
modified_lines=[],
)
)
elif tag == "replace":
self._line_changes.append(
LineChange(
operation="replace",
original_start=orig_start,
original_end=orig_end,
modified_start=mod_start,
modified_end=mod_end,
original_lines=original_lines[orig_start:orig_end],
modified_lines=modified_lines[mod_start:mod_end],
)
)
@property
def line_changes(self) -> list[LineChange]:
"""Get all line changes in the diff."""
return self._line_changes
@property
def has_changes(self) -> bool:
"""Check if there are any changes."""
return len(self._line_changes) > 0
@property
def added_lines(self) -> list[tuple[int, str]]:
"""Get all added lines with their line numbers (0-based) in the modified file."""
result = []
for change in self._line_changes:
if change.operation in ("insert", "replace"):
for i, line in enumerate(change.modified_lines):
result.append((change.modified_start + i, line))
return result
@property
def deleted_lines(self) -> list[tuple[int, str]]:
"""Get all deleted lines with their line numbers (0-based) in the original file."""
result = []
for change in self._line_changes:
if change.operation in ("delete", "replace"):
for i, line in enumerate(change.original_lines):
result.append((change.original_start + i, line))
return result
@property
def modified_line_numbers(self) -> list[int]:
"""Get all line numbers (0-based) that were modified in the modified file."""
line_nums: set[int] = set()
for change in self._line_changes:
if change.operation in ("insert", "replace"):
line_nums.update(range(change.modified_start, change.modified_end))
return sorted(line_nums)
@property
def affected_original_line_numbers(self) -> list[int]:
"""Get all line numbers (0-based) that were affected in the original file."""
line_nums: set[int] = set()
for change in self._line_changes:
if change.operation in ("delete", "replace"):
line_nums.update(range(change.original_start, change.original_end))
return sorted(line_nums)
def get_unified_diff(self, context_lines: int = 3) -> str:
"""Get the unified diff as a string."""
import difflib
original_lines = self.original_content.splitlines(keepends=True)
modified_lines = self.modified_content.splitlines(keepends=True)
diff = difflib.unified_diff(
original_lines, modified_lines, fromfile=f"a/{self.relative_path}", tofile=f"b/{self.relative_path}", n=context_lines
)
return "".join(diff)
def get_context_diff(self, context_lines: int = 3) -> str:
"""Get the context diff as a string."""
import difflib
original_lines = self.original_content.splitlines(keepends=True)
modified_lines = self.modified_content.splitlines(keepends=True)
diff = difflib.context_diff(
original_lines, modified_lines, fromfile=f"a/{self.relative_path}", tofile=f"b/{self.relative_path}", n=context_lines
)
return "".join(diff)
@dataclass
class SymbolLocation:
class LanguageServerSymbolLocation:
"""
Represents the (start) location of a symbol identifier, which, within Serena, uniquely identifies the symbol.
"""
@@ -199,7 +55,64 @@ class SymbolLocation:
return self.relative_path is not None and self.line is not None and self.column is not None
class Symbol(ToStringMixin):
@dataclass
class PositionInFile:
"""
Represents a character position within a file
"""
line: int
"""
the 0-based line number in the file
"""
col: int
"""
the 0-based column
"""
def to_lsp_position(self) -> Position:
"""
Convert to LSP Position.
"""
return Position(line=self.line, character=self.col)
class Symbol(ABC):
@abstractmethod
def get_body_start_position(self) -> PositionInFile | None:
pass
@abstractmethod
def get_body_end_position(self) -> PositionInFile | None:
pass
def get_body_start_position_or_raise(self) -> PositionInFile:
"""
Get the start position of the symbol body, raising an error if it is not defined.
"""
pos = self.get_body_start_position()
if pos is None:
raise ValueError(f"Body start position is not defined for {self}")
return pos
def get_body_end_position_or_raise(self) -> PositionInFile:
"""
Get the end position of the symbol body, raising an error if it is not defined.
"""
pos = self.get_body_end_position()
if pos is None:
raise ValueError(f"Body end position is not defined for {self}")
return pos
@abstractmethod
def is_neighbouring_definition_separated_by_empty_line(self) -> bool:
"""
:return: whether a symbol definition of this symbol's kind is usually separated from the
previous/next definition by at least one empty line.
"""
class LanguageServerSymbol(Symbol, ToStringMixin):
_NAME_PATH_SEP = "/"
@staticmethod
@@ -214,7 +127,7 @@ class Symbol(ToStringMixin):
"""
assert name_path, "name_path must not be empty"
assert symbol_name_path_parts, "symbol_name_path_parts must not be empty"
name_path_sep = Symbol._NAME_PATH_SEP
name_path_sep = LanguageServerSymbol._NAME_PATH_SEP
is_absolute_pattern = name_path.startswith(name_path_sep)
pattern_parts = name_path.lstrip(name_path_sep).rstrip(name_path_sep).split(name_path_sep)
@@ -260,10 +173,6 @@ class Symbol(ToStringMixin):
return self.symbol_root["kind"]
def is_neighbouring_definition_separated_by_empty_line(self) -> bool:
"""
:return: whether a symbol definition of this symbol's kind is usually separated from the
previous/next definition by at least one empty line.
"""
return self.symbol_kind in (SymbolKind.Function, SymbolKind.Method, SymbolKind.Class, SymbolKind.Interface, SymbolKind.Struct)
@property
@@ -274,11 +183,11 @@ class Symbol(ToStringMixin):
return None
@property
def location(self) -> SymbolLocation:
def location(self) -> LanguageServerSymbolLocation:
"""
:return: the start location of the actual symbol identifier
"""
return SymbolLocation(relative_path=self.relative_path, line=self.line, column=self.column)
return LanguageServerSymbolLocation(relative_path=self.relative_path, line=self.line, column=self.column)
@property
def body_start_position(self) -> Position | None:
@@ -302,6 +211,18 @@ class Symbol(ToStringMixin):
return end_pos
return None
def get_body_start_position(self) -> PositionInFile | None:
start_pos = self.body_start_position
if start_pos is None:
return None
return PositionInFile(line=start_pos["line"], col=start_pos["character"])
def get_body_end_position(self) -> PositionInFile | None:
end_pos = self.body_end_position
if end_pos is None:
return None
return PositionInFile(line=end_pos["line"], col=end_pos["character"])
def get_body_line_numbers(self) -> tuple[int | None, int | None]:
start_pos = self.body_start_position
end_pos = self.body_end_position
@@ -408,18 +329,18 @@ class Symbol(ToStringMixin):
"""
result = []
def should_include(s: "Symbol") -> bool:
def should_include(s: "LanguageServerSymbol") -> bool:
if include_kinds is not None and s.symbol_kind not in include_kinds:
return False
if exclude_kinds is not None and s.symbol_kind in exclude_kinds:
return False
return Symbol.match_name_path(
return LanguageServerSymbol.match_name_path(
name_path=name_path,
symbol_name_path_parts=s.get_name_path_parts(),
substring_matching=substring_matching,
)
def traverse(s: "Symbol") -> None:
def traverse(s: "LanguageServerSymbol") -> None:
if should_include(s):
result.append(s)
for c in s.iter_children():
@@ -490,24 +411,36 @@ class Symbol(ToStringMixin):
@dataclass
class ReferenceInSymbol(ToStringMixin):
"""Same as the class of the same name in the language server, but using Serena's Symbol class.
Be careful to not confuse it with counterpart!
class ReferenceInLanguageServerSymbol(ToStringMixin):
"""
Represents the location of a reference to another symbol within a symbol/file.
The contained symbol is the symbol within which the reference is located,
not the symbol that is referenced.
"""
symbol: Symbol
symbol: LanguageServerSymbol
"""
the symbol within which the reference is located
"""
line: int
"""
the line number in which the reference is located (0-based)
"""
character: int
"""
the column number in which the reference is located (0-based)
"""
@classmethod
def from_lsp_reference(cls, reference: LSPReferenceInSymbol) -> Self:
return cls(symbol=LanguageServerSymbol(reference.symbol), line=reference.line, character=reference.character)
def get_relative_path(self) -> str | None:
return self.symbol.location.relative_path
@classmethod
def from_lsp_reference(cls, reference: LSPReferenceInSymbol) -> Self:
return cls(symbol=Symbol(reference.symbol), line=reference.line, character=reference.character)
class SymbolManager:
class LanguageServerSymbolRetriever:
def __init__(self, lang_server: SolidLanguageServer, agent: Union["SerenaAgent", None] = None) -> None:
"""
:param lang_server: the language server to use for symbol retrieval as well as editing operations.
@@ -524,6 +457,9 @@ class SymbolManager:
"""
self._lang_server = lang_server
def get_language_server(self) -> SolidLanguageServer:
return self._lang_server
def find_by_name(
self,
name_path: str,
@@ -532,33 +468,33 @@ class SymbolManager:
exclude_kinds: Sequence[SymbolKind] | None = None,
substring_matching: bool = False,
within_relative_path: str | None = None,
) -> list[Symbol]:
) -> list[LanguageServerSymbol]:
"""
Find all symbols that match the given name. See docstring of `Symbol.find` for more details.
The only parameter not mentioned there is `within_relative_path`, which can be used to restrict the search
to symbols within a specific file or directory.
"""
symbols: list[Symbol] = []
symbols: list[LanguageServerSymbol] = []
symbol_roots = self._lang_server.request_full_symbol_tree(within_relative_path=within_relative_path, include_body=include_body)
for root in symbol_roots:
symbols.extend(
Symbol(root).find(
LanguageServerSymbol(root).find(
name_path, include_kinds=include_kinds, exclude_kinds=exclude_kinds, substring_matching=substring_matching
)
)
return symbols
def get_document_symbols(self, relative_path: str) -> list[Symbol]:
def get_document_symbols(self, relative_path: str) -> list[LanguageServerSymbol]:
symbol_dicts, roots = self._lang_server.request_document_symbols(relative_path, include_body=False)
symbols = [Symbol(s) for s in symbol_dicts]
symbols = [LanguageServerSymbol(s) for s in symbol_dicts]
return symbols
def find_by_location(self, location: SymbolLocation) -> Symbol | None:
def find_by_location(self, location: LanguageServerSymbolLocation) -> LanguageServerSymbol | None:
if location.relative_path is None:
return None
symbol_dicts, roots = self._lang_server.request_document_symbols(location.relative_path, include_body=False)
for symbol_dict in symbol_dicts:
symbol = Symbol(symbol_dict)
symbol = LanguageServerSymbol(symbol_dict)
if symbol.location == location:
return symbol
return None
@@ -570,7 +506,7 @@ class SymbolManager:
include_body: bool = False,
include_kinds: Sequence[SymbolKind] | None = None,
exclude_kinds: Sequence[SymbolKind] | None = None,
) -> list[ReferenceInSymbol]:
) -> list[ReferenceInLanguageServerSymbol]:
"""
Find all symbols that reference the symbol with the given name.
If multiple symbols fit the name (e.g. for variables that are overwritten), will use the first one.
@@ -600,11 +536,11 @@ class SymbolManager:
def find_referencing_symbols_by_location(
self,
symbol_location: SymbolLocation,
symbol_location: LanguageServerSymbolLocation,
include_body: bool = False,
include_kinds: Sequence[SymbolKind] | None = None,
exclude_kinds: Sequence[SymbolKind] | None = None,
) -> list[ReferenceInSymbol]:
) -> list[ReferenceInLanguageServerSymbol]:
"""
Find all symbols that reference the symbol at the given location.
@@ -641,284 +577,65 @@ class SymbolManager:
if exclude_kinds is not None:
references = [s for s in references if s.symbol["kind"] not in exclude_kinds]
return [ReferenceInSymbol.from_lsp_reference(r) for r in references]
return [ReferenceInLanguageServerSymbol.from_lsp_reference(r) for r in references]
@contextmanager
def _edited_file(self, relative_path: str) -> Iterator[None]:
with self._lang_server.open_file(relative_path) as file_buffer:
yield
root_path = self._lang_server.language_server.repository_root_path
abs_path = os.path.join(root_path, relative_path)
with open(abs_path, "w", encoding="utf-8") as f:
f.write(file_buffer.contents)
if self.agent is not None:
self.agent.mark_file_modified(relative_path)
@dataclass
class SymbolOverviewElement:
name_path: str
kind: int
@contextmanager
def _edited_symbol_location(self, location: SymbolLocation) -> Iterator[Symbol]:
def get_symbol_overview(self, relative_path: str) -> dict[str, list[SymbolOverviewElement]]:
path_to_symbol_infos = self._lang_server.request_overview(relative_path)
result = {}
for file_path, symbols in path_to_symbol_infos.items():
# TODO: maybe include not just top-level symbols? We could filter by kind to exclude variables
# The language server methods would need to be adjusted for this.
result[file_path] = [self.SymbolOverviewElement(name_path=symbol[0], kind=int(symbol[1])) for symbol in symbols]
return result
class JetBrainsSymbol(Symbol):
def __init__(self, symbol_dict: dict, project: Project) -> None:
"""
Context manager for locating and editing a symbol in a file.
:param symbol_dict: dictionary as returned by the JetBrains plugin client.
"""
symbol = self.find_by_location(location)
if symbol is None:
raise ValueError("Symbol not found/has no defined location within a file")
assert location.relative_path is not None
with self._edited_file(location.relative_path):
yield symbol
self._project = project
self._dict = symbol_dict
self._cached_file_content: str | None = None
self._cached_body_start_position: PositionInFile | None = None
self._cached_body_end_position: PositionInFile | None = None
def _get_code_file_content(self, relative_path: str) -> str:
"""Get the content of a file using the language server."""
return self._lang_server.language_server.retrieve_full_file_content(relative_path)
def get_relative_path(self) -> str:
return self._dict["relative_path"]
def replace_body(self, name_path: str, relative_file_path: str, body: str, *, use_same_indentation: bool = True) -> None:
"""
Replace the body of the symbol with the given name_path in the given file.
def get_file_content(self) -> str:
if self._cached_file_content is None:
path = os.path.join(self._project.project_root, self.get_relative_path())
with open(path, encoding=self._project.project_config.encoding) as f:
self._cached_file_content = f.read()
return self._cached_file_content
:param name_path: the name path of the symbol to replace.
:param relative_file_path: the relative path of the file in which the symbol is defined.
:param body: the new body
:param use_same_indentation: whether to use the same indentation as the original body. This means that
the user doesn't have to provide the correct indentation, but can just write the body.
"""
symbol_candidates = self.find_by_name(name_path, within_relative_path=relative_file_path)
if len(symbol_candidates) == 0:
raise ValueError(f"No symbol with name {name_path} found in file {relative_file_path}")
if len(symbol_candidates) > 1:
raise ValueError(
f"Found multiple {len(symbol_candidates)} symbols with name {name_path} in file {relative_file_path}. "
"Will not replace the body of any of them, but you can use `replace_body_at_location`, the replace lines tool or other editing "
"tools to perform your edits. Their locations are: \n "
+ json.dumps([s.location.to_dict() for s in symbol_candidates], indent=2)
)
symbol = symbol_candidates[0]
return self.replace_body_at_location(symbol.location, body, use_same_indentation=use_same_indentation)
def is_position_in_file_available(self) -> bool:
return "text_range" in self._dict
def replace_body_at_location(self, location: SymbolLocation, body: str, *, use_same_indentation: bool = True) -> None:
"""
Replace the body of the symbol at the given location with the given body
def get_body_start_position(self) -> PositionInFile | None:
if not self.is_position_in_file_available():
return None
if self._cached_body_start_position is None:
pos = self._dict["text_range"]["start_pos"]
line, col = pos["line"], pos["col"]
self._cached_body_start_position = PositionInFile(line=line, col=col)
return self._cached_body_start_position
:param location: the location of the symbol to replace.
:param body: the new body
:param use_same_indentation: whether to use the same indentation as the original body. This means that
the user doesn't have to provide the correct indentation, but can just write the body.
"""
with self._edited_symbol_location(location) as symbol:
assert location.relative_path is not None
start_pos = symbol.body_start_position
end_pos = symbol.body_end_position
if start_pos is None or end_pos is None:
raise ValueError(f"Symbol at {location} does not have a defined body range.")
start_line, start_col = start_pos["line"], start_pos["character"]
def get_body_end_position(self) -> PositionInFile | None:
if not self.is_position_in_file_available():
return None
if self._cached_body_end_position is None:
pos = self._dict["text_range"]["end_pos"]
line, col = pos["line"], pos["col"]
self._cached_body_end_position = PositionInFile(line=line, col=col)
return self._cached_body_end_position
if use_same_indentation:
indent = " " * start_col
body_lines = body.splitlines()
body = body_lines[0] + "\n" + "\n".join(indent + line for line in body_lines[1:])
# make sure the replacement adds no additional newlines (before or after) - all newlines
# and whitespace before/after should remain the same, so we strip it entirely
body = body.strip()
self._lang_server.delete_text_between_positions(location.relative_path, start_pos, end_pos)
self._lang_server.insert_text_at_position(location.relative_path, start_line, start_col, body)
@staticmethod
def _count_leading_newlines(text: Iterable) -> int:
cnt = 0
for c in text:
if c == "\n":
cnt += 1
elif c == "\r":
continue
else:
break
return cnt
@classmethod
def _count_trailing_newlines(cls, text: Reversible) -> int:
return cls._count_leading_newlines(reversed(text))
def insert_after_symbol(self, name_path: str, relative_file_path: str, body: str, *, use_same_indentation: bool = True) -> None:
"""
Inserts content after the symbol with the given name in the given file.
"""
symbol_candidates = self.find_by_name(name_path, within_relative_path=relative_file_path)
if len(symbol_candidates) == 0:
raise ValueError(f"No symbol with name {name_path} found in file {relative_file_path}")
if len(symbol_candidates) > 1:
raise ValueError(
f"Found multiple {len(symbol_candidates)} symbols with name {name_path} in file {relative_file_path}. "
f"May be an overwritten variable, in which case you can ignore this error. Proceeding with the last one. "
f"Found symbols at locations: \n" + json.dumps([s.location.to_dict() for s in symbol_candidates], indent=2)
)
symbol = symbol_candidates[-1]
return self.insert_after_symbol_at_location(symbol.location, body, use_same_indentation=use_same_indentation)
def insert_after_symbol_at_location(self, location: SymbolLocation, body: str, *, use_same_indentation: bool = True) -> None:
"""
Appends content after the given symbol
:param location: the location of the symbol after which to add new lines
:param body: the body of the entity to append
"""
# make sure body always ends with at least one newline
if not body.endswith("\n"):
body += "\n"
assert location.relative_path is not None
# Find the symbol to get its end position
symbol = self.find_by_location(location)
if symbol is None:
raise ValueError("Symbol not found/has no defined location within a file")
pos = symbol.body_end_position
if pos is None:
raise ValueError(f"Symbol at {location} does not have a defined end position.")
# start at the beginning of the next line
col = 0
line = pos["line"] + 1
# make sure a suitable number of leading empty lines is used (at least 0/1 depending on the symbol type,
# otherwise as many as the caller wanted to insert)
original_leading_newlines = self._count_leading_newlines(body)
body = body.lstrip("\r\n")
min_empty_lines = 0
if symbol.is_neighbouring_definition_separated_by_empty_line():
min_empty_lines = 1
num_leading_empty_lines = max(min_empty_lines, original_leading_newlines)
if num_leading_empty_lines:
body = ("\n" * num_leading_empty_lines) + body
# make sure the one line break succeeding the original symbol, which we repurposed as prefix via
# `line += 1`, is replaced
body = body.rstrip("\r\n") + "\n"
if use_same_indentation:
symbol_start_pos = symbol.body_start_position
assert symbol_start_pos is not None, f"Symbol at {location=} does not have a defined start position."
symbol_identifier_col = symbol_start_pos["character"]
indent = " " * (symbol_identifier_col)
body = "\n".join(indent + line for line in body.splitlines())
# IMPORTANT: without this, the insertion does the wrong thing. See implementation of insert_text_at_position in TextUtils,
# it is somewhat counterintuitive (never inserts whitespace)
# I am not 100% sure whether col=0 is always the best choice here.
#
# Without col=0, inserting after dataclass_instance in variables.py:
# > dataclass_instance = VariableDataclass(id=1, name="Test")
# > test test
# > dataclass_instancetest test
# > second line
# > .status = "active" # Reassign dataclass field
#
# With col=0:
# > dataclass_instance = VariableDataclass(id=1, name="Test")
# > test test
# > second line
# > dataclass_instance.status = "active" # Reassign dataclass field
with self._edited_symbol_location(location):
self._lang_server.insert_text_at_position(location.relative_path, line=line, column=col, text_to_be_inserted=body)
def insert_before_symbol(self, name_path: str, relative_file_path: str, body: str, *, use_same_indentation: bool = True) -> None:
"""
Inserts content before the symbol with the given name in the given file.
"""
symbol_candidates = self.find_by_name(name_path, within_relative_path=relative_file_path)
if len(symbol_candidates) == 0:
raise ValueError(f"No symbol with name {name_path} found in file {relative_file_path}")
if len(symbol_candidates) > 1:
raise ValueError(
f"Found multiple {len(symbol_candidates)} symbols with name {name_path} in file {relative_file_path}. "
f"May be an overwritten variable, in which case you can ignore this error. Proceeding with the first one. "
f"Found symbols at locations: \n" + json.dumps([s.location.to_dict() for s in symbol_candidates], indent=2)
)
symbol = symbol_candidates[0]
self.insert_before_symbol_at_location(symbol.location, body, use_same_indentation=use_same_indentation)
def insert_before_symbol_at_location(self, location: SymbolLocation, body: str, *, use_same_indentation: bool = True) -> None:
"""
Inserts content before the given symbol
:param location: the location of the symbol before which to add new lines
:param body: the body of the entity to insert
"""
with self._edited_symbol_location(location) as symbol:
symbol_start_pos = symbol.body_start_position
if symbol_start_pos is None:
raise ValueError(f"Symbol at {location} does not have a defined start position.")
if use_same_indentation:
indent = " " * (symbol_start_pos["character"])
body = "\n".join(indent + line for line in body.splitlines())
# insert position is the start of line where the symbol is defined
line = symbol_start_pos["line"]
col = 0
original_trailing_empty_lines = self._count_trailing_newlines(body) - 1
# ensure eol is present at end
body = body.rstrip() + "\n"
# add suitable number of trailing empty lines after the body (at least 0/1 depending on the symbol type,
# otherwise as many as the caller wanted to insert)
min_trailing_empty_lines = 0
if symbol.is_neighbouring_definition_separated_by_empty_line():
min_trailing_empty_lines = 1
num_trailing_newlines = max(min_trailing_empty_lines, original_trailing_empty_lines)
body += "\n" * num_trailing_newlines
assert location.relative_path is not None
self._lang_server.insert_text_at_position(location.relative_path, line=line, column=col, text_to_be_inserted=body)
def insert_at_line(self, relative_path: str, line: int, content: str) -> None:
"""
Inserts content at the given line in the given file.
:param line: the 0-based index of the line to insert content at
:param content: the content to insert
"""
with self._edited_file(relative_path):
self._lang_server.insert_text_at_position(relative_path, line, 0, content)
def delete_lines(self, relative_path: str, start_line: int, end_line: int) -> None:
"""
Deletes lines in the given file.
:param start_line: the 0-based index of the first line to delete (inclusive)
:param end_line: the 0-based index of the last line to delete (inclusive)
"""
start_col = 0
end_line_for_delete = end_line + 1
end_col = 0
with self._edited_file(relative_path):
start_pos = Position(line=start_line, character=start_col)
end_pos = Position(line=end_line_for_delete, character=end_col)
self._lang_server.delete_text_between_positions(relative_path, start_pos, end_pos)
def delete_symbol_at_location(self, location: SymbolLocation) -> None:
"""
Deletes the symbol at the given location.
"""
with self._edited_symbol_location(location) as symbol:
assert location.relative_path is not None
assert symbol.body_start_position is not None
assert symbol.body_end_position is not None
self._lang_server.delete_text_between_positions(location.relative_path, symbol.body_start_position, symbol.body_end_position)
def delete_symbol(self, name_path: str, relative_file_path: str) -> None:
"""
Deletes the symbol with the given name in the given file.
"""
symbol_candidates = self.find_by_name(name_path, within_relative_path=relative_file_path)
if len(symbol_candidates) == 0:
raise ValueError(f"No symbol with name {name_path} found in file {relative_file_path}")
if len(symbol_candidates) > 1:
raise ValueError(
f"Found multiple {len(symbol_candidates)} symbols with name {name_path} in file {relative_file_path}. "
"Will not delete any of them, but you can use `delete_symbol_at_location` or a corresponding tool to perform your edits. "
"Their locations are: \n " + json.dumps([s.location.to_dict() for s in symbol_candidates], indent=2)
)
symbol = symbol_candidates[0]
self.delete_symbol_at_location(symbol.location)
def is_neighbouring_definition_separated_by_empty_line(self) -> bool:
# NOTE: Symbol types cannot really be differentiated, because types are not handled in a language-agnostic way.
return False
+2 -1
View File
@@ -1,8 +1,9 @@
# ruff: noqa
from .tools_base import *
from .file_tools import *
from .ls_tools import *
from .symbol_tools import *
from .memory_tools import *
from .cmd_tools import *
from .config_tools import *
from .workflow_tools import *
from .jetbrains_tools import *
+17 -15
View File
@@ -39,9 +39,9 @@ class ReadFileTool(Tool):
required for the task.
:return: the full text of the file at the given relative path
"""
self.agent.validate_relative_path(relative_path)
self.project.validate_relative_path(relative_path)
result = self.language_server.retrieve_full_file_content(relative_path)
result = self.project.read_file(relative_path)
result_lines = result.splitlines()
if end_line is None:
result_lines = result_lines[start_line:]
@@ -73,7 +73,7 @@ class CreateTextFileTool(Tool, ToolMarkerCanEdit):
:param content: the (utf-8-encoded) content to write to the file
:return: a message indicating success or failure
"""
self.agent.validate_relative_path(relative_path)
self.project.validate_relative_path(relative_path)
abs_path = (Path(self.get_project_root()) / relative_path).resolve()
will_overwrite_existing = abs_path.exists()
@@ -102,14 +102,14 @@ class ListDirTool(Tool):
required for the task.
:return: a JSON object with the names of directories and files within the given directory
"""
self.agent.validate_relative_path(relative_path)
self.project.validate_relative_path(relative_path)
dirs, files = scan_directory(
os.path.join(self.get_project_root(), relative_path),
relative_to=self.get_project_root(),
recursive=recursive,
is_ignored_dir=self.agent.path_is_gitignored,
is_ignored_file=self.agent.path_is_gitignored,
is_ignored_dir=self.project.is_ignored_path,
is_ignored_file=self.project.is_ignored_path,
)
result = json.dumps({"dirs": dirs, "files": files})
@@ -129,13 +129,13 @@ class FindFileTool(Tool):
:param relative_path: the relative path to the directory to search in; pass "." to scan the project root
:return: a JSON object with the list of matching files
"""
self.agent.validate_relative_path(relative_path)
self.project.validate_relative_path(relative_path)
dir_to_scan = os.path.join(self.get_project_root(), relative_path)
# find the files by ignoring everything that doesn't match
def is_ignored_file(abs_path: str) -> bool:
if self.agent.path_is_gitignored(abs_path):
if self.project.is_ignored_path(abs_path):
return True
filename = os.path.basename(abs_path)
return not fnmatch(filename, file_mask)
@@ -143,7 +143,7 @@ class FindFileTool(Tool):
dirs, files = scan_directory(
path=dir_to_scan,
recursive=True,
is_ignored_dir=self.agent.path_is_gitignored,
is_ignored_dir=self.project.is_ignored_path,
is_ignored_file=is_ignored_file,
relative_to=self.get_project_root(),
)
@@ -185,7 +185,7 @@ class ReplaceRegexTool(Tool, ToolMarkerCanEdit):
If this is set to False and the regex matches multiple occurrences, an error will be returned
(and you may retry with a revised, more specific regex).
"""
self.agent.validate_relative_path(relative_path)
self.project.validate_relative_path(relative_path)
with EditedFileContext(relative_path, self.agent) as context:
original_content = context.get_original_content()
updated_content, n = re.subn(regex, repl, original_content, flags=re.DOTALL | re.MULTILINE)
@@ -223,7 +223,8 @@ class DeleteLinesTool(Tool, ToolMarkerCanEdit):
if not self.lines_read.were_lines_read(relative_path, (start_line, end_line)):
read_lines_tool = self.agent.get_tool(ReadFileTool)
return f"Error: Must call `{read_lines_tool.get_name_from_cls()}` first to read exactly the affected lines."
self.symbol_manager.delete_lines(relative_path, start_line, end_line)
code_editor = self.create_code_editor()
code_editor.delete_lines(relative_path, start_line, end_line)
return SUCCESS_RESULT
@@ -281,7 +282,8 @@ class InsertAtLineTool(Tool, ToolMarkerCanEdit):
"""
if not content.endswith("\n"):
content += "\n"
self.symbol_manager.insert_at_line(relative_path, line, content)
code_editor = self.create_code_editor()
code_editor.insert_at_line(relative_path, line, content)
return SUCCESS_RESULT
@@ -356,7 +358,7 @@ class SearchForPatternTool(Tool):
raise FileNotFoundError(f"Relative path {relative_path} does not exist.")
if restrict_search_to_code_files:
matches = self.language_server.search_files_for_pattern(
matches = self.project.search_source_files_for_pattern(
pattern=substring_pattern,
relative_path=relative_path,
context_lines_before=context_lines_before,
@@ -371,8 +373,8 @@ class SearchForPatternTool(Tool):
dirs, rel_paths_to_search = scan_directory(
path=abs_path,
recursive=True,
is_ignored_dir=self.agent.path_is_gitignored,
is_ignored_file=self.agent.path_is_gitignored,
is_ignored_dir=self.project.is_ignored_path,
is_ignored_file=self.project.is_ignored_path,
relative_to=self.get_project_root(),
)
# TODO (maybe): not super efficient to walk through the files again and filter if glob patterns are provided
+165
View File
@@ -0,0 +1,165 @@
"""
Client for the Serena JetBrains Plugin
"""
import json
from typing import Any, Optional, Self, TypeVar
import requests
T = TypeVar("T")
class SerenaClientError(Exception):
"""Base exception for Serena client errors."""
class ConnectionError(SerenaClientError):
"""Raised when connection to the service fails."""
class APIError(SerenaClientError):
"""Raised when the API returns an error response."""
class JetBrainsPluginClient:
"""
Python client for the Serena Backend Service.
Provides simple methods to interact with all available endpoints.
"""
def __init__(self, base_url: str = "http://localhost:8080", timeout: int = 30):
self.base_url = base_url.rstrip("/")
self.timeout = timeout
self.session = requests.Session()
self.session.headers.update({"Content-Type": "application/json", "Accept": "application/json"})
def _make_request(self, method: str, endpoint: str, data: Optional[dict] = None) -> dict[str, Any]:
url = f"{self.base_url}{endpoint}"
try:
if method.upper() == "GET":
response = self.session.get(url, timeout=self.timeout)
elif method.upper() == "POST":
json_data = json.dumps(data) if data else None
response = self.session.post(url, data=json_data, timeout=self.timeout)
else:
raise ValueError(f"Unsupported HTTP method: {method}")
response.raise_for_status()
# Try to parse JSON response
try:
return self._pythonify_response(response.json())
except json.JSONDecodeError:
# If response is not JSON, return raw text
return {"response": response.text}
except requests.exceptions.ConnectionError as e:
raise ConnectionError(f"Failed to connect to Serena service at {url}: {e}")
except requests.exceptions.Timeout as e:
raise ConnectionError(f"Request to {url} timed out: {e}")
except requests.exceptions.HTTPError:
raise APIError(f"API request failed with status {response.status_code}: {response.text}")
except requests.exceptions.RequestException as e:
raise SerenaClientError(f"Request failed: {e}")
@staticmethod
def _pythonify_response(response: T) -> T:
"""
Converts dictionary keys from camelCase to snake_case recursively.
:response: the response in which to convert keys (dictionary or list)
"""
to_snake_case = lambda s: "".join(["_" + c.lower() if c.isupper() else c for c in s])
def convert(x): # type: ignore
if isinstance(x, dict):
return {to_snake_case(k): convert(v) for k, v in x.items()}
elif isinstance(x, list):
return [convert(item) for item in x]
else:
return x
return convert(response)
def heartbeat(self) -> dict[str, Any]:
return self._make_request("GET", "/heartbeat")
def find_symbol(
self, name_path: str, relative_path: str | None = None, include_body: bool = False, depth: int = 0, include_location: bool = False
) -> dict[str, Any]:
"""
Find symbols by name.
:param name_path: the name path to match
:param relative_path: the relative path to which to restrict the search
:param include_body: whether to include symbol body content
:param depth: depth of children to include (0 = no children)
:return: Dictionary containing 'symbols' list with matching symbols
"""
request_data = {
"namePath": name_path,
"relativePath": relative_path,
"includeBody": include_body,
"depth": depth,
"includeLocation": include_location,
}
return self._make_request("POST", "/findSymbol", request_data)
def find_references(self, name_path: str, relative_path: str) -> dict[str, Any]:
"""
Find references to a symbol.
:param name_path: the name path of the symbol
:param relative_path: the relative path
:return: dictionary containing 'symbols' list with symbol references
"""
request_data = {"namePath": name_path, "relativePath": relative_path}
return self._make_request("POST", "/findReferences", request_data)
def get_symbols_overview(self, relative_path: str) -> dict[str, Any]:
"""
:param relative_path: the relative path to a source file
"""
request_data = {"relativePath": relative_path}
return self._make_request("POST", "/getSymbolsOverview", request_data)
def is_service_available(self) -> bool:
try:
response = self.heartbeat()
return response.get("status") == "OK"
except (ConnectionError, APIError):
return False
def close(self) -> None:
self.session.close()
def __enter__(self) -> Self:
return self
def __exit__(self, exc_type, exc_val, exc_tb): # type: ignore
self.close()
if __name__ == "__main__":
with JetBrainsPluginClient() as client:
# check heartbeat
heartbeat_response = client.heartbeat()
print(f"Heartbeat: {heartbeat_response}")
# find symbol
symbols_response = client.find_symbol("DQN", include_body=False, depth=1)
symbols = symbols_response.get("symbols", [])
print(f"Found {len(symbols)} symbols")
from pprint import pprint
pprint(symbols_response)
# find references
if symbols:
first_symbol = symbols[0]
refs_response = client.find_references(name_path=first_symbol["name_path"], relative_path=first_symbol["relative_path"])
pprint(refs_response)
+126
View File
@@ -0,0 +1,126 @@
import json
from serena.tools import TOOL_DEFAULT_MAX_ANSWER_LENGTH, Tool, ToolMarkerOptional
from serena.tools.jetbrains_plugin_client import JetBrainsPluginClient
class JetBrainsFindSymbolTool(Tool, ToolMarkerOptional):
"""
Performs a global (or local) search for symbols with/containing a given name/substring (optionally filtered by type).
"""
def apply(
self,
name_path: str,
depth: int = 0,
relative_path: str | None = None,
include_body: bool = False,
max_answer_chars: int = TOOL_DEFAULT_MAX_ANSWER_LENGTH,
) -> str:
"""
Retrieves information on all symbols/code entities (classes, methods, etc.) based on the given `name_path`,
which represents a pattern for the symbol's path within the symbol tree of a single file.
The returned symbol location can be used for edits or further queries.
Specify `depth > 0` to retrieve children (e.g., methods of a class).
The matching behavior is determined by the structure of `name_path`, which can
either be a simple name (e.g. "method") or a name path like "class/method" (relative name path)
or "/class/method" (absolute name path).
Note that the name path is not a path in the file system but rather a path in the symbol tree
**within a single file**. Thus, file or directory names should never be included in the `name_path`.
For restricting the search to a single file or directory, pass the `relative_path` parameter.
The retrieved symbols' `name_path` attribute will always be composed of symbol names, never file
or directory names.
Key aspects of the name path matching behavior:
- The name of the retrieved symbols will match the last segment of `name_path`, while preceding segments
will restrict the search to symbols that have a desired sequence of ancestors.
- If there is no `/` in `name_path`, there is no restriction on the ancestor symbols.
For example, passing `method` will match against all symbols with name paths like `method`,
`class/method`, `class/nested_class/method`, etc.
- If `name_path` contains at least one `/`, the matching is restricted to symbols
with the respective ancestors. For example, passing `class/method` will match against
`class/method` as well as `nested_class/class/method` but not `other_class/method`.
- If `name_path` starts with a `/`, it will be treated as an absolute name path pattern, i.e.
all ancestors are provided and must match.
For example, passing `/class` will match only against top-level symbols named `class` but
will not match `nested_class/class`. Passing `/class/method` will match `class/method` but
not `outer_class/class/method`.
:param name_path: The name path pattern to search for, see above for details.
:param depth: Depth to retrieve descendants (e.g., 1 for class methods/attributes).
:param relative_path: Optional. Restrict search to this file or directory.
If None, searches entire codebase.
If a directory is passed, the search will be restricted to the files in that directory.
If a file is passed, the search will be restricted to that file.
If you have some knowledge about the codebase, you should use this parameter, as it will significantly
speed up the search as well as reduce the number of results.
:param include_body: If True, include the symbol's source code. Use judiciously.
:param max_answer_chars: max characters for the JSON result. If exceeded, no content is returned.
:return: JSON string: a list of symbols (with locations) matching the name.
"""
with JetBrainsPluginClient() as client:
response_dict = client.find_symbol(
name_path=name_path,
relative_path=relative_path,
depth=depth,
include_body=include_body,
)
result = json.dumps(response_dict)
return self._limit_length(result, max_answer_chars)
class JetBrainsFindReferencingSymbolsTool(Tool, ToolMarkerOptional):
"""
Finds symbols that reference the given symbol
"""
def apply(
self,
name_path: str,
relative_path: str,
max_answer_chars: int = TOOL_DEFAULT_MAX_ANSWER_LENGTH,
) -> str:
"""
Finds symbols that reference the symbol at the given `name_path`.
The result will contain metadata about the referencing symbols.
:param name_path: name path of the symbol for which to find references; matching logic as described in find symbol tool.
:param relative_path: the relative path to the file containing the symbol for which to find references.
Note that here you can't pass a directory but must pass a file.
:param max_answer_chars: max characters for the JSON result. If exceeded, no content is returned.
:return: a list of JSON objects with the symbols referencing the requested symbol
"""
with JetBrainsPluginClient() as client:
response_dict = client.find_references(
name_path=name_path,
relative_path=relative_path,
)
result = json.dumps(response_dict)
return self._limit_length(result, max_answer_chars)
class JetBrainsGetSymbolsOverviewTool(Tool, ToolMarkerOptional):
"""
Retrieves an overview of the top-level symbols within a specified file
"""
def apply(
self,
relative_path: str,
max_answer_chars: int = TOOL_DEFAULT_MAX_ANSWER_LENGTH,
) -> str:
"""
Gets an overview of the top-level symbols in the given file.
Calling this is often a good idea before more targeted reading, searching or editing operations on the code symbols.
:param relative_path: the relative path to the file to get the overview of
:param max_answer_chars: max characters for the JSON result. If exceeded, no content is returned.
:return: a JSON object containing the symbols
"""
with JetBrainsPluginClient() as client:
response_dict = client.get_symbols_overview(
relative_path=relative_path,
)
result = json.dumps(response_dict)
return self._limit_length(result, max_answer_chars)
@@ -2,6 +2,7 @@
Language server-related tools
"""
import dataclasses
import json
from collections.abc import Sequence
from copy import copy
@@ -60,14 +61,9 @@ class GetSymbolsOverviewTool(Tool):
(e.g. a subdirectory).
:return: a JSON object mapping relative paths of all contained files to info about top-level symbols in the file (name_path, kind).
"""
path_to_symbol_infos = self.language_server.request_overview(relative_path)
result = {}
for file_path, symbols in path_to_symbol_infos.items():
# TODO: maybe include not just top-level symbols? We could filter by kind to exclude variables
# The language server methods would need to be adjusted for this.
result[file_path] = [{"name_path": symbol[0], "kind": int(symbol[1])} for symbol in symbols]
result_json_str = json.dumps(result)
symbol_retriever = self.create_language_server_symbol_retriever()
result = symbol_retriever.get_symbol_overview(relative_path)
result_json_str = json.dumps({k: [dataclasses.asdict(i) for i in l] for k, l in result.items()})
return self._limit_length(result_json_str, max_answer_chars)
@@ -137,7 +133,8 @@ class FindSymbolTool(Tool):
"""
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
symbols = self.symbol_manager.find_by_name(
symbol_retriever = self.create_language_server_symbol_retriever()
symbols = symbol_retriever.find_by_name(
name_path,
include_body=include_body,
include_kinds=parsed_include_kinds,
@@ -180,7 +177,8 @@ class FindReferencingSymbolsTool(Tool):
include_body = False # It is probably never a good idea to include the body of the referencing symbols
parsed_include_kinds: Sequence[SymbolKind] | None = [SymbolKind(k) for k in include_kinds] if include_kinds else None
parsed_exclude_kinds: Sequence[SymbolKind] | None = [SymbolKind(k) for k in exclude_kinds] if exclude_kinds else None
references_in_symbols = self.symbol_manager.find_referencing_symbols(
symbol_retriever = self.create_language_server_symbol_retriever()
references_in_symbols = symbol_retriever.find_referencing_symbols(
name_path,
relative_file_path=relative_path,
include_body=include_body,
@@ -194,7 +192,7 @@ class FindReferencingSymbolsTool(Tool):
if not include_body:
ref_relative_path = ref.symbol.location.relative_path
assert ref_relative_path is not None, f"Referencing symbol {ref.symbol.name} has no relative path, this is likely a bug."
content_around_ref = self.language_server.retrieve_content_around_line(
content_around_ref = self.project.retrieve_content_around_line(
relative_file_path=ref_relative_path, line=ref.line, context_lines_before=1, context_lines_after=1
)
ref_dict["content_around_reference"] = content_around_ref.to_display_string()
@@ -222,11 +220,11 @@ class ReplaceSymbolBodyTool(Tool, ToolMarkerCanEdit):
:param body: the new symbol body. Important: Begin directly with the symbol definition and provide no
leading indentation for the first line (but do indent the rest of the body according to the context).
"""
self.symbol_manager.replace_body(
code_editor = self.create_code_editor()
code_editor.replace_body(
name_path,
relative_file_path=relative_path,
body=body,
use_same_indentation=False,
)
return SUCCESS_RESULT
@@ -251,7 +249,8 @@ class InsertAfterSymbolTool(Tool, ToolMarkerCanEdit):
:param body: the body/content to be inserted. The inserted code shall begin with the next line after
the symbol.
"""
self.symbol_manager.insert_after_symbol(name_path, relative_file_path=relative_path, body=body, use_same_indentation=False)
code_editor = self.create_code_editor()
code_editor.insert_after_symbol(name_path, relative_file_path=relative_path, body=body)
return SUCCESS_RESULT
@@ -275,5 +274,6 @@ class InsertBeforeSymbolTool(Tool, ToolMarkerCanEdit):
: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
"""
self.symbol_manager.insert_before_symbol(name_path, relative_file_path=relative_path, body=body, use_same_indentation=False)
code_editor = self.create_code_editor()
code_editor.insert_before_symbol(name_path, relative_file_path=relative_path, body=body)
return SUCCESS_RESULT
+66 -83
View File
@@ -1,9 +1,9 @@
import inspect
import os
import traceback
from abc import ABC, abstractmethod
from collections.abc import Callable, Generator, Iterable
from copy import copy
from abc import ABC
from collections.abc import Callable, Iterable
from dataclasses import dataclass
from types import TracebackType
from typing import TYPE_CHECKING, Any, Self, TypeVar
@@ -11,13 +11,15 @@ from mcp.server.fastmcp.utilities.func_metadata import FuncMetadata, func_metada
from sensai.util import logging
from sensai.util.string import dict_string
from serena.project import Project
from serena.prompt_factory import PromptFactory
from serena.symbol import SymbolManager
from serena.symbol import LanguageServerSymbolRetriever
from serena.util.class_decorators import singleton
from serena.util.inspection import iter_subclasses
from solidlsp import SolidLanguageServer
if TYPE_CHECKING:
from serena.agent import LinesRead, MemoriesManager, SerenaAgent
from serena.code_editor import CodeEditor
log = logging.getLogger(__name__)
T = TypeVar("T")
@@ -28,11 +30,6 @@ class Component(ABC):
def __init__(self, agent: "SerenaAgent"):
self.agent = agent
@property
def language_server(self) -> SolidLanguageServer:
assert self.agent.language_server is not None
return self.agent.language_server
def get_project_root(self) -> str:
"""
:return: the root directory of the active project, raises a ValueError if no active project configuration is set
@@ -48,10 +45,24 @@ class Component(ABC):
assert self.agent.memories_manager is not None
return self.agent.memories_manager
def create_language_server_symbol_retriever(self) -> LanguageServerSymbolRetriever:
if not self.agent.is_using_language_server():
raise Exception("Cannot create LanguageServerSymbolRetriever; agent is not in language server mode.")
language_server = self.agent.language_server
assert language_server is not None
return LanguageServerSymbolRetriever(language_server, agent=self.agent)
@property
def symbol_manager(self) -> SymbolManager:
assert self.agent.symbol_manager is not None
return self.agent.symbol_manager
def project(self) -> Project:
return self.agent.get_active_project_or_raise()
def create_code_editor(self) -> "CodeEditor":
from ..code_editor import JetBrainsCodeEditor, LanguageServerCodeEditor
if self.agent.is_using_language_server():
return LanguageServerCodeEditor(self.create_language_server_symbol_retriever(), agent=self.agent)
else:
return JetBrainsCodeEditor(project=self.project, agent=self.agent)
@property
def lines_read(self) -> "LinesRead":
@@ -72,31 +83,13 @@ class ToolMarkerDoesNotRequireActiveProject:
pass
class ToolInterface(ABC):
"""Protocol defining the complete interface that make_tool() expects from a tool."""
@abstractmethod
def get_name(self) -> str:
"""Get the tool name."""
...
@abstractmethod
def get_apply_docstring(self) -> str:
"""Get the docstring for the tool application, used by the MCP server."""
...
@abstractmethod
def get_apply_fn_metadata(self) -> FuncMetadata:
"""Get the metadata for the tool application function, used by the MCP server."""
...
@abstractmethod
def apply_ex(self, log_call: bool = True, catch_exceptions: bool = True, **kwargs: Any) -> str:
"""Apply the tool with logging and exception handling."""
...
class ToolMarkerOptional:
"""
Marker class for optional tools that are disabled by default.
"""
class Tool(Component, ToolInterface):
class Tool(Component):
# NOTE: each tool should implement the apply method, which is then used in
# the central method of the Tool class `apply_ex`.
# Failure to do so will result in a RuntimeError at tool execution time.
@@ -161,11 +154,11 @@ class Tool(Component, ToolInterface):
return docstring.strip()
def get_apply_docstring(self) -> str:
"""Get the docstring for the apply method (instance method implementing ToolProtocol)."""
"""Gets the docstring for the tool application, used by the MCP server."""
return self.get_apply_docstring_from_cls()
def get_apply_fn_metadata(self) -> FuncMetadata:
"""Get the metadata for the apply method (instance method implementing ToolProtocol)."""
"""Gets the metadata for the tool application function, used by the MCP server."""
return self.get_apply_fn_metadata_from_cls()
@classmethod
@@ -210,7 +203,7 @@ class Tool(Component, ToolInterface):
def apply_ex(self, log_call: bool = True, catch_exceptions: bool = True, **kwargs) -> str: # type: ignore
"""
Applies the tool with the given arguments
Applies the tool with logging and exception handling, using the given keyword arguments
"""
def task() -> str:
@@ -232,7 +225,7 @@ class Tool(Component, ToolInterface):
"Error: No active project. Ask to user to select a project from this list: "
+ f"{self.agent.serena_config.project_names}"
)
if not self.agent.is_language_server_running():
if self.agent.is_using_language_server() and not self.agent.is_language_server_running():
log.info("Language server is not running. Starting it ...")
self.agent.reset_language_server()
@@ -254,7 +247,8 @@ class Tool(Component, ToolInterface):
log.info(f"Result: {result}")
try:
self.language_server.save_cache()
if self.agent.language_server is not None:
self.agent.language_server.save_cache()
except Exception as e:
log.error(f"Error saving language server cache: {e}")
@@ -311,58 +305,44 @@ class EditedFileContext:
# If they do not, we may have to add a call to notify it.
@dataclass(kw_only=True)
class RegisteredTool:
tool_class: type[Tool]
is_optional: bool
tool_name: str
@singleton
class ToolRegistry:
_tool_dict: dict[str, type[Tool]] | None = None
"""maps tool name to the corresponding tool class"""
@staticmethod
def _iter_tool_classes() -> Generator[type[Tool], None, None]:
"""
Iterate over Tool subclasses.
"""
def __init__(self) -> None:
self._tool_dict: dict[str, RegisteredTool] = {}
for cls in iter_subclasses(Tool):
if cls.__module__.startswith("serena.tools"):
yield cls
if not cls.__module__.startswith("serena.tools"):
continue
is_optional = issubclass(cls, ToolMarkerOptional)
name = cls.get_name_from_cls()
if name in self._tool_dict:
raise ValueError(f"Duplicate tool name found: {name}. Tool classes must have unique names.")
self._tool_dict[name] = RegisteredTool(tool_class=cls, is_optional=is_optional, tool_name=name)
@classmethod
def _get_tool_dict(cls) -> dict[str, type[Tool]]:
if cls._tool_dict is None:
cls._tool_dict = {}
for tool_class in cls._iter_tool_classes():
name = tool_class.get_name_from_cls()
if name in cls._tool_dict:
raise ValueError(f"Duplicate tool name found: {name}. Tool classes must have unique names.")
cls._tool_dict[name] = tool_class
return cls._tool_dict
def get_tool_class_by_name(self, tool_name: str) -> type[Tool]:
return self._tool_dict[tool_name].tool_class
@classmethod
def get_tool_class_by_name(cls, tool_name: str) -> type[Tool]:
try:
return cls._get_tool_dict()[tool_name]
except KeyError as e:
available_tools = "\n".join(ToolRegistry.get_tool_names())
raise ValueError(f"Tool with name {tool_name} not found. Available tools:\n{available_tools}") from e
def get_all_tool_classes(self) -> list[type[Tool]]:
return list(t.tool_class for t in self._tool_dict.values())
@classmethod
def get_all_tool_classes(cls) -> list[type[Tool]]:
return list(cls._get_tool_dict().values())
def get_tool_names_default_enabled(self) -> list[str]:
"""
:return: the list of tool names that are enabled by default (i.e. non-optional tools).
"""
return [t.tool_name for t in self._tool_dict.values() if not t.is_optional]
@classmethod
def get_tool_names(cls) -> list[str]:
return list(cls._get_tool_dict().keys())
@classmethod
def tool_dict(cls) -> dict[str, type[Tool]]:
"""Maps tool name to the corresponding tool class"""
return copy(cls._get_tool_dict())
@classmethod
def print_tool_overview(cls, tools: Iterable[type[Tool] | Tool] | None = None) -> None:
def print_tool_overview(self, tools: Iterable[type[Tool] | Tool] | None = None) -> None:
"""
Print a summary of the tools. If no tools are passed, a summary of all tools is printed.
"""
if tools is None:
tools = cls._get_tool_dict().values()
tools = [tool.tool_class for tool in self._tool_dict.values() if not tool.is_optional]
tool_dict: dict[str, type[Tool] | Tool] = {}
for tool_class in tools:
@@ -370,3 +350,6 @@ class ToolRegistry:
for tool_name in sorted(tool_dict.keys()):
tool_class = tool_dict[tool_name]
print(f" * `{tool_name}`: {tool_class.get_tool_description().strip()}")
def is_valid_tool_name(self, tool_name: str) -> bool:
return tool_name in self._tool_dict
+1 -104
View File
@@ -16,9 +16,8 @@ from pathlib import Path, PurePath
from typing import Self, Union, cast
import pathspec
import tqdm
from serena.text_utils import MatchedConsecutiveLines, search_files
from serena.text_utils import MatchedConsecutiveLines
from serena.util.file_system import match_path
from solidlsp import ls_types
from solidlsp.ls_config import Language, LanguageServerConfig
@@ -710,26 +709,6 @@ class SolidLanguageServer(ABC):
return ret
def request_references_with_content(
self, relative_file_path: str, line: int, column: int, context_lines_before: int = 0, context_lines_after: int = 0
) -> list[MatchedConsecutiveLines]:
"""
Like request_references, but returns the content of the lines containing the references, not just the locations.
:param relative_file_path: The relative path of the file that has the symbol for which references should be looked up
:param line: The line number of the symbol
:param column: The column number of the symbol
:param context_lines_before: The number of lines to include in the context before the line containing the reference
:param context_lines_after: The number of lines to include in the context after the line containing the reference
:return: A list of MatchedConsecutiveLines objects, one for each reference.
"""
references = self.request_references(relative_file_path, line, column)
return [
self.retrieve_content_around_line(ref["relativePath"], ref["range"]["start"]["line"], context_lines_before, context_lines_after)
for ref in references
]
def retrieve_full_file_content(self, file_path: str) -> str:
"""
Retrieve the full content of the given file.
@@ -1223,70 +1202,6 @@ class SolidLanguageServer(ABC):
symbol_body = symbol_body[symbol_start_column:]
return symbol_body
def request_parsed_files(self, relative_path: str = "") -> list[str]:
"""Retrieves relative paths of all files analyzed by the Language Server.
:param relative_path: will only retrieve files that are subpaths of this.
"""
if not self.server_started:
self.logger.log(
"request_parsed_files called before Language Server started",
logging.ERROR,
)
raise LanguageServerException("Language Server not started")
rel_file_paths = []
start_path = os.path.join(self.repository_root_path, relative_path)
if not os.path.exists(start_path):
raise FileNotFoundError(f"Relative path {start_path} not found.")
if os.path.isfile(start_path):
return [relative_path]
else:
for root, dirs, files in os.walk(start_path, followlinks=True):
dirs[:] = [d for d in dirs if not self.is_ignored_path(os.path.join(root, d))]
for file in files:
rel_file_path = os.path.relpath(os.path.join(root, file), start=self.repository_root_path)
try:
if not self.is_ignored_path(rel_file_path):
rel_file_paths.append(rel_file_path)
except FileNotFoundError:
self.logger.log(
f"File {rel_file_path} not found (possibly due it being a symlink), skipping it in request_parsed_files",
logging.WARNING,
)
return rel_file_paths
def search_files_for_pattern(
self,
pattern: str,
relative_path: str = "",
context_lines_before: int = 0,
context_lines_after: int = 0,
paths_include_glob: str | None = None,
paths_exclude_glob: str | None = None,
) -> list[MatchedConsecutiveLines]:
"""
Search for a pattern across all files analyzed by the Language Server.
:param pattern: Regular expression pattern to search for, either as a compiled Pattern or string
:param relative_path:
:param context_lines_before: Number of lines of context to include before each match
:param context_lines_after: Number of lines of context to include after each match
:param paths_include_glob: Glob pattern to filter which files to include in the search
:param paths_exclude_glob: Glob pattern to filter which files to exclude from the search. Takes precedence over paths_include_glob.
:return: List of matched consecutive lines with context
"""
relative_file_paths = self.request_parsed_files(relative_path=relative_path)
return search_files(
relative_file_paths,
pattern,
file_reader=self.retrieve_full_file_content,
root_path=self.repository_root_path,
context_lines_before=context_lines_before,
context_lines_after=context_lines_after,
paths_include_glob=paths_include_glob,
paths_exclude_glob=paths_exclude_glob,
)
def request_referencing_symbols(
self,
relative_file_path: str,
@@ -1617,24 +1532,6 @@ class SolidLanguageServer(ABC):
"""
return Path(self.repository_root_path) / ".serena" / "cache" / self.language_id / "document_symbols_cache_v23-06-25.pkl"
def index_repository(self, progress_bar: bool = True, save_after_n_files: int = 10) -> None:
"""Will go through the entire repository and "index" all files, meaning save their symbols to the cache.
:param progress_bar: Whether to show a progress bar while indexing the repository.
:param save_after_n_files: How many files to process before saving a checkpoint of the cache.
"""
parsed_files = self.request_parsed_files()
files_processed = 0
pbar = tqdm.tqdm(parsed_files, disable=not progress_bar)
for relative_file_path in pbar:
pbar.set_description(f"Indexing ({os.path.basename(relative_file_path)})")
self.request_document_symbols(relative_file_path, include_body=False)
self.request_document_symbols(relative_file_path, include_body=True)
files_processed += 1
if files_processed % save_after_n_files == 0:
self.save_cache()
self.save_cache()
def save_cache(self):
with self._cache_lock:
if not self._cache_has_changed:
+37 -1
View File
@@ -4,6 +4,7 @@ from pathlib import Path
import pytest
from sensai.util.logging import configure
from serena.project import Project
from serena.util.file_system import GitignoreParser
from solidlsp.ls import SolidLanguageServer
from solidlsp.ls_config import Language, LanguageServerConfig
@@ -50,6 +51,11 @@ def create_default_ls(language: Language) -> SolidLanguageServer:
return create_ls(language, repo_path)
def create_default_project(language: Language) -> Project:
repo_path = str(get_repo_path(language))
return Project.load(repo_path)
@pytest.fixture(scope="session")
def repo_path(request: LanguageParamRequest) -> Path:
"""Get the repository path for a specific language.
@@ -73,7 +79,7 @@ def repo_path(request: LanguageParamRequest) -> Path:
@pytest.fixture(scope="session")
def language_server(request: LanguageParamRequest):
"""Create a SyncLanguageServer instance configured for the specified language.
"""Create a language server instance configured for the specified language.
This fixture requires a language parameter via pytest.mark.parametrize:
@@ -104,3 +110,33 @@ def language_server(request: LanguageParamRequest):
yield server
finally:
server.stop()
@pytest.fixture(scope="session")
def project(request: LanguageParamRequest):
"""Create a Project for the specified language.
This fixture requires a language parameter via pytest.mark.parametrize:
Example:
```
@pytest.mark.parametrize("project", [Language.PYTHON], indirect=True)
def test_python_project(project: Project) -> None:
# Use the Python project to test something
pass
```
You can also test multiple languages in a single test:
```
@pytest.mark.parametrize("project", [Language.PYTHON, Language.TYPESCRIPT], indirect=True)
def test_multiple_languages(project: SyncLanguageServer) -> None:
# This test will run once for each language
pass
```
"""
if not hasattr(request, "param"):
raise ValueError("Language parameter must be provided via pytest.mark.parametrize")
language = request.param
yield create_default_project(language)
+1 -1
View File
@@ -271,7 +271,7 @@ def is_test_mock_class(tool_class: type) -> bool:
)
@pytest.mark.parametrize("tool_class", ToolRegistry.get_all_tool_classes())
@pytest.mark.parametrize("tool_class", ToolRegistry().get_all_tool_classes())
def test_make_tool_all_tools(tool_class) -> None:
"""Test that make_tool works for all tools in the codebase."""
+2 -1
View File
@@ -6,7 +6,8 @@ import pytest
import test.solidlsp.clojure as clj
from serena.agent import SerenaAgent
from serena.config.serena_config import Project, ProjectConfig, SerenaConfig
from serena.config.serena_config import ProjectConfig, SerenaConfig
from serena.project import Project
from serena.tools import FindReferencingSymbolsTool, FindSymbolTool
from solidlsp.ls_config import Language
from test.conftest import get_repo_path
+3 -3
View File
@@ -1,6 +1,6 @@
import pytest
from src.serena.symbol import Symbol
from src.serena.symbol import LanguageServerSymbol
class TestSymbolNameMatching:
@@ -77,7 +77,7 @@ class TestSymbolNameMatching:
)
def test_match_simple_name(self, name_path_pattern, symbol_name_path_parts, is_substring_match, expected):
"""Tests matching for simple names (no '/' in pattern)."""
result = Symbol.match_name_path(name_path_pattern, symbol_name_path_parts, is_substring_match)
result = LanguageServerSymbol.match_name_path(name_path_pattern, symbol_name_path_parts, is_substring_match)
error_msg = self._create_assertion_error_message(name_path_pattern, symbol_name_path_parts, is_substring_match, expected, result)
assert result == expected, error_msg
@@ -157,6 +157,6 @@ class TestSymbolNameMatching:
)
def test_match_name_path_pattern_path_len_2(self, name_path_pattern, symbol_name_path_parts, is_substring_match, expected):
"""Tests matching for qualified names (e.g. 'module/class/func')."""
result = Symbol.match_name_path(name_path_pattern, symbol_name_path_parts, is_substring_match)
result = LanguageServerSymbol.match_name_path(name_path_pattern, symbol_name_path_parts, is_substring_match)
error_msg = self._create_assertion_error_message(name_path_pattern, symbol_name_path_parts, is_substring_match, expected, result)
assert result == expected, error_msg
+163 -15
View File
@@ -6,14 +6,16 @@ import time
from abc import abstractmethod
from collections.abc import Iterator
from contextlib import contextmanager
from dataclasses import dataclass, field
from difflib import SequenceMatcher
from pathlib import Path
from typing import Literal
from typing import Literal, NamedTuple
import pytest
from serena.symbol import CodeDiff
from serena.code_editor import CodeEditor, LanguageServerCodeEditor
from solidlsp.ls_config import Language
from src.serena.symbol import SymbolManager
from src.serena.symbol import LanguageServerSymbolRetriever
from test.conftest import create_ls, get_repo_path
pytestmark = pytest.mark.snapshot
@@ -21,6 +23,151 @@ pytestmark = pytest.mark.snapshot
log = logging.getLogger(__name__)
class LineChange(NamedTuple):
"""Represents a change to a specific line or range of lines."""
operation: Literal["insert", "delete", "replace"]
original_start: int
original_end: int
modified_start: int
modified_end: int
original_lines: list[str]
modified_lines: list[str]
@dataclass
class CodeDiff:
"""
Represents the difference between original and modified code.
Provides object-oriented access to diff information including line numbers.
"""
relative_path: str
original_content: str
modified_content: str
_line_changes: list[LineChange] = field(init=False)
def __post_init__(self) -> None:
"""Compute the diff using difflib's SequenceMatcher."""
original_lines = self.original_content.splitlines(keepends=True)
modified_lines = self.modified_content.splitlines(keepends=True)
matcher = SequenceMatcher(None, original_lines, modified_lines)
self._line_changes = []
for tag, orig_start, orig_end, mod_start, mod_end in matcher.get_opcodes():
if tag == "equal":
continue
if tag == "insert":
self._line_changes.append(
LineChange(
operation="insert",
original_start=orig_start,
original_end=orig_start,
modified_start=mod_start,
modified_end=mod_end,
original_lines=[],
modified_lines=modified_lines[mod_start:mod_end],
)
)
elif tag == "delete":
self._line_changes.append(
LineChange(
operation="delete",
original_start=orig_start,
original_end=orig_end,
modified_start=mod_start,
modified_end=mod_start,
original_lines=original_lines[orig_start:orig_end],
modified_lines=[],
)
)
elif tag == "replace":
self._line_changes.append(
LineChange(
operation="replace",
original_start=orig_start,
original_end=orig_end,
modified_start=mod_start,
modified_end=mod_end,
original_lines=original_lines[orig_start:orig_end],
modified_lines=modified_lines[mod_start:mod_end],
)
)
@property
def line_changes(self) -> list[LineChange]:
"""Get all line changes in the diff."""
return self._line_changes
@property
def has_changes(self) -> bool:
"""Check if there are any changes."""
return len(self._line_changes) > 0
@property
def added_lines(self) -> list[tuple[int, str]]:
"""Get all added lines with their line numbers (0-based) in the modified file."""
result = []
for change in self._line_changes:
if change.operation in ("insert", "replace"):
for i, line in enumerate(change.modified_lines):
result.append((change.modified_start + i, line))
return result
@property
def deleted_lines(self) -> list[tuple[int, str]]:
"""Get all deleted lines with their line numbers (0-based) in the original file."""
result = []
for change in self._line_changes:
if change.operation in ("delete", "replace"):
for i, line in enumerate(change.original_lines):
result.append((change.original_start + i, line))
return result
@property
def modified_line_numbers(self) -> list[int]:
"""Get all line numbers (0-based) that were modified in the modified file."""
line_nums: set[int] = set()
for change in self._line_changes:
if change.operation in ("insert", "replace"):
line_nums.update(range(change.modified_start, change.modified_end))
return sorted(line_nums)
@property
def affected_original_line_numbers(self) -> list[int]:
"""Get all line numbers (0-based) that were affected in the original file."""
line_nums: set[int] = set()
for change in self._line_changes:
if change.operation in ("delete", "replace"):
line_nums.update(range(change.original_start, change.original_end))
return sorted(line_nums)
def get_unified_diff(self, context_lines: int = 3) -> str:
"""Get the unified diff as a string."""
import difflib
original_lines = self.original_content.splitlines(keepends=True)
modified_lines = self.modified_content.splitlines(keepends=True)
diff = difflib.unified_diff(
original_lines, modified_lines, fromfile=f"a/{self.relative_path}", tofile=f"b/{self.relative_path}", n=context_lines
)
return "".join(diff)
def get_context_diff(self, context_lines: int = 3) -> str:
"""Get the context diff as a string."""
import difflib
original_lines = self.original_content.splitlines(keepends=True)
modified_lines = self.modified_content.splitlines(keepends=True)
diff = difflib.context_diff(
original_lines, modified_lines, fromfile=f"a/{self.relative_path}", tofile=f"b/{self.relative_path}", n=context_lines
)
return "".join(diff)
class EditingTest:
def __init__(self, language: Language, rel_path: str):
"""
@@ -33,7 +180,7 @@ class EditingTest:
self.repo_path: Path | None = None
@contextmanager
def _setup(self) -> Iterator[SymbolManager]:
def _setup(self) -> Iterator[LanguageServerSymbolRetriever]:
"""Context manager for setup/teardown with a temporary directory, providing the symbol manager."""
temp_dir = Path(tempfile.mkdtemp())
self.repo_path = temp_dir / self.original_repo_path.name
@@ -50,7 +197,7 @@ class EditingTest:
log.info(f"Starting language server for {self.language} {self.rel_path}")
language_server.start()
log.info(f"Language server started for {self.language} {self.rel_path}")
yield SymbolManager(lang_server=language_server)
yield LanguageServerSymbolRetriever(lang_server=language_server)
finally:
if language_server is not None and language_server.is_running():
log.info(f"Stopping language server for {self.language} {self.rel_path}")
@@ -74,15 +221,16 @@ class EditingTest:
return f.read()
def run_test(self, content_after_ground_truth: str) -> None:
with self._setup() as symbol_manager:
with self._setup() as symbol_retriever:
content_before = self._read_file(self.rel_path)
self._apply_edit(symbol_manager)
code_editor = LanguageServerCodeEditor(symbol_retriever)
self._apply_edit(code_editor)
content_after = self._read_file(self.rel_path)
code_diff = CodeDiff(self.rel_path, original_content=content_before, modified_content=content_after)
self._test_diff(code_diff, content_after_ground_truth)
@abstractmethod
def _apply_edit(self, symbol_manager: SymbolManager) -> None:
def _apply_edit(self, code_editor: CodeEditor) -> None:
pass
def _test_diff(self, code_diff: CodeDiff, snapshot: str) -> None:
@@ -102,8 +250,8 @@ class DeleteSymbolTest(EditingTest):
self.deleted_symbol = deleted_symbol
self.rel_path = rel_path
def _apply_edit(self, symbol_manager: SymbolManager) -> None:
symbol_manager.delete_symbol(self.deleted_symbol, self.rel_path)
def _apply_edit(self, code_editor: CodeEditor) -> None:
code_editor.delete_symbol(self.deleted_symbol, self.rel_path)
@pytest.mark.parametrize(
@@ -170,12 +318,12 @@ class InsertInRelToSymbolTest(EditingTest):
def set_mode(self, mode: Literal["before", "after"]):
self.mode = mode
def _apply_edit(self, symbol_manager: SymbolManager) -> None:
def _apply_edit(self, code_editor: CodeEditor) -> None:
assert self.mode is not None
if self.mode == "before":
symbol_manager.insert_before_symbol(self.symbol_name, self.rel_path, self.new_content, use_same_indentation=False)
code_editor.insert_before_symbol(self.symbol_name, self.rel_path, self.new_content)
elif self.mode == "after":
symbol_manager.insert_after_symbol(self.symbol_name, self.rel_path, self.new_content, use_same_indentation=False)
code_editor.insert_after_symbol(self.symbol_name, self.rel_path, self.new_content)
@pytest.mark.parametrize("mode", ["before", "after"])
@@ -266,8 +414,8 @@ class ReplaceBodyTest(EditingTest):
self.symbol_name = symbol_name
self.new_body = new_body
def _apply_edit(self, symbol_manager: SymbolManager) -> None:
symbol_manager.replace_body(self.symbol_name, self.rel_path, self.new_body, use_same_indentation=False)
def _apply_edit(self, code_editor: CodeEditor) -> None:
code_editor.replace_body(self.symbol_name, self.rel_path, self.new_body)
@pytest.mark.parametrize(
+41 -35
View File
@@ -1,5 +1,6 @@
import pytest
from serena.project import Project
from solidlsp.ls import SolidLanguageServer
from solidlsp.ls_config import Language
from solidlsp.ls_types import UnifiedSymbolInformation
@@ -90,24 +91,6 @@ class TestLanguageServerBasics:
symbol_names = [symbol["name"] for symbol in result]
assert any("add" in name.lower() for name in symbol_names), f"Should find 'add' function in symbols: {symbol_names}"
@pytest.mark.parametrize("language_server", [Language.CLOJURE], indirect=True)
def test_retrieve_content_around_line(self, language_server: SolidLanguageServer):
"""Test retrieving content around specific lines"""
# Test retrieving content around the greet function definition (line 2)
result = language_server.retrieve_content_around_line(CORE_PATH, 2, 2)
assert result is not None, "Should retrieve content around line 2"
content_str = result.to_display_string()
assert "greet" in content_str, "Should contain the greet function definition"
assert "defn" in content_str, "Should contain defn keyword"
# Test retrieving content around multiply function (around line 13)
result = language_server.retrieve_content_around_line(CORE_PATH, 13, 1)
assert result is not None, "Should retrieve content around line 13"
content_str = result.to_display_string()
assert "multiply" in content_str, "Should contain multiply function"
@pytest.mark.parametrize("language_server", [Language.CLOJURE], indirect=True)
def test_namespace_functions(self, language_server: SolidLanguageServer):
"""Test definition lookup for core/greet usage in utils.clj"""
@@ -120,26 +103,13 @@ class TestLanguageServerBasics:
definition = result[0]
assert definition["relativePath"] == CORE_PATH, "Should find the definition of greet in core.clj"
@pytest.mark.parametrize("language_server", [Language.CLOJURE], indirect=True)
def test_search_files_for_pattern(self, language_server: SolidLanguageServer):
result = language_server.search_files_for_pattern("defn.*greet")
assert result is not None, "Pattern search should return results"
assert len(result) > 0, "Should find at least one match for 'defn.*greet'"
core_matches = [match for match in result if match.source_file_path and "core.clj" in match.source_file_path]
assert len(core_matches) > 0, "Should find greet function in core.clj"
result = language_server.search_files_for_pattern(":require")
assert result is not None, "Should find require statements"
utils_matches = [match for match in result if match.source_file_path and "utils.clj" in match.source_file_path]
assert len(utils_matches) > 0, "Should find require statement in utils.clj"
@pytest.mark.parametrize("language_server", [Language.CLOJURE], indirect=True)
def test_request_references_with_content(self, language_server: SolidLanguageServer):
"""Test references to multiply function with content"""
result = language_server.request_references_with_content(CORE_PATH, 12, 6, 3)
references = language_server.request_references(CORE_PATH, 12, 6)
result = [
language_server.retrieve_content_around_line(ref1["relativePath"], ref1["range"]["start"]["line"], 3, 0) for ref1 in references
]
assert result is not None, "Should find references with content"
assert isinstance(result, list)
@@ -214,3 +184,39 @@ class TestLanguageServerBasics:
break
assert found_relevant_references, f"Should have found calculate-area referencing multiply, but got: {result}"
class TestProjectBasics:
@pytest.mark.parametrize("project", [Language.CLOJURE], indirect=True)
def test_retrieve_content_around_line(self, project: Project):
"""Test retrieving content around specific lines"""
# Test retrieving content around the greet function definition (line 2)
result = project.retrieve_content_around_line(CORE_PATH, 2, 2)
assert result is not None, "Should retrieve content around line 2"
content_str = result.to_display_string()
assert "greet" in content_str, "Should contain the greet function definition"
assert "defn" in content_str, "Should contain defn keyword"
# Test retrieving content around multiply function (around line 13)
result = project.retrieve_content_around_line(CORE_PATH, 13, 1)
assert result is not None, "Should retrieve content around line 13"
content_str = result.to_display_string()
assert "multiply" in content_str, "Should contain multiply function"
@pytest.mark.parametrize("project", [Language.CLOJURE], indirect=True)
def test_search_files_for_pattern(self, project: Project) -> None:
result = project.search_source_files_for_pattern("defn.*greet")
assert result is not None, "Pattern search should return results"
assert len(result) > 0, "Should find at least one match for 'defn.*greet'"
core_matches = [match for match in result if match.source_file_path and "core.clj" in match.source_file_path]
assert len(core_matches) > 0, "Should find greet function in core.clj"
result = project.search_source_files_for_pattern(":require")
assert result is not None, "Should find require statements"
utils_matches = [match for match in result if match.source_file_path and "utils.clj" in match.source_file_path]
assert len(utils_matches) > 0, "Should find require statement in utils.clj"
+35 -32
View File
@@ -10,6 +10,7 @@ from pathlib import Path
import pytest
from serena.project import Project
from solidlsp import SolidLanguageServer
from solidlsp.ls_config import Language
@@ -67,34 +68,6 @@ class TestElixirIntegration:
# Should point to models.ex
assert "models.ex" in defining_symbol["location"]["uri"]
@pytest.mark.parametrize("language_server", [Language.ELIXIR], indirect=True)
def test_comprehensive_symbol_search(self, language_server: SolidLanguageServer):
"""Test comprehensive symbol search across the entire project."""
# Search for all function definitions
function_pattern = r"def\s+\w+\s*[\(\s]"
function_matches = language_server.search_files_for_pattern(function_pattern)
# Should find functions across multiple files
if function_matches:
files_with_functions = set()
for match in function_matches:
if match.source_file_path:
files_with_functions.add(os.path.basename(match.source_file_path))
# Should find functions in multiple files
expected_files = {"models.ex", "services.ex", "examples.ex", "utils.ex", "test_repo.ex"}
found_files = expected_files.intersection(files_with_functions)
assert len(found_files) > 0, f"Expected functions in {expected_files}, found in {files_with_functions}"
# Search for struct definitions
struct_pattern = r"defstruct\s+\["
struct_matches = language_server.search_files_for_pattern(struct_pattern)
if struct_matches:
# Should find structs primarily in models.ex
models_structs = [m for m in struct_matches if m.source_file_path and "models.ex" in m.source_file_path]
assert len(models_structs) > 0, "Should find struct definitions in models.ex"
@pytest.mark.parametrize("language_server", [Language.ELIXIR], indirect=True)
def test_module_hierarchy_understanding(self, language_server: SolidLanguageServer):
"""Test that the language server understands Elixir module hierarchy."""
@@ -136,12 +109,42 @@ class TestElixirIntegration:
assert not matcher.is_relevant_filename("package.json")
assert not matcher.is_relevant_filename("Cargo.toml")
@pytest.mark.parametrize("language_server", [Language.ELIXIR], indirect=True)
def test_protocol_and_implementation_understanding(self, language_server: SolidLanguageServer):
class TestElixirProject:
@pytest.mark.parametrize("project", [Language.ELIXIR], indirect=True)
def test_comprehensive_symbol_search(self, project: Project):
"""Test comprehensive symbol search across the entire project."""
# Search for all function definitions
function_pattern = r"def\s+\w+\s*[\(\s]"
function_matches = project.search_source_files_for_pattern(function_pattern)
# Should find functions across multiple files
if function_matches:
files_with_functions = set()
for match in function_matches:
if match.source_file_path:
files_with_functions.add(os.path.basename(match.source_file_path))
# Should find functions in multiple files
expected_files = {"models.ex", "services.ex", "examples.ex", "utils.ex", "test_repo.ex"}
found_files = expected_files.intersection(files_with_functions)
assert len(found_files) > 0, f"Expected functions in {expected_files}, found in {files_with_functions}"
# Search for struct definitions
struct_pattern = r"defstruct\s+\["
struct_matches = project.search_source_files_for_pattern(struct_pattern)
if struct_matches:
# Should find structs primarily in models.ex
models_structs = [m for m in struct_matches if m.source_file_path and "models.ex" in m.source_file_path]
assert len(models_structs) > 0, "Should find struct definitions in models.ex"
@pytest.mark.parametrize("project", [Language.ELIXIR], indirect=True)
def test_protocol_and_implementation_understanding(self, project: Project):
"""Test that the language server understands Elixir protocols and implementations."""
# Search for protocol definitions
protocol_pattern = r"defprotocol\s+\w+"
protocol_matches = language_server.search_files_for_pattern(protocol_pattern, paths_include_glob="**/models.ex")
protocol_matches = project.search_source_files_for_pattern(protocol_pattern, paths_include_glob="**/models.ex")
if protocol_matches:
# Should find the Serializable protocol
@@ -150,7 +153,7 @@ class TestElixirIntegration:
# Search for protocol implementations
impl_pattern = r"defimpl\s+\w+"
impl_matches = language_server.search_files_for_pattern(impl_pattern, paths_include_glob="**/models.ex")
impl_matches = project.search_source_files_for_pattern(impl_pattern, paths_include_glob="**/models.ex")
if impl_matches:
# Should find multiple implementations
+21 -18
View File
@@ -9,6 +9,7 @@ import os
import pytest
from serena.project import Project
from serena.text_utils import LineType
from solidlsp import SolidLanguageServer
from solidlsp.ls_config import Language
@@ -78,20 +79,22 @@ class TestLanguageServerBasics:
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", [Language.PYTHON], indirect=True)
def test_retrieve_content_around_line(self, language_server: SolidLanguageServer) -> None:
class TestProjectBasics:
@pytest.mark.parametrize("project", [Language.PYTHON], indirect=True)
def test_retrieve_content_around_line(self, project: Project) -> None:
"""Test retrieve_content_around_line functionality with various scenarios."""
file_path = os.path.join("test_repo", "models.py")
# Scenario 1: Just a single line (User class definition)
line_31 = language_server.retrieve_content_around_line(file_path, 31)
line_31 = project.retrieve_content_around_line(file_path, 31)
assert len(line_31.lines) == 1
assert "class User(BaseModel):" in line_31.lines[0].line_content
assert line_31.lines[0].line_number == 31
assert line_31.lines[0].match_type == LineType.MATCH
# Scenario 2: Context above and below
with_context_around_user = language_server.retrieve_content_around_line(file_path, 31, 2, 2)
with_context_around_user = project.retrieve_content_around_line(file_path, 31, 2, 2)
assert len(with_context_around_user.lines) == 5
# Check line content
assert "class User(BaseModel):" in with_context_around_user.matched_lines[0].line_content
@@ -111,7 +114,7 @@ class TestLanguageServerBasics:
assert with_context_around_user.lines[4].match_type == LineType.AFTER_MATCH
# Scenario 3a: Only context above
with_context_above = language_server.retrieve_content_around_line(file_path, 31, 3, 0)
with_context_above = project.retrieve_content_around_line(file_path, 31, 3, 0)
assert len(with_context_above.lines) == 4
assert "return cls(id=id, name=name)" in with_context_above.lines[0].line_content
assert "class User(BaseModel):" in with_context_above.matched_lines[0].line_content
@@ -128,7 +131,7 @@ class TestLanguageServerBasics:
assert with_context_above.lines[3].match_type == LineType.MATCH
# Scenario 3b: Only context below
with_context_below = language_server.retrieve_content_around_line(file_path, 31, 0, 3)
with_context_below = project.retrieve_content_around_line(file_path, 31, 0, 3)
assert len(with_context_below.lines) == 4
assert "class User(BaseModel):" in with_context_below.matched_lines[0].line_content
assert with_context_below.num_matched_lines == 1
@@ -143,7 +146,7 @@ class TestLanguageServerBasics:
assert with_context_below.lines[3].match_type == LineType.AFTER_MATCH
# Scenario 4a: Edge case - context above but line is at 0
first_line_with_context_around = language_server.retrieve_content_around_line(file_path, 0, 2, 1)
first_line_with_context_around = project.retrieve_content_around_line(file_path, 0, 2, 1)
assert len(first_line_with_context_around.lines) <= 4 # Should have at most 4 lines (line 0 + 1 below + up to 2 above)
assert first_line_with_context_around.lines[0].line_number <= 2 # First line should be at most line 2
# Check match type for the target line
@@ -156,7 +159,7 @@ class TestLanguageServerBasics:
assert line.match_type == LineType.AFTER_MATCH
# Scenario 4b: Edge case - context above but line is at 1
second_line_with_context_above = language_server.retrieve_content_around_line(file_path, 1, 3, 1)
second_line_with_context_above = project.retrieve_content_around_line(file_path, 1, 3, 1)
assert len(second_line_with_context_above.lines) <= 5 # Should have at most 5 lines (line 1 + 1 below + up to 3 above)
assert second_line_with_context_above.lines[0].line_number <= 1 # First line should be at most line 1
# Check match type for the target line
@@ -170,10 +173,10 @@ class TestLanguageServerBasics:
# Scenario 4c: Edge case - context below but line is at the end of file
# First get the total number of lines in the file
all_content = language_server.retrieve_full_file_content(file_path)
all_content = project.read_file(file_path)
total_lines = len(all_content.split("\n"))
last_line_with_context_around = language_server.retrieve_content_around_line(file_path, total_lines - 1, 1, 3)
last_line_with_context_around = project.retrieve_content_around_line(file_path, total_lines - 1, 1, 3)
assert len(last_line_with_context_around.lines) <= 5 # Should have at most 5 lines (last line + 1 above + up to 3 below)
assert last_line_with_context_around.lines[-1].line_number >= total_lines - 4 # Last line should be at least total_lines - 4
# Check match type for the target line
@@ -185,33 +188,33 @@ class TestLanguageServerBasics:
else:
assert line.match_type == LineType.AFTER_MATCH
@pytest.mark.parametrize("language_server", [Language.PYTHON], indirect=True)
def test_search_files_for_pattern(self, language_server: SolidLanguageServer) -> None:
@pytest.mark.parametrize("project", [Language.PYTHON], indirect=True)
def test_search_files_for_pattern(self, project: Project) -> None:
"""Test search_files_for_pattern with various patterns and glob filters."""
# Test 1: Search for class definitions across all files
class_pattern = r"class\s+\w+\s*(?:\([^{]*\)|:)"
matches = language_server.search_files_for_pattern(class_pattern)
matches = project.search_source_files_for_pattern(class_pattern)
assert len(matches) > 0
# Should find multiple classes like User, Item, BaseModel, etc.
assert len(matches) >= 5
# Test 2: Search for specific class with include glob
user_class_pattern = r"class\s+User\s*(?:\([^{]*\)|:)"
matches = language_server.search_files_for_pattern(user_class_pattern, paths_include_glob="**/models.py")
matches = project.search_source_files_for_pattern(user_class_pattern, paths_include_glob="**/models.py")
assert len(matches) == 1 # Should only find User class in models.py
assert matches[0].source_file_path is not None
assert "models.py" in matches[0].source_file_path
# Test 3: Search for method definitions with exclude glob
method_pattern = r"def\s+\w+\s*\([^)]*\):"
matches = language_server.search_files_for_pattern(method_pattern, paths_exclude_glob="**/models.py")
matches = project.search_source_files_for_pattern(method_pattern, paths_exclude_glob="**/models.py")
assert len(matches) > 0
# Should find methods in services.py but not in models.py
assert all(match.source_file_path is not None and "models.py" not in match.source_file_path for match in matches)
# Test 4: Search for specific method with both include and exclude globs
create_user_pattern = r"def\s+create_user\s*\([^)]*\)(?:\s*->[^:]+)?:"
matches = language_server.search_files_for_pattern(
matches = project.search_source_files_for_pattern(
create_user_pattern, paths_include_glob="**/*.py", paths_exclude_glob="**/models.py"
)
assert len(matches) == 1 # Should only find create_user in services.py
@@ -220,7 +223,7 @@ class TestLanguageServerBasics:
# Test 5: Search for a pattern that should appear in multiple files
init_pattern = r"def\s+__init__\s*\([^)]*\):"
matches = language_server.search_files_for_pattern(init_pattern)
matches = project.search_source_files_for_pattern(init_pattern)
assert len(matches) > 1 # Should find __init__ in multiple classes
# Should find __init__ in both models.py and services.py
assert any(match.source_file_path is not None and "models.py" in match.source_file_path for match in matches)
@@ -228,5 +231,5 @@ class TestLanguageServerBasics:
# Test 6: Search with a pattern that should have no matches
no_match_pattern = r"def\s+this_method_does_not_exist\s*\([^)]*\):"
matches = language_server.search_files_for_pattern(no_match_pattern)
matches = project.search_source_files_for_pattern(no_match_pattern)
assert len(matches) == 0