Refactoring and fixes in agent

1. Separated out a method create_ls_for_project
2. Fix serialization and IPC of config objects
This commit is contained in:
Michael Panchenko
2025-06-16 15:44:38 +02:00
parent 3d65db8286
commit 2e3eda72ba
4 changed files with 171 additions and 166 deletions
+113 -106
View File
@@ -15,7 +15,7 @@ from abc import ABC, abstractmethod
from collections import defaultdict
from collections.abc import Callable, Generator, Iterable, Sequence
from copy import copy, deepcopy
from dataclasses import dataclass, field
from dataclasses import asdict, dataclass, field
from fnmatch import fnmatch
from functools import cached_property
from logging import Logger
@@ -159,45 +159,47 @@ class ProjectConfig(ToStringMixin):
config_with_comments["language"] = dominant_language
if save_to_disk:
save_yaml(str(project_root / cls.rel_path_to_project_yml()), config_with_comments, preserve_comments=True)
return cls._from_yml_data(config_with_comments)
return cls.from_json_dict(config_with_comments)
@classmethod
def rel_path_to_project_yml(cls) -> str:
return os.path.join(SERENA_MANAGED_DIR_NAME, cls.SERENA_DEFAULT_PROJECT_FILE)
@classmethod
def _from_yml_data(cls, yaml_data: dict[str, Any]) -> Self:
def from_json_dict(cls, data: dict[str, Any]) -> Self:
"""
Create a ProjectConfig instance from a configuration dictionary
"""
data = copy(data)
try:
yaml_data["language"] = Language(yaml_data["language"].lower())
data["language"] = Language(data["language"].lower())
except ValueError as e:
raise ValueError(f"Invalid language: {yaml_data['language']}.\nValid languages are: {[l.value for l in Language]}") from e
return cls(
project_name=yaml_data["project_name"],
language=yaml_data["language"],
ignored_paths=yaml_data.get("ignored_paths", []),
excluded_tools=set(yaml_data.get("excluded_tools", [])),
read_only=yaml_data.get("read_only", False),
ignore_all_files_in_gitignore=yaml_data.get("ignore_all_files_in_gitignore", True),
initial_prompt=yaml_data.get("initial_prompt", ""),
)
raise ValueError(f"Invalid language: {data['language']}.\nValid languages are: {[l.value for l in Language]}") from e
return cls(**data)
def to_json_dict(self) -> dict[str, Any]:
result = asdict(self)
result["language"] = result["language"].value
result["excluded_tools"] = list(result["excluded_tools"])
return result
@classmethod
def load(cls, project_root: Path | str) -> Self:
def load(cls, project_root: Path | str, autogenerate: bool = True) -> Self:
"""
Load a ProjectConfig instance from the path to the project root.
"""
project_root = Path(project_root)
yaml_path = project_root / cls.rel_path_to_project_yml()
if not yaml_path.exists():
raise FileNotFoundError(f"Project configuration file not found: {yaml_path}")
if autogenerate:
return cls.autogenerate(project_root)
else:
raise FileNotFoundError(f"Project configuration file not found: {yaml_path}")
with open(yaml_path, encoding="utf-8") as f:
yaml_data = yaml.safe_load(f)
if "project_name" not in yaml_data:
yaml_data["project_name"] = project_root.name
return cls._from_yml_data(yaml_data)
return cls.from_json_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)
@@ -217,13 +219,20 @@ class Project:
return self.project_config.language
@classmethod
def load(cls, project_root: str | Path) -> Self:
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)
project_config = ProjectConfig.load(project_root, autogenerate=autogenerate)
return cls(project_root=str(project_root), project_config=project_config)
@classmethod
def from_json_dict(cls, data: dict) -> Self:
return cls(project_root=data["project_root"], project_config=ProjectConfig.from_json_dict(data["project_config"]))
def to_json_dict(self) -> dict:
return {"project_root": self.project_root, "project_config": self.project_config.to_json_dict()}
def path_to_project_yml(self) -> str:
return os.path.join(self.project_root, self.project_config.rel_path_to_project_yml())
@@ -318,58 +327,18 @@ class SerenaConfigBase(ABC):
else:
raise ValueError(f"Project '{project_name}' not found in Serena configuration; valid project names: {self.project_names}")
def to_dict(self) -> dict:
def to_json_dict(self) -> dict:
"""Convert configuration to dictionary for serialization."""
return {
"projects": [
{
"project_root": project.project_root,
"project_config": {
"project_name": project.project_config.project_name,
"language": project.project_config.language,
"ignored_paths": project.project_config.ignored_paths,
"ignore_all_files_in_gitignore": project.project_config.ignore_all_files_in_gitignore,
},
}
for project in self.projects
],
"gui_log_window_enabled": self.gui_log_window_enabled,
"log_level": self.log_level,
"trace_lsp_communication": self.trace_lsp_communication,
"web_dashboard": self.web_dashboard,
"tool_timeout": self.tool_timeout,
}
result = asdict(self)
result["projects"] = [project.to_json_dict() for project in self.projects]
return result
@classmethod
def from_dict(cls, data: dict) -> "SerenaConfigBase":
def from_json_dict(cls, data: dict) -> Self:
"""Create configuration from dictionary."""
projects = []
for project_data in data.get("projects", []):
project_config = ProjectConfig(
project_name=project_data["project_config"]["project_name"],
language=project_data["project_config"]["language"],
ignored_paths=project_data["project_config"]["ignored_paths"],
ignore_all_files_in_gitignore=project_data["project_config"]["ignore_all_files_in_gitignore"],
)
project = Project(project_root=project_data["project_root"], project_config=project_config)
projects.append(project)
# Create a basic config class that can be instantiated
from dataclasses import dataclass
@dataclass(kw_only=True)
class DeserializedConfig(SerenaConfigBase):
def _add_new_project(self, project: Project) -> None:
self.projects.append(project)
return DeserializedConfig(
projects=projects,
gui_log_window_enabled=data.get("gui_log_window_enabled", False),
log_level=data.get("log_level", logging.INFO),
trace_lsp_communication=data.get("trace_lsp_communication", False),
web_dashboard=data.get("web_dashboard", True),
tool_timeout=data.get("tool_timeout", DEFAULT_TOOL_TIMEOUT),
)
data = copy(data)
data["projects"] = [Project.from_json_dict(project_data) for project_data in data["projects"]]
return cls(**data)
@dataclass(kw_only=True)
@@ -400,23 +369,26 @@ class SerenaConfig(SerenaConfigBase):
config_file = cls.CONFIG_FILE_DOCKER if is_running_in_docker() else cls.CONFIG_FILE
return os.path.join(REPO_ROOT, config_file)
@classmethod
def _load_commented_yaml(cls, config_file: str, generate_if_missing: bool = True) -> CommentedMap:
if not os.path.exists(config_file):
if not generate_if_missing:
raise FileNotFoundError(f"Serena configuration file not found: {config_file}")
log.info(f"Serena configuration file not found at {config_file}, autogenerating...")
cls.autogenerate()
try:
return load_yaml(config_file, preserve_comments=True)
except Exception as e:
raise ValueError(f"Error loading Serena configuration from {config_file}: {e}") from e
@classmethod
def from_config_file(cls, generate_if_missing: bool = True) -> "SerenaConfig":
"""
Static constructor to create SerenaConfig from the configuration file
"""
config_file = cls.get_config_file_path()
if not os.path.exists(config_file):
if not generate_if_missing:
raise FileNotFoundError(f"Serena configuration file not found: {config_file}")
cls.autogenerate()
log.info(f"Loading Serena configuration from {config_file}")
try:
loaded_commented_yaml = load_yaml(config_file, preserve_comments=True)
except Exception as e:
raise ValueError(f"Error loading Serena configuration from {config_file}: {e}") from e
loaded_commented_yaml = cls._load_commented_yaml(config_file, generate_if_missing)
# Create instance
instance = cls(loaded_commented_yaml=loaded_commented_yaml)
@@ -500,6 +472,16 @@ class SerenaConfig(SerenaConfigBase):
super().remove_project(project_name)
self.save()
def to_json_dict(self) -> dict:
result = super().to_json_dict()
result.pop("loaded_commented_yaml", None)
return result
@classmethod
def from_json_dict(cls, data: dict) -> Self:
data["loaded_commented_yaml"] = cls._load_commented_yaml(cls.get_config_file_path())
return super().from_json_dict(data)
class LinesRead:
def __init__(self) -> None:
@@ -570,8 +552,6 @@ class MemoriesManagerMDFilesInProject(MemoriesManager):
def create_serena_config(
serena_config: SerenaConfigBase | None = None,
context: SerenaAgentContext | None = None,
modes: list[SerenaAgentMode] | None = None,
enable_web_dashboard: bool | None = None,
enable_gui_log_window: bool | None = None,
log_level: Literal["DEBUG", "INFO", "WARNING", "ERROR", "CRITICAL"] | None = None,
@@ -585,8 +565,6 @@ def create_serena_config(
to allow creating configurations independently for process isolation and other use cases.
:param serena_config: the base Serena configuration or None to read from default location
:param context: the context in which the agent will operate
:param modes: list of modes in which the agent will operate
:param enable_web_dashboard: Whether to enable the web dashboard
:param enable_gui_log_window: Whether to enable the GUI log window
:param log_level: Log level
@@ -637,6 +615,56 @@ def create_serena_config(
return config
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,
) -> SyncLanguageServer:
"""
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 = MultilspyConfig(
code_language=project_instance.language,
ignored_paths=ignored_paths,
trace_lsp_communication=trace_lsp_communication,
)
ls_logger = MultilspyLogger(log_level=log_level)
log.info(f"Creating language server instance for {project_instance.project_root}.")
return SyncLanguageServer.create(
multilspy_config,
ls_logger,
project_instance.project_root,
timeout=ls_timeout,
)
class SerenaAgent:
def __init__(
self,
@@ -669,8 +697,6 @@ class SerenaAgent:
# obtain serena configuration using the decoupled factory function
self.serena_config = create_serena_config(
serena_config=serena_config,
context=context,
modes=modes,
enable_web_dashboard=enable_web_dashboard,
enable_gui_log_window=enable_gui_log_window,
log_level=log_level,
@@ -1005,31 +1031,12 @@ class SerenaAgent:
# instantiate and start the language server
assert self._active_project is not None
ignored_paths = self._active_project.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 self._active_project.project_config.ignore_all_files_in_gitignore:
log.info(f"Parsing all gitignore files in {self._active_project.project_root}")
gitignore_parser = GitignoreParser(self._active_project.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 = MultilspyConfig(
code_language=self._active_project.project_config.language,
ignored_paths=ignored_paths,
self.language_server = create_ls_for_project(
self._active_project,
log_level=self.serena_config.log_level,
ls_timeout=ls_timeout,
trace_lsp_communication=self.serena_config.trace_lsp_communication,
)
ls_logger = MultilspyLogger(log_level=self.serena_config.log_level)
log.info(f"Starting language server for {self._active_project.project_root}.")
self.language_server = SyncLanguageServer.create(
multilspy_config,
ls_logger,
self._active_project.project_root,
timeout=ls_timeout,
)
self.language_server.start()
if not self.language_server.is_running():
raise RuntimeError(
+24 -1
View File
@@ -3,7 +3,8 @@ Context and Mode configuration loader
"""
import os
from dataclasses import dataclass, field
from copy import copy
from dataclasses import asdict, dataclass, field
from pathlib import Path
from typing import TYPE_CHECKING, Self
@@ -30,6 +31,17 @@ class SerenaAgentMode:
description: str = ""
excluded_tools: set[str] = field(default_factory=set)
def to_json_dict(self) -> dict[str, str | list[str]]:
result = asdict(self)
result["excluded_tools"] = list(result["excluded_tools"])
return result
@classmethod
def from_json_dict(cls, data: dict) -> Self:
data = copy(data)
data["excluded_tools"] = set(data["excluded_tools"])
return cls(**data)
def print_overview(self) -> None:
"""Print an overview of the mode."""
print(f"{self.name}:\n {self.description}")
@@ -91,6 +103,17 @@ class SerenaAgentContext:
description: str = ""
excluded_tools: set[str] = field(default_factory=set)
def to_json_dict(self) -> dict[str, str | list[str]]:
result = asdict(self)
result["excluded_tools"] = list(result["excluded_tools"])
return result
@classmethod
def from_json_dict(cls, data: dict) -> Self:
data = copy(data)
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
+8 -7
View File
@@ -99,6 +99,7 @@ def make_tool(
fn_metadata=func_arg_metadata,
is_async=is_async,
context_kwarg=None,
annotations=None,
)
@@ -141,15 +142,15 @@ def create_mcp_server_and_agent(
try:
serena_config = create_serena_config(
context=context_instance,
modes=modes_instances,
enable_web_dashboard=enable_web_dashboard,
enable_gui_log_window=enable_gui_log_window,
log_level=log_level,
trace_lsp_communication=trace_lsp_communication,
tool_timeout=tool_timeout,
)
serena_agent_process = ProcessIsolatedSerenaAgent(project=project, serena_config=serena_config)
serena_agent_process = ProcessIsolatedSerenaAgent(
project=project, serena_config=serena_config, modes=modes_instances, context=context_instance
)
# Start process-isolated dashboard if enabled
serena_dashboard_process = None
@@ -261,7 +262,7 @@ PROJECT_TYPE = ProjectType()
@click.command()
@click.option(
"--project",
"project_file_opt",
"project",
type=PROJECT_TYPE,
default=None,
help="Either an absolute path to the project directory or a name of an already registered project. "
@@ -270,7 +271,7 @@ PROJECT_TYPE = ProjectType()
# Keep --project-file for backwards compatibility
@click.option(
"--project-file",
"project_file_opt", # Use same destination variable to avoid conflicts
"project", # Use same destination variable to avoid conflicts
type=PROJECT_TYPE,
default=None,
help="[DEPRECATED] Use --project instead.",
@@ -357,7 +358,7 @@ PROJECT_TYPE = ProjectType()
help="Timeout in seconds for tool execution. If not specified, will take the value from the serena configuration.",
)
def start_mcp_server(
project_file_opt: str | None,
project: str | None,
project_file_arg: str | None,
context: str = DEFAULT_CONTEXT,
modes: tuple[str, ...] = DEFAULT_MODES,
@@ -378,7 +379,7 @@ def start_mcp_server(
"""
# Prioritize the positional argument if provided
# This is for backward compatibility with the old CLI, should be removed in the future!
project_file = project_file_arg if project_file_arg is not None else project_file_opt
project_file = project_file_arg if project_file_arg is not None else project
# Use process isolation by default to prevent asyncio event loop contamination
mcp_server, agent = create_mcp_server_and_agent(
+26 -52
View File
@@ -5,7 +5,6 @@ import multiprocessing
import os
import traceback
import webbrowser
from collections.abc import Callable
from enum import StrEnum
from logging.handlers import QueueHandler
from multiprocessing.connection import Connection
@@ -270,27 +269,19 @@ class SerenaAgentWorker:
"""Initialize the SerenaAgent."""
try:
# Extract all possible initialization parameters
context_param = params.get("context")
project = params.get("project")
project_activation_callback = params.get("project_activation_callback")
serena_config = params.get("serena_config")
context = params.get("context")
modes = params.get("modes")
serena_config = SerenaConfig.from_json_dict(params["serena_config"])
context = SerenaAgentContext.from_json_dict(context_param) if context_param is not None else None
modes = [SerenaAgentMode.from_json_dict(m) for m in params["modes"]]
enable_web_dashboard = params.get("enable_web_dashboard")
enable_gui_log_window = params.get("enable_gui_log_window")
log_level = params.get("log_level")
trace_lsp_communication = params.get("trace_lsp_communication")
tool_timeout = params.get("tool_timeout")
# Handle legacy config_dict parameter for backward compatibility
if "config" in params and serena_config is None:
config_dict = params["config"]
serena_config = SerenaConfig.from_dict(config_dict)
elif serena_config is not None and isinstance(serena_config, dict):
serena_config = SerenaConfig.from_dict(serena_config)
self.agent = SerenaAgent(
project=project,
project_activation_callback=project_activation_callback,
serena_config=serena_config,
context=context,
modes=modes,
@@ -405,7 +396,6 @@ class ProcessIsolatedSerenaAgent:
def __init__(
self,
project: str | None = None,
project_activation_callback: Callable[[], None] | None = None,
serena_config: SerenaConfigBase | None = None,
context: SerenaAgentContext | None = None,
modes: list[SerenaAgentMode] | None = None,
@@ -415,28 +405,17 @@ class ProcessIsolatedSerenaAgent:
trace_lsp_communication: bool | None = None,
tool_timeout: float | None = None,
):
# Store all initialization parameters to pass to worker
self._init_params = {
"project": project,
"project_activation_callback": project_activation_callback,
"serena_config": serena_config,
"context": context,
"modes": modes,
"enable_web_dashboard": enable_web_dashboard,
"enable_gui_log_window": enable_gui_log_window,
"log_level": log_level,
"trace_lsp_communication": trace_lsp_communication,
"tool_timeout": tool_timeout,
}
# Keep serena_config for compatibility
if serena_config is not None:
self.serena_config = serena_config
else:
# Use the default config loader
self.serena_config = SerenaConfig.from_config_file()
self.process: multiprocessing.Process | None = None
self.project = project
self.serena_config = serena_config or SerenaConfig.from_config_file()
self.context = context
self.modes = modes or []
self.enable_web_dashboard = enable_web_dashboard
self.enable_gui_log_window = enable_gui_log_window
self.log_level = log_level
self.trace_lsp_communication = trace_lsp_communication
self.tool_timeout = tool_timeout
self.process: multiprocessing.Process | None = None
self.conn: Connection | None = None
def start(self) -> None:
@@ -455,25 +434,20 @@ class ProcessIsolatedSerenaAgent:
self.process = multiprocessing.Process(target=worker.run, args=[_global_log_queue])
self.process.start()
# Prepare initialization parameters, converting complex objects to dict if present
init_params = {
"project": self.project,
"serena_config": self.serena_config.to_json_dict(),
"context": self.context.to_json_dict() if self.context is not None else None,
"modes": [m.to_json_dict() for m in self.modes],
"enable_web_dashboard": self.enable_web_dashboard,
"enable_gui_log_window": self.enable_gui_log_window,
"log_level": self.log_level,
"trace_lsp_communication": self.trace_lsp_communication,
"tool_timeout": self.tool_timeout,
}
# Initialize the agent in the worker process
try:
# TODO: instead pass config and project to worker, use at startup of agent there
# Prepare initialization parameters, converting complex objects to dict if present
init_params = self._init_params.copy()
if init_params["serena_config"] is not None and hasattr(init_params["serena_config"], "to_dict"):
# Convert SerenaConfigBase to dict for serialization
serena_config_obj = init_params["serena_config"]
init_params["serena_config"] = serena_config_obj.to_dict() # type: ignore
# init_params["project"] = self.project
# Note: context and modes are not serializable across processes,
# so they will be None and SerenaAgent will use defaults
init_params["context"] = None
init_params["modes"] = None
# project_activation_callback cannot be serialized across processes
init_params["project_activation_callback"] = None
# Disable web dashboard in worker process - it will run in MCP process instead
init_params["enable_web_dashboard"] = False
self._make_request_with_result(SerenaAgentWorker.RequestMethod.INITIALIZE, init_params)
except Exception as e:
self.stop()