diff --git a/.serena/memories/critical_info.md b/.serena/memories/critical_info.md index d7e9023e..700c57b6 100644 --- a/.serena/memories/critical_info.md +++ b/.serena/memories/critical_info.md @@ -25,18 +25,39 @@ Snapshot tests use syrupy. # Docstrings & Comments +Documentation style: + * You consistently use reStructuredText. * You structure function implementations into functional blocks that are separated by blank lines. Atop each functional block, you write an elliptical phrase (starting with lower-case letter) that describes the purpose of the block in a concise manner. * When describing parameters, methods/functions and classes, you use a precise style, where the initial (elliptical) phrase clearly defines *what* it is. Any details then follow in subsequent sentences. + +General documentation principles: + * Each piece of information appears exactly once, at the element that owns it: callers do not explain callees' internals, and callees do not describe their callers. +* Code changes are documented exclusively in commit messages, not in comments. +* Things you consequently avoid: + - For a function/class, you do not describe who calls/uses it. A call/usage site may explain why a usage occurs, but only if it is non-obvious. + - When a function is called, you do not describe what it does. That belongs in the function's docstring. + - You do not describe in comments how an implementation differs from a previous state of the code or why a change was made. That belongs in the commit message. -# Pull requests +# REPL & facades -Read `mem:creating_pull_requests` when asked to participate in the creation of a pull request. +Read `mem:repl` before working on `serena.repl` (the code-execution paradigm and its facade APIs) or on tools +delegating to it: structure, exposure/naming principles, configuration of the API scope and the availability policy. + +# Commits & pull requests + +* Commit messages: + * The subject line must cover the *entire* change; make it suitably abstract if necessary + * Details are presented in concise bullet items, one point per item; no prose paragraphs. + Where the change spans several topics, group the items by topic, each group with a short heading. + * Wrap all lines at ~100 characters. + * Write the message to a file and commit with `-F`; do not pass long text via `-m`. +* Read `mem:creating_pull_requests` when asked to participate in the creation of a pull request. # Memories diff --git a/.serena/memories/repl.md b/.serena/memories/repl.md new file mode 100644 index 00000000..9be43b9e --- /dev/null +++ b/.serena/memories/repl.md @@ -0,0 +1,120 @@ +# REPL (`serena.repl`) + +Alternative interaction paradigm: one tool (`serena_repl`) executes Python code against entrypoint `s`, +whose attributes are facades (`s.lsp`, `s.edit`, `s.fs`, `s.mem`, `s.shell`, `s.jb`, `s.cfg`). +Code runs like a notebook cell (module-level exec in the session namespace); the value of a trailing +expression is the result. No `return` (a top-level `return` yields a SyntaxError with a hint). + +IMPORTANT: The REPL interface is a BETA feature. If you encounter any issues, please report them +but do not submit PRs for it (except for trivial fixes); the implementation is still evolving. + +## Structure + +- `repl/api/*_api.py`: `FacadeApi` implementations = the single implementation of each operation. + The classic tools are thin adapters delegating to the APIs (via `*ApiMixin`); tools retain only + transport concerns (input sanitisation, diagnostics context, tool-level output shaping). +- `repl/facade.py`: `Facade` = indirection over an API instance; `FacadeMethod` (enabled flag + `FacadeMethodInfo`); + `ApiScope` = which facades/methods are enabled. +- `repl/repl.py`: `SerenaRepl` (execution, error formatting), `SerenaReplEntrypoint` (`s`, `info`). +- `repl/representable.py`: `Representable`/`Renderer`; result objects carry their rendering policy. + +## Design principles + +- Exposure is explicit: a method is exposed iff decorated with `@facade_method(...)`, which carries + `optional`, `beta`, `can_edit`, `corresponding_tool` (mirroring the tool markers; the tool correspondence + is recorded for optional derivation of exclusions and prompt conditions, never applied automatically). +- Naming: on result objects and non-exposed API helpers, a trailing underscore (`symbols_`, `to_dict_`) + marks members that are Serena-public but not LLM-facing. +- Facades group by *domain*, not by read vs. write; mutation is expressed via `can_edit` (read-only projects + exclude editing methods). Boundary `fs`/`edit`: files as units vs. modifying content within existing files. +- Facade descriptions describe the domain only; never list operations (the method list is always shown alongside). +- Output parameters (depth, include_body, max_answer_chars, ...) are passed at retrieval time so that the + rendering policy is fixed once and inherited by derived results. +- Results expose data to code (`.symbols`, `.occurrences`, `.lines`, ...) and render like the classic tool output. +- Progressive disclosure: a priori only facade names, descriptions and method names; `s.info("")` / + `s.info(".")` give signature + docstring together, never a signature alone. + `info(*items)` documents several items at once; unknown items are reported inline. +- Disclosure tiers: tier 0 (tool description) = facades, descriptions, method names with navigable return types + (`find_symbol -> LspSymbolCollection`); tier 1 (`s.info("")`) = all common methods in full, + `niche` methods (`@facade_method(niche=True)`, rarely needed + long docs) only as summary + pointer, result + types by name only; tier 2 = types on request. Types are never pushed (`provide_info_with_facade` exists but + is set nowhere); the tool description tells the model to request type docs only when processing results in code. +- Result types: every user-defined class reachable through annotations (method parameters/returns, and the members + of reachable types, transitively) is automatically documentable (`Facade._discover_referenced_types`); builtins, + typing constructs and stdlib classes are excluded. Explicit `ReferencedType` declarations (constructor arg + `types=`) exist for curation only: an optional `members` whitelist (foreign/large classes such as + `LanguageServerSymbol`; listed methods are shown even if undocumented, convention-derived ones only if + documented) and flags. Enums render with members/values, TypedDicts with their keys. Types are documented via + `s.info(".")` or bare `s.info("")`; method docs point to their referenced types. + Result classes declare attribute annotations at class level (attributes set only in `__init__` are not + discoverable). Annotations are rendered without module paths, so signature names equal lookup names. + A type's documentation transitively includes the declared types its members reference; within a session + (`SerenaSession.described_type_names`), a contained type is documented once and afterwards only pointed to + (explicit requests always yield full documentation). + +## Sessions (`serena.session`) + +- MCP provides no reliable session identification (newer protocol versions drop it), and clients keep a stdio + server across conversations. Hence the REPL's session identity is LLM-supplied: `create_system_prompt` creates + a `SerenaSession` and states its id; `serena_repl` takes a required `session_id`. `SessionRegistry` creates + unknown ids on demand (benign: at worst docs are repeated) and evicts LRU and idle (TTL) sessions. + Sessions survive REPL rebuilds. +- Persistence (notebook semantics): `SerenaSession.repl_namespace` is the globals of the session's executions, + which run at module level (statements exec'd, a trailing expression eval'd; line numbers preserved), so all + top-level bindings persist across calls. `s` is re-bound in the namespace before every execution, so persisted + functions always use the current entrypoint. `s.vars()`/`s.clear()` list/remove persisted items. Data is tied to + the session's lifetime (not cleared on REPL rebuild); stored facades/project objects may go stale. + Deferred idea if memory becomes an issue: hybrid — implicit items expire after N turns, explicit store + (e.g. `s.d`) for indefinite retention. +- Only tools whose use presupposes having read the instructions may require the id. `activate_project` and + `initial_instructions` may be called first and keep the existing MCP-context-derived session handling (prompt + provision status); migrating that to LLM-supplied ids is a separate, future change. +- APIs must not import `serena.tools` at module level except for tool classes in decorators; tools import + APIs locally in `_api()` (API modules refer to tool classes). + +## Configuration + +- `agent_interface: tools | REPL` (`AgentInterface`; global config, overridable per project; CLI `--agent-interface`). + `None` = Serena's default (`tools`). Fixed for the session. In REPL mode the toolset is *fixed* + (`serena_repl`, `initial_instructions`, `activate_project` unless single-project); tool inclusion/exclusion + definitions do not apply — each interface has its own configuration vocabulary (tool definitions ↔ tools, + API definitions ↔ REPL). Contexts do not influence the interface. + In REPL mode, the language backend may change upon project activation (a project's backend override is + applied; background modes, facades, prompt params and backend initialisation are recomputed), whereas the + tool interface forbids this (the toolset depends on the backend and is fixed). + Idea (not implemented, considered over-engineered for now): contexts could declare *supported* interfaces + (a capability constraint, e.g. clients that handle the REPL badly), with the user's preference choosing among them. +- `included_apis`/`excluded_apis` (references `facade` or `facade.method`) in global config, context, modes, + project config; applied in that order via `ApiScope` (exclusions first, then inclusions; later definitions win). +- Opt-in rule: methods of a facade that is not included (excluded, or optional without explicit inclusion) and + optional methods are enabled only if included explicitly; all other methods are enabled unless excluded. + Facades can be optional (`Facade.from_api(..., is_optional=True)`, e.g. `ext`, mirroring optional tools); + `Facade.is_enabled()` is derived: a facade is available iff it has at least one enabled method. +- The REPL is rebuilt whenever the active tools are updated (mode switch, project activation). + +## Availability policy + +- Keep as much functionality as possible in the REPL; exclude nothing by default. + * Do not derive API exclusions from tool exclusions automatically (not via the `corresponding_tool` + correspondence, not via an option): contexts exclude tools mostly because the *host* provides equivalents + (`read_file`, `find_file`, `replace_content`, shell). In the REPL those reads are what makes operations + composable (read → filter → return a summary), and a host tool cannot participate in REPL code. + * The host's better-integrated edit tools (diff view, undo) are a matter of guidance in the context prompt, + not of availability: exclusions can only steer the model, never enforce anything. + * For users migrating from tool mode: prefer a startup hint listing the `excluded_apis` entries corresponding + to their own (global/project, not context/mode) tool exclusions over any automatic derivation. +- Python code can always modify the system; the REPL tool is inherently fully privileged, regardless of + facade scope or the project's `read_only` setting (which only makes Serena's own API refuse edits). + A "read-only REPL" is not feasible and must not be promised. +- External projects (`s.ext`): `list_projects()`, `project_context(name)` (a `with`-able context; not nestable). + Within it, the agent's active project is temporarily switched (`active_project_context`) and the facades are + read-only (`can_edit` methods raise). Methods marked `@facade_method(uses_project_server=True)` (all of `lsp`) + are executed in the project server via `/call_facade_method` ({facade, method, args, kwargs} as JSON, result + pickled; the server is a trusted local process) when the LSP backend is active; with JetBrains they run locally + (the IDE serves all projects). Result objects must be self-contained/picklable: renderers hold no agent (only the + default length limit), LSP results carry eagerly retrieved info and reference contexts, no lambdas in output + params. Replaces the query_project/list_queryable_projects tools in the REPL. +- Project activation (activate_project) and initial_instructions stay tool-only (activation rebuilds the REPL); + Serena's configuration/session state (config overview, dashboard; later e.g. modes) lives in the `cfg` facade. + Computed conditions (read-only project, dashboard not openable) are applied to the API scope in + `SerenaAgent.get_repl` via `exclude_editing()`/`NamedApiInclusionDefinition`, mirroring the tool side. diff --git a/.serena/project.yml b/.serena/project.yml index 7c5257c1..8745f0bc 100644 --- a/.serena/project.yml +++ b/.serena/project.yml @@ -1,24 +1,24 @@ # the name by which the project can be referenced within Serena/when chatting with the LLM. project_name: "serena" - # list of language servers to start when using the LSP backend; choose from: -# ada al angular ansible bash -# bsl clojure cpp cpp_ccls crystal -# csharp csharp_omnisharp cue dart elixir -# elm erlang fortran fsharp gdscript -# go groovy haskell haxe hlsl -# html java json julia kotlin -# latex lean4 lua luau markdown -# matlab msl nix ocaml pascal -# perl php php_phpactor php_phpantom powershell -# python python_jedi python_pyrefly python_ty r -# rego ruby ruby_solargraph rust scala -# scss solidity svelte swift systemverilog -# terraform toml typescript typescript_vts vue -# yaml zig +# ada al angular ansible bash +# bsl clojure cpp cpp_ccls crystal +# csharp csharp_omnisharp cue dart deno +# elixir elm erlang fortran fsharp +# gdscript gleam go groovy haskell +# haxe hlsl html java json +# julia julia_fatou kotlin latex lean4 +# lua luau markdown matlab msl +# nextflow nix ocaml pascal perl +# php php_phpactor php_phpantom powershell python +# python_basedpyright python_jedi python_pyrefly python_ty qml +# r rego ruby ruby_solargraph rust +# scala scss solidity svelte swift +# systemverilog terraform toml typescript typescript_vts +# vue wolfram yaml zig # (This list may be outdated; generated with scripts/print_language_list.py; -# For the current list, see values of Language enum here: +# For the current list, see values of the LanguageServerId enum here: # https://github.com/oraios/serena/blob/main/src/solidlsp/ls_config.py) # For some languages, there are several alternative language servers, e.g. csharp_omnisharp, ruby_solargraph.) # Note: @@ -26,6 +26,7 @@ project_name: "serena" # - For JavaScript, use typescript # - For Angular projects, use angular (subsumes typescript+html; requires `npm install` in the project root) # - For Svelte projects, use svelte (subsumes typescript/javascript for .svelte projects; requires npm) +# - For Deno projects, use deno (serves the same .ts/.js files as typescript; requires the deno CLI on PATH) # - For SCSS / Sass / plain CSS, use scss (some-sass-language-server handles all three) # - For Free Pascal/Lazarus, use pascal # Special requirements: @@ -35,8 +36,8 @@ project_name: "serena" # The first language server is the default language and the respective language server will be used as a fallback. # Note that when using the JetBrains backend, language servers are not used and this list is correspondingly ignored. language_servers: - - python - - typescript +- python +- typescript # whether to use project's .gitignore files to ignore files ignore_all_files_in_gitignore: true @@ -68,9 +69,8 @@ excluded_tools: [] # Find the list of tools here: https://oraios.github.io/serena/01-about/035_tools.html included_optional_tools: [] -# initial prompt for the project, which will be provided to the LLM upon project activation -# (or within Serena's initial instructions if the project is activated at startup). -## See: https://oraios.github.io/serena/02-usage/050_configuration.html#prompt-templates +# initial prompt for the project. It will always be given to the LLM upon activating the project +# (contrary to the memories, which are loaded on demand). initial_prompt: | {{ embed_memory("critical_info") }} @@ -177,3 +177,17 @@ ls_workspace_folders: # - ../sibling-package # - ../shared-lib ls_additional_workspace_folders: [] + +# list of APIs (facades or facade methods, e.g. "lsp" or "lsp.get_diagnostics_for_symbol") to include in the REPL +# that would otherwise be disabled (particularly optional methods, which are disabled by default). +# This extends the existing inclusions (e.g. from the global configuration). +included_apis: [] + +# list of APIs (facades or facade methods, e.g. "lsp" or "lsp.find_symbol") to exclude from the REPL. +# This extends the existing exclusions (e.g. from the global configuration). +excluded_apis: [] + +# The interface through which the agent (LLM) accesses Serena's functionality (overrides the global setting). +# Valid values: tools, REPL (see the global configuration for details); leave empty to use the global setting. +# Note: the interface is fixed at startup. If a project is activated post-init, its setting is not applied. +agent_interface: diff --git a/CHANGELOG.md b/CHANGELOG.md index 9eeb6567..ff6a60af 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -12,7 +12,16 @@ Status of the `main` branch. Changes prior to the next official version change w see `CONTRIBUTING.md` * General: - - Fix: MCP `initialize` now reports Serena's version instead of the installed mcp SDK version (#1889) + - **Major**: Add the Serena REPL as a new agent interface, reducing the tool set to a minimum and providing + a general code execution environment for all Serena operations. + This has several significant advantages over regular tool executions. + Please refer to our [documentation](https://oraios.github.io/serena/01-about/035_tools.html) for details. + - Add `auth_secret` to `serena_config.yml` for authenticating communication between Serena components + and services. When missing, null, or empty, a random UUID is generated and persisted; existing values + are preserved + - Fix: MCP server now reports Serena's version instead of the installed MCP SDK version (#1889) + - Fix: importing Serena no longer loads the `anthropic` package unless the Anthropic token counter is + actually used; the unconditional import added seconds to CLI/MCP startup on some machines (#2012) - Fix: Parallel agents auto-registering projects could overwrite each other's changes to the global project list in `serena_config.yml` - Perf: `search_for_pattern` resolved each match's line number by rescanning the file from the @@ -24,23 +33,56 @@ Status of the `main` branch. Changes prior to the next official version change w - Fix: process-tree cleanup signaled descendant language-server processes without waiting for them, which could leave grandchildren as zombies; cleanup now waits for the discovered descendants (#1464) - Fix: `read_only` restriction in project definition was not applied to base tool set when in single-project context (#1938) + - Fix: `SerenaConfig.project_names` / `project_paths` were cached and never invalidated after + projects were added or removed mid-session, so user-facing project lists and error messages + stayed stale; the lists are no longer cached - Docs: `trusted_project_path_patterns` now documents how to trust a single project. Trust is decided by the project's root path, so a `/**` entry matches only paths below the root and therefore trusts no project at all; the template now shows the bare root form alongside the parent-directory glob (#2001) + - Session IDs are now created and tracked internally by Serena instead of being derived from the + MCP session, since the MCP SDK v2 no longer provides session identifiers and client session usage + was inconsistent anyway. Tools that need a session id (e.g. `activate_project`, the REPL tool) now + take it as an explicit parameter, obtained from `initial_instructions` + - Performance: `Project.gather_source_files` transitively re-derived from the filesystem, for every path, + whether that path was a file or a directory; related methods/functions now receive the information + as a parameter where it is already known (#2077) * CLI: - Fix: `project health-check` reported `Health check passed - All tools working correctly` and exited 0 even when `FindReferencingSymbolsTool` had raised, because that failure was logged as a warning while the verdict checked `FindSymbolTool` only. A reference-search failure now fails the check; a symbol with no references is still a pass + - Add `project remove`, which unregisters a project from the project list in `serena_config.yml`, + addressed either by name or by path. Only the registry entry is removed; the project's own files, + including its project configuration, are left untouched (#2029) + +* Tools: + - Fix: `$!N` backreferences in regex-mode replacements expanded to the literal template text + (e.g. `EA_INPUT$!1(...)`) when the referenced group existed but did not participate in the + match (e.g. a group inside an optional construct that was skipped); unmatched groups now expand + to the empty string, and a reference to a group that the search expression does not define + raises a clear error instead of a raw `IndexError`. In literal mode, the replacement is now + used verbatim (`$!N` sequences need no escaping) instead of failing with a backreference error + - Fix: the file-editing tools saved the edited file with `open(path, "w")`, which truncates it + before the new content is complete, so a crash, an OOM kill or a full disk partway through the + write could leave a source file empty or half-written. Saves now go through the same atomic + temp-file-plus-`os.replace` helper that the memory writes already use. The helper resolves + symlinks first, so a symlinked file is still written through to its target rather than being + replaced by a regular file (#1958) * Memories: + - Fix: `move_memory` / rename only checked write access on the destination name, so a tool-context + rename could relocate a read-only memory; both source and destination are now checked - Fix: `save_memory`/`edit_memory` wrote directly to the memory file with `open(path, "w")`, which truncates it before the new content is written; a crash, OOM kill, or full disk partway through the write could destroy the previous, valid content instead of just losing the update. Both now write through a temp-file-plus-`os.replace` helper, matching the approach `save_yaml()` already uses for settings files (#1958) + - Fix: renaming a memory through the `rename_memory` tool raised `PermissionError` when another memory + marked read-only by `read_only_memory_patterns` referenced it, after the rename had already been + applied, leaving the memory graph half-updated; reference propagation in tool contexts now covers + only writable memories, as documented, while the CLI still propagates into read-only ones * JetBrains: - Fix: Concurrent Serena sessions activating different projects at the same time with @@ -54,7 +96,35 @@ Status of the `main` branch. Changes prior to the next official version change w successful Serena call. Add a `serena-hooks reset` command and a `PostToolUse` example matched to Serena's own tools to close the gap (#1852) +* Dashboard: + - Fix: DashboardManager's unsupported-mode fallback warning logged the literal text + `{fallback_mode.value}` because only the first string fragment was an f-string + - Fix: On macOS, the tray manager refreshed the tray menu straight from the Flask request handlers + for `/register`, `/update_project` and `/unregister` and from the alive-check thread. That reaches + `NSStatusItem.setMenu_()` off the main thread, which AppKit forbids and which recent macOS + versions punish with SIGTRAP, so the tray-manager process died within seconds of every agent + start and the tray icon never became usable. Menu refreshes are now marshalled onto the main + thread (#2038) + * Language Servers: + - Fix: Dart analysis server no longer receives rootUri/rootPath, which added the monorepo root as an extra analysis root and could pin a CPU core at idle (#2045) + - Fix: The C# language server opened every `.csproj` found anywhere under the repository root, + without consulting the project's ignore settings. On repositories that vendor third-party or + sample C# projects, this loads projects the server cannot restore on every start, and their + restore failures bury the diagnostics of the projects the user actually works on. Project + discovery now skips `.csproj` files matched by the project's ignore patterns + - Kotlin: update the managed Kotlin LSP from `262.9593.0` to `263.4702.0`; the `262.9593.0` build + has expired and fails on startup with "This build of intellij-server has expired" (#2008) + - Fix: Godot's GDScript parser can report a symbol's end column one column past the + line-end convention every other language server follows (closing a node's range from + the next lookahead token instead of the last consumed one, when that lookahead is a + synthesized newline); `replace_symbol_body` on the last function in a file silently + consumed the separating blank line as a result. `GodotLanguageServer` now corrects this + specific, measured overshoot when building its high-level document symbols (#1974) + - Fix: High-level document symbol cache was not invalidated when the LS-specific low-level result + version changed + - Fix: A language server's cache directory was determined by the language_id rather than + the language server identifier's key. The two identifiers coincided in most cases. - Fix: TypeScript and VTS now disable automatic type acquisition as intended, while VTS preserves explicit user settings across initialization and configuration requests (#1989) VTS initialization options now override defaults per top-level key rather than replacing the @@ -67,6 +137,10 @@ Status of the `main` branch. Changes prior to the next official version change w its global state under ``~/Library``; Serena now gives the child process an isolated home-directory view via ``solidity_state_dir`` without changing the parent process's ``HOME`` (#1817) - Add Fatou support as an alternative Julia language server (`julia_fatou`) + - Fix: C# properties/fields whose type contains a literal `(`, e.g. a tuple type like + `(int X, string Y)`, had their name corrupted to include a trailing `:` because the + parenthesis in the type was mistaken for a method's parameter list; `find_symbol` on + the real name then returned nothing - Fix: Nextflow's `_flush_deferred_workspace_scan` marked the workspace scan flushed even when both of its `completion` probes failed, permanently skipping the flush (and silencing retries) for the rest of the session (#1871) @@ -111,7 +185,10 @@ CLI: - Fix `project index-file` command not using only the relevant language server to index the given file (#1965) * Dependencies: + - Fix: declare `click` as a direct dependency; all three console scripts (`serena`, `serena-agent`, + `serena-hooks`) import it but it was only available transitively - Remove the redundant `dotenv` dependency; the `dotenv` module is provided by `python-dotenv` + - Upgrade the `mcp` SDK from 1.28.1 to 2.2.0 # v1.7.0 (2026-08-09) @@ -264,7 +341,6 @@ CLI: `target_file`/`targetFile` file-path keys (shared payload parsing, applies to all hook clients). - Fix hook input parsing for clients that emit raw control characters in JSON string values #1743. - # v1.6.1 (2026-07-21) * General: diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index 2d8dae3b..d9df20eb 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -11,6 +11,9 @@ The following types of contributions can be submitted directly via pull requests For other changes, please open an issue first to discuss your ideas with the maintainers. +Do not submit pull requests for beta features (unless they are trivial bug fixes); instead, provide feedback via issues or discussions. +At present, the Serena REPL is a beta feature. + ## Licensing and Contributor License Agreement (CLA) Serena is multi-licensed by component (see [LICENSE](LICENSE)): @@ -44,7 +47,8 @@ See the corresponding [memory](.serena/memories/adding_new_language_support_guid ## Submitting Pull Requests -Before submitting a PR, be sure to document your relevant changes (i.e. new features, fixes) in `CHANGELOG.md`. +Before submitting a PR, be sure to document your relevant changes (i.e. new features, fixes) in `CHANGELOG.md`; +documentation changes should not be included. Use a concise style and add your change to the appropriate section ("Language Servers", "Tools", "JetBrains", "CLI", "Memories", "Dashboard", "Hooks", "General", "Security"). diff --git a/README-dev.md b/README-dev.md index 6207719a..d4c23de1 100644 --- a/README-dev.md +++ b/README-dev.md @@ -8,12 +8,18 @@ and tools for formatting and type checking. ## Release Process 1. Ensure clean git status. -2. Set the version for release, e.g. - - python scripts/bump_version.py --patch - python scripts/bump_version.py --minor +2. Set the version for release. Normally, the version to be released is the one already reserved by the + current `.dev0` version (e.g. `1.8.0` when the repository is at `1.8.0.dev0`): - This also creates the git tag. + python scripts/bump_version.py release current + + To release a version beyond the reserved one, name the part to bump instead, e.g. + + python scripts/bump_version.py release patch + python scripts/bump_version.py release minor + + This updates `CHANGELOG.md`, commits the release version, creates the git tag, and then + commits the subsequent `.dev0` version for the next iteration. 3. Push to GitHub: git push @@ -26,4 +32,19 @@ and tools for formatting and type checking. [GitHub Releases page](https://github.com/oraios/serena/releases). When ready, publish it (click *Publish release*). This triggers the `publish` workflow, which builds and publishes the - package to PyPI. \ No newline at end of file + package to PyPI. + +### Bumping the Development Version + +Independently of a release, the development version can be bumped, e.g. when work on `main` +begins to target a new minor or major version: + + python scripts/bump_version.py dev minor + python scripts/bump_version.py dev major + +This sets the version to the respective new `.dev0` version (e.g. `1.8.0.dev0`) and commits it as +"Set version to vX"; it creates no tag and does not modify `CHANGELOG.md`. + +The subsequent release of that version is then performed with `release current`. + +Both commands require a clean git status and support `--dry-run` to preview the changes. \ No newline at end of file diff --git a/docs/01-about/020_programming-languages.md b/docs/01-about/020_programming-languages.md index ede55fab..ce0db2db 100644 --- a/docs/01-about/020_programming-languages.md +++ b/docs/01-about/020_programming-languages.md @@ -42,7 +42,7 @@ Some languages require additional installations or setup steps, as noted. subsumes `typescript` and `html` for `.ts`/`.html` files, so do not also list those) * **Ansible** (experimental; requires Node.js and npm; automatically installs `@ansible/ansible-language-server`; - must be explicitly specified in the `languages` entry in the `project.yml`; requires `ansible` in PATH for full functionality) + must be explicitly specified in the `language_servers` entry in the `project.yml`; requires `ansible` in PATH for full functionality; the upstream `@ansible/ansible-language-server@1.2.3` supports hover, completion, definition, semantic tokens, and validation; document symbols, workspace symbols, references, and rename are not supported by this version) @@ -51,7 +51,7 @@ Some languages require additional installations or setup steps, as noted. (requires Java 21+ on PATH; uses [bsl-language-server](https://github.com/1c-syntax/bsl-language-server) by 1c-syntax; the JAR is auto-downloaded and SHA-256-verified for the bundled default version; supports `.bsl` and `.os` files; configure optional `ls_path` or `bsl_ls_version` under `ls_specific_settings.bsl`) * **C#** (by default, uses the Roslyn language server (language `csharp`), requiring [.NET v10+](https://dotnet.microsoft.com/en-us/download/dotnet) and, on Windows, `pwsh` ([PowerShell 7+](https://learn.microsoft.com/en-us/powershell/scripting/install/install-powershell-on-windows?view=powershell-7.5)); - set language to `csharp_omnisharp` to use OmiSharp instead) + set language to `csharp_omnisharp` to use OmniSharp instead) * **C/C++** (by default, uses the clangd language server (language `cpp`) but we also support ccls (language `cpp_ccls`); for best results, provide a `compile_commands.json` at the repository root; diff --git a/docs/01-about/070_privacy.md b/docs/01-about/070_privacy.md new file mode 100644 index 00000000..7160f262 --- /dev/null +++ b/docs/01-about/070_privacy.md @@ -0,0 +1,17 @@ +(privacy)= +# Privacy Policy + +Serena respects your privacy and is committed to protecting your personal information. + +When using Serena, no personal data or data about the project being worked on is sent to any external servers. +We collect only the following anonymous usage data whenever Serena is started: + + * the version of Serena being used, + * the operating system being used, + * the language backend being used, + * the enabled status of the Serena Dashboard, + * the enabled [Serena agent context](contexts) + +This data is collected strictly to help us understand Serena usage. + +If you want to opt out of usage data reporting, set the environment variable `SERENA_USAGE_REPORTING` to `false`. diff --git a/docs/02-usage/030_clients.md b/docs/02-usage/030_clients.md index da204dba..7f5bddad 100644 --- a/docs/02-usage/030_clients.md +++ b/docs/02-usage/030_clients.md @@ -124,8 +124,8 @@ When using Serena, we highly recommend that you start CC as claude --system-prompt="$(serena prompts print-cc-system-prompt-override)" ``` -You can also consider adding the content of `serena cc-system-prompt-override` to your `CLAUDE.md` files, -but the effect be insufficient for counteracting Claude Code's bias towards internal tools. +You can also consider adding the content of `serena prompts print-cc-system-prompt-override` to your `CLAUDE.md` files, +but the effect may be insufficient for counteracting Claude Code's bias towards internal tools. ::: **Global Configuration**. To add the Serena MCP server for all your projects, use the user-level configuration of claude code and the `--project-from-cwd` flag: diff --git a/docs/02-usage/045_memories.md b/docs/02-usage/045_memories.md index 1aa8e3b9..5c6b3129 100644 --- a/docs/02-usage/045_memories.md +++ b/docs/02-usage/045_memories.md @@ -80,7 +80,9 @@ This convention has two practical consequences: - **Renames keep references intact.** When you rename or move a memory with the `rename_memory` tool, Serena rewrites every `` `mem:OLD_NAME` `` occurrence across all memories to point to - the new name. References that do not use the `mem:` prefix will not be updated automatically. + the new name, except in memories matched by `read_only_memory_patterns`, which the agent cannot + write; `serena memories check` reports such a reference as stale. + References that do not use the `mem:` prefix will not be updated automatically. - **Integrity checks** (see [below](memory-cli)) report any `` `mem:NAME` `` whose target does not resolve to an existing memory, and propose similarly-named candidates as likely intended targets. diff --git a/docs/02-usage/050_configuration.md b/docs/02-usage/050_configuration.md index 980cc7a8..583c0174 100644 --- a/docs/02-usage/050_configuration.md +++ b/docs/02-usage/050_configuration.md @@ -26,7 +26,7 @@ Some of the configurable settings include: * the language backend to use by default (i.e., the JetBrains plugin or language servers); this can also be [overridden per project](per-project-language-backend) * UI settings affecting the [Serena Dashboard and GUI tool](060_dashboard.md) - * the set of tools to enable/disable by default + * the set of tools or REPL API functions to enable/disable by default * the set of [modes](modes) to use by default * tool execution parameters (timeout, max. answer length) * global ignore rules @@ -55,6 +55,37 @@ You can access it ```shell serena config edit ``` + +(agent-interfaces)= +### Agent Interfaces + +Serena provides its functionality to the agent (LLM) through one of two interfaces +(see [Tools and APIs](../01-about/035_tools) for the operations they offer): + +* **tools**: every operation is a separate tool of the MCP server. +* **REPL** (new in Serena v2): a single tool executes Python code, through which the agent accesses the operations + programmatically, being able to combine several of them in one call. + +The interface is selected via the `agent_interface` setting in the global configuration. +It can be overridden in the project configuration or via the `--agent-interface` command-line option, +and it is fixed for the duration of a session. + +The two interfaces are configured differently: + +* With the **tool interface**, the set of tools results from the tool inclusion/exclusion settings + (`excluded_tools`, `included_optional_tools`, `fixed_tools`) of the global configuration, the context, + the modes and the project configuration. +* With the **REPL interface**, the set of tools is fixed (the REPL tool and the tools which have no + counterpart within the REPL, e.g. for project activation); the tool settings above consequently do not apply. + The operations available *within* the REPL are configured via `included_apis`/`excluded_apis` instead, + which are supported in the same configuration layers and reference either a group of operations + (e.g. `lsp`) or an individual operation (e.g. `lsp.find_symbol`). + +```{note} +Restricting the operations available in the REPL is a means of steering the agent, not a security mechanism: +the Python code that is executed can, in principle, do anything the Serena process can do. +See [Security](070_security) for isolation options. +``` ## Modes and Contexts @@ -238,7 +269,7 @@ This ensures backward compatibility: existing projects that already have a `.ser Most users will not need to adjust these settings. ::: -Under the key `ls_specific_settings` in `serena_config.yml`, you can you pass global per-language, +Under the key `ls_specific_settings` in `serena_config.yml`, you can pass global per-language, language server-specific configuration. You can use the same key in the project configuration files (`project.yml` @@ -796,10 +827,10 @@ Supported settings: | Setting | Default | Description | |---|---|---| | `ls_path` | managed download | Override the Kotlin Language Server executable path. | -| `kotlin_lsp_version` | `262.9593.0` | Override the Kotlin Language Server version Serena downloads when `ls_path` is not set. | +| `kotlin_lsp_version` | `263.4702.0` | Override the Kotlin Language Server version Serena downloads when `ls_path` is not set. | | `jvm_options` | `-Xmx2G` | Value assigned to `JAVA_TOOL_OPTIONS` for the Kotlin LS process. Set to `""` to disable JVM options entirely. | -The managed `262.9593.0` packages include a bundled JBR. For a custom `ls_path`, point directly to +The managed `263.4702.0` packages include a bundled JBR. For a custom `ls_path`, point directly to `bin/intellij-server` (`bin/intellij-server.exe` on Windows). Serena also retains the legacy download layout for custom Kotlin LSP versions older than `262.4739.0`. The pinned current and frozen initial releases are checksum-verified; arbitrary custom versions are downloaded without checksum verification. @@ -809,7 +840,7 @@ Example: ```yaml ls_specific_settings: kotlin: - kotlin_lsp_version: "262.9593.0" + kotlin_lsp_version: "263.4702.0" jvm_options: "-Xmx4G -XX:+UseG1GC" ``` @@ -1225,6 +1256,19 @@ Supported settings: | `server_ready_timeout` | `10.0` | Timeout in seconds for waiting on the server-ready signal after initialization. If the signal does not arrive within this window, Serena logs a message and proceeds anyway. | | `indexing_start_grace` | `5.0` | Timeout in seconds to wait for tsserver to *start* reporting `$/progress` before the first cross-file reference query. tsserver must resolve the project graph before it can emit the first progress token, and that can take longer than the default on a very large project; if it takes longer than this window, Serena assumes no indexing was needed and may return incomplete cross-file references. Raising `indexing_timeout` alone does not help here, since this grace elapses first. Increase this for very large projects if `find_referencing_symbols`/`request_references` returns incomplete results shortly after project load. | +##### TypeScript monorepos and cross-package references + +In a monorepo, `find_referencing_symbols` / `find_references` only include consumers in other packages when tsserver can walk from a package's declaration file back to its sources. That walk requires [TypeScript project references](https://www.typescriptlang.org/docs/handbook/project-references.html) (`composite` + `references`), not merely a solution-style root `tsconfig.json` or `package.json` `exports`. + +Without those edges, results are **silently partial**: a symbol may show only same-package references (or none) even though other packages import it. This is tsserver behaviour Serena inherits, not a Serena bug ([microsoft/TypeScript#30823](https://github.com/microsoft/TypeScript/issues/30823); oraios/serena#1939). + +What to do in a TypeScript monorepo: + +- Declare `composite: true` in each library package's `tsconfig.json` and list dependent projects under `references` in the consumer (or a solution-style root). +- Prefer source imports (or generate declaration maps) so tsserver can map `dist/*.d.ts` back to sources. +- After changing the project graph, restart Serena (or the TypeScript language server) so tsserver rebuilds the program. +- If cross-package references still look short, verify with grep before treating the LSP answer as complete; same-package results being complete does not imply the package boundary was crossed. + #### Svelte Serena uses `svelte-language-server` for the `svelte` language key. Use `svelte` for Svelte projects instead of also listing `typescript`, unless you intentionally want multiple language servers active for the same files. @@ -1334,8 +1378,6 @@ It is advisable to use the default prompt as a starting point and modify it to s ### Usage Reporting -On startup, Serena reports anonymous usage data to help us understand Serena usage. -Specifically, we collect the Serena version, the operating system & language backend being used as well as the dashboard enabled status. -No personally identifiable information or project-specific information is collected. +On startup, Serena reports anonymous usage data to help us understand Serena usage, as explained in our [privacy policy](privacy). -If you want to opt out of usage reporting, set the environment variable `SERENA_USAGE_REPORTING` to `false`. +If you want to opt out of usage data reporting, set the environment variable `SERENA_USAGE_REPORTING` to `false`. diff --git a/docs/02-usage/070_security.md b/docs/02-usage/070_security.md index 249409cf..c8e6b652 100644 --- a/docs/02-usage/070_security.md +++ b/docs/02-usage/070_security.md @@ -23,18 +23,22 @@ However, reports which amount to noting that Serena's tools can execute commands describe intended functionality rather than vulnerabilities, and we will reject advisories that fail to recognise this or otherwise ignore the above assumptions. Sandboxing is the *only* way to fully protect against unintended consequences when using coding agents; -constraints on the tools themselves cannot achieve this and are therefore not an approach we pursue. +constraints on the tools or REPL APIs cannot achieve this and are therefore not an approach we pursue. ::: ## General Recommendations for Risk Reduction To reduce the risk of unintended consequences, we recommend that you: - back up your work regularly (keep the project being worked on under version control), -- restrict the set of allowed tools via the [configuration](050_configuration), +- restrict the set of allowed tools via the [configuration](050_configuration) + (note that restrictions are effective for the tool interface only, see [below](repl-security)), - do not expose [Serena's network services](network-security) to untrusted networks. If you do not fully trust the client/the LLM, we additionally recommend to monitor tool executions carefully (provided that your MCP client supports this). +Note that with the REPL interface, such monitoring is necessarily coarser: every action appears as the same tool +being called, and the client cannot tell a read-only call from a modifying one by the tool's name alone. +What needs to be reviewed is the submitted code. (sandboxing)= ## Sandboxing @@ -97,6 +101,28 @@ the introduction of this setting retain a pattern that trusts all projects, ensu not broken, whereas newly created configurations trust no project by default. The applicable value can be inspected in the dashboard. +(repl-security)= +## The REPL Interface + +With the [REPL interface](agent-interfaces), the agent does not invoke individual tools; it submits Python code, +which Serena executes. +In the default configuration, Serena is equally capable either way, as shell execution and file modification are +available in both interfaces. +The difference lies in what restrictions can achieve: + +- With the **tool interface**, excluding a tool removes the respective capability: a tool that is not exposed + cannot be invoked, so forbidding shell execution or file modification is effective. +- With the **REPL interface**, there is no such guarantee. + Restricting the available operations (`included_apis`/`excluded_apis`) steers the agent towards the intended + way of working, but the submitted code can, in principle, do anything the Serena process can do — irrespective + of the operations Serena itself provides. + Excluding shell execution, for example, does not prevent the code from achieving the same effect by other means. +- The `read_only` project setting is subject to the same limitation: Serena's own editing operations are refused, + yet code executed in the REPL is not prevented from modifying files. + +The assumptions stated above therefore apply unchanged, but if you require actual constraints rather than +guidance, [sandboxing](sandboxing) is the answer. + (network-security)= ## Network Security diff --git a/docs/03-special-guides/cpp_setup.md b/docs/03-special-guides/cpp_setup.md index dd5622dd..606c26a3 100644 --- a/docs/03-special-guides/cpp_setup.md +++ b/docs/03-special-guides/cpp_setup.md @@ -40,8 +40,10 @@ You can customize this location via project settings: ```yaml # .serena/project.yml language_servers: + - cpp +ls_specific_settings: cpp: - compile_commands_dir: custom/rel/path (defaults to .serena) + compile_commands_dir: custom/rel/path # defaults to .serena ``` ### With ccls @@ -76,7 +78,7 @@ choco install ccls #### Configuration After installing ccls, configure Serena to use it via project settings (in `.serena/project.yml`) -by adding `cpp_ccls` to the `languages` list. Replace `cpp` with `cpp_ccls` if you already have the `cpp` entry. +by adding `cpp_ccls` to the `language_servers` list. Replace `cpp` with `cpp_ccls` if you already have the `cpp` entry. ccls can handle relative paths in `compile_commands.json`, so no transformation is necessary and no transformed `compile_commands.json` file will be created. diff --git a/docs/autogen_docs.py b/docs/autogen_docs.py index 43976fd9..3d613d4a 100644 --- a/docs/autogen_docs.py +++ b/docs/autogen_docs.py @@ -150,30 +150,112 @@ def autogen_tool_list(target_filename = "01-about/035_tools.md"): from serena.tools import ToolRegistry target_file = Path(__file__).parent / target_filename - with open(target_file, "w") as f: + with open(target_file, "w", encoding="utf-8") as f: f.write("\n\n") - f.write("# Tools\n\n") - f.write("Find the full list of Serena's tools below.\n\n") - f.write("Note that in most configurations, only a subset of these tools will be enabled simultaneously.\n") - f.write("Tools marked as *optional* are disabled by default.\n\n") - f.write("Tools marked as [BETA] were recently introduced and may not be fully robust yet.\n\n") - tools_by_module = ToolRegistry().get_registered_tools_by_module() - priority_modules = {"serena.tools.symbol_tools": 1, "serena.tools.jetbrains_tools": 2} + f.write("# Tools and APIs\n\n") + f.write( + "Serena provides an agent (LLM) with its functionality through one of two interfaces (configured in Serena's [global configuration](global-config)):\n\n" + "* **Tools** (the classic interface): every operation is a separate tool of the MCP server.\n" + "* **REPL** (new in Serena v2): a single tool executes Python code, through which the agent accesses the operations\n" + " programmatically. The agent can thus combine several operations in one call, process the results\n" + " in code and return only the information it actually needs.\n\n" + "While both interfaces offer the same general functionality for the most part, " + "the REPL interface addresses several limitations inherent in the tool-based approach (see [advantages](repl-advantages) below).\n\n" + ) + f.write("\n\n:::{note}\nThe Serena REPL is an unreleased BETA feature. Please provide feedback; if you encounter issues, report them.\n:::\n\n") - text = TextBuilder() - sorted_modules = sorted(tools_by_module.keys(), key=lambda m: (priority_modules.get(m, 3), m)) - for module in sorted_modules: - tools = tools_by_module[module] - module = module.replace("serena.tools.", "") - text.with_line(f"* **{module}**") - for tool in tools: - info = "" - if tool.is_optional: - info += " *(optional)*" - if tool.is_beta: - info += " [BETA]" - text.with_line(f"* `{tool.tool_name}`{info}: {tool.class_docstring}", indent=2) - f.write(text.build()) + + def tools_section(): + f.write("## Tools (Classic Interface)\n\n") + f.write("Find the full list of Serena's tools below.\n\n") + f.write("Note that in most configurations, only a subset of these tools will be enabled simultaneously.\n") + f.write("Tools marked as *optional* are disabled by default.\n\n") + f.write("Tools marked as [BETA] were recently introduced and may not be fully robust yet.\n\n") + tools_by_module = ToolRegistry().get_registered_tools_by_module() + priority_modules = {"serena.tools.symbol_tools": 1, "serena.tools.jetbrains_tools": 2} + + text = TextBuilder() + sorted_modules = sorted(tools_by_module.keys(), key=lambda m: (priority_modules.get(m, 3), m)) + for module in sorted_modules: + tools = tools_by_module[module] + module = module.replace("serena.tools.", "") + text.with_line(f"* **{module}**") + for tool in tools: + info = "" + if tool.is_optional: + info += " *(optional)*" + if tool.is_beta: + info += " [BETA]" + text.with_line(f"* `{tool.tool_name}`{info}: {tool.class_docstring}", indent=2) + f.write(text.build()) + + def facades_section(): + from serena.repl.facade import ApiScope + from serena.agent import SerenaAgent, SerenaConfig + from serena.language_backend import BuiltinLanguageBackend + + f.write("\n\n## Serena's REPL (Code Execution-Based Interface)\n\n") + f.write( + "With the REPL interface, the agent uses a single tool, which executes Python code. The code accesses\n" + "Serena's functionality through the entrypoint object `s`, whose attributes are *facades*, each of which\n" + "groups the operations of one domain. For instance, this code finds a class and returns the names of its\n" + "members in a single call:\n\n" + "```python\n" + 'result = s.lsp.find_symbol("SerenaAgent", depth=1)\n' + "[member.name for member in result.symbols[0].iter_children()]\n" + "```\n\n" + "The agent retrieves the details it requires (signatures, documentation, result types) at runtime via\n" + "`s.info(...)` (principle of *progressive disclosure*).\n\n" + "As with tools, only a subset of the operations is available in a given configuration:\n\n" + "* The facades `lsp` and `jb` are mutually exclusive, being tied to the respective language backend.\n" + "* Operations marked as *optional*, as well as all operations of facades marked as *optional*,\n" + " are disabled by default and must be enabled explicitly.\n\n" + "(repl-advantages)=\n\n" + "### Advantages\n\n" + "* **Composition**: The agent can combine several operations in a single call, using control flow,\n" + " filtering and aggregation. Multi-step retrievals that would otherwise require a series of\n" + " round-trips (find a symbol, inspect its members, find their references) become one call.\n" + "* **Context economy**: Only the result the agent actually needs enters the conversation.\n" + " Intermediate results remain in the Python runtime instead of consuming the context window\n" + " — which, for large result sets, is the difference between a summary and thousands of tokens.\n" + "* **Results as objects**: Operations return objects which can be processed programmatically\n" + " rather than plain text.\n" + "* **Progressive disclosure**: Instead of the schemas of dozens of tools, the agent is given a list of\n" + " the available operations, retrieving the details it requires (signatures, documentation, result\n" + " types) on demand. This eliminates the need for a client-specific tool search/dynamic tool discovery mechanism.\n" + "* **Reusability within a session**: Variables and helper functions defined by the agent persist\n" + " across calls, allowing intermediate results to be revisited and recurring logic to be applied\n" + " repeatedly.\n" + "* **Dynamic adaptation**: The set of available operations can change at runtime.\n" + " For MCP tools, changes to the toolset are not widely supported by clients and are therefore not\n" + " applied; within the REPL, the operations can follow the current situation.\n" + " This is what allows the language backend to be switched during a session, e.g. when a project\n" + " whose configuration demands a different backend is activated..\n\n" + "### List of Facades\n\n" + ) + api_scope = ApiScope() + agent = SerenaAgent(serena_config=SerenaConfig().with_headless_mode_overrides()) + facades = [] + for backend in BuiltinLanguageBackend: + facades.extend(backend.get_instance().create_facades(agent, api_scope)) + facades.extend(agent.create_default_facade_list(api_scope)) + + text = TextBuilder() + for facade in facades: + facade_info = " *(optional)*" if facade.is_optional() else "" + text.with_line(f"* **{facade.name}**{facade_info}: {facade.description}") + for method in facade.get_methods(): + method_info = "" + if method.info.optional: + method_info += " *(optional)*" + if method.info.beta: + method_info += " [BETA]" + summary = method.get_summary().replace("`", "") + text.with_line(f"* `{method.qualified_name}`{method_info}: {summary}", indent=2) + f.write(text.build()) + + tools_section() + facades_section() def autogen_about_intro_features(): diff --git a/pyproject.toml b/pyproject.toml index 74a7cc1f..e83145bf 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,7 +4,7 @@ requires = ["hatchling"] [project] name = "serena-agent" -version = "1.7.1.dev0" +version = "2.0.0.dev0" description = "A powerful MCP toolkit for coding, providing semantic retrieval and editing capabilities - the IDE for your agent" authors = [{ name = "Oraios AI", email = "info@oraios-ai.de" }] readme = "README.md" @@ -21,7 +21,8 @@ dependencies = [ "requests==2.33.0", "overrides==7.7.0", "python-dotenv==1.2.2", - "mcp==1.28.1", + "mcp==2.2.0", + "click==8.3.1", "flask==3.1.3", # bumped from 3.1.1 for CVE fix (also fixes werkzeug alert) "sensai-utils==1.5.0", "pydantic==2.12.5", @@ -107,17 +108,15 @@ packages = ["src/serena", "src/interprompt", "src/solidlsp"] max-line-length = 1000 [tool.ty.environment] -python-version = "3.11" +python-version = "3.11" # We configure the oldest Python version supported by Serena # Analyze for all platforms rather than defaulting to the OS ty happens to run on. This keeps the # check deterministic across the CI matrix (Linux/Windows/macOS) and lets platform-conditional stdlib # members (e.g. subprocess.CREATE_NO_WINDOW, ctypes.windll, pwd) resolve without per-OS type-ignores. python-platform = "all" [tool.ty.rules] -# Mirror mypy's ignore_missing_imports=true: optional extras (e.g. agno, google-genai) and -# platform-specific modules (e.g. AppKit on macOS) are not installed in the default dev environment, -# so we do not want unresolvable imports to fail the type check. -unresolved-import = "ignore" +# Unresolvable imports do fail the type check (catching broken first-party imports); the exceptions for +# optional extras and platform-specific modules are handled per file/line below. possibly-missing-submodule = "ignore" [tool.ty.src] @@ -127,6 +126,15 @@ possibly-missing-submodule = "ignore" # test/resources is instead excluded from the `ty check test` CLI task via its --exclude flag. exclude = ["build/", "docs/"] +[[tool.ty.overrides]] +# Modules whose imports cannot be resolved in the default dev environment: the agno integration depends +# on the optional extra `agno`, and the pywebview integration imports macOS-only modules (AppKit, PyObjCTools) +# in several places. Elsewhere, individual platform-specific imports are suppressed inline. +include = ["src/serena/agno.py", "src/serena/util/pywebview.py"] + +[tool.ty.overrides.rules] +unresolved-import = "ignore" + [[tool.ty.overrides]] # Test code is heavily dynamic (pytest fixtures, MagicMock, intentionally loose Optionals). ty models # pytest's fail/skip helpers and MagicMock far more strictly than mypy did (mypy inferred `Any` for @@ -145,7 +153,6 @@ no-matching-overload = "ignore" not-subscriptable = "ignore" parameter-already-assigned = "ignore" too-many-positional-arguments = "ignore" -unresolved-attribute = "ignore" unsupported-operator = "ignore" [tool.poe.env] diff --git a/scripts/bump_version.py b/scripts/bump_version.py index 08af6b33..ac58973f 100644 --- a/scripts/bump_version.py +++ b/scripts/bump_version.py @@ -17,98 +17,110 @@ from serena.util.git import get_git_status log = logging.getLogger(__name__) VersionPart = Literal["major", "minor", "patch"] +#: a version part to bump or, in the case of "current", the version already reserved by the current .dev version +VersionTarget = Literal["major", "minor", "patch", "current"] _VERSION_PATTERN = re.compile(r"^(?P\d+)\.(?P\d+)\.(?P\d+)(\.\w+)?$") _INIT_VERSION_PATTERN = re.compile(r'^(?P__version__\s*=\s*")(?P\d+\.\d+\.\d+(?:\.\w+)?)(?P"\s*)$', re.MULTILINE) _PYPROJECT_VERSION_PATTERN = re.compile( r'(?m)^(?P\[project\]\n(?:.*\n)*?^version\s*=\s*")(?P\d+\.\d+\.\d+(?:\.\w+)?)(?P"\s*)$' ) +_VERSION_SUFFIX_PATTERN = re.compile(r"^\d+\.\d+\.\d+\.(?P\w+)$") _UNRELEASED_HEADER = "# Unreleased (main)\n" -@click.command() -@click.option("--major", "major", is_flag=True, help="Bump the major version and reset minor and patch to 0.") -@click.option("--minor", "minor", is_flag=True, help="Bump the minor version and reset patch to 0.") -@click.option("--patch", "patch", is_flag=True, help="Bump the patch version.") -@click.option("--version", "-v", "target_version", metavar="X.Y.Z", help="Set an explicit version instead of bumping.") -@click.option("--dry-run", is_flag=True, help="Show what would change without writing any files.") -def bump_version(major: bool, minor: bool, patch: bool, target_version: str | None, dry_run: bool) -> None: - git_status = get_git_status() - if not git_status.is_clean: - raise click.ClickException("Working directory is not clean. Please commit or stash your changes first.") +_version_target_argument = click.argument("version_target", type=click.Choice(["current", "major", "minor", "patch"])) +_version_part_argument = click.argument("version_part", type=click.Choice(["major", "minor", "patch"])) +_dry_run_option = click.option("--dry-run", is_flag=True, help="Show what would change without writing any files.") - log.info("bump_version called: major=%s, minor=%s, patch=%s, target_version=%s", major, minor, patch, target_version) - # determine part to bump - version_part = resolve_version_selection(major=major, minor=minor, patch=patch, target_version=target_version) - log.info("Resolved version_part=%s", version_part) +@click.group() +def cli() -> None: + """Manages the Serena version.""" + + +@cli.command() +@_version_target_argument +@_dry_run_option +def release(version_target: VersionTarget, dry_run: bool) -> None: + """Bumps the version for a release and starts the next dev iteration. + + Bumps the version, updates the changelog, commits and tags the release, and then commits + the subsequent .dev0 version. + + VERSION_TARGET is either "current", releasing the version already reserved by the current .dev version + (the usual case), or the part of the version to bump beyond it (major, minor or patch). + """ + require_clean_working_directory() + log.info("release called: version_target=%s", version_target) - # bump it (never incrementing patch because it was already updated with the last .dev version) repo_root = find_repo_root() log.info("Repo root: %s", repo_root) - new_version = bump_repo_version( - repo_root, version_part=version_part, target_version=target_version, dry_run=dry_run, increment_patch=False - ) - log.info("New version: %s", new_version) - # commit and tag for new version + # bump to the release version + new_version = bump_repo_version(repo_root, version_target=version_target, dry_run=dry_run) + log.info("New version: %s", new_version) if dry_run: click.echo(f"Dry run complete. Version would be bumped to {new_version}") return - else: - os.system("uv lock") - click.echo(f"Bumped version to {new_version}") - os.system("git add -u") - os.system(f'git commit -m "Release v{new_version}"') - os.system(f"git tag v{new_version}") - # bump patch and add suffix for next dev iteration - new_snapshot_version = bump_repo_version( - repo_root, - version_part="patch", - target_version=None, - dry_run=dry_run, - target_version_suffix=".dev0", - increment_patch=True, - ) + # commit and tag the release version + commit_version_change(new_version, message=f"Release v{new_version}") + os.system(f"git tag v{new_version}") + + # bump patch and add the suffix for the next dev iteration + new_snapshot_version = bump_repo_version(repo_root, version_target="patch", dry_run=dry_run, target_version_suffix=".dev0") log.info("New snapshot version: %s", new_snapshot_version) + commit_version_change(new_snapshot_version, message=f"Set version to v{new_snapshot_version}") - # commit the new snapshot version + +@cli.command() +@_version_part_argument +@_dry_run_option +def dev(version_part: VersionPart, dry_run: bool) -> None: + """Bumps the development version without creating a release. + + Sets the version to a new .dev0 version and commits it; no tag is created and the changelog + is not modified. + + VERSION_PART is the part of the version to bump (major, minor or patch). + """ + require_clean_working_directory() + log.info("dev called: version_part=%s", version_part) + + repo_root = find_repo_root() + log.info("Repo root: %s", repo_root) + + new_version = bump_repo_version(repo_root, version_target=version_part, dry_run=dry_run, target_version_suffix=".dev0") + log.info("New version: %s", new_version) + if dry_run: + click.echo(f"Dry run complete. Version would be bumped to {new_version}") + return + + commit_version_change(new_version, message=f"Set version to v{new_version}") + + +def require_clean_working_directory() -> None: + if not get_git_status().is_clean: + raise click.ClickException("Working directory is not clean. Please commit or stash your changes first.") + + +def commit_version_change(new_version: str, *, message: str) -> None: os.system("uv lock") - click.echo(f"Bumped version to {new_snapshot_version}") + click.echo(f"Bumped version to {new_version}") os.system("git add -u") - os.system(f'git commit -m "Set version to v{new_snapshot_version}"') + os.system(f'git commit -m "{message}"') def find_repo_root() -> Path: return Path(REPO_ROOT) -def resolve_version_selection(*, major: bool, minor: bool, patch: bool, target_version: str | None) -> VersionPart | None: - bump_flags_selected = sum([major, minor, patch]) - if target_version is not None and bump_flags_selected > 0: - raise click.ClickException("Use either --version or one of --major/--minor/--patch, not both.") - if bump_flags_selected > 1: - raise click.ClickException("Use only one of --major, --minor, or --patch.") - if target_version is not None: - validate_version_string(target_version) - return None - if major: - return "major" - if minor: - return "minor" - if patch: - return "patch" - raise click.ClickException("No version bump selected. Use --major, --minor, --patch or --version.") - - def bump_repo_version( repo_root: Path, *, - version_part: VersionPart | None, - target_version: str | None, + version_target: VersionTarget, dry_run: bool = False, target_version_suffix: str | None = None, - increment_patch: bool = True, ) -> str: pyproject_path = repo_root / "pyproject.toml" init_path = repo_root / "src" / "serena" / "__init__.py" @@ -130,12 +142,12 @@ def bump_repo_version( f"Version mismatch between pyproject.toml and src/serena/__init__.py: {current_version} != {init_version}" ) - if target_version is not None: - new_version = validate_version_string(target_version) - else: - if version_part is None: - raise click.ClickException("No version target specified.") - new_version = increment_version(current_version, version_part, increment_patch=increment_patch) + if version_target == "current" and _VERSION_SUFFIX_PATTERN.search(current_version) is None: + raise click.ClickException( + f"The current version {current_version} is not a development version, so there is no reserved version to release. " + f"Use major, minor or patch to bump the version instead." + ) + new_version = increment_version(current_version, version_target) if target_version_suffix is not None: new_version += target_version_suffix log.info("New version will be: %s", new_version) @@ -199,7 +211,14 @@ def replace_version(text: str, pattern: re.Pattern[str], new_version: str, file_ return f"{text[: match.start('version')]}{new_version}{text[match.end('version') :]}" -def increment_version(version: str, version_part: VersionPart, increment_patch: bool) -> str: +def increment_version(version: str, version_target: VersionTarget) -> str: + """ + Computes the new version, dropping any development suffix of the given version. + + :param version: the current version + :param version_target: the part of the version to bump or "current" to keep the version as is + :return: the new version + """ match = _VERSION_PATTERN.fullmatch(version) if match is None: raise click.ClickException(f"Unsupported version format: {version}") @@ -208,22 +227,17 @@ def increment_version(version: str, version_part: VersionPart, increment_patch: minor = int(match.group("minor")) patch = int(match.group("patch")) - if version_part == "major": - return f"{major + 1}.0.0" - if version_part == "minor": - return f"{major}.{minor + 1}.0" - elif version_part == "patch": - if increment_patch: - patch += 1 - return f"{major}.{minor}.{patch}" - else: - raise ValueError(version_part) - - -def validate_version_string(version: str) -> str: - if _VERSION_PATTERN.fullmatch(version) is None: - raise click.ClickException(f"Unsupported version format: {version}") - return version + match version_target: + case "major": + return f"{major + 1}.0.0" + case "minor": + return f"{major}.{minor + 1}.0" + case "patch": + return f"{major}.{minor}.{patch + 1}" + case "current": + return f"{major}.{minor}.{patch}" + case _: + raise ValueError(version_target) def update_changelog(changelog_text: str, new_version: str) -> str: @@ -278,4 +292,4 @@ def split_unreleased_body(unreleased_body: str) -> tuple[str, str]: if __name__ == "__main__": logging.basicConfig(level=logging.DEBUG, format="%(levelname)s %(name)s: %(message)s") log.info("Script starting") - bump_version() + cli() diff --git a/scripts/demo_diagnostics.py b/scripts/demo_diagnostics.py index 02d12929..2ead12a4 100644 --- a/scripts/demo_diagnostics.py +++ b/scripts/demo_diagnostics.py @@ -14,8 +14,9 @@ from pathlib import Path from pprint import pprint from serena.agent import SerenaAgent -from serena.config.serena_config import LanguageBackend, ProjectConfig, RegisteredProject, SerenaConfig +from serena.config.serena_config import ProjectConfig, RegisteredProject, SerenaConfig from serena.constants import REPO_ROOT +from serena.language_backend import BuiltinLanguageBackend from serena.project import Project from serena.tools import ( CreateTextFileTool, @@ -35,7 +36,7 @@ def make_agent() -> SerenaAgent: """Create an LSP-backed Serena agent for the Serena repository.""" serena_config = SerenaConfig.from_config_file() serena_config.web_dashboard = False - serena_config.language_backend = LanguageBackend.LSP + serena_config.set_builtin_language_backend(BuiltinLanguageBackend.LSP) project = Project( project_root=str(REPO_PATH), diff --git a/scripts/demo_find_defining_symbol.py b/scripts/demo_find_defining_symbol.py index 001f001d..92107fe7 100644 --- a/scripts/demo_find_defining_symbol.py +++ b/scripts/demo_find_defining_symbol.py @@ -9,8 +9,9 @@ from pathlib import Path from pprint import pprint from serena.agent import SerenaAgent -from serena.config.serena_config import LanguageBackend, ProjectConfig, RegisteredProject, SerenaConfig +from serena.config.serena_config import ProjectConfig, RegisteredProject, SerenaConfig from serena.constants import REPO_ROOT +from serena.language_backend import BuiltinLanguageBackend from serena.project import Project from serena.tools import FindDeclarationTool from solidlsp.ls_config import LanguageServerId @@ -24,7 +25,7 @@ def make_agent(project_root: Path, language: LanguageServerId, project_name: str """Create an LSP-backed Serena agent for a single explicit project.""" serena_config = SerenaConfig.from_config_file() serena_config.web_dashboard = False - serena_config.language_backend = LanguageBackend.LSP + serena_config.set_builtin_language_backend(BuiltinLanguageBackend.LSP) project = Project( project_root=str(project_root), diff --git a/scripts/demo_find_implementing_symbol.py b/scripts/demo_find_implementing_symbol.py index 30402f25..87ea6c62 100644 --- a/scripts/demo_find_implementing_symbol.py +++ b/scripts/demo_find_implementing_symbol.py @@ -8,8 +8,9 @@ from pathlib import Path from pprint import pprint from serena.agent import SerenaAgent -from serena.config.serena_config import LanguageBackend, ProjectConfig, RegisteredProject, SerenaConfig +from serena.config.serena_config import ProjectConfig, RegisteredProject, SerenaConfig from serena.constants import REPO_ROOT +from serena.language_backend import BuiltinLanguageBackend from serena.project import Project from serena.tools import FindImplementationsTool from solidlsp.ls_config import LanguageServerId @@ -22,7 +23,7 @@ def make_agent(project_root: Path, language: LanguageServerId, project_name: str """Create an LSP-backed Serena agent for a single explicit project.""" serena_config = SerenaConfig.from_config_file() serena_config.web_dashboard = False - serena_config.language_backend = LanguageBackend.LSP + serena_config.set_builtin_language_backend(BuiltinLanguageBackend.LSP) project = Project( project_root=str(project_root), diff --git a/scripts/demo_progressive_tool_shortening.py b/scripts/demo_progressive_tool_shortening.py index 9af742ef..0e06f151 100644 --- a/scripts/demo_progressive_tool_shortening.py +++ b/scripts/demo_progressive_tool_shortening.py @@ -11,8 +11,9 @@ import json from pprint import pprint from serena.agent import SerenaAgent -from serena.config.serena_config import LanguageBackend, SerenaConfig +from serena.config.serena_config import SerenaConfig from serena.constants import REPO_ROOT +from serena.language_backend import BuiltinLanguageBackend from serena.tools import ( FindReferencingSymbolsTool, FindSymbolTool, @@ -165,16 +166,16 @@ def run_jb_tools(agent: SerenaAgent) -> None: ) -def make_agent(backend: LanguageBackend) -> SerenaAgent: +def make_agent(backend: BuiltinLanguageBackend) -> SerenaAgent: config = SerenaConfig.from_config_file() config.web_dashboard = False - config.language_backend = backend + config.set_builtin_language_backend(backend) return SerenaAgent(project=REPO_ROOT, serena_config=config) if __name__ == "__main__": # LSP backend - lsp_agent = make_agent(LanguageBackend.LSP) + lsp_agent = make_agent(BuiltinLanguageBackend.LSP) try: run_lsp_tools(lsp_agent) run_backend_independent_tools(lsp_agent) @@ -183,7 +184,7 @@ if __name__ == "__main__": # JetBrains backend (requires a running IDE) try: - jb_agent = make_agent(LanguageBackend.JETBRAINS) + jb_agent = make_agent(BuiltinLanguageBackend.JETBRAINS) try: run_jb_tools(jb_agent) finally: diff --git a/scripts/demo_run_tools.py b/scripts/demo_run_tools.py index d6c61f0e..2dc6efbe 100644 --- a/scripts/demo_run_tools.py +++ b/scripts/demo_run_tools.py @@ -9,8 +9,9 @@ from pathlib import Path from pprint import pprint from serena.agent import SerenaAgent -from serena.config.serena_config import LanguageBackend, SerenaConfig +from serena.config.serena_config import SerenaConfig from serena.constants import REPO_ROOT +from serena.language_backend import BuiltinLanguageBackend from serena.tools import ( FindFileTool, FindReferencingSymbolsTool, @@ -26,7 +27,7 @@ from serena.tools import ( if __name__ == "__main__": serena_config = SerenaConfig.from_config_file() serena_config.web_dashboard = False - serena_config.language_backend = LanguageBackend.LSP + serena_config.set_builtin_language_backend(BuiltinLanguageBackend.LSP) # project = Path(REPO_ROOT).parent / "serena-jetbrains-plugin-copy" project = Path(REPO_ROOT) agent = SerenaAgent(project=str(project), serena_config=serena_config) diff --git a/src/serena/__init__.py b/src/serena/__init__.py index 3b043bab..0d2c577e 100644 --- a/src/serena/__init__.py +++ b/src/serena/__init__.py @@ -1,6 +1,6 @@ # SPDX-License-Identifier: GPL-3.0-or-later -__version__ = "1.7.1.dev0" +__version__ = "2.0.0.dev0" import logging diff --git a/src/serena/agent.py b/src/serena/agent.py index c43d7198..d7e45fb7 100644 --- a/src/serena/agent.py +++ b/src/serena/agent.py @@ -17,7 +17,7 @@ from dataclasses import dataclass from datetime import datetime from enum import Enum from logging import Logger -from typing import TYPE_CHECKING, Optional, TypeVar +from typing import TYPE_CHECKING, Optional, TypeVar, cast import requests import webview @@ -31,10 +31,11 @@ from serena import serena_version from serena.analytics import RegisteredTokenCountEstimator, ToolUsageStats from serena.config.context_mode import SerenaAgentContext, SerenaAgentMode from serena.config.serena_config import ( - LanguageBackend, + AgentInterface, ModeSelectionDefinition, ModeSelectionDefinitionWithAddedModes, ModeSelectionDefinitionWithBaseModes, + NamedApiInclusionDefinition, NamedToolInclusionDefinition, RegisteredProject, SerenaConfig, @@ -42,19 +43,30 @@ from serena.config.serena_config import ( ToolInclusionDefinition, ) from serena.dashboard import SerenaDashboardAPI, SerenaDashboardTrayManager, SerenaDashboardViewer, open_url_in_browser -from serena.jetbrains import launch_coordinator as jetbrains_launch_coordinator +from serena.language_backend import BuiltinLanguageBackend, LanguageBackend from serena.ls_manager import LanguageServerManager from serena.memories.memory_manager import MemoryManager from serena.project import Project from serena.prompt_factory import SerenaPromptFactory +from serena.repl.api.cfg_api import ConfigApi +from serena.repl.api.edit_api import EditApi +from serena.repl.api.ext_api import ExternalProjectsApi +from serena.repl.api.fs_api import FsApi +from serena.repl.api.mem_api import MemoryApi +from serena.repl.api.shell_api import ShellApi +from serena.repl.facade import ApiScope, Facade +from serena.repl.repl import SerenaRepl +from serena.session import SerenaSession, SessionRegistry from serena.task_executor import TaskExecutor from serena.tools import ( ActivateProjectTool, GetCurrentConfigTool, + InitialInstructionsTool, OnboardingTool, OpenDashboardTool, ReadMemoryTool, ReplaceContentTool, + SerenaReplTool, Tool, ToolMarker, ToolRegistry, @@ -427,7 +439,7 @@ class DashboardManager: fallback_mode = self.Mode.from_platform() log.warning( f"Dashboard interface mode '{mode.value}' is not supported on the current platform; " - "falling back to '{fallback_mode.value}'." + f"falling back to '{fallback_mode.value}'." ) mode = fallback_mode @@ -567,9 +579,12 @@ class SerenaAgent: self._gui_log_viewer: Optional["GuiLogViewer"] = None self._dashboard_manager: DashboardManager | None = None self._project_prompt_status = ProjectPromptProvisionStatus() + self._session_registry = SessionRegistry() self._session_mode_selection_definition = modes self.version = serena_version() self._config_changed_callbacks: list[Callable[[], None]] = [] + self._repl: SerenaRepl | None = None + self._prompt_params: SerenaAgent.PromptParams | None = None # obtain serena configuration using the decoupled factory function self.serena_config = serena_config or SerenaConfig.from_config_file() @@ -645,14 +660,19 @@ class SerenaAgent: # determine the effective language backend for this session. # If a startup project is provided and has a per-project override, use it; otherwise use the global config. - # Since we don't want to change the toolset after startup, the language backend cannot be changed within a running Serena session + # With the tool interface, the backend cannot change within a session (the toolset depends on it and is fixed); + # with the REPL interface, it may change upon project activation (see _activate_project). self._language_backend = self.serena_config.determine_language_backend( project_config=registered_project_to_activate.project_config if registered_project_to_activate is not None else None, log_choice=True, ) - # create the tool names mapping for prompts - self._prompt_tool_names_mapping = self._create_prompt_tool_names_mapping(self._language_backend) + # determine the effective agent interface for this session (project configuration > global configuration). + # Like the language backend, it is fixed for the session, since the set of exposed tools cannot change after startup. + self._agent_interface = self.serena_config.determine_agent_interface( + project_config=registered_project_to_activate.project_config if registered_project_to_activate is not None else None, + log_choice=True, + ) # create executor for starting the language server and running tools in another thread # This executor is used to achieve linear task execution @@ -676,8 +696,14 @@ class SerenaAgent: self._project_activation_error = str(e) self._update_active_modes() + # determine whether we are operating in a single-project session, i.e. the project that was activated at startup + # (if any) is the only project that will be worked with throughout the session (no project switching) + self._is_single_project = self._context.single_project and self._active_project is not None + # determine the base toolset defining the set of exposed tools (which e.g. the MCP shall see), - self._base_toolset = self._create_base_toolset(self.serena_config, self._context, self._active_modes, self._active_project) + self._base_toolset = self._create_base_toolset( + self.serena_config, self._context, self._active_modes, self._active_project, self._agent_interface, self._is_single_project + ) self._exposed_tools = self._base_toolset.to_available_tools(self._all_tools) log.info(f"Number of exposed tools: {len(self._exposed_tools)}. Exposed tools: {self._exposed_tools.tool_names}") @@ -729,7 +755,7 @@ class SerenaAgent: "os": platform.system(), "dashboard": int(self.serena_config.web_dashboard), "version": self.version, - "backend": self._language_backend.value, + "backend": self._language_backend.get_key(), "context": self._context.name, } try: @@ -737,6 +763,15 @@ class SerenaAgent: except Exception as e: log.debug(f"Failed to send usage info: {e}") + @staticmethod + def _is_dashboard_openable(serena_config: SerenaConfig) -> bool: + """ + :param serena_config: the configuration + :return: whether the web dashboard is available and opening it is a meaningful operation + (i.e. it is enabled and not opened automatically) + """ + return serena_config.web_dashboard and not serena_config.web_dashboard_open_on_launch and not serena_config.gui_log_window + @classmethod def _create_base_toolset( cls, @@ -744,10 +779,12 @@ class SerenaAgent: context: SerenaAgentContext, modes: ActiveModes, project: Project | None, + agent_interface: AgentInterface, + is_single_project: bool, ) -> ToolSet: """ Determines the base toolset defining the set of exposed tools (which e.g. the MCP shall see). - It depends on ... + In REPL mode, the toolset is fixed. Otherwise, it depends on ... * dashboard availability/opening on launch * Serena config * the context (which is fixed for the session) @@ -755,9 +792,16 @@ class SerenaAgent: * the optional tools enabled by initial dynamic modes * single-project mode reductions (if applicable) """ + # when in REPL mode, the toolset is fixed and does not depend on the configuration, context, modes or project + if agent_interface.is_repl(): + tool_classes: list[type[Tool]] = [SerenaReplTool, InitialInstructionsTool] + if not is_single_project: + tool_classes.append(ActivateProjectTool) + return ToolSet({tool_class.get_name_from_cls() for tool_class in tool_classes}) + # determine whether to include the OpenDashboardTool based on the Serena configuration tool_inclusion_definitions: list[ToolInclusionDefinition] = [] - if serena_config.web_dashboard and not serena_config.web_dashboard_open_on_launch and not serena_config.gui_log_window: + if cls._is_dashboard_openable(serena_config): tool_inclusion_definitions.append( NamedToolInclusionDefinition(name="OpenDashboard", included_optional_tools=[OpenDashboardTool.get_name_from_cls()]) ) @@ -766,10 +810,6 @@ class SerenaAgent: tool_inclusion_definitions.append(serena_config) tool_inclusion_definitions.append(context) - # determine whether we are operating in a single-project context - # (i.e. the project that is activated at startup is the only project that will be worked with throughout the session) - is_single_project = context.single_project and project is not None - # consider modes # * base modes: These cannot be changed, so they are fully applied for base_mode in modes.get_base_modes(include_background_base_modes=True): @@ -821,6 +861,17 @@ class SerenaAgent: def get_language_backend(self) -> LanguageBackend: return self._language_backend + def is_single_project(self) -> bool: + """ + :return: whether this is a single-project session, i.e. the project activated at startup is the only project + that will be worked with throughout the session (no project switching); requires a single-project context + and a project at startup + """ + return self._is_single_project + + def get_agent_interface(self) -> AgentInterface: + return self._agent_interface + def get_current_tasks(self) -> list[TaskExecutor.TaskInfo]: """ Gets the list of tasks currently running or queued for execution. @@ -939,27 +990,97 @@ class SerenaAgent: """ return self._active_modes - @staticmethod - def _create_prompt_tool_names_mapping(language_backend: LanguageBackend) -> dict[str, str]: + @dataclass + class PromptParams: """ - Creates a mapping from tool names to new tool names, which take into consideration - - * legacy tool names, where the name was changed and - * LSP tools which are functionally replaced by other tools due to the active language backend - (e.g. "find_symbol" being replaced by "jet_brains_find_symbol" in JetBrains mode). - - The mapping is intended to be used for the generation of prompts, such that prompts can - refer to tool names as `{{ tool_names["find_symbol"] }}`, and the mapping will ensure that - the correct tool name is used in the prompt based on the active language backend. - - :return: the mapping from tool names to new tool names + Holds parameters for prompt rendering """ - result = dict(ToolSet.LEGACY_TOOL_NAME_MAPPING) - class_replacements = language_backend.get_lsp_tool_class_replacements() - for tool_class in ToolRegistry().get_all_tool_classes(): - new_tool_class: type[Tool] = class_replacements.get(tool_class, tool_class) - result[tool_class.get_name_from_cls()] = new_tool_class.get_name_from_cls() - return result + + available_tools: set[str] + """ + available tool names or, in REPL mode, the names of the raw facade methods (without facade name prefix) and + the names of the corresponding tools + """ + available_markers: set[str] + """ + names of the ToolMarkers (class names) that the available tools inherit from + """ + tool_names_mapping: dict[str, str] + """ + mapping from standard tool names to currently used and replacement tool/API method names. + In particular, this maps + * legacy tool names to current tool names + * LSP tool names to their backend- and interface-specific counterparts + (e.g. "find_symbol" to "jet_brains_find_symbol" in JetBrains mode, "find_symbol" to the corresponding API method name + when using the REPL interface). + """ + + def get_function_name(self, tool_class: type[Tool]) -> str: + """ + :param tool_class: the tool for which to get the function name + :return: the function name to use for this tool in prompts, which may be different from the tool's standard name + (e.g. when using a different language backend or when using the REPL interface) + """ + tool_name = tool_class.get_name_from_cls() + return self.tool_names_mapping.get(tool_name, tool_name) + + def _get_prompt_params(self) -> PromptParams: + """ + :return: parameters for prompt rendering depending on the current agent interface, language backend and active tools/methods + """ + if self._prompt_params is not None: + return self._prompt_params + + if self._agent_interface == AgentInterface.TOOLS: + # available tool names are simply the exposed tools + available_tool_names = set(self._exposed_tools.tool_names) + available_tool_marker_names = set(self._exposed_tools.tool_marker_names) + + tool_name_mapping = dict(ToolSet.LEGACY_TOOL_NAME_MAPPING) + class_replacements = self._language_backend.get_lsp_tool_class_replacements() + for tool_class in ToolRegistry().get_all_tool_classes(): + new_tool_class: type[Tool] = class_replacements.get(tool_class, tool_class) + tool_name_mapping[tool_class.get_name_from_cls()] = new_tool_class.get_name_from_cls() + + elif self._agent_interface == AgentInterface.REPL: + # available tool names include both the names of the facade methods and the names of the corresponding tools + repl = self.get_repl() + enabled_methods = repl.entrypoint.get_enabled_methods() + corresponding_tool_classes = [m.info.corresponding_tool for m in enabled_methods if m.info.corresponding_tool is not None] + available_tools = AvailableTools( + self._exposed_tools.tools + [self._all_tools[tool_class] for tool_class in corresponding_tool_classes] + ) + available_tool_names = set(available_tools.tool_names).union({m.info.name for m in enabled_methods}) + available_tool_marker_names = set(available_tools.tool_marker_names) + + tool_class_replacements = self._language_backend.get_lsp_tool_class_replacements() + methods_by_tool_class = {m.info.corresponding_tool: m for m in enabled_methods if m.info.corresponding_tool is not None} + + def get_name(tool_class: type[Tool]) -> str: + # if there is a corresponding method in the API, return its qualified name (as used in REPL code) + method = methods_by_tool_class.get(tool_class) + if method is not None: + return method.qualified_name + # if there is a corresponding method for the replacement class, return its qualified name + replacement_class = tool_class_replacements.get(tool_class) + if replacement_class is not None: + replacement_method = methods_by_tool_class.get(replacement_class) + if replacement_method is not None: + return replacement_method.qualified_name + # otherwise, keep the tool's name + return tool_class.get_name_from_cls() + + tool_name_mapping = {} + for legacy_name, new_name in ToolSet.LEGACY_TOOL_NAME_MAPPING.items(): + tool_name_mapping[legacy_name] = get_name(ToolRegistry().get_tool_class_by_name(new_name)) + for tool_class in ToolRegistry().get_all_tool_classes(): + tool_name_mapping[tool_class.get_name_from_cls()] = get_name(tool_class) + else: + raise ValueError() + + return self.PromptParams( + available_tools=available_tool_names, available_markers=available_tool_marker_names, tool_names_mapping=tool_name_mapping + ) @staticmethod def _format_prompt_tag(text: str, tag: str, tag_name_attr: str | None = None) -> str: @@ -986,10 +1107,11 @@ class SerenaAgent: return "" template = JinjaTemplate(prompt_template) + prompt_params = self._get_prompt_params() text = template.render( - available_tools=self._exposed_tools.tool_names, - available_markers=self._exposed_tools.tool_marker_names, - tool_names=self._prompt_tool_names_mapping, + available_tools=prompt_params.available_tools, + available_markers=prompt_params.available_markers, + tool_names=prompt_params.tool_names_mapping, embed_memory=embed_memory, ) @@ -1021,18 +1143,33 @@ class SerenaAgent: else: return self._create_global_memory_manager() - def create_system_prompt(self, session_id: str = "global") -> str: + def create_session(self) -> SerenaSession: + """ + :return: a new client session (with a random id) + """ + return self._session_registry.create_session() + + def get_session(self, session_id: str) -> SerenaSession: + """ + :param session_id: the session id (as supplied by the LLM) + :return: the session, which is created if it is unknown + """ + return self._session_registry.get_session(session_id) + + def create_system_prompt(self) -> str: """ Returns the 'Serena Instructions Manual', i.e. Serena's system prompt. + The prompt also establishes a new Serena session (see `SerenaSession`), stating its id for use with tools + which require it (e.g. the REPL tool and project activation tool). - :param session_id: the client session ID for the case where this is run from a tool; "global" for the connection time case :return: the prompt """ - available_tools = self._active_tools - available_markers = available_tools.tool_marker_names + # establish a Serena session + serena_session = self.create_session() + session_id = serena_session.session_id + global_memories = self._create_global_memory_manager().list_global_memories() global_memories_str = dict_string(global_memories.to_dict()) if len(global_memories) > 0 else "" - log.info("Generating system prompt with available_tools=(see active tools), available_markers=%s", available_markers) # determine modes for which prompts must (still) be provided, excluding modes that were already provided in a # previously provided project activation message (if any) @@ -1043,13 +1180,14 @@ class SerenaAgent: relevant_modes.append(mode) self._project_prompt_status.mark_mode_prompts_as_provided(session_id) + prompt_params = self._get_prompt_params() system_prompt = self.prompt_factory.create_system_prompt( context_system_prompt=self._render_prompt(self._context.prompt, tag="context"), mode_system_prompts=[self._render_prompt(mode.prompt, tag="mode", tag_name_attr=mode.name) for mode in relevant_modes], - available_tools=available_tools.tool_names, - available_markers=available_markers, + available_tools=prompt_params.available_tools, + available_markers=prompt_params.available_markers, global_memories_list=global_memories_str, - tool_names=self._prompt_tool_names_mapping, + tool_names=prompt_params.tool_names_mapping, ) # provide the project activation message if it hasn't yet been provided @@ -1058,6 +1196,12 @@ class SerenaAgent: elif self._project_activation_error: system_prompt += f"\n\nNo project is active ({self._project_activation_error})." + # inform about the session id + system_prompt += "\n\n" + self._format_prompt_tag( + f"Your Serena session id is `{session_id}`. Pass it as the `session_id` parameter to tools which require it.", + tag="session", + ) + return self._format_prompt_tag(system_prompt, tag="serena") def get_project_activation_message(self, session_id: str) -> str: @@ -1068,6 +1212,8 @@ class SerenaAgent: proj = self._active_project assert proj is not None, "A project must be active before calling this." + prompt_params = self._get_prompt_params() + # Note: The activation message is always returned in full, even if it was already provided in the current session, # because some clients (e.g. Claude Desktop) will use the same session across multiple chats. # So while we don't want the activation message to be additionally included in the system prompt @@ -1081,22 +1227,20 @@ class SerenaAgent: msg = f"Created and activated a new project with name '{proj.project_name}' at {proj.project_root}.\n" else: msg = f"The project with name '{proj.project_name}' at {proj.project_root} is activated.\n" - if self._language_backend == LanguageBackend.LSP: - language_servers_str = ", ".join([ls.get_key() for ls in proj.project_config.language_servers]) - msg += f"Active language servers: {language_servers_str}.\n" + msg += self._language_backend.get_project_activation_statement(proj) msg += f"File encoding: {proj.project_config.encoding}.\n" # add list of memories (if memories are enabled) - include_memories = self._active_tools.contains_tool_class(ReadMemoryTool) + include_memories = self.is_tool_function_available(ReadMemoryTool) if include_memories: project_memories = proj.memory_manager.list_project_memories() if project_memories: msg += ( f"{json.dumps(project_memories.to_dict())}\n" - + f"Use the `{ReadMemoryTool.get_name_from_cls()}` tool to read these memories later if they are relevant to the task.\n" + + f"Use `{prompt_params.get_function_name(ReadMemoryTool)}` to read these memories later if they are relevant to the task.\n" ) - elif self._active_tools.contains_tool_class(OnboardingTool): - msg += f"Onboarding has not been performed yet. Ask the user whether to perform onboarding via the `{OnboardingTool.get_name_from_cls()}` tool.\n" + elif self.is_tool_function_available(OnboardingTool): + msg += f"Onboarding has not been performed yet. Ask the user whether to perform onboarding and if so, call `{prompt_params.get_function_name(OnboardingTool)}`.\n" # add prompts for modes that were dynamically activated by the project modes_with_prompts = self._project_prompt_status.get_modes_with_prompts_to_be_provided_for_project_activation(session_id) @@ -1105,10 +1249,15 @@ class SerenaAgent: msg += self._render_prompt(mode.prompt, tag="mode", tag_name_attr=mode.name) + "\n" self._project_prompt_status.mark_mode_prompts_as_provided(session_id) - # add project-specific prompt + # add the project's prompt (if any) if proj.project_config.initial_prompt: msg += "\n" + self._render_prompt(proj.project_config.initial_prompt, tag="project-instructions") + # when the REPL is active, add information on available facades if the agent is not in single-project mode + # (for single-project mode where the facades can't change, they are provided in the tool's description) + if self._active_tools.contains_tool_class(SerenaReplTool) and not self.is_single_project(): + msg += f"\n\nAvailable facades for the `{SerenaReplTool.get_name_from_cls()}` tool:\n" + self.get_repl().entrypoint.overview() + self._project_prompt_status.mark_project_activation_message_as_provided(session_id) return msg @@ -1133,22 +1282,31 @@ class SerenaAgent: def _update_active_tools(self) -> None: """ - Updates the active tools based on the active modes and the active project. + Updates the active tools (and the REPL, which depends on the same configuration) based on the active modes + and the active project. Must be called whenever the active modes or the active project change. The base tool set already takes the Serena configuration and the context into account (as well as many other aspects, such as JetBrains mode). """ - # apply modes - tool_set = self._base_toolset.apply(*self._active_modes.get_modes()) + if self._agent_interface.is_repl(): + # the REPL toolset is fixed; tool inclusion/exclusion definitions do not apply + tool_set = self._base_toolset + else: + # apply modes + tool_set = self._base_toolset.apply(*self._active_modes.get_modes()) - # apply active project configuration (if any) - if self._active_project is not None: - tool_set = tool_set.apply(self._active_project.project_config) - if self._active_project.project_config.read_only: - tool_set = tool_set.without_editing_tools() + # apply active project configuration (if any) + if self._active_project is not None: + tool_set = tool_set.apply(self._active_project.project_config) + if self._active_project.project_config.read_only: + tool_set = tool_set.without_editing_tools() self._active_tools = tool_set.to_available_tools(self._all_tools) log.info(f"Active tools ({len(self._active_tools)}): {', '.join(self._active_tools.tool_names)}") + # reset members that depend on the active tools, so that they are re-created on demand with the new active tools + self._repl = None + self._prompt_params = None + # check if a tool was activated that is not in the exposed tool set and issue a warning if so active_tools_not_exposed = set(self._active_tools.tool_names) - set(self._exposed_tools.tool_names) if active_tools_not_exposed: @@ -1158,6 +1316,41 @@ class SerenaAgent: "Consider adjusting your configuration to include these tools if you want to use them." ) + def create_default_facade_list(self, api_scope: ApiScope) -> list[Facade]: + """ + :return: the default list of facades provided by Serena itself, not including any language backend-specific facades + """ + return [ + Facade.from_api(ConfigApi(self), api_scope), + Facade.from_api(FsApi(self), api_scope), + Facade.from_api(EditApi(self), api_scope), + Facade.from_api(MemoryApi(self), api_scope), + Facade.from_api(ShellApi(self), api_scope), + Facade.from_api(ExternalProjectsApi(self), api_scope, is_optional=True), + ] + + def get_repl(self) -> SerenaRepl: + """ + :return: the REPL instance for this agent, creating it if necessary + """ + if self._repl is None: + # determine API scope + api_scope = ApiScope() + api_scope.process(self.serena_config) + api_scope.process(self._context) + for mode in self._active_modes.get_modes(): + api_scope.process(mode) + if self._active_project: + api_scope.process(self._active_project.project_config) + if self._active_project.project_config.read_only: + api_scope.exclude_editing() + if not self._is_dashboard_openable(self.serena_config): + api_scope.process(NamedApiInclusionDefinition(name="Dashboard", excluded_apis=["cfg.open_dashboard"])) + + facades = self.create_default_facade_list(api_scope) + self._language_backend.create_facades(self, api_scope) + self._repl = SerenaRepl(facades, api_scope) + return self._repl + def issue_task( self, task: Callable[[], T], name: str | None = None, logged: bool = True, timeout: float | None = None ) -> TaskExecutor.Task[T]: @@ -1211,7 +1404,7 @@ class SerenaAgent: """ :return: whether this agent uses language server-based code analysis """ - return self._language_backend == LanguageBackend.LSP + return self._language_backend == BuiltinLanguageBackend.LSP def _activate_project(self, project: Project, update_active_modes: bool = True, update_active_tools: bool = True) -> bool: """ @@ -1225,15 +1418,22 @@ class SerenaAgent: self._project_activation_error = None - # check if the project requires a different language backend than the one initialized at startup + # handle the case where the project requires a different language backend than the current one. + # With the tool interface, the backend cannot change, since the set of exposed tools depends on it and is fixed + # for the session. With the REPL interface, the backend can be switched, as all backend-dependent state + # (background modes, REPL facades, prompt parameters, the project's language backend initialisation) is + # recomputed upon activation. project_backend = project.project_config.language_backend if project_backend is not None and project_backend != self._language_backend: - raise ValueError( - f"Cannot activate project '{project.project_name}': it requires the {project_backend.value} backend, " - f"but this session was initialized with {self._language_backend.value}. " - f"Workarounds: (1) Use project activation at startup via the --project flag, " - f"(2) Configure one MCP server per backend in your client." - ) + if self._agent_interface.is_tools(): + raise ValueError( + f"Cannot activate project '{project.project_name}': it requires the {project_backend} backend, " + f"but this session was initialized with {self._language_backend}. " + f"Workarounds: (1) Use project activation at startup via the --project flag, " + f"(2) Configure one MCP server per backend in your client, (3) use the REPL interface." + ) + log.info(f"Switching language backend from {self._language_backend} to {project_backend} for project '{project.project_name}'") + self._language_backend = project_backend # shut down the previously active project to release its language server processes if self._active_project is not None: @@ -1257,7 +1457,7 @@ class SerenaAgent: def init_project_services() -> None: self._run_project_activation_command(project) - self._init_active_project_language_backend() + self._language_backend.init_active_project(self) # initialise the project's language backend in the background self.issue_task(init_project_services) @@ -1313,28 +1513,6 @@ class SerenaAgent: except Exception: log.exception(f"Unexpected error running activation_command for project '{project.project_name}'") - def _init_active_project_language_backend(self) -> None: - """ - Initialises the active project's language backend - """ - project = self._active_project - assert project is not None - - # for LSP mode, start the language server manager - if self.get_language_backend().is_lsp(): - with LogTime("Language server initialization", logger=log): - self.reset_language_server_manager() - - # for JetBrains mode, search for plugin server and spawn IDE (if not found and launch command provided) - elif self.get_language_backend().is_jetbrains(): - client = jetbrains_launch_coordinator.find_plugin_server(project) - if client is not None: - log.info("Found Serena JetBrains Plugin server: %s", client) - else: - log.info("Serena JetBrains Plugin server not found for project %s", project.project_name) - if self.serena_config.jetbrains_launch_command: - jetbrains_launch_coordinator.launch_and_wait_for_plugin_server(project, self.serena_config.jetbrains_launch_command) - def activate_project_from_path_or_name( self, project_root_or_name: str, update_active_modes: bool = True, update_active_tools: bool = True ) -> bool: @@ -1367,19 +1545,27 @@ class SerenaAgent: """ return self._active_tools.tool_names - def tool_is_active(self, tool_name: str) -> bool: + def get_active_tools(self) -> AvailableTools: """ - :param tool_class: the name of the tool to check - :return: True if the tool is active, False otherwise + :return: the set of active tools """ - return self._active_tools.contains_tool_name(tool_name) + return self._active_tools - def tool_is_exposed(self, tool_name: str) -> bool: + def is_tool_function_available(self, tool_class: type[Tool]) -> bool: """ - :param tool_name: the name of the tool to check - :return: True if the tool is in the exposed tool set, False otherwise + Checks whether the functionality offered by a tool is available - either through the tool + itself being enabled or through the corresponding function being exposed in the REPL. + + :param tool_class: the tool class + :return: whether the function is available """ - return self._exposed_tools.contains_tool_name(tool_name) + is_active_tool = self._active_tools.contains_tool_class(tool_class) + if self._agent_interface == AgentInterface.TOOLS: + return is_active_tool + elif self._agent_interface == AgentInterface.REPL: + return is_active_tool or self.get_repl().entrypoint.is_tool_function_available(tool_class) + else: + raise NotImplementedError def get_current_config_overview(self) -> str: """ @@ -1392,12 +1578,13 @@ class SerenaAgent: result_str += f"Active project: {self._active_project.project_name}\n" else: result_str += "No active project\n" - result_str += f"Language backend: {self._language_backend.value}" + result_str += f"Agent interface: {self._agent_interface.value}\n" + result_str += f"Language backend: {self._language_backend.get_key()}" if self._active_project and self._active_project.project_config.language_backend is not None: result_str += " (project override)" - result_str += f" (global default: {self.serena_config.language_backend.value})\n" - if self._language_backend.is_lsp() and self._active_project: - result_str += f"Language server status: {self._active_project.get_language_server_manager_status()}\n" + result_str += f" (global default: {self.serena_config.language_backend.get_key()})\n" + if self._active_project: + result_str += self._language_backend.get_config_overview_statement(self._active_project) result_str += "Available projects:\n" + "\n".join(list(self.serena_config.project_names)) + "\n" result_str += f"Active context: {self._context.name}\n" @@ -1457,7 +1644,7 @@ class SerenaAgent: self.issue_task(lambda: self.get_active_project_or_raise().remove_language_server(ls_id), name=f"RemoveLanguage:{ls_id.get_key()}") def get_tool(self, tool_class: type[TTool]) -> TTool: - return self._all_tools[tool_class] + return cast(TTool, self._all_tools[tool_class]) def print_tool_overview(self) -> None: ToolRegistry().print_tool_overview(self._active_tools.tools) diff --git a/src/serena/analytics.py b/src/serena/analytics.py index 315cd57d..8fcc5636 100644 --- a/src/serena/analytics.py +++ b/src/serena/analytics.py @@ -9,10 +9,15 @@ from collections import defaultdict from copy import copy from dataclasses import asdict, dataclass from enum import Enum +from typing import TYPE_CHECKING -from anthropic.types import MessageParam, MessageTokensCount from dotenv import load_dotenv +if TYPE_CHECKING: + # Imported for annotations only: loading the anthropic package costs seconds on some + # machines (see #2012) and is only needed when the Anthropic token counter is used. + from anthropic.types import MessageTokensCount + log = logging.getLogger(__name__) @@ -64,7 +69,7 @@ class AnthropicTokenCount(TokenCountEstimator): def _send_count_tokens_request(self, text: str) -> MessageTokensCount: return self._anthropic_client.messages.count_tokens( model=self._model_name, - messages=[MessageParam(role="user", content=text)], + messages=[{"role": "user", "content": text}], ) def estimate_token_count(self, text: str) -> int: diff --git a/src/serena/cli.py b/src/serena/cli.py index ab41b61a..b44a2ac3 100644 --- a/src/serena/cli.py +++ b/src/serena/cli.py @@ -23,7 +23,7 @@ from serena import serena_version from serena.config.client_setup import client_setup_handlers from serena.config.context_mode import SerenaAgentContext, SerenaAgentMode from serena.config.serena_config import ( - LanguageBackend, + AgentInterface, ModeSelectionDefinition, ModeSelectionDefinitionWithAddedModes, ProjectConfig, @@ -38,11 +38,12 @@ from serena.constants import ( SERENAS_OWN_CONTEXT_YAMLS_DIR, SERENAS_OWN_MODE_YAMLS_DIR, ) +from serena.language_backend import BuiltinLanguageBackend, LanguageBackendRegistry from serena.prompt_factory import SerenaPromptFactory from serena.tools import ActivateProjectTool from serena.util.cli_util import AutoRegisteringGroup from serena.util.logging import MemoryLogHandler -from solidlsp.ls_config import LanguageServerId, LanguageServerIdLike +from solidlsp.ls_config import LanguageServerIdLike, LanguageServerRegistry from solidlsp.ls_types import SymbolKind from solidlsp.util.subprocess_util import subprocess_kwargs @@ -180,14 +181,14 @@ class TopLevelCommands(AutoRegisteringGroup): @click.option( "--language-backend", "-b", - type=click.Choice([b.value for b in LanguageBackend]), - default=LanguageBackend.LSP.value, + type=click.Choice([b.value for b in BuiltinLanguageBackend]), + default=BuiltinLanguageBackend.LSP.value, show_default=True, help="Default code intelligence backend (can be overridden in the project config).", ) def init(language_backend: Literal["LSP", "JetBrains"] = "LSP") -> None: click.echo(f"\nSerena version: {serena_version()}\n") - serena_config = SerenaConfig.init(language_backend=LanguageBackend(language_backend)) + serena_config = SerenaConfig.init(builtin_language_backend=BuiltinLanguageBackend(language_backend)) click.echo(f"Configuration file: {serena_config.config_file_path}") click.echo(f"Language backend: {language_backend}") @@ -259,10 +260,17 @@ class TopLevelCommands(AutoRegisteringGroup): ) @click.option( "--language-backend", - type=click.Choice([lb.value for lb in LanguageBackend]), + type=click.Choice(LanguageBackendRegistry.get_instance().get_keys()), default=None, help="Override the configured language backend.", ) + @click.option( + "--agent-interface", + type=click.Choice([i.value for i in AgentInterface], case_sensitive=False), + default=None, + help="Override the configured agent interface: 'tools' (one tool per operation) or " + "'REPL' (Python code execution via the serena_repl tool, with a fixed set of tools).", + ) @click.option( "--transport", type=click.Choice(["stdio", "sse", "streamable-http"]), @@ -325,6 +333,7 @@ class TopLevelCommands(AutoRegisteringGroup): default_modes: Sequence[str], added_modes: Sequence[str], language_backend: str | None, + agent_interface: str | None, transport: Literal["stdio", "sse", "streamable-http"], host: str, port: int, @@ -381,10 +390,9 @@ class TopLevelCommands(AutoRegisteringGroup): factory = SerenaMCPFactory(transport=transport, context=context, project=project_file, memory_log_handler=memory_log_handler) server = factory.create_mcp_server( - host=host, - port=port, mode_selection_def=mode_selection_def, - language_backend=LanguageBackend.from_str(language_backend) if language_backend else None, + language_backend=LanguageBackendRegistry.get_instance().resolve(language_backend) if language_backend else None, + agent_interface=AgentInterface.from_str(agent_interface) if agent_interface else None, enable_web_dashboard=enable_web_dashboard, open_web_dashboard=open_web_dashboard, enable_gui_log_window=enable_gui_log_window, @@ -399,7 +407,11 @@ class TopLevelCommands(AutoRegisteringGroup): project_file, ) log.info("Starting MCP server …") - server.run(transport=transport) + kwargs = {} + if transport != "stdio": + kwargs["host"] = host + kwargs["port"] = port + server.run(transport=transport, **kwargs) @staticmethod @click.command( @@ -708,14 +720,15 @@ class ProjectCommands(AutoRegisteringGroup): if os.path.exists(yml_path): raise FileExistsError(f"Project file {yml_path} already exists.") - languages: list[LanguageServerId] = [] + languages: list[LanguageServerIdLike] = [] if language: + registry = LanguageServerRegistry.get_instance() for lang in language: + ls_key = lang.lower() try: - languages.append(LanguageServerId(lang.lower())) + languages.append(registry.resolve(ls_key)) except ValueError: - all_langs = [l.value for l in LanguageServerId] - raise ValueError(f"Unknown language '{lang}'. Supported: {all_langs}") + raise ValueError(f"Unknown language '{lang}'. Supported: {registry.get_keys()}") generated_conf = ProjectConfig.autogenerate( project_root=project_path, @@ -766,6 +779,27 @@ class ProjectCommands(AutoRegisteringGroup): except ValueError as e: raise click.ClickException(str(e)) + @staticmethod + @click.command( + "remove", + help="Remove a project from Serena's project registry. " + "The project's own files, including its project configuration, are left untouched.", + context_settings={"max_content_width": _MAX_CONTENT_WIDTH}, + ) + @click.argument("project", type=PROJECT_TYPE) + def remove(project: str) -> None: + serena_config = SerenaConfig.from_config_file() + registered_project_names = serena_config.project_names + try: + registered_project = serena_config.get_registered_project(project) + except ValueError as e: + # raised when the name is ambiguous; the message names the candidate locations + raise click.ClickException(str(e)) + if registered_project is None: + raise click.ClickException(f"No registered project found for '{project}'; registered project names: {registered_project_names}") + serena_config.remove_registered_project(registered_project) + click.echo(f"Removed project '{registered_project.project_name}' ({registered_project.project_root}) from the project registry.") + @staticmethod @click.command( "index", @@ -932,12 +966,12 @@ class ProjectCommands(AutoRegisteringGroup): # NOTE: completely written by Claude Code, only functionality was reviewed, not implementation from serena.agent import SerenaAgent from serena.project import Project - from serena.tools import FindReferencingSymbolsTool, FindSymbolTool, GetSymbolsOverviewTool + from serena.repl.api.lsp_api import LspApi logging.configure(level=logging.INFO) project_path = os.path.abspath(project) serena_config = SerenaConfig.from_config_file().with_headless_mode_overrides() - serena_config.language_backend = LanguageBackend.LSP + serena_config.set_builtin_language_backend(BuiltinLanguageBackend.LSP) proj = Project.load(project_path, serena_config=serena_config) # Create log file with timestamp @@ -977,61 +1011,55 @@ class ProjectCommands(AutoRegisteringGroup): if not target_file: raise ProjectCommands._HealthCheckFailure("No analyzable files found") - # Get tools from agent - overview_tool = agent.get_tool(GetSymbolsOverviewTool) - find_symbol_tool = agent.get_tool(FindSymbolTool) - find_refs_tool = agent.get_tool(FindReferencingSymbolsTool) + api = LspApi(agent) - # Test 1: Get symbols overview - log.info("Testing GetSymbolsOverviewTool on file: %s", target_file) - overview_data = agent.execute_task(lambda: overview_tool.get_symbol_overview(target_file)) - log.info(f"GetSymbolsOverviewTool returned: {overview_data}") + # Test 1: symbols overview + log.info("Testing get_symbols_overview on file: %s", target_file) + overview = agent.execute_task(lambda: api.get_symbols_overview(target_file)) + log.info(f"get_symbols_overview returned: {overview.represent()}") - if not overview_data: + if len(overview) == 0: raise ProjectCommands._HealthCheckFailure(f"No symbols found in target file {target_file}") # Extract suitable symbol (prefer class or function over variables) - preferred_kinds = {SymbolKind.Class.name, SymbolKind.Function.name, SymbolKind.Method.name, SymbolKind.Constructor.name} - selected_symbol = None - for symbol in overview_data: - if symbol.get("kind") in preferred_kinds: - selected_symbol = symbol - break + preferred_kinds = {SymbolKind.Class, SymbolKind.Function, SymbolKind.Method, SymbolKind.Constructor} + selected_symbol = next((s for s in overview.symbols if s.symbol_kind in preferred_kinds), None) # If no preferred symbol found, use first available - if not selected_symbol: - selected_symbol = overview_data[0] + if selected_symbol is None: + selected_symbol = overview.symbols[0] log.info("No class or function found, using first available symbol") - symbol_name = selected_symbol["name"] - symbol_kind = selected_symbol["kind"] - log.info("Using symbol for testing: %s (kind: %s)", symbol_name, symbol_kind) + symbol_name = selected_symbol.name + log.info("Using symbol for testing: %s (kind: %s)", symbol_name, selected_symbol.symbol_kind_name) - # Test 2: FindSymbolTool - log.info("Testing FindSymbolTool for symbol: %s", symbol_name) - with find_symbol_tool.symbol_dict_grouper.disabled_context(): + # Test 2: find_symbol + log.info("Testing find_symbol for symbol: %s", symbol_name) + with LspApi.find_symbol_dict_grouper_.disabled_context(): find_symbol_result = agent.execute_task( - lambda: find_symbol_tool.apply(symbol_name, relative_path=target_file, include_body=True) + lambda: api.find_symbol(symbol_name, relative_path=target_file, include_body=True).represent() ) find_symbol_data = json.loads(find_symbol_result) - log.info("FindSymbolTool found %d matches for symbol %s", len(find_symbol_data), symbol_name) + log.info("find_symbol found %d matches for symbol %s", len(find_symbol_data), symbol_name) if not find_symbol_data: raise ProjectCommands._HealthCheckFailure("FindSymbolTool returned no results") - # Test 3: FindReferencingSymbolsTool - log.info("Testing FindReferencingSymbolsTool for symbol: %s", symbol_name) + # Test 3: find_referencing_symbols + log.info("Testing find_referencing_symbols for symbol: %s", symbol_name) try: - with find_refs_tool.symbol_dict_grouper.disabled_context(): - find_refs_result = agent.execute_task(lambda: find_refs_tool.apply(symbol_name, relative_path=target_file)) + with LspApi.references_grouper_.disabled_context(): + find_refs_result = agent.execute_task( + lambda: api.find_referencing_symbols(symbol_name, relative_path=target_file).represent() + ) find_refs_data = json.loads(find_refs_result) - log.info("FindReferencingSymbolsTool found %d references for symbol %s", len(find_refs_data), symbol_name) + log.info("find_referencing_symbols found %d references for symbol %s", len(find_refs_data), symbol_name) except Exception as e: # A symbol with no references at all is a legitimate result, so the number of # references is not asserted - but a *failure* of the reference search means the # language server is not functional, which is the single thing this command is # asked to determine. Logging it as a warning let the command print # "All tools working correctly" and exit 0 after the search had already failed. - raise ProjectCommands._HealthCheckFailure(f"FindReferencingSymbolsTool failed for symbol {symbol_name}: {e}") from e + raise ProjectCommands._HealthCheckFailure(f"find_referencing_symbols failed for symbol {symbol_name}: {e}") from e log.info("Health check completed successfully") diff --git a/src/serena/code_editor.py b/src/serena/code_editor.py index b9091e72..85aa8b96 100644 --- a/src/serena/code_editor.py +++ b/src/serena/code_editor.py @@ -6,7 +6,8 @@ import os from abc import ABC, abstractmethod from collections.abc import Iterable, Iterator, Reversible from contextlib import contextmanager -from typing import Generic, TypeVar, cast +from types import TracebackType +from typing import Any, Generic, Self, TypeVar, cast from serena.jetbrains.jetbrains_plugin_client import JetBrainsPluginClient from serena.symbol import JetBrainsSymbol, LanguageServerSymbol, LanguageServerSymbolRetriever, PositionInFile, Symbol @@ -16,6 +17,7 @@ from solidlsp.ls_utils import PathUtils, TextStepper, TextUtils from .project import Project from .util.file_proxy import FileProxy +from .util.file_system import write_file_atomic log = logging.getLogger(__name__) TSymbol = TypeVar("TSymbol", bound=Symbol) @@ -23,6 +25,7 @@ TSymbol = TypeVar("TSymbol", bound=Symbol) class CodeEditor(Generic[TSymbol], ABC): def __init__(self, project: Project) -> None: + self.project = project self.project_root = project.project_root self.encoding = project.project_config.encoding self.newline = project.line_ending.newline_str @@ -81,7 +84,7 @@ class CodeEditor(Generic[TSymbol], ABC): """ Context manager for editing a file. """ - if FileProxy.is_external_path(relative_path): + if FileProxy.is_external_path(relative_path, self.project): raise ValueError(f"Cannot edit external file: {relative_path}") with self._open_file_context(relative_path) as edited_file: yield edited_file @@ -91,8 +94,7 @@ class CodeEditor(Generic[TSymbol], ABC): def _save_edited_file(self, edited_file: "CodeEditor.EditedFile") -> None: abs_path = os.path.join(self.project_root, edited_file.relative_path) new_contents = edited_file.get_contents() - with open(abs_path, "w", encoding=self.encoding, newline=self.newline) as f: - f.write(new_contents) + write_file_atomic(abs_path, new_contents, encoding=self.encoding, newline=self.newline) @abstractmethod def _find_unique_symbol(self, name_path: str, relative_file_path: str) -> TSymbol: @@ -493,3 +495,45 @@ class JetBrainsCodeEditor(CodeEditor[JetBrainsSymbol]): rename_in_text_occurrences=rename_in_text_occurrences, ) return "Success" + + +class EditedFileContext: + """ + Context manager for file editing. + + Create the context, then use `set_updated_content` to set the new content, the original content + being provided in `original_content`. + When exiting the context without an exception, the updated content will be written back to the file. + """ + + def __init__(self, relative_path: str, code_editor: CodeEditor): + self._relative_path = relative_path + self._code_editor = code_editor + self._edited_file: CodeEditor.EditedFile | None = None + self._edited_file_context: Any = None + + def __enter__(self) -> Self: + self._edited_file_context = self._code_editor.edited_file_context(self._relative_path) + self._edited_file = self._edited_file_context.__enter__() + return self + + def get_original_content(self) -> str: + """ + :return: the original content of the file before any modifications. + """ + assert self._edited_file is not None + return self._edited_file.get_contents() + + def set_updated_content(self, content: str) -> None: + """ + Sets the updated content of the file, which will be written back to the file + when the context is exited without an exception. + + :param content: the updated content of the file + """ + assert self._edited_file is not None + self._edited_file.set_contents(content) + + def __exit__(self, exc_type: type[BaseException] | None, exc_value: BaseException | None, traceback: TracebackType | None) -> None: + assert self._edited_file_context is not None + self._edited_file_context.__exit__(exc_type, exc_value, traceback) diff --git a/src/serena/config/context_mode.py b/src/serena/config/context_mode.py index 7e08afb8..eae373ed 100644 --- a/src/serena/config/context_mode.py +++ b/src/serena/config/context_mode.py @@ -12,7 +12,7 @@ import yaml from sensai.util import logging from sensai.util.string import ToStringMixin -from serena.config.serena_config import SerenaPaths, ToolInclusionDefinition +from serena.config.serena_config import ApiInclusionDefinition, SerenaPaths, ToolInclusionDefinition from serena.constants import ( DEFAULT_CONTEXT, INTERNAL_MODE_YAMLS_DIR, @@ -32,7 +32,7 @@ def looks_like_yaml_path(s: str) -> bool: @dataclass(kw_only=True) -class SerenaAgentMode(ToolInclusionDefinition, ToStringMixin): +class SerenaAgentMode(ToolInclusionDefinition, ApiInclusionDefinition, 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. @@ -148,7 +148,7 @@ class SerenaAgentMode(ToolInclusionDefinition, ToStringMixin): @dataclass(kw_only=True) -class SerenaAgentContext(ToolInclusionDefinition, ToStringMixin): +class SerenaAgentContext(ToolInclusionDefinition, ApiInclusionDefinition, 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. diff --git a/src/serena/config/serena_config.py b/src/serena/config/serena_config.py index bc975d73..b7a78417 100644 --- a/src/serena/config/serena_config.py +++ b/src/serena/config/serena_config.py @@ -7,15 +7,16 @@ import dataclasses import os import re import shutil +import stat import threading from collections.abc import Iterator, Sequence from copy import deepcopy from dataclasses import dataclass, field from datetime import UTC, datetime from enum import Enum -from functools import cached_property from pathlib import Path from typing import TYPE_CHECKING, Any, Optional, Self, TypeVar +from uuid import uuid4 import yaml from ruamel.yaml.comments import CommentedMap @@ -36,16 +37,16 @@ from serena.constants import ( from serena.util.inspection import compute_language_server_support_composition from serena.util.text_utils import GlobMatcher from serena.util.yaml import YamlCommentNormalisation, load_yaml, normalise_yaml_comments, save_yaml, transfer_yaml_comments -from solidlsp.ls_config import LanguageServerId, LanguageServerIdLike, LanguageServerRegistry +from solidlsp.ls_config import LanguageServerIdLike, LanguageServerRegistry from ..analytics import RegisteredTokenCountEstimator +from ..language_backend import BuiltinLanguageBackend, LanguageBackend, LanguageBackendRegistry from ..util.class_decorators import singleton from ..util.cli_util import ask_yes_no from ..util.dataclass import get_dataclass_default if TYPE_CHECKING: from ..project import Project - from ..tools.tools_base import Tool log = logging.getLogger(__name__) T = TypeVar("T") @@ -178,6 +179,26 @@ class NamedToolInclusionDefinition(ToolInclusionDefinition): return f"ToolInclusionDefinition[{self.name}]" +@dataclass +class ApiInclusionDefinition: + """ + Defines which APIs to include/exclude in Serena's operation. + A single API inclusion/exclusion can either be a full facade (facade name, which encompasses all of its methods, e.g. "lsp") + or a method of a facade (facade name + method name, e.g. "lsp.find_symbol"). + """ + + included_apis: Sequence[str] = () + excluded_apis: Sequence[str] = () + + +@dataclass +class NamedApiInclusionDefinition(ApiInclusionDefinition): + name: str | None = None + + def __str__(self) -> str: + return f"ApiInclusionDefinition[{self.name}]" + + @dataclass class ModeSelectionDefinition: default_modes: Sequence[str] | None = None @@ -196,52 +217,35 @@ class ModeSelectionDefinitionWithAddedModes(ModeSelectionDefinition): added_modes: Sequence[str] | None = None -class LanguageBackend(Enum): - LSP = "LSP" +class AgentInterface(Enum): """ - Use the language server protocol (LSP), spawning freely available language servers - via the SolidLSP library that is part of Serena + The interface through which the agent (LLM) accesses Serena's functionality. """ - JETBRAINS = "JetBrains" + + TOOLS = "tools" """ - Use the Serena plugin in your JetBrains IDE. - (requires the plugin to be installed and the project being worked on to be open in your IDE) + The classic tool interface: each operation is a separate tool, and the set of tools is configurable + (via tool inclusions/exclusions in the configuration, context, modes and project). + """ + REPL = "REPL" + """ + The REPL interface: operations are accessed programmatically via the serena_repl tool, which executes Python code. + The set of tools is fixed (the REPL tool and the tools required for session management) and tool inclusions/exclusions + do not apply; the operations available in the REPL are configured via API inclusions/exclusions instead. """ @staticmethod - def from_str(backend_str: str) -> "LanguageBackend": - for backend in LanguageBackend: - if backend.value.lower() == backend_str.lower(): - return backend - raise ValueError(f"Unknown language backend '{backend_str}': valid values are {[b.value for b in LanguageBackend]}") + def from_str(interface_str: str) -> "AgentInterface": + for interface in AgentInterface: + if interface.value.lower() == interface_str.lower(): + return interface + raise ValueError(f"Unknown agent interface '{interface_str}': valid values are {[i.value for i in AgentInterface]}") - def is_lsp(self) -> bool: - return self == LanguageBackend.LSP + def is_tools(self) -> bool: + return self == AgentInterface.TOOLS - def is_jetbrains(self) -> bool: - return self == LanguageBackend.JETBRAINS - - def get_lsp_tool_class_replacements(self) -> "dict[type[Tool], type[Tool]]": - """ - :return: mapping from LSP tool classes to replacement tool classes (functional replacements) - """ - match self: - case LanguageBackend.LSP: - return {} - case LanguageBackend.JETBRAINS: - from ..tools import jetbrains_tools, symbol_tools - - return { - symbol_tools.FindSymbolTool: jetbrains_tools.JetBrainsFindSymbolTool, - symbol_tools.GetSymbolsOverviewTool: jetbrains_tools.JetBrainsGetSymbolsOverviewTool, - symbol_tools.FindReferencingSymbolsTool: jetbrains_tools.JetBrainsFindReferencingSymbolsTool, - symbol_tools.FindImplementationsTool: jetbrains_tools.JetBrainsFindImplementationsTool, - symbol_tools.FindDeclarationTool: jetbrains_tools.JetBrainsFindDeclarationTool, - symbol_tools.RenameSymbolTool: jetbrains_tools.JetBrainsRenameTool, - symbol_tools.SafeDeleteSymbol: jetbrains_tools.JetBrainsSafeDeleteTool, - } - case _: - raise NotImplementedError() + def is_repl(self) -> bool: + return self == AgentInterface.REPL class LineEnding(Enum): @@ -274,7 +278,7 @@ class LineEnding(Enum): @dataclass -class SharedConfig(ToolInclusionDefinition, ToStringMixin): +class SharedConfig(ToolInclusionDefinition, ApiInclusionDefinition, ToStringMixin): """Shared between SerenaConfig and ProjectConfig, the latter used to override values in the form (same as in ModeSelectionDefinition). The defaults here shall be none and should be set to the global default values in SerenaConfig. @@ -282,6 +286,7 @@ class SharedConfig(ToolInclusionDefinition, ToStringMixin): symbol_info_budget: float | None = None language_backend: LanguageBackend | None = None + agent_interface: AgentInterface | None = None line_ending: LineEnding | None = None read_only_memory_patterns: list[str] = field(default_factory=list) ignored_memory_patterns: list[str] = field(default_factory=list) @@ -363,11 +368,15 @@ class ProjectConfig(SharedConfig, ModeSelectionDefinitionWithAddedModes): log.info("Determining suitable language servers for the project") # determine language servers to be considered and their priorities - ls_priorities = {} - for language in LanguageServerId: - priority = serena_config.get_ls_priority(language) + # the registry is the single source of truth — it includes both built-in enum members + # and externally-registered adapters (via solidlsp.language_server_registration entry points). + # priorities are user-configurable per-key via serena_config.ls_priorities (works for both kinds). + ls_priorities: dict[LanguageServerIdLike, int] = {} + registry = LanguageServerRegistry.get_instance() + for ls_id in registry.iter_registered_ls_ids(): + priority = serena_config.get_ls_priority(ls_id) if priority > 0: - ls_priorities[language] = priority + ls_priorities[ls_id] = priority log.debug("Language server priorities: %s", ls_priorities) ls_composition = compute_language_server_support_composition(project_root, list(ls_priorities.keys())) @@ -394,7 +403,7 @@ class ProjectConfig(SharedConfig, ModeSelectionDefinitionWithAddedModes): if len(other_language_pairs) > 0 and interactive: print( "Detected and enabled main language server '%s' (%.2f%% of source files)." - % (top_language_pair[0].value, top_language_pair[1]) + % (top_language_pair[0].get_key(), top_language_pair[1]) ) print(f"Additionally detected {len(other_language_pairs)} other applicable language servers.\n") print("Note: Enable only servers for languages you need symbolic retrieval/editing capabilities for.") @@ -402,7 +411,7 @@ class ProjectConfig(SharedConfig, ModeSelectionDefinitionWithAddedModes): print(" system-level installations/configuration (see Serena documentation).") print("\nWhich additional language servers do you want to enable?") for ls_id, perc in other_language_pairs: - enable = ask_yes_no("Enable %s (%.2f%% of source files)?" % (ls_id.value, perc), default=False) + enable = ask_yes_no("Enable %s (%.2f%% of source files)?" % (ls_id.get_key(), perc), default=False) if enable: language_servers_to_use.append(ls_id) print() @@ -416,7 +425,7 @@ class ProjectConfig(SharedConfig, ModeSelectionDefinitionWithAddedModes): project_root: str | Path, serena_config: "SerenaConfig", project_name: str | None = None, - languages: list[LanguageServerId] | None = None, + languages: list[LanguageServerIdLike] | None = None, save_to_disk: bool = True, interactive: bool = False, asynchronous: bool = False, @@ -455,7 +464,7 @@ class ProjectConfig(SharedConfig, ModeSelectionDefinitionWithAddedModes): ) languages_to_use = [l.get_key() for l in determined_languages] else: - languages_to_use = [lang.value for lang in languages] + languages_to_use = [lang.get_key() for lang in languages] config_with_comments, _ = cls._load_yaml_dict(PROJECT_TEMPLATE_FILE) config_with_comments["project_name"] = project_name config_with_comments["language_servers"] = languages_to_use @@ -611,7 +620,9 @@ class ProjectConfig(SharedConfig, ModeSelectionDefinitionWithAddedModes): raise ValueError(f"symbol_info_budget cannot be negative, got: {symbol_info_budget}") language_backend_value = data.get("language_backend") - language_backend = LanguageBackend.from_str(language_backend_value) if language_backend_value else None + language_backend = LanguageBackendRegistry.get_instance().resolve(language_backend_value) if language_backend_value else None + agent_interface_value = data.get("agent_interface") + agent_interface = AgentInterface.from_str(agent_interface_value) if agent_interface_value else None line_ending_value = data.get("line_ending") line_ending = LineEnding.from_str(line_ending_value) if line_ending_value else None @@ -621,6 +632,8 @@ class ProjectConfig(SharedConfig, ModeSelectionDefinitionWithAddedModes): fixed_tools = data["fixed_tools"] or [] excluded_tools = data["excluded_tools"] or [] included_optional_tools = data["included_optional_tools"] or [] + excluded_apis = data.get("excluded_apis") or [] + included_apis = data.get("included_apis") or [] additional_workspace_folders = data.get("ls_additional_workspace_folders") or [] if "base_modes" in data and data["base_modes"] is not None: @@ -635,6 +648,8 @@ class ProjectConfig(SharedConfig, ModeSelectionDefinitionWithAddedModes): excluded_tools=excluded_tools, fixed_tools=fixed_tools, included_optional_tools=included_optional_tools, + excluded_apis=excluded_apis, + included_apis=included_apis, read_only=data["read_only"], read_only_memory_patterns=data.get("read_only_memory_patterns", []), ignored_memory_patterns=data.get("ignored_memory_patterns", []), @@ -643,6 +658,7 @@ class ProjectConfig(SharedConfig, ModeSelectionDefinitionWithAddedModes): encoding=data["encoding"], line_ending=line_ending, language_backend=language_backend, + agent_interface=agent_interface, added_modes=data["added_modes"], default_modes=data["default_modes"], symbol_info_budget=symbol_info_budget, @@ -666,7 +682,8 @@ class ProjectConfig(SharedConfig, ModeSelectionDefinitionWithAddedModes): # map fields using non-primitive types to a YAML-compatible representation d["language_servers"] = [lang.get_key() for lang in self.language_servers] - d["language_backend"] = self.language_backend.value if self.language_backend is not None else None + d["language_backend"] = self.language_backend.get_key() if self.language_backend is not None else None + d["agent_interface"] = self.agent_interface.value if self.agent_interface is not None else None d["line_ending"] = self.line_ending.value if self.line_ending is not None else None return d @@ -870,6 +887,11 @@ class SerenaConfig(SharedConfig, ModeSelectionDefinitionWithBaseModes): # *** fields that are mapped directly to/from the configuration file (DO NOT RENAME) *** projects: list[RegisteredProject] = field(default_factory=list) + auth_secret: str = field(default_factory=lambda: str(uuid4()), repr=False) + """ + shared secret for authenticating communication between Serena components and services. + A random UUID is generated and persisted when the configuration setting is missing or empty. + """ gui_log_window: bool = False log_level: int = logging.INFO trace_lsp_communication: bool = False @@ -932,7 +954,13 @@ class SerenaConfig(SharedConfig, ModeSelectionDefinitionWithBaseModes): # settings with overridden defaults - language_backend: LanguageBackend = LanguageBackend.LSP + agent_interface: AgentInterface = AgentInterface.TOOLS + """ + the agent interface to use (unless overridden by the active project's configuration). + Defaults to TOOLS for backward compatibility (as users without this settings will get this default). + The default for new users is defined in the template file. + """ + language_backend: LanguageBackend = field(default_factory=lambda: BuiltinLanguageBackend.LSP.get_instance()) """ the language backend to use for code understanding features """ @@ -960,7 +988,7 @@ class SerenaConfig(SharedConfig, ModeSelectionDefinitionWithBaseModes): # *** static members *** CONFIG_FILE = "serena_config.yml" - CONFIG_FIELDS_WITH_TYPE_CONVERSION = {"projects", "language_backend", "line_ending"} + CONFIG_FIELDS_WITH_TYPE_CONVERSION = {"projects", "language_backend", "agent_interface", "line_ending"} # *** methods *** @classmethod @@ -1035,6 +1063,17 @@ class SerenaConfig(SharedConfig, ModeSelectionDefinitionWithBaseModes): log.info(f"Serena configuration file not found at {config_file_path}, autogenerating...") cls._generate_config_file(config_file_path) + # restrict access to the owner's read/write permissions (as the config file contains secrets) + if os.name == "posix": + current_mode = stat.S_IMODE(os.stat(config_file_path).st_mode) + if current_mode != 0o600: + try: + os.chmod(config_file_path, 0o600) + except Exception as e: + log.error("Failed to restrict permissions of Serena configuration %s to 0600: %s", config_file_path, e) + else: + log.info("Changed permissions of Serena configuration %s from %04o to 0600", config_file_path, current_mode) + # load the configuration log.info(f"Loading Serena configuration from {config_file_path}") try: @@ -1057,6 +1096,11 @@ class SerenaConfig(SharedConfig, ModeSelectionDefinitionWithBaseModes): assert hasattr(instance, field_name) setattr(instance, field_name, get_value_or_default(field_name)) + # generate a persistent authentication secret for explicitly unset settings + if not instance.auth_secret: + instance.auth_secret = str(uuid4()) + num_migrations += 1 + # read projects if "projects" not in loaded_commented_yaml: raise SerenaConfigError("`projects` key not found in Serena configuration. Please update your `serena_config.yml` file.") @@ -1100,16 +1144,25 @@ class SerenaConfig(SharedConfig, ModeSelectionDefinitionWithBaseModes): language_backend = get_dataclass_default(SerenaConfig, "language_backend") if "language_backend" in loaded_commented_yaml: backend_str = loaded_commented_yaml["language_backend"] - language_backend = LanguageBackend.from_str(backend_str) + language_backend = LanguageBackendRegistry.get_instance().resolve(backend_str) else: # backward compatibility (migrate Boolean field "jetbrains") if "jetbrains" in loaded_commented_yaml: num_migrations += 1 if loaded_commented_yaml["jetbrains"]: - language_backend = LanguageBackend.JETBRAINS + language_backend = BuiltinLanguageBackend.JETBRAINS.get_instance() del loaded_commented_yaml["jetbrains"] instance.language_backend = language_backend + # determine agent interface + agent_interface: AgentInterface | None = get_dataclass_default(SerenaConfig, "agent_interface") + if "agent_interface" in loaded_commented_yaml: + agent_interface_value = loaded_commented_yaml["agent_interface"] + agent_interface = AgentInterface.from_str(agent_interface_value) + else: + num_migrations += 1 + instance.agent_interface = agent_interface + # determine line ending line_ending_value = loaded_commented_yaml.get("line_ending") if line_ending_value: @@ -1164,17 +1217,25 @@ class SerenaConfig(SharedConfig, ModeSelectionDefinitionWithBaseModes): log.error(f"Error migrating configuration file: {e}") return None + def set_builtin_language_backend(self, backend: BuiltinLanguageBackend) -> None: + """ + Sets the built-in language backend to use for code understanding features. + + :param backend: the language backend to set + """ + self.language_backend = backend.get_instance() + @classmethod - def init(cls, language_backend: LanguageBackend) -> "SerenaConfig": + def init(cls, builtin_language_backend: BuiltinLanguageBackend) -> "SerenaConfig": """ Supports the config initialisation CLI command, allowing the user to configure fundamental settings before the first launch. - :param language_backend: the language backend to use + :param builtin_language_backend: the language backend to use :return: the created SerenaConfig instance """ config = cls.from_config_file() - config.language_backend = language_backend + config.language_backend = builtin_language_backend.get_instance() config._save() return config @@ -1191,11 +1252,11 @@ class SerenaConfig(SharedConfig, ModeSelectionDefinitionWithBaseModes): self.jetbrains_launch_command = None return self - @cached_property + @property def project_paths(self) -> list[str]: return sorted(str(project.project_root) for project in self.projects) - @cached_property + @property def project_names(self) -> list[str]: return sorted(project.project_config.project_name for project in self.projects) @@ -1247,6 +1308,21 @@ class SerenaConfig(SharedConfig, ModeSelectionDefinitionWithBaseModes): self.projects.append(registered_project) self._persist_projects() + def remove_registered_project(self, registered_project: RegisteredProject) -> None: + """ + Removes the given registered project, persisting the updated project list. + Only the registry entry is removed; the project's own files, including its project + configuration file, are left untouched. + + Unlike :meth:`remove_project`, which resolves the project by name, this removes the + given entry itself and is therefore unambiguous when several registered projects + share a name. + + :param registered_project: the project to remove, which must be an element of :attr:`projects` + """ + self.projects.remove(registered_project) + self._persist_projects() + def add_project_from_path(self, project_root: Path | str, asynchronous_autogen: bool = False) -> "Project": """ Adds a new project to the Serena configuration from a given path, auto-generating the project @@ -1351,7 +1427,10 @@ class SerenaConfig(SharedConfig, ModeSelectionDefinitionWithBaseModes): commented_yaml["projects"] = sorted({str(project.project_root) for project in self.projects}) # convert language backend to string - commented_yaml["language_backend"] = self.language_backend.value + commented_yaml["language_backend"] = self.language_backend.get_key() + + # convert agent interface to string (None if not configured) + commented_yaml["agent_interface"] = self.agent_interface.value if self.agent_interface is not None else None # convert line ending to string commented_yaml["line_ending"] = self.line_ending.value @@ -1442,7 +1521,7 @@ class SerenaConfig(SharedConfig, ModeSelectionDefinitionWithBaseModes): """ Propagate settings from this configuration to individual components that are statically configured """ - from serena.tools import JetBrainsPluginClient + from serena.jetbrains.jetbrains_plugin_client import JetBrainsPluginClient JetBrainsPluginClient.set_server_address(self.jetbrains_plugin_server_address) @@ -1459,7 +1538,26 @@ class SerenaConfig(SharedConfig, ModeSelectionDefinitionWithBaseModes): return True return False - def determine_language_backend(self, project_config: ProjectConfig | None = None, log_choice: bool = False): + def determine_agent_interface(self, project_config: "ProjectConfig | None" = None, log_choice: bool = False) -> AgentInterface: + """ + Determines the effective agent interface: the project configuration takes precedence over the global configuration; + if neither configures an interface, the tool interface is used. + + :param project_config: the configuration of the project to be activated, if any + :param log_choice: whether to log the choice + :return: the effective agent interface + """ + if project_config is not None and project_config.agent_interface is not None: + agent_interface, source = project_config.agent_interface, "project configuration" + elif self.agent_interface is not None: + agent_interface, source = self.agent_interface, "global configuration" + else: + agent_interface, source = AgentInterface.TOOLS, "default" + if log_choice: + log.info(f"Using agent interface '{agent_interface.value}' ({source})") + return agent_interface + + def determine_language_backend(self, project_config: ProjectConfig | None = None, log_choice: bool = False) -> LanguageBackend: language_backend = self.language_backend if project_config and project_config.language_backend is not None: language_backend = project_config.language_backend @@ -1470,7 +1568,7 @@ class SerenaConfig(SharedConfig, ModeSelectionDefinitionWithBaseModes): log.info(f"Using language backend from global configuration: {language_backend.name}") return language_backend - def get_ls_priority(self, ls_id: LanguageServerId) -> int: + def get_ls_priority(self, ls_id: LanguageServerIdLike) -> int: """ Gets the priority value associated with a language server @@ -1479,9 +1577,9 @@ class SerenaConfig(SharedConfig, ModeSelectionDefinitionWithBaseModes): """ if self.ls_priorities is not None: try: - configured_value = self.ls_priorities.get(ls_id.value) + configured_value = self.ls_priorities.get(ls_id.get_key()) if configured_value is not None: return int(configured_value) except Exception as e: - log.error("Error reading language priority for %s: %s. Using default priority.", ls_id.value, e) + log.error("Error reading language priority for %s: %s. Using default priority.", ls_id.get_key(), e) return ls_id.get_priority() diff --git a/src/serena/dashboard.py b/src/serena/dashboard.py index 4d0574b8..117482b4 100644 --- a/src/serena/dashboard.py +++ b/src/serena/dashboard.py @@ -27,6 +27,7 @@ from serena.analytics import ToolUsageStats from serena.config.serena_config import SerenaConfig, SerenaPaths from serena.constants import SERENA_DASHBOARD_DIR, SerenaPorts from serena.task_executor import TaskExecutor +from serena.tools import ReadMemoryTool from serena.util.logging import MemoryLogHandler from serena.util.pypi import PyPIPackageInfo from serena.util.pywebview import WebViewWithTray @@ -60,11 +61,25 @@ class ResponseToolStats(BaseModel): stats: dict[str, dict[str, int]] +class ResponseFacadeMethod(BaseModel): + name: str + is_enabled: bool + + +class ResponseFacade(BaseModel): + name: str + is_enabled: bool + methods: list[ResponseFacadeMethod] + + class ResponseConfigOverview(BaseModel): active_project: dict[str, str | None] context: dict[str, str] modes: list[dict[str, str]] active_tools: list[str] + agent_interface: str + language_backend: str + facades: list[ResponseFacade] | None tool_stats_summary: dict[str, dict[str, int]] registered_projects: list[dict[str, str | bool]] available_tools: list[dict[str, str | bool]] @@ -604,9 +619,22 @@ class SerenaDashboardAPI: # Get available memories if ReadMemoryTool is active available_memories = None - if self._agent.tool_is_active("read_memory") and project is not None: + if self._agent.is_tool_function_available(ReadMemoryTool) and project is not None: available_memories = project.memory_manager.list_memories().get_full_list() + # Get the availability of the REPL's facades and their methods (REPL interface only) + facades = None + if self._agent.get_agent_interface().is_repl(): + availability_info = self._agent.get_repl().entrypoint.get_facade_availability_info() + facades = [ + ResponseFacade( + name=facade_info.name, + is_enabled=facade_info.is_enabled, + methods=[ResponseFacadeMethod(name=m.name, is_enabled=m.is_enabled) for m in facade_info.methods], + ) + for facade_info in availability_info.facades + ] + # Get list of languages for the active project ls_ids = [] if project is not None: @@ -622,6 +650,9 @@ class SerenaDashboardAPI: context=context_info, modes=modes_info, active_tools=active_tools, + agent_interface=self._agent.get_agent_interface().value, + language_backend=self._agent.get_language_backend().get_key(), + facades=facades, tool_stats_summary=tool_stats_summary, registered_projects=registered_projects, available_tools=available_tools, @@ -1030,9 +1061,28 @@ class SerenaDashboardTrayManager: log.info("Unregistered instance on port %d", port) return {"status": "unregistered"} + @staticmethod + def _run_in_ui_thread(fn: Callable[[], None]) -> None: + """ + Runs a UI mutation in the thread in which the platform's UI toolkit requires it to run (where necessary). + + On macOS, AppKit demands that mutations of the status item happen on the main thread, and + recent macOS versions terminate the process with SIGTRAP when they do not. The tray manager + reaches such mutations from Flask request handlers and from the alive-check thread, so the + call has to be marshalled. On other platforms it is made directly. + + :param fn: the UI mutation to run + """ + if sys.platform == "darwin": + from PyObjCTools import AppHelper # ty: ignore[unresolved-import] + + AppHelper.callAfter(fn) + else: + fn() + def _update_menu(self) -> None: if self._tray_icon: - self._tray_icon.update_menu() + self._run_in_ui_thread(self._tray_icon.update_menu) def _build_menu_items(self) -> tuple[Any, ...]: """ @@ -1190,7 +1240,7 @@ class SerenaDashboardTrayManager: # set up tray icon with a dynamic menu (callable returns items on each open) kwargs: dict[str, Any] = {} if sys.platform == "darwin": - from AppKit import NSApplication, NSApplicationActivationPolicyAccessory + from AppKit import NSApplication, NSApplicationActivationPolicyAccessory # ty: ignore[unresolved-import] (macOS only) nsapp = NSApplication.sharedApplication() # run as an accessory app so that only the menu bar icon is shown (no Dock icon) diff --git a/src/serena/jetbrains/jetbrains_backend.py b/src/serena/jetbrains/jetbrains_backend.py new file mode 100644 index 00000000..6e62fce9 --- /dev/null +++ b/src/serena/jetbrains/jetbrains_backend.py @@ -0,0 +1,105 @@ +# SPDX-License-Identifier: GPL-3.0-or-later +import logging +from typing import TYPE_CHECKING + +from overrides import override + +from serena.code_editor import JetBrainsCodeEditor +from serena.jetbrains import launch_coordinator +from serena.language_backend import BuiltinLanguageBackend, LanguageBackend + +from ..util.file_proxy import FileProxy, LocalProjectFileProxy +from . import jetbrains_types as jb + +if TYPE_CHECKING: + from serena.agent import SerenaAgent + from serena.code_editor import CodeEditor + from serena.project import Project + from serena.repl.facade import ApiScope, Facade + from serena.tools import Tool + +log = logging.getLogger(__name__) + + +class LanguageBackendJetBrains(LanguageBackend): + def __init__(self): + super().__init__(BuiltinLanguageBackend.JETBRAINS.value) + + @override + def get_lsp_tool_class_replacements(self) -> "dict[type[Tool], type[Tool]]": + from ..tools import jetbrains_tools, symbol_tools + + return { + symbol_tools.FindSymbolTool: jetbrains_tools.JetBrainsFindSymbolTool, + symbol_tools.GetSymbolsOverviewTool: jetbrains_tools.JetBrainsGetSymbolsOverviewTool, + symbol_tools.FindReferencingSymbolsTool: jetbrains_tools.JetBrainsFindReferencingSymbolsTool, + symbol_tools.FindImplementationsTool: jetbrains_tools.JetBrainsFindImplementationsTool, + symbol_tools.FindDeclarationTool: jetbrains_tools.JetBrainsFindDeclarationTool, + symbol_tools.RenameSymbolTool: jetbrains_tools.JetBrainsRenameTool, + symbol_tools.SafeDeleteSymbol: jetbrains_tools.JetBrainsSafeDeleteTool, + } + + @override + def create_facades(self, agent: "SerenaAgent", api_scope: "ApiScope") -> list["Facade"]: + from ..repl.api.jb_api import JetBrainsApi + from ..repl.facade import Facade + + return [Facade.from_api(JetBrainsApi(agent), api_scope)] + + @override + def init_active_project(self, agent: "SerenaAgent") -> None: + project = agent.get_active_project_or_raise() + client = launch_coordinator.find_plugin_server(project) + if client is not None: + log.info("Found Serena JetBrains Plugin server: %s", client) + else: + log.info("Serena JetBrains Plugin server not found for project %s", project.project_name) + launch_command = agent.serena_config.jetbrains_launch_command + if launch_command: + launch_coordinator.launch_and_wait_for_plugin_server(project, launch_command) + + @override + def shutdown_active_project(self, project: "Project", timeout: float) -> None: + pass + + @override + def create_code_editor(self, project: "Project") -> "CodeEditor": + return JetBrainsCodeEditor(project) + + @override + def is_source_file(self, abs_path: str, project: "Project") -> bool: + # no distinction is made; every file is potentially a source file + return True + + @override + def is_external_path(self, relative_path: str) -> bool: + return relative_path.startswith(jb.JB_EXTERNAL_FILE_PREFIX) + + @override + def create_file_proxy(self, relative_path: str, project: "Project") -> FileProxy: + if self.is_external_path(relative_path): + return JetBrainsFileProxy(relative_path, project) + return LocalProjectFileProxy(relative_path, project) + + +class JetBrainsFileProxy(FileProxy): + """ + Retrieves the contents of a file from the JetBrains plugin via the plugin client, given its relative path, + which may be an external path (e.g., "") + """ + + def __init__(self, relative_path: str, project: "Project"): + self._relative_path = relative_path + self._project = project + + def get_contents(self) -> str: + from serena.jetbrains.jetbrains_plugin_client import JetBrainsPluginClient + + client = JetBrainsPluginClient.from_project(self._project) + return client.read_file(self._relative_path) + + def get_relative_path(self) -> str: + return self._relative_path + + def is_glob_supported(self): + return False diff --git a/src/serena/jetbrains/jetbrains_types.py b/src/serena/jetbrains/jetbrains_types.py index 1739e607..5ccd561b 100644 --- a/src/serena/jetbrains/jetbrains_types.py +++ b/src/serena/jetbrains/jetbrains_types.py @@ -8,14 +8,6 @@ Prefix used for in relative paths of symbols that are from external libraries (i """ -def is_external_path(relative_path: str): - """ - :param relative_path: a relative path (e.g., from a symbol's `relative_path` field) - :return: whether the path is an external path (i.e., from a library, not the user's codebase) - """ - return relative_path.startswith(JB_EXTERNAL_FILE_PREFIX) - - class PluginStatusDTO(TypedDict): project_root: str plugin_version: str diff --git a/src/serena/language_backend.py b/src/serena/language_backend.py new file mode 100644 index 00000000..54b98525 --- /dev/null +++ b/src/serena/language_backend.py @@ -0,0 +1,251 @@ +# SPDX-License-Identifier: GPL-3.0-or-later +import importlib +import logging +import threading +from abc import ABC, abstractmethod +from enum import Enum +from functools import cache +from typing import TYPE_CHECKING + +from serena.util.file_proxy import FileProxy + +log = logging.getLogger(__name__) + +if TYPE_CHECKING: + from serena.agent import SerenaAgent + from serena.code_editor import CodeEditor + from serena.project import Project + from serena.repl.facade import ApiScope, Facade + from serena.tools import Tool + + +class LanguageBackend(ABC): + def __init__(self, key: str): + """ + :param key: the key by which the backend is identified in the registry and in configuration + """ + self._key = key + + def __str__(self): + return self._key + + def get_key(self) -> str: + """ + :return: the key by which the backend is identified in the registry and in configuration + """ + return self._key + + def is_lsp(self): + return self.get_key() == BuiltinLanguageBackend.LSP.value + + def is_jetbrains(self): + return self.get_key() == BuiltinLanguageBackend.JETBRAINS.value + + @property + def name(self): + return self.get_key() + + @abstractmethod + def get_lsp_tool_class_replacements(self) -> "dict[type[Tool], type[Tool]]": + """ + :return: mapping from LSP tool classes to replacement tool classes (functional replacements) + """ + + @abstractmethod + def create_facades(self, agent: "SerenaAgent", api_scope: "ApiScope") -> list["Facade"]: + """ + Creates backend-specific facades for the given agent and API scope. + + :param agent: the agent + :param api_scope: the API scope defining active facade methods + :return: the list of facades to be used by the agent for this backend + """ + + @abstractmethod + def init_active_project(self, agent: "SerenaAgent") -> None: + """ + Initialises the backend for the given agent's newly activated project. + + :param agent: the agent, which has just set a new active project + """ + + @abstractmethod + def shutdown_active_project(self, project: "Project", timeout: float) -> None: + """ + Cleans up, freeing resources, after a project has been deactivated. + + :param project: the project + :param timeout: the timeout, in seconds, after which to give up on graceful shutdown + """ + + def get_project_activation_statement(self, project: "Project") -> str: + """ + :return: a statement to add to the project activation message + """ + return "" + + def get_config_overview_statement(self, project: "Project") -> str: + """ + :return: a statement to add to the project configuration overview + """ + return "" + + @abstractmethod + def create_code_editor(self, project: "Project") -> "CodeEditor": + pass + + @abstractmethod + def is_source_file(self, abs_path: str, project: "Project") -> bool: + """ + Determines whether the given absolute path corresponds to a source file that can (potentially) be processed/understood by the backend. + + :param abs_path: the absolute path to an existing file + :param project: the project in which the file is located + :return: True if the file is a source file for this backend (or the backend does not specifically make distinctions), + False otherwise + """ + + @abstractmethod + def is_external_path(self, relative_path: str) -> bool: + """ + Determines whether the given relative path corresponds to a file that is external to the project (e.g. a dependency file). + Virtually all of Serena's interfaces use `relative_path` (relative to the project root) to refer to files, but some backends + may need to support project-external files. In this case, the external path should be encoded in the `relative_path` parameter + (e.g. "") rather than this being an actual relative path that points outside the project root. + Therefore, information about the project in question is deliberately not provided to this method. + + :param relative_path: the relative path to a file within the project or an encoded external path. + The path can be assumed to have been provided by the backend itself. + :return: whether the file is considered external to the project by this backend + """ + + @abstractmethod + def create_file_proxy(self, relative_path: str, project: "Project") -> FileProxy: + """ + Creates a file proxy for the given relative path in the given project. + + :param relative_path: the relative path to a file within the project or an encoded external path. + :param project: the project + :return: a file proxy for the given file + """ + + +class BuiltinLanguageBackend(Enum): + LSP = "LSP" + """ + Use the language server protocol (LSP), spawning freely available language servers + via the SolidLSP library that is part of Serena + """ + JETBRAINS = "JetBrains" + """ + Use the Serena plugin in your JetBrains IDE. + (requires the plugin to be installed and the project being worked on to be open in your IDE) + """ + + @staticmethod + def from_str(backend_str: str) -> "BuiltinLanguageBackend": + for backend in BuiltinLanguageBackend: + if backend.value.lower() == backend_str.lower(): + return backend + raise ValueError(f"Unknown language backend '{backend_str}': valid values are {[b.value for b in BuiltinLanguageBackend]}") + + @cache + def get_instance(self) -> LanguageBackend: + if self == BuiltinLanguageBackend.LSP: + from .lsp.lsp_backend import LanguageBackendLSP + + return LanguageBackendLSP() + elif self == BuiltinLanguageBackend.JETBRAINS: + from .jetbrains.jetbrains_backend import LanguageBackendJetBrains + + return LanguageBackendJetBrains() + else: + raise NotImplementedError + + +class LanguageBackendRegistry: + """ + Registry of language backends + """ + + REGISTRATION_ENTRY_POINT_GROUP = "serena.language_backend_registration" + """ + entry point group for language backend registration functions; each function should call use + `LanguageBackendRegistry.get_instance().register(...)` to register a backend + """ + + _instance = None + _instance_lock = threading.Lock() + + @classmethod + def get_instance(cls): + if cls._instance is None: + with cls._instance_lock: + if cls._instance is None: + cls._instance = cls(True) + cls._discover_backends_from_entry_points() + return cls._instance + + def __init__(self, _singleton: bool): + if not _singleton: + raise RuntimeError("LanguageServerRegistry is a singleton. Use get_instance() to access it.") + self._registered_backends: dict[str, LanguageBackend] = {} + + # auto-register built-in language backends + for builtin_backend in BuiltinLanguageBackend: + self._registered_backends[builtin_backend.value] = builtin_backend.get_instance() + + @classmethod + def _discover_backends_from_entry_points(cls) -> None: + """ + Discover and execute language server adapter registration functions from entry points. + """ + log.debug("Discovering language backend registration entry points ...") + try: + entry_points = importlib.metadata.entry_points(group=cls.REGISTRATION_ENTRY_POINT_GROUP) + except Exception as error: + log.exception("Failed to discover language server registration entry points: %s", error) + return + + def get_distribution_name(ep: importlib.metadata.EntryPoint) -> str: + distribution = getattr(ep, "dist", None) + if distribution is None: + return "unknown distribution" + return distribution.name or "unknown distribution" + + log.debug("Found %d language server registration entry points", len(entry_points)) + for entry_point in entry_points: + try: + registration = entry_point.load() + if not callable(registration): + raise TypeError("Entry point must resolve to a callable registration function") + registration() + except Exception as error: + log.exception( + "Failed to load language backend entry point '%s' from %s: %s", + entry_point.name, + get_distribution_name(entry_point), + error, + ) + + def resolve(self, key: str) -> LanguageBackend: + if key in self._registered_backends: + return self._registered_backends[key] + raise ValueError(f"Unknown language backend key: '{key}'; Valid keys: {self.get_keys()}") + + def register(self, backend: LanguageBackend, allow_override: bool = False) -> None: + """ + :param backend: the backend to register + :param allow_override: whether to allow overriding an existing registration with the same key + """ + key = backend.get_key() + log.info("Registering language backend: %s (class=%s)", key, backend.__class__.__name__) + if backend.get_key() in self._registered_backends and not allow_override: + raise ValueError(f"Language backend already registered: {key}") + self._registered_backends[key] = backend + + def get_keys(self) -> list[str]: + """ + :return: the sorted list of all registered string keys + """ + return sorted(self._registered_backends.keys()) diff --git a/src/serena/lsp/lsp_backend.py b/src/serena/lsp/lsp_backend.py new file mode 100644 index 00000000..1d672195 --- /dev/null +++ b/src/serena/lsp/lsp_backend.py @@ -0,0 +1,80 @@ +# SPDX-License-Identifier: GPL-3.0-or-later +import logging +from typing import TYPE_CHECKING + +from overrides import override +from sensai.util.logging import LogTime + +from serena.language_backend import BuiltinLanguageBackend, LanguageBackend +from serena.util.file_proxy import FileProxy, LocalProjectFileProxy + +if TYPE_CHECKING: + from serena.agent import SerenaAgent + from serena.code_editor import LanguageServerCodeEditor + from serena.project import Project + from serena.repl.facade import ApiScope, Facade + from serena.tools import Tool + +log = logging.getLogger(__name__) + + +class LanguageBackendLSP(LanguageBackend): + def __init__(self): + super().__init__(BuiltinLanguageBackend.LSP.value) + + @override + def get_lsp_tool_class_replacements(self) -> "dict[type[Tool], type[Tool]]": + return {} + + @override + def create_facades(self, agent: "SerenaAgent", api_scope: "ApiScope") -> list["Facade"]: + from ..repl.api.lsp_api import LspApi + from ..repl.facade import Facade + + return [Facade.from_api(LspApi(agent), api_scope)] + + @override + def init_active_project(self, agent: "SerenaAgent") -> None: + with LogTime("Language server initialization", logger=log): + agent.reset_language_server_manager() + + @override + def shutdown_active_project(self, project: "Project", timeout: float) -> None: + # nothing to do; the language server manager is already shut down by the project itself + pass + + @override + def get_project_activation_statement(self, project: "Project") -> str: + language_servers_str = ", ".join([ls.get_key() for ls in project.project_config.language_servers]) + return f"Active language servers: {language_servers_str}.\n" + + @override + def get_config_overview_statement(self, project: "Project") -> str: + return f"Language server status: {project.get_language_server_manager_status()}\n" + + @override + def create_code_editor(self, project: "Project") -> "LanguageServerCodeEditor": + from serena.code_editor import LanguageServerCodeEditor + from serena.symbol import LanguageServerSymbolRetriever + + symbol_retriever = LanguageServerSymbolRetriever(project) + return LanguageServerCodeEditor(symbol_retriever) + + @override + def is_source_file(self, abs_path: str, project: "Project") -> bool: + is_file_in_supported_languages = False + for language in project.project_config.language_servers: + fn_matcher = language.get_source_fn_matcher() + if fn_matcher.is_relevant_filename(abs_path): + is_file_in_supported_languages = True + break + return is_file_in_supported_languages + + @override + def is_external_path(self, relative_path: str) -> bool: + # LSP backend currently uses only true project-relative paths + return False + + @override + def create_file_proxy(self, relative_path: str, project: "Project") -> FileProxy: + return LocalProjectFileProxy(relative_path, project) diff --git a/src/serena/util/ls_diagnostics.py b/src/serena/lsp/lsp_diagnostics.py similarity index 76% rename from src/serena/util/ls_diagnostics.py rename to src/serena/lsp/lsp_diagnostics.py index c9f80d3f..4d0a0844 100644 --- a/src/serena/util/ls_diagnostics.py +++ b/src/serena/lsp/lsp_diagnostics.py @@ -3,12 +3,14 @@ import json from collections.abc import Iterable from dataclasses import dataclass -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING, Any, Optional, Self +from serena.util.text_utils import TextOutputUtils from solidlsp import ls_types from solidlsp.lsp_protocol_handler.lsp_types import DiagnosticSeverity if TYPE_CHECKING: + from serena.agent import SerenaAgent from serena.symbol import LanguageServerSymbolRetriever @@ -203,3 +205,55 @@ class DiagnosticsDiff: def get_grouped_diagnostics(self) -> GroupedDiagnostics: return self._grouped_diagnostics + + +class DiagnosticsContext: + ENABLE_DIAGNOSTICS_DEFAULT: bool = False + """ + Global flag to enable/disable diagnostics for LSP-based editing tools derived from this class. + The feature is currently disabled, because per-edit diagnostics are a questionable feature, since individual + edits often intentionally introduce diagnostics (e.g. function signature mismatches or even syntax errors) that + are then resolved in subsequent edits. + """ + + DIAGNOSTICS_KEY = "diagnostics[warning-or-higher]" + + def __init__(self, agent: "SerenaAgent", *edited_relative_paths: str, enable: bool = ENABLE_DIAGNOSTICS_DEFAULT) -> None: + self._is_diagnostics_enabled = enable and agent.get_language_backend() + self._edited_files = [EditedFilePath(path, path) for path in edited_relative_paths] + self._before_edit_diagnostics_snapshot: PublishedDiagnosticsSnapshot | None = None + self._symbol_retriever: Optional["LanguageServerSymbolRetriever"] | None = None + if self._is_diagnostics_enabled: + from serena.symbol import LanguageServerSymbolRetriever # local import to avoid a circular dependency + + self._symbol_retriever = LanguageServerSymbolRetriever(agent.get_active_project_or_raise()) + self._before_edit_diagnostics_snapshot = PublishedDiagnosticsSnapshot(self._edited_files, self._symbol_retriever) + + def __enter__(self) -> Self: + return self + + def __exit__(self, exc_type, exc_val, exc_tb): + pass + + def format_result( + self, + base_result: str, + ) -> str: + if not self._is_diagnostics_enabled: + return base_result + + if self._before_edit_diagnostics_snapshot is None: + return base_result + + assert self._symbol_retriever is not None + diagnostics_diff = DiagnosticsDiff(self._before_edit_diagnostics_snapshot, self._edited_files, self._symbol_retriever) + grouped_diagnostics = diagnostics_diff.get_grouped_diagnostics().get_dict() + + if not grouped_diagnostics: + return base_result + else: + result_dict = { + "result": base_result, + self.DIAGNOSTICS_KEY: grouped_diagnostics, + } + return TextOutputUtils.to_json(result_dict) diff --git a/src/serena/mcp.py b/src/serena/mcp.py index 1bd2844b..db8bff75 100644 --- a/src/serena/mcp.py +++ b/src/serena/mcp.py @@ -11,23 +11,23 @@ from dataclasses import dataclass from typing import Any, Literal, cast import docstring_parser -from mcp.server.fastmcp import server -from mcp.server.fastmcp.exceptions import ToolError -from mcp.server.fastmcp.server import Context, FastMCP, Settings -from mcp.server.fastmcp.tools.base import Tool as FastMCPTool -from mcp.server.session import ServerSessionT -from mcp.shared.context import LifespanContextT, RequestT +from mcp.server.mcpserver import server +from mcp.server.mcpserver.context import LifespanContextT, RequestT +from mcp.server.mcpserver.exceptions import ToolError +from mcp.server.mcpserver.server import Context +from mcp.server.mcpserver.server import MCPServer as FastMCP +from mcp.server.mcpserver.tools.base import Tool as FastMCPTool from mcp.types import ToolAnnotations -from pydantic_settings import SettingsConfigDict from sensai.util import logging -from serena import __version__ +from serena import __version__ as serena_version_str from serena.agent import ( SerenaAgent, ) from serena.config.context_mode import SerenaAgentContext -from serena.config.serena_config import LanguageBackend, ModeSelectionDefinition, SerenaConfig +from serena.config.serena_config import AgentInterface, ModeSelectionDefinition, SerenaConfig from serena.constants import DEFAULT_CONTEXT, SERENA_LOG_FORMAT +from serena.language_backend import LanguageBackend from serena.tools import Tool, ToolCallError from serena.util.exception import show_fatal_exception_safe from serena.util.logging import MemoryLogHandler @@ -109,8 +109,8 @@ class SerenaFastMCPTool(FastMCPTool): can_edit = tool.can_edit() annotations = ToolAnnotations( title=tool_title, - readOnlyHint=not can_edit, - destructiveHint=can_edit, + read_only_hint=not can_edit, + destructive_hint=can_edit, ) super().__init__( @@ -132,7 +132,7 @@ class SerenaFastMCPTool(FastMCPTool): async def run( self, arguments: dict[str, Any], - context: Context[ServerSessionT, LifespanContextT, RequestT] | None = None, + context: Context[LifespanContextT, RequestT], convert_result: bool = False, ) -> Any: # apply parameter aliases @@ -322,10 +322,9 @@ class SerenaMCPFactory: def create_mcp_server( self, - host: str = "127.0.0.1", - port: int = 8000, mode_selection_def: ModeSelectionDefinition | None = None, language_backend: LanguageBackend | None = None, + agent_interface: AgentInterface | None = None, enable_web_dashboard: bool | None = None, enable_gui_log_window: bool | None = None, open_web_dashboard: bool | None = None, @@ -337,10 +336,9 @@ class SerenaMCPFactory: """ Create an MCP server with process-isolated SerenaAgent to prevent asyncio contamination. - :param host: The host to bind to - :param port: The port to bind to :param mode_selection_def: the mode selection definition to apply :param language_backend: the language backend to use, overriding the configuration setting. + :param agent_interface: the agent interface to use, overriding the configuration setting. :param enable_web_dashboard: Whether to enable the web dashboard. If not specified, will take the value from the serena configuration. :param enable_gui_log_window: Whether to enable the GUI log window. It currently does not work on macOS, and setting this to True will be ignored then. If not specified, will take the value from the serena configuration. @@ -371,6 +369,8 @@ class SerenaMCPFactory: config.tool_timeout = tool_timeout if language_backend is not None: config.language_backend = language_backend + if agent_interface is not None: + config.agent_interface = agent_interface self.agent = self._create_serena_agent(config, modes=mode_selection_def, project_activation_error=project_activation_error) @@ -378,23 +378,15 @@ class SerenaMCPFactory: show_fatal_exception_safe(e) raise - # Override model_config to disable the use of `.env` files for reading settings, because user projects are likely to contain - # `.env` files (e.g. containing LOG_LEVEL) that are not supposed to override the MCP settings; - # retain only FASTMCP_ prefix for already set environment variables. - Settings.model_config = SettingsConfigDict(env_prefix="FASTMCP_") instructions = self._get_initial_instructions() log.info("MCP server initial instructions:\n%s", instructions) mcp = FastMCP( name="Serena", + version=serena_version_str, lifespan=self.server_lifespan, website_url="https://oraios.github.io/serena", - host=host, - port=port, instructions=instructions, ) - # FastMCP currently falls back to the installed mcp SDK version when no version is set. - # Set the low-level server value explicitly so MCP clients identify Serena correctly. - mcp._mcp_server.version = __version__ return mcp @asynccontextmanager diff --git a/src/serena/memories/memory_manager.py b/src/serena/memories/memory_manager.py index a112180e..c185ff18 100644 --- a/src/serena/memories/memory_manager.py +++ b/src/serena/memories/memory_manager.py @@ -330,6 +330,7 @@ class MemoryManager: new_name = self._sanitize_name(new_name) self._check_not_ignored(old_name) self._check_not_ignored(new_name) + self._check_write_access(old_name, is_tool_context) self._check_write_access(new_name, is_tool_context) old_path = self.get_memory_file_path(old_name) @@ -347,22 +348,32 @@ class MemoryManager: def rename_memory_and_propagate_references(self, old_name: str, new_name: str, is_tool_context: bool) -> tuple[str, int]: """ - Renames a memory and updates every ``mem:OLD_NAME`` reference across all memories. + Renames a memory and updates every ``mem:OLD_NAME`` reference in the memories which + accept writes in the given context. Memories whose content does not contain a reference to ``old_name`` are left - untouched (no spurious mtime changes). Memories that do are rewritten via - :meth:`save_memory`. + untouched (no spurious mtime changes); those that do are rewritten via + :meth:`save_memory`. References in a memory which does not accept writes (a read-only + memory in a tool context) are not affected and remain reported as stale by + :meth:`validate_referential_integrity`. :param old_name: the current memory name (the source of the rename) :param new_name: the target memory name :param is_tool_context: forwarded to :meth:`save_memory` for read-only enforcement :return: a tuple of (rename message returned by :meth:`move_memory`, total number of - ``mem:`` reference occurrences rewritten across all memories). + ``mem:`` reference occurrences rewritten in those memories). """ renaming_message = self.move_memory(old_name, new_name, is_tool_context=is_tool_context) + # propagate the reference, enumerating after the move such that the renamed memory + # itself is covered; the read-only memories are excluded in a tool context because + # writing to one would raise after the move was already applied, leaving the memory + # graph half-updated + memories_list = self.list_memories() + target_names = sorted(memories_list.memories) if is_tool_context else memories_list.get_full_list() + total_updates = 0 - for memory_name in self.list_memories().get_full_list(): + for memory_name in target_names: content = self.load_memory(memory_name) updated_content, n_replacements = self.rename_references_to_memory(content, old_name, new_name) if n_replacements > 0: diff --git a/src/serena/project.py b/src/serena/project.py index 7341d891..5f2570e5 100644 --- a/src/serena/project.py +++ b/src/serena/project.py @@ -12,11 +12,11 @@ from sensai.util.logging import LogTime from sensai.util.string import TextBuilder, ToStringMixin from serena.config.serena_config import ( - LanguageBackend, ProjectConfig, ProjectConfigAutoGenerationMode, SerenaConfig, ) +from serena.language_backend import LanguageBackend from serena.ls_manager import LanguageServerFactory, LanguageServerManager from serena.memories.memory_manager import MemoryManager from serena.util.file_proxy import FileCollection, FileProxy @@ -128,7 +128,7 @@ class Project(ToStringMixin): @property def language_backend(self) -> LanguageBackend: # The backend configuration is fundamentally owned by the agent, so it takes - # precedence. (Note: The agent does not necessary honour the project's choice, + # precedence. (Note: The agent does not necessarily honour the project's choice, # as it may be invalid.) if self._agent is not None: return self._agent.get_language_backend() @@ -212,7 +212,9 @@ class Project(ToStringMixin): ) return self.__ignored_patterns - def _is_ignored_relative_path(self, relative_path: str | Path, ignore_non_source_files: bool = True) -> bool: + def _is_ignored_relative_path( + self, relative_path: str | Path, ignore_non_source_files: bool = True, is_file: bool | None = None + ) -> bool: """ Determine whether a path should be ignored based on file type and ignore patterns. Returns False for non-existent paths since they cannot be matched by ignore patterns. @@ -220,6 +222,7 @@ class Project(ToStringMixin): :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 + :param is_file: whether the path exists and is a file, for callers that already know :return: whether the path should be ignored """ @@ -230,24 +233,19 @@ class Project(ToStringMixin): return False abs_path = os.path.join(self.project_root, relative_path) - if not os.path.exists(abs_path): - log.debug(f"Path {abs_path} does not exist, skipping ignore check") - return False + if is_file is None: + if not os.path.exists(abs_path): + log.debug(f"Path {abs_path} does not exist, skipping ignore check") + return False # check code file restriction (depending on backend) if ignore_non_source_files: - # apply restriction only for LSP backend, which enumerates known languages - # and therefore can determine whether a file is a source file or not - if self.language_backend.is_lsp(): - if os.path.isfile(abs_path): - is_file_in_supported_language = False - for language in self.project_config.language_servers: - fn_matcher = language.get_source_fn_matcher() - if fn_matcher.is_relevant_filename(abs_path): - is_file_in_supported_language = True - break - if not is_file_in_supported_language: - return True + if is_file is None: + is_file = os.path.isfile(abs_path) + if is_file: + # non-source files are ignored + if not self.language_backend.is_source_file(abs_path, self): + return True # Create normalized path for consistent handling rel_path = Path(relative_path) @@ -256,15 +254,18 @@ class Project(ToStringMixin): if len(rel_path.parts) > 0 and ".git" in rel_path.parts: return True - return match_path(str(relative_path), self._ignore_spec, root_path=self.project_root) + is_dir = None if is_file is None else not is_file + return match_path(str(relative_path), self._ignore_spec, root_path=self.project_root, is_dir=is_dir) - def is_ignored_path(self, path: str | Path, ignore_non_source_files: bool = False) -> bool: + def is_ignored_path(self, path: str | Path, ignore_non_source_files: bool = False, is_file: bool | None = None) -> bool: """ Checks whether the given path is ignored :param path: the path to check, can be absolute or relative :param ignore_non_source_files: whether to ignore files that are not source files (according to the file masks determined by the project's programming language) + :param is_file: whether the path exists and is a file, for callers that already know; + see :meth:`_is_ignored_relative_path`. `None` determines it from the filesystem. """ path = Path(path) if path.is_absolute(): @@ -278,7 +279,7 @@ class Project(ToStringMixin): else: relative_path = path - return self._is_ignored_relative_path(str(relative_path), ignore_non_source_files=ignore_non_source_files) + return self._is_ignored_relative_path(str(relative_path), ignore_non_source_files=ignore_non_source_files, is_file=is_file) def get_is_ignored_path_fn(self, base_path: str, skip_ignored_paths: bool) -> Callable[[str], bool]: """ @@ -336,7 +337,7 @@ class Project(ToStringMixin): :param relative_path: the path to validate, relative to the project root :param require_not_ignored: if True, the path must not be ignored according to the project's ignore settings """ - if FileProxy.is_external_path(relative_path): + if FileProxy.is_external_path(relative_path, self): return if not self.is_path_in_project(relative_path): @@ -358,15 +359,17 @@ class Project(ToStringMixin): if os.path.isfile(start_path): return [relative_path] else: + # os.walk hands back directories and files separately, so `is_file` is already known here and + # does not have to be re-derived from the filesystem for every one of them. for root, dirs, files in os.walk(start_path, followlinks=True): # prevent recursion into ignored directories - dirs[:] = [d for d in dirs if not self.is_ignored_path(os.path.join(root, d))] + dirs[:] = [d for d in dirs if not self.is_ignored_path(os.path.join(root, d), is_file=False)] # collect non-ignored files for file in files: abs_file_path = os.path.join(root, file) try: - if not self.is_ignored_path(abs_file_path, ignore_non_source_files=True): + if not self.is_ignored_path(abs_file_path, ignore_non_source_files=True, is_file=True): try: rel_file_path = os.path.relpath(abs_file_path, start=self.project_root) except Exception: @@ -383,7 +386,7 @@ class Project(ToStringMixin): ) return rel_file_paths - def _create_file_collection(self, relative_path: str, *, code_files_only: bool, skip_ignored_files: bool) -> FileCollection: + def create_file_collection(self, relative_path: str, *, code_files_only: bool, skip_ignored_files: bool) -> FileCollection: """ Creates the file collection for the given relative path. @@ -392,7 +395,7 @@ class Project(ToStringMixin): :param skip_ignored_files: whether to skip ignored files; has no effect if `code_files_only` is True :return: """ - if FileProxy.is_external_path(relative_path): + if FileProxy.is_external_path(relative_path, self): # single external path: create appropriate proxy file_collection = FileCollection([FileProxy.from_project_relative_path(self, relative_path)]) else: @@ -446,9 +449,7 @@ class Project(ToStringMixin): :param skip_ignored_files: whether to skip ignored files; has no effect if `code_files_only` is True :return: list of matches """ - file_collection = self._create_file_collection( - relative_path, code_files_only=code_files_only, skip_ignored_files=skip_ignored_files - ) + file_collection = self.create_file_collection(relative_path, code_files_only=code_files_only, skip_ignored_files=skip_ignored_files) return search_files( file_collection, pattern, @@ -623,6 +624,14 @@ class Project(ToStringMixin): return 0 def shutdown(self, timeout: float = 2.0) -> None: + """ + Shuts down the project, calling the language backend-specific shutdown of the active project. + + :param timeout: the timeout, in seconds + """ + # clean up internal resources if self.language_server_manager is not None: self.language_server_manager.stop_all(save_cache=True, timeout=timeout) self.language_server_manager = None + # trigger additional backend-specific shutdown + self.language_backend.shutdown_active_project(self, timeout=timeout) diff --git a/src/serena/project_server.py b/src/serena/project_server.py index 279f8523..abc2e3c0 100644 --- a/src/serena/project_server.py +++ b/src/serena/project_server.py @@ -2,16 +2,19 @@ import json import logging +import pickle +import secrets import threading -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, Any import requests as requests_lib -from flask import Flask, request +from flask import Flask, Response, abort, request from pydantic import BaseModel from sensai.util.logging import LogTime -from serena.config.serena_config import LanguageBackend, SerenaConfig +from serena.config.serena_config import SerenaConfig from serena.constants import SerenaPorts +from serena.language_backend import BuiltinLanguageBackend if TYPE_CHECKING: from serena.project import Project @@ -33,6 +36,19 @@ class QueryProjectRequest(BaseModel): tool_params_json: str +class CallFacadeMethodRequest(BaseModel): + """ + Request model for the /call_facade_method endpoint: the execution of a REPL facade method + in the context of a project. + """ + + project_name: str + facade_name: str + method_name: str + args: list[Any] + kwargs: dict[str, Any] + + class ProjectServer: """ A lightweight Flask server that exposes a SerenaAgent's project querying @@ -58,7 +74,7 @@ class ProjectServer: port = self.PORT serena_config = SerenaConfig.from_config_file().with_headless_mode_overrides() - serena_config.language_backend = LanguageBackend.LSP + serena_config.set_builtin_language_backend(BuiltinLanguageBackend.LSP) self._agent = SerenaAgent(serena_config=serena_config) self._loaded_projects_by_root: dict[str, "Project"] = {} @@ -76,7 +92,22 @@ class ProjectServer: self._setup_routes() + def get_serena_config(self) -> SerenaConfig: + return self._agent.serena_config + + def get_auth_secret(self) -> str: + """Returns the authentication secret used by the server.""" + return self._agent.serena_config.auth_secret + def _setup_routes(self) -> None: + @self._app.before_request + def authenticate() -> None: + # authenticate every request before parsing input or accessing projects + secret = self.get_auth_secret() + provided = request.headers.get("Authorization", "") + if not secret or not secrets.compare_digest(provided.encode("utf-8"), f"Bearer {secret}".encode()): + abort(401) + @self._app.route("/heartbeat", methods=["GET"]) def heartbeat() -> dict[str, str]: return {"status": "alive"} @@ -86,6 +117,18 @@ class ProjectServer: query_request = QueryProjectRequest.model_validate(request.get_json()) return self._query_project(query_request) + @self._app.route("/call_facade_method", methods=["POST"]) + def call_facade_method() -> Response: + call_request = CallFacadeMethodRequest.model_validate(request.get_json()) + try: + result = self._call_facade_method(call_request) + except Exception as e: + # report the error to the client (which raises it in the REPL) instead of a generic server error page + log.warning("Facade method call failed: %s", e) + return Response(f"{type(e).__name__}: {e}", status=400, mimetype="text/plain") + # NOTE: the result is pickled; the client (a Serena instance on the same machine) unpickles it + return Response(pickle.dumps(result), mimetype="application/octet-stream") + def _get_project(self, project_root_or_name: str) -> "Project": """Gets the project with the given name, loading it if necessary.""" serena_config = self._agent.serena_config @@ -136,6 +179,17 @@ class ProjectServer: params = json.loads(req.tool_params_json) return tool.apply_ex(**params) + def _call_facade_method(self, req: CallFacadeMethodRequest) -> Any: + """ + Handles a /call_facade_method request by executing the facade method on the agent's REPL facades in the + context of the specified project (see `_query_project` regarding the lock). + """ + project = self._get_project(req.project_name) + with self._active_project_lock, self._agent.active_project_context(project): + facade = self._agent.get_repl().entrypoint.get_facade_(req.facade_name) + method = facade.get_method(req.method_name) + return self._agent.execute_task(lambda: method(*req.args, **req.kwargs)) + def run(self) -> None: """ Run the server on the given host and port. @@ -158,18 +212,23 @@ class ProjectServerClient: :class:`ConnectionError` is raised. """ - def __init__(self, host: str = "127.0.0.1", port: int = ProjectServer.PORT, timeout: int = 300) -> None: + def __init__(self, serena_config: SerenaConfig, host: str = "127.0.0.1", port: int | None = None) -> None: """ :param host: the host address of the project server. - :param port: the port of the project server. + :param port: the port of the project server; if None, use default. + :param auth_secret: the shared authentication secret; defaults to the secret in Serena's configuration. :raises ConnectionError: if the project server is not reachable. """ + if port is None: + port = ProjectServer.PORT self._base_url = f"http://{host}:{port}" - self._timeout = timeout + self._timeout = serena_config.tool_timeout - 1 + auth_secret = serena_config.auth_secret + self._headers = {"Authorization": f"Bearer {auth_secret}"} # verify that the server is running try: - response = requests_lib.get(f"{self._base_url}/heartbeat", timeout=5) + response = requests_lib.get(f"{self._base_url}/heartbeat", headers=self._headers, timeout=5) response.raise_for_status() except requests_lib.ConnectionError: raise ConnectionError(f"ProjectServer is not reachable at {self._base_url}. Make sure the server is running.") @@ -194,6 +253,25 @@ class ProjectServerClient: tool_params_json=tool_params_json, ).model_dump() - response = requests_lib.post(f"{self._base_url}/query_project", json=payload, timeout=self._timeout) + response = requests_lib.post(f"{self._base_url}/query_project", json=payload, headers=self._headers, timeout=self._timeout) response.raise_for_status() return response.text + + def call_facade_method(self, project_name: str, facade_name: str, method_name: str, args: list[Any], kwargs: dict[str, Any]) -> Any: + """ + Executes a (read-only) REPL facade method in the context of a project. + + :param project_name: the name of the project to query + :param facade_name: the facade's name + :param method_name: the method's name + :param args: the positional arguments (JSON-serialisable) + :param kwargs: the keyword arguments (JSON-serialisable) + :return: the method's result, as returned by the server (unpickled; the server is a trusted local process) + """ + payload = CallFacadeMethodRequest( + project_name=project_name, facade_name=facade_name, method_name=method_name, args=args, kwargs=kwargs + ).model_dump() + response = requests_lib.post(f"{self._base_url}/call_facade_method", json=payload, headers=self._headers, timeout=self._timeout) + if not response.ok: + raise ValueError(f"Project server error ({response.status_code}): {response.text[:2000]}") + return pickle.loads(response.content) diff --git a/src/serena/repl/api/cfg_api.py b/src/serena/repl/api/cfg_api.py new file mode 100644 index 00000000..61053727 --- /dev/null +++ b/src/serena/repl/api/cfg_api.py @@ -0,0 +1,40 @@ +# SPDX-License-Identifier: GPL-3.0-or-later +""" +The implementation of operations concerning Serena's configuration and session state. +""" + +from typing import TYPE_CHECKING + +from serena.tools import GetCurrentConfigTool, OpenDashboardTool + +from ..facade import FacadeApi, facade_method + +if TYPE_CHECKING: + from serena.agent import SerenaAgent + + +class ConfigApi(FacadeApi): + def __init__(self, agent: "SerenaAgent") -> None: + super().__init__(agent, name="cfg", description="Serena's configuration and session state (incl. the dashboard)") + + @facade_method(corresponding_tool=GetCurrentConfigTool) + def get_current_config(self) -> str: + """ + Provides the current configuration of the agent, including the active and available projects, tools, contexts, and modes. + + :return: the configuration overview + """ + return self._agent.get_current_config_overview() + + @facade_method(corresponding_tool=OpenDashboardTool) + def open_dashboard(self) -> str: + """ + Opens the Serena web dashboard in the default web browser. + The dashboard provides logs, session information, and tool usage statistics. + + :return: a message indicating whether the dashboard could be opened + """ + if self._agent.open_dashboard(): + return f"Serena web dashboard has been opened in the user's default web browser: {self._agent.get_dashboard_url()}" + else: + return f"Serena web dashboard could not be opened automatically; tell the user to open it via {self._agent.get_dashboard_url()}" diff --git a/src/serena/repl/api/edit_api.py b/src/serena/repl/api/edit_api.py new file mode 100644 index 00000000..be09770e --- /dev/null +++ b/src/serena/repl/api/edit_api.py @@ -0,0 +1,302 @@ +# SPDX-License-Identifier: GPL-3.0-or-later +""" +The implementation of editing operations, which are independent of the language backend. +""" + +from typing import TYPE_CHECKING, Literal + +from serena.code_editor import EditedFileContext +from serena.tools import ( + DeleteLinesTool, + InsertAfterSymbolTool, + InsertAtLineTool, + InsertBeforeSymbolTool, + ReplaceContentTool, + ReplaceInFilesTool, + ReplaceLinesTool, + ReplaceSymbolBodyTool, +) +from serena.util.text_utils import ContentReplacer, MultiFileReplacement, ReplacementOccurrence, ReplacementRejectedError + +from ..facade import SUCCESS_RESULT, FacadeApi, ReferencedType, facade_method +from ..representable import Renderer, RepresentableViaRenderer + +if TYPE_CHECKING: + from serena.agent import SerenaAgent + + +class ReplacementPreview(RepresentableViaRenderer): + """ + The prospective changes of a multi-file replacement (nothing has been modified). + Each entry of `occurrences` (`ReplacementOccurrence`) has an `occurrence_id` (to be passed to `replace_in_files` + in order to apply exactly that occurrence), `relative_path`, `start_line`, `end_line`, `matched_text`, `replacement` + and `is_ambiguous`. + """ + + def __init__(self, replacement: MultiFileReplacement, renderer: "ReplacementPreviewRenderer"): + """ + :param replacement: the replacement + :param renderer: the renderer to use for representing the preview + """ + super().__init__(renderer) + self.replacement_ = replacement + + @property + def occurrences(self) -> list[ReplacementOccurrence]: + return self.replacement_.occurrences + + @property + def affected_files(self) -> list[str]: + return self.replacement_.affected_files + + +class ReplacementPreviewRenderer(Renderer[ReplacementPreview]): + """ + Renders the listing of prospective changes (minimal line diffs with occurrence ids), subject to the length limit. + """ + + def __init__(self, agent: "SerenaAgent", max_answer_chars: int, dry_run: bool): + """ + :param agent: the agent + :param max_answer_chars: the maximum number of characters; -1 for the configured default + :param dry_run: whether the listing is the result of a dry run (adding instructions on how to proceed) + """ + super().__init__(agent, max_answer_chars) + self._dry_run = dry_run + + def render(self, obj: ReplacementPreview) -> str: + return obj.replacement_.render_listing(self._get_max_answer_chars(), dry_run=self._dry_run) + + +class EditApi(FacadeApi): + def __init__(self, agent: "SerenaAgent") -> None: + super().__init__( + agent, + name="edit", + description="modifying content within existing files (independent of the language backend)", + types=[ + ReferencedType( + ReplacementOccurrence, + members=["occurrence_id", "relative_path", "start_line", "end_line", "matched_text", "replacement", "is_ambiguous"], + ), + ], + ) + + # file-level operations + + @facade_method(can_edit=True, corresponding_tool=ReplaceContentTool) + def replace_content( + self, + relative_path: str, + needle: str, + repl: str, + mode: Literal["literal", "regex"], + allow_multiple_occurrences: bool = False, + ) -> str: + r""" + Replaces one or more occurrences of a given pattern in a file with new content. + + VERY IMPORTANT: The "regex" mode allows very large sections of code to be replaced WITHOUT + quoting them fully: use a needle of the form "beginning.*?end-of-text-to-be-replaced" with + wildcards instead of pasting the exact original text — shorter, cheaper, and you cannot make + mistakes, because an ambiguous match returns an error you can refine, so wildcards are safe. + Prefer regex mode with suitable wildcards for long multi-line replacements; use the + symbol-level editors when replacing a whole method/class. + + :param relative_path: the relative path to the file + :param needle: the string or regex pattern to search for. + If `mode` is "literal", this string will be matched exactly. + If `mode` is "regex", this string will be treated as a regular expression (syntax of Python's `re` module, + with flags DOTALL and MULTILINE enabled). + :param repl: the replacement string (verbatim). + If mode is "regex", the string can contain backreferences to matched groups in the needle regex, + specified using the syntax $!1, $!2, etc. for groups 1, 2, etc. + :param mode: either "literal" or "regex", specifying how the `needle` parameter is to be interpreted. + :param allow_multiple_occurrences: whether to allow matching and replacing multiple occurrences. + If false and multiple occurrences are found, an error will be raised + :return: a success message + """ + self._get_project().validate_relative_path(relative_path) + with EditedFileContext(relative_path, self._create_code_editor()) as context: + replacer = ContentReplacer(mode=mode, allow_multiple_occurrences=allow_multiple_occurrences) + context.set_updated_content(replacer.replace(context.get_original_content(), needle, repl)) + return SUCCESS_RESULT + + @facade_method(can_edit=True, corresponding_tool=ReplaceInFilesTool) + def replace_in_files( + self, + needle: str, + repl: str, + mode: Literal["literal", "regex"], + relative_path: str = "", + paths_include_glob: str = "", + paths_exclude_glob: str = "", + dry_run: bool = False, + occurrence_ids: list[str] | None = None, + expected_count: int = -1, + max_answer_chars: int = -1, + ) -> ReplacementPreview | str: + r""" + Replaces occurrences of a pattern across multiple files in ONE call. + + This is the preferred operation for repeated small edits (renames, import swaps, annotation changes, + path prefixes) spanning several files or many places in one file: one call with a SHORT pattern + replaces many single-file replacements with long disambiguating needles. + + Recommended protocol whenever there is ANY risk of unintended replacements: + 1. Call with dry_run=True: every prospective change is returned as a minimal line diff with an + occurrence id; nothing is modified. + 2. Call again with dry_run=False, passing the ids you want in occurrence_ids (omit it to apply + all). You pick the desired replacements from the list - no counting, no needle-crafting. + + For clearly unambiguous bulk replacements you may skip the dry run; pass expected_count as a + guard. If the actual number of matches differs, NOTHING is changed and an error containing the + prospective changes is raised, so a failed guard costs one call and gives you the dry-run output to select from. + + :param needle: the string (mode "literal") or regular expression (mode "regex"; Python `re` + syntax with DOTALL and MULTILINE) to search for + :param repl: the replacement string. In regex mode, backreferences to matched groups can be + specified as $!1, $!2, etc. + :param mode: either "literal" or "regex", specifying how `needle` is to be interpreted + :param relative_path: only consider this file or directory (default: the whole project) + :param paths_include_glob: optional glob (relative to the project root, e.g. "src/**/*.java") + restricting which files are considered + :param paths_exclude_glob: optional glob of files to exclude; takes precedence over the include glob + :param dry_run: if True, do not modify anything; return the prospective changes with occurrence ids + :param occurrence_ids: optional list of occurrence ids (obtained from a dry run) to which the + replacement is restricted; if any id is unknown or stale, NOTHING is changed. If omitted, + all occurrences are replaced. + :param expected_count: optional guard for calls without occurrence_ids: the number of + occurrences you expect to be replaced. If the actual count differs, nothing is changed and + an error containing the prospective changes is raised. -1 disables the guard. + :return: in a dry run, the prospective changes (`ReplacementPreview`); otherwise a summary of the applied replacements + """ + replacement = MultiFileReplacement( + self._get_project(), + needle, + repl, + mode, + relative_path=relative_path, + paths_include_glob=paths_include_glob, + paths_exclude_glob=paths_exclude_glob, + ) + if dry_run: + return ReplacementPreview(replacement, ReplacementPreviewRenderer(self._agent, max_answer_chars, dry_run=True)) + + # select the occurrences to replace + try: + if occurrence_ids is not None: + occurrences = replacement.select(occurrence_ids) + else: + occurrences = replacement.select_all_guarded(expected_count) + except ReplacementRejectedError as e: + message = str(e) + if e.show_prospective_changes: + preview = ReplacementPreview(replacement, ReplacementPreviewRenderer(self._agent, max_answer_chars, dry_run=False)) + message += "\n" + preview.represent() + raise ValueError(message) from e + + return replacement.apply(self._create_code_editor(), occurrences).to_display_string() + + # line-level operations + + @facade_method(optional=True, can_edit=True, corresponding_tool=DeleteLinesTool) + def delete_lines(self, relative_path: str, start_line: int, end_line: int) -> str: + """ + Deletes the given lines in the file. + Requires that the same range of lines was previously read to verify correctness of the operation. + + :param relative_path: the relative path to the file + :param start_line: the 0-based index of the first line to be deleted + :param end_line: the 0-based index of the last line to be deleted + :return: a success message + """ + self._create_code_editor().delete_lines(relative_path, start_line, end_line) + return SUCCESS_RESULT + + @facade_method(optional=True, can_edit=True, corresponding_tool=ReplaceLinesTool) + def replace_lines(self, relative_path: str, start_line: int, end_line: int, content: str) -> str: + """ + Replaces the given range of lines in the given file. + Requires that the same range of lines was previously read to verify correctness of the operation. + + :param relative_path: the relative path to the file + :param start_line: the 0-based index of the first line to be replaced + :param end_line: the 0-based index of the last line to be replaced + :param content: the content to insert + :return: a success message + """ + code_editor = self._create_code_editor() + code_editor.delete_lines(relative_path, start_line, end_line) + code_editor.insert_at_line(relative_path, start_line, self._normalize_inserted_content(content)) + return SUCCESS_RESULT + + @facade_method(optional=True, can_edit=True, corresponding_tool=InsertAtLineTool) + def insert_at_line(self, relative_path: str, line: int, content: str) -> str: + """ + Inserts the given content at the given line in the file, pushing existing content of the line down. + In general, symbolic insert operations like insert_after_symbol or insert_before_symbol should be preferred if you know which + symbol you are looking for. + However, this can also be useful for small targeted edits of the body of a longer symbol (without replacing the entire body). + + :param relative_path: the relative path to the file + :param line: the 0-based index of the line to insert content at + :param content: the content to be inserted + :return: a success message + """ + self._create_code_editor().insert_at_line(relative_path, line, self._normalize_inserted_content(content)) + return SUCCESS_RESULT + + @staticmethod + def _normalize_inserted_content(content: str) -> str: + return content if content.endswith("\n") else content + "\n" + + # symbol-level operations + + @facade_method(can_edit=True, corresponding_tool=ReplaceSymbolBodyTool) + def replace_symbol_body(self, name_path: str, relative_path: str, body: str) -> str: + """ + Replaces the body of the given symbol. + + IMPORTANT: Only replace symbol bodies if you have previously made a retrieval with include_body=True and thus know what + constitutes the body! + + :param name_path: name path of the symbol whose body to replace + :param relative_path: the relative path to the file containing the symbol + :param body: the new symbol body. The symbol body is the definition of a symbol + in the programming language, including e.g. the signature line for functions. + Depending on the language, it may or may not include a preceding docstring or other preceding annotations. + :return: a success message + """ + self._create_code_editor().replace_body(name_path, relative_file_path=relative_path, body=body) + return SUCCESS_RESULT + + @facade_method(can_edit=True, corresponding_tool=InsertAfterSymbolTool) + def insert_after_symbol(self, name_path: str, relative_path: str, body: str) -> str: + """ + Inserts code after a class/method/function definition. + Don't use this to insert after assignments (constants, fields). + + :param name_path: name path of the symbol after which to insert content + :param relative_path: the relative path to the file containing the symbol + :param body: the body/content to be inserted. The inserted code shall begin with the next line after + the symbol. + :return: a success message + """ + self._create_code_editor().insert_after_symbol(name_path, relative_file_path=relative_path, body=body) + return SUCCESS_RESULT + + @facade_method(can_edit=True, corresponding_tool=InsertBeforeSymbolTool) + def insert_before_symbol(self, name_path: str, relative_path: str, body: str) -> str: + """ + Inserts the given content before the beginning of the definition of the given symbol (via the symbol's location). + A typical use case is to insert a new class, function, method, field or variable assignment; or + a new import statement before the first symbol in the file. + + :param name_path: name path of the symbol before which to insert content + :param relative_path: the relative path to the file containing the symbol + :param body: the body/content to be inserted before the line in which the referenced symbol is defined + :return: a success message + """ + self._create_code_editor().insert_before_symbol(name_path, relative_file_path=relative_path, body=body) + return SUCCESS_RESULT diff --git a/src/serena/repl/api/ext_api.py b/src/serena/repl/api/ext_api.py new file mode 100644 index 00000000..eed18827 --- /dev/null +++ b/src/serena/repl/api/ext_api.py @@ -0,0 +1,115 @@ +# SPDX-License-Identifier: GPL-3.0-or-later +""" +The implementation of access to external projects (projects other than the active one). +""" + +from types import TracebackType +from typing import TYPE_CHECKING + +from serena.jetbrains.jetbrains_plugin_client import JetBrainsPluginClientManager +from serena.tools import ListQueryableProjectsTool, QueryProjectTool + +from ..external_project import ExternalProjectExecution +from ..facade import FacadeApi, facade_method +from ..representable import JsonObject, JsonObjectRenderer + +if TYPE_CHECKING: + from serena.agent import SerenaAgent + + +class ExternalProjectContextManager: + """ + A context manager (for use in a `with` statement) within which the facades operate on an external project + (read-only): the external project is temporarily activated, and operations requiring language servers are + executed in the project's server. Contexts cannot be nested. + """ + + def __init__(self, agent: "SerenaAgent", project_name: str, read_only: bool) -> None: + """ + :param agent: the agent + :param project_name: the name (or root path) of the registered external project + :param read_only: whether the context is read-only + """ + self._agent = agent + self._project_name = project_name + self._active_project_context = None + self._read_only = read_only + + def __enter__(self) -> None: + entrypoint = self._agent.get_repl().entrypoint + if entrypoint.get_external_project_() is not None: + raise ValueError("External project contexts cannot be nested") + + # temporarily activate the external project + registered_project = self._agent.serena_config.get_registered_project(self._project_name) + if registered_project is None: + raise ValueError(f"Project '{self._project_name}' is not registered and cannot be queried") + project = registered_project.get_project_instance(self._agent.serena_config) + self._active_project_context = self._agent.active_project_context(project) + self._active_project_context.__enter__() + + # switch the facades to the external project + entrypoint.set_external_project_(ExternalProjectExecution(registered_project.project_name, self._read_only, self._agent)) + + def __exit__(self, exc_type: type[BaseException] | None, exc_value: BaseException | None, traceback: TracebackType | None) -> None: + self._agent.get_repl().entrypoint.set_external_project_(None) + assert self._active_project_context is not None + self._active_project_context.__exit__(exc_type, exc_value, traceback) + + +class ExternalProjectsApi(FacadeApi): + def __init__(self, agent: "SerenaAgent") -> None: + super().__init__(agent, name="ext", description="read-only access to external projects (projects other than the active one)") + + @facade_method(corresponding_tool=ListQueryableProjectsTool) + def list_projects(self, symbol_access: bool = True) -> JsonObject: + """ + Lists the registered projects which can be queried. + + :param symbol_access: whether to list only projects for which symbol-level access is available + :return: the project names mapped to their root directories + """ + registered_projects = self._agent.serena_config.projects + if symbol_access and self._agent.get_language_backend().is_jetbrains(): + # only projects with open IDE instances can be queried + matched_clients = JetBrainsPluginClientManager().match_clients(registered_projects) + relevant_projects = [mc.registered_project for mc in matched_clients] + else: + # all projects can be queried (the project server instantiates projects dynamically) + relevant_projects = registered_projects + result = {p.project_name: str(p.project_root) for p in relevant_projects} + return JsonObject(result, JsonObjectRenderer(self._agent, -1)) + + @facade_method(corresponding_tool=QueryProjectTool) + def read_project_context(self, project_name: str) -> ExternalProjectContextManager: + """ + Provides a context (for use in a `with` statement) within which all facades operate on the given external project + instead of the active one, with read-only access. + + Example: + `with s.ext.project_context("other"): result = s.lsp.find_symbol("Foo")` + + Results obtained within the context can be used after it (they are self-contained). + + :param project_name: the name (or root path) of the project, as listed by `list_projects` + :return: the context manager + + """ + return ExternalProjectContextManager(self._agent, project_name, read_only=True) + + @facade_method(optional=True) + def project_context(self, project_name: str) -> ExternalProjectContextManager: + """ + Provides a context (for use in a `with` statement) within which all facades operate on the given external project + instead of the active one (read and write operations are possible). + + Example: + `with s.ext.project_context("other"): result = s.lsp.find_symbol("Foo")` + + Results obtained within the context can be used after it (they are self-contained). + + :param project_name: the name (or root path) of the project, as listed by `list_projects` + :return: the context manager + + """ + return ExternalProjectContextManager(self._agent, project_name, read_only=False) diff --git a/src/serena/repl/api/fs_api.py b/src/serena/repl/api/fs_api.py new file mode 100644 index 00000000..b457035e --- /dev/null +++ b/src/serena/repl/api/fs_api.py @@ -0,0 +1,330 @@ +# SPDX-License-Identifier: GPL-3.0-or-later +""" +The implementation of operations on the project's files. +""" + +import os +from collections import defaultdict +from fnmatch import fnmatch +from pathlib import Path +from typing import TYPE_CHECKING + +from serena.tools import CreateTextFileTool, FindFileTool, ListDirTool, ReadFileTool, SearchForPatternTool +from serena.util.file_system import scan_directory +from serena.util.text_utils import MatchedConsecutiveLines +from solidlsp.ls_utils import TextUtils + +from ..facade import FacadeApi, ReferencedType, facade_method +from ..representable import Renderer, RepresentableViaRenderer + +if TYPE_CHECKING: + from serena.agent import SerenaAgent + + +class FileContent(RepresentableViaRenderer): + """ + The content of a file (or of a range of its lines): `text` (the joined lines) and `lines`. + """ + + def __init__(self, lines: list[str], renderer: "FileContentRenderer"): + """ + :param lines: the lines (without line breaks) + :param renderer: the renderer to use for representing the content + """ + super().__init__(renderer) + self.lines = lines + + lines: list[str] + + @property + def text(self) -> str: + return "\n".join(self.lines) + + +class FileContentRenderer(Renderer[FileContent]): + def render(self, obj: FileContent) -> str: + return self._limit_length(obj.text) + + +class DirectoryListing(RepresentableViaRenderer): + """ + The entries of a directory: `dirs` and `files` (relative paths). + """ + + def __init__(self, dirs: list[str], files: list[str], renderer: "DirectoryListingRenderer"): + """ + :param dirs: the relative paths of the directories + :param files: the relative paths of the files + :param renderer: the renderer to use for representing the listing + """ + super().__init__(renderer) + self.dirs = dirs + self.files = files + + dirs: list[str] + files: list[str] + + +class DirectoryListingRenderer(Renderer[DirectoryListing]): + def render(self, obj: DirectoryListing) -> str: + return self._limit_length(self._to_json({"dirs": obj.dirs, "files": obj.files})) + + +class PatternMatches(RepresentableViaRenderer): + """ + The matches of a pattern search (`MatchedConsecutiveLines`). + """ + + def __init__(self, matches: list[MatchedConsecutiveLines], renderer: "PatternMatchesRenderer"): + """ + :param matches: the matches + :param renderer: the renderer to use for representing the matches + """ + super().__init__(renderer) + self.matches = matches + + matches: list[MatchedConsecutiveLines] + + def __len__(self) -> int: + return len(self.matches) + + def matches_by_file_(self) -> dict[str, list[MatchedConsecutiveLines]]: + result: defaultdict[str, list[MatchedConsecutiveLines]] = defaultdict(list) + for match in self.matches: + assert match.source_file_path is not None + result[match.source_file_path].append(match) + return result + + +class PatternMatchesRenderer(Renderer[PatternMatches]): + """ + Renders matches as a mapping from file paths to matched line blocks (with context), falling back to progressively + shorter representations (first lines, truncated first lines, line numbers, per-file counts, a summary) if the + length limit is exceeded. + """ + + _TEXT_TRUNCATE = 60 + + def render(self, obj: PatternMatches) -> str: + matches_by_file = obj.matches_by_file_() + file_to_matches = {path: [m.to_display_string() for m in matches] for path, matches in matches_by_file.items()} + + # capture lightweight match data for shortening before serialization + match_lines_by_file = { + path: [{"line": m.matched_lines[0].line_number, "text": m.matched_lines[0].line_content.strip()} for m in matches] + for path, matches in matches_by_file.items() + } + + # shortened result closures, from least to most aggressive shortening + def render_first_lines(truncate: bool) -> str: + """Render each match's first line, either in full or truncated to a fixed length.""" + + def entry_text(text: str) -> str: + if truncate and len(text) > self._TEXT_TRUNCATE: + return text[: self._TEXT_TRUNCATE] + "..." + return text + + compact = { + path: [{"line": m["line"], "text": entry_text(str(m["text"]))} for m in lines] + for path, lines in match_lines_by_file.items() + } + if truncate: + header = ( + f"Matched lines (text over {self._TEXT_TRUNCATE} chars is truncated, marked with a trailing '...'); " + "use read_file with the line numbers for full content:" + ) + else: + header = "Matched lines per file; use read_file with the line numbers for surrounding context:" + return f"{header}\n{self._to_json(compact)}" + + def make_first_lines_full() -> str: + return render_first_lines(truncate=False) + + def make_first_lines_truncated() -> str: + return render_first_lines(truncate=True) + + def make_line_numbers_only() -> str: + numbers = {path: [m["line"] for m in lines] for path, lines in match_lines_by_file.items()} + return f"Match lines per file:\n{self._to_json(numbers)}" + + def make_per_file_counts() -> str: + counts = {path: len(lines) for path, lines in match_lines_by_file.items()} + return f"Match counts per file:\n{self._to_json(counts)}" + + def make_summary() -> str: + return f"Found {len(obj)} matches in {len(match_lines_by_file)} files." + + return self._limit_length( + self._to_json(file_to_matches), + shortened_result_factories=[ + make_first_lines_full, + make_first_lines_truncated, + make_line_numbers_only, + make_per_file_counts, + make_summary, + ], + ) + + +class FsApi(FacadeApi): + def __init__(self, agent: "SerenaAgent") -> None: + super().__init__( + agent, + name="fs", + description="the project's files as units (as opposed to their content, see `edit`)", + types=[ + ReferencedType( + MatchedConsecutiveLines, members=["source_file_path", "matched_lines", "start_line", "end_line", "to_display_string"] + ), + ], + ) + + @facade_method(corresponding_tool=ReadFileTool) + def read_file(self, relative_path: str, start_line: int = 0, end_line: int | None = None, max_answer_chars: int = -1) -> FileContent: + """ + Reads the given file or a range of its lines. + + :param relative_path: the relative path to the file to read + :param start_line: the 0-based index of the first line to be retrieved, negative values count from the end of the file. + :param end_line: the 0-based index of the last line to be retrieved (inclusive). If None, read until the end of the file. + :return: the content + """ + project = self._get_project() + project.validate_relative_path(relative_path) + + # read lines, using the same (LSP-compliant) notion of line breaks as the line-based editing operations + lines = TextUtils.split_lines(project.read_file(relative_path)) + lines = lines[start_line:] if end_line is None else lines[start_line : end_line + 1] + return FileContent(lines, FileContentRenderer(self._agent, max_answer_chars)) + + @facade_method(can_edit=True, corresponding_tool=CreateTextFileTool) + def create_text_file(self, relative_path: str, content: str) -> str: + """ + Writes a new file or overwrites an existing file with the given content. + + :param relative_path: the relative path to the file to create + :param content: the (appropriately encoded) content to write to the file + :return: a message indicating success + """ + project = self._get_project() + project_root = Path(project.project_root) + abs_path = (project_root / relative_path).resolve() + will_overwrite_existing = abs_path.exists() + + # validate the destination path + if will_overwrite_existing: + project.validate_relative_path(relative_path) + else: + assert abs_path.is_relative_to(project_root), f"Cannot create file outside of the project directory, got {relative_path=}" + + # write the file + abs_path.parent.mkdir(parents=True, exist_ok=True) + abs_path.write_text(content, encoding=project.project_config.encoding, newline=project.line_ending.newline_str) + answer = f"File created: {relative_path}." + if will_overwrite_existing: + answer += " Overwrote existing file." + return answer + + @facade_method(corresponding_tool=ListDirTool) + def list_dir( + self, relative_path: str, recursive: bool, skip_ignored_files: bool = False, max_answer_chars: int = -1 + ) -> DirectoryListing: + """ + Lists files and directories in the given directory (optionally with recursion). + + :param relative_path: the relative path to the directory to list; pass "." to scan the project root + :param recursive: whether to scan subdirectories recursively + :param skip_ignored_files: whether to skip files and directories that are ignored + :return: the listing + """ + project = self._get_project() + if not project.relative_path_exists(relative_path): + raise FileNotFoundError(f"Directory not found: {relative_path} (check if the path is correct relative to the project root)") + project.validate_relative_path(relative_path) + + is_ignored_path_fn = project.get_is_ignored_path_fn(relative_path, skip_ignored_files) + dirs, files = scan_directory( + os.path.join(project.project_root, relative_path), + relative_to=project.project_root, + recursive=recursive, + is_ignored_dir=is_ignored_path_fn, + is_ignored_file=is_ignored_path_fn, + ) + return DirectoryListing(dirs, files, DirectoryListingRenderer(self._agent, max_answer_chars)) + + @facade_method(corresponding_tool=FindFileTool) + def find_file(self, file_mask: str, relative_path: str) -> list[str]: + """ + Finds files matching the given file mask within the given relative path. + + :param file_mask: the filename or file mask (using the wildcards * or ?) to search for + :param relative_path: the relative path to the directory to search in; pass "." to scan the project root + :return: the relative paths of the matching files + """ + project = self._get_project() + project.validate_relative_path(relative_path) + + is_ignored_path_fn = project.get_is_ignored_path_fn(relative_path, skip_ignored_paths=False) + + # find the files by ignoring everything that doesn't match + def is_ignored_file(abs_path: str) -> bool: + if is_ignored_path_fn(abs_path): + return True + return not fnmatch(os.path.basename(abs_path), file_mask) + + _dirs, files = scan_directory( + path=os.path.join(project.project_root, relative_path), + recursive=True, + is_ignored_dir=is_ignored_path_fn, + is_ignored_file=is_ignored_file, + relative_to=project.project_root, + ) + return files + + @facade_method(corresponding_tool=SearchForPatternTool) + def search_for_pattern( + self, + substring_pattern: str, + context_lines_before: int = 0, + context_lines_after: int = 0, + paths_include_glob: str = "", + paths_exclude_glob: str = "", + relative_path: str = "", + restrict_search_to_code_files: bool = False, + skip_ignored_files: bool = True, + multiline: bool = True, + max_answer_chars: int = -1, + ) -> PatternMatches: + """ + Searches for a regex pattern across project files, returning whole matched lines (plus optional context). + Prefer symbolic operations if you know which symbols you are looking for! + + :param substring_pattern: regular expression to search for. + :param context_lines_before: number of context lines to include before each match. + :param context_lines_after: number of context lines to include after each match. + :param paths_include_glob: optional glob (relative to project root, e.g. ``"src/**/*.ts"``) restricting which files are searched. + :param paths_exclude_glob: optional glob to exclude files; takes precedence over `paths_include_glob`. + :param relative_path: restricts the search to this file or subdirectory of the project root + :param restrict_search_to_code_files: whether to search only (non-ignored) files containing analyzable code symbols + (useful when looking for class/method definitions); otherwise also search non-code files. + :param skip_ignored_files: whether to skip ignored sub-paths (default: True) + :param multiline: whether to apply multi-line matching (default: True), enabling the flags re.DOTALL and re.MULTILINE + :return: the matches, rendered as a mapping from file paths to matched consecutive lines (0-based line numbers) + """ + project = self._get_project() + relative_path = relative_path.strip() + if relative_path: + project.validate_relative_path(relative_path) + + matches = project.search_project_files_for_pattern( + pattern=substring_pattern, + relative_path=relative_path, + context_lines_before=context_lines_before, + context_lines_after=context_lines_after, + paths_include_glob=paths_include_glob.strip(), + paths_exclude_glob=paths_exclude_glob.strip(), + multiline=multiline, + code_files_only=restrict_search_to_code_files, + skip_ignored_files=skip_ignored_files, + ) + return PatternMatches(matches, PatternMatchesRenderer(self._agent, max_answer_chars)) diff --git a/src/serena/repl/api/jb_api.py b/src/serena/repl/api/jb_api.py new file mode 100644 index 00000000..5cf8fc27 --- /dev/null +++ b/src/serena/repl/api/jb_api.py @@ -0,0 +1,642 @@ +# SPDX-License-Identifier: GPL-3.0-or-later +""" +The implementation of JetBrains IDE-backed operations. +""" + +from collections import Counter, defaultdict +from collections.abc import Iterator +from contextlib import contextmanager +from typing import TYPE_CHECKING, Any, Literal + +import serena.jetbrains.jetbrains_types as jb +from serena.code_editor import JetBrainsCodeEditor +from serena.jetbrains.jetbrains_plugin_client import JetBrainsPluginClient +from serena.jetbrains.jetbrains_types import SymbolDTO, SymbolDTOUtil +from serena.symbol import JetBrainsSymbolDictGrouper +from serena.tools import ( + JetBrainsDebugTool, + JetBrainsFindDeclarationTool, + JetBrainsFindImplementationsTool, + JetBrainsFindReferencingSymbolsTool, + JetBrainsFindSymbolTool, + JetBrainsGetSymbolsOverviewTool, + JetBrainsInlineSymbol, + JetBrainsListInspectionsTool, + JetBrainsMoveTool, + JetBrainsRenameTool, + JetBrainsRunInspectionsTool, + JetBrainsSafeDeleteTool, + JetBrainsTypeHierarchyTool, +) +from serena.util.text_utils import find_text_coordinates + +from ..facade import FacadeApi, facade_method +from ..representable import JsonObject, JsonObjectRenderer, Renderer, RepresentableViaRenderer + +if TYPE_CHECKING: + from serena.agent import SerenaAgent + + +class JetBrainsSymbolCollection(RepresentableViaRenderer): + """ + A collection of symbols retrieved via the JetBrains backend. + Each symbol is a dict with keys such as `name_path`, `relative_path` and `type`, and optionally + `children`, `body`, `quick_info`, `documentation` and (for references) `context`. + """ + + symbols: list[SymbolDTO] + + def __init__(self, symbols: list[SymbolDTO], renderer: "JetBrainsSymbolCollectionRenderer"): + """ + :param symbols: the symbols + :param renderer: the renderer to use for representing the collection + """ + super().__init__(renderer) + self.symbols = symbols + + def __len__(self) -> int: + return len(self.symbols) + + def relative_paths_(self) -> list[str]: + return [s.get("relative_path", "unknown") for s in self.symbols] + + def identifiers_(self) -> list[SymbolDTO]: + """ + :return: dicts containing only the identifying information (name_path, type, relative_path) of the symbols + """ + return [{"name_path": s["name_path"], "type": s["type"], "relative_path": s["relative_path"]} for s in self.symbols] + + +class JetBrainsSymbolCollectionRenderer(Renderer[JetBrainsSymbolCollection]): + """ + Renders a symbol collection as (optionally grouped) JSON, falling back to a listing of symbol identifiers + if the length limit is exceeded. + """ + + def __init__(self, agent: "SerenaAgent", max_answer_chars: int, grouper: JetBrainsSymbolDictGrouper | None = None): + super().__init__(agent, max_answer_chars) + self._grouper = grouper + + def _group(self, symbols: list[SymbolDTO]) -> Any: + return self._grouper.group(symbols) if self._grouper is not None else symbols + + def render_identifiers(self, obj: JetBrainsSymbolCollection) -> str: + """ + :return: a shortened representation containing symbol types and identifiers (path + name_path) only, without children + """ + return f"Names with paths:\n{self._to_json(self._group(obj.identifiers_()))}" + + def render(self, obj: JetBrainsSymbolCollection) -> str: + result = self._to_json(self._group(obj.symbols)) + return self._limit_length(result, shortened_result_factories=[lambda: self.render_identifiers(obj)]) + + +class JetBrainsReferencesRenderer(JetBrainsSymbolCollectionRenderer): + """ + Renders a collection of referencing symbols, falling back to per-file counts and finally the total count + if the length limit is exceeded. + """ + + def render(self, obj: JetBrainsSymbolCollection) -> str: + ref_paths = obj.relative_paths_() + result = self._to_json(self._group(obj.symbols)) + return self._limit_length( + result, + shortened_result_factories=[ + lambda: f"Reference counts per file:\n{self._to_json(Counter(ref_paths))}", + lambda: f"Found {len(ref_paths)} references.", + ], + ) + + +class JetBrainsSymbolsOverview(RepresentableViaRenderer): + """ + The overview of the symbols defined in a file, i.e. the top-level symbols (each a dict with keys such as + `name_path` and `type`, optionally with `children`) and, if requested, the file's documentation. + """ + + def __init__(self, symbols: list[SymbolDTO], documentation: str | None, renderer: "JetBrainsSymbolsOverviewRenderer"): + """ + :param symbols: the top-level symbols + :param documentation: the file's documentation, if requested and present + :param renderer: the renderer to use for representing the overview + """ + super().__init__(renderer) + self.symbols = symbols + self.documentation = documentation + + symbols: list[SymbolDTO] + documentation: str | None + + +class JetBrainsSymbolsOverviewRenderer(Renderer[JetBrainsSymbolsOverview]): + """ + Renders an overview in the compact grouped format, dropping (in order) the documentation, the children + and finally everything but symbol counts by type if the length limit is exceeded. + """ + + def __init__(self, agent: "SerenaAgent", max_answer_chars: int, grouper: JetBrainsSymbolDictGrouper, depth: int): + super().__init__(agent, max_answer_chars) + self._grouper = grouper + self._depth = depth + + def render(self, obj: JetBrainsSymbolsOverview) -> str: + grouped_symbols = self._grouper.group(obj.symbols) + shortened_result_factories = [] + + # create the full result + result: dict[str, Any] = {"symbols": grouped_symbols} + if obj.documentation: + result["docstring"] = obj.documentation + shortened_result_factories.append(lambda: self._to_json(grouped_symbols)) # shortened result without docstring + json_result = self._to_json(result) + + # create shortened results + if self._depth > 0: + + def create_short_result_depth_0() -> str: + depth_0_symbols = [d.copy() for d in obj.symbols] + for d in depth_0_symbols: + d.pop("children", None) + return "Depth 0 overview:\n" + self._to_json(self._grouper.group(depth_0_symbols)) + + shortened_result_factories.append(create_short_result_depth_0) + + def create_short_result_type_counts() -> str: + type_names = [d.get("type", "unknown") for d in obj.symbols] + return f"Symbol counts by type:\n{self._to_json(Counter(type_names))}" + + shortened_result_factories.append(create_short_result_type_counts) + + return self._limit_length(json_result, shortened_result_factories=shortened_result_factories) + + +class JetBrainsApi(FacadeApi): + # groupers for the various symbol collections; top-level symbols are grouped by the first key list, + # children by the second + find_symbol_grouper_ = JetBrainsSymbolDictGrouper( + ["relative_path", "type"], ["type"], collapse_singleton=True, map_name_path_to_name=True + ) + references_grouper_ = JetBrainsSymbolDictGrouper(["relative_path", "type"], ["type"], collapse_singleton=True) + overview_grouper_ = JetBrainsSymbolDictGrouper(["type"], ["type"], collapse_singleton=True, map_name_path_to_name=True) + + def __init__(self, agent: "SerenaAgent") -> None: + super().__init__( + agent, + name="jb", + description="operations on the codebase backed by the JetBrains IDE's code intelligence", + ) + + @contextmanager + def _client(self) -> Iterator[JetBrainsPluginClient]: + with JetBrainsPluginClient.from_project(self._get_project()) as client: + yield client + + def _json_object(self, data: Any, max_answer_chars: int = -1) -> JsonObject: + return JsonObject(data, JsonObjectRenderer(self._agent, max_answer_chars)) + + # read operations + + @facade_method(corresponding_tool=JetBrainsFindSymbolTool) + def find_symbol( + self, + name_path_pattern: str, + depth: int = 0, + relative_path: str | None = None, + include_body: bool = False, + include_info: bool = False, + search_deps: bool = False, + max_matches: int = -1, + max_answer_chars: int = -1, + ) -> JetBrainsSymbolCollection: + """ + Finds symbols and code entities (classes, methods, etc.) based on the given name path pattern. + + The returned symbol information can be used for edits or further queries. + Specify `depth > 0` to retrieve children (e.g., methods of a class). + Important: through `search_deps=True` dependencies can be searched, which + should be preferred to web search or other less sophisticated approaches to analyzing dependencies. + You will always receive at least quick info for returned symbols (even if `include_info=False`). + + A name path is a path in the symbol tree *within a source file*. + For example, the method `my_method` defined in class `MyClass` would have the name path `MyClass/my_method`. + If a symbol is overloaded (e.g., in Java), a 0-based index is appended (e.g. "MyClass/my_method[0]") to + uniquely identify it. + + To search for a symbol, you provide a name path pattern that is used to match against name paths. + It can be + * a simple name (e.g. "method"), which will match any symbol with that name + * a relative path like "class/method", which will match any symbol with that name path suffix + * an absolute name path "/class/method" (absolute name path), which requires an exact match of the full name path within the source file. + Append an index `[i]` to match a specific overload only, e.g. "MyClass/my_method[1]". + In any path component, using `*` will match any sequence of characters (excluding /), e.g. "Class/*substring*" matches a member substring. + A pattern must not contain only wildcards (e.g. "*" or "/*"); use `get_symbols_overview` to list a file's symbols. + + :param name_path_pattern: the name path matching pattern (see above) + :param depth: depth up to which descendants shall be retrieved (e.g. use 1 to also retrieve immediate children; + for the case where the symbol is a class, this will return its methods). + Ignored if `include_body=True`. Default 0. + :param relative_path: Optional. Restrict search to this file or directory. If not specified, searches entire codebase. + Note: for external dependencies, this must be an identifier starting with `{max_matches=} symbols.\n" + renderer.render_identifiers(collection)) + return collection + + @facade_method(corresponding_tool=JetBrainsFindReferencingSymbolsTool) + def find_referencing_symbols(self, name_path: str, relative_path: str, max_answer_chars: int = -1) -> JetBrainsSymbolCollection: + """ + Finds all symbols that reference the given symbol (its callers / usages / dependents) + i.e. the symbols whose own definition (e.g. a method body) contains a reference to it. + For each, returns its name path, file, and the surrounding line of code. + + :param name_path: name path of the symbol for which to find references + :param relative_path: the relative path to the file containing the symbol (must be a file, not a directory) + Note: for external dependencies, this must be an identifier starting with `= 0: + content_around_ref = project.retrieve_content_around_line( + relative_file_path=symbol_dict["relative_path"], line=ref_line, context_lines_before=1, context_lines_after=1 + ) + symbol_dict["context"] = content_around_ref.to_display_string() + del symbol_dict["reference_line_no"] + + renderer = JetBrainsReferencesRenderer(self._agent, max_answer_chars, grouper=self.references_grouper_) + return JetBrainsSymbolCollection(symbol_dicts, renderer) + + @facade_method(corresponding_tool=JetBrainsGetSymbolsOverviewTool) + def get_symbols_overview( + self, relative_path: str, depth: int = -1, max_answer_chars: int = -1, include_file_documentation: bool = False + ) -> JetBrainsSymbolsOverview: + """ + Gets an overview of the top-level symbols defined in the given file (classes, methods, fields). + + Returns STRUCTURE only, without bodies. This is the cheap, structure-first way to learn what a file + contains: it costs far less context than reading the whole file. + + :param relative_path: the relative path to the file to get the overview of + :param depth: depth up to which descendants shall be retrieved. + Default (-1) results in a language specific choice: 1 for java and kotlin and 0 for other languages + :param include_file_documentation: whether to include the file's docstring. Default False. + :return: the overview + """ + if depth == -1: + depth = 1 if relative_path.endswith((".java", ".kt")) else 0 + + with self._client() as client: + response = client.get_symbols_overview( + relative_path=relative_path, depth=depth, include_file_documentation=include_file_documentation + ) + renderer = JetBrainsSymbolsOverviewRenderer(self._agent, max_answer_chars, grouper=self.overview_grouper_, depth=depth) + return JetBrainsSymbolsOverview(response["symbols"], response.get("documentation"), renderer) + + @staticmethod + def _transform_hierarchy_nodes(nodes: list[jb.TypeHierarchyNodeDTO] | None) -> dict[str, list]: + """ + Transforms a list of hierarchy nodes into a file-grouped compact format. + + :return: a dict where keys are relative paths and values are lists of either a name path (for a leaf node) + or a dict mapping the name path to the (recursively transformed) children + """ + result: defaultdict[str, list] = defaultdict(list) + for node in nodes or []: + symbol = node["symbol"] + name_path = symbol["name_path"] + rel_path = symbol["relative_path"] + children = node.get("children", []) + if children: + result[rel_path].append({name_path: JetBrainsApi._transform_hierarchy_nodes(children)}) + else: + result[rel_path].append(name_path) + return dict(result) + + @facade_method(corresponding_tool=JetBrainsTypeHierarchyTool) + def get_type_hierarchy( + self, + name_path: str, + relative_path: str, + hierarchy_type: Literal["super", "sub", "both"] = "both", + depth: int | None = 1, + max_answer_chars: int = -1, + ) -> JsonObject: + """ + Gets the type hierarchy of a symbol (supertypes, subtypes, or both). + + :param name_path: name path of the symbol for which to get the type hierarchy. + :param relative_path: the relative path to the file containing the symbol. + :param hierarchy_type: which hierarchy to retrieve: "super" for parent classes/interfaces, + "sub" for subclasses/implementations, or "both" for both directions. Default is "both". + :param depth: depth limit for hierarchy traversal (None or 0 for unlimited). Default is 1. + :return: the file-grouped hierarchy, with keys "supertypes" and/or "subtypes" (and "levels_not_included" + if the depth limit truncated the hierarchy) + """ + result: dict[str, dict | list] = {} + levels_not_included = {} + with self._client() as client: + if hierarchy_type in ("super", "both"): + response = client.get_supertypes(name_path=name_path, relative_path=relative_path, depth=depth) + if "num_levels_not_included" in response: + levels_not_included["supertypes"] = response["num_levels_not_included"] + result["supertypes"] = self._transform_hierarchy_nodes(response.get("hierarchy")) + if hierarchy_type in ("sub", "both"): + response = client.get_subtypes(name_path=name_path, relative_path=relative_path, depth=depth) + if "num_levels_not_included" in response: + levels_not_included["subtypes"] = response["num_levels_not_included"] + result["subtypes"] = self._transform_hierarchy_nodes(response.get("hierarchy")) + if levels_not_included: + result["levels_not_included"] = levels_not_included + return self._json_object(result, max_answer_chars) + + @facade_method(corresponding_tool=JetBrainsFindDeclarationTool) + def find_declaration(self, relative_path: str, regex: str, include_body: bool = False) -> JetBrainsSymbolCollection: + r""" + Finds the declaration of a symbol based on an occurrence of the symbol in a source file, specified by a regex. + + :param relative_path: the relative path to the source file containing the symbol for which to find the declaration. + :param regex: a regular expression with one group, where the group matches the symbol for which to perform the lookup. + For example, to find the declaration of the `process` method in a call like `obj.process()`, + pass an expression like "obj\.(process)\(process_input_arg=37\)". + Prefer regexes with sufficiently large context around the group to render the match unambiguous. + Uses Python syntax with MULTILINE and DOTALL flags enabled. + :param include_body: whether to include the symbol's body in the result. Default False. + :return: the declaring symbol(s) + """ + content = self._get_project().read_file(relative_path) + coords = find_text_coordinates(content, regex, require_unique=True) + assert coords is not None + with self._client() as client: + response = client.find_declaration( + relative_path=relative_path, line=coords.line, col=coords.col, include_quick_info=False, include_body=include_body + ) + return JetBrainsSymbolCollection(response["symbols"], JetBrainsSymbolCollectionRenderer(self._agent, -1)) + + @facade_method(corresponding_tool=JetBrainsFindImplementationsTool) + def find_implementations(self, relative_path: str, name_path: str) -> JetBrainsSymbolCollection: + """ + Finds the implementations of a symbol. + + :param relative_path: the relative path to the source file containing the symbol for which to find implementations. + :param name_path: name path of the symbol for which to find implementations + :return: the implementing symbols + """ + with self._client() as client: + response = client.find_implementations(relative_path=relative_path, name_path=name_path, include_quick_info=False) + return JetBrainsSymbolCollection(response["symbols"], JetBrainsSymbolCollectionRenderer(self._agent, -1)) + + # edit operations + + @facade_method(can_edit=True, corresponding_tool=JetBrainsRenameTool) + def rename( + self, + relative_path: str, + new_name: str, + name_path: str | None = None, + rename_in_comments: bool = False, + rename_in_text_occurrences: bool = False, + ) -> JsonObject: + """ + Renames a symbol, file or directory throughout the codebase. + + Note: renaming in comments/text is on a best-effort basis by the IDE; if the symbol name is non-unique, further + verification is recommended. + + :param relative_path: if `name_path` is passed, the relative path of the file containing the symbol. + Otherwise, the path to the directory or file to rename. + :param new_name: the new name + :param name_path: the name path of the symbol to rename or None if renaming a file or directory. + :param rename_in_comments: whether to also rename occurrences in comments. Default False. + :param rename_in_text_occurrences: whether to also rename occurrences in text. Default False. + :return: the result of the operation + """ + result = JetBrainsCodeEditor(self._get_project()).rename_symbol( + name_path=name_path, + relative_path=relative_path, + new_name=new_name, + rename_in_comments=rename_in_comments, + rename_in_text_occurrences=rename_in_text_occurrences, + ) + return self._json_object(result) + + @facade_method(beta=True, can_edit=True, corresponding_tool=JetBrainsMoveTool) + def move( + self, + relative_path: str, + name_path: str | None = None, + target_relative_path: str | None = None, + target_parent_name_path: str | None = None, + ) -> JsonObject: + """ + Moves a symbol, file or directory to a different location and automatically updates all references to affected symbols. + + **Important**: this should always be preferred to naive moving (e.g. via file system operations or edits) + as it is much more reliable and efficient. It is always safe to use. For some symbols, moving may not be applicable, + and will result in no edits and a suitable error message. + The target location is the new parent of the symbol, + i.e. the moved entity is never renamed by the operation, only moved. + + Valid moves: + - Symbol: + * (relative_path, name_path) -> new parent symbol (target_relative_path, target_parent_name_path) + * (relative_path, name_path) -> top level of target file or directory (target_relative_path) + Always consider the concrete language-specific semantics! + - target is a file: valid for languages like Python, where files are modules + - target is a directory: valid for languages like Java, where directories are packages and can contain classes + - File or directory: + * relative_path -> new parent directory (target_relative_path) + + :param relative_path: the relative path to the file containing the symbol to move. + :param name_path: the name path of the symbol to move (empty for moving file or dir). + :param target_relative_path: the relative path of the target directory or file. + :param target_parent_name_path: the name path of the target parent symbol. + :return: the result of the operation + """ + with self._client() as client: + result = client.move( + name_path=name_path or None, + relative_path=relative_path, + target_parent_name_path=target_parent_name_path or None, + target_relative_path=target_relative_path or None, + ) + return self._json_object(result) + + @facade_method(beta=True, can_edit=True, corresponding_tool=JetBrainsSafeDeleteTool) + def safe_delete( + self, relative_path: str, name_path: str | None = None, delete_even_if_used: bool = False, propagate: bool = False + ) -> JsonObject: + """ + Safely deletes a symbol, file, or directory, checking for usages first and propagating deletion, if desired. + + Propagation means it is possible to request deleting of usages and cleaning up of unused code. + Propagation is powerful for cleaning up code but should be used with care. + **Important**: this should always be preferred to naive deleting (e.g. via file system operations or edits). + When using it, you don't have to search for usages first, as the operation will do it for you. + + :param relative_path: the relative path to the file containing the symbol to delete. + :param name_path: the name path of the symbol to delete. + A name path identifies a symbol within a source file, e.g. "MyClass/my_method". + Omit for deleting a file or directory. + :param delete_even_if_used: whether to force deletion even if the symbol still has usages. + Default is False (safe mode: will report usages instead of deleting). + :param propagate: whether to propagate the deletion to usages of the symbol and also + remove symbols that become unused after the deletion. Default is False. + :return: the result of the operation + """ + with self._client() as client: + result = client.safe_delete( + name_path=name_path or None, relative_path=relative_path, delete_even_if_used=delete_even_if_used, propagate=propagate + ) + return self._json_object(result) + + @facade_method(beta=True, can_edit=True, corresponding_tool=JetBrainsInlineSymbol) + def inline_symbol(self, name_path: str, relative_path: str, keep_definition: bool = False) -> JsonObject: + """ + Inlines a symbol, replacing all call sites with the symbol's body. + + **Important**: this should always be preferred to naive inlining (e.g. via searching for references and + editing them). + + :param name_path: the name path of the symbol to inline (usually a method/function, but also classes may be amenable to inlining, + which turns invocation into anonymous class creation) + :param relative_path: the relative path to the file containing the symbol to inline. + :param keep_definition: whether to keep the original method definition after inlining all call sites. + May be ignored in some cases (e.g. when inlining a class). + :return: the result of the operation + """ + with self._client() as client: + result = client.inline_symbol(name_path=name_path, relative_path=relative_path, keep_definition=keep_definition) + return self._json_object(result) + + # inspections + + @facade_method(corresponding_tool=JetBrainsRunInspectionsTool) + def run_inspections( + self, + relative_path: str, + min_severity: str | None = None, + inspection_names: list[str] | None = None, + start_line: int | None = None, + end_line: int | None = None, + max_answer_chars: int = -1, + ) -> JsonObject: + """ + Runs IDE inspections (code analysis) on the given file and returns the problems found. + + This leverages the full power of JetBrains' static analysis engine, including language-specific + inspections, type checking, potential bugs, code style issues, and more. + + :param relative_path: the relative path to the file to inspect. + :param min_severity: minimum severity level to include in results (e.g. "ERROR", "WARNING", "WEAK_WARNING", "INFO"). + If not specified, all severities are returned. + :param inspection_names: optional list of specific inspection names to run (e.g. ["UnusedImport", "TypeMismatch"]). + If not specified, all applicable inspections are run. + :param start_line: optional 1-based start line to restrict the inspection range. + :param end_line: optional 1-based end line to restrict the inspection range. + :return: the inspection results including severity, message, and location. + """ + with self._client() as client: + result = client.run_inspections( + relative_path=relative_path, + min_severity=min_severity, + inspection_names=inspection_names, + start_line=start_line, + end_line=end_line, + ) + return self._json_object(result, max_answer_chars) + + @facade_method(corresponding_tool=JetBrainsListInspectionsTool) + def list_inspections( + self, language: str | None = None, group_path_contains: str | None = None, max_answer_chars: int = -1 + ) -> JsonObject: + """ + Lists available IDE inspections. + + Use this to discover which inspections can be passed to `run_inspections` via `inspection_names`. + + :param language: optional language to filter by (e.g. "Java", "Python", "Kotlin"). + :param group_path_contains: optional substring to match against the inspection group path + (e.g. "probable bugs", "code style"). + :return: the list of available inspections including name, group path, and language. + """ + with self._client() as client: + result = client.list_inspections(language=language, group_path_contains=group_path_contains) + return self._json_object(result, max_answer_chars) + + # debugging + + @facade_method() + def debug_eval_info(self) -> str: + """ + Provides usage information for the debug REPL (method `debug_eval`) + + :return: the usage information + """ + return self._agent.prompt_factory.create_info_jet_brains_debug_repl() + + @facade_method(corresponding_tool=JetBrainsDebugTool) + def debug_eval(self, expression: str, repl_key: str = "default") -> str: + """ + Provides debugging functionality (run configs, breakpoints, stepping, inspection, and evaluation) + via a persistent debug REPL connected to the JetBrains IDE. + + Call `debug_eval_info()` first for usage information. + + :param expression: a Groovy/Java expression/statement to evaluate in the REPL. + If empty, closes the REPL with the given key. + :param repl_key: identifier for the REPL instance. State persists across calls with the same key. + :return: the string representation of the result + """ + with self._client() as client: + if expression: + response = client.debug_eval(repl_key=repl_key, expression=expression) + else: + response = client.debug_close(repl_key=repl_key) + return response.get("result", str(response)) diff --git a/src/serena/repl/api/lsp_api.py b/src/serena/repl/api/lsp_api.py new file mode 100644 index 00000000..8b86bff7 --- /dev/null +++ b/src/serena/repl/api/lsp_api.py @@ -0,0 +1,789 @@ +# SPDX-License-Identifier: GPL-3.0-or-later +""" +The implementation of language server (LSP)-backed operations. +""" + +import os +from collections import Counter, defaultdict +from collections.abc import Callable, Sequence +from dataclasses import dataclass +from typing import TYPE_CHECKING, Any + +from serena.code_editor import LanguageServerCodeEditor +from serena.lsp.lsp_diagnostics import GroupedDiagnostics +from serena.symbol import ( + LanguageServerSymbol, + LanguageServerSymbolDictGrouper, + LanguageServerSymbolRetriever, + ReferenceInLanguageServerSymbol, + SymbolDictGrouper, +) +from serena.tools import ( + FindDeclarationTool, + FindImplementationsTool, + FindReferencingSymbolsTool, + FindSymbolTool, + GetDiagnosticsForFileTool, + GetDiagnosticsForSymbolTool, + GetSymbolsOverviewTool, + RenameSymbolTool, + RestartLanguageServerTool, + SafeDeleteSymbol, +) +from serena.util.text_utils import TextOutputUtils, find_text_coordinates +from solidlsp.lsp_protocol_handler.lsp_types import SymbolKind + +from ..facade import SUCCESS_RESULT, FacadeApi, ReferencedType, facade_method +from ..representable import Renderer, RepresentableViaRenderer + +if TYPE_CHECKING: + from serena.agent import SerenaAgent + + +class LspSymbolCollection(RepresentableViaRenderer): + """ + A collection of symbols (`LanguageServerSymbol`) retrieved via the language server. + """ + + symbols: list[LanguageServerSymbol] + + def __init__( + self, + symbols: list[LanguageServerSymbol], + renderer: "LspSymbolCollectionRenderer", + info_by_symbol: dict[LanguageServerSymbol, str] | None = None, + ): + """ + :param symbols: the list of symbols + :param renderer: the renderer to use for representing the collection + :param info_by_symbol: additional (hover-like) info per symbol, if requested + """ + super().__init__(renderer) + self.symbols = symbols + self.info_by_symbol_ = info_by_symbol or {} + + def __len__(self) -> int: + return len(self.symbols) + + def relative_path_to_name_paths_(self) -> dict[str, list[str]]: + result: defaultdict[str, list[str]] = defaultdict(list) + for s in self.symbols: + result[s.location.relative_path or "unknown"].append(s.get_name_path()) + return result + + +class LspSymbol(RepresentableViaRenderer): + """ + A single symbol retrieved via the language server (see `LspSymbolCollection` for the symbol's interface). + """ + + symbol: LanguageServerSymbol + + def __init__(self, symbol: LanguageServerSymbol, renderer: "LspSymbolRenderer", info: str | None = None): + """ + :param symbol: the symbol + :param renderer: the renderer to use for representing the symbol + :param info: additional (hover-like) info on the symbol, if requested + """ + super().__init__(renderer) + self.symbol = symbol + self.info_ = info + + +@dataclass(kw_only=True) +class SymbolOutputParams: + name_path: bool = True + name: bool = False + kind: bool = False + location: bool = False + depth: int = 0 + body_location: bool = False + children_body: bool = False + children_name: bool | None = None + children_name_path: bool | None = None + relative_path: bool = False + include_body: bool = False + include_info: bool = False + child_inclusion_predicate: Callable[[LanguageServerSymbol], bool] | None = None + + +class LspSymbolCollectionRenderer(Renderer[LspSymbolCollection]): + """ + Renders a symbol collection as (optionally grouped) JSON according to the output parameters, falling back + to a mapping from files to name paths if the length limit is exceeded. + """ + + def __init__( + self, + agent: "SerenaAgent", + max_answer_chars: int, + output_params: SymbolOutputParams, + grouper: SymbolDictGrouper | None = None, + ): + super().__init__(agent, max_answer_chars) + self._output_params = output_params + self._grouper = grouper + + def symbol_dicts_( + self, symbols: list[LanguageServerSymbol], info_by_symbol: dict[LanguageServerSymbol, str] + ) -> list[LanguageServerSymbol.OutputDict]: + """ + :param symbols: the symbols to convert + :param info_by_symbol: additional info to include per symbol, if any + :return: the dict representations of the symbols according to the output parameters (including the info) + """ + p = self._output_params + symbol_dicts = [ + s.to_dict( + kind=p.kind, + name_path=p.name_path, + name=p.name, + location=p.location, + relative_path=p.relative_path, + body_location=p.body_location, + depth=p.depth, + body=p.include_body, + children_body=p.children_body, + children_name=p.children_name, + children_name_path=p.children_name_path, + child_inclusion_predicate=p.child_inclusion_predicate, + ) + for s in symbols + ] + for s, s_dict in zip(symbols, symbol_dicts, strict=True): + if symbol_info := info_by_symbol.get(s): + # In python 3.15 we could specify extra_items=True in the TypedDict definition, + # https://peps.python.org/pep-0728/ + # If we ever upgrade to 3.15, we can remove the type: ignore[typeddict-unknown-key] + s_dict["info"] = symbol_info + return symbol_dicts + + def _group(self, symbol_dicts: list[LanguageServerSymbol.OutputDict]) -> Any: + return self._grouper.group(symbol_dicts) if self._grouper is not None else symbol_dicts + + def render(self, obj: LspSymbolCollection) -> str: + def create_short_result_relative_path_to_name_paths() -> str: + return f"Shortened result:\n{TextOutputUtils.to_json(obj.relative_path_to_name_paths_())}" + + result = self._to_json(self._group(self.symbol_dicts_(obj.symbols, obj.info_by_symbol_))) + return self._limit_length(result, shortened_result_factories=[create_short_result_relative_path_to_name_paths]) + + +class LspSymbolRenderer(Renderer[LspSymbol]): + """ + Renders a single symbol as JSON, using a collection renderer for the conversion. + """ + + def __init__(self, agent: "SerenaAgent", max_answer_chars: int, collection_renderer: LspSymbolCollectionRenderer): + super().__init__(agent, max_answer_chars) + self._collection_renderer = collection_renderer + + def render(self, obj: LspSymbol) -> str: + info_by_symbol = {obj.symbol: obj.info_} if obj.info_ else {} + symbol_dict = self._collection_renderer.symbol_dicts_([obj.symbol], info_by_symbol)[0] + return self._limit_length(self._to_json(symbol_dict)) + + +class LspSymbolsOverviewRenderer(LspSymbolCollectionRenderer): + """ + Renders a file's symbol overview, falling back to a depth-0 overview and finally symbol counts by kind + if the length limit is exceeded. + """ + + def render(self, obj: LspSymbolCollection) -> str: + symbol_dicts = self.symbol_dicts_(obj.symbols, obj.info_by_symbol_) + result = self._to_json(self._group(symbol_dicts)) + + def make_kind_counts() -> str: + kind_names = [d.get("kind", "unknown") for d in symbol_dicts] + return f"Symbol counts by kind:\n{self._to_json(Counter(kind_names))}" + + shortened_results: list[Callable[[], str]] = [make_kind_counts] + if self._output_params.depth > 0: + + def make_depth_0_result() -> str: + depth_0_dicts = [d.copy() for d in symbol_dicts] + for d in depth_0_dicts: + d.pop("children", None) + return "Depth 0 overview:\n" + self._to_json(self._group(depth_0_dicts)) + + shortened_results.insert(0, make_depth_0_result) + + return self._limit_length(result, shortened_result_factories=shortened_results) + + +class LspReferenceCollection(RepresentableViaRenderer): + """ + The references to a symbol (`ReferenceInLanguageServerSymbol`). + """ + + references: list[ReferenceInLanguageServerSymbol] + + def __init__( + self, + references: list[ReferenceInLanguageServerSymbol], + contents_around_references: list[str], + renderer: "LspReferenceCollectionRenderer", + ): + """ + :param references: the references + :param contents_around_references: for each reference, the code around it (for display) + :param renderer: the renderer to use for representing the collection + """ + super().__init__(renderer) + self.references = references + self.contents_around_references_ = contents_around_references + + def __len__(self) -> int: + return len(self.references) + + +class LspReferenceCollectionRenderer(Renderer[LspReferenceCollection]): + """ + Renders references as grouped JSON including the code around each reference, falling back to + references without code, per-file counts and finally the total count if the length limit is exceeded. + """ + + def __init__(self, agent: "SerenaAgent", max_answer_chars: int, grouper: SymbolDictGrouper): + super().__init__(agent, max_answer_chars) + self._grouper = grouper + + def render(self, obj: LspReferenceCollection) -> str: + reference_dicts = [] + ref_summaries = [] + for ref, content_around_ref in zip(obj.references, obj.contents_around_references_, strict=True): + ref_dict = dict(ref.symbol.to_dict(kind=True, relative_path=True, depth=0, body=False, body_location=True)) + ref_dict["content_around_reference"] = content_around_ref + reference_dicts.append(ref_dict) + ref_summaries.append( + { + "name_path": ref_dict.get("name_path"), + "kind": ref_dict.get("kind"), + "relative_path": ref_dict.get("relative_path"), + "reference_line": ref.line, + } + ) + + result = self._to_json(self._grouper.group(reference_dicts)) + + # shortened result closures, from least to most aggressive shortening + def make_refs_without_context() -> str: + return f"References without surrounding lines:\n{self._to_json(self._grouper.group([dict(s) for s in ref_summaries]))}" + + def make_per_file_counts() -> str: + counts = Counter(str(r["relative_path"]) for r in ref_summaries) + return f"Reference counts per file:\n{self._to_json(counts)}" + + def make_summary() -> str: + return f"Found {len(ref_summaries)} references." + + return self._limit_length(result, shortened_result_factories=[make_refs_without_context, make_per_file_counts, make_summary]) + + +class LspDiagnostics(RepresentableViaRenderer): + """ + Diagnostics grouped as `relative_path -> severity -> name_path -> diagnostics`; see `grouped.get_dict()`. + """ + + grouped: GroupedDiagnostics + + def __init__(self, grouped: GroupedDiagnostics, renderer: "LspDiagnosticsRenderer"): + """ + :param grouped: the grouped diagnostics + :param renderer: the renderer to use for representing the diagnostics + """ + super().__init__(renderer) + self.grouped = grouped + + +class LspDiagnosticsRenderer(Renderer[LspDiagnostics]): + def render(self, obj: LspDiagnostics) -> str: + return self._limit_length(self._to_json(obj.grouped.get_dict())) + + +def _is_not_low_level(symbol: LanguageServerSymbol) -> bool: + return not symbol.is_low_level() + + +class LspApi(FacadeApi): + FILE_LEVEL_DIAGNOSTIC_BUCKET = "" + """the name path under which diagnostics that cannot be mapped to a symbol are grouped""" + + # groupers for the various symbol collections; top-level symbols are grouped by the first key list, + # children by the second. + # For find_symbol, we group children by kind, keeping just the name (the parent's name_path makes it unambiguous); + # we don't group the top-level result list because many tests rely on it being a flat list of symbol dicts + find_symbol_dict_grouper_ = LanguageServerSymbolDictGrouper([], ["kind"], collapse_singleton=True) + references_grouper_ = LanguageServerSymbolDictGrouper(["relative_path", "kind"], ["kind"], collapse_singleton=True) + overview_grouper_ = LanguageServerSymbolDictGrouper(["kind"], ["kind"], collapse_singleton=True) + + def __init__(self, agent: "SerenaAgent") -> None: + super().__init__( + agent, + name="lsp", + description="symbol-level operations on the codebase backed by language servers", + types=[ + ReferencedType( + LanguageServerSymbol, + members=[ + "name", + "get_name_path", + "relative_path", + "symbol_kind_name", + "line", + "column", + "body", + "get_body_line_numbers", + "iter_children", + "iter_ancestors", + "get_parent", + ], + ), + ], + ) + + def _create_symbol_retriever(self) -> LanguageServerSymbolRetriever: + assert self._agent.get_language_backend().is_lsp(), "Language server operations require the language server backend" + return LanguageServerSymbolRetriever(self._get_project()) + + def _create_ls_code_editor(self, symbol_retriever: LanguageServerSymbolRetriever | None = None) -> LanguageServerCodeEditor: + return LanguageServerCodeEditor(symbol_retriever or self._create_symbol_retriever()) + + @staticmethod + def _request_info( + symbol_retriever: LanguageServerSymbolRetriever, symbols: list[LanguageServerSymbol], output_params: SymbolOutputParams + ) -> dict[LanguageServerSymbol, str]: + """ + :return: additional (hover-like) info per symbol, if the output parameters request it (and not the body, which + supersedes it); requested eagerly, such that results are self-contained + """ + if output_params.include_info and not output_params.include_body: + return {s: info for s, info in symbol_retriever.request_info_for_symbol_batch(symbols).items() if info} + return {} + + def _retrieve_content_around_reference(self, reference: ReferenceInLanguageServerSymbol) -> str: + relative_path = reference.symbol.location.relative_path + assert relative_path is not None, f"Referencing symbol {reference.symbol.name} has no relative path, this is likely a bug." + content = self._get_project().retrieve_content_around_line( + relative_file_path=relative_path, line=reference.line, context_lines_before=1, context_lines_after=1 + ) + return content.to_display_string() + + @staticmethod + def _parse_kinds(kinds: Sequence[int]) -> Sequence[SymbolKind] | None: + return [SymbolKind(k) for k in kinds] if kinds else None + + def _create_diagnostics(self, grouped: GroupedDiagnostics, max_answer_chars: int) -> LspDiagnostics: + return LspDiagnostics(grouped, LspDiagnosticsRenderer(self._agent, max_answer_chars)) + + # language server management + + @facade_method(uses_project_server=True, optional=True, corresponding_tool=RestartLanguageServerTool) + def restart_language_server(self) -> str: + """ + Restarts the language server(s). + + Use this only on explicit user request or after confirmation; it may be necessary if a language server hangs. + + :return: a success message + """ + self._agent.reset_language_server_manager() + return SUCCESS_RESULT + + # read operations + + @facade_method(uses_project_server=True, corresponding_tool=GetSymbolsOverviewTool) + def get_symbols_overview(self, relative_path: str, depth: int = -1, max_answer_chars: int = -1) -> LspSymbolCollection: + """ + Gets an overview of the symbols defined in the given file (classes, methods, fields, functions, etc.) + + Returns STRUCTURE only, without bodies. This is the cheap, structure-first way to learn what a file + contains: it costs far less context than reading the whole file. + + :param relative_path: the relative path to the file to get the overview of + :param depth: depth up to which descendants shall be retrieved. + Default (-1) results in a language specific choice: 1 for java and kotlin and 0 for other languages + :return: the top-level symbols of the file + """ + # Note: file system sync not required (relevant file is opened in the language server explicitly) + if depth == -1: + depth = 1 if relative_path.endswith((".java", ".kt")) else 0 + + symbol_retriever = self._create_symbol_retriever() + + # the symbol overview is capable of working with both files and directories, but we require a file + file_path = os.path.join(self._get_project().project_root, relative_path) + if not os.path.exists(file_path): + raise FileNotFoundError(f"File or directory {relative_path} does not exist in the project.") + if os.path.isdir(file_path): + raise ValueError(f"Expected a file path, but got a directory path: {relative_path}. ") + if not symbol_retriever.can_analyze_file(relative_path): + raise ValueError( + f"Cannot extract symbols from file {relative_path}. " + f"Active language servers: {[l.get_key() for l in self._agent.get_active_language_server_ids()]}" + ) + + symbols = symbol_retriever.get_symbol_overview(relative_path)[relative_path] + output_params = SymbolOutputParams( + name_path=False, + name=True, + depth=depth, + kind=True, + relative_path=False, + location=False, + child_inclusion_predicate=_is_not_low_level, + ) + renderer = LspSymbolsOverviewRenderer(self._agent, max_answer_chars, output_params, grouper=self.overview_grouper_) + return LspSymbolCollection(symbols, renderer) + + @facade_method(uses_project_server=True, corresponding_tool=FindSymbolTool) + def find_symbol( + self, + name_path_pattern: str, + depth: int = 0, + relative_path: str = "", + include_body: bool = False, + include_info: bool = False, + include_kinds: Sequence[int] = (), + exclude_kinds: Sequence[int] = (), + substring_matching: bool = False, + max_matches: int = -1, + max_answer_chars: int = -1, + ) -> LspSymbolCollection: + """ + Finds symbols and code entities (classes, methods, etc.) based on the given name path pattern. + + The returned symbol information can be used for edits or further queries. + Specify `depth > 0` to also retrieve children/descendants (e.g., methods of a class). + + A name path is a path in the symbol tree *within a source file*. + For example, the method `my_method` defined in class `MyClass` would have the name path `MyClass/my_method`. + If a symbol is overloaded (e.g., in Java), a 0-based index is appended (e.g. "MyClass/my_method[0]") to + uniquely identify it. + + To search for a symbol, you provide a name path pattern that is used to match against name paths. + It can be + * a simple name (e.g. "method"), which will match any symbol with that name + * a relative path like "class/method", which will match any symbol with that name path suffix + * an absolute name path "/class/method" (absolute name path), which requires an exact match of the full name path within the source file. + Append an index `[i]` to match a specific overload only, e.g. "MyClass/my_method[1]". + + :param name_path_pattern: the name path matching pattern (see above) + :param depth: depth up to which descendants shall be retrieved (e.g. use 1 to also retrieve immediate children; + for the case where the symbol is a class, this will return its methods). + Ignored if `include_body=True`. Default 0. + :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 include_info: whether to include additional info (hover-like, typically including docstring and signature), + about the symbol (ignored if include_body is True). Info is never included for child symbols. + Note: Depending on the language, this can be slow (e.g., C/C++). + :param include_kinds: (optional) limits results to the given LSP symbol kinds (integers, i.e. values of `SymbolKind`) + :param exclude_kinds: (optional) list of LSP symbol kinds (integers, i.e. values of `SymbolKind`) to exclude. + :param substring_matching: If True, use substring matching for the last segment of `name_path_pattern` + (i.e. the name of the symbol, e.g. "foo" in "Class/foo" or "my_method" in "my_method"). + :param max_matches: Maximum number of permitted matches. If exceeded, an error containing a shortened result is raised, + which allows refining the search. -1 (default) means no limit. Set to 1 to search for a unique symbol. + :return: the symbols (with locations) matching the name path pattern + """ + # Note: file system sync not required; the symbol finder opens all relevant source files explicitly in the case of changes + + if include_body: + depth = 0 # ignore user-specified depth if include_body is True + assert max_matches != 0, "max_matches must be > 0 or equal to -1." + symbol_retriever = self._create_symbol_retriever() + symbols = symbol_retriever.find( + name_path_pattern, + include_kinds=self._parse_kinds(include_kinds), + exclude_kinds=self._parse_kinds(exclude_kinds), + substring_matching=substring_matching, + within_relative_path=relative_path, + ) + + output_params = SymbolOutputParams( + kind=True, + name_path=True, + name=False, + relative_path=True, + body_location=True, + depth=depth, + include_body=include_body, + children_name=True, + children_name_path=False, + include_info=include_info, + ) + renderer = LspSymbolCollectionRenderer(self._agent, max_answer_chars, output_params, grouper=self.find_symbol_dict_grouper_) + symbol_collection = LspSymbolCollection(symbols, renderer, self._request_info(symbol_retriever, symbols, output_params)) + + # check for max_matches limit exceeded + n_matches = len(symbols) + if 0 < max_matches < n_matches: + raise ValueError( + f"Matched {n_matches}>{max_matches=} symbols.\n" + TextOutputUtils.to_json(symbol_collection.relative_path_to_name_paths_()) + ) + + return symbol_collection + + @facade_method(uses_project_server=True, corresponding_tool=FindReferencingSymbolsTool) + def find_referencing_symbols( + self, + name_path: str, + relative_path: str, + include_kinds: Sequence[int] = (), + exclude_kinds: Sequence[int] = (), + max_answer_chars: int = -1, + ) -> LspReferenceCollection: + """ + Finds references to the symbol at the given `name_path`. + + The result will contain metadata about the referencing symbols as well as a short code snippet around the reference. + + :param name_path: name path of the symbol + :param relative_path: the relative path to the file containing the symbol for which to find references. + :param include_kinds: (optional) limits results to the given LSP symbol kinds (integers, i.e. values of `SymbolKind`) + :param exclude_kinds: optional list of LSP symbol kinds (integers, i.e. values of `SymbolKind`) to exclude. + :return: the references to the symbol + """ + # file system sync needed for case where symbol finder does not perform a global search, updating everything + if relative_path: + self._get_project().ls_sync_file_system_changes() + + symbol_retriever = self._create_symbol_retriever() + references = symbol_retriever.find_referencing_symbols( + name_path, + relative_file_path=relative_path, + include_body=False, # it is probably never a good idea to include the body of the referencing symbols + include_kinds=self._parse_kinds(include_kinds), + exclude_kinds=self._parse_kinds(exclude_kinds), + ) + contents_around_references = [self._retrieve_content_around_reference(ref) for ref in references] + renderer = LspReferenceCollectionRenderer(self._agent, max_answer_chars, self.references_grouper_) + return LspReferenceCollection(references, contents_around_references, renderer) + + @facade_method(uses_project_server=True, corresponding_tool=FindImplementationsTool) + def find_implementations( + self, + name_path: str, + relative_path: str, + include_info: bool = False, + include_kinds: Sequence[int] = (), + exclude_kinds: Sequence[int] = (), + max_answer_chars: int = -1, + ) -> LspSymbolCollection: + """ + Finds implementations of the symbol at the given `name_path`. + + :param name_path: the symbol's name path + :param relative_path: the relative path to the file containing the symbol for which to find implementations. + Note that here you can't pass a directory but must pass a file. + :param include_info: whether to include additional info (hover-like, typically including docstring and signature), + about the implementing symbols. + :param include_kinds: (optional) limits results to the given LSP symbol kinds (integers, i.e. values of `SymbolKind`) + :param exclude_kinds: (optional) list of LSP symbol kinds (integers, i.e. values of `SymbolKind`) to exclude. + :return: the symbols implementing the given symbol + """ + self._get_project().ls_sync_file_system_changes() + + symbol_retriever = self._create_symbol_retriever() + symbols = symbol_retriever.find_implementing_symbols( + name_path, + relative_file_path=relative_path, + include_body=False, + include_kinds=self._parse_kinds(include_kinds), + exclude_kinds=self._parse_kinds(exclude_kinds), + ) + output_params = SymbolOutputParams(kind=True, relative_path=True, body_location=True, include_info=include_info) + renderer = LspSymbolCollectionRenderer(self._agent, max_answer_chars, output_params) + return LspSymbolCollection(symbols, renderer, self._request_info(symbol_retriever, symbols, output_params)) + + @facade_method(uses_project_server=True, corresponding_tool=FindDeclarationTool) + def find_declaration( + self, + relative_path: str, + regex: str, + containing_symbol_name_path: str | None = None, + include_body: bool = False, + include_info: bool = False, + ) -> LspSymbol: + r""" + Finds the declaration of a symbol based on an occurrence of the symbol in a source file, specified by a regex. + + :param relative_path: the relative path to the source file containing the symbol for which to find the declaration. + :param regex: a regular expression with one group, where the group matches the symbol for which to perform the lookup. + For example, to find the declaration of the `process` method in a call like `obj.process()`, + pass an expression like "obj\.(process)\(process_input_arg=37\)". + Prefer regexes with sufficiently large context around the group to render the match unambiguous. + Uses Python syntax with MULTILINE and DOTALL flags enabled. + :param containing_symbol_name_path: optional name path of a containing symbol whose body shall be searched instead of the full file. + :param include_body: whether to include the symbol's body in the result. Default False. + :param include_info: whether to include additional info (hover-like). Default False. + :return: the declaring symbol + """ + self._get_project().ls_sync_file_system_changes() + symbol_retriever = self._create_symbol_retriever() + + # find relevant location for lookup + editor = self._create_ls_code_editor(symbol_retriever) + if not containing_symbol_name_path: + content = editor.read_file(relative_path) + coords = find_text_coordinates(content, regex, require_unique=True) + assert coords is not None + else: + symbol = symbol_retriever.find_unique(name_path_pattern=containing_symbol_name_path, within_relative_path=relative_path) + body_line_numbers = symbol.get_body_line_numbers_or_raise() + content = editor.read_file(relative_path, lines=body_line_numbers) + coords = find_text_coordinates(content, regex, require_unique=True) + assert coords is not None + coords.line += body_line_numbers[0] + + # retrieve declaration + defining_symbol = symbol_retriever.find_declaration( + relative_file_path=relative_path, line=coords.line, column=coords.col, include_body=include_body + ) + if defining_symbol is None: + raise ValueError( + f"No symbol declaration found at the location of the regex match. Location: {relative_path}:{coords.line}:{coords.col}." + ) + + output_params = SymbolOutputParams( + kind=True, relative_path=True, body_location=True, include_body=include_body, include_info=include_info + ) + collection_renderer = LspSymbolCollectionRenderer(self._agent, -1, output_params) + info = self._request_info(symbol_retriever, [defining_symbol], output_params).get(defining_symbol) + return LspSymbol(defining_symbol, LspSymbolRenderer(self._agent, -1, collection_renderer), info) + + @facade_method(uses_project_server=True, corresponding_tool=GetDiagnosticsForFileTool) + def get_diagnostics_for_file( + self, relative_path: str, start_line: int = 0, end_line: int = -1, min_severity: int = 4, max_answer_chars: int = -1 + ) -> LspDiagnostics: + """ + Gets diagnostics for a file. + + Diagnostics are grouped as `relative_path -> severity -> name_path -> diagnostics_results`. + If a diagnostic cannot be mapped to a symbol, it is grouped under the special name path ``. + + :param relative_path: the relative path to the file to inspect. + :param start_line: the first 0-based line to include. Defaults to 0. + :param end_line: the last 0-based line to include. Defaults to -1, which means until the end of the file. + :param min_severity: minimum LSP severity to include, where 1=Error, 2=Warning, 3=Information, 4=Hint. + Diagnostics with lower-or-equal numeric severity are returned. + :return: the grouped diagnostics for the requested file. + """ + self._get_project().ls_sync_file_system_changes() + + symbol_retriever = self._create_symbol_retriever() + diagnostics = symbol_retriever.get_file_diagnostics( + relative_file_path=relative_path, start_line=start_line, end_line=end_line, min_severity=min_severity + ) + + grouped_diagnostics = GroupedDiagnostics() + for diagnostic in diagnostics: + diag_start = diagnostic["range"]["start"] + owner_symbol = symbol_retriever.find_diagnostic_owner_symbol( + relative_file_path=relative_path, line=diag_start["line"], column=diag_start["character"] + ) + name_path = owner_symbol.get_name_path() if owner_symbol is not None else self.FILE_LEVEL_DIAGNOSTIC_BUCKET + grouped_diagnostics.add(relative_path, name_path, diagnostic) + + return self._create_diagnostics(grouped_diagnostics, max_answer_chars) + + @facade_method(uses_project_server=True, optional=True, corresponding_tool=GetDiagnosticsForSymbolTool) + def get_diagnostics_for_symbol( + self, + name_path: str, + reference_file: str = "", + check_symbol_references: bool = False, + min_severity: int = 4, + max_answer_chars: int = -1, + ) -> LspDiagnostics: + """ + Gets diagnostics for the specified symbol. + + When `check_symbol_references` is true, diagnostics for all referencing symbols are also included. + The result is grouped as `relative_path -> severity -> name_path -> diagnostics_results`. + + :param name_path: the name path of the symbol to inspect. + :param reference_file: optional file path used to disambiguate the symbol search. + :param check_symbol_references: whether to additionally collect diagnostics for symbols that reference the symbol. + :param min_severity: minimum LSP severity to include, where 1=Error, 2=Warning, 3=Information, 4=Hint. + Diagnostics with lower-or-equal numeric severity are returned. + :return: the grouped diagnostics for the requested symbol and, optionally, its referencing symbols. + """ + self._get_project().ls_sync_file_system_changes() + + symbol_retriever = self._create_symbol_retriever() + diagnostics_by_symbol = symbol_retriever.get_symbol_diagnostics( + name_path=name_path, + reference_file=reference_file or None, + check_symbol_references=check_symbol_references, + min_severity=min_severity, + ) + + grouped_diagnostics = GroupedDiagnostics() + for symbol, diagnostics in diagnostics_by_symbol.items(): + relative_path = symbol.relative_path + if relative_path is None: + continue + for diagnostic in diagnostics: + grouped_diagnostics.add(relative_path, symbol.get_name_path(), diagnostic) + + return self._create_diagnostics(grouped_diagnostics, max_answer_chars) + + # edit operations + + @facade_method(uses_project_server=True, can_edit=True, corresponding_tool=RenameSymbolTool) + def rename_symbol(self, name_path: str, relative_path: str, new_name: str) -> str: + """ + Renames the symbol with the given `name_path` to `new_name` throughout the entire codebase. + Note: for languages with method overloading, like Java, name_path may have to include a method's + signature to uniquely identify a method. + + :param name_path: name path of the symbol to rename + :param relative_path: the relative path to the file containing the symbol to rename + :param new_name: the new name for the symbol + :return: a result summary indicating success or failure + """ + self._get_project().ls_sync_file_system_changes() + return self._create_ls_code_editor().rename_symbol(name_path, relative_path=relative_path, new_name=new_name) + + @facade_method(uses_project_server=True, can_edit=True, corresponding_tool=SafeDeleteSymbol) + def safe_delete_symbol(self, name_path_pattern: str, relative_path: str) -> str: + """ + Deletes the symbol if it is safe to do so (i.e., if there are no references to it) + or returns a list of references to it. + + :param name_path_pattern: name path of the symbol to delete + :param relative_path: the relative path to the file containing the symbol to delete + :return: a success message, or a message listing the references preventing deletion + """ + self._get_project().ls_sync_file_system_changes() + + symbol_retriever = self._create_symbol_retriever() + symbol = symbol_retriever.find_unique(name_path_pattern, substring_matching=False, within_relative_path=relative_path) + symbol_rel_path = symbol.relative_path + assert symbol_rel_path is not None, f"Symbol {name_path_pattern} has no relative path, this is likely a bug." + assert symbol_rel_path == relative_path, f"Symbol {name_path_pattern} is not in the expected relative path {relative_path}." + symbol_name_path = symbol.get_name_path() + + # check for references + symbol_line = symbol.line + symbol_col = symbol.column + assert symbol_line is not None and symbol_col is not None, ( + f"Symbol {name_path_pattern} has no identifier position, this is likely a bug." + ) + lang_server = symbol_retriever.get_language_server(symbol_rel_path) + references_locations = lang_server.request_references(symbol_rel_path, symbol_line, symbol_col) + file_to_lines: dict[str, list[int]] = defaultdict(list) + for ref_loc in references_locations or []: + ref_relative_path = ref_loc.get("relativePath") + if ref_relative_path is None: + continue + file_to_lines[ref_relative_path].append(ref_loc["range"]["start"]["line"]) + if file_to_lines: + return f"Cannot delete, the symbol {symbol_name_path} is referenced in: {TextOutputUtils.to_json(file_to_lines)}" + + self._create_ls_code_editor(symbol_retriever).delete_symbol(symbol_name_path, relative_file_path=symbol_rel_path) + return SUCCESS_RESULT diff --git a/src/serena/repl/api/mem_api.py b/src/serena/repl/api/mem_api.py new file mode 100644 index 00000000..91f2803a --- /dev/null +++ b/src/serena/repl/api/mem_api.py @@ -0,0 +1,181 @@ +# SPDX-License-Identifier: GPL-3.0-or-later +""" +The implementation of memory operations (and onboarding, which creates the initial memories). +""" + +import logging +import platform +from typing import TYPE_CHECKING, Literal + +from serena.memories.memory_manager import MemoryManager +from serena.tools import ( + DeleteMemoryTool, + EditMemoryTool, + ListMemoriesTool, + OnboardingTool, + ReadMemoryTool, + RenameMemoryTool, + WriteMemoryTool, +) + +from ..facade import FacadeApi, facade_method +from ..representable import Renderer, RepresentableViaRenderer + +if TYPE_CHECKING: + from serena.agent import SerenaAgent + +log = logging.getLogger(__name__) + + +class MemoryList(RepresentableViaRenderer): + """ + The available memories: `memories` (writable) and `read_only_memories` (e.g. global memories), each a list of names. + """ + + def __init__(self, memories_list: MemoryManager.MemoriesList, renderer: "MemoryListRenderer"): + """ + :param memories_list: the list of memories + :param renderer: the renderer to use for representing the list + """ + super().__init__(renderer) + self.memories_list_ = memories_list + + @property + def memories(self) -> list[str]: + return sorted(self.memories_list_.memories) + + @property + def read_only_memories(self) -> list[str]: + return sorted(self.memories_list_.read_only_memories) + + +class MemoryListRenderer(Renderer[MemoryList]): + def render(self, obj: MemoryList) -> str: + return self._limit_length(self._to_json(obj.memories_list_.to_dict())) + + +class MemoryApi(FacadeApi): + def __init__(self, agent: "SerenaAgent") -> None: + super().__init__( + agent, + name="mem", + description="project memories, i.e. persistent notes for future tasks", + ) + + def _get_memory_manager(self) -> MemoryManager: + return self._get_project().memory_manager + + @facade_method(corresponding_tool=ListMemoriesTool) + def list_memories(self, topic: str = "") -> MemoryList: + """ + Lists the available memories, optionally filtered by topic. + + :param topic: the topic (prefix of the memory name, e.g. "frontend") to restrict the listing to; empty for all memories + :return: the memories + """ + return MemoryList(self._get_memory_manager().list_memories(topic), MemoryListRenderer(self._agent, -1)) + + @facade_method(corresponding_tool=ReadMemoryTool) + def read_memory(self, memory_name: str) -> str: + """ + Reads a memory. + + :param memory_name: the name of the memory + :return: the memory's content + """ + return self._get_memory_manager().load_memory(memory_name) + + @facade_method(can_edit=True, corresponding_tool=WriteMemoryTool) + def write_memory(self, memory_name: str, content: str, max_chars: int = -1) -> str: + """ + Writes information (about the active project) to a memory. + + The name should be meaningful and can include "/" to organize into topics. + If explicitly instructed, use the "global/" prefix for writing a memory that is shared across projects. + References to other memories should be inside backticks and prefixed with mem:, + e.g., `mem:auth`. + + :param memory_name: the memory name + :param content: the memory content (utf-8-encoded markdown) + :param max_chars: the maximum content length; -1 for the configured default + :return: a message indicating the result + """ + if max_chars == -1: + max_chars = self._agent.serena_config.default_max_tool_answer_chars + if len(content) > max_chars: + raise ValueError( + f"Content for {memory_name} is too long. Max length is {max_chars} characters. Please make the content shorter." + ) + return self._get_memory_manager().save_memory(memory_name, content, is_tool_context=True) + + @facade_method(can_edit=True, corresponding_tool=EditMemoryTool) + def edit_memory( + self, + memory_name: str, + needle: str, + repl: str, + mode: Literal["literal", "regex"], + allow_multiple_occurrences: bool = False, + ) -> str: + """ + Replaces content matching a pattern in a memory. + + :param memory_name: the name of the memory + :param needle: the string or regex pattern to search for. In regex mode, be careful to not replace too much! + If `mode` is "literal", this string will be matched exactly. + If `mode` is "regex", this string will be treated as a regular expression (syntax of Python's `re` module, + with the MULTILINE and DOTALL flags enabled). + :param repl: the replacement string (verbatim). + :param mode: either "literal" or "regex", specifying how the `needle` parameter is to be interpreted. + :param allow_multiple_occurrences: whether to allow matching and replacing multiple occurrences. + If false and multiple occurrences are found, an error will be raised. + :return: a message indicating the result + """ + return self._get_memory_manager().edit_memory( + memory_name, needle, repl, mode, allow_multiple_occurrences, is_tool_context=True, regex_multiline=True + ) + + @facade_method(can_edit=True, corresponding_tool=RenameMemoryTool) + def rename_memory(self, old_name: str, new_name: str) -> str: + """ + Renames or moves a memory; use "/" in the name to organize into topics. + The "global" topic should only be used if explicitly instructed. + References to other memories that are marked with the `mem:` prefix will be updated accordingly. + References in read-only memories are not affected. + + :param old_name: the current name of the memory + :param new_name: the new name of the memory + :return: a message indicating the result + """ + renaming_message, n_references_updated = self._get_memory_manager().rename_memory_and_propagate_references( + old_name, new_name, is_tool_context=True + ) + if n_references_updated > 0: + log.info(f"Updated {n_references_updated} references to memory {old_name} to {new_name}") + return renaming_message + + @facade_method(can_edit=True, corresponding_tool=DeleteMemoryTool) + def delete_memory(self, memory_name: str) -> str: + """ + Deletes a memory; only call this if instructed explicitly or permission was granted by the user. + + :param memory_name: the name of the memory + :return: a message indicating the result + """ + return self._get_memory_manager().delete_memory(memory_name, is_tool_context=True) + + @facade_method(corresponding_tool=OnboardingTool) + def onboarding(self) -> str: + """ + Provides the instructions for performing onboarding (identifying the project structure and essential tasks, + e.g. for testing or building, and recording the findings in memories). + Call this if onboarding was not performed yet, at most once per conversation. + + :return: the instructions on how to create the onboarding information + """ + # seed the project-local memory-maintenance memory (or detect a global override) so + # the prompt can point the agent at the conventions before it writes anything + memory_maintenance_name = self._get_memory_manager().ensure_memory_maintenance_memory() + return self._agent.prompt_factory.create_onboarding_prompt( + system=platform.system(), memory_maintenance_name=memory_maintenance_name + ) diff --git a/src/serena/repl/api/shell_api.py b/src/serena/repl/api/shell_api.py new file mode 100644 index 00000000..9f4db097 --- /dev/null +++ b/src/serena/repl/api/shell_api.py @@ -0,0 +1,88 @@ +# SPDX-License-Identifier: GPL-3.0-or-later +""" +The implementation of shell command execution. +""" + +import os.path +from typing import TYPE_CHECKING + +from serena.tools import ExecuteShellCommandTool +from serena.util.shell import ShellCommandResult, execute_shell_command + +from ..facade import FacadeApi, facade_method +from ..representable import Renderer, RepresentableViaRenderer + +if TYPE_CHECKING: + from serena.agent import SerenaAgent + + +class ShellCommandOutput(RepresentableViaRenderer): + """ + The outcome of a shell command: `stdout`, `stderr` (None if not captured), `return_code` and `cwd`. + """ + + def __init__(self, result: ShellCommandResult, renderer: "ShellCommandOutputRenderer"): + """ + :param result: the result of the command execution + :param renderer: the renderer to use for representing the output + """ + super().__init__(renderer) + self.result_ = result + + @property + def stdout(self) -> str: + return self.result_.stdout + + @property + def stderr(self) -> str | None: + return self.result_.stderr + + @property + def return_code(self) -> int: + return self.result_.return_code + + @property + def cwd(self) -> str: + return self.result_.cwd + + +class ShellCommandOutputRenderer(Renderer[ShellCommandOutput]): + def render(self, obj: ShellCommandOutput) -> str: + return self._limit_length(obj.result_.model_dump_json()) + + +class ShellApi(FacadeApi): + def __init__(self, agent: "SerenaAgent") -> None: + super().__init__(agent, name="shell", description="execution of shell commands") + + @facade_method(can_edit=True, corresponding_tool=ExecuteShellCommandTool) + def execute_shell_command( + self, command: str, cwd: str | None = None, capture_stderr: bool = True, max_answer_chars: int = -1 + ) -> ShellCommandOutput: + """ + Executes a shell command and returns its output. If there is a memory about suggested commands, read that first. + Never execute unsafe shell commands! + IMPORTANT: Do not use this to start + * long-running processes (e.g. servers) that are not intended to terminate quickly, + * processes that require user interaction. + + :param command: the shell command to execute + :param cwd: the working directory to execute the command in (absolute, or relative to the project root). + If None, the project root will be used. + :param capture_stderr: whether to capture and return stderr output + :return: the output (object with properties stdout, stderr, return_code and cwd) + """ + project_root = self._get_project().project_root + if cwd is None: + _cwd = project_root + elif os.path.isabs(cwd): + _cwd = cwd + else: + _cwd = os.path.join(project_root, cwd) + if not os.path.isdir(_cwd): + raise FileNotFoundError( + f"Specified a relative working directory ({cwd}), but the resulting path is not a directory: {_cwd}" + ) + + result = execute_shell_command(command, cwd=_cwd, capture_stderr=capture_stderr) + return ShellCommandOutput(result, ShellCommandOutputRenderer(self._agent, max_answer_chars)) diff --git a/src/serena/repl/external_project.py b/src/serena/repl/external_project.py new file mode 100644 index 00000000..401c2d19 --- /dev/null +++ b/src/serena/repl/external_project.py @@ -0,0 +1,67 @@ +# SPDX-License-Identifier: GPL-3.0-or-later +""" +Execution of facade methods in the context of an external project (i.e. a project other than the active one). +""" + +from typing import TYPE_CHECKING, Any + +if TYPE_CHECKING: + from serena.project_server import ProjectServerClient + + from ..agent import SerenaAgent + from .facade import FacadeMethod + + +class ExternalProjectExecution: + """ + The context in which facade methods are executed while an external project is being queried: + methods which use the project server (see `FacadeMethodInfo.uses_project_server`) are executed remotely + in the project server (if remote execution applies to the language backend), all other methods are executed + locally against the temporarily switched project. Editing methods are not permitted. + """ + + def __init__(self, project_name: str, read_only: bool, agent: "SerenaAgent") -> None: + """ + :param project_name: the name of the external project + :param read_only: whether the external project is to be treated as read-only (editing methods are not permitted) + """ + self.project_name = project_name + self._client: ProjectServerClient | None = None + self._read_only = read_only + self._agent = agent + + def is_called_remotely(self, method: "FacadeMethod") -> bool: + """ + :param method: the method to check + :return: whether the given method must be executed remotely + """ + # Any method that uses the project server must be executed remotely, + # as does any edit operation when using the LSP backend (as edit operations indirectly + # use the language server via the CodeEditor abstraction) + return method.info.uses_project_server or (self._agent.get_language_backend().is_lsp() and method.info.can_edit) + + def check_call_permission(self, method: "FacadeMethod") -> None: + """ + Checks whether the given method is permitted to be called in the context of this external project execution. + Raises an exception if the method is not permitted. + + :param method: the facade method to check + """ + if self._read_only and method.info.can_edit: + raise PermissionError(f"Editing methods are not permitted in read-only external project execution: {method.qualified_name}") + + def call_remotely(self, facade_name: str, method_name: str, args: tuple[Any, ...], kwargs: dict[str, Any]) -> Any: + """ + Executes the given facade method remotely via the project server. + + :param facade_name: the facade's name + :param method_name: the method's name + :param args: the positional arguments (must be JSON-serialisable) + :param kwargs: the keyword arguments (must be JSON-serialisable) + :return: the method's result (unpickled) + """ + if self._client is None: + from serena.project_server import ProjectServerClient + + self._client = ProjectServerClient(self._agent.serena_config) + return self._client.call_facade_method(self.project_name, facade_name, method_name, list(args), kwargs) diff --git a/src/serena/repl/facade.py b/src/serena/repl/facade.py new file mode 100644 index 00000000..96828682 --- /dev/null +++ b/src/serena/repl/facade.py @@ -0,0 +1,747 @@ +""" +The facade, i.e. the object through which REPL code accesses a group of related operations. +""" + +# SPDX-License-Identifier: GPL-3.0-or-later + +import inspect +import logging +import re +import typing +from abc import ABC +from collections.abc import Callable, Sequence +from dataclasses import dataclass +from enum import Enum +from typing import TYPE_CHECKING, Any, TypeVar + +import typing_extensions + +from serena.config.serena_config import ApiInclusionDefinition +from serena.project import Project + +from .representable import RepresentableViaRenderer + +if TYPE_CHECKING: + from serena.agent import SerenaAgent + from serena.code_editor import CodeEditor + from serena.tools import Tool + + from .external_project import ExternalProjectExecution + +log = logging.getLogger(__name__) +TCallable = TypeVar("TCallable", bound=Callable[..., Any]) + +SUCCESS_RESULT = "OK" +"""the result returned by operations which have no result other than their success""" + + +def format_annotation(annotation: Any) -> str: + """ + :param annotation: a type annotation (or a signature/annotation string) + :return: the annotation rendered without module paths (e.g. `list[LanguageServerSymbol]`), such that type names + match the names by which the types can be looked up + """ + text = annotation if isinstance(annotation, str) else inspect.formatannotation(annotation) + return re.sub(r"\b(?:[A-Za-z_]\w*\.)+([A-Za-z_]\w*)", r"\1", text) + + +def format_signature(callable_: Callable[..., Any]) -> str: + """ + :param callable_: the callable + :return: the signature rendered without module paths in annotations + """ + return format_annotation(str(inspect.signature(callable_))) + + +def extract_referenced_classes(annotation: Any) -> list[type]: + """ + :param annotation: a (resolved) type annotation + :return: the user-defined classes appearing in the annotation (recursively, e.g. in `list[X] | None`), in order of + appearance; builtins, typing constructs and classes from the standard library are excluded + """ + classes: list[type] = [] + + def visit(a: Any) -> None: + if isinstance(a, type): + module = getattr(a, "__module__", "") + if module not in ("builtins", "typing", "collections.abc", "abc") and not module.startswith("_") and a not in classes: + classes.append(a) + for arg in typing.get_args(a): + visit(arg) + + visit(annotation) + return classes + + +def get_annotated_classes(callable_: Callable[..., Any]) -> list[type]: + """ + :param callable_: a function or method + :return: the user-defined classes appearing in the annotations of its parameters and return type + """ + try: + hints = typing.get_type_hints(callable_) + except Exception: # unresolvable forward references + return [] + classes: list[type] = [] + for hint in hints.values(): + classes.extend(c for c in extract_referenced_classes(hint) if c not in classes) + return classes + + +@dataclass +class ReferencedType: + """ + A type that is referenced by a facade's methods (returned by them, contained in their results or used as a parameter + type), whose interface the LLM can inspect via `info`. + All types reachable through annotations are discovered automatically; an explicit declaration is needed only in order + to curate the type's presentation (the members to describe, inclusion in the facade's description). + """ + + cls: type + """the type""" + provide_info_with_facade: bool = False + """whether the type's full description is included in the facade's description (rather than just its name)""" + members: Sequence[str] | None = None + """ + the LLM-facing members (attributes, properties, methods) to describe; if None, all members admitted by the naming + convention (no leading or trailing underscore) which are documented are described. Explicitly listed methods are + described even if undocumented, as the listing is the documentation decision. + """ + + _CAPABILITIES: typing.ClassVar[dict[str, str]] = { + "__len__": "len()", + "__iter__": "iteration", + "__getitem__": "indexing", + "__enter__": "use in a `with` statement", + } + + @property + def name(self) -> str: + return self.cls.__name__ + + def _get_member_names(self) -> list[str]: + if self.members is not None: + return list(self.members) + names = set(typing.get_type_hints(self.cls)) + names.update(n for n in dir(self.cls) if not n.startswith("_")) + # exclude the representation mechanism, which is not meant to be used from REPL code + names.difference_update(dir(RepresentableViaRenderer)) + return sorted(n for n in names if not n.startswith("_") and not n.endswith("_")) + + def get_referenced_classes(self) -> list[type]: + """ + :return: the user-defined classes appearing in the annotations of the described members (attributes, properties, + method parameters and return types), in order of appearance (each class at most once, excluding the type itself) + """ + if self.is_enum(): + return [] + classes: list[type] = [] + type_hints = typing.get_type_hints(self.cls) + if self.is_typed_dict(): + for hint in type_hints.values(): + classes.extend(c for c in extract_referenced_classes(hint) if c is not self.cls and c not in classes) + return classes + for member_name in self._get_member_names(): + member = inspect.getattr_static(self.cls, member_name, None) + if isinstance(member, property) and member.fget is not None: + found = get_annotated_classes(member.fget) + elif inspect.isfunction(member): + found = get_annotated_classes(member) + elif member_name in type_hints: + found = extract_referenced_classes(type_hints[member_name]) + else: + found = [] + classes.extend(c for c in found if c is not self.cls and c not in classes) + return classes + + def is_enum(self) -> bool: + return isinstance(self.cls, type) and issubclass(self.cls, Enum) + + def is_typed_dict(self) -> bool: + # NOTE: TypedDicts defined via typing_extensions are not recognised by typing.is_typeddict + return typing.is_typeddict(self.cls) or typing_extensions.is_typeddict(self.cls) + + def _describe_typed_dict(self) -> str: + parts = [f"type {self.name} (a dict with the following keys)"] + if self.cls.__doc__ and not self.cls.__doc__.startswith(self.name + "("): # NOTE: the default docstring is uninformative + parts.append(f" {inspect.cleandoc(self.cls.__doc__).replace(chr(10), chr(10) + ' ')}") + type_hints = typing.get_type_hints(self.cls) + member_names = self.members if self.members is not None else list(type_hints) + parts.append( + "keys:\n" + "\n".join(f" {name}: {format_annotation(type_hints[name])}" for name in member_names if name in type_hints) + ) + return "\n".join(parts) + "\n" + + def _describe_enum(self) -> str: + parts = [f"enum {self.name}"] + if self.cls.__doc__: + parts.append(f" {inspect.cleandoc(self.cls.__doc__).replace(chr(10), chr(10) + ' ')}") + parts.append("members:\n" + "\n".join(f" {self.name}.{member.name} = {member.value!r}" for member in self.cls)) # type: ignore[attr-defined] + return "\n".join(parts) + "\n" + + @staticmethod + def _first_doc_line(obj: Any) -> str: + doc = inspect.getdoc(obj) or "" + first_line = doc.splitlines()[0] if doc else "" + return first_line.removeprefix(":return:").strip() + + def describe(self) -> str: + """ + :return: the type's documentation: its docstring, attributes/properties with their types and methods with their + signatures and documentation; for enums, the members with their values + """ + if self.is_enum(): + return self._describe_enum() + if self.is_typed_dict(): + return self._describe_typed_dict() + attributes: list[str] = [] + methods: list[str] = [] + type_hints = typing.get_type_hints(self.cls) + for member_name in self._get_member_names(): + member = inspect.getattr_static(self.cls, member_name, None) + if isinstance(member, property): + fget = member.fget + annotation = inspect.signature(fget).return_annotation if fget is not None else inspect.Signature.empty + type_str = f": {format_annotation(annotation)}" if annotation is not inspect.Signature.empty else "" + doc = self._first_doc_line(member) + attributes.append(f" {member_name}{type_str}" + (f" # {doc}" if doc else "")) + elif inspect.isfunction(member): + doc = inspect.getdoc(member) + if doc is None and self.members is None: + continue + signature = format_signature(member).replace("(self, ", "(", 1).replace("(self)", "()", 1) + methods.append(f" {member_name}{signature}" + (f"\n {doc.replace(chr(10), chr(10) + ' ')}" if doc else "")) + elif member_name in type_hints: + attributes.append(f" {member_name}: {format_annotation(type_hints[member_name])}") + else: + attributes.append(f" {member_name}") + + # assemble the description + parts = [f"type {self.name}"] + if self.cls.__doc__: # NOTE: the class' own docstring (inspect.getdoc would fall back to base class docstrings) + doc = inspect.cleandoc(self.cls.__doc__) + parts.append(f" {doc.replace(chr(10), chr(10) + ' ')}") + if attributes: + parts.append("attributes:\n" + "\n".join(attributes)) + if methods: + parts.append("methods:\n" + "\n".join(methods)) + capabilities = [text for dunder, text in self._CAPABILITIES.items() if dunder in dir(self.cls) and dunder not in dir(object)] + if capabilities: + parts.append("supports: " + ", ".join(capabilities)) + return "\n".join(parts) + "\n" + + +@dataclass(kw_only=True, frozen=True) +class FacadeMethodInfo: + """ + The metadata of a method exposed through a facade (see `facade_method`), mirroring the tool markers. + """ + + name: str + """the name of the method""" + optional: bool = False + """whether the method is disabled by default and must be enabled explicitly""" + beta: bool = False + """whether the method is in beta (not yet fully stable)""" + can_edit: bool = False + """whether the method can modify the codebase (relevant for read-only contexts)""" + niche: bool = False + """ + whether the method is rarely needed, such that the facade's description only summarises it (first line of its + documentation and a pointer to its full documentation) in order to keep the description compact + """ + uses_project_server: bool = False + """ + whether the method requires the project's language servers and must therefore be executed in the project server + when an external project is queried (see `ExternalProjectContext`). + Edit operations are always executed in the project server, regardless of this flag, since they implicitly + use the CodeEditor, which requires language servers when using the LSP backend. + Polymorphic edit operations therefore must not set this flag to True. + """ + corresponding_tool: "type[Tool] | None" = None + """the classic tool offering the same functionality, if any""" + + def get_corresponding_tool_name(self) -> str | None: + """ + :return: the name of the corresponding tool, or None if there is none + """ + return self.corresponding_tool.get_name_from_cls() if self.corresponding_tool is not None else None + + +_FACADE_METHOD_INFO_ATTR = "__facade_method_info__" + + +def facade_method( + *, + optional: bool = False, + beta: bool = False, + can_edit: bool = False, + niche: bool = False, + uses_project_server: bool = False, + corresponding_tool: "type[Tool] | None" = None, +) -> Callable[[TCallable], TCallable]: + """ + Marks a method of a `FacadeApi` as exposed through the facade, attaching the given metadata. + The decorator only annotates the method (it does not wrap it), such that signature and docstring remain intact. + + :param optional: whether the method is disabled by default and must be enabled explicitly + :param beta: whether the method is in beta + :param can_edit: whether the method can modify the codebase + :param niche: whether the method is rarely needed (its documentation is then only summarised in the facade's description) + :param uses_project_server: whether the method must be executed remotely in the project server when an external project is queried + :param corresponding_tool: the classic tool offering the same functionality, if any + :return: the decorator + """ + + def decorator(method: TCallable) -> TCallable: + info = FacadeMethodInfo( + name=method.__name__, + optional=optional, + beta=beta, + can_edit=can_edit, + niche=niche, + uses_project_server=uses_project_server, + corresponding_tool=corresponding_tool, + ) + setattr(method, _FACADE_METHOD_INFO_ATTR, info) + return method + + return decorator + + +def get_facade_method_info(method: Callable[..., Any]) -> FacadeMethodInfo | None: + """ + :param method: a (bound or unbound) method + :return: the metadata attached via `facade_method`, or None if the method is not exposed + """ + return getattr(method, _FACADE_METHOD_INFO_ATTR, None) + + +class FacadeApi(ABC): + """ + The implementation of a facade's functionality. + + API design principles: + + * A method is exposed to the LLM if and only if it is decorated with `facade_method`, which also carries + the method's metadata (optional, beta, can_edit). Undecorated methods are never exposed, regardless of their name. + * On the objects returned by API methods (which are not decorated), the name determines visibility: + names with a trailing underscore (e.g. `symbols_`, `to_dict_`) are public within Serena (e.g. for use by + classic tools or other facade implementations) but are not meant to be called from REPL code, whereas + names without leading or trailing underscore constitute the LLM-facing interface. + The same convention applies to non-exposed helper methods of API classes. + * Names with a leading underscore are private, as usual. + """ + + def __init__(self, agent: "SerenaAgent", name: str, description: str, types: Sequence[ReferencedType] = ()) -> None: + """ + :param agent: the agent providing access to the project and its resources + :param name: the attribute name under which the facade is accessible from the REPL entrypoint + :param description: a one-line description of the functionality offered by the facade + :param types: declarations for referenced types whose presentation is to be curated (see `ReferencedType`); + types reachable through annotations need not be declared in order to be documentable + """ + self._agent = agent + self._name = name + self._description = description + self._types = list(types) + + def get_name_(self) -> str: + return self._name + + def get_description_(self) -> str: + return self._description + + def get_referenced_types_(self) -> list[ReferencedType]: + return self._types + + def _get_project(self) -> Project: + return self._agent.get_active_project_or_raise() + + def _create_code_editor(self) -> "CodeEditor": + """ + :return: a code editor for the active project, using the active language backend + """ + project = self._get_project() + backend = self._agent.get_language_backend() + return backend.create_code_editor(project) + + +class FacadeMethod: + """ + A method of a facade, which delegates to a method of the underlying implementation and which can be + enabled or disabled; only enabled methods are accessible from REPL code. + """ + + def __init__(self, parent: "Facade", implementation: Callable[..., Any], info: FacadeMethodInfo, enabled: bool) -> None: + """ + :param parent: the facade the method belongs to + :param implementation: the implementation to delegate to + :param info: the method's metadata (including the method's name) + :param enabled: whether the method is initially enabled + """ + self.parent = parent + self._implementation = implementation + self.info = info + self.enabled = enabled + + @property + def name(self) -> str: + return self.info.name + + @property + def facade_name(self) -> str: + return self.parent.name + + @property + def qualified_name(self) -> str: + """ + :return: the name under which the method is accessible from REPL code (facade name and method name) + """ + return f"{self.facade_name}.{self.name}" + + def __call__(self, *args: Any, **kwargs: Any) -> Any: + external_project_execution = self.parent.get_external_project_() + if external_project_execution is not None: + external_project_execution.check_call_permission(self) + if external_project_execution.is_called_remotely(self): + return external_project_execution.call_remotely(self.facade_name, self.name, args, kwargs) + return self._implementation(*args, **kwargs) + + def get_implementation_(self) -> Callable[..., Any]: + return self._implementation + + def get_referenced_return_types(self) -> list[ReferencedType]: + """ + :return: the types referenced by the facade which appear in the method's return type annotation + """ + return_annotation = format_annotation(inspect.signature(self._implementation).return_annotation) + return self._find_referenced_types(return_annotation) + + def get_referenced_parameter_types(self) -> list[ReferencedType]: + """ + :return: the types referenced by the facade which appear in the annotations of the method's parameters + """ + parameters = inspect.signature(self._implementation).parameters.values() + return self._find_referenced_types(" ".join(format_annotation(p.annotation) for p in parameters)) + + def _find_referenced_types(self, annotation_text: str) -> list[ReferencedType]: + return [t for t in self.parent.get_types() if re.search(rf"\b{re.escape(t.name)}\b", annotation_text)] + + def get_summary(self) -> str: + """ + :return: the first line of the method's documentation + """ + doc = inspect.getdoc(self._implementation) + return doc.splitlines()[0] if doc else "(no documentation)" + + def describe(self) -> str: + """ + :return: the method's signature and documentation, with pointers to the documentation of the referenced types + appearing in its return type and parameter annotations + """ + signature = format_signature(self._implementation) + doc = inspect.getdoc(self._implementation) or "(no documentation)" + text = f"{self.qualified_name}{signature}\n{doc}\n" + referenced = self.get_referenced_return_types() + referenced += [t for t in self.get_referenced_parameter_types() if t not in referenced] + if referenced: + pointers = ", ".join(f'`s.info("{t.name}")`' for t in referenced) + text += f"Type documentation: {pointers}\n" + return text + + def describe_summary(self) -> str: + """ + :return: the method's name and the first line of its documentation, with a pointer to its full documentation + """ + return f'{self.qualified_name}: {self.get_summary()} [full documentation: `s.info("{self.qualified_name}")`]\n' + + +class ApiScope: + """ + The scope of APIs available to the LLM, i.e. which facade methods are enabled, as determined by + applying a sequence of inclusion/exclusion definitions (from the global configuration, the context, + the active modes and the project configuration) to the methods' default enablement. + """ + + class FacadeScope: + """ + The scope of a single facade: whether the facade as a whole is included, and which of its methods + were explicitly included/excluded (a method is never in both sets). + If the facade is not included, it is opt-in, i.e. only explicitly included methods are enabled. + """ + + def __init__(self) -> None: + self._is_included: bool | None = None + self.method_inclusions: set[str] = set() + self.method_exclusions: set[str] = set() + + def exclude_facade(self) -> None: + self._is_included = False + self.method_inclusions = set() + self.method_exclusions = set() + + def include_facade(self) -> None: + self._is_included = True + + def is_facade_included(self, is_facade_optional: bool) -> bool: + if is_facade_optional: + return self._is_included is True + else: + return self._is_included is not False + + def exclude_method(self, method_name: str) -> None: + self.method_inclusions.discard(method_name) + self.method_exclusions.add(method_name) + + def include_method(self, method_name: str) -> None: + self.method_exclusions.discard(method_name) + self.method_inclusions.add(method_name) + + def __init__(self) -> None: + self._facade_scopes: dict[str, ApiScope.FacadeScope] = {} + self._editing_excluded = False + + def _get_facade_scope(self, facade_name: str) -> "ApiScope.FacadeScope": + if facade_name not in self._facade_scopes: + self._facade_scopes[facade_name] = ApiScope.FacadeScope() + return self._facade_scopes[facade_name] + + def process(self, definition: ApiInclusionDefinition) -> None: + """ + Applies the given definition, exclusions first, then inclusions (such that inclusions take precedence + within a definition; across definitions, later definitions take precedence). + + :param definition: the definition to apply + """ + + def apply(api_ref: str, *, excluded: bool) -> None: + components = api_ref.split(".") + if len(components) > 2: + log.warning("Ignoring invalid API reference '%s' in %s (expected 'facade' or 'facade.method')", api_ref, definition) + return + facade_scope = self._get_facade_scope(components[0]) + if len(components) == 1: + facade_scope.exclude_facade() if excluded else facade_scope.include_facade() + else: + facade_scope.exclude_method(components[1]) if excluded else facade_scope.include_method(components[1]) + + for api_exclusion in definition.excluded_apis: + apply(api_exclusion, excluded=True) + for api_inclusion in definition.included_apis: + apply(api_inclusion, excluded=False) + + def exclude_editing(self) -> None: + """ + Excludes all methods which can edit the codebase (read-only operation), regardless of other inclusions. + """ + self._editing_excluded = True + + def is_method_enabled(self, facade_name: str, method_info: FacadeMethodInfo, is_facade_optional: bool) -> bool: + """ + :param facade_name: the name of the facade + :param method_info: the method's metadata + :param is_facade_optional: whether the facade is optional (disabled by default and must be enabled explicitly) + :return: whether the method is enabled: optional methods (and all methods of a facade which is not included, + i.e. an excluded facade or an optional facade that was not explicitly included) must be explicitly + included, other methods are enabled unless explicitly excluded; if editing is excluded, editing + methods are always disabled + """ + if self._editing_excluded and method_info.can_edit: + return False + facade_scope = self._get_facade_scope(facade_name) + # A method that would be disabled because the facade it is part of is not included + # or the method itself is optional must be explicitly included in order to be enabled. + if not facade_scope.is_facade_included(is_facade_optional) or method_info.optional: + return method_info.name in facade_scope.method_inclusions + # A method that is not optional and whose facade is included is enabled unless it is explicitly excluded. + else: + return method_info.name not in facade_scope.method_exclusions + + +class Facade: + """ + A named group of related operations which an LLM can invoke from REPL code. + """ + + def __init__(self, name: str, description: str, is_optional: bool = False, types: Sequence[ReferencedType] = ()) -> None: + # NOTE: attributes are set via object.__setattr__ because __getattr__ is overridden + object.__setattr__(self, "_is_optional", is_optional) + object.__setattr__(self, "_name", name) + object.__setattr__(self, "_description", description) + object.__setattr__(self, "_methods", {}) + object.__setattr__(self, "_types", {t.name: t for t in types}) + object.__setattr__(self, "_external_project", None) + + def set_external_project_(self, external_project: "ExternalProjectExecution | None") -> None: + """ + :param external_project: the context of the external project being queried (None if the active project is used) + """ + object.__setattr__(self, "_external_project", external_project) + + def get_external_project_(self) -> "ExternalProjectExecution | None": + return self._external_project + + def _add_method(self, method: FacadeMethod) -> None: + assert method.parent is self + self._methods[method.name] = method + + @staticmethod + def from_api(api: FacadeApi, api_scope: ApiScope, *, is_optional: bool = False) -> "Facade": + """ + Creates a facade wrapping the given implementation. + + :param api: the implementation; each of its methods decorated with `facade_method` becomes a facade method + :param api_scope: API scope definition determining which methods are enabled + :param is_optional: whether the facade is optional (disabled by default and must be enabled explicitly) + :return: the facade + """ + facade = Facade(api.get_name_(), api.get_description_(), is_optional, api.get_referenced_types_()) + for name, member in inspect.getmembers(api, predicate=inspect.ismethod): + method_info = get_facade_method_info(member) + if method_info is None: + continue + is_enabled = api_scope.is_method_enabled(facade.name, method_info, is_facade_optional=is_optional) + facade._add_method(FacadeMethod(facade, member, method_info, enabled=is_enabled)) + facade._discover_referenced_types() + return facade + + def is_enabled(self) -> bool: + """ + :return: whether the facade is enabled + """ + return len(self.get_enabled_methods()) > 0 + + def is_optional(self) -> bool: + """ + :return: whether the facade is optional, i.e. disabled unless it is explicitly included + """ + return self._is_optional + + def _discover_referenced_types(self) -> None: + """ + Adds referenced types for all classes reachable (transitively) through the annotations of the facade's methods + and of the referenced types' members, such that every type an LLM may encounter can be documented. + Explicitly declared types take precedence (they may curate members and carry flags). + """ + # seed the worklist with the classes referenced by the declared types and by the methods + pending: list[type] = [] + for referenced_type in self._types.values(): + pending.extend(referenced_type.get_referenced_classes()) + for method in self._methods.values(): + pending.extend(get_annotated_classes(method.get_implementation_())) + + # add undeclared classes, following their references in turn + while pending: + cls = pending.pop(0) + if cls.__name__ in self._types: + continue + referenced_type = ReferencedType(cls) + self._types[cls.__name__] = referenced_type + pending.extend(referenced_type.get_referenced_classes()) + + @property + def name(self) -> str: + return self._name + + @property + def description(self) -> str: + return self._description + + @property + def enabled_method_names(self) -> list[str]: + return [m.name for m in self._methods.values() if m.enabled] + + def get_enabled_methods(self) -> list[FacadeMethod]: + """ + :return: the list of enabled methods + """ + return [m for m in self._methods.values() if m.enabled] + + def get_methods(self) -> list[FacadeMethod]: + """ + :return: the list of all methods, regardless of whether they are enabled (e.g. for changing their enabled state) + """ + return list(self._methods.values()) + + def get_types(self) -> list[ReferencedType]: + """ + :return: the types referenced by the facade's methods + """ + return list(self._types.values()) + + def get_type(self, type_name: str) -> ReferencedType | None: + """ + :param type_name: the name of the type + :return: the referenced type, or None if the facade does not reference a type of that name + """ + return self._types.get(type_name) + + def get_method(self, method_name: str) -> FacadeMethod: + """ + :param method_name: the name of the method + :return: the method, regardless of whether it is enabled (e.g. for changing its enabled state) + """ + if method_name not in self._methods: + raise ValueError(f"Facade '{self._name}' has no method '{method_name}'") + return self._methods[method_name] + + def _get_enabled_method(self, name: str) -> FacadeMethod | None: + method = self._methods.get(name) + return method if method is not None and method.enabled else None + + def _no_such_method_message(self, name: str) -> str: + return f"Facade '{self._name}' has no method '{name}'. Available methods: {self.enabled_method_names}" + + def __getattr__(self, name: str) -> Any: + # delegate attribute access to enabled methods only (called only if regular attribute lookup fails) + method = self._get_enabled_method(name) + if method is None: + raise AttributeError(self._no_such_method_message(name)) + return method + + def __setattr__(self, name: str, value: Any) -> None: + raise AttributeError(f"Facade '{self._name}' is read-only") + + def describe(self) -> str: + """ + :return: a description of the facade listing all of its enabled methods with their signatures and documentation + as well as its referenced types (in full if so declared, otherwise by name) + """ + parts = [f"Facade '{self._name}': {self._description}", ""] + enabled_methods = self.get_enabled_methods() + for method in enabled_methods: + if not method.info.niche: + parts.append(method.describe()) + niche_methods = [m for m in enabled_methods if m.info.niche] + if niche_methods: + parts.append("Rarely needed methods (documented on request):\n" + "".join(m.describe_summary() for m in niche_methods)) + for referenced_type in self._types.values(): + if referenced_type.provide_info_with_facade: + parts.append(referenced_type.describe()) + result_type_names = sorted( + {t.name for m in enabled_methods for t in m.get_referenced_return_types() if not t.provide_info_with_facade} + ) + if result_type_names: + parts.append( + "Result types: " + + ", ".join(result_type_names) + + ' (request documentation via `s.info("")` only if you intend to process results in code)' + ) + return "\n".join(parts) + + def describe_member(self, member_name: str) -> str: + """ + :param member_name: the name of one of the facade's enabled methods or referenced types + :return: the member's documentation + """ + method = self._get_enabled_method(member_name) + if method is not None: + return method.describe() + referenced_type = self._types.get(member_name) + if referenced_type is not None: + return referenced_type.describe() + raise ValueError( + f"Facade '{self._name}' has no method or type '{member_name}'. " + f"Available methods: {self.enabled_method_names}; types: {list(self._types)}" + ) diff --git a/src/serena/repl/repl.py b/src/serena/repl/repl.py new file mode 100644 index 00000000..9dc1b7b1 --- /dev/null +++ b/src/serena/repl/repl.py @@ -0,0 +1,384 @@ +""" +The REPL through which an LLM executes Python code against Serena's facades. +""" + +# SPDX-License-Identifier: GPL-3.0-or-later + +import ast +import logging +import re +import traceback +from dataclasses import dataclass +from typing import TYPE_CHECKING, Any + +from ..session import SerenaSession +from .external_project import ExternalProjectExecution +from .facade import ApiScope, Facade, FacadeMethod, ReferencedType +from .representable import Representable + +if TYPE_CHECKING: + from ..tools.tools_base import Tool + +log = logging.getLogger(__name__) + + +class FacadeAvailabilityInfo: + """ + Represents information on the availability of facades and the methods therein + """ + + @dataclass + class FacadeInfo: + name: str + is_enabled: bool + methods: list["FacadeAvailabilityInfo.MethodInfo"] + + @dataclass + class MethodInfo: + name: str + is_enabled: bool + + def __init__(self): + self.facades: list[FacadeAvailabilityInfo.FacadeInfo] = [] + + def add_facade(self, facade: Facade): + methods_info = [FacadeAvailabilityInfo.MethodInfo(name=method.name, is_enabled=method.enabled) for method in facade.get_methods()] + self.facades.append(FacadeAvailabilityInfo.FacadeInfo(name=facade.name, is_enabled=facade.is_enabled(), methods=methods_info)) + + +class SerenaReplEntrypoint: + """ + Represents the entrypoint object for the REPL. It holds the configured facades as attributes + and offers progressive disclosure of their interfaces via `info`. + """ + + def __init__(self, facades: list[Facade], api_scope: ApiScope) -> None: + """ + :param facades: the candidate facades + :param api_scope: the API scope, which determines which of the facades are made available + """ + self._facades: dict[str, Facade] = {} + self._current_session: SerenaSession | None = None + self._current_namespace: dict[str, Any] | None = None + self._facade_availability_info = FacadeAvailabilityInfo() + registered_facade_names = [] + for facade in facades: + self._facade_availability_info.add_facade(facade) + if facade.is_enabled(): + if facade.name in self._facades: + raise ValueError(f"Duplicate facade name: {facade.name}") + self._facades[facade.name] = facade + setattr(self, facade.name, facade) + registered_facade_names.append(facade.name) + log.info("Registered %d/%d facades: %s", len(registered_facade_names), len(facades), registered_facade_names) + + def get_facade_availability_info(self) -> FacadeAvailabilityInfo: + """ + :return: the availability of all facades and their methods (enabled or disabled) + """ + return self._facade_availability_info + + def get_enabled_methods(self) -> list[FacadeMethod]: + """ + :return: the list of all enabled methods across all facades + """ + return [method for facade in self._facades.values() for method in facade.get_enabled_methods()] + + def is_tool_function_available(self, tool_class: "type[Tool]") -> bool: + """ + Checks whether any of the enabled methods corresponds to the given tool class. + + :param tool_class: the tool class to check for + :return: whether any enabled method corresponds to the given tool class + """ + for method in self.get_enabled_methods(): + if method.info.corresponding_tool == tool_class: + return True + return False + + def set_external_project_(self, external_project: "ExternalProjectExecution | None") -> None: + """ + :param external_project: the context of the external project being queried by the currently executing code + (None if the active project is used); propagated to all facades + """ + for facade in self._facades.values(): + facade.set_external_project_(external_project) + + def get_external_project_(self) -> "ExternalProjectExecution | None": + external_projects = {facade.get_external_project_() for facade in self._facades.values()} + return next(iter(external_projects)) if external_projects else None + + def set_current_session_(self, session: SerenaSession | None, namespace: dict[str, Any] | None) -> None: + """ + :param session: the session on whose behalf code is being executed (None if no code is being executed) + :param namespace: the namespace of the execution (None if no code is being executed) + """ + self._current_session = session + self._current_namespace = namespace + + def _get_persisted_items(self) -> dict[str, Any]: + assert self._current_namespace is not None, "No code execution in progress" + return { + name: value + for name, value in self._current_namespace.items() + if name != SerenaRepl.ENTRYPOINT_NAME and SerenaRepl.is_persisted_name(name) + } + + def vars(self) -> str: + """ + Lists the variables and functions which persist in the session's namespace across executions. + + :return: the listing (name, type and a short representation per item) + """ + items = self._get_persisted_items() + if not items: + return "No persisted variables." + lines = [] + for name, value in items.items(): + summary = value.__name__ if callable(value) and hasattr(value, "__name__") else repr(value) + if len(summary) > 80: + summary = summary[:77] + "..." + lines.append(f"{name}: {type(value).__name__} = {summary}") + return "\n".join(lines) + + def clear(self) -> str: + """ + Removes all persisted variables and functions from the session's namespace. + + :return: a message indicating the number of removed items + """ + items = self._get_persisted_items() + assert self._current_namespace is not None + for name in items: + del self._current_namespace[name] + return f"Removed {len(items)} persisted item(s)." + + def _get_facade(self, name: str) -> Facade: + if name not in self._facades: + raise ValueError(f"Unknown facade '{name}'. Available facades: {list(self._facades)}") + return self._facades[name] + + def get_facade_(self, name: str) -> Facade: + """ + :param name: the facade's name + :return: the facade + """ + return self._get_facade(name) + + def overview(self) -> str: + """ + :return: the list of available facades, each with a one-line description and the names of its methods + (with the result type of methods returning objects that can be processed in code) + """ + + def method_entry(method: FacadeMethod) -> str: + return_types = method.get_referenced_return_types() + return method.name + (f" -> {'|'.join(t.name for t in return_types)}" if return_types else "") + + return "\n".join( + f"s.{facade.name}: {facade.description}\n methods: {', '.join(method_entry(m) for m in facade.get_enabled_methods())}" + for facade in self._facades.values() + ) + + def info(self, *items: str) -> str: + """ + Provides documentation on the available functionality. + + :param items: the items to document; if none are given, an overview of all facades is provided. + Each item is either a facade name (e.g. "lsp") for the documentation of all of the facade's methods and types, + a dotted path (e.g. "lsp.find_symbol" or "lsp.LspSymbolCollection") for the documentation of a single method + or type, or a bare type name (e.g. "LanguageServerSymbol"), which is looked up across all facades. + Unknown items are reported without affecting the documentation of the other items. + :return: the requested documentation + """ + if not items: + return self.overview() + described_in_call: set[str] = set() + return "\n\n".join(self._describe_item(item, described_in_call) for item in items) + + def _describe_item(self, item: str, described_in_call: set[str]) -> str: + facade_name, _, member_name = item.partition(".") + try: + if member_name: + facade = self._get_facade(facade_name) + referenced_type = facade.get_type(member_name) + if referenced_type is not None: + return self._describe_type(referenced_type, described_in_call) + return facade.describe_member(member_name) + if facade_name in self._facades: + return self._facades[facade_name].describe() + # not a facade: look up the item as a type across all facades + referenced_type = self._find_type(item) + if referenced_type is None: + raise ValueError(f"Unknown item '{item}': neither a facade nor a type. Available facades: {list(self._facades)}") + return self._describe_type(referenced_type, described_in_call) + except ValueError as e: + return str(e) + + def _find_type(self, type_name: str) -> ReferencedType | None: + for facade in self._facades.values(): + referenced_type = facade.get_type(type_name) + if referenced_type is not None: + return referenced_type + return None + + def _describe_type(self, referenced_type: ReferencedType, described_in_call: set[str]) -> str: + """ + Describes the given (explicitly requested) type along with the types it references (transitively), each at most + once per call. Referenced types whose documentation was already provided earlier in the session are not repeated + but pointed to (an explicit request always yields the full documentation). + + :param referenced_type: the requested type + :param described_in_call: the names of the types already described in the current `info` call (updated) + :return: the documentation + """ + session = self._current_session + parts = [] + if referenced_type.name not in described_in_call: + parts.append(referenced_type.describe()) + described_in_call.add(referenced_type.name) + if session is not None: + session.described_type_names.add(referenced_type.name) + + # append the referenced types (breadth-first), unless already described in this call or earlier in the session + pending = [cls.__name__ for cls in referenced_type.get_referenced_classes()] + while pending: + type_name = pending.pop(0) + if type_name in described_in_call: + continue + contained_type = self._find_type(type_name) + if contained_type is None: + continue + described_in_call.add(type_name) + if session is not None and type_name in session.described_type_names: + parts.append(f'type {type_name}: documented earlier in this session (request `s.info("{type_name}")` to see it again)\n') + else: + parts.append(contained_type.describe()) + if session is not None: + session.described_type_names.add(type_name) + pending.extend(cls.__name__ for cls in contained_type.get_referenced_classes()) + return "\n".join(parts) + + +class SerenaRepl: + """ + Executes Python code submitted by an LLM, binding the configured facades to the entrypoint object `s` + and rendering the result of the execution as a string for the LLM. + + The code is executed like the cell of a notebook: it is executed at module level in the session's namespace, + such that the names it binds (variables, functions, classes, imports) persist across executions, and if its last + statement is an expression, the expression's value is the result of the execution. + The entrypoint `s` is (re)bound in the namespace before every execution, such that persisted functions always + access the current entrypoint. + """ + + SOURCE_NAME = "" + ENTRYPOINT_NAME = "s" + _PERSISTED_NAME_PATTERN = re.compile(r"^(?!__)[A-Za-z_]\w*$") + + @classmethod + def is_persisted_name(cls, name: str) -> bool: + """ + :param name: a name in a session namespace + :return: whether the name denotes a persisted item of the LLM's (as opposed to an implementation detail + such as `__builtins__`) + """ + return cls._PERSISTED_NAME_PATTERN.match(name) is not None + + def __init__(self, facades: list[Facade], api_scope: ApiScope) -> None: + """ + :param facades: the candidate facades + :param api_scope: the API scope, which determines which of the facades are made available + """ + self._entrypoint = SerenaReplEntrypoint(facades, api_scope) + + @property + def entrypoint(self) -> SerenaReplEntrypoint: + return self._entrypoint + + @classmethod + def _represent(cls, obj: Any) -> str: + """ + Renders an arbitrary object as a string for the LLM. Representables render themselves, + lists and tuples are rendered element-wise (one element per line), everything else via `str`. + + :param obj: the object to render + :return: the textual representation + """ + if isinstance(obj, Representable): + return obj.represent() + if isinstance(obj, list | tuple): + if len(obj) == 0: + return "[]" + return "\n".join(cls._represent(item) for item in obj) + return str(obj) + + def execute(self, code: str, session: SerenaSession | None = None) -> str: + """ + Executes the given code and renders its result. + Executions are expected to be serialised (the entrypoint holds the current session during execution). + + :param code: the Python code to execute + :param session: the client session on whose behalf the code is executed (None for session-less execution, + e.g. in tests), which determines e.g. which type documentation has already been provided + :return: the representation of the code's result, or a description of the error if execution failed + """ + namespace = session.repl_namespace if session is not None else {} + self._entrypoint.set_current_session_(session, namespace) + try: + result = self._run(code, namespace) + except Exception as e: + return self._format_error(e, code) + finally: + self._entrypoint.set_current_session_(None, None) + return self._represent(result) + + def _run(self, code: str, namespace: dict[str, Any]) -> Any: + """ + Runs the given code at module level in the given namespace (as globals) with the entrypoint bound. + + :return: the value of the code's last statement if it is an expression, None otherwise + """ + namespace[self.ENTRYPOINT_NAME] = self._entrypoint + module = ast.parse(code, self.SOURCE_NAME) + + # separate a trailing expression, whose value is the result + statements = module.body + trailing_expression: ast.expr | None = None + if statements and isinstance(statements[-1], ast.Expr): + trailing_expression = statements[-1].value + statements = statements[:-1] + + # execute the statements, then evaluate the trailing expression (both retain the original line numbers) + if statements: + exec(compile(ast.Module(body=statements, type_ignores=[]), self.SOURCE_NAME, "exec"), namespace) + if trailing_expression is not None: + return eval(compile(ast.Expression(body=trailing_expression), self.SOURCE_NAME, "eval"), namespace) + return None + + def _format_error(self, e: Exception, code: str) -> str: + """ + :param e: the exception raised during execution + :param code: the code that was executed + :return: an error message which locates the failure within the executed code + """ + code_lines = code.splitlines() + + def location_line(line_number: int) -> str: + line_text = code_lines[line_number - 1].strip() if 0 < line_number <= len(code_lines) else "" + return f" line {line_number}: {line_text}" + + # report syntax errors in the executed code (which carry no traceback frames of their own) + if isinstance(e, SyntaxError) and e.filename == self.SOURCE_NAME and e.lineno is not None: + message = f"SyntaxError: {e.msg}\n" + location_line(e.lineno) + if "return" in (e.msg or "") and "outside function" in e.msg: + message += "\nNote: the code is executed like a notebook cell; the value of the last expression is the result (do not use `return`)." + return message + + # report runtime errors, locating them within the executed code + location_lines = [] + for frame in traceback.extract_tb(e.__traceback__): + if frame.filename != self.SOURCE_NAME or frame.lineno is None: + continue + location_lines.append(location_line(frame.lineno)) + return "\n".join([f"{type(e).__name__}: {e}", *location_lines]) diff --git a/src/serena/repl/representable.py b/src/serena/repl/representable.py new file mode 100644 index 00000000..069d712c --- /dev/null +++ b/src/serena/repl/representable.py @@ -0,0 +1,106 @@ +""" +The representation protocol through which objects returned from REPL code are rendered for the LLM. +""" + +# SPDX-License-Identifier: GPL-3.0-or-later + +from abc import ABC, abstractmethod +from collections.abc import Callable +from typing import TYPE_CHECKING, Any, Generic, TypeVar + +from serena.util.text_utils import TextOutputUtils + +if TYPE_CHECKING: + from serena.agent import SerenaAgent + + +T = TypeVar("T") + + +class Renderer(Generic[T], ABC): + def __init__(self, agent: "SerenaAgent", max_answer_chars: int = -1): + """ + :param agent: the agent, from which the configured default length limit is taken (the agent itself is not + retained, such that results remain picklable) + :param max_answer_chars: the maximum number of characters; -1 for the configured default + """ + self._default_max_answer_chars = agent.serena_config.default_max_tool_answer_chars + self._max_answer_chars = max_answer_chars + + def _limit_length( + self, + result: str, + shortened_result_factories: list[Callable[[], str]] | None = None, + ) -> str: + """Limit the length of the result string, optionally trying progressively shorter versions. + + :param result: the full result string + :param max_answer_chars: maximum allowed characters. -1 means use the default from config. + :param shortened_result_factories: optional list of closures, each producing a progressively shorter + version of the result. They are tried in order until one fits within ``max_answer_chars``. + :return: the result string, potentially replaced by a shortened version + """ + return TextOutputUtils.limit_length( + result=result, max_answer_chars=self._get_max_answer_chars(), shortened_result_factories=shortened_result_factories + ) + + def _get_max_answer_chars(self) -> int: + """ + :return: the effective maximum number of characters, resolving the default from the configuration + """ + return self._default_max_answer_chars if self._max_answer_chars == -1 else self._max_answer_chars + + def _to_json(self, x: Any) -> str: + return TextOutputUtils.to_json(x) + + @abstractmethod + def render(self, obj: T) -> str: + """ + :return: a textual representation of this object for the LLM + """ + + +class Representable(ABC): + """ + An object which can render itself as a string suitable for consumption by an LLM. + """ + + @abstractmethod + def represent(self) -> str: + """ + :return: a textual representation of this object for the LLM + """ + + +class RepresentableViaRenderer(Representable): + """ + A representable object which uses a renderer to render itself. + """ + + def __init__(self, renderer: Renderer): + self._renderer = renderer + + def represent(self) -> str: + return self._renderer.render(self) + + +class JsonObject(RepresentableViaRenderer): + """ + A JSON-serializable result (dict, list, etc.) which is rendered as JSON, subject to length limitation. + """ + + def __init__(self, data: Any, renderer: "JsonObjectRenderer"): + """ + :param data: the JSON-serializable data + :param renderer: the renderer to use for representing the data + """ + super().__init__(renderer) + self.data = data + + data: Any + """the JSON-serializable data (dict, list, etc.)""" + + +class JsonObjectRenderer(Renderer[JsonObject]): + def render(self, obj: JsonObject) -> str: + return self._limit_length(self._to_json(obj.data)) diff --git a/src/serena/resources/config/contexts/context.template.yml b/src/serena/resources/config/contexts/context.template.yml index 1f8e37ab..71eefa51 100644 --- a/src/serena/resources/config/contexts/context.template.yml +++ b/src/serena/resources/config/contexts/context.template.yml @@ -23,6 +23,13 @@ included_optional_tools: [] # Find the list of tools here: https://oraios.github.io/serena/01-about/035_tools.html fixed_tools: [] +# APIs (facades or facade methods, e.g. "lsp" or "lsp.find_symbol") to be excluded from the REPL in this context. +excluded_apis: [] + +# included APIs that would otherwise be excluded (particularly optional facade methods, which are disabled by default), +# e.g. "lsp.get_diagnostics_for_symbol". +included_apis: [] + # mapping of tool names to an override of their descriptions (the default description is the docstring of the Tool's apply method). # Sometimes, tool descriptions are too long (e.g., for ChatGPT), or users may want to override them for another reason. tool_description_overrides: {} diff --git a/src/serena/resources/config/modes/benchmark.yml b/src/serena/resources/config/modes/benchmark.yml index 636f5f4b..04e7d8d7 100644 --- a/src/serena/resources/config/modes/benchmark.yml +++ b/src/serena/resources/config/modes/benchmark.yml @@ -18,3 +18,5 @@ excluded_tools: - delete_memory - rename_memory - onboarding +excluded_apis: + - mem diff --git a/src/serena/resources/config/modes/editing.yml b/src/serena/resources/config/modes/editing.yml index f877fc51..1602ca51 100644 --- a/src/serena/resources/config/modes/editing.yml +++ b/src/serena/resources/config/modes/editing.yml @@ -4,7 +4,7 @@ prompt: | **Refactoring tools** For operations on existing symbols, prefer the dedicated refactoring tools over hand-edits: - `{{ tool_names['rename_symbol'] }}` and `{{ tool_names['safe_delete_symbol'] }}`{% if 'jet_brains_move' in available_tools %}, plus `jet_brains_move` and `jet_brains_inline_symbol`{% endif %} — + `{{ tool_names['rename_symbol'] }}` and `{{ tool_names['safe_delete_symbol'] }}`{% if 'jet_brains_move' in available_tools %}, plus `{{ tool_names['jet_brains_move'] }}` and `{{ tool_names['jet_brains_inline_symbol'] }}`{% endif %} — they are reference-aware and update or check all usages atomically. When such a tool returns success, the refactoring is already complete and consistent across all declarations, references, overrides and imports — trust it: do not re-read the changed files or re-run the build / test suite just to confirm the refactor @@ -16,9 +16,9 @@ prompt: | **Symbolic editing** Use symbolic retrieval tools to identify the symbols you need to edit. - If you need to replace the definition of a symbol, use the `replace_symbol_body` tool. - If you want to add some new code at the end of the file, use the `insert_after_symbol` tool with the last top-level symbol in the file. - Similarly, you can use `insert_before_symbol` with the first top-level symbol in the file to insert code at the beginning of a file. + If you need to replace the definition of a symbol, use the `{{ tool_names['replace_symbol_body'] }}` tool. + If you want to add some new code at the end of the file, use the `{{ tool_names['insert_after_symbol'] }}` tool with the last top-level symbol in the file. + Similarly, you can use `{{ tool_names['insert_before_symbol'] }}` with the first top-level symbol in the file to insert code at the beginning of a file. You can understand relationships between symbols by using the `{{ tool_names['find_referencing_symbols'] }}` tool. If not explicitly requested otherwise by the user, you make sure that when you edit a symbol, the change is either backward-compatible or you find and update all references as needed. The `{{ tool_names['find_referencing_symbols'] }}` tool will give you code snippets around the references as well as symbolic information. @@ -26,14 +26,14 @@ prompt: | {% if 'replace_content' in available_tools %} **File-based editing** - The `replace_content` tool allows you to perform regex-based replacements within files (as well as simple string replacements). + The `{{ tool_names['replace_content'] }}` tool allows you to perform regex-based replacements within files (as well as simple string replacements). This is your primary tool for editing code whenever replacing or deleting a whole symbol would be a more expensive operation, e.g. if you need to adjust just a few lines of code within a method. In `regex` mode, wildcards like `start.*?end` let you match a span without quoting its full text; an ambiguous match returns an error you can refine, so a tight wildcard pattern is both cheaper and safe. - For several small edits within one file, prefer a batch of targeted `replace_content` calls over rewriting the whole + For several small edits within one file, prefer a batch of targeted `{{ tool_names['replace_content'] }}` calls over rewriting the whole file: rewriting has to re-emit the file's entire contents, whereas each targeted edit emits only the changed text. - {% if 'replace_in_files' in available_tools %}For the SAME edit across many files, use `replace_in_files`, which applies + {% if 'replace_in_files' in available_tools %}For the SAME edit across many files, use `{{ tool_names['replace_in_files'] }}`, which applies it everywhere in one call — equally safe and transparent: its `dry_run` mode first previews every change as a diff with a per-occurrence id, so you can then apply all of them or just a chosen subset.{% endif %} {% endif %} diff --git a/src/serena/resources/config/modes/mode.template.yml b/src/serena/resources/config/modes/mode.template.yml index c62573d7..7efb6f9b 100644 --- a/src/serena/resources/config/modes/mode.template.yml +++ b/src/serena/resources/config/modes/mode.template.yml @@ -22,4 +22,11 @@ included_optional_tools: [] # fixed set of tools to use as the base tool set (if non-empty), replacing Serena's default set of tools. # This cannot be combined with non-empty excluded_tools or included_optional_tools. # Find the list of tools here: https://oraios.github.io/serena/01-about/035_tools.html -fixed_tools: [] \ No newline at end of file +fixed_tools: [] + +# APIs (facades or facade methods, e.g. "lsp" or "lsp.find_symbol") to be excluded from the REPL in this mode. +excluded_apis: [] + +# included APIs that would otherwise be excluded (particularly optional facade methods, which are disabled by default), +# e.g. "lsp.get_diagnostics_for_symbol". +included_apis: [] diff --git a/src/serena/resources/config/modes/no-memories.yml b/src/serena/resources/config/modes/no-memories.yml index d4175eef..0d6c50a7 100644 --- a/src/serena/resources/config/modes/no-memories.yml +++ b/src/serena/resources/config/modes/no-memories.yml @@ -9,3 +9,5 @@ excluded_tools: - rename_memory - list_memories - onboarding +excluded_apis: + - mem diff --git a/src/serena/resources/config/modes/no-onboarding.yml b/src/serena/resources/config/modes/no-onboarding.yml index d441b5d3..bae5cb36 100644 --- a/src/serena/resources/config/modes/no-onboarding.yml +++ b/src/serena/resources/config/modes/no-onboarding.yml @@ -3,3 +3,5 @@ prompt: | The onboarding process is not applied. excluded_tools: - onboarding +excluded_apis: + - mem.onboarding diff --git a/src/serena/resources/config/modes/query-projects.yml b/src/serena/resources/config/modes/query-projects.yml index ed140de9..9ceea8fa 100644 --- a/src/serena/resources/config/modes/query-projects.yml +++ b/src/serena/resources/config/modes/query-projects.yml @@ -6,3 +6,5 @@ excluded_tools: [] included_optional_tools: - list_queryable_projects - query_project +included_apis: + - ext diff --git a/src/serena/resources/config/prompt_templates/system_prompt.yml b/src/serena/resources/config/prompt_templates/system_prompt.yml index 99604bf1..95d5bae9 100644 --- a/src/serena/resources/config/prompt_templates/system_prompt.yml +++ b/src/serena/resources/config/prompt_templates/system_prompt.yml @@ -5,29 +5,28 @@ prompts: connection_prompt: | CRITICAL: Before starting to work on a coding task, call the `initial_instructions` tool to read the 'Serena Instructions Manual'. system_prompt: | - You have semantic coding tools that you rely on heavily. Work resource-efficiently: don't read or + You have semantic coding functions/tools that you rely on heavily. Work resource-efficiently: don't read or generate content the task doesn't need. Some tasks require understanding a large part of the codebase; others need only a few symbols or a single file. Avoid reading whole files unless necessary — acquire information step by step, using - the symbolic tools to get an overview of symbols and their relations, then reading only the bodies + the symbolic functions to get an overview of symbols and their relations, then reading only the bodies you need. {% if 'ToolMarkerSymbolicRead' in available_markers %}Once you have read a full file, there is no point re-analysing it with the symbolic read tools — you already have it.{% endif %} {% if 'search_for_pattern' in available_tools %} - If you are unsure about a symbol's name or location{% if 'find_symbol' in available_tools %} (to the extent that substring_matching for the symbol name is not enough){% endif %}, you can use the `search_for_pattern` tool, which allows fast - and flexible search for patterns in the codebase.{% if 'ToolMarkerSymbolicRead' in available_markers %} In this way, you can first find candidates for symbols or files, - and then proceed with the symbolic tools.{% endif %} + For searches beyond named symbols, you can use `{{ tool_names['search_for_pattern'] }}` and flexible search for patterns in the codebase.{% if 'ToolMarkerSymbolicRead' in available_markers %} In this way, you can first find candidates for symbols or files to explore, + and then proceed with symbolic operations.{% endif %} {% endif %} {% if 'ToolMarkerSymbolicRead' in available_markers %} Symbols are identified by their `name_path` and `relative_path`. - You can get an overview of the symbols in a file by using the `{{ tool_names['get_symbols_overview'] }}` tool, or search for a specific symbol with `{{ tool_names['find_symbol'] }}`. + You can get an overview of the symbols in a file by using `{{ tool_names['get_symbols_overview'] }}`, or search for a specific symbol with `{{ tool_names['find_symbol'] }}`. You only read the bodies of symbols when you need to (e.g. if you want to fully understand or edit it). For example, if you are working with Python code and already know that you need to read the body of the constructor of the class Foo, you can directly use `{{ tool_names['find_symbol'] }}` with name path pattern `Foo/__init__` and `include_body=True`. If you don't know yet which methods in `Foo` you need to read or edit, - you can use `{{ tool_names['find_symbol'] }}` with name path pattern `Foo`, `include_body=False` and `depth=1` to get all (top-level) methods of `Foo` before proceeding + you can use `{{ tool_names['find_symbol'] }}` with name path pattern `Foo`, `include_body=False` and `depth=1` to get all (top-level) members of `Foo` before proceeding to read the desired methods with `include_body=True`. - You can understand relationships between symbols by using the `{{ tool_names['find_referencing_symbols'] }}` tool. + You can understand relationships between symbols by using `{{ tool_names['find_referencing_symbols'] }}`. {% endif %} {% if 'read_memory' in available_tools -%} diff --git a/src/serena/resources/dashboard/dashboard.css b/src/serena/resources/dashboard/dashboard.css index f8e4e1af..08514e77 100644 --- a/src/serena/resources/dashboard/dashboard.css +++ b/src/serena/resources/dashboard/dashboard.css @@ -603,6 +603,23 @@ code, pre, kbd, samp, cursor: default; } +.facade-block { + margin-bottom: 12px; +} + +.facade-name { + font-size: 13px; + font-weight: 600; + color: var(--text-primary); + margin-bottom: 4px; +} + +.facade-name.disabled, +.tool-item.disabled { + color: var(--text-secondary); + opacity: 0.5; +} + /* Projects List */ .project-item { padding: 10px 12px; diff --git a/src/serena/resources/dashboard/dashboard.js b/src/serena/resources/dashboard/dashboard.js index 74308c93..45330466 100644 --- a/src/serena/resources/dashboard/dashboard.js +++ b/src/serena/resources/dashboard/dashboard.js @@ -614,6 +614,7 @@ class Dashboard { const $existingToolsContent = $('#tools-content'); const $existingMemoriesContent = $('#memories-content'); const wasToolsExpanded = $existingToolsContent.is(':visible'); + const wasFunctionsExpanded = $('#functions-content').is(':visible'); const wasMemoriesExpanded = $existingMemoriesContent.is(':visible'); let html = '
'; @@ -634,10 +635,14 @@ class Dashboard { html += '
' + (config.active_project.name || 'None') + '
'; } - html += '
Languages:
'; - if (this.jetbrainsMode) { - html += '
Using JetBrains backend
'; - } else { + html += '
Interface:
'; + html += '
' + config.agent_interface + '
'; + + html += '
Backend:
'; + html += '
' + config.language_backend + '
'; + + if (!this.jetbrainsMode) { + html += '
Languages:
'; html += '
'; if (config.languages && config.languages.length > 0) { html += '
'; @@ -705,6 +710,32 @@ class Dashboard { html += '
'; html += '
'; + // Active functions of the REPL's facades - collapsible (REPL interface only) + if (config.facades) { + const enabledMethodCount = config.facades.reduce(function (count, facade) { + return count + facade.methods.filter(function (method) { return method.is_enabled; }).length; + }, 0); + html += '
'; + html += '

'; + html += 'Active Functions (' + enabledMethodCount + ')'; + html += '▼'; + html += '

'; + html += '
'; + config.facades.forEach(function (facade) { + html += '
'; + html += '
s.' + facade.name + '
'; + html += '
'; + facade.methods.forEach(function (method) { + const title = facade.name + '.' + method.name + (method.is_enabled ? '' : ' (disabled)'); + html += '
' + method.name + '
'; + }); + html += '
'; + html += '
'; + }); + html += '
'; + html += '
'; + } + // Available memories - collapsible (show if memories exist or if project exists) if (config.active_project && config.active_project.name) { html += '
'; @@ -773,6 +804,15 @@ class Dashboard { $('#create-memory-btn').click(this.openCreateMemoryModal.bind(this)); // Re-attach collapsible handler for the newly created tools header + $('#functions-header').click(function () { + const $header = $(this); + const $content = $('#functions-content'); + const $icon = $header.find('.toggle-icon'); + + $content.slideToggle(300); + $icon.toggleClass('expanded'); + }); + $('#tools-header').click(function () { const $header = $(this); const $content = $('#tools-content'); diff --git a/src/serena/resources/project.template.yml b/src/serena/resources/project.template.yml index a4f2a0fb..c72f7ad0 100644 --- a/src/serena/resources/project.template.yml +++ b/src/serena/resources/project.template.yml @@ -68,6 +68,11 @@ line_ending: # is activated post-init, an error will be returned. language_backend: +# The interface through which the agent (LLM) accesses Serena's functionality (overrides the global setting). +# Valid values: tools, REPL (see the global configuration for details); leave empty to use the global setting. +# Note: the interface is fixed at startup. If a project is activated post-init, its setting is not applied. +agent_interface: + # whether to use project's .gitignore files to ignore files ignore_all_files_in_gitignore: true @@ -131,6 +136,15 @@ included_optional_tools: [] # Find the list of tools here: https://oraios.github.io/serena/01-about/035_tools.html fixed_tools: [] +# list of APIs (facades or facade methods, e.g. "lsp" or "lsp.find_symbol") to exclude from the REPL. +# This extends the existing exclusions (e.g. from the global configuration). +excluded_apis: [] + +# list of APIs (facades or facade methods, e.g. "lsp" or "lsp.get_diagnostics_for_symbol") to include in the REPL +# that would otherwise be disabled (particularly optional methods, which are disabled by default). +# This extends the existing inclusions (e.g. from the global configuration). +included_apis: [] + # list of mode names that are to be activated by default, overriding the setting in the global configuration. # The full set of modes to be activated is base_modes (from global config) + default_modes + added_modes. # If the setting is undefined/empty, the default_modes from the global configuration (serena_config.yml) apply. diff --git a/src/serena/resources/serena_config.template.yml b/src/serena/resources/serena_config.template.yml index 864d037c..12f70b9f 100644 --- a/src/serena/resources/serena_config.template.yml +++ b/src/serena/resources/serena_config.template.yml @@ -7,6 +7,18 @@ # in your IDE). language_backend: LSP +# The interface through which the agent (LLM) accesses Serena's functionality: +# * tools: the classic tool interface, in which each operation is a separate tool. +# The set of tools is configurable via the tool inclusion/exclusion settings (excluded_tools etc.) +# of the configuration, the context, modes and the project. +# * REPL: operations are accessed programmatically via the serena_repl tool, which executes Python code. +# The set of tools is fixed (the REPL tool and the tools required for session management), and the tool +# inclusion/exclusion settings do not apply; instead, the operations available in the REPL are configured +# via the API inclusion/exclusion settings (excluded_apis etc.). +# The interface is fixed at startup; it can be overridden by the project activated at startup. +# IMPORTANT: The REPL interface is a BETA feature. Please provide feedback; if you encounter issues, report them. +agent_interface: tools + # line ending convention to use when writing source files. # Possible values: "lf" (Unix), "crlf" (Windows), "native" (platform default). # Note that Serena's own files (e.g. memories and configuration files) always use native line endings. @@ -151,6 +163,13 @@ included_optional_tools: [] # This cannot be combined with non-empty excluded_tools or included_optional_tools. fixed_tools: [] +# list of APIs (facades or facade methods, e.g. "lsp" or "lsp.find_symbol") to be globally excluded from the REPL +excluded_apis: [] + +# list of APIs (facades or facade methods, e.g. "lsp.get_diagnostics_for_symbol") to be included in the REPL +# (particularly optional methods, which are disabled by default) +included_apis: [] + # list of mode names to that are always to be included in the set of active modes. # The full set of modes to be activated is base_modes + default_modes + added_modes, # where added_modes can be defined by projects/CLI parameters. @@ -221,5 +240,9 @@ project_serena_folder_location: "$projectDir/.serena" # The pattern "**" matches any project path, so it can be used to trust all projects. trusted_project_path_patterns: [] +# shared secret for authenticating communication between Serena components and services. +# Keep this value private. When missing, null, or empty, a random UUID is generated and saved on load. +auth_secret: + # the list of registered project paths (updated automatically). projects: [] diff --git a/src/serena/session.py b/src/serena/session.py new file mode 100644 index 00000000..01d668a6 --- /dev/null +++ b/src/serena/session.py @@ -0,0 +1,89 @@ +# SPDX-License-Identifier: GPL-3.0-or-later +""" +Client sessions (conversations) and their state. +""" + +import logging +import secrets +import time +from collections import OrderedDict +from typing import Any + +log = logging.getLogger(__name__) + + +class SerenaSession: + """ + A client session, i.e. a conversation between an LLM and Serena, and the state pertaining to it. + + Session identity is supplied by the LLM: the session id is issued as part of Serena's instructions + (system prompt/initial instructions) and passed by the LLM to tools which require it. + """ + + def __init__(self, session_id: str) -> None: + self.session_id = session_id + self.described_type_names: set[str] = set() + """the names of the REPL's result types whose documentation has already been provided in this session""" + self.repl_namespace: dict[str, Any] = {} + """ + the namespace (globals) of the session's REPL code executions: variables and functions defined at the top level + of executed code persist here across executions for the lifetime of the session + """ + self.last_access_time = time.time() + + +class SessionRegistry: + """ + Holds the sessions of an agent, creating them on demand and evicting sessions which have been idle for too long + as well as the least recently used ones if their number exceeds the limit (session ids being supplied by LLMs, + the set of ids is not controlled). + """ + + def __init__(self, max_sessions: int = 100, idle_ttl_seconds: float = 6 * 3600) -> None: + """ + :param max_sessions: the maximum number of sessions to keep + :param idle_ttl_seconds: the time after which an idle session is evicted (releasing its REPL namespace) + """ + self._max_sessions = max_sessions + self._idle_ttl_seconds = idle_ttl_seconds + self._sessions: OrderedDict[str, SerenaSession] = OrderedDict() + + @staticmethod + def _next_session_id() -> str: + return secrets.token_hex(4) + + def create_session(self) -> SerenaSession: + """ + :return: a new session with a random id + """ + return self.get_session(self._next_session_id()) + + def get_session(self, session_id: str) -> SerenaSession: + """ + :param session_id: the session id + :return: the session, which is created if it is unknown (an unknown id may e.g. stem from a session that has been + evicted or from an earlier run of the server) + """ + self._evict_idle_sessions() + session = self._sessions.get(session_id) + if session is None: + session = SerenaSession(session_id) + self._sessions[session_id] = session + log.info("Created session %s (%d sessions)", session_id, len(self._sessions)) + while len(self._sessions) > self._max_sessions: + evicted_id, _ = self._sessions.popitem(last=False) + log.info("Evicted session %s (session limit)", evicted_id) + else: + self._sessions.move_to_end(session_id) + session.last_access_time = time.time() + return session + + def _evict_idle_sessions(self) -> None: + # sessions are ordered by last access, so the idle ones are at the front + now = time.time() + while self._sessions: + oldest_id, oldest = next(iter(self._sessions.items())) + if now - oldest.last_access_time <= self._idle_ttl_seconds: + break + del self._sessions[oldest_id] + log.info("Evicted session %s (idle)", oldest_id) diff --git a/src/serena/tools/__init__.py b/src/serena/tools/__init__.py index 2475dfca..0ed2e6bf 100644 --- a/src/serena/tools/__init__.py +++ b/src/serena/tools/__init__.py @@ -10,3 +10,4 @@ from .config_tools import * from .workflow_tools import * from .jetbrains_tools import * from .query_project_tools import * +from .repl_tools import * diff --git a/src/serena/tools/cmd_tools.py b/src/serena/tools/cmd_tools.py index ace9fab1..51357995 100644 --- a/src/serena/tools/cmd_tools.py +++ b/src/serena/tools/cmd_tools.py @@ -3,13 +3,28 @@ Tools supporting the execution of (external) commands """ # SPDX-License-Identifier: GPL-3.0-or-later -import os.path +from typing import TYPE_CHECKING, cast from serena.tools import Tool, ToolMarkerCanEdit -from serena.util.shell import execute_shell_command + +if TYPE_CHECKING: + from serena.repl.api.shell_api import ShellApi -class ExecuteShellCommandTool(Tool, ToolMarkerCanEdit): +class ShellApiMixin: + """ + Mixin for tools which delegate to the shell API. + The API is imported locally, since the API module refers to the tools (as corresponding tools). + """ + + def _api(self) -> "ShellApi": + from serena.repl.api.shell_api import ShellApi + + tool = cast(Tool, cast(object, self)) + return ShellApi(tool.agent) + + +class ExecuteShellCommandTool(Tool, ToolMarkerCanEdit, ShellApiMixin): """ Executes a shell command. """ @@ -36,18 +51,4 @@ class ExecuteShellCommandTool(Tool, ToolMarkerCanEdit): required for the task. :return: a JSON object containing the command's stdout and optionally stderr output """ - if cwd is None: - _cwd = self.get_project_root() - else: - if os.path.isabs(cwd): - _cwd = cwd - else: - _cwd = os.path.join(self.get_project_root(), cwd) - if not os.path.isdir(_cwd): - raise FileNotFoundError( - f"Specified a relative working directory ({cwd}), but the resulting path is not a directory: {_cwd}" - ) - - result = execute_shell_command(command, cwd=_cwd, capture_stderr=capture_stderr) - result = result.model_dump_json() - return self._limit_length(result, max_answer_chars) + return self._api().execute_shell_command(command, cwd, capture_stderr, max_answer_chars).represent() diff --git a/src/serena/tools/config_tools.py b/src/serena/tools/config_tools.py index aab6b13e..0fdd7f56 100644 --- a/src/serena/tools/config_tools.py +++ b/src/serena/tools/config_tools.py @@ -1,11 +1,29 @@ # SPDX-License-Identifier: GPL-3.0-or-later +from typing import TYPE_CHECKING, cast + from sensai.util.helper import mark_used from serena.tools import Tool, ToolMarkerDoesNotRequireActiveProject, ToolMarkerOptional +if TYPE_CHECKING: + from serena.repl.api.cfg_api import ConfigApi -class OpenDashboardTool(Tool, ToolMarkerOptional, ToolMarkerDoesNotRequireActiveProject): + +class ConfigApiMixin: + """ + Mixin for tools which delegate to the configuration API. + The API is imported locally, since the API module refers to the tools (as corresponding tools). + """ + + def _api(self) -> "ConfigApi": + from serena.repl.api.cfg_api import ConfigApi + + tool = cast(Tool, cast(object, self)) + return ConfigApi(tool.agent) + + +class OpenDashboardTool(Tool, ToolMarkerOptional, ToolMarkerDoesNotRequireActiveProject, ConfigApiMixin): """ Opens the Serena web dashboard in the default web browser. The dashboard provides logs, session information, and tool usage statistics. @@ -15,10 +33,7 @@ class OpenDashboardTool(Tool, ToolMarkerOptional, ToolMarkerDoesNotRequireActive """ Opens the Serena web dashboard in the default web browser. """ - if self.agent.open_dashboard(): - return f"Serena web dashboard has been opened in the user's default web browser: {self.agent.get_dashboard_url()}" - else: - return f"Serena web dashboard could not be opened automatically; tell the user to open it via {self.agent.get_dashboard_url()}" + return self._api().open_dashboard() class ActivateProjectTool(Tool, ToolMarkerDoesNotRequireActiveProject): @@ -26,13 +41,12 @@ class ActivateProjectTool(Tool, ToolMarkerDoesNotRequireActiveProject): Activates a project based on the project name or path. """ - # noinspection PyIncorrectDocstring - # (session_id is injected via apply_ex) def apply(self, project: str, session_id: str) -> str: """ Activates the project with the given name or path. :param project: the name of a registered project to activate or a path to a project directory + :param session_id: your Serena session id, as provided in Serena's instructions (call `initial_instructions` if you do not have one) """ is_new_activation = self.agent.activate_project_from_path_or_name(project) mark_used(is_new_activation) @@ -56,7 +70,7 @@ class RemoveProjectTool(Tool, ToolMarkerDoesNotRequireActiveProject, ToolMarkerO return f"Successfully removed project '{project_name}' from configuration." -class GetCurrentConfigTool(Tool): +class GetCurrentConfigTool(Tool, ConfigApiMixin): """ Prints the current configuration of the agent, including the active and available projects, tools, contexts, and modes. """ @@ -65,4 +79,4 @@ class GetCurrentConfigTool(Tool): """ Print the current configuration of the agent, including the active and available projects, tools, contexts, and modes. """ - return self.agent.get_current_config_overview() + return self._api().get_current_config() diff --git a/src/serena/tools/file_tools.py b/src/serena/tools/file_tools.py index 67c09fe5..fd61ed54 100644 --- a/src/serena/tools/file_tools.py +++ b/src/serena/tools/file_tools.py @@ -7,24 +7,42 @@ File and file system-related tools, specifically for """ # SPDX-License-Identifier: GPL-3.0-or-later -import os -from collections import defaultdict -from fnmatch import fnmatch -from pathlib import Path -from typing import Literal +from typing import TYPE_CHECKING, Literal, cast -from serena.tools import SUCCESS_RESULT, EditedFileContext, EditingToolWithDiagnostics, Tool, ToolMarkerOptional -from serena.util.file_system import scan_directory -from serena.util.text_utils import ( - ContentReplacer, - GlobMatcher, - MultiFileContentReplacer, - ReplacementOccurrence, -) -from solidlsp.ls_utils import TextUtils +from serena.tools import EditingToolWithDiagnostics, Tool, ToolMarkerOptional + +if TYPE_CHECKING: + from serena.repl.api.edit_api import EditApi + from serena.repl.api.fs_api import FsApi -class ReadFileTool(Tool): +class EditApiMixin: + """ + Mixin for tools which delegate to the editing API. + The API is imported locally, since the API module refers to the tools (as corresponding tools). + """ + + def _api(self) -> "EditApi": + from serena.repl.api.edit_api import EditApi + + tool = cast(Tool, cast(object, self)) + return EditApi(tool.agent) + + +class FsApiMixin: + """ + Mixin for tools which delegate to the file system API. + The API is imported locally, since the API module refers to the tools (as corresponding tools). + """ + + def _api(self) -> "FsApi": + from serena.repl.api.fs_api import FsApi + + tool = cast(Tool, cast(object, self)) + return FsApi(tool.agent) + + +class ReadFileTool(Tool, FsApiMixin): """ Reads a file within the project directory. """ @@ -41,22 +59,10 @@ class ReadFileTool(Tool): required for the task. :return: the full text of the file at the given relative path """ - self.project.validate_relative_path(relative_path) - - # read lines, using the same (LSP-compliant) notion of line breaks as the line-based editing tools - result = self.project.read_file(relative_path) - result_lines = TextUtils.split_lines(result) - - if end_line is None: - result_lines = result_lines[start_line:] - else: - result_lines = result_lines[start_line : end_line + 1] - result = "\n".join(result_lines) - - return self._limit_length(result, max_answer_chars) + return self._api().read_file(relative_path, start_line, end_line, max_answer_chars).represent() -class CreateTextFileTool(EditingToolWithDiagnostics): +class CreateTextFileTool(EditingToolWithDiagnostics, FsApiMixin): """ Creates/overwrites a file in the project directory. """ @@ -69,30 +75,11 @@ class CreateTextFileTool(EditingToolWithDiagnostics): :param content: the (appropriately encoded) content to write to the file :return: a message indicating success or failure """ - with self.DiagnosticsContext(self, relative_path) as diagnostics_context: - # validating the destination path - project_root = self.get_project_root() - abs_path = (Path(project_root) / relative_path).resolve() - will_overwrite_existing = abs_path.exists() - - if will_overwrite_existing: - self.project.validate_relative_path(relative_path) - else: - assert abs_path.is_relative_to(self.get_project_root()), ( - f"Cannot create file outside of the project directory, got {relative_path=}" - ) - - # writing the file - abs_path.parent.mkdir(parents=True, exist_ok=True) - abs_path.write_text(content, encoding=self.project.project_config.encoding, newline=self.project.line_ending.newline_str) - answer = f"File created: {relative_path}." - if will_overwrite_existing: - answer += " Overwrote existing file." - - return diagnostics_context.format_result(answer) + with self.diagnostics_context(relative_path) as diagnostics_context: + return diagnostics_context.format_result(self._api().create_text_file(relative_path, content)) -class ListDirTool(Tool): +class ListDirTool(Tool, FsApiMixin): """ Lists files and directories in the given directory (optionally with recursion). """ @@ -109,31 +96,13 @@ class ListDirTool(Tool): Don't adjust unless there is really no other way to get the content required for the task. :return: a JSON object with the names of directories and files within the given directory """ - # Check if the directory exists before validation - if not self.project.relative_path_exists(relative_path): - error_info = { - "error": f"Directory not found: {relative_path}", - "project_root": self.get_project_root(), - "hint": "Check if the path is correct relative to the project root", - } - return self._to_json(error_info) - - self.project.validate_relative_path(relative_path) - - is_ignored_path_fn = self.project.get_is_ignored_path_fn(relative_path, skip_ignored_files) - dirs, files = scan_directory( - os.path.join(self.get_project_root(), relative_path), - relative_to=self.get_project_root(), - recursive=recursive, - is_ignored_dir=is_ignored_path_fn, - is_ignored_file=is_ignored_path_fn, - ) - - result = self._to_json({"dirs": dirs, "files": files}) - return self._limit_length(result, max_answer_chars) + try: + return self._api().list_dir(relative_path, recursive, skip_ignored_files, max_answer_chars).represent() + except FileNotFoundError as e: + return self._to_json({"error": str(e), "project_root": self.get_project_root()}) -class FindFileTool(Tool): +class FindFileTool(Tool, FsApiMixin): """ Finds files in the given relative paths """ @@ -147,31 +116,10 @@ class FindFileTool(Tool): :param skip_ignored_files: whether to skip ignored files/directories :return: a JSON object with the list of matching files """ - self.project.validate_relative_path(relative_path) - - is_ignored_path_fn = self.project.get_is_ignored_path_fn(relative_path, skip_ignored_paths=False) - 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 is_ignored_path_fn(abs_path): - return True - filename = os.path.basename(abs_path) - return not fnmatch(filename, file_mask) - - _dirs, files = scan_directory( - path=dir_to_scan, - recursive=True, - is_ignored_dir=is_ignored_path_fn, - is_ignored_file=is_ignored_file, - relative_to=self.get_project_root(), - ) - - result = self._to_json({"files": files}) - return result + return self._to_json({"files": self._api().find_file(file_mask, relative_path)}) -class ReplaceContentTool(EditingToolWithDiagnostics): +class ReplaceContentTool(EditingToolWithDiagnostics, EditApiMixin): """ Replaces content in a file (optionally using regular expressions). """ @@ -206,17 +154,13 @@ class ReplaceContentTool(EditingToolWithDiagnostics): :param allow_multiple_occurrences: whether to allow matching and replacing multiple occurrences. If false and multiple occurrences are found, an error will be returned """ - with self.DiagnosticsContext(self, relative_path) as diagnostics_context: - self.project.validate_relative_path(relative_path) - with EditedFileContext(relative_path, self.create_code_editor()) as context: - original_content = context.get_original_content() - replacer = ContentReplacer(mode=mode, allow_multiple_occurrences=allow_multiple_occurrences) - updated_content = replacer.replace(original_content, needle, repl) - context.set_updated_content(updated_content) - return diagnostics_context.format_result(SUCCESS_RESULT) + with self.diagnostics_context(relative_path) as diagnostics_context: + return diagnostics_context.format_result( + self._api().replace_content(relative_path, needle, repl, mode, allow_multiple_occurrences=allow_multiple_occurrences) + ) -class ReplaceInFilesTool(EditingToolWithDiagnostics): +class ReplaceInFilesTool(EditingToolWithDiagnostics, EditApiMixin): """ Replaces occurrences of a pattern across multiple files, with dry-run preview and per-occurrence selection. """ @@ -272,184 +216,28 @@ class ReplaceInFilesTool(EditingToolWithDiagnostics): returned. -1 uses the configured default. :return: in a dry run, the prospective changes; otherwise a summary of the applied replacements """ - replacer = MultiFileContentReplacer(mode=mode) - files = self._collect_files(relative_path, paths_include_glob, paths_exclude_glob) - occurrences = replacer.find_occurrences(files, needle, repl) - contents = dict(files) - + api = self._api() if dry_run: - return self._render_listing(replacer, occurrences, contents, max_answer_chars, dry_run=True) - - if occurrence_ids is not None: - selected, problems = self._resolve_occurrence_ids(occurrence_ids, occurrences) - if problems: - problem_lines = "\n".join(f" {p}" for p in problems) - raise ValueError( - f"{len(problems)} of the given occurrence_ids could not be resolved - NO changes were applied:\n" - f"{problem_lines}\n" - "Re-run with dry_run=True to obtain current occurrence ids." - ) - if not selected: - raise ValueError("occurrence_ids is empty - pass at least one id from a dry run, or omit the parameter to replace all.") - return self._apply_occurrences(replacer, selected, contents, needle, repl) - - # blind apply (no ids) - if not occurrences: - raise ValueError( - "No occurrences of the pattern were found - NO changes were applied. " - "Check the mode (a literal needle containing regex metacharacters must use mode 'literal'; " - "wildcards require mode 'regex') and the path/glob restrictions, " - "or locate the content with search_for_pattern first." + return api.replace_in_files( + needle, repl, mode, relative_path, paths_include_glob, paths_exclude_glob, dry_run=True, max_answer_chars=max_answer_chars + ).represent() + with self.diagnostics_context() as diagnostics_context: + result = api.replace_in_files( + needle, + repl, + mode, + relative_path, + paths_include_glob, + paths_exclude_glob, + occurrence_ids=occurrence_ids, + expected_count=expected_count, + max_answer_chars=max_answer_chars, ) - if expected_count >= 0 and len(occurrences) != expected_count: - listing = self._render_listing(replacer, occurrences, contents, max_answer_chars, dry_run=False) - raise ValueError( - f"expected_count={expected_count}, but the pattern matches {len(occurrences)} occurrence(s) - " - f"NO changes were applied. Review the prospective changes below; re-issue with the corrected " - f"expectation, a refined pattern, or occurrence_ids selecting the intended subset.\n{listing}" - ) - ambiguous = [o for o in occurrences if o.is_ambiguous] - if ambiguous: - listing = self._render_listing(replacer, occurrences, contents, max_answer_chars, dry_run=False) - raise ValueError( - f"{len(ambiguous)} occurrence(s) are ambiguous (the pattern matches again inside the matched text, " - f"indicating possible over-matching) - NO changes were applied. Review the prospective changes below " - f"and either refine the pattern or explicitly select occurrences via occurrence_ids.\n{listing}" - ) - return self._apply_occurrences(replacer, occurrences, contents, needle, repl) - - def _collect_files(self, relative_path: str, paths_include_glob: str, paths_exclude_glob: str) -> list[tuple[str, str]]: - """Collects (relative_path, content) pairs of the non-ignored files in scope, in sorted path order.""" - relative_path = relative_path.strip() - if relative_path: - self.project.validate_relative_path(relative_path, require_not_ignored=True) - abs_path = os.path.join(self.get_project_root(), relative_path) - if not os.path.exists(abs_path): - raise FileNotFoundError(f"Relative path {relative_path} does not exist.") - if os.path.isfile(abs_path): - rel_paths = [relative_path] - else: - _dirs, rel_paths = scan_directory( - path=abs_path, - recursive=True, - is_ignored_dir=self.project.is_ignored_path, - is_ignored_file=self.project.is_ignored_path, - relative_to=self.get_project_root(), - ) - include_glob_matcher = GlobMatcher(paths_include_glob.strip()) if paths_include_glob.strip() else None - exclude_glob_matcher = GlobMatcher(paths_exclude_glob.strip()) if paths_exclude_glob.strip() else None - files: list[tuple[str, str]] = [] - for path in sorted(rel_paths): - if include_glob_matcher and not include_glob_matcher.matches(path): - continue - if exclude_glob_matcher and exclude_glob_matcher.matches(path): - continue - try: - files.append((path, self.project.read_file(path))) - except Exception: - continue # skip unreadable (e.g. binary) files - return files - - def _render_listing( - self, - replacer: MultiFileContentReplacer, - occurrences: list[ReplacementOccurrence], - contents: dict[str, str], - max_answer_chars: int, - dry_run: bool, - ) -> str: - affected_files = sorted({o.relative_path for o in occurrences}) - header = f"Found {len(occurrences)} occurrence(s) in {len(affected_files)} file(s)." - if dry_run: - header += ( - " DRY RUN - no changes were applied.\n" - "Re-issue with dry_run=False to replace all of them, or additionally pass occurrence_ids " - "with the ids of the occurrences to replace." - ) - parts = [header] - for path in affected_files: - file_occurrences = [o for o in occurrences if o.relative_path == path] - parts.append(f"\n{path} ({len(file_occurrences)} occurrence(s)):") - for occ in file_occurrences: - parts.append(replacer.render_occurrence_diff(occ, contents[path])) - result = "\n".join(parts) - - def make_locations_only() -> str: - lines = [header] + [f" [{o.occurrence_id}] line {o.start_line}" for o in occurrences] - return "\n".join(lines) - - def make_per_file_counts() -> str: - counts = {path: sum(1 for o in occurrences if o.relative_path == path) for path in affected_files} - return f"{header}\nOccurrence counts per file:\n{self._to_json(counts)}" - - def make_summary() -> str: - return header - - return self._limit_length( - result, max_answer_chars, shortened_result_factories=[make_locations_only, make_per_file_counts, make_summary] - ) - - @staticmethod - def _resolve_occurrence_ids( - occurrence_ids: list[str], occurrences: list[ReplacementOccurrence] - ) -> tuple[list[ReplacementOccurrence], list[str]]: - """Resolves the requested ids against the current occurrences, diagnosing each failure.""" - occurrences_by_id = {o.occurrence_id: o for o in occurrences} - indices_by_path: dict[str, set[int]] = {} - for o in occurrences: - indices_by_path.setdefault(o.relative_path, set()).add(o.index_in_file) - selected: dict[str, ReplacementOccurrence] = {} - problems: list[str] = [] - for oid in occurrence_ids: - occurrence = occurrences_by_id.get(oid) - if occurrence is not None: - selected[oid] = occurrence - continue - id_match = MultiFileContentReplacer.OCCURRENCE_ID_REGEX.match(oid) - if id_match is None: - problems.append(f"{oid}: malformed id (expected ':@' as returned by a dry run)") - elif id_match.group("path") not in indices_by_path: - problems.append(f"{oid}: the pattern currently has no matches in this file") - elif int(id_match.group("index")) not in indices_by_path[id_match.group("path")]: - problems.append(f"{oid}: the file now has fewer matches than at dry-run time (content changed)") - else: - problems.append(f"{oid}: the matched text changed since the dry run (content changed)") - return list(selected.values()), problems - - def _apply_occurrences( - self, - replacer: MultiFileContentReplacer, - occurrences: list[ReplacementOccurrence], - contents: dict[str, str], - needle: str, - repl: str, - ) -> str: - occurrences_by_file: dict[str, list[ReplacementOccurrence]] = {} - for occ in occurrences: - occurrences_by_file.setdefault(occ.relative_path, []).append(occ) - with self.DiagnosticsContext(self, *occurrences_by_file.keys()) as diagnostics_context: - code_editor = self.create_code_editor() - for path, file_occurrences in occurrences_by_file.items(): - with EditedFileContext(path, code_editor) as context: - original_content = context.get_original_content() - if original_content != contents[path]: - # the editor's view differs from what was scanned (e.g. line-ending normalization); - # re-derive the occurrences from the authoritative content and re-validate by id - fresh_by_id = {o.occurrence_id: o for o in replacer.find_occurrences([(path, original_content)], needle, repl)} - try: - file_occurrences = [fresh_by_id[o.occurrence_id] for o in file_occurrences] - except KeyError as e: - raise ValueError( - f"The content of {path} changed while replacing (occurrence {e} no longer resolves); " - f"the file was NOT modified. Re-run with dry_run=True for current ids." - ) from e - context.set_updated_content(replacer.apply_to_content(original_content, file_occurrences)) - per_file = "\n".join(f" {path}: {len(occs)}" for path, occs in occurrences_by_file.items()) - summary = f"Replaced {len(occurrences)} occurrence(s) in {len(occurrences_by_file)} file(s):\n{per_file}" - return diagnostics_context.format_result(summary) + assert isinstance(result, str) + return diagnostics_context.format_result(result) -class DeleteLinesTool(EditingToolWithDiagnostics, ToolMarkerOptional): +class DeleteLinesTool(EditingToolWithDiagnostics, ToolMarkerOptional, EditApiMixin): """ Deletes a range of lines within a file. """ @@ -469,13 +257,11 @@ class DeleteLinesTool(EditingToolWithDiagnostics, ToolMarkerOptional): :param start_line: the 0-based index of the first line to be deleted :param end_line: the 0-based index of the last line to be deleted """ - with self.DiagnosticsContext(self, relative_path) as diagnostics_context: - code_editor = self.create_code_editor() - code_editor.delete_lines(relative_path, start_line, end_line) - return diagnostics_context.format_result(SUCCESS_RESULT) + with self.diagnostics_context(relative_path) as diagnostics_context: + return diagnostics_context.format_result(self._api().delete_lines(relative_path, start_line, end_line)) -class ReplaceLinesTool(EditingToolWithDiagnostics, ToolMarkerOptional): +class ReplaceLinesTool(EditingToolWithDiagnostics, ToolMarkerOptional, EditApiMixin): """ Replaces a range of lines within a file with new content. """ @@ -497,19 +283,11 @@ class ReplaceLinesTool(EditingToolWithDiagnostics, ToolMarkerOptional): :param end_line: the 0-based index of the last line to be deleted :param content: the content to insert """ - # normalizing the replacement content - if not content.endswith("\n"): - content += "\n" - - with self.DiagnosticsContext(self, relative_path) as diagnostics_context: - code_editor = self.create_code_editor() - code_editor.delete_lines(relative_path, start_line, end_line) - code_editor.insert_at_line(relative_path, start_line, content) - - return diagnostics_context.format_result(SUCCESS_RESULT) + with self.diagnostics_context(relative_path) as diagnostics_context: + return diagnostics_context.format_result(self._api().replace_lines(relative_path, start_line, end_line, content)) -class InsertAtLineTool(EditingToolWithDiagnostics, ToolMarkerOptional): +class InsertAtLineTool(EditingToolWithDiagnostics, ToolMarkerOptional, EditApiMixin): """ Inserts content at a given line in a file. """ @@ -530,18 +308,11 @@ class InsertAtLineTool(EditingToolWithDiagnostics, ToolMarkerOptional): :param line: the 0-based index of the line to insert content at :param content: the content to be inserted """ - # normalizing the inserted content - if not content.endswith("\n"): - content += "\n" - - with self.DiagnosticsContext(self, relative_path) as diagnostics_context: - code_editor = self.create_code_editor() - code_editor.insert_at_line(relative_path, line, content) - - return diagnostics_context.format_result(SUCCESS_RESULT) + with self.diagnostics_context(relative_path) as diagnostics_context: + return diagnostics_context.format_result(self._api().insert_at_line(relative_path, line, content)) -class SearchForPatternTool(Tool): +class SearchForPatternTool(Tool, FsApiMixin): def apply( self, substring_pattern: str, @@ -573,92 +344,19 @@ class SearchForPatternTool(Tool): ``-1`` uses the configured default. :return: A mapping from file paths to matched consecutive lines (0-based line numbers). """ - relative_path = relative_path.strip() - if relative_path: - self.project.validate_relative_path(relative_path) - - matches = self.project.search_project_files_for_pattern( - pattern=substring_pattern, - relative_path=relative_path, - context_lines_before=context_lines_before, - context_lines_after=context_lines_after, - paths_include_glob=paths_include_glob.strip(), - paths_exclude_glob=paths_exclude_glob.strip(), - multiline=multiline, - code_files_only=restrict_search_to_code_files, - skip_ignored_files=skip_ignored_files, + return ( + self._api() + .search_for_pattern( + substring_pattern, + context_lines_before=context_lines_before, + context_lines_after=context_lines_after, + paths_include_glob=paths_include_glob, + paths_exclude_glob=paths_exclude_glob, + relative_path=relative_path, + restrict_search_to_code_files=restrict_search_to_code_files, + skip_ignored_files=skip_ignored_files, + multiline=multiline, + max_answer_chars=max_answer_chars, + ) + .represent() ) - - # group matches by file - file_to_matches: dict[str, list[str]] = defaultdict(list) - for match in matches: - assert match.source_file_path is not None - file_to_matches[match.source_file_path].append(match.to_display_string()) - - # capture lightweight match data for shortening before serialization - match_lines_by_file: dict[str, list[dict[str, int | str]]] = defaultdict(list) - for match in matches: - assert match.source_file_path is not None - first = match.matched_lines[0] - match_lines_by_file[match.source_file_path].append({"line": first.line_number, "text": first.line_content.strip()}) - - # shortened result closures, from least to most aggressive shortening - _TEXT_TRUNCATE = 60 - - def render_first_lines(truncate: bool) -> str: - """Render each match's first line, either in full or truncated to a fixed length.""" - - def entry_text(text: str) -> str: - if truncate and len(text) > _TEXT_TRUNCATE: - return text[:_TEXT_TRUNCATE] + "..." - return text - - compact = { - path: [{"line": m["line"], "text": entry_text(str(m["text"]))} for m in lines] - for path, lines in match_lines_by_file.items() - } - if truncate: - header = ( - f"Matched lines (text over {_TEXT_TRUNCATE} chars is truncated, marked with a trailing '...'); " - "use read_file with the line numbers for full content:" - ) - else: - header = "Matched lines per file; use read_file with the line numbers for surrounding context:" - return f"{header}\n{self._to_json(compact)}" - - def make_first_lines_full() -> str: - """Match locations with each match's full first line.""" - return render_first_lines(truncate=False) - - def make_first_lines_truncated() -> str: - """Match locations with each match's first line truncated to a fixed length.""" - return render_first_lines(truncate=True) - - def make_line_numbers_only() -> str: - """Match locations as bare line numbers (no text).""" - numbers = {path: [m["line"] for m in lines] for path, lines in match_lines_by_file.items()} - return f"Match lines per file:\n{self._to_json(numbers)}" - - def make_per_file_counts() -> str: - counts = {path: len(lines) for path, lines in match_lines_by_file.items()} - return f"Match counts per file:\n{self._to_json(counts)}" - - def make_summary() -> str: - return f"Found {len(matches)} matches in {len(match_lines_by_file)} files." - - result = self._to_json(file_to_matches) - return self._limit_length( - result, - max_answer_chars, - shortened_result_factories=[ - make_first_lines_full, - make_first_lines_truncated, - make_line_numbers_only, - make_per_file_counts, - make_summary, - ], - ) - - """ - Performs a search for a pattern in the project. - """ diff --git a/src/serena/tools/jetbrains_tools.py b/src/serena/tools/jetbrains_tools.py index 6ed276de..6f856ebc 100644 --- a/src/serena/tools/jetbrains_tools.py +++ b/src/serena/tools/jetbrains_tools.py @@ -1,29 +1,40 @@ # SPDX-License-Identifier: GPL-3.0-or-later import logging -from collections import Counter -from typing import Any, Literal +from typing import TYPE_CHECKING, Literal, cast -import serena.jetbrains.jetbrains_types as jb -from serena.code_editor import JetBrainsCodeEditor -from serena.jetbrains.jetbrains_plugin_client import JetBrainsPluginClient -from serena.jetbrains.jetbrains_types import SymbolDTO, SymbolDTOUtil -from serena.symbol import JetBrainsSymbolDictGrouper -from serena.tools import Tool, ToolMarkerBeta, ToolMarkerOptional, ToolMarkerSymbolicEdit, ToolMarkerSymbolicRead -from serena.util.text_utils import find_text_coordinates +from serena.symbol import SymbolDictGrouper +from serena.tools import Tool, ToolMarkerOptional, ToolMarkerSymbolicEdit, ToolMarkerSymbolicRead + +if TYPE_CHECKING: + from serena.repl.api.jb_api import JetBrainsApi log = logging.getLogger(__name__) -class JetBrainsFindSymbolTool(Tool, ToolMarkerSymbolicRead, ToolMarkerOptional): +class JetBrainsApiMixin: + """ + Mixin for tools which delegate to the JetBrains API. + The API is imported locally, since the API module refers to the tools (as corresponding tools). + """ + + def _api(self) -> "JetBrainsApi": + from serena.repl.api.jb_api import JetBrainsApi + + tool = cast(Tool, cast(object, self)) + return JetBrainsApi(tool.agent) + + +class JetBrainsFindSymbolTool(Tool, ToolMarkerSymbolicRead, ToolMarkerOptional, JetBrainsApiMixin): """ Performs a global (or local) search for symbols using the JetBrains backend """ - # groups top-level symbols only; children are grouped separately by _group_children_by_type - symbol_dict_grouper = JetBrainsSymbolDictGrouper( - ["relative_path", "type"], ["type"], collapse_singleton=True, map_name_path_to_name=True - ) + @property + def symbol_dict_grouper(self) -> SymbolDictGrouper: + from serena.repl.api.jb_api import JetBrainsApi + + return JetBrainsApi.find_symbol_grouper_ def apply( self, @@ -75,76 +86,40 @@ class JetBrainsFindSymbolTool(Tool, ToolMarkerSymbolicRead, ToolMarkerOptional): :param max_answer_chars: max characters for the result (-1 for default). If exceeded, no content/a shortened result is returned. :return: symbols matching the name. """ - # check input - # - pattern with only wildcards is invalid, but in some cases we delegate to the overview tool - if name_path_pattern.replace("*", "").replace("/", "") == "": - if relative_path: - if self.project.relative_path_exists(relative_path, require_file=True): - overview_tool = self.agent.get_tool(JetBrainsGetSymbolsOverviewTool) - overview_response = overview_tool.apply(relative_path, depth=depth) - return self._wrapped_tool_response( - overview_response, f"Wildcard-only pattern not admitted; used {overview_tool.get_name()} instead" - ) - raise ValueError("name_path_pattern must not be empty or contain only wildcards; consider using the overview tool") - - if include_body: - depth = 0 # ignore user-specified depth if body is requested - name_path_pattern = self._sanitize_input_param(name_path_pattern) - if relative_path: relative_path = self._sanitize_input_param(relative_path) - if relative_path == ".": - relative_path = None - if relative_path is not None and relative_path.startswith(jb.JB_EXTERNAL_FILE_PREFIX): - search_deps = True + # for a wildcard-only pattern restricted to a file, delegate to the overview tool + if name_path_pattern.replace("*", "").replace("/", "") == "" and relative_path: + if self.project.relative_path_exists(relative_path, require_file=True): + overview_tool = self.agent.get_tool(JetBrainsGetSymbolsOverviewTool) + overview_response = overview_tool.apply(relative_path, depth=depth) + return self._wrapped_tool_response( + overview_response, f"Wildcard-only pattern not admitted; used {overview_tool.get_name()} instead" + ) - with JetBrainsPluginClient.from_project(self.project) as client: - if include_body: - include_quick_info = False - include_documentation = False - else: - if include_info: - include_documentation = True - include_quick_info = False - else: - # If no additional information is requested, we still include the quick info (type signature) - include_documentation = False - include_quick_info = True - symbol_collection_response = client.find_symbol( - name_path=name_path_pattern, - relative_path=relative_path, + return ( + self._api() + .find_symbol( + name_path_pattern, depth=depth, + relative_path=relative_path, include_body=include_body, - include_documentation=include_documentation, - include_quick_info=include_quick_info, + include_info=include_info, search_deps=search_deps, + max_matches=max_matches, + max_answer_chars=max_answer_chars, ) - symbols = symbol_collection_response["symbols"] - - def create_shortened_result() -> str: - """Shortened results containing symbol types and identifiers (path + name_path) only, without children""" - dicts: list[SymbolDTO] = [ - {"name_path": s["name_path"], "type": s["type"], "relative_path": s["relative_path"]} for s in symbols - ] - grouped = self.symbol_dict_grouper.group(dicts) - return f"Names with paths:\n{self._to_json(grouped)}" - - n_matches = len(symbols) - if 0 < max_matches < n_matches: - return f"Matched {n_matches}>{max_matches=} symbols.\n" + create_shortened_result() - - grouped_symbols = self.symbol_dict_grouper.group(symbols) - result = self._to_json(grouped_symbols) - return self._limit_length(result, max_answer_chars, shortened_result_factories=[create_shortened_result]) + .represent() + ) @classmethod def get_param_aliases(cls) -> dict[str, str]: return {"name_path": "name_path_pattern"} -class JetBrainsMoveTool(Tool, ToolMarkerSymbolicEdit, ToolMarkerOptional, ToolMarkerBeta): +class JetBrainsMoveTool(Tool, ToolMarkerSymbolicEdit, ToolMarkerOptional, JetBrainsApiMixin): """ Moves a symbol, file or directory to a new location using the JetBrains backend, updating all references """ @@ -180,21 +155,11 @@ class JetBrainsMoveTool(Tool, ToolMarkerSymbolicEdit, ToolMarkerOptional, ToolMa :param target_relative_path: the relative path of the target directory or file. :param target_parent_name_path: the name path of the target parent symbol. """ - name_path = name_path or None - target_relative_path = target_relative_path or None - target_parent_name_path = target_parent_name_path or None relative_path = self._sanitize_input_param(relative_path) - with JetBrainsPluginClient.from_project(self.project) as client: - response_dict = client.move( - name_path=name_path, - relative_path=relative_path, - target_parent_name_path=target_parent_name_path, - target_relative_path=target_relative_path, - ) - return self._to_json(response_dict) + return self._api().move(relative_path, name_path, target_relative_path, target_parent_name_path).represent() -class JetBrainsSafeDeleteTool(Tool, ToolMarkerSymbolicEdit, ToolMarkerOptional, ToolMarkerBeta): +class JetBrainsSafeDeleteTool(Tool, ToolMarkerSymbolicEdit, ToolMarkerOptional, JetBrainsApiMixin): """ Safely deletes a symbol using the JetBrains backend, checking for remaining usages first """ @@ -223,18 +188,10 @@ class JetBrainsSafeDeleteTool(Tool, ToolMarkerSymbolicEdit, ToolMarkerOptional, remove symbols that become unused after the deletion. Default is False. """ relative_path = self._sanitize_input_param(relative_path) - name_path = name_path or None - with JetBrainsPluginClient.from_project(self.project) as client: - response_dict = client.safe_delete( - name_path=name_path, - relative_path=relative_path, - delete_even_if_used=delete_even_if_used, - propagate=propagate, - ) - return self._to_json(response_dict) + return self._api().safe_delete(relative_path, name_path, delete_even_if_used, propagate).represent() -class JetBrainsInlineSymbol(Tool, ToolMarkerSymbolicEdit, ToolMarkerOptional, ToolMarkerBeta): +class JetBrainsInlineSymbol(Tool, ToolMarkerSymbolicEdit, ToolMarkerOptional, JetBrainsApiMixin): """ Inlines a symbol using the JetBrains backend, replacing all call sites with the symbol's body """ @@ -258,21 +215,19 @@ class JetBrainsInlineSymbol(Tool, ToolMarkerSymbolicEdit, ToolMarkerOptional, To May be ignored in some cases (e.g. when inlining a class). """ relative_path = self._sanitize_input_param(relative_path) - with JetBrainsPluginClient.from_project(self.project) as client: - response_dict = client.inline_symbol( - name_path=name_path, - relative_path=relative_path, - keep_definition=keep_definition, - ) - return self._to_json(response_dict) + return self._api().inline_symbol(name_path, relative_path, keep_definition).represent() -class JetBrainsFindReferencingSymbolsTool(Tool, ToolMarkerSymbolicRead, ToolMarkerOptional): +class JetBrainsFindReferencingSymbolsTool(Tool, ToolMarkerSymbolicRead, ToolMarkerOptional, JetBrainsApiMixin): """ Finds symbols that reference the given symbol using the JetBrains backend """ - symbol_dict_grouper = JetBrainsSymbolDictGrouper(["relative_path", "type"], ["type"], collapse_singleton=True) + @property + def symbol_dict_grouper(self) -> SymbolDictGrouper: + from serena.repl.api.jb_api import JetBrainsApi + + return JetBrainsApi.references_grouper_ def apply( self, @@ -292,52 +247,19 @@ class JetBrainsFindReferencingSymbolsTool(Tool, ToolMarkerSymbolicRead, ToolMark :param max_answer_chars: max characters for the result (-1 for default). If exceeded, no content/a shortened result is returned. """ relative_path = self._sanitize_input_param(relative_path) - with JetBrainsPluginClient.from_project(self.project) as client: - response_dict = client.find_references( - name_path=name_path, - relative_path=relative_path, - include_quick_info=False, - ) - symbol_dicts = response_dict["symbols"] - - # replace reference line number (if present) by actual line/context - for symbol_dict in symbol_dicts: - if "reference_line_no" in symbol_dict: - ref_line = symbol_dict["reference_line_no"] - ref_relative_path = symbol_dict["relative_path"] - if not SymbolDTOUtil.is_external_symbol(symbol_dict) and ref_line is not None and ref_line >= 0: - 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 - ) - symbol_dict["context"] = content_around_ref.to_display_string() - del symbol_dict["reference_line_no"] - - # capture file paths before grouping - ref_paths = [s.get("relative_path", "unknown") for s in symbol_dicts] - - result = self.symbol_dict_grouper.group(symbol_dicts) - - def create_shortened_result_counts_per_file() -> str: - return f"Reference counts per file:\n{self._to_json(Counter(ref_paths))}" - - def create_shortened_result_num_results() -> str: - return f"Found {len(ref_paths)} references." - - result_json = self._to_json(result) - return self._limit_length( - result_json, - max_answer_chars, - shortened_result_factories=[create_shortened_result_counts_per_file, create_shortened_result_num_results], - ) + return self._api().find_referencing_symbols(name_path, relative_path, max_answer_chars).represent() -class JetBrainsGetSymbolsOverviewTool(Tool, ToolMarkerSymbolicRead, ToolMarkerOptional): +class JetBrainsGetSymbolsOverviewTool(Tool, ToolMarkerSymbolicRead, ToolMarkerOptional, JetBrainsApiMixin): """ Retrieves an overview of the top-level symbols within a specified file using the JetBrains backend """ - USE_COMPACT_FORMAT = True - symbol_dict_grouper = JetBrainsSymbolDictGrouper(["type"], ["type"], collapse_singleton=True, map_name_path_to_name=True) + @property + def symbol_dict_grouper(self) -> SymbolDictGrouper: + from serena.repl.api.jb_api import JetBrainsApi + + return JetBrainsApi.overview_grouper_ def apply( self, @@ -357,95 +279,15 @@ class JetBrainsGetSymbolsOverviewTool(Tool, ToolMarkerSymbolicRead, ToolMarkerOp :param max_answer_chars: max characters for the result (-1 for default). If exceeded, no content/a shortened result is returned. :param include_file_documentation: whether to include the file's docstring. Default False. """ - if depth == -1: - if relative_path.endswith((".java", ".kt")): - depth = 1 - else: - depth = 0 - relative_path = self._sanitize_input_param(relative_path) - with JetBrainsPluginClient.from_project(self.project) as client: - symbol_overview = client.get_symbols_overview( - relative_path=relative_path, depth=depth, include_file_documentation=include_file_documentation - ) - - if self.USE_COMPACT_FORMAT: - symbols = symbol_overview["symbols"] - - grouped_symbols = self.symbol_dict_grouper.group(symbols) - - shortened_result_factories = [] - - # create full result - result: dict[str, Any] = {"symbols": grouped_symbols} - documentation = symbol_overview.pop("documentation", None) - if documentation: - result["docstring"] = documentation - shortened_result_factories.append(lambda: self._to_json(grouped_symbols)) # shortened result without docstring - json_result = self._to_json(result) - - if depth > 0: - - def create_short_result_depth_0() -> str: - depth_0_symbols = [d.copy() for d in symbols] - for d in depth_0_symbols: - d.pop("children", None) - compact_depth_0_result = self.symbol_dict_grouper.group(depth_0_symbols) - return "Depth 0 overview:\n" + self._to_json(compact_depth_0_result) - - shortened_result_factories.append(create_short_result_depth_0) - - def create_short_result_type_counts() -> str: - type_names = [d.get("type", "unknown") for d in symbols] - return f"Symbol counts by type:\n{self._to_json(Counter(type_names))}" - - shortened_result_factories.append(create_short_result_type_counts) - else: - # this path is currently abandoned, consider introducing shortened results if ever needed - shortened_result_factories = None - json_result = self._to_json(symbol_overview) - - return self._limit_length(json_result, max_answer_chars, shortened_result_factories=shortened_result_factories) + return self._api().get_symbols_overview(relative_path, depth, max_answer_chars, include_file_documentation).represent() -class JetBrainsTypeHierarchyTool(Tool, ToolMarkerSymbolicRead, ToolMarkerOptional): +class JetBrainsTypeHierarchyTool(Tool, ToolMarkerSymbolicRead, ToolMarkerOptional, JetBrainsApiMixin): """ Retrieves the type hierarchy (supertypes and/or subtypes) of a symbol using the JetBrains backend """ - @staticmethod - def _transform_hierarchy_nodes(nodes: list[jb.TypeHierarchyNodeDTO] | None) -> dict[str, list]: - """ - Transform a list of TypeHierarchyNode into a file-grouped compact format. - - Returns a dict where keys are relative_paths and values are lists of either: - - "SymbolNamePath" (leaf node) - - {"SymbolNamePath": {nested_file_grouped_children}} (node with children) - """ - if not nodes: - return {} - - result: dict[str, list] = {} - - for node in nodes: - symbol = node["symbol"] - name_path = symbol["name_path"] - rel_path = symbol["relative_path"] - children = node.get("children", []) - - if rel_path not in result: - result[rel_path] = [] - - if children: - # Node with children - recurse - nested = JetBrainsTypeHierarchyTool._transform_hierarchy_nodes(children) - result[rel_path].append({name_path: nested}) - else: - # Leaf node - result[rel_path].append(name_path) - - return result - def apply( self, name_path: str, @@ -460,51 +302,17 @@ class JetBrainsTypeHierarchyTool(Tool, ToolMarkerSymbolicRead, ToolMarkerOptiona :param name_path: name path of the symbol for which to get the type hierarchy. :param relative_path: the relative path to the file containing the symbol. :param hierarchy_type: which hierarchy to retrieve: "super" for parent classes/interfaces, - "sub" for subclasses/implementations, or "both" for both directions. Default is "sub". + "sub" for subclasses/implementations, or "both" for both directions. Default is "both". :param depth: depth limit for hierarchy traversal (None or 0 for unlimited). Default is 1. :param max_answer_chars: max characters for the JSON result. If exceeded, no content is returned. -1 means the default value from the config will be used. :return: Compact JSON with file-grouped hierarchy. Error string if not applicable. """ relative_path = self._sanitize_input_param(relative_path) - with JetBrainsPluginClient.from_project(self.project) as client: - subtypes = None - supertypes = None - levels_not_included = {} - - if hierarchy_type in ("super", "both"): - supertypes_response = client.get_supertypes( - name_path=name_path, - relative_path=relative_path, - depth=depth, - ) - if "num_levels_not_included" in supertypes_response: - levels_not_included["supertypes"] = supertypes_response["num_levels_not_included"] - supertypes = self._transform_hierarchy_nodes(supertypes_response.get("hierarchy")) - - if hierarchy_type in ("sub", "both"): - subtypes_response = client.get_subtypes( - name_path=name_path, - relative_path=relative_path, - depth=depth, - ) - if "num_levels_not_included" in subtypes_response: - levels_not_included["subtypes"] = subtypes_response["num_levels_not_included"] - subtypes = self._transform_hierarchy_nodes(subtypes_response.get("hierarchy")) - - result_dict: dict[str, dict | list] = {} - if supertypes is not None: - result_dict["supertypes"] = supertypes - if subtypes is not None: - result_dict["subtypes"] = subtypes - if levels_not_included: - result_dict["levels_not_included"] = levels_not_included - - result = self._to_json(result_dict) - return self._limit_length(result, max_answer_chars) + return self._api().get_type_hierarchy(name_path, relative_path, hierarchy_type, depth, max_answer_chars).represent() -class JetBrainsFindDeclarationTool(Tool, ToolMarkerSymbolicRead, ToolMarkerOptional): +class JetBrainsFindDeclarationTool(Tool, ToolMarkerSymbolicRead, ToolMarkerOptional, JetBrainsApiMixin): """ Finds the declaration of a symbol using the JetBrains backend """ @@ -523,20 +331,10 @@ class JetBrainsFindDeclarationTool(Tool, ToolMarkerSymbolicRead, ToolMarkerOptio """ relative_path = self._sanitize_input_param(relative_path) regex = self._sanitize_input_param(regex) - - editor = self.create_code_editor() - content = editor.read_file(relative_path) - coords = find_text_coordinates(content, regex, require_unique=True) - assert coords is not None - with JetBrainsPluginClient.from_project(self.project) as client: - symbol_collection = client.find_declaration( - relative_path=relative_path, line=coords.line, col=coords.col, include_quick_info=False, include_body=include_body - ) - result = self._to_json(symbol_collection) - return result + return self._api().find_declaration(relative_path, regex, include_body).represent() -class JetBrainsFindImplementationsTool(Tool, ToolMarkerSymbolicRead, ToolMarkerOptional): +class JetBrainsFindImplementationsTool(Tool, ToolMarkerSymbolicRead, ToolMarkerOptional, JetBrainsApiMixin): """ Finds the implementations of a symbol using the JetBrains backend """ @@ -548,17 +346,10 @@ class JetBrainsFindImplementationsTool(Tool, ToolMarkerSymbolicRead, ToolMarkerO :param relative_path: the relative path to the source file containing the symbol for which to find implementations. :param name_path: name path of the symbol for which to find implementations """ - with JetBrainsPluginClient.from_project(self.project) as client: - symbol_collection = client.find_implementations( - relative_path=relative_path, - name_path=name_path, - include_quick_info=False, - ) - result = self._to_json(symbol_collection) - return result + return self._api().find_implementations(relative_path, name_path).represent() -class JetBrainsRenameTool(Tool, ToolMarkerSymbolicEdit, ToolMarkerOptional): +class JetBrainsRenameTool(Tool, ToolMarkerSymbolicEdit, ToolMarkerOptional, JetBrainsApiMixin): """ Renames a symbol, file or directory throughout the codebase using the JetBrains backend. """ @@ -584,18 +375,10 @@ class JetBrainsRenameTool(Tool, ToolMarkerSymbolicEdit, ToolMarkerOptional): :param rename_in_text_occurrences: whether to also rename occurrences in text. Default True. :return: a status message """ - code_editor = JetBrainsCodeEditor(self.project) - result = code_editor.rename_symbol( - name_path=name_path, - relative_path=relative_path, - new_name=new_name, - rename_in_comments=rename_in_comments, - rename_in_text_occurrences=rename_in_text_occurrences, - ) - return self._to_json(result) + return self._api().rename(relative_path, new_name, name_path, rename_in_comments, rename_in_text_occurrences).represent() -class JetBrainsDebugTool(Tool, ToolMarkerOptional, ToolMarkerBeta): +class JetBrainsDebugTool(Tool, ToolMarkerOptional, JetBrainsApiMixin): """ Provides debugging functionality (run configs, breakpoints, stepping, inspection, and evaluation) via a persistent debug REPL connected to the JetBrains IDE. @@ -617,15 +400,10 @@ class JetBrainsDebugTool(Tool, ToolMarkerOptional, ToolMarkerBeta): :param repl_key: identifier for the REPL instance. State persists across calls with the same key. :return: string representation of the result """ - with JetBrainsPluginClient.from_project(self.project) as client: - if expression: - response = client.debug_eval(repl_key=repl_key, expression=expression) - else: - response = client.debug_close(repl_key=repl_key) - return response.get("result", str(response)) + return self._api().debug_eval(expression, repl_key) -class JetBrainsRunInspectionsTool(Tool, ToolMarkerSymbolicRead, ToolMarkerOptional): +class JetBrainsRunInspectionsTool(Tool, ToolMarkerSymbolicRead, ToolMarkerOptional, JetBrainsApiMixin): """ Runs JetBrains IDE inspections on a file and returns the results. """ @@ -655,19 +433,12 @@ class JetBrainsRunInspectionsTool(Tool, ToolMarkerSymbolicRead, ToolMarkerOption -1 means the default value from the config will be used. :return: JSON string with inspection results including severity, message, and location. """ - with JetBrainsPluginClient.from_project(self.project) as client: - response_dict = client.run_inspections( - relative_path=relative_path, - min_severity=min_severity, - inspection_names=inspection_names, - start_line=start_line, - end_line=end_line, - ) - result = self._to_json(response_dict) - return self._limit_length(result, max_answer_chars) + return ( + self._api().run_inspections(relative_path, min_severity, inspection_names, start_line, end_line, max_answer_chars).represent() + ) -class JetBrainsListInspectionsTool(Tool, ToolMarkerSymbolicRead, ToolMarkerOptional): +class JetBrainsListInspectionsTool(Tool, ToolMarkerSymbolicRead, ToolMarkerOptional, JetBrainsApiMixin): """ Lists available JetBrains IDE inspections, optionally filtered by language or group. """ @@ -689,10 +460,4 @@ class JetBrainsListInspectionsTool(Tool, ToolMarkerSymbolicRead, ToolMarkerOptio -1 means the default value from the config will be used. :return: JSON string with the list of available inspections including name, group path, and language. """ - with JetBrainsPluginClient.from_project(self.project) as client: - response_dict = client.list_inspections( - language=language, - group_path_contains=group_path_contains, - ) - result = self._to_json(response_dict) - return self._limit_length(result, max_answer_chars) + return self._api().list_inspections(language, group_path_contains, max_answer_chars).represent() diff --git a/src/serena/tools/memory_tools.py b/src/serena/tools/memory_tools.py index a6e78103..cfbbf3e4 100644 --- a/src/serena/tools/memory_tools.py +++ b/src/serena/tools/memory_tools.py @@ -1,14 +1,27 @@ # SPDX-License-Identifier: GPL-3.0-or-later -import logging -from typing import Literal +from typing import TYPE_CHECKING, Literal, cast from serena.tools import Tool, ToolMarkerCanEdit -log = logging.getLogger(__name__) +if TYPE_CHECKING: + from serena.repl.api.mem_api import MemoryApi -class WriteMemoryTool(Tool, ToolMarkerCanEdit): +class MemoryApiMixin: + """ + Mixin for tools which delegate to the memory API. + The API is imported locally, since the API module refers to the tools (as corresponding tools). + """ + + def _api(self) -> "MemoryApi": + from serena.repl.api.mem_api import MemoryApi + + tool = cast(Tool, cast(object, self)) + return MemoryApi(tool.agent) + + +class WriteMemoryTool(Tool, ToolMarkerCanEdit, MemoryApiMixin): """ Write some information (utf-8-encoded) about this project that can be useful for future tasks to a memory in md format. The memory name should be meaningful. @@ -26,18 +39,10 @@ class WriteMemoryTool(Tool, ToolMarkerCanEdit): :param content: memory content, utf8-encoded :param max_chars: see other tools """ - # NOTE: utf-8 encoding is configured in the MemoriesManager - if max_chars == -1: - max_chars = self.agent.serena_config.default_max_tool_answer_chars - if len(content) > max_chars: - raise ValueError( - f"Content for {memory_name} is too long. Max length is {max_chars} characters. " + "Please make the content shorter." - ) - - return self.memory_manager.save_memory(memory_name, content, is_tool_context=True) + return self._api().write_memory(memory_name, content, max_chars) -class ReadMemoryTool(Tool): +class ReadMemoryTool(Tool, MemoryApiMixin): """ Reads the content of a memory file. """ @@ -46,10 +51,10 @@ class ReadMemoryTool(Tool): """ Use to read a memory that is likely to be relevant to the current task, inferring relevance e.g. from the name. """ - return self.memory_manager.load_memory(memory_name) + return self._api().read_memory(memory_name) -class ListMemoriesTool(Tool): +class ListMemoriesTool(Tool, MemoryApiMixin): """ Lists available memories. """ @@ -58,10 +63,10 @@ class ListMemoriesTool(Tool): """ Lists available memories, optionally filtered by topic. """ - return self._to_json(self.memory_manager.list_memories(topic).to_dict()) + return self._api().list_memories(topic).represent() -class DeleteMemoryTool(Tool, ToolMarkerCanEdit): +class DeleteMemoryTool(Tool, ToolMarkerCanEdit, MemoryApiMixin): """ Delete a memory file. """ @@ -70,10 +75,10 @@ class DeleteMemoryTool(Tool, ToolMarkerCanEdit): """ Delete a memory, only call if instructed explicitly or permission was granted by the user. """ - return self.memory_manager.delete_memory(memory_name, is_tool_context=True) + return self._api().delete_memory(memory_name) -class RenameMemoryTool(Tool, ToolMarkerCanEdit): +class RenameMemoryTool(Tool, ToolMarkerCanEdit, MemoryApiMixin): """ Renames or moves a memory, updating references that are marked with the `mem:` prefix. """ @@ -85,15 +90,10 @@ class RenameMemoryTool(Tool, ToolMarkerCanEdit): References to other memories that are marked with the `mem:` prefix will be updated accordingly. References in read-only memories are not affected. """ - renaming_message, n_references_updated = self.memory_manager.rename_memory_and_propagate_references( - old_name, new_name, is_tool_context=True - ) - if n_references_updated > 0: - log.info(f"Updated {n_references_updated} references to memory {old_name} to {new_name}") - return renaming_message + return self._api().rename_memory(old_name, new_name) -class EditMemoryTool(Tool, ToolMarkerCanEdit): +class EditMemoryTool(Tool, ToolMarkerCanEdit, MemoryApiMixin): """ Replaces content matching a regular expression in a memory. """ @@ -119,6 +119,4 @@ class EditMemoryTool(Tool, ToolMarkerCanEdit): :param allow_multiple_occurrences: whether to allow matching and replacing multiple occurrences. If false and multiple occurrences are found, an error will be returned. """ - return self.memory_manager.edit_memory( - memory_name, needle, repl, mode, allow_multiple_occurrences, is_tool_context=True, regex_multiline=True - ) + return self._api().edit_memory(memory_name, needle, repl, mode, allow_multiple_occurrences) diff --git a/src/serena/tools/query_project_tools.py b/src/serena/tools/query_project_tools.py index 514454b8..781ccc27 100644 --- a/src/serena/tools/query_project_tools.py +++ b/src/serena/tools/query_project_tools.py @@ -2,7 +2,6 @@ import json -from serena.config.serena_config import LanguageBackend from serena.jetbrains.jetbrains_plugin_client import JetBrainsPluginClientManager from serena.project_server import ProjectServerClient from serena.tools import Tool, ToolMarkerDoesNotRequireActiveProject, ToolMarkerOptional @@ -57,7 +56,7 @@ class QueryProjectTool(Tool, ToolMarkerOptional, ToolMarkerDoesNotRequireActiveP tool = self.agent.get_tool_by_name(tool_name) assert tool.is_readonly(), f"Tool {tool_name} is not read-only and cannot be executed in another project." if self._is_project_server_required(tool): - client = ProjectServerClient() + client = ProjectServerClient(self.agent.serena_config) return client.query_project(project_name, tool_name, tool_params_json) else: registered_project = self.agent.serena_config.get_registered_project(project_name) @@ -67,13 +66,10 @@ class QueryProjectTool(Tool, ToolMarkerOptional, ToolMarkerDoesNotRequireActiveP return tool.apply(**json.loads(tool_params_json)) def _is_project_server_required(self, tool: Tool) -> bool: - match self.agent.get_language_backend(): - case LanguageBackend.JETBRAINS: - return False - case LanguageBackend.LSP: - # Note: As long as only read-only tools are considered, only symbolic tools require the project server. - # But if we were to allow non-read-only tools, then tools using a CodeEditor also indirectly require language servers. - assert tool.is_readonly() - return tool.is_symbolic() - case _: - raise NotImplementedError + # The project server is relevant to the LSP backend only + if not self.agent.get_language_backend().is_lsp(): + return False + # Note: As long as only read-only tools are considered, only symbolic tools require the project server. + # But if we were to allow non-read-only tools, then tools using a CodeEditor also indirectly require language servers. + assert tool.is_readonly() + return tool.is_symbolic() diff --git a/src/serena/tools/repl_tools.py b/src/serena/tools/repl_tools.py new file mode 100644 index 00000000..14df5c61 --- /dev/null +++ b/src/serena/tools/repl_tools.py @@ -0,0 +1,54 @@ +""" +Tools which provide access to Serena's functionality through Python code execution +""" + +# SPDX-License-Identifier: GPL-3.0-or-later + +from serena.tools.tools_base import Tool, ToolMarkerBeta, ToolMarkerOptional + + +class SerenaReplTool(Tool, ToolMarkerOptional, ToolMarkerBeta): + """ + Executes Python code which accesses Serena's functionality programmatically. + """ + + def get_apply_docstring(self) -> str: + docs = self.get_apply_docstring_from_cls() + if self.agent.is_single_project(): + docs += "\n\nAvailable facades:\n" + self.agent.get_repl().entrypoint.overview() + else: + docs += "\n\nAvailable facades are provided at project activation" + return docs + + def apply(self, session_id: str, code: str) -> str: + """ + Executes the given Python code, which has access to Serena's functionality through the object `s`. + The functionality is organised in facades, which are attributes of `s` (e.g. `s.myfacade`). + + Documentation: Use `s.info("")` when you will use a facade's functionality (it documents all common + operations at once) and `s.info(".")` for a single or a rarely needed operation. Several items + can be requested in one call, e.g. `s.info("lsp", "edit.replace_content")`. + `s.info("")` documents the facade's operations only, not their result types. The facade listing + provides result types (`method -> Type`); request their documentation via `s.info("")`, which + includes the types they contain, ONLY if you intend to process results in code (filter, aggregate, chain + calls). + + The code is executed like a notebook cell: if its last statement is an expression, the expression's value is + the result. Results are rendered in a form suitable for you; lists are rendered element-wise, + strings are passed through unchanged. + + Output size: methods with a `max_answer_chars` parameter limit the size of the rendered result (-1 uses the + configured default). If the limit is exceeded, a shortened result (or no content) is rendered instead; + adjust the limit only if there is no other way to obtain the content required for the task (e.g. by + narrowing the query or processing the result in code and returning only what is needed). + + Persistence: variables, functions and classes defined at the top level of your code persist across calls + within your session (like the cells of a notebook), so you can reuse results and define helper functions once. + `s.vars()` lists the persisted items, `s.clear()` removes them. Do not store facades (`s.`) in + variables; access them via `s` at call time. Do not keep large results longer than needed. + + :param session_id: your Serena session id, as provided in Serena's instructions (call `initial_instructions` if you do not have one) + :param code: the Python code to execute + :return: the representation of the returned value, or the error if execution failed + """ + return self.agent.get_repl().execute(code, self.agent.get_session(session_id)) diff --git a/src/serena/tools/symbol_tools.py b/src/serena/tools/symbol_tools.py index 1ab40e0a..9712b821 100644 --- a/src/serena/tools/symbol_tools.py +++ b/src/serena/tools/symbol_tools.py @@ -3,43 +3,55 @@ Language server-related tools """ # SPDX-License-Identifier: GPL-3.0-or-later -import copy -import os -from collections import Counter, defaultdict -from collections.abc import Sequence -from typing import Any +from typing import TYPE_CHECKING, cast -from serena.symbol import LanguageServerSymbol, LanguageServerSymbolDictGrouper +from serena.symbol import SymbolDictGrouper from serena.tools import ( - SUCCESS_RESULT, EditingToolWithDiagnostics, Tool, ToolMarkerSymbolicEdit, ToolMarkerSymbolicRead, ) +from serena.tools.file_tools import EditApiMixin from serena.tools.tools_base import ToolMarkerOptional -from serena.util.ls_diagnostics import GroupedDiagnostics -from serena.util.text_utils import find_text_coordinates -from solidlsp.ls_types import SymbolKind + +if TYPE_CHECKING: + from serena.repl.api.lsp_api import LspApi -class RestartLanguageServerTool(Tool, ToolMarkerOptional): +class LspApiMixin: + """ + Mixin for tools which delegate to the language server API. + The API is imported locally, since the API module refers to the tools (as corresponding tools). + """ + + def _api(self) -> "LspApi": + from serena.repl.api.lsp_api import LspApi + + tool = cast(Tool, cast(object, self)) + return LspApi(tool.agent) + + +class RestartLanguageServerTool(Tool, ToolMarkerOptional, LspApiMixin): """Restarts the language server(s).""" def apply(self) -> str: """Use this tool only on explicit user request or after confirmation. It may be necessary to restart the language server if it hangs. """ - self.agent.reset_language_server_manager() - return SUCCESS_RESULT + return self._api().restart_language_server() -class GetSymbolsOverviewTool(Tool, ToolMarkerSymbolicRead): +class GetSymbolsOverviewTool(Tool, ToolMarkerSymbolicRead, LspApiMixin): """ Gets an overview of the top-level symbols defined in a given file. """ - symbol_dict_grouper = LanguageServerSymbolDictGrouper(["kind"], ["kind"], collapse_singleton=True) + @property + def symbol_dict_grouper(self) -> SymbolDictGrouper: + from serena.repl.api.lsp_api import LspApi + + return LspApi.overview_grouper_ def apply(self, relative_path: str, depth: int = -1, max_answer_chars: int = -1) -> str: """ @@ -55,93 +67,20 @@ class GetSymbolsOverviewTool(Tool, ToolMarkerSymbolicRead): Don't adjust unless there is really no other way to get the content required for the task. :return: a JSON object containing symbols grouped by kind in a compact format. """ - # Note: file system sync not required (relevant file is opened in the language server explicitly) - - if depth == -1: - if relative_path.endswith((".java", ".kt")): - depth = 1 - else: - depth = 0 - - result = self.get_symbol_overview(relative_path, depth=depth) - - # capture kind names and depth-0 snapshots before grouping, which mutates the dicts - kind_names = [d.get("kind", "unknown") for d in result] - if depth > 0: - depth_0_result = [d.copy() for d in result] - for d in depth_0_result: - d.pop("children", None) - - compact_result = self.symbol_dict_grouper.group(result) - result_json_str = self._to_json(compact_result) - - # shortened result closures - def make_kind_counts() -> str: - return f"Symbol counts by kind:\n{self._to_json(Counter(kind_names))}" - - if depth == 0: - shortened_results = [make_kind_counts] - else: - - def make_depth_0_result() -> str: - compact_depth_0_result = self.symbol_dict_grouper.group(depth_0_result) - return "Depth 0 overview:\n" + self._to_json(compact_depth_0_result) - - shortened_results = [make_depth_0_result, make_kind_counts] - - return self._limit_length(result_json_str, max_answer_chars, shortened_result_factories=shortened_results) - - def get_symbol_overview(self, relative_path: str, depth: int = 0) -> list[LanguageServerSymbol.OutputDict]: - """ - :param relative_path: relative path to a source file - :param depth: the depth up to which descendants shall be retrieved - :return: a list of symbol dictionaries representing the symbol overview of the file - """ - symbol_retriever = self.create_language_server_symbol_retriever() - - # The symbol overview is capable of working with both files and directories, - # but we want to ensure that the user provides a file path. - file_path = os.path.join(self.project.project_root, relative_path) - if not os.path.exists(file_path): - raise FileNotFoundError(f"File or directory {relative_path} does not exist in the project.") - if os.path.isdir(file_path): - raise ValueError(f"Expected a file path, but got a directory path: {relative_path}. ") - if not symbol_retriever.can_analyze_file(relative_path): - raise ValueError( - f"Cannot extract symbols from file {relative_path}. Active language servers: {[l.value for l in self.agent.get_active_language_server_ids()]}" - ) - - symbols = symbol_retriever.get_symbol_overview(relative_path)[relative_path] - - def child_inclusion_predicate(s: LanguageServerSymbol) -> bool: - return not s.is_low_level() - - symbol_dicts = [] - for symbol in symbols: - symbol_dicts.append( - symbol.to_dict( - name_path=False, - name=True, - depth=depth, - kind=True, - relative_path=False, - location=False, - child_inclusion_predicate=child_inclusion_predicate, - ) - ) - return symbol_dicts + return self._api().get_symbols_overview(relative_path, depth=depth, max_answer_chars=max_answer_chars).represent() -class FindSymbolTool(Tool, ToolMarkerSymbolicRead): +class FindSymbolTool(Tool, ToolMarkerSymbolicRead, LspApiMixin): """ Performs a global (or local) search using the language server backend. """ - # group children by kind, keeping just the name (the parent's name_path makes it unambiguous); - # we don't group the top-level result list because many tests rely on it being a flat list of symbol dicts - symbol_dict_grouper = LanguageServerSymbolDictGrouper([], ["kind"], collapse_singleton=True) + @property + def symbol_dict_grouper(self) -> SymbolDictGrouper: + from serena.repl.api.lsp_api import LspApi + + return LspApi.find_symbol_dict_grouper_ - # noinspection PyDefaultArgument def apply( self, name_path_pattern: str, @@ -179,85 +118,54 @@ class FindSymbolTool(Tool, ToolMarkerSymbolicRead): :param relative_path: (optional) restrict search to this file or directory. If None, searches entire codebase. If a directory is passed, the search will be restricted to the files in that directory. If a file is passed, the search will be restricted to that file. - :param include_body: whether to include the symbol's source code. Use judiciously. + If you have some knowledge about the codebase, you should use this parameter, as it will significantly + speed up the search as well as reduce the number of results. + :param include_body: If True, include the symbol's source code. Use judiciously. :param include_info: whether to include additional info (hover-like, typically including docstring and signature), about the symbol (ignored if include_body is True). Info is never included for child symbols. Note: Depending on the language, this can be slow (e.g., C/C++). :param include_kinds: (optional) limits results to the given LSP symbol kinds (integers) :param exclude_kinds: (optional) list of LSP symbol kinds (integers) to exclude. - :param substring_matching: If True, use substring matching for the last element of the pattern, such that - "Foo/get" would match "Foo/getValue" and "Foo/getData". - :param max_matches: maximum number of permitted matches. If exceeded, a shortened result is returned - which allows refining the search. -1 (default) means no limit. Set to 1 if you search for a single symbol. + :param substring_matching: If True, use substring matching for the last segment of `name_path_pattern` + (i.e. the name of the symbol, e.g. "foo" in "Class/foo" or "my_method" in "my_method"). + :param max_matches: Maximum number of permitted matches. If exceeded, a shortened result is returned + which allows refining the search. -1 (default) means no limit. Set to 1 to search for a unique symbol. :param max_answer_chars: max result length; -1 for default :return: symbols (with locations) matching the name. """ - # Note: file system sync not required; the symbol finder opens all relevant source files explicitly in the case of changes - - if include_body: - depth = 0 # ignore user-specified depth if include_body is True - assert max_matches != 0, "max_matches must be > 0 or equal to -1." - parsed_include_kinds: Sequence[SymbolKind] | None = [SymbolKind(k) for k in include_kinds] if include_kinds else None - parsed_exclude_kinds: Sequence[SymbolKind] | None = [SymbolKind(k) for k in exclude_kinds] if exclude_kinds else None - symbol_retriever = self.create_language_server_symbol_retriever() - symbols = symbol_retriever.find( - name_path_pattern, - include_kinds=parsed_include_kinds, - exclude_kinds=parsed_exclude_kinds, - substring_matching=substring_matching, - within_relative_path=relative_path, - ) - n_matches = len(symbols) - - def create_short_result_relative_path_to_name_paths() -> str: - relative_path_to_name_paths: defaultdict[str, list[str]] = defaultdict(list) - for s in symbols: - relative_path_to_name_paths[s.location.relative_path or "unknown"].append(s.get_name_path()) - return f"Shortened result:\n{self._to_json(relative_path_to_name_paths)}" - - if 0 < max_matches < n_matches: - return f"Matched {n_matches}>{max_matches=} symbols.\n" + create_short_result_relative_path_to_name_paths() - - symbol_dicts = [ - s.to_dict( - kind=True, - name_path=True, - name=False, - relative_path=True, - body_location=True, + return ( + self._api() + .find_symbol( + name_path_pattern, depth=depth, - body=include_body, - children_name=True, - children_name_path=False, + relative_path=relative_path, + include_body=include_body, + include_info=include_info, + include_kinds=include_kinds, + exclude_kinds=exclude_kinds, + substring_matching=substring_matching, + max_matches=max_matches, + max_answer_chars=max_answer_chars, ) - for s in symbols - ] - if not include_body and include_info: - info_by_symbol = symbol_retriever.request_info_for_symbol_batch(symbols) - for s, s_dict in zip(symbols, symbol_dicts, strict=True): - if symbol_info := info_by_symbol.get(s): - # In python 3.15 we could specify extra_items=True in the TypedDict definition, - # https://peps.python.org/pep-0728/ - # If we ever upgrade to 3.15, we can remove the type: ignore[typeddict-unknown-key] - s_dict["info"] = symbol_info - - grouped_symbol_dicts = self.symbol_dict_grouper.group(symbol_dicts) - result = self._to_json(grouped_symbol_dicts) - return self._limit_length(result, max_answer_chars, shortened_result_factories=[create_short_result_relative_path_to_name_paths]) + .represent() + ) @classmethod def get_param_aliases(cls) -> dict[str, str]: return {"name_path": "name_path_pattern"} -class FindReferencingSymbolsTool(Tool, ToolMarkerSymbolicRead): +class FindReferencingSymbolsTool(Tool, ToolMarkerSymbolicRead, LspApiMixin): """ - Finds symbols that reference the given symbol using the language server backend + Finds symbols that reference the given symbol """ - symbol_dict_grouper = LanguageServerSymbolDictGrouper(["relative_path", "kind"], ["kind"], collapse_singleton=True) + @property + def symbol_dict_grouper(self) -> SymbolDictGrouper: + from serena.repl.api.lsp_api import LspApi + + return LspApi.references_grouper_ - # noinspection PyDefaultArgument def apply( self, name_path: str, @@ -277,75 +185,20 @@ class FindReferencingSymbolsTool(Tool, ToolMarkerSymbolicRead): :param max_answer_chars: max result length; -1 for default :return: a list of JSON objects with the symbols referencing the requested symbol """ - # file system sync needed for case where symbol finder does not perform a global search, updating everything - if relative_path: - self.project.ls_sync_file_system_changes() - - include_body = False # It is probably never a good idea to include the body of the referencing symbols - parsed_include_kinds: Sequence[SymbolKind] | None = [SymbolKind(k) for k in include_kinds] if include_kinds else None - parsed_exclude_kinds: Sequence[SymbolKind] | None = [SymbolKind(k) for k in exclude_kinds] if exclude_kinds else None - - symbol_retriever = self.create_language_server_symbol_retriever() - references_in_symbols = symbol_retriever.find_referencing_symbols( - name_path, - relative_file_path=relative_path, - include_body=include_body, - include_kinds=parsed_include_kinds, - exclude_kinds=parsed_exclude_kinds, + return ( + self._api() + .find_referencing_symbols( + name_path, relative_path, include_kinds=include_kinds, exclude_kinds=exclude_kinds, max_answer_chars=max_answer_chars + ) + .represent() ) - reference_dicts = [] - for ref in references_in_symbols: - ref_dict_orig = ref.symbol.to_dict(kind=True, relative_path=True, depth=0, body=include_body, body_location=True) - ref_dict = dict(ref_dict_orig) - if not include_body: - ref_relative_path = ref.symbol.location.relative_path - assert ref_relative_path is not None, f"Referencing symbol {ref.symbol.name} has no relative path, this is likely a bug." - content_around_ref = self.project.retrieve_content_around_line( - relative_file_path=ref_relative_path, line=ref.line, context_lines_before=1, context_lines_after=1 - ) - ref_dict["content_around_reference"] = content_around_ref.to_display_string() - reference_dicts.append(ref_dict) - # capture lightweight reference data before grouping - ref_summaries = [] - for ref, d in zip(references_in_symbols, reference_dicts, strict=True): - ref_summaries.append( - { - "name_path": d.get("name_path"), - "kind": d.get("kind"), - "relative_path": d.get("relative_path"), - "reference_line": ref.line, - } - ) - - result = self.symbol_dict_grouper.group(reference_dicts) - - # shortened result closures, from least to most aggressive shortening - def make_refs_without_context() -> str: - """References with name_path and reference line, without surrounding code lines""" - grouped = self.symbol_dict_grouper.group(copy.deepcopy(ref_summaries)) - return f"References without surrounding lines:\n{self._to_json(grouped)}" - - def make_per_file_counts() -> str: - counts = Counter(str(r["relative_path"]) for r in ref_summaries) - return f"Reference counts per file:\n{self._to_json(counts)}" - - def make_summary() -> str: - return f"Found {len(ref_summaries)} references." - - shortened_results = [make_refs_without_context, make_per_file_counts, make_summary] - - result_json = self._to_json(result) - return self._limit_length(result_json, max_answer_chars, shortened_result_factories=shortened_results) - - -class FindImplementationsTool(Tool, ToolMarkerSymbolicRead): +class FindImplementationsTool(Tool, ToolMarkerSymbolicRead, LspApiMixin): """ - Finds symbols that implement the given symbol using the language server backend. + Finds the implementations of a symbol """ - # noinspection PyDefaultArgument def apply( self, name_path: str, @@ -368,36 +221,21 @@ class FindImplementationsTool(Tool, ToolMarkerSymbolicRead): :param max_answer_chars: max result length; -1 for default :return: a list of JSON objects with the symbols implementing the requested symbol """ - self.project.ls_sync_file_system_changes() - - include_body = False - parsed_include_kinds: Sequence[SymbolKind] | None = [SymbolKind(k) for k in include_kinds] if include_kinds else None - parsed_exclude_kinds: Sequence[SymbolKind] | None = [SymbolKind(k) for k in exclude_kinds] if exclude_kinds else None - symbol_retriever = self.create_language_server_symbol_retriever() - - implementing_symbols = symbol_retriever.find_implementing_symbols( - name_path, - relative_file_path=relative_path, - include_body=include_body, - include_kinds=parsed_include_kinds, - exclude_kinds=parsed_exclude_kinds, + return ( + self._api() + .find_implementations( + name_path, + relative_path, + include_info=include_info, + include_kinds=include_kinds, + exclude_kinds=exclude_kinds, + max_answer_chars=max_answer_chars, + ) + .represent() ) - symbol_dicts = [ - dict(s.to_dict(kind=True, relative_path=True, depth=0, body=include_body, body_location=True)) for s in implementing_symbols - ] - if include_info: - info_by_symbol = symbol_retriever.request_info_for_symbol_batch(implementing_symbols) - for s, s_dict in zip(implementing_symbols, symbol_dicts, strict=True): - if symbol_info := info_by_symbol.get(s): - s_dict["info"] = symbol_info - s_dict.pop("name", None) # name is included in the info - result = self._to_json(symbol_dicts) - return self._limit_length(result, max_answer_chars) - - -class FindDeclarationTool(Tool, ToolMarkerSymbolicRead): +class FindDeclarationTool(Tool, ToolMarkerSymbolicRead, LspApiMixin): """ Finds the declaration/definition of a symbol """ @@ -423,70 +261,26 @@ class FindDeclarationTool(Tool, ToolMarkerSymbolicRead): :param include_body: whether to include the symbol's body in the result. Default False. :param include_info: whether to include additional info (hover-like). Default False. """ - self.project.ls_sync_file_system_changes() - - symbol_retriever = self.create_language_server_symbol_retriever() relative_path = self._sanitize_input_param(relative_path) regex = self._sanitize_input_param(regex) - - # find relevant location for lookup - editor = self.create_code_editor() - if not containing_symbol_name_path: - content = editor.read_file(relative_path) - coords = find_text_coordinates(content, regex, require_unique=True) - assert coords is not None - else: - symbol = symbol_retriever.find_unique(name_path_pattern=containing_symbol_name_path, within_relative_path=relative_path) - body_line_numers = symbol.get_body_line_numbers_or_raise() - content = editor.read_file(relative_path, lines=body_line_numers) - coords = find_text_coordinates(content, regex, require_unique=True) - assert coords is not None - coords.line += body_line_numers[0] - - # retrieve declaration - defining_symbol = symbol_retriever.find_declaration( - relative_file_path=relative_path, - line=coords.line, - column=coords.col, - include_body=include_body, - ) - if defining_symbol is None: - raise ValueError( - f"No symbol declaration found at the location of the regex match. Location: {relative_path}:{coords.line}:{coords.col}." + return ( + self._api() + .find_declaration( + relative_path, + regex, + containing_symbol_name_path=containing_symbol_name_path, + include_body=include_body, + include_info=include_info, ) - - # create output - symbol_dict = self._defining_symbol_to_result_dict( - symbol_retriever, - defining_symbol, - include_body, - include_info, + .represent() ) - result = self._to_json(symbol_dict) - return result - - @staticmethod - def _defining_symbol_to_result_dict( - symbol_retriever: Any, - defining_symbol: LanguageServerSymbol, - include_body: bool, - include_info: bool, - ) -> dict[str, Any]: - symbol_dict = dict(defining_symbol.to_dict(kind=True, relative_path=True, depth=0, body=include_body, body_location=True)) - if not include_body and include_info: - if symbol_info := symbol_retriever.request_info_for_symbol(defining_symbol): - symbol_dict["info"] = symbol_info - symbol_dict.pop("name", None) - return symbol_dict -class GetDiagnosticsForFileTool(Tool, ToolMarkerSymbolicRead): +class GetDiagnosticsForFileTool(Tool, ToolMarkerSymbolicRead, LspApiMixin): """ - Gets diagnostics for a file, optionally restricted to a line range, grouped by file, severity, and containing symbol. + Gets diagnostics for a file, grouped by symbol. """ - FILE_LEVEL_DIAGNOSTIC_BUCKET = "" - def apply( self, relative_path: str, @@ -507,34 +301,16 @@ class GetDiagnosticsForFileTool(Tool, ToolMarkerSymbolicRead): :param max_answer_chars: max result length; -1 for default :return: grouped diagnostics for the requested file. """ - self.project.ls_sync_file_system_changes() - - symbol_retriever = self.create_language_server_symbol_retriever() - diagnostics = symbol_retriever.get_file_diagnostics( - relative_file_path=relative_path, - start_line=start_line, - end_line=end_line, - min_severity=min_severity, + return ( + self._api() + .get_diagnostics_for_file( + relative_path, start_line=start_line, end_line=end_line, min_severity=min_severity, max_answer_chars=max_answer_chars + ) + .represent() ) - grouped_diagnostics = GroupedDiagnostics() - for diagnostic in diagnostics: - diag_range = diagnostic["range"]["start"] - name_path = self.FILE_LEVEL_DIAGNOSTIC_BUCKET - owner_symbol = symbol_retriever.find_diagnostic_owner_symbol( - relative_file_path=relative_path, - line=diag_range["line"], - column=diag_range["character"], - ) - if owner_symbol is not None: - name_path = owner_symbol.get_name_path() - grouped_diagnostics.add(relative_path, name_path, diagnostic) - result = self._to_json(grouped_diagnostics.get_dict()) - return self._limit_length(result, max_answer_chars) - - -class GetDiagnosticsForSymbolTool(Tool, ToolMarkerSymbolicRead, ToolMarkerOptional): +class GetDiagnosticsForSymbolTool(Tool, ToolMarkerSymbolicRead, ToolMarkerOptional, LspApiMixin): """ Gets diagnostics for a symbol and, optionally, for symbols that reference it. """ @@ -560,30 +336,20 @@ class GetDiagnosticsForSymbolTool(Tool, ToolMarkerSymbolicRead, ToolMarkerOption :param max_answer_chars: max result length; -1 for default :return: grouped diagnostics for the requested symbol and, optionally, its referencing symbols. """ - self.project.ls_sync_file_system_changes() - - symbol_retriever = self.create_language_server_symbol_retriever() - diagnostics_by_symbol = symbol_retriever.get_symbol_diagnostics( - name_path=name_path, - reference_file=reference_file or None, - check_symbol_references=check_symbol_references, - min_severity=min_severity, + return ( + self._api() + .get_diagnostics_for_symbol( + name_path, + reference_file=reference_file, + check_symbol_references=check_symbol_references, + min_severity=min_severity, + max_answer_chars=max_answer_chars, + ) + .represent() ) - grouped_diagnostics = GroupedDiagnostics() - for symbol, diagnostics in diagnostics_by_symbol.items(): - relative_path = symbol.relative_path - if relative_path is None: - continue - symbol_name_path = symbol.get_name_path() - for diagnostic in diagnostics: - grouped_diagnostics.add(relative_path, symbol_name_path, diagnostic) - result = self._to_json(grouped_diagnostics.get_dict()) - return self._limit_length(result, max_answer_chars) - - -class ReplaceSymbolBodyTool(EditingToolWithDiagnostics): +class ReplaceSymbolBodyTool(EditingToolWithDiagnostics, EditApiMixin): """ Replaces the full definition of a symbol using the language server backend. """ @@ -606,17 +372,12 @@ class ReplaceSymbolBodyTool(EditingToolWithDiagnostics): in the programming language, including e.g. the signature line for functions. Depending on the language, it may or may not include a preceding docstring or other preceding annotations. """ - with self.DiagnosticsContext(self, relative_path) as diagnostics_context: - code_editor = self.create_code_editor() - code_editor.replace_body( - name_path, - relative_file_path=relative_path, - body=body, - ) - return diagnostics_context.format_result(SUCCESS_RESULT) + with self.diagnostics_context(relative_path) as diagnostics_context: + result = self._api().replace_symbol_body(name_path, relative_path, body) + return diagnostics_context.format_result(result) -class InsertAfterSymbolTool(EditingToolWithDiagnostics): +class InsertAfterSymbolTool(EditingToolWithDiagnostics, EditApiMixin): """ Inserts content after the end of the definition of a given symbol. """ @@ -636,13 +397,12 @@ class InsertAfterSymbolTool(EditingToolWithDiagnostics): :param body: the body/content to be inserted. The inserted code shall begin with the next line after the symbol. """ - with self.DiagnosticsContext(self, relative_path) as diagnostics_context: - code_editor = self.create_code_editor() - code_editor.insert_after_symbol(name_path, relative_file_path=relative_path, body=body) - return diagnostics_context.format_result(SUCCESS_RESULT) + with self.diagnostics_context(relative_path) as diagnostics_context: + result = self._api().insert_after_symbol(name_path, relative_path, body) + return diagnostics_context.format_result(result) -class InsertBeforeSymbolTool(EditingToolWithDiagnostics): +class InsertBeforeSymbolTool(EditingToolWithDiagnostics, EditApiMixin): """ Inserts content before the beginning of the definition of a given symbol. """ @@ -662,13 +422,12 @@ class InsertBeforeSymbolTool(EditingToolWithDiagnostics): :param relative_path: the relative path to the file containing the symbol :param body: the body/content to be inserted before the line in which the referenced symbol is defined """ - with self.DiagnosticsContext(self, relative_path) as diagnostics_context: - code_editor = self.create_code_editor() - code_editor.insert_before_symbol(name_path, relative_file_path=relative_path, body=body) - return diagnostics_context.format_result(SUCCESS_RESULT) + with self.diagnostics_context(relative_path) as diagnostics_context: + result = self._api().insert_before_symbol(name_path, relative_path, body) + return diagnostics_context.format_result(result) -class RenameSymbolTool(Tool, ToolMarkerSymbolicEdit): +class RenameSymbolTool(Tool, ToolMarkerSymbolicEdit, LspApiMixin): """ Renames a symbol throughout the codebase using language server refactoring capabilities. For JB, we use a separate tool. @@ -690,13 +449,10 @@ class RenameSymbolTool(Tool, ToolMarkerSymbolicEdit): :param new_name: the new name for the symbol :return: result summary indicating success or failure """ - self.project.ls_sync_file_system_changes() - code_editor = self.create_ls_code_editor() - status_message = code_editor.rename_symbol(name_path, relative_path=relative_path, new_name=new_name) - return status_message + return self._api().rename_symbol(name_path, relative_path, new_name) -class SafeDeleteSymbol(Tool, ToolMarkerSymbolicEdit): +class SafeDeleteSymbol(Tool, ToolMarkerSymbolicEdit, LspApiMixin): def apply( self, name_path_pattern: str, @@ -709,31 +465,4 @@ class SafeDeleteSymbol(Tool, ToolMarkerSymbolicEdit): :param name_path_pattern: name path of the symbol to delete :param relative_path: the relative path to the file containing the symbol to delete """ - self.project.ls_sync_file_system_changes() - - ls_symbol_retriever = self.create_language_server_symbol_retriever() - symbol = ls_symbol_retriever.find_unique(name_path_pattern, substring_matching=False, within_relative_path=relative_path) - symbol_rel_path = symbol.relative_path - assert symbol_rel_path is not None, f"Symbol {name_path_pattern} has no relative path, this is likely a bug." - assert symbol_rel_path == relative_path, f"Symbol {name_path_pattern} is not in the expected relative path {relative_path}." - symbol_name_path = symbol.get_name_path() - - symbol_line = symbol.line - symbol_col = symbol.column - assert symbol_line is not None and symbol_col is not None, ( - f"Symbol {name_path_pattern} has no identifier position, this is likely a bug." - ) - lang_server = ls_symbol_retriever.get_language_server(symbol_rel_path) - references_locations = lang_server.request_references(symbol_rel_path, symbol_line, symbol_col) - file_to_lines: dict[str, list[int]] = defaultdict(list) - if references_locations: - for ref_loc in references_locations: - ref_relative_path = ref_loc.get("relativePath") - if ref_relative_path is None: - continue - file_to_lines[ref_relative_path].append(ref_loc["range"]["start"]["line"]) - if file_to_lines: - return f"Cannot delete, the symbol {symbol_name_path} is referenced in: {self._to_json(file_to_lines)}" - code_editor = self.create_ls_code_editor() - code_editor.delete_symbol(symbol_name_path, relative_file_path=symbol_rel_path) - return SUCCESS_RESULT + return self._api().safe_delete_symbol(name_path_pattern, relative_path) diff --git a/src/serena/tools/tools_base.py b/src/serena/tools/tools_base.py index 60e81ad9..a57656e8 100644 --- a/src/serena/tools/tools_base.py +++ b/src/serena/tools/tools_base.py @@ -5,33 +5,35 @@ import json from abc import ABC from collections.abc import Callable, Iterable from dataclasses import dataclass -from functools import cached_property -from types import TracebackType -from typing import TYPE_CHECKING, Any, Optional, Protocol, Self, TypeVar, cast +from typing import TYPE_CHECKING, Any, Protocol, TypeVar, cast from mcp import Implementation -from mcp.server.fastmcp import Context -from mcp.server.fastmcp.utilities.func_metadata import FuncMetadata, func_metadata +from mcp.server.mcpserver import Context +from mcp.server.mcpserver.utilities.func_metadata import FuncMetadata, func_metadata from sensai.util import logging +from sensai.util.helper import mark_used from sensai.util.string import dict_string -from serena.config.serena_config import LanguageBackend +from serena.code_editor import EditedFileContext +from serena.lsp.lsp_diagnostics import DiagnosticsContext from serena.memories.memory_manager import MemoryManager from serena.project import Project from serena.prompt_factory import PromptFactory +from serena.repl.facade import SUCCESS_RESULT from serena.util.class_decorators import singleton from serena.util.inspection import iter_subclasses -from serena.util.ls_diagnostics import DiagnosticsDiff, EditedFilePath, PublishedDiagnosticsSnapshot +from serena.util.text_utils import TextOutputUtils from solidlsp.ls_exceptions import SolidLSPException if TYPE_CHECKING: from serena.agent import SerenaAgent - from serena.code_editor import CodeEditor, LanguageServerCodeEditor + from serena.code_editor import CodeEditor from serena.symbol import LanguageServerSymbolRetriever + +mark_used(SUCCESS_RESULT, EditedFileContext) # backward compatibility log = logging.getLogger(__name__) T = TypeVar("T") -SUCCESS_RESULT = "OK" class Component(ABC): @@ -63,22 +65,7 @@ class Component(ABC): return self.agent.get_active_project_or_raise() def create_code_editor(self) -> "CodeEditor": - from ..code_editor import JetBrainsCodeEditor - - match self.agent.get_language_backend(): - case LanguageBackend.LSP: - return self.create_ls_code_editor() - case LanguageBackend.JETBRAINS: - return JetBrainsCodeEditor(project=self.project) - case _: - raise ValueError - - def create_ls_code_editor(self) -> "LanguageServerCodeEditor": - from ..code_editor import LanguageServerCodeEditor - - if not self.agent.is_using_language_server(): - raise Exception("Cannot create LanguageServerCodeEditor; agent is not in language server mode.") - return LanguageServerCodeEditor(self.create_language_server_symbol_retriever()) + return self.agent.get_language_backend().create_code_editor(self.project) class ToolMarker: @@ -152,32 +139,12 @@ class Tool(Component): # (which is use by the LLM, so a good description is important) # and to validate the tool call arguments. - SESSION_ID_PARAM_NAME = "session_id" - """ - parameter name to use in apply method for the client session ID. - This parameter will be ignored by the MCP interface but will be populated with the session ID of the current client session - when the tool is called, allowing tools to be session-aware if needed. - """ - _last_tool_call_client_str: str | None = None """We can only get the client info from within a tool call. Each tool call will update this variable.""" def __init__(self, agent: "SerenaAgent"): super().__init__(agent) - @cached_property - def _is_session_aware(self) -> bool: - """ - :return: whether the tool is session-aware, i.e. whether the apply method expects a session_id (str) parameter. - """ - # check apply method for session_id arg - apply_fn = self.get_apply_fn() - sig = inspect.signature(apply_fn) - for param in sig.parameters.values(): - if param.name == self.SESSION_ID_PARAM_NAME: - return True - return False - @staticmethod def _sanitize_input_param(raw_param: str) -> str: # some clients replace < and > with their escaped html versions, we need to counteract this @@ -266,9 +233,9 @@ class Tool(Component): if apply_fn is None: raise AttributeError(f"apply method not defined in {cls}. Did you forget to implement it?") - return func_metadata(apply_fn, skip_names=["self", "cls", cls.SESSION_ID_PARAM_NAME], structured_output=structured_output) + return func_metadata(apply_fn, skip_names=["self", "cls"], structured_output=structured_output) - def _log_tool_application(self, frame: Any, session_id: str) -> None: + def _log_tool_application(self, frame: Any) -> None: params = {} ignored_params = {"self", "log_call", "catch_exceptions", "args", "apply_fn"} for param, value in frame.f_locals.items(): @@ -278,7 +245,14 @@ class Tool(Component): params.update(value) else: params[param] = value - log.info(f"{self.get_name_from_cls()}: {dict_string(params)}; session_id: {session_id}") + log.info(f"{self.get_name_from_cls()}: {dict_string(params)}") + + def _resolve_max_answer_chars(self, max_answer_chars: int) -> int: + """ + :param max_answer_chars: the maximum number of answer characters as passed to the tool; -1 for the configured default + :return: the effective maximum + """ + return self.agent.serena_config.default_max_tool_answer_chars if max_answer_chars == -1 else max_answer_chars def _limit_length( self, @@ -294,26 +268,13 @@ class Tool(Component): version of the result. They are tried in order until one fits within ``max_answer_chars``. :return: the result string, potentially replaced by a shortened version """ - if max_answer_chars == -1: - max_answer_chars = self.agent.serena_config.default_max_tool_answer_chars - if max_answer_chars <= 0: - raise ValueError(f"Must be positive or the default (-1), got: {max_answer_chars=}") - if (n_chars := len(result)) > max_answer_chars: - too_long_msg = ( - f"The answer is too long ({n_chars} characters). " + "You can adjust your query or raise the max_answer_chars parameter." - ) - if shortened_result_factories is not None: - # try each shortening closure in order; - for make_shorter in shortened_result_factories: - shortened = make_shorter() - candidate = f"{too_long_msg}\n{shortened}" - if len(candidate) <= max_answer_chars: - return candidate - result = too_long_msg - return result + max_answer_chars = self._resolve_max_answer_chars(max_answer_chars) + return TextOutputUtils.limit_length( + result=result, max_answer_chars=max_answer_chars, shortened_result_factories=shortened_result_factories + ) def is_active(self) -> bool: - return self.agent.tool_is_active(self.get_name()) + return self.agent.get_active_tools().contains_tool_name(self.get_name()) def is_readonly(self) -> bool: return not self.can_edit() @@ -339,10 +300,8 @@ class Tool(Component): :param catch_exceptions: whether to catch exceptions and return their messages as strings, instead of raising a ToolCallError """ # obtain session ID and client info - session_id = "global" if mcp_ctx is not None: try: - session_id = "%x" % id(mcp_ctx.session) client_params = mcp_ctx.session.client_params if client_params is not None: client_info = cast(Implementation, client_params.clientInfo) @@ -363,7 +322,7 @@ class Tool(Component): ) if log_call: - self._log_tool_application(inspect.currentframe(), session_id) + self._log_tool_application(inspect.currentframe()) # check whether the tool requires an active project and language server if not isinstance(self, ToolMarkerDoesNotRequireActiveProject): @@ -375,8 +334,6 @@ class Tool(Component): # construct apply kwargs, adding session_id if the tool is session-aware apply_kwargs = dict(kwargs) - if self._is_session_aware: - apply_kwargs["session_id"] = session_id # apply the actual tool try: @@ -444,7 +401,7 @@ class Tool(Component): @staticmethod def _to_json(x: Any) -> str: - return json.dumps(x, ensure_ascii=False) + return TextOutputUtils.to_json(x) def _wrapped_tool_response(self, response: Any, message: str) -> str: """ @@ -478,89 +435,16 @@ class EditingToolWithDiagnostics(Tool, ToolMarkerCanEdit): are then resolved in subsequent edits. """ - DIAGNOSTICS_KEY = "diagnostics[warning-or-higher]" - - class DiagnosticsContext: - def __init__(self, tool: "EditingToolWithDiagnostics", *edited_relative_paths: str) -> None: - self._tool = tool - self._is_diagnostics_enabled = tool.ENABLE_DIAGNOSTICS and tool.agent.is_using_language_server() - self._edited_files = [EditedFilePath(path, path) for path in edited_relative_paths] - self._before_edit_diagnostics_snapshot: PublishedDiagnosticsSnapshot | None = None - self._symbol_retriever: Optional["LanguageServerSymbolRetriever"] | None = None - if self._is_diagnostics_enabled: - self._symbol_retriever = tool.create_language_server_symbol_retriever() - self._before_edit_diagnostics_snapshot = PublishedDiagnosticsSnapshot(self._edited_files, self._symbol_retriever) - - def __enter__(self) -> Self: - return self - - def __exit__(self, exc_type, exc_val, exc_tb): - pass - - def format_result( - self, - base_result: str, - ) -> str: - if not self._is_diagnostics_enabled: - return base_result - - if self._before_edit_diagnostics_snapshot is None: - return base_result - - assert self._symbol_retriever is not None - diagnostics_diff = DiagnosticsDiff(self._before_edit_diagnostics_snapshot, self._edited_files, self._symbol_retriever) - grouped_diagnostics = diagnostics_diff.get_grouped_diagnostics().get_dict() - - if not grouped_diagnostics: - return base_result - else: - result_dict = { - "result": base_result, - EditingToolWithDiagnostics.DIAGNOSTICS_KEY: grouped_diagnostics, - } - return self._tool._to_json(result_dict) - - -class EditedFileContext: - """ - Context manager for file editing. - - Create the context, then use `set_updated_content` to set the new content, the original content - being provided in `original_content`. - When exiting the context without an exception, the updated content will be written back to the file. - """ - - def __init__(self, relative_path: str, code_editor: "CodeEditor"): - self._relative_path = relative_path - self._code_editor = code_editor - self._edited_file: CodeEditor.EditedFile | None = None - self._edited_file_context: Any = None - - def __enter__(self) -> Self: - self._edited_file_context = self._code_editor.edited_file_context(self._relative_path) - self._edited_file = self._edited_file_context.__enter__() - return self - - def get_original_content(self) -> str: + def diagnostics_context(self, *edited_relative_paths: str) -> DiagnosticsContext: """ - :return: the original content of the file before any modifications. - """ - assert self._edited_file is not None - return self._edited_file.get_contents() + Creates a context for use with the `with` statement, which captures the diagnostics before the edit, + such that changes can be reported - def set_updated_content(self, content: str) -> None: + :param edited_relative_paths: the relative paths of the files that are to be edited within the context + :return: a context which captures the diagnostics before the edit, such that changes can be reported + via `format_result` """ - Sets the updated content of the file, which will be written back to the file - when the context is exited without an exception. - - :param content: the updated content of the file - """ - assert self._edited_file is not None - self._edited_file.set_contents(content) - - def __exit__(self, exc_type: type[BaseException] | None, exc_value: BaseException | None, traceback: TracebackType | None) -> None: - assert self._edited_file_context is not None - self._edited_file_context.__exit__(exc_type, exc_value, traceback) + return DiagnosticsContext(self.agent, *edited_relative_paths, enable=self.ENABLE_DIAGNOSTICS) @dataclass(kw_only=True) diff --git a/src/serena/tools/workflow_tools.py b/src/serena/tools/workflow_tools.py index 8991eea1..4aeff204 100644 --- a/src/serena/tools/workflow_tools.py +++ b/src/serena/tools/workflow_tools.py @@ -3,12 +3,11 @@ Tools supporting the general workflow of the agent """ # SPDX-License-Identifier: GPL-3.0-or-later -import platform - from serena.tools import Tool, ToolMarkerDoesNotRequireActiveProject, ToolMarkerOptional, WriteMemoryTool +from serena.tools.memory_tools import MemoryApiMixin -class OnboardingTool(Tool): +class OnboardingTool(Tool, MemoryApiMixin): """ Performs onboarding (identifying the project structure and essential tasks, e.g. for testing or building). """ @@ -20,14 +19,10 @@ class OnboardingTool(Tool): :return: instructions on how to create the onboarding information """ - write_memory_tool_available = self.agent.tool_is_exposed(WriteMemoryTool.get_name_from_cls()) + write_memory_tool_available = self.agent.is_tool_function_available(WriteMemoryTool) if not write_memory_tool_available: return "Memory writing tool not activated, skipping onboarding." - system = platform.system() - # seed the project-local memory-maintenance memory (or detect a global override) so - # the prompt can point the agent at the conventions before it writes anything - memory_maintenance_name = self.memory_manager.ensure_memory_maintenance_memory() - return self.prompt_factory.create_onboarding_prompt(system=system, memory_maintenance_name=memory_maintenance_name) + return self._api().onboarding() class InitialInstructionsTool(Tool, ToolMarkerDoesNotRequireActiveProject): @@ -36,15 +31,13 @@ class InitialInstructionsTool(Tool, ToolMarkerDoesNotRequireActiveProject): for clients that do not read the initial instructions when the MCP server is connected. """ - # noinspection PyIncorrectDocstring - # (session_id is injected via apply_ex) - def apply(self, session_id: str) -> str: + def apply(self) -> str: """ Provides the 'Serena Instructions Manual', which contains essential information on how to use the Serena toolbox. IMPORTANT: If you have not yet read the manual, call this tool immediately after you are given your task by the user, as it will critically inform you! """ - return self.agent.create_system_prompt(session_id=session_id) + return self.agent.create_system_prompt() class SerenaInfoTool(Tool, ToolMarkerOptional, ToolMarkerDoesNotRequireActiveProject): diff --git a/src/serena/util/file_proxy.py b/src/serena/util/file_proxy.py index df6ab3a4..14889a1a 100644 --- a/src/serena/util/file_proxy.py +++ b/src/serena/util/file_proxy.py @@ -6,8 +6,6 @@ from abc import ABC, abstractmethod from collections.abc import Iterator from typing import TYPE_CHECKING, Self -from serena.jetbrains import jetbrains_types as jb - if TYPE_CHECKING: from serena.project import Project @@ -30,19 +28,15 @@ class FileProxy(ABC): """ @staticmethod - def is_external_path(relative_path: str) -> bool: + def is_external_path(relative_path: str, project: "Project") -> bool: """ :return: whether the given relative path is an encoded external path (not a local project file) """ - # This is intended to be extended once we also support external paths in other backends - return jb.is_external_path(relative_path) + return project.language_backend.is_external_path(relative_path) @classmethod def from_project_relative_path(cls, project: "Project", relative_path: str) -> "FileProxy": - if cls.is_external_path(relative_path): - if project.language_backend.is_jetbrains(): - return JetBrainsFileProxy(relative_path, project) - return LocalProjectFileProxy(relative_path, project) + return project.language_backend.create_file_proxy(relative_path, project) class LocalProjectFileProxy(FileProxy): @@ -62,29 +56,6 @@ class LocalProjectFileProxy(FileProxy): return True -class JetBrainsFileProxy(FileProxy): - """ - Retrieves the contents of a file from the JetBrains plugin via the plugin client, given its relative path, - which may be an external path (e.g., "") - """ - - def __init__(self, relative_path: str, project: "Project"): - self._relative_path = relative_path - self._project = project - - def get_contents(self) -> str: - from serena.jetbrains.jetbrains_plugin_client import JetBrainsPluginClient - - client = JetBrainsPluginClient.from_project(self._project) - return client.read_file(self._relative_path) - - def get_relative_path(self) -> str: - return self._relative_path - - def is_glob_supported(self): - return False - - class FileCollection: def __init__(self, file_proxies: list[FileProxy]): self._file_proxies = file_proxies diff --git a/src/serena/util/file_system.py b/src/serena/util/file_system.py index ad8b7b84..b78af51f 100644 --- a/src/serena/util/file_system.py +++ b/src/serena/util/file_system.py @@ -31,6 +31,11 @@ def write_file_atomic(path: str, content: str, *, encoding: str, newline: str | :param encoding: the encoding to use for the write :param newline: passed through to the underlying ``open()`` call to control newline translation """ + # ``open(path, "w")`` follows symlinks and writes through to the target, whereas replacing the + # link path itself would swap the link out for a regular file and leave its target holding the + # old content. Resolving first keeps this a drop-in replacement, and puts the temporary file in + # the destination's real directory, which is where it has to be for the rename to be atomic. + path = os.path.realpath(path) target_dir = os.path.dirname(path) or "." try: existing_mode: int | None = stat.S_IMODE(os.stat(path).st_mode) @@ -435,7 +440,7 @@ class GitignoreParser: self._load_gitignore_files() -def match_path(relative_path: str, path_spec: PathSpec, root_path: str = "") -> bool: +def match_path(relative_path: str, path_spec: PathSpec, root_path: str = "", is_dir: bool | None = None) -> bool: """ Match a relative path against a given pathspec. Just pathspec.match_file() is not enough, we need to do some massaging to fix issues with pathspec matching. @@ -443,6 +448,8 @@ def match_path(relative_path: str, path_spec: PathSpec, root_path: str = "") -> :param relative_path: relative path to match against the pathspec :param path_spec: the pathspec to match against :param root_path: the root path from which the relative path is derived + :param is_dir: whether the path is a directory, where the caller already knows; passing it avoids + an `os.path.isdir` call. `None` determines it from the filesystem. :return: """ if str(relative_path) in {"", "."}: @@ -460,7 +467,9 @@ def match_path(relative_path: str, path_spec: PathSpec, root_path: str = "") -> # pathspec can't handle the matching of directories if they don't end with a slash! # see https://github.com/cpburnz/python-pathspec/issues/89 - abs_path = os.path.abspath(os.path.join(root_path, relative_path)) - if os.path.isdir(abs_path) and not normalized_path.endswith("/"): + if is_dir is None: + abs_path = os.path.abspath(os.path.join(root_path, relative_path)) + is_dir = os.path.isdir(abs_path) + if is_dir and not normalized_path.endswith("/"): normalized_path = normalized_path + "/" return path_spec.match_file(normalized_path) diff --git a/src/serena/util/inspection.py b/src/serena/util/inspection.py index 5ab17aa7..73362089 100644 --- a/src/serena/util/inspection.py +++ b/src/serena/util/inspection.py @@ -6,7 +6,7 @@ from collections.abc import Callable, Iterator from typing import TypeVar from serena.util.file_system import find_all_non_ignored_files -from solidlsp.ls_config import LanguageServerId +from solidlsp.ls_config import LanguageServerId, LanguageServerIdLike T = TypeVar("T") @@ -16,22 +16,30 @@ log = logging.getLogger(__name__) def iter_subclasses( cls: type[T], recursive: bool = True, inclusion_predicate: Callable[[type[T]], bool] = lambda t: True ) -> Iterator[type[T]]: - """Iterate over all subclasses of a class. + """Iterate over all subclasses of a class, yielding each subclass once (even if it is reachable via multiple base classes). :param cls: The class whose subclasses to iterate over. :param recursive: If True, also iterate over all subclasses of all subclasses. :param inclusion_predicate: a predicate function to decide whether to include a subclass in the result """ - for subclass in cls.__subclasses__(): - if inclusion_predicate(subclass): - yield subclass - if recursive: - yield from iter_subclasses(subclass, recursive, inclusion_predicate) + seen: set[type] = set() + + def iterate(c: type[T]) -> Iterator[type[T]]: + for subclass in c.__subclasses__(): + if subclass in seen: + continue + seen.add(subclass) + if inclusion_predicate(subclass): + yield subclass + if recursive: + yield from iterate(subclass) + + yield from iterate(cls) def compute_language_server_support_composition( - repo_path: str, ls_ids: list[LanguageServerId] | None = None -) -> dict[LanguageServerId, float]: + repo_path: str, ls_ids: list[LanguageServerIdLike] | None = None +) -> dict[LanguageServerIdLike, float]: """ Determine the composition of a repository in terms of the language servers that can be used to analyze it. @@ -56,7 +64,7 @@ def compute_language_server_support_composition( matchers = {lang: lang.get_source_fn_matcher() for lang in ls_ids} # count files per language in a single pass over the files - ls_file_counts: dict[LanguageServerId, int] = {} + ls_file_counts: dict[LanguageServerIdLike, int] = {} recognised_files = 0 for file_path in all_files: # Use just the filename for matching, not the full path diff --git a/src/serena/util/text_utils.py b/src/serena/util/text_utils.py index 9c030cf7..75d4de63 100644 --- a/src/serena/util/text_utils.py +++ b/src/serena/util/text_utils.py @@ -1,12 +1,13 @@ # SPDX-License-Identifier: GPL-3.0-or-later import hashlib +import json import logging import re from collections.abc import Callable from dataclasses import dataclass, field from enum import StrEnum -from typing import Any, Literal, Self +from typing import TYPE_CHECKING, Any, Literal, Self from bs4 import BeautifulSoup from joblib import Parallel, delayed @@ -15,6 +16,10 @@ from sensai.util.string import ToStringMixin from serena.util.file_proxy import FileCollection, FileProxy from solidlsp.ls_utils import TextCoordinateProvider, TextCoordinates, TextUtils +if TYPE_CHECKING: + from serena.code_editor import CodeEditor + from serena.project import Project + log = logging.getLogger(__name__) @@ -406,13 +411,18 @@ class ContentReplacer: self.regex_multiline = regex_multiline @staticmethod - def _create_replacement_function(regex_pattern: str, repl_template: str, regex_flags: int) -> Callable[[re.Match], str]: + def _create_replacement_function( + regex_pattern: str, repl_template: str, regex_flags: int, expand_backrefs: bool + ) -> Callable[[re.Match], str]: """ Creates a replacement function that validates for ambiguity and handles backreferences. :param regex_pattern: The regex pattern being used for matching - :param repl_template: The replacement template with $!1, $!2, etc. for backreferences + :param repl_template: The replacement template; in regex mode, it may contain $!1, $!2, etc. for + backreferences; in literal mode, it is used verbatim :param regex_flags: The flags to use when searching (e.g., re.DOTALL | re.MULTILINE) + :param expand_backrefs: Whether $!N backreferences are expanded in the template; false in literal mode, + mirroring the mode gate in MultiFileContentReplacer.find_occurrences :return: A function suitable for use with re.sub() or re.subn() """ @@ -434,11 +444,19 @@ class ContentReplacer: "e.g. by matching specific context after the match, or try using the literal mode." ) - # Handle backreferences: replace $!1, $!2, etc. with actual matched groups + # in literal mode, the template is the final replacement; $!N sequences need no escaping + if not expand_backrefs: + return repl_template + + # Handle backreferences: replace $!1, $!2, etc. with actual matched groups; groups that + # exist but did not participate in the match expand to the empty string def expand_backreference(m: re.Match) -> str: group_num = int(m.group(1)) - group_value = match.group(group_num) - return group_value if group_value is not None else m.group(0) + try: + group_value = match.group(group_num) + except IndexError as e: + raise ValueError(f"Backreference $!{group_num} refers to a group that does not exist in the search expression") from e + return group_value if group_value is not None else "" result = re.sub(r"\$!(\d+)", expand_backreference, repl_template) return result @@ -458,8 +476,8 @@ class ContentReplacer: :param content: the content in which to perform the replacement :param needle: the search expression, which is either a literal string or a regular expression, depending on the mode - :param repl: the replacement string, which, in regex mode, may contain backreferences in the form of $!1, $!2, etc. to - refer to matched groups in the search expression + :param repl: the replacement string; in regex mode, it may contain backreferences in the form of $!1, $!2, etc. + to refer to matched groups in the search expression; in literal mode, it is used verbatim :return: the updated content after performing the replacement """ if self.mode == "literal": @@ -471,8 +489,8 @@ class ContentReplacer: regex_flags = (re.MULTILINE | re.DOTALL) if self.regex_multiline else 0 - # create replacement function with validation and backreference handling - repl_fn = self._create_replacement_function(regex, repl, regex_flags=regex_flags) + # create replacement function with ambiguity validation and, in regex mode, backreference handling + repl_fn = self._create_replacement_function(regex, repl, regex_flags=regex_flags, expand_backrefs=self.mode == "regex") # perform replacement updated_content, n = re.subn(regex, repl_fn, content, flags=regex_flags) @@ -548,8 +566,12 @@ class MultiFileContentReplacer: """Expands $!1, $!2, ... in the replacement template (same syntax as :class:`ContentReplacer`).""" def expand(m: re.Match) -> str: - group_value = match.group(int(m.group(1))) - return group_value if group_value is not None else m.group(0) + group_num = int(m.group(1)) + try: + group_value = match.group(group_num) + except IndexError as e: + raise ValueError(f"Backreference $!{group_num} refers to a group that does not exist in the search expression") from e + return group_value if group_value is not None else "" return re.sub(r"\$!(\d+)", expand, repl_template) @@ -643,6 +665,253 @@ class MultiFileContentReplacer: return "\n".join(diff_lines) +class ReplacementRejectedError(ValueError): + """ + Raised when a replacement is not applied because a safety check failed; no changes have been made. + """ + + def __init__(self, message: str, show_prospective_changes: bool) -> None: + """ + :param message: the reason for the rejection + :param show_prospective_changes: whether the listing of the prospective changes should be presented along + with the message (the message refers to it) + """ + super().__init__(message) + self.show_prospective_changes = show_prospective_changes + + +@dataclass +class MultiFileReplacementResult: + """ + The result of an applied replacement. + """ + + num_occurrences_by_file: dict[str, int] + """the number of replaced occurrences per file (relative path), in application order""" + + @property + def num_occurrences(self) -> int: + return sum(self.num_occurrences_by_file.values()) + + def to_display_string(self) -> str: + per_file = "\n".join(f" {path}: {n}" for path, n in self.num_occurrences_by_file.items()) + return f"Replaced {self.num_occurrences} occurrence(s) in {len(self.num_occurrences_by_file)} file(s):\n{per_file}" + + +class MultiFileReplacement: + """ + A prospective replacement of a pattern across the files of a project within a given scope: holds the + occurrences found, supports selecting occurrences by their ids, guards blind application against + unintended replacements, renders a preview listing and applies the replacement via a code editor. + """ + + def __init__( + self, + project: "Project", + needle: str, + repl: str, + mode: Literal["literal", "regex"], + relative_path: str = "", + paths_include_glob: str = "", + paths_exclude_glob: str = "", + ) -> None: + """ + :param project: the project whose files are to be searched + :param needle: the string (mode "literal") or regular expression (mode "regex") to search for + :param repl: the replacement string (may contain $!N backreferences in regex mode) + :param mode: how `needle` is to be interpreted + :param relative_path: only consider this file or directory (default: the whole project) + :param paths_include_glob: optional glob restricting which files are considered + :param paths_exclude_glob: optional glob of files to exclude; takes precedence over the include glob + """ + self._replacer = MultiFileContentReplacer(mode=mode) + self._needle = needle + self._repl = repl + files = self._collect_files(project, relative_path.strip(), paths_include_glob.strip(), paths_exclude_glob.strip()) + self._contents = dict(files) + self.occurrences: list[ReplacementOccurrence] = self._replacer.find_occurrences(files, needle, repl) + + @staticmethod + def _collect_files(project: "Project", relative_path: str, paths_include_glob: str, paths_exclude_glob: str) -> list[tuple[str, str]]: + """ + :return: the (relative_path, content) pairs of the readable, non-ignored files in scope, in sorted path order + """ + if relative_path: + project.validate_relative_path(relative_path, require_not_ignored=True) + file_collection = project.create_file_collection(relative_path, code_files_only=False, skip_ignored_files=True).filter_glob( + paths_include_glob or None, paths_exclude_glob or None + ) + files: list[tuple[str, str]] = [] + for file_proxy in sorted(file_collection, key=lambda f: f.get_relative_path()): + try: + files.append((file_proxy.get_relative_path(), file_proxy.get_contents())) + except Exception: + continue # skip unreadable (e.g. binary) files + return files + + @property + def affected_files(self) -> list[str]: + """ + :return: the relative paths of the files containing occurrences, sorted + """ + return sorted({o.relative_path for o in self.occurrences}) + + def select(self, occurrence_ids: list[str]) -> list[ReplacementOccurrence]: + """ + Resolves the given occurrence ids (as obtained from a previous listing). + + :param occurrence_ids: the ids of the occurrences to select + :return: the selected occurrences + :raises ReplacementRejectedError: if the selection is empty or any id cannot be resolved + """ + occurrences_by_id = {o.occurrence_id: o for o in self.occurrences} + indices_by_path: dict[str, set[int]] = {} + for o in self.occurrences: + indices_by_path.setdefault(o.relative_path, set()).add(o.index_in_file) + + # resolve each id, diagnosing failures + selected: dict[str, ReplacementOccurrence] = {} + problems: list[str] = [] + for oid in occurrence_ids: + occurrence = occurrences_by_id.get(oid) + if occurrence is not None: + selected[oid] = occurrence + continue + id_match = MultiFileContentReplacer.OCCURRENCE_ID_REGEX.match(oid) + if id_match is None: + problems.append(f"{oid}: malformed id (expected ':@' as returned by a dry run)") + elif id_match.group("path") not in indices_by_path: + problems.append(f"{oid}: the pattern currently has no matches in this file") + elif int(id_match.group("index")) not in indices_by_path[id_match.group("path")]: + problems.append(f"{oid}: the file now has fewer matches than at dry-run time (content changed)") + else: + problems.append(f"{oid}: the matched text changed since the dry run (content changed)") + + if problems: + problem_lines = "\n".join(f" {p}" for p in problems) + raise ReplacementRejectedError( + f"{len(problems)} of the given occurrence_ids could not be resolved - NO changes were applied:\n" + f"{problem_lines}\n" + "Re-run with dry_run=True to obtain current occurrence ids.", + show_prospective_changes=False, + ) + if not selected: + raise ReplacementRejectedError( + "occurrence_ids is empty - pass at least one id from a dry run, or omit the parameter to replace all.", + show_prospective_changes=False, + ) + return list(selected.values()) + + def select_all_guarded(self, expected_count: int = -1) -> list[ReplacementOccurrence]: + """ + Selects all occurrences for a blind application (without explicit selection), applying safety checks. + + :param expected_count: the number of occurrences expected; -1 disables the check + :return: all occurrences + :raises ReplacementRejectedError: if there are no occurrences, the count differs from the expectation + or any occurrence is ambiguous + """ + if not self.occurrences: + raise ReplacementRejectedError( + "No occurrences of the pattern were found - NO changes were applied. " + "Check the mode (a literal needle containing regex metacharacters must use mode 'literal'; " + "wildcards require mode 'regex') and the path/glob restrictions, " + "or locate the content with search_for_pattern first.", + show_prospective_changes=False, + ) + if expected_count >= 0 and len(self.occurrences) != expected_count: + raise ReplacementRejectedError( + f"expected_count={expected_count}, but the pattern matches {len(self.occurrences)} occurrence(s) - " + "NO changes were applied. Review the prospective changes below; re-issue with the corrected " + "expectation, a refined pattern, or occurrence_ids selecting the intended subset.", + show_prospective_changes=True, + ) + num_ambiguous = sum(1 for o in self.occurrences if o.is_ambiguous) + if num_ambiguous: + raise ReplacementRejectedError( + f"{num_ambiguous} occurrence(s) are ambiguous (the pattern matches again inside the matched text, " + "indicating possible over-matching) - NO changes were applied. Review the prospective changes below " + "and either refine the pattern or explicitly select occurrences via occurrence_ids.", + show_prospective_changes=True, + ) + return list(self.occurrences) + + def render_listing(self, max_answer_chars: int, dry_run: bool) -> str: + """ + Renders the prospective changes as a list of minimal line diffs with occurrence ids, subject to the given length limit + (falling back to locations only, per-file counts and finally a summary). + + :param max_answer_chars: the maximum number of characters (must be positive) + :param dry_run: whether the listing is the result of a dry run (adding instructions on how to proceed) + :return: the listing + """ + affected_files = self.affected_files + header = f"Found {len(self.occurrences)} occurrence(s) in {len(affected_files)} file(s)." + if dry_run: + header += ( + " DRY RUN - no changes were applied.\n" + "Re-issue with dry_run=False to replace all of them, or additionally pass occurrence_ids " + "with the ids of the occurrences to replace." + ) + parts = [header] + for path in affected_files: + file_occurrences = [o for o in self.occurrences if o.relative_path == path] + parts.append(f"\n{path} ({len(file_occurrences)} occurrence(s)):") + for occ in file_occurrences: + parts.append(self._replacer.render_occurrence_diff(occ, self._contents[path])) + result = "\n".join(parts) + + # shortened result closures, from least to most aggressive shortening + def make_locations_only() -> str: + return "\n".join([header] + [f" [{o.occurrence_id}] line {o.start_line}" for o in self.occurrences]) + + def make_per_file_counts() -> str: + counts = {path: sum(1 for o in self.occurrences if o.relative_path == path) for path in affected_files} + return f"{header}\nOccurrence counts per file:\n{TextOutputUtils.to_json(counts)}" + + def make_summary() -> str: + return header + + shortened_result_factories: list[Callable[[], str]] = [make_locations_only, make_per_file_counts, make_summary] + return TextOutputUtils.limit_length(result, max_answer_chars, shortened_result_factories) + + def apply(self, code_editor: "CodeEditor", occurrences: list[ReplacementOccurrence]) -> MultiFileReplacementResult: + """ + Applies the given (selected) occurrences. + + :param code_editor: the code editor through which to modify the files + :param occurrences: the occurrences to replace (obtained from `select` or `select_all_guarded`) + :return: the result + :raises ValueError: if a file's content changed such that a selected occurrence no longer resolves + (the file is then not modified) + """ + from serena.code_editor import EditedFileContext + + occurrences_by_file: dict[str, list[ReplacementOccurrence]] = {} + for occ in occurrences: + occurrences_by_file.setdefault(occ.relative_path, []).append(occ) + + for path, file_occurrences in occurrences_by_file.items(): + with EditedFileContext(path, code_editor) as context: + original_content = context.get_original_content() + if original_content != self._contents[path]: + # the editor's view differs from what was scanned (e.g. line-ending normalization); + # re-derive the occurrences from the authoritative content and re-validate by id + fresh_by_id = { + o.occurrence_id: o for o in self._replacer.find_occurrences([(path, original_content)], self._needle, self._repl) + } + try: + file_occurrences = [fresh_by_id[o.occurrence_id] for o in file_occurrences] + except KeyError as e: + raise ValueError( + f"The content of {path} changed while replacing (occurrence {e} no longer resolves); " + f"the file was NOT modified. Re-run with dry_run=True for current ids." + ) from e + context.set_updated_content(self._replacer.apply_to_content(original_content, file_occurrences)) + + return MultiFileReplacementResult({path: len(occs) for path, occs in occurrences_by_file.items()}) + + def find_text_coordinates(content: str, regex: str, require_unique: bool = False) -> TextCoordinates | None: """ Finds the line and column number of the first match of a regex pattern in the given content. @@ -669,3 +938,39 @@ def find_text_coordinates(content: str, regex: str, require_unique: bool = False index_in_content = match.start(1) line, col = TextUtils.get_line_col_from_index(content, index_in_content) return TextCoordinates(line, col) + + +class TextOutputUtils: + @staticmethod + def to_json(x: Any) -> str: + return json.dumps(x, ensure_ascii=False) + + @staticmethod + def limit_length( + result: str, + max_answer_chars: int, + shortened_result_factories: list[Callable[[], str]] | None = None, + ) -> str: + """Limit the length of the result string, optionally trying progressively shorter versions. + + :param result: the full result string + :param max_answer_chars: maximum allowed characters; if exceeded, attempt to use shortened versions + :param shortened_result_factories: optional list of closures, each producing a progressively shorter + version of the result. They are tried in order until one fits within ``max_answer_chars``. + :return: the result string, potentially replaced by a shortened version + """ + if max_answer_chars <= 0: + raise ValueError(f"max_answer_chars must be positive; got: {max_answer_chars=}") + if (n_chars := len(result)) > max_answer_chars: + too_long_msg = ( + f"The answer is too long ({n_chars} characters). " + "You can adjust your query or raise the max_answer_chars parameter." + ) + if shortened_result_factories is not None: + # try each shortening closure in order; + for make_shorter in shortened_result_factories: + shortened = make_shorter() + candidate = f"{too_long_msg}\n{shortened}" + if len(candidate) <= max_answer_chars: + return candidate + result = too_long_msg + return result diff --git a/src/solidlsp/initialize_params.py b/src/solidlsp/initialize_params.py index f71f6e5d..87fd36d3 100644 --- a/src/solidlsp/initialize_params.py +++ b/src/solidlsp/initialize_params.py @@ -40,10 +40,11 @@ class InitializeParamsBuilder(ABC): class DefaultInitializeParamsBuilder(InitializeParamsBuilder): - def __init__(self, ls: "SolidLanguageServer", set_workspace_folders: bool = True): + def __init__(self, ls: "SolidLanguageServer", set_workspace_folders: bool = True, set_root_uri: bool = True): super().__init__() self._ls = ls self._set_workspace_folders = set_workspace_folders + self._set_root_uri = set_root_uri @staticmethod def _create_workspace_folder_entry(path: str) -> WorkspaceFolder: @@ -54,10 +55,22 @@ class DefaultInitializeParamsBuilder(InitializeParamsBuilder): root_abs_path = self._ls.repository_root_path self._set("processId", os.getpid()) - self._set("rootPath", root_abs_path) - self._set("rootUri", pathlib.Path(root_abs_path).as_uri()) self._set("clientInfo", {"name": "Serena"}) + # Some language servers treat rootUri as an additional analysis root on top of + # workspaceFolders, with no de-duplication, which can cause unbounded indexing. + # When set_root_uri is False, rootUri/rootPath are omitted so that workspaceFolders + # alone determine the analysis roots. + if self._set_root_uri: + self._set("rootPath", root_abs_path) + self._set("rootUri", pathlib.Path(root_abs_path).as_uri()) + else: + # Some servers reject initialize when the key is absent + # ("params.rootUri must not be undefined"). Send explicit null so the + # field is present but not used as an analysis root (#2045). + self._set("rootPath", None) + self._set("rootUri", None) + if self._set_workspace_folders: abs_workspace_paths = self._ls.config.get_absolute_workspace_folders(root_abs_path) log.info("Workspace folders: %s", abs_workspace_paths) diff --git a/src/solidlsp/language_servers/csharp_language_server.py b/src/solidlsp/language_servers/csharp_language_server.py index bef929d3..a411cc78 100644 --- a/src/solidlsp/language_servers/csharp_language_server.py +++ b/src/solidlsp/language_servers/csharp_language_server.py @@ -257,7 +257,7 @@ class CSharpLanguageServer(SolidLanguageServer): return hover def _document_symbols_cache_fingerprint(self) -> Hashable | None: - normalize_symbol_name_version = 1 + normalize_symbol_name_version = 2 return normalize_symbol_name_version def _normalize_symbol_name(self, symbol: RawDocumentSymbol, relative_file_path: str) -> str: @@ -301,15 +301,19 @@ class CSharpLanguageServer(SolidLanguageServer): "Add(int, int) : int" -> ("Add", "(int, int) : int") "ToString()" -> ("ToString", "()") "SimpleMethod" -> ("SimpleMethod", "") + "Position : (int X, string Y)" -> ("Position", ": (int X, string Y)") Returns: Tuple of (base_name, type_info) """ - # Check for property pattern: "Name : Type" - if " : " in roslyn_name and "(" not in roslyn_name: + # Check for property pattern: "Name : Type". The '(' guard must look only at the + # name segment before the first " : ", not the whole string, since a tuple type + # ("(int X, string Y)") legitimately contains parentheses. + if " : " in roslyn_name: base_name, type_part = roslyn_name.split(" : ", 1) - return base_name.strip(), f": {type_part.strip()}" + if "(" not in base_name: + return base_name.strip(), f": {type_part.strip()}" # Check for method pattern: "MethodName(params) : ReturnType" if "(" in roslyn_name: @@ -744,11 +748,24 @@ class CSharpLanguageServer(SolidLanguageServer): self.server.notify.send_notification("solution/open", {"solution": solution_uri}) log.debug(f"Opened solution file: {solution_file}") - # Find and open project files + # Find and open project files, skipping any that the project's ignore settings exclude. + # Vendored, third-party and sample trees routinely contain .csproj files that the language + # server cannot restore or build. Each one costs a project load on every server start, and + # the resulting restore failures bury the diagnostics of the projects the user cares about. project_files = [] + skipped = 0 for filename in breadth_first_file_scan(self.repository_root_path): - if filename.endswith(".csproj"): - project_files.append(filename) + if not filename.endswith(".csproj"): + continue + relative_path = os.path.relpath(filename, self.repository_root_path) + # ignore_unsupported_files=False, because a .csproj is not itself a C# source file and + # would otherwise be excluded on file type rather than by the ignore patterns. + if self.is_ignored_path(relative_path, ignore_unsupported_files=False): + skipped += 1 + continue + project_files.append(filename) + if skipped: + log.debug(f"Skipped {skipped} .csproj file(s) matched by the project's ignore settings") # Send project/open notifications for each project file if project_files: diff --git a/src/solidlsp/language_servers/dart_language_server.py b/src/solidlsp/language_servers/dart_language_server.py index 6bd890e9..c668d8ba 100644 --- a/src/solidlsp/language_servers/dart_language_server.py +++ b/src/solidlsp/language_servers/dart_language_server.py @@ -7,6 +7,7 @@ from collections.abc import Hashable from overrides import override +from solidlsp.initialize_params import DefaultInitializeParamsBuilder, InitializeParamsBuilder from solidlsp.ls import RawDocumentSymbol, SolidLanguageServer from solidlsp.lsp_protocol_handler.server import ProcessLaunchInfo from solidlsp.settings import SolidLSPSettings @@ -73,6 +74,13 @@ class DartLanguageServer(SolidLanguageServer): # via either notification it sends for this (see _start_server). self.analysis_complete = threading.Event() + def _create_initialize_params_builder(self) -> InitializeParamsBuilder: + # The Dart analysis server treats rootUri as an additional analysis root on top of + # workspaceFolders, with no de-duplication, so on a monorepo root that is not a Dart + # package the whole tree is analysed and the server burns CPU at idle (oraios/serena#2045). + # Omit rootUri/rootPath and rely on workspaceFolders alone. + return DefaultInitializeParamsBuilder(self, set_root_uri=False) + @override def _document_symbols_cache_fingerprint(self) -> Hashable: normalize_symbol_name_version = 1 diff --git a/src/solidlsp/language_servers/godot_language_server.py b/src/solidlsp/language_servers/godot_language_server.py index 52995988..6ade0c88 100644 --- a/src/solidlsp/language_servers/godot_language_server.py +++ b/src/solidlsp/language_servers/godot_language_server.py @@ -10,8 +10,9 @@ The editor must be open with its built-in language server enabled (default). import logging import os from collections.abc import Callable +from typing import Any -from solidlsp.ls import SolidLanguageServer +from solidlsp.ls import DocumentSymbols, LSPFileBuffer, SolidLanguageServer from solidlsp.ls_config import LanguageServerConfig from solidlsp.ls_process import LanguageServerInterface, TCPConnectionInfo, TCPLanguageServer from solidlsp.lsp_protocol_handler.server import StringDict @@ -38,6 +39,10 @@ class GodotLanguageServer(SolidLanguageServer): - ``request_timeout`` (float): seconds to wait for an LSP response (default: 30.0). """ + # Bump whenever _fix_range_end/_fix_symbol_ranges below changes, so a stale cached + # high-level result (from before this fix existed) is not served back to callers. + _DOCUMENT_SYMBOLS_CACHE_VERSION = 1 + def __init__(self, config: LanguageServerConfig, repository_root_path: str, solidlsp_settings: SolidLSPSettings) -> None: self._godot_version = self._detect_godot_version(repository_root_path) if self._godot_version is not None: @@ -137,3 +142,66 @@ class GodotLanguageServer(SolidLanguageServer): self.server.send.initialize(initialize_params) self.server.notify.initialized({}) log.info("Godot LSP initialized") + + def _build_document_symbols_from_raw_symbols(self, relative_file_path: str, file_buffer: LSPFileBuffer) -> DocumentSymbols: + """Override to correct a Godot GDScript parser off-by-one in reported end columns. + + See :meth:`_fix_range_end` for the mechanism (oraios/serena#1974). Applied here, on the + converted high-level symbols, rather than in ``_request_raw_document_symbols`` or in the + generic ``TextUtils.get_index_from_line_col`` (used by every language server): the + overshoot is specific to Godot's own parser, not a property of LSP position math in + general, and post-processing at this level only needs to invalidate the high-level + symbol cache, not the (expensive to rebuild) raw one the language server itself answers. + + TODO: gate this behind a Godot version check once the upstream parser bug is fixed + (tracked at https://github.com/godotengine/godot/issues, not yet filed there). + """ + document_symbols = super()._build_document_symbols_from_raw_symbols(relative_file_path, file_buffer) + lines = file_buffer.split_lines() + for root_symbol in document_symbols.root_symbols: + self._fix_symbol_ranges(root_symbol, lines) + return document_symbols + + def _document_symbols_cache_fingerprint(self) -> int: + return self._DOCUMENT_SYMBOLS_CACHE_VERSION + + @staticmethod + def _fix_range_end(rng: Any, lines: list[str]) -> None: + """Correct a Godot GDScript parser off-by-one in an LSP ``Range``'s end position, in place. + + Godot's ``gdscript_parser.cpp`` closes a node's range using the *next* lookahead token + instead of the *last consumed* one. When that lookahead is a synthesized NEWLINE token, + ``gdscript_tokenizer.cpp``'s ``newline()`` sets the token's own ``end_column`` to the + column reached *after* consuming the newline character, one past where a real content + token would end. Stacked on top of the usual one-past-the-end range convention, a symbol + whose body ends at that line reports an end column two past its last character instead of + one past it (oraios/serena#1974). + + Only that exact, measured overshoot is corrected; a larger one is not this bug and is left + alone rather than guessed at. + """ + end = rng.get("end") + if end is None: + return + end_line, end_char = end.get("line"), end.get("character") + if end_line is None or end_char is None or not (0 <= end_line < len(lines)): + return + # The correct one-past-the-end column for this line is len(lines[end_line]). + correct_end_char = len(lines[end_line]) + if end_char == correct_end_char + 1: + end["character"] = correct_end_char + + @staticmethod + def _fix_symbol_ranges(symbol: Any, lines: list[str]) -> None: + """Recursively apply :meth:`_fix_range_end` to a (raw or unified) symbol and its children.""" + location = symbol.get("location") + if location is not None: + GodotLanguageServer._fix_range_end(location.get("range", {}), lines) + symbol_range = symbol.get("range") + if symbol_range is not None: + GodotLanguageServer._fix_range_end(symbol_range, lines) + selection_range = symbol.get("selectionRange") + if selection_range is not None: + GodotLanguageServer._fix_range_end(selection_range, lines) + for child in symbol.get("children") or []: + GodotLanguageServer._fix_symbol_ranges(child, lines) diff --git a/src/solidlsp/language_servers/kotlin_language_server.py b/src/solidlsp/language_servers/kotlin_language_server.py index d9e35d47..5dca5d56 100644 --- a/src/solidlsp/language_servers/kotlin_language_server.py +++ b/src/solidlsp/language_servers/kotlin_language_server.py @@ -6,7 +6,7 @@ You can configure the following options in ls_specific_settings (in serena_confi ls_specific_settings: kotlin: ls_path: '/path/to/bin/intellij-server' # Custom path to Kotlin Language Server executable - kotlin_lsp_version: '262.9593.0' # Kotlin Language Server version (default: current bundled version) + kotlin_lsp_version: '263.4702.0' # Kotlin Language Server version (default: current bundled version) jvm_options: '-Xmx2G' # JVM options for Kotlin Language Server (default: -Xmx2G) Example configuration for large projects: @@ -50,7 +50,7 @@ KOTLIN_LSP_ALLOWED_HOSTS = ("download-cdn.jetbrains.com",) # DEFAULT_* — bumped on upgrades; goes into a versioned subdir. # NOTE: After changing either pinned version, run scripts/update_downloaded_dependency_hashes.py. INITIAL_KOTLIN_LSP_VERSION = "261.13587.0" -DEFAULT_KOTLIN_LSP_VERSION = "262.9593.0" +DEFAULT_KOTLIN_LSP_VERSION = "263.4702.0" # Versions before this one use kotlin-lsp-{version}-{platform}.zip and a kotlin-lsp script. # Starting with 262.4739.0, JetBrains publishes kotlin-server archives with platform-specific diff --git a/src/solidlsp/ls.py b/src/solidlsp/ls.py index 7c20af86..26dda831 100644 --- a/src/solidlsp/ls.py +++ b/src/solidlsp/ls.py @@ -546,7 +546,7 @@ class SolidLanguageServer(ABC): self._published_diagnostics_condition = threading.Condition() # initialise symbol caches - self.cache_dir = Path(self._solidlsp_settings.project_data_path) / self.CACHE_FOLDER_NAME / self.language_id + self.cache_dir = Path(self._solidlsp_settings.project_data_path) / self.CACHE_FOLDER_NAME / self.ls_id.get_key() self.cache_dir.mkdir(parents=True, exist_ok=True) # * raw document symbols cache self._ls_specific_raw_document_symbols_cache_version = cache_version_raw_document_symbols @@ -2984,10 +2984,8 @@ class SolidLanguageServer(ABC): high_level_fingerprint = self._document_symbols_cache_fingerprint() if high_level_fingerprint is not None: version.append(high_level_fingerprint) - raw_fingerprint = self._raw_document_symbols_cache_fingerprint() - if raw_fingerprint is not None: - version.append(raw_fingerprint) - return version[0] if len(version) == 1 else tuple(version) + version.append(self._raw_document_symbols_cache_version()) + return tuple(version) def _save_raw_document_symbols_cache(self) -> None: cache_file = self.cache_dir / self.RAW_DOCUMENT_SYMBOL_CACHE_FILENAME diff --git a/src/solidlsp/ls_config.py b/src/solidlsp/ls_config.py index 80afde88..49c295ea 100644 --- a/src/solidlsp/ls_config.py +++ b/src/solidlsp/ls_config.py @@ -10,7 +10,7 @@ import logging import os import re import threading -from collections.abc import Iterable +from collections.abc import Iterable, Iterator from dataclasses import dataclass, field from enum import Enum from functools import cache @@ -1066,6 +1066,14 @@ class LanguageServerRegistry: return self._registered_language_servers[key] raise ValueError(f"Unknown language server key: '{key}'; Valid keys: {self.get_keys()}") + def iter_registered_ls_ids(self) -> Iterator[LanguageServerIdLike]: + """ + Iterate over all registered language servers (built-in + externally-registered via + entry points). Order follows ``get_keys()`` (alphabetical). + """ + for key in self.get_keys(): + yield self._registered_language_servers[key] + def register(self, ls_id: LanguageServerIdLike, allow_override: bool = False) -> None: """ :param ls_id: the identifier to register diff --git a/src/solidlsp/resources/downloaded_dependency_hashes.json b/src/solidlsp/resources/downloaded_dependency_hashes.json index d1006ff8..94758740 100644 --- a/src/solidlsp/resources/downloaded_dependency_hashes.json +++ b/src/solidlsp/resources/downloaded_dependency_hashes.json @@ -16,5 +16,11 @@ "https://download-cdn.jetbrains.com/language-server/kotlin-server/262.9593.0/kotlin-server-262.9593.0.tar.gz": "2d99d8e198fbe4aa8f4481e37799724ce94803b4ea12a60b416040e3fcd7cc5e", "https://download-cdn.jetbrains.com/language-server/kotlin-server/262.9593.0/kotlin-server-262.9593.0-aarch64.tar.gz": "2317831c6e5607d05b7ebc1da655330125ce0e3d66fbf24517dfce442debc14e", "https://download-cdn.jetbrains.com/language-server/kotlin-server/262.9593.0/kotlin-server-262.9593.0.sit": "17369fda97c85418ac24ab38a9df56b21522a3468dfe193832fe455c13920745", - "https://download-cdn.jetbrains.com/language-server/kotlin-server/262.9593.0/kotlin-server-262.9593.0-aarch64.sit": "6ba6021a706b21e64cef33f7e2b79f187c0910320722bb2d3ed05ad1115ec43f" -} + "https://download-cdn.jetbrains.com/language-server/kotlin-server/262.9593.0/kotlin-server-262.9593.0-aarch64.sit": "6ba6021a706b21e64cef33f7e2b79f187c0910320722bb2d3ed05ad1115ec43f", + "https://download-cdn.jetbrains.com/language-server/kotlin-server/263.4702.0/kotlin-server-263.4702.0.win.zip": "a9b471b16025b1bfb3b0a097862580abb40e3c35406c44242c18b1d70f5d0e44", + "https://download-cdn.jetbrains.com/language-server/kotlin-server/263.4702.0/kotlin-server-263.4702.0-aarch64.win.zip": "3bf008d8c94fa70eb13fc998eaa42f29b9d13f368984d4cec46277808f94e1de", + "https://download-cdn.jetbrains.com/language-server/kotlin-server/263.4702.0/kotlin-server-263.4702.0.tar.gz": "1e11d2e5fefbf9ea215ad8dd6be95f2222897cd086e8cb7a661a52084a590405", + "https://download-cdn.jetbrains.com/language-server/kotlin-server/263.4702.0/kotlin-server-263.4702.0-aarch64.tar.gz": "ec7cb254a6662a07fff9f10e4365226afab6c40008f8a974c10ac5e785d6510f", + "https://download-cdn.jetbrains.com/language-server/kotlin-server/263.4702.0/kotlin-server-263.4702.0.sit": "62ab735947b1c855b505f64f5db8fbd7ff0b52a35ab1897938c6dbfc7b24c8a3", + "https://download-cdn.jetbrains.com/language-server/kotlin-server/263.4702.0/kotlin-server-263.4702.0-aarch64.sit": "95da3fc6d3b9092c7616345044a05edb85e5408dc648d081e4e433595c892bec" +} \ No newline at end of file diff --git a/test/serena/config/test_serena_config.py b/test/serena/config/test_serena_config.py index e9a3c46d..c0f21d9e 100644 --- a/test/serena/config/test_serena_config.py +++ b/test/serena/config/test_serena_config.py @@ -4,19 +4,21 @@ import shutil import tempfile from copy import deepcopy from pathlib import Path +from uuid import UUID import pytest from serena.agent import SerenaAgent from serena.config.serena_config import ( DEFAULT_PROJECT_SERENA_FOLDER_LOCATION, - LanguageBackend, + AgentInterface, ProjectConfig, RegisteredProject, SerenaConfig, SerenaConfigError, ) from serena.constants import PROJECT_TEMPLATE_FILE, SERENA_MANAGED_DIR_NAME +from serena.language_backend import BuiltinLanguageBackend from serena.project import MemoryManager, Project from solidlsp.ls_config import LanguageServerId from test.conftest import create_default_serena_config @@ -176,15 +178,16 @@ class TestProjectConfigLanguageBackend: config = ProjectConfig( project_name="test", language_servers=[LanguageServerId.PYTHON], - language_backend=LanguageBackend.JETBRAINS, + language_backend=BuiltinLanguageBackend.JETBRAINS.get_instance(), ) - assert config.language_backend == LanguageBackend.JETBRAINS + assert config.language_backend is not None + assert config.language_backend.is_jetbrains() def test_language_backend_roundtrips_through_yaml(self): config = ProjectConfig( project_name="test", language_servers=[LanguageServerId.PYTHON], - language_backend=LanguageBackend.JETBRAINS, + language_backend=BuiltinLanguageBackend.JETBRAINS.get_instance(), ) d = config._to_yaml_dict() assert d["language_backend"] == "JetBrains" @@ -205,7 +208,8 @@ class TestProjectConfigLanguageBackend: data["languages"] = ["python"] data["language_backend"] = "JetBrains" config = ProjectConfig._from_dict(data, local_override_keys=[]) - assert config.language_backend == LanguageBackend.JETBRAINS + assert config.language_backend is not None + assert config.language_backend.is_jetbrains() def test_language_backend_none_when_missing_from_dict(self): """Test that _from_dict handles missing language_backend gracefully.""" @@ -218,22 +222,84 @@ class TestProjectConfigLanguageBackend: assert config.language_backend is None +class TestAgentInterface: + """Tests for the agent_interface setting (global and per project).""" + + @staticmethod + def _project_config(agent_interface: AgentInterface | None) -> ProjectConfig: + return ProjectConfig(project_name="test", language_servers=[LanguageServerId.PYTHON], agent_interface=agent_interface) + + def test_agent_interface_roundtrips_through_project_yaml(self): + assert self._project_config(AgentInterface.REPL)._to_yaml_dict()["agent_interface"] == "REPL" + assert self._project_config(None)._to_yaml_dict()["agent_interface"] is None + + def test_agent_interface_parsed_from_project_dict(self): + data, _ = ProjectConfig._load_yaml_dict(PROJECT_TEMPLATE_FILE) + data["project_name"] = "test" + data["languages"] = ["python"] + data["agent_interface"] = "repl" # case-insensitive + assert ProjectConfig._from_dict(data, local_override_keys=[]).agent_interface == AgentInterface.REPL + data.pop("agent_interface") + assert ProjectConfig._from_dict(data, local_override_keys=[]).agent_interface is None + + def test_determine_agent_interface_precedence(self): + # default + assert SerenaConfig().determine_agent_interface() == AgentInterface.TOOLS + assert SerenaConfig().determine_agent_interface(self._project_config(None)) == AgentInterface.TOOLS + # global configuration + assert SerenaConfig(agent_interface=AgentInterface.REPL).determine_agent_interface() == AgentInterface.REPL + # project configuration takes precedence + config = SerenaConfig(agent_interface=AgentInterface.REPL) + assert config.determine_agent_interface(self._project_config(AgentInterface.TOOLS)) == AgentInterface.TOOLS + assert config.determine_agent_interface(self._project_config(None)) == AgentInterface.REPL + + def test_repl_toolset_is_fixed_and_repl_follows_project_activation(self): + """ + In REPL mode, neither the exposed nor the active toolset is affected by tool inclusion/exclusion definitions + (here: the project's exclusions and read-only setting), whereas the REPL's API scope follows the active project. + """ + config, name = _make_config_with_project("test_proj") + config.agent_interface = AgentInterface.REPL + project_config = config.projects[0].project_config + project_config.excluded_tools = ["initial_instructions", "serena_repl"] + project_config.excluded_apis = ["mem"] + project_config.read_only = True + + agent = SerenaAgent(project=None, serena_config=config) + try: + # before activation: the fixed toolset and the full set of facades + fixed_toolset = {"serena_repl", "initial_instructions", "activate_project"} + assert {t.get_name() for t in agent.get_exposed_tool_instances()} == fixed_toolset + assert set(agent.get_active_tool_names()) == fixed_toolset + overview = agent.get_repl().entrypoint.overview() + assert "s.mem" in overview + # the dashboard is disabled in the test configuration, so opening it is not offered + assert "s.cfg" in overview and "open_dashboard" not in overview + + # after activation: the toolset is unchanged, the REPL reflects the project's API exclusions + agent.activate_project_from_path_or_name(name) + assert set(agent.get_active_tool_names()) == fixed_toolset + assert "s.mem" not in agent.get_repl().entrypoint.overview() + finally: + agent.on_shutdown(timeout=5) + + def _make_config_with_project( project_name: str, - language_backend: LanguageBackend | None = None, - global_backend: LanguageBackend = LanguageBackend.LSP, + language_backend: BuiltinLanguageBackend | None = None, + global_backend: BuiltinLanguageBackend = BuiltinLanguageBackend.LSP, ) -> tuple[SerenaConfig, str]: """Create a SerenaConfig with a single registered project and return (config, project_name).""" config = SerenaConfig( log_level=logging.ERROR, - language_backend=global_backend, + language_backend=global_backend.get_instance(), ).with_headless_mode_overrides() project = Project( project_root=str(Path(__file__).parent.parent / "resources" / "repos" / "python" / "test_repo"), project_config=ProjectConfig( project_name=project_name, language_servers=[LanguageServerId.PYTHON], - language_backend=language_backend, + language_backend=language_backend.get_instance() if language_backend is not None else None, ), serena_config=config, ) @@ -246,7 +312,7 @@ class TestEffectiveLanguageBackend: def test_default_backend_is_global(self): """When no project override, effective backend matches global config.""" - config, name = _make_config_with_project("test_proj", language_backend=None, global_backend=LanguageBackend.LSP) + config, name = _make_config_with_project("test_proj", language_backend=None, global_backend=BuiltinLanguageBackend.LSP) agent = SerenaAgent(project=name, serena_config=config) try: assert agent.get_language_backend().is_lsp() @@ -256,7 +322,7 @@ class TestEffectiveLanguageBackend: def test_project_overrides_global_backend(self): """When startup project has language_backend set, it overrides the global.""" config, name = _make_config_with_project( - "test_jetbrains", language_backend=LanguageBackend.JETBRAINS, global_backend=LanguageBackend.LSP + "test_jetbrains", language_backend=BuiltinLanguageBackend.JETBRAINS, global_backend=BuiltinLanguageBackend.LSP ) agent = SerenaAgent(project=name, serena_config=config) try: @@ -268,18 +334,18 @@ class TestEffectiveLanguageBackend: """When no startup project is provided, effective backend is the global one.""" config = SerenaConfig( log_level=logging.ERROR, - language_backend=LanguageBackend.LSP, + language_backend=BuiltinLanguageBackend.LSP.get_instance(), ).with_headless_mode_overrides() agent = SerenaAgent(project=None, serena_config=config) try: - assert agent.get_language_backend() == LanguageBackend.LSP + assert agent.get_language_backend().is_lsp() finally: agent.on_shutdown(timeout=5) def test_activate_project_rejects_backend_mismatch(self): """Post-init activation of a project with mismatched backend raises ValueError.""" # Start with LSP backend - config, name = _make_config_with_project("lsp_proj", language_backend=None, global_backend=LanguageBackend.LSP) + config, name = _make_config_with_project("lsp_proj", language_backend=None, global_backend=BuiltinLanguageBackend.LSP) # Add a second project that requires JetBrains jb_project = Project( @@ -287,7 +353,7 @@ class TestEffectiveLanguageBackend: project_config=ProjectConfig( project_name="jb_proj", language_servers=[LanguageServerId.JAVA], - language_backend=LanguageBackend.JETBRAINS, + language_backend=BuiltinLanguageBackend.JETBRAINS.get_instance(), ), serena_config=config, ) @@ -300,9 +366,38 @@ class TestEffectiveLanguageBackend: finally: agent.on_shutdown(timeout=5) + def test_activate_project_switches_backend_with_repl_interface(self): + """With the REPL interface, post-init activation of a project with a different backend switches the backend.""" + config, name = _make_config_with_project("lsp_proj", language_backend=None, global_backend=BuiltinLanguageBackend.LSP) + config.agent_interface = AgentInterface.REPL + jb_project = Project( + project_root=str(Path(__file__).parent.parent / "resources" / "repos" / "java" / "test_repo"), + project_config=ProjectConfig( + project_name="jb_proj", + language_servers=[LanguageServerId.JAVA], + language_backend=BuiltinLanguageBackend.JETBRAINS.get_instance(), + ), + serena_config=config, + ) + config.projects.append(RegisteredProject.from_project_instance(jb_project)) + + agent = SerenaAgent(project=name, serena_config=config) + try: + assert agent.get_language_backend().is_lsp() + assert "s.lsp" in agent.get_repl().entrypoint.overview() + + # the backend and everything depending on it follow the activated project + agent.activate_project_from_path_or_name("jb_proj") + assert agent.get_language_backend().is_jetbrains() + overview = agent.get_repl().entrypoint.overview() + assert "s.jb" in overview and "s.lsp" not in overview + assert "jetbrains" in [m.name for m in agent.get_active_modes().get_modes(include_background_base_modes=True)] + finally: + agent.on_shutdown(timeout=5) + def test_activate_project_allows_matching_backend(self): """Post-init activation of a project with matching backend succeeds.""" - config, name = _make_config_with_project("lsp_proj", language_backend=None, global_backend=LanguageBackend.LSP) + config, name = _make_config_with_project("lsp_proj", language_backend=None, global_backend=BuiltinLanguageBackend.LSP) # Add a second project that also uses LSP lsp_project2 = Project( @@ -310,7 +405,7 @@ class TestEffectiveLanguageBackend: project_config=ProjectConfig( project_name="lsp_proj2", language_servers=[LanguageServerId.PYTHON], - language_backend=LanguageBackend.LSP, + language_backend=BuiltinLanguageBackend.LSP.get_instance(), ), serena_config=config, ) @@ -325,7 +420,7 @@ class TestEffectiveLanguageBackend: def test_activate_project_allows_none_backend(self): """Post-init activation of a project with no backend override succeeds.""" - config, name = _make_config_with_project("lsp_proj", language_backend=None, global_backend=LanguageBackend.LSP) + config, name = _make_config_with_project("lsp_proj", language_backend=None, global_backend=BuiltinLanguageBackend.LSP) # Add a second project with no backend override proj2 = Project( @@ -543,6 +638,36 @@ class TestSerenaConfigLoadSave: config = SerenaConfig.from_config_file(generate_if_missing=False) assert config.projects == [] + @pytest.mark.parametrize("setting", ["", "auth_secret: null\n", 'auth_secret: ""\n']) + def test_unset_auth_secret_is_generated_and_persisted(self, setting: str) -> None: + # load an existing configuration without a usable secret + self.master_config_path.write_text("projects: []\n" + setting) + config = SerenaConfig.from_config_file(generate_if_missing=False) + + # subsequent loads retain the generated random UUID + assert UUID(config.auth_secret).version == 4 + assert SerenaConfig.from_config_file(generate_if_missing=False).auth_secret == config.auth_secret + + def test_configured_auth_secret_is_preserved(self) -> None: + # retain a user-provided secret across loading and migration + self.master_config_path.write_text("projects: []\nauth_secret: custom-secret\n") + assert SerenaConfig.from_config_file(generate_if_missing=False).auth_secret == "custom-secret" + assert SerenaConfig.from_config_file(generate_if_missing=False).auth_secret == "custom-secret" + + def test_new_config_has_persistent_auth_secret(self) -> None: + # generate the configuration from the template and retain its secret + config = SerenaConfig.from_config_file() + assert UUID(config.auth_secret).version == 4 + assert SerenaConfig.from_config_file().auth_secret == config.auth_secret + + def test_direct_config_instances_have_distinct_auth_secrets(self) -> None: + # directly constructed configurations receive independent secrets + first = SerenaConfig() + second = SerenaConfig() + assert UUID(first.auth_secret).version == 4 + assert UUID(second.auth_secret).version == 4 + assert first.auth_secret != second.auth_secret + def test_malformed_project_is_skipped_with_warning(self, caplog): """A malformed project.yml must not abort loading of the others.""" good_project = self._make_project_dir( diff --git a/test/serena/test_analytics_imports.py b/test/serena/test_analytics_imports.py new file mode 100644 index 00000000..ec18f378 --- /dev/null +++ b/test/serena/test_analytics_imports.py @@ -0,0 +1,29 @@ +"""The anthropic package must only be loaded when the Anthropic token counter is used (#2012).""" + +import subprocess +import sys +from unittest.mock import MagicMock + +import pytest + +from serena.analytics import AnthropicTokenCount + + +@pytest.mark.parametrize("module", ["serena.cli", "serena.agent", "serena.analytics"]) +def test_importing_serena_does_not_load_anthropic(module: str) -> None: + code = f"import sys, {module}; print('anthropic' in sys.modules)" + result = subprocess.run([sys.executable, "-c", code], capture_output=True, text=True, check=True, timeout=120) + assert result.stdout.strip() == "False", f"importing {module} loaded the anthropic package" + + +def test_anthropic_token_count_sends_a_plain_user_message() -> None: + estimator = AnthropicTokenCount.__new__(AnthropicTokenCount) + estimator._model_name = "claude-sonnet-4-20250514" + estimator._anthropic_client = MagicMock() + estimator._anthropic_client.messages.count_tokens.return_value = MagicMock(input_tokens=7) + + assert estimator.estimate_token_count("hello") == 7 + estimator._anthropic_client.messages.count_tokens.assert_called_once_with( + model="claude-sonnet-4-20250514", + messages=[{"role": "user", "content": "hello"}], + ) diff --git a/test/serena/test_cfg_api.py b/test/serena/test_cfg_api.py new file mode 100644 index 00000000..46ea3cbd --- /dev/null +++ b/test/serena/test_cfg_api.py @@ -0,0 +1,25 @@ +""" +Tests for the configuration facade API. +""" + +from unittest.mock import MagicMock + +from serena.repl.api.cfg_api import ConfigApi +from serena.repl.facade import ApiScope, Facade + + +def test_facade_exposes_config_operations() -> None: + facade = Facade.from_api(ConfigApi(MagicMock()), ApiScope()) + assert facade.name == "cfg" + assert set(facade.enabled_method_names) == {"get_current_config", "open_dashboard"} + assert not any(facade.get_method(name).info.can_edit for name in facade.enabled_method_names) + + +def test_operations_delegate_to_agent() -> None: + agent = MagicMock() + agent.get_current_config_overview.return_value = "overview" + agent.open_dashboard.return_value = False + agent.get_dashboard_url.return_value = "http://localhost:1" + api = ConfigApi(agent) + assert api.get_current_config() == "overview" + assert "http://localhost:1" in api.open_dashboard() diff --git a/test/serena/test_cli_project_remove.py b/test/serena/test_cli_project_remove.py new file mode 100644 index 00000000..a0e64066 --- /dev/null +++ b/test/serena/test_cli_project_remove.py @@ -0,0 +1,123 @@ +"""Tests for the CLI's ``project remove`` command.""" + +import shutil +import tempfile +from pathlib import Path + +import pytest +from click.testing import CliRunner + +from serena.cli import ProjectCommands +from serena.config.serena_config import SerenaConfig +from serena.constants import SERENA_MANAGED_DIR_NAME + + +class TestProjectRemove: + """ + Drives ``serena project remove`` against a temporary Serena configuration file, so the + user's real project registry is never touched. + """ + + @pytest.fixture(autouse=True) + def setup(self, monkeypatch): + self.test_dir = Path(tempfile.mkdtemp()) + self.master_config_path = self.test_dir / "serena_config.yml" + monkeypatch.setattr( + SerenaConfig, + "_determine_config_file_path", + classmethod(lambda cls: str(self.master_config_path)), + ) + self.runner = CliRunner() + + def teardown_method(self): + shutil.rmtree(self.test_dir, ignore_errors=True) + + def _make_project_dir(self, dir_name: str, project_name: str | None = None) -> Path: + project_dir = self.test_dir / dir_name + (project_dir / SERENA_MANAGED_DIR_NAME).mkdir(parents=True) + (project_dir / SERENA_MANAGED_DIR_NAME / "project.yml").write_text( + f'project_name: "{project_name or dir_name}"\nlanguages: ["python"]\n' + ) + return project_dir + + def _write_master_config(self, project_paths: list[Path]) -> None: + self.master_config_path.write_text("projects:\n" + "".join(f" - {p}\n" for p in project_paths)) + + def _registered_roots(self) -> set[Path]: + config = SerenaConfig.from_config_file(generate_if_missing=False) + return {Path(project.project_root) for project in config.projects} + + def test_remove_by_name_unregisters_only_that_project(self): + keep = self._make_project_dir("keep") + drop = self._make_project_dir("drop") + self._write_master_config([keep, drop]) + + result = self.runner.invoke(ProjectCommands.remove, ["drop"]) + + assert result.exit_code == 0, f"Command failed: {result.output}" + assert self._registered_roots() == {keep.resolve()} + + def test_remove_by_path_unregisters_only_that_project(self): + keep = self._make_project_dir("keep") + drop = self._make_project_dir("drop") + self._write_master_config([keep, drop]) + + result = self.runner.invoke(ProjectCommands.remove, [str(drop)]) + + assert result.exit_code == 0, f"Command failed: {result.output}" + assert self._registered_roots() == {keep.resolve()} + + def test_remove_names_the_project_it_removed(self): + drop = self._make_project_dir("drop") + self._write_master_config([drop]) + + result = self.runner.invoke(ProjectCommands.remove, [str(drop)]) + + assert result.exit_code == 0, f"Command failed: {result.output}" + assert "drop" in result.output + assert str(drop.resolve()) in result.output + + def test_remove_keeps_the_project_configuration_file_on_disk(self): + """Unregistering must not delete the project's own files; re-registering it must remain possible.""" + drop = self._make_project_dir("drop") + self._write_master_config([drop]) + + result = self.runner.invoke(ProjectCommands.remove, [str(drop)]) + + assert result.exit_code == 0, f"Command failed: {result.output}" + assert (drop / SERENA_MANAGED_DIR_NAME / "project.yml").is_file() + + def test_remove_unknown_project_fails_without_changing_the_registry(self): + keep = self._make_project_dir("keep") + self._write_master_config([keep]) + + result = self.runner.invoke(ProjectCommands.remove, ["no_such_project"]) + + assert result.exit_code != 0 + assert "no_such_project" in result.output + assert self._registered_roots() == {keep.resolve()} + + def test_remove_by_path_picks_the_entry_at_that_path_when_names_collide(self): + """Two directories may carry the same ``project_name``; a path must remove the entry at that path.""" + first = self._make_project_dir("first_dir", project_name="twin") + second = self._make_project_dir("second_dir", project_name="twin") + self._write_master_config([first, second]) + + result = self.runner.invoke(ProjectCommands.remove, [str(second)]) + + assert result.exit_code == 0, f"Command failed: {result.output}" + assert self._registered_roots() == {first.resolve()} + + def test_remove_by_ambiguous_name_fails_with_a_message_rather_than_a_traceback(self): + first = self._make_project_dir("first_dir", project_name="twin") + second = self._make_project_dir("second_dir", project_name="twin") + self._write_master_config([first, second]) + + result = self.runner.invoke(ProjectCommands.remove, ["twin"]) + + assert result.exit_code != 0 + assert result.exception is None or isinstance(result.exception, SystemExit), ( + f"Expected a handled CLI error, got: {result.exception!r}" + ) + assert "twin" in result.output + assert self._registered_roots() == {first.resolve(), second.resolve()} diff --git a/test/serena/test_code_editor_atomic_writes.py b/test/serena/test_code_editor_atomic_writes.py new file mode 100644 index 00000000..77b55237 --- /dev/null +++ b/test/serena/test_code_editor_atomic_writes.py @@ -0,0 +1,199 @@ +"""Tests that saving an edited source file cannot destroy the previous content (issue #1958). + +``CodeEditor._save_edited_file`` is the third of the three direct ``open(path, "w")`` writes the +issue lists; the two in ``MemoryManager`` were addressed in #1969, which is also where +``write_file_atomic`` comes from. + +The editor is exercised through a stub subclass so that the inherited save path can be driven +without a language server. +""" + +import os +import stat +import sys +from collections.abc import Iterator +from contextlib import contextmanager +from typing import Any + +import pytest + +from serena.code_editor import CodeEditor +from serena.language_backend import BuiltinLanguageBackend +from serena.util import file_system + + +class _InMemoryEditedFile(CodeEditor.EditedFile): + """An ``EditedFile`` that holds its contents in memory.""" + + def __init__(self, relative_path: str, contents: str) -> None: + super().__init__(relative_path) + self._contents = contents + + def get_contents(self) -> str: + return self._contents + + def set_contents(self, contents: str) -> None: + self._contents = contents + + def delete_text_between_positions(self, start_pos: Any, end_pos: Any) -> None: + raise NotImplementedError + + def insert_text_at_position(self, pos: Any, text: str) -> None: + raise NotImplementedError + + +class _StubCodeEditor(CodeEditor[Any]): + """A ``CodeEditor`` whose only inherited behaviour under test is the file-saving path.""" + + class DummyProject: + """A dummy project object with only the attributes needed to construct a ``CodeEditor``.""" + + def __init__(self) -> None: + self.language_backend = BuiltinLanguageBackend.LSP.get_instance() + + def __init__(self, project_root: str, encoding: str = "utf-8", newline: str | None = None) -> None: + self.project = self.DummyProject() + self.project_root = project_root + self.encoding = encoding + self.newline = newline + + @contextmanager + def _open_file_context(self, relative_path: str) -> Iterator[CodeEditor.EditedFile]: + abs_path = os.path.join(self.project_root, relative_path) + with open(abs_path, encoding=self.encoding) as f: + contents = f.read() + yield _InMemoryEditedFile(relative_path, contents) + + def _find_unique_symbol(self, name_path: str, relative_file_path: str) -> Any: + raise NotImplementedError + + def rename_symbol(self, name_path: str, relative_path: str, new_name: str) -> str: + raise NotImplementedError + + +class TestSourceFileSaveIsAtomic: + def _editor(self, tmp_path: Any, **kwargs: Any) -> _StubCodeEditor: + return _StubCodeEditor(str(tmp_path), **kwargs) + + def test_interrupted_save_keeps_the_previous_file_content(self, tmp_path, monkeypatch): + """A crash partway through the write must leave the file holding its old content. + + The crash is injected into the temp-file write that ``write_file_atomic`` performs, so on + a non-atomic implementation nothing raises at all and this fails with ``DID NOT RAISE``, + which is exactly the state the issue describes: the real file is truncated first and there + is no intermediate copy to fall back to. + """ + source = tmp_path / "module.py" + original = "def original():\n return 1\n" * 40 + source.write_text(original, encoding="utf-8") + + real_fdopen = os.fdopen + + def crashing_fdopen(fd: int, *args: Any, **kwargs: Any) -> Any: + f = real_fdopen(fd, *args, **kwargs) + real_write = f.write + + def crashing_write(data: str) -> int: + real_write(data[: len(data) // 4]) + f.flush() + raise RuntimeError("simulated crash mid-write") + + f.write = crashing_write + return f + + monkeypatch.setattr(file_system.os, "fdopen", crashing_fdopen) + + editor = self._editor(tmp_path) + with pytest.raises(RuntimeError, match="simulated crash mid-write"): + with editor.edited_file_context("module.py") as edited: + edited.set_contents("def replacement():\n return 2\n" * 40) + + assert source.read_text(encoding="utf-8") == original + assert list(tmp_path.iterdir()) == [source], "the partial temp file must not be left behind" + + def test_successful_save_writes_the_new_content(self, tmp_path): + """Control: the ordinary path still writes what was asked for.""" + source = tmp_path / "module.py" + source.write_text("old\n", encoding="utf-8") + + editor = self._editor(tmp_path) + with editor.edited_file_context("module.py") as edited: + edited.set_contents("new\n") + + assert source.read_text(encoding="utf-8") == "new\n" + assert list(tmp_path.iterdir()) == [source] + + def test_save_into_a_subdirectory(self, tmp_path): + """Every other test writes at the project root; the relative path is joined and resolved, + so a nested file has to work the same way. + """ + package = tmp_path / "pkg" / "sub" + package.mkdir(parents=True) + source = package / "module.py" + source.write_text("old\n", encoding="utf-8") + + editor = self._editor(tmp_path) + with editor.edited_file_context("pkg/sub/module.py") as edited: + edited.set_contents("new\n") + + assert source.read_text(encoding="utf-8") == "new\n" + assert list(package.iterdir()) == [source], "no temp file may be left beside the source" + + def test_save_writes_through_a_symlinked_source_file(self, tmp_path): + """A symlinked source file must keep being written through to its target, as + ``open(path, "w")`` did; the link itself must not be replaced by a regular file. + """ + target_dir = tmp_path / "shared" + target_dir.mkdir() + target = target_dir / "shared.py" + target.write_text("old\n", encoding="utf-8") + project = tmp_path / "project" + project.mkdir() + link = project / "module.py" + try: + link.symlink_to(target) + except OSError as e: + pytest.skip(f"cannot create symlinks on this platform/permissions: {e}") + + editor = self._editor(project) + with editor.edited_file_context("module.py") as edited: + edited.set_contents("new\n") + + assert link.is_symlink(), "the source file's symlink must survive the edit" + assert target.read_text(encoding="utf-8") == "new\n", "the edit must reach the link's target" + + @pytest.mark.skipif( + sys.platform == "win32", reason="Windows does not model POSIX permission bits; chmod only toggles the read-only flag" + ) + def test_save_preserves_the_executable_bit(self, tmp_path): + script = tmp_path / "run.sh" + script.write_text("#!/bin/sh\necho old\n", encoding="utf-8") + os.chmod(script, 0o755) + + editor = self._editor(tmp_path) + with editor.edited_file_context("run.sh") as edited: + edited.set_contents("#!/bin/sh\necho new\n") + + assert stat.S_IMODE(os.stat(script).st_mode) == 0o755 + + def test_save_respects_the_configured_newline(self, tmp_path): + source = tmp_path / "module.py" + source.write_bytes(b"old\n") + + editor = self._editor(tmp_path, newline="\r\n") + with editor.edited_file_context("module.py") as edited: + edited.set_contents("a\nb\n") + + assert source.read_bytes() == b"a\r\nb\r\n" + + def test_save_respects_the_configured_encoding(self, tmp_path): + source = tmp_path / "module.py" + source.write_text("alt\n", encoding="latin-1") + + # newline is pinned so this test is about the encoding alone: LineEnding.NATIVE yields + # newline=None, under which Python translates "\n" to os.linesep on write + editor = self._editor(tmp_path, encoding="latin-1", newline="\n") + with editor.edited_file_context("module.py") as edited: + edited.set_contents("café\n") + + assert source.read_bytes() == "café\n".encode("latin-1") diff --git a/test/serena/test_edit_api.py b/test/serena/test_edit_api.py new file mode 100644 index 00000000..d961857c --- /dev/null +++ b/test/serena/test_edit_api.py @@ -0,0 +1,68 @@ +""" +Tests for the editing facade API (backend-independent parts, without a language server). +""" + +from pathlib import Path +from unittest.mock import MagicMock + +import pytest + +from serena.config.serena_config import SerenaConfig +from serena.project import Project +from serena.repl.api.edit_api import EditApi, ReplacementPreview +from serena.repl.facade import ApiScope, Facade + + +@pytest.fixture +def project(tmp_path: Path) -> Project: + (tmp_path / "a.py").write_text("x = foo(1)\ny = foo(2)\n", encoding="utf-8") + (tmp_path / "b.py").write_text("z = foo(3)\n", encoding="utf-8") + return Project.load(str(tmp_path), serena_config=SerenaConfig(gui_log_window=False, web_dashboard=False)) + + +@pytest.fixture +def api(project: Project) -> EditApi: + agent = MagicMock() + agent.get_active_project_or_raise.return_value = project + agent.serena_config.default_max_tool_answer_chars = 10000 + return EditApi(agent) + + +def test_facade_exposes_editing_operations(api: EditApi) -> None: + default_methods = { + "replace_content", + "replace_in_files", + "replace_symbol_body", + "insert_after_symbol", + "insert_before_symbol", + } + optional_methods = {"delete_lines", "replace_lines", "insert_at_line"} # the line-level operations are optional (as are the tools) + + facade = Facade.from_api(api, ApiScope()) + assert facade.name == "edit" + assert set(facade.enabled_method_names) == default_methods + for name in optional_methods: + assert facade.get_method(name).info.optional + assert all(facade.get_method(name).info.can_edit for name in default_methods | optional_methods) + + +def test_replace_in_files_dry_run_returns_inspectable_preview(api: EditApi, project: Project) -> None: + preview = api.replace_in_files("foo", "bar", mode="literal", dry_run=True) + assert isinstance(preview, ReplacementPreview) + + # the occurrences are accessible from code + assert [o.relative_path for o in preview.occurrences] == ["a.py", "a.py", "b.py"] + assert preview.affected_files == ["a.py", "b.py"] + assert all(o.replacement == "bar" for o in preview.occurrences) + + # the rendering lists the occurrence ids and diffs; nothing was modified + rendered = preview.represent() + assert "DRY RUN" in rendered + assert all(o.occurrence_id in rendered for o in preview.occurrences) + assert (Path(project.project_root) / "a.py").read_text(encoding="utf-8") == "x = foo(1)\ny = foo(2)\n" + + +def test_replace_in_files_guard_failure_includes_preview(api: EditApi) -> None: + with pytest.raises(ValueError, match="expected_count=1") as exc_info: + api.replace_in_files("foo", "bar", mode="literal", expected_count=1) + assert "b.py" in str(exc_info.value) # the listing of prospective changes is included diff --git a/test/serena/test_external_projects.py b/test/serena/test_external_projects.py new file mode 100644 index 00000000..ce74565c --- /dev/null +++ b/test/serena/test_external_projects.py @@ -0,0 +1,107 @@ +""" +End-to-end test of querying an external project through the REPL: a project server executes the language server +operations of the queried project, and the (pickled) results are usable in the querying agent's REPL. +""" + +import socket +import threading +from collections.abc import Iterator + +import pytest +from werkzeug.serving import make_server + +from serena.agent import SerenaAgent +from serena.config.serena_config import SerenaConfig +from serena.project_server import ProjectServer, ProjectServerClient +from serena.repl.api.lsp_api import LspSymbolCollection +from serena.tools import SerenaReplTool +from solidlsp.ls_config import LanguageServerId +from test.conftest import language_server_tests_enabled +from test.serena.test_serena_agent import serena_config # noqa: F401 (fixture) + + +def _free_port() -> int: + with socket.socket() as sock: + sock.bind(("127.0.0.1", 0)) + return sock.getsockname()[1] + + +@pytest.fixture +def project_server(serena_config: SerenaConfig) -> Iterator[tuple[ProjectServer, int]]: # noqa: F811 + """ + Runs a project server (backed by a real agent) on a free port for the duration of the test. + """ + config = serena_config + port = _free_port() + + # construct the server around an agent with the test configuration (the constructor would load the user's configuration) + server = ProjectServer.__new__(ProjectServer) + server._agent = SerenaAgent(serena_config=config) + server._loaded_projects_by_root = {} + server._project_load_locks_by_root = {} + server._active_project_lock = threading.Lock() + server._loaded_projects_lock = threading.Lock() + server._port = port + server._host = "127.0.0.1" + from flask import Flask + + server._app = Flask(__name__) + server._setup_routes() + + http_server = make_server("127.0.0.1", port, server._app, threaded=True) + thread = threading.Thread(target=http_server.serve_forever, daemon=True) + thread.start() + try: + yield server, port + finally: + http_server.shutdown() + server._agent.on_shutdown(timeout=5) + + +@pytest.mark.python +@pytest.mark.skipif(not language_server_tests_enabled(LanguageServerId.PYTHON), reason="python tests are disabled in this environment") +def test_facade_method_results_are_transferred_from_the_project_server(project_server: tuple[ProjectServer, int]) -> None: + server, port = project_server + client = ProjectServerClient(server.get_serena_config(), port=port) + result = client.call_facade_method("test_repo_python", "lsp", "find_symbol", ["create_user"], {"include_body": True}) + + # the result is a self-contained object which can be processed and rendered locally + assert isinstance(result, LspSymbolCollection) + assert [s.name for s in result.symbols] == ["create_user"] + assert result.symbols[0].body.startswith("def create_user") + assert "create_user" in result.represent() + + +@pytest.mark.python +@pytest.mark.skipif(not language_server_tests_enabled(LanguageServerId.PYTHON), reason="python tests are disabled in this environment") +def test_external_project_context_in_repl( + project_server: tuple[ProjectServer, int], + serena_config: SerenaConfig, # noqa: F811 + monkeypatch: pytest.MonkeyPatch, +) -> None: + server, port = project_server + monkeypatch.setattr(ProjectServer, "PORT", port) # let the REPL's external project context use the test server + + # enable the optional "ext" facade + serena_config.included_apis = ["ext"] + + # the querying agent has another project active and queries the python test project + serena_config.auth_secret = server.get_auth_secret() + agent = SerenaAgent(project="test_repo_typescript", serena_config=serena_config) + agent.execute_task(lambda: None) + try: + tool = agent.get_tool(SerenaReplTool) + session_id = agent.create_session().session_id + code = ( + 'with s.ext.read_project_context("test_repo_python"):\n' + ' result = s.lsp.find_symbol("create_user")\n' + "[s.name for s in result.symbols]" + ) + assert tool.apply(session_id, code) == "create_user" + assert server._loaded_projects_by_root, "the operation was not executed by the project server" + + # the result persists and can be used after the context; the active project is restored + assert "services.py" in tool.apply(session_id, "result.represent()") + assert "test_repo_typescript" in agent.get_current_config_overview() + finally: + agent.on_shutdown(timeout=5) diff --git a/test/serena/test_file_tools.py b/test/serena/test_file_tools.py index 24c1649d..2704c1ae 100644 --- a/test/serena/test_file_tools.py +++ b/test/serena/test_file_tools.py @@ -15,10 +15,8 @@ def read_file_tool(tmp_path: Path) -> ReadFileTool: project = Project.load(str(tmp_path), serena_config=SerenaConfig(gui_log_window=False, web_dashboard=False)) agent = MagicMock() agent.get_active_project_or_raise.return_value = project - tool = ReadFileTool(agent) - # bypass the length limit, which would otherwise depend on the agent configuration - tool._limit_length = lambda result, max_answer_chars: result - return tool + agent.serena_config.default_max_tool_answer_chars = 10000 + return ReadFileTool(agent) def _deleted_by_delete_lines(content: str, line: int) -> str: diff --git a/test/serena/test_fs_api.py b/test/serena/test_fs_api.py new file mode 100644 index 00000000..65f29b3b --- /dev/null +++ b/test/serena/test_fs_api.py @@ -0,0 +1,81 @@ +""" +Tests for the file system facade API. +""" + +from pathlib import Path +from unittest.mock import MagicMock + +import pytest + +from serena.config.serena_config import SerenaConfig +from serena.project import Project +from serena.repl.api.fs_api import FsApi +from serena.repl.facade import ApiScope, Facade + + +@pytest.fixture +def project(tmp_path: Path) -> Project: + (tmp_path / "src").mkdir() + (tmp_path / "src" / "a.py").write_text("x = foo(1)\ny = foo(2)\nz = 3\n", encoding="utf-8") + (tmp_path / "src" / "b.txt").write_text("foo in text\n", encoding="utf-8") + (tmp_path / "README.md").write_text("# readme\n", encoding="utf-8") + return Project.load(str(tmp_path), serena_config=SerenaConfig(gui_log_window=False, web_dashboard=False)) + + +@pytest.fixture +def api(project: Project) -> FsApi: + agent = MagicMock() + agent.get_active_project_or_raise.return_value = project + agent.serena_config.default_max_tool_answer_chars = 10000 + return FsApi(agent) + + +def test_facade_exposes_file_operations(api: FsApi) -> None: + facade = Facade.from_api(api, ApiScope()) + assert facade.name == "fs" + assert set(facade.enabled_method_names) == {"read_file", "create_text_file", "list_dir", "find_file", "search_for_pattern"} + assert {name for name in facade.enabled_method_names if facade.get_method(name).info.can_edit} == {"create_text_file"} + + +def test_read_file(api: FsApi) -> None: + content = api.read_file("src/a.py") + assert content.lines == ["x = foo(1)", "y = foo(2)", "z = 3", ""] + assert content.represent() == content.text + + assert api.read_file("src/a.py", start_line=1, end_line=1).text == "y = foo(2)" + assert api.read_file("src/a.py", start_line=-2).lines == ["z = 3", ""] + + +def test_create_text_file(api: FsApi, project: Project) -> None: + result = api.create_text_file("sub/new.txt", "hello\n") + assert "new.txt" in result + assert (Path(project.project_root) / "sub" / "new.txt").read_text(encoding="utf-8") == "hello\n" + + result = api.create_text_file("sub/new.txt", "changed\n") + assert "Overwrote" in result + + with pytest.raises(AssertionError): + api.create_text_file("../outside.txt", "nope") + + +def test_list_dir_and_find_file(api: FsApi) -> None: + listing = api.list_dir(".", recursive=True) + assert "src" in listing.dirs + assert {"src/a.py", "src/b.txt", "README.md"} <= {f.replace("\\", "/") for f in listing.files} + assert '"dirs"' in listing.represent() and '"files"' in listing.represent() + + with pytest.raises(FileNotFoundError): + api.list_dir("missing", recursive=False) + + assert [f.replace("\\", "/") for f in api.find_file("*.py", ".")] == ["src/a.py"] + + +def test_search_for_pattern(api: FsApi) -> None: + matches = api.search_for_pattern("foo", relative_path="src") + assert len(matches) == 3 + assert {m.source_file_path.replace("\\", "/") for m in matches.matches} == {"src/a.py", "src/b.txt"} # type: ignore + + # restricting to code files excludes the text file; the rendering maps files to matched lines + code_matches = api.search_for_pattern("foo", restrict_search_to_code_files=True) + assert all(m.source_file_path.endswith("a.py") for m in code_matches.matches) # type: ignore + assert "foo(1)" in code_matches.represent() diff --git a/test/serena/test_jetbrains_api.py b/test/serena/test_jetbrains_api.py new file mode 100644 index 00000000..884c0d45 --- /dev/null +++ b/test/serena/test_jetbrains_api.py @@ -0,0 +1,79 @@ +""" +Tests for the JetBrains facade API, using a mocked plugin client (no IDE required). +""" + +from unittest.mock import MagicMock, patch + +import pytest + +from serena.repl.api.jb_api import JetBrainsApi +from serena.repl.facade import ApiScope, Facade + + +@pytest.fixture +def agent() -> MagicMock: + agent = MagicMock() + agent.serena_config.default_max_tool_answer_chars = 10000 + return agent + + +@pytest.fixture +def client() -> MagicMock: + client = MagicMock() + with patch("serena.repl.api.jb_api.JetBrainsPluginClient.from_project") as from_project: + from_project.return_value.__enter__.return_value = client + yield client + + +def test_facade_exposes_all_jetbrains_operations(agent: MagicMock) -> None: + facade = Facade.from_api(JetBrainsApi(agent), ApiScope()) + assert facade.name == "jb" + assert set(facade.enabled_method_names) == { + "find_symbol", + "find_referencing_symbols", + "get_symbols_overview", + "get_type_hierarchy", + "find_declaration", + "find_implementations", + "rename", + "move", + "safe_delete", + "inline_symbol", + "run_inspections", + "list_inspections", + "debug_eval", + "debug_eval_info", + } + + +def test_find_symbol_renders_grouped_and_exposes_symbols(agent: MagicMock, client: MagicMock) -> None: + symbols = [ + {"name_path": "Foo", "type": "class", "relative_path": "a.py", "quick_info": "class Foo"}, + {"name_path": "Bar/foo", "type": "method", "relative_path": "b.py", "quick_info": "def foo()"}, + ] + client.find_symbol.return_value = {"symbols": symbols} + + result = JetBrainsApi(agent).find_symbol("foo") + + # the underlying symbols are accessible from code + assert [s["name_path"] for s in result.symbols] == ["Foo", "Bar/foo"] + # and the rendering contains them, grouped by file + rendered = result.represent() + assert '"a.py"' in rendered and '"b.py"' in rendered + assert "class Foo" in rendered + + +def test_find_symbol_rejects_too_many_matches_with_identifiers(agent: MagicMock, client: MagicMock) -> None: + client.find_symbol.return_value = { + "symbols": [{"name_path": f"foo{i}", "type": "function", "relative_path": "a.py", "body": "..."} for i in range(3)] + } + with pytest.raises(ValueError, match="Matched 3>max_matches=1") as exc_info: + JetBrainsApi(agent).find_symbol("foo*", max_matches=1) + assert "foo2" in str(exc_info.value) + assert '"body"' not in str(exc_info.value) # only identifiers, no content + + +def test_find_symbol_rejects_wildcard_only_pattern(agent: MagicMock, client: MagicMock) -> None: + with pytest.raises(ValueError, match="get_symbols_overview"): + JetBrainsApi(agent).find_symbol("*") + client.find_symbol.assert_not_called() diff --git a/test/serena/test_ls_file_sync.py b/test/serena/test_ls_file_sync.py index 618559b3..381e31de 100644 --- a/test/serena/test_ls_file_sync.py +++ b/test/serena/test_ls_file_sync.py @@ -62,7 +62,8 @@ class FileSystemSyncTestCase: symbol_names = [ref["name_path"].split("/")[-1] for ref in ref_symbols] return symbol_names else: - ls = next(iter(agent.get_active_project_or_raise().language_server_manager.iter_language_servers())) + ls_manager = agent.get_active_project_or_raise().get_language_server_manager_or_raise() + ls = next(iter(ls_manager.iter_language_servers())) document_symbols = ls.request_document_symbols(self._TARGET_FILE).get_all_symbols_and_roots() target = next((s for s in document_symbols[0] if s.get("name") == self._TARGET_SYMBOL), None) assert target is not None and "selectionRange" in target, f"{self._TARGET_SYMBOL} not found in {self._TARGET_FILE}" @@ -162,7 +163,9 @@ class SymbolPositionStaleAfterExternalEditTestCase: with agent_for_project_context(LanguageServerId.PYTHON, str(repo_root)) as agent: project = agent.get_active_project_or_raise() - ls = next(iter(project.language_server_manager.iter_language_servers())) + ls_manager = project.language_server_manager + assert ls_manager is not None + ls = next(iter(ls_manager.iter_language_servers())) tool = agent.get_tool(FindSymbolTool) # Hold the file's buffer open across the external edit, mirroring the state left diff --git a/test/serena/test_mcp.py b/test/serena/test_mcp.py index db1923fd..c37e9fc9 100644 --- a/test/serena/test_mcp.py +++ b/test/serena/test_mcp.py @@ -1,13 +1,14 @@ """Tests for the mcp.py module in serena.""" import pytest -from mcp.server.fastmcp.tools.base import Tool as MCPTool +from mcp.server.mcpserver import Context +from mcp.server.mcpserver.tools.base import Tool as MCPTool -from serena import __version__ from serena.agent import Tool, ToolRegistry from serena.config.context_mode import SerenaAgentContext -from serena.config.serena_config import SerenaConfig from serena.mcp import SerenaMCPFactory +from serena.repl.facade import ApiScope +from serena.repl.repl import SerenaRepl make_tool = SerenaMCPFactory.make_mcp_tool @@ -22,6 +23,14 @@ class MockAgent: def get_context() -> SerenaAgentContext: return SerenaAgentContext.load_default() + @staticmethod + def get_repl() -> SerenaRepl: + return SerenaRepl([], ApiScope()) + + @staticmethod + def is_single_project() -> bool: + return False + class BaseMockTool(Tool): """A mock Tool class for testing.""" @@ -46,30 +55,13 @@ class BasicTool(BaseMockTool): self, log_call: bool = True, catch_exceptions: bool = True, + mcp_ctx: Context | None = None, **kwargs, ) -> str: """Mock implementation of apply_ex.""" return self.apply(**kwargs) -def test_create_mcp_server_reports_serena_version(monkeypatch: pytest.MonkeyPatch) -> None: - """MCP initialize must report Serena's version, not the installed mcp SDK version.""" - - class MinimalAgent: - def create_connection_prompt(self) -> str: - return "" - - monkeypatch.setattr(SerenaConfig, "from_config_file", classmethod(lambda cls: SerenaConfig())) - factory = SerenaMCPFactory(transport="stdio") - monkeypatch.setattr(factory, "_create_serena_agent", lambda *args, **kwargs: MinimalAgent()) - - mcp_server = factory.create_mcp_server() - initialization_options = mcp_server._mcp_server.create_initialization_options() - - assert initialization_options.server_name == "Serena" - assert initialization_options.server_version == __version__ - - def test_make_tool_basic() -> None: """Test that make_tool correctly creates an MCP tool from a Tool object.""" mock_tool = BasicTool() diff --git a/test/serena/test_mem_api.py b/test/serena/test_mem_api.py new file mode 100644 index 00000000..f3f133f2 --- /dev/null +++ b/test/serena/test_mem_api.py @@ -0,0 +1,68 @@ +""" +Tests for the memory facade API. +""" + +from pathlib import Path +from unittest.mock import MagicMock + +import pytest + +from serena.config.serena_config import SerenaConfig +from serena.project import Project +from serena.repl.api.mem_api import MemoryApi +from serena.repl.facade import ApiScope, Facade + + +@pytest.fixture +def api(tmp_path: Path) -> MemoryApi: + project = Project.load(str(tmp_path), serena_config=SerenaConfig(gui_log_window=False, web_dashboard=False)) + agent = MagicMock() + agent.get_active_project_or_raise.return_value = project + agent.serena_config.default_max_tool_answer_chars = 10000 + return MemoryApi(agent) + + +def test_facade_exposes_memory_operations(api: MemoryApi) -> None: + facade = Facade.from_api(api, ApiScope()) + assert facade.name == "mem" + assert set(facade.enabled_method_names) == { + "list_memories", + "read_memory", + "write_memory", + "edit_memory", + "rename_memory", + "delete_memory", + "onboarding", + } + assert {name for name in facade.enabled_method_names if facade.get_method(name).info.can_edit} == { + "write_memory", + "edit_memory", + "rename_memory", + "delete_memory", + } + + +def test_memory_lifecycle(api: MemoryApi) -> None: + api.write_memory("topic/first", "# First\nhello") + api.write_memory("second", "see `mem:topic/first`") + + # listing exposes the names to code and renders as JSON (global memories of the machine may be present, too) + memory_list = api.list_memories() + assert {"second", "topic/first"} <= set(memory_list.memories) + assert api.list_memories("topic").memories == ["topic/first"] + assert '"memories"' in memory_list.represent() + + # reading, editing, renaming (with reference propagation) and deleting + assert api.read_memory("topic/first") == "# First\nhello" + api.edit_memory("topic/first", "hello", "world", mode="literal") + assert api.read_memory("topic/first") == "# First\nworld" + api.rename_memory("topic/first", "topic/renamed") + assert "mem:topic/renamed" in api.read_memory("second") + api.delete_memory("second") + assert "second" not in api.list_memories().memories + assert api.list_memories("topic").memories == ["topic/renamed"] + + +def test_write_memory_rejects_overlong_content(api: MemoryApi) -> None: + with pytest.raises(ValueError, match="too long"): + api.write_memory("big", "x" * 100, max_chars=10) diff --git a/test/serena/test_memories_manager.py b/test/serena/test_memories_manager.py index 46c1b47b..513f9057 100644 --- a/test/serena/test_memories_manager.py +++ b/test/serena/test_memories_manager.py @@ -844,3 +844,41 @@ class TestAutoPrefixBareReferences: # idempotent: the second run should not touch anything assert second.total_replacements == 0 assert fs_manager.load_memory("docs") == "the mem:auth/login process" + + +class TestRenameMemorySparesReadOnlyMemories: + """Regression: a tool-context rename enumerated read-only memories, so propagating the + reference into one raised ``PermissionError`` after the rename itself had already been applied. + """ + + @staticmethod + def _manager(tmp_path, monkeypatch) -> MemoryManager: + manager = MemoryManager(serena_data_folder=tmp_path, read_only_memory_patterns=[r"frozen/.*"]) + # the global memories of the machine would otherwise join the enumeration as well + global_dir = tmp_path / "global" + global_dir.mkdir() + monkeypatch.setattr(manager, "_global_memory_dir", global_dir) + _write(manager, "auth/login", "# login notes") + _write(manager, "frozen/notes", "see `mem:auth/login`") + _write(manager, "docs", "first `mem:auth/login`, then `mem:auth/login`") + return manager + + def test_tool_context_rename_completes_and_leaves_read_only_reference_alone(self, tmp_path, monkeypatch) -> None: + manager = self._manager(tmp_path, monkeypatch) + + message, n_updated = manager.rename_memory_and_propagate_references("auth/login", "auth/signin", is_tool_context=True) + + assert "auth/signin" in message + assert manager.load_memory("auth/signin") == "# login notes" + assert manager.load_memory("docs") == "first `mem:auth/signin`, then `mem:auth/signin`" + assert manager.load_memory("frozen/notes") == "see `mem:auth/login`" + assert n_updated == 2 + + def test_cli_context_rename_still_propagates_into_read_only_memories(self, tmp_path, monkeypatch) -> None: + manager = self._manager(tmp_path, monkeypatch) + + _, n_updated = manager.rename_memory_and_propagate_references("auth/login", "auth/signin", is_tool_context=False) + + assert manager.load_memory("frozen/notes") == "see `mem:auth/signin`" + assert manager.load_memory("docs") == "first `mem:auth/signin`, then `mem:auth/signin`" + assert n_updated == 3 diff --git a/test/serena/test_project_server.py b/test/serena/test_project_server.py index 1288e177..845603b1 100644 --- a/test/serena/test_project_server.py +++ b/test/serena/test_project_server.py @@ -8,8 +8,11 @@ from typing import Any, cast from unittest.mock import MagicMock import pytest +from flask import Flask +from werkzeug.serving import make_server -from serena.project_server import ProjectServer, QueryProjectRequest +from serena.config.serena_config import SerenaConfig +from serena.project_server import ProjectServer, ProjectServerClient, QueryProjectRequest @pytest.fixture @@ -23,6 +26,53 @@ def project_server() -> ProjectServer: return server +@pytest.fixture +def authenticated_server(project_server: ProjectServer, monkeypatch: pytest.MonkeyPatch) -> ProjectServer: + # expose the real HTTP routes with a query handler that needs no language servers + project_server._agent.serena_config.auth_secret = "test-shared-secret" + project_server._app = Flask(__name__) + monkeypatch.setattr(project_server, "_query_project", lambda req: req.project_name) + project_server._setup_routes() + return project_server + + +@pytest.mark.parametrize("authorization", [None, "Bearer wrong-secret", "test-shared-secret", "Bearer café"]) +@pytest.mark.parametrize("path", ["/heartbeat", "/query_project"]) +def test_project_server_rejects_invalid_credentials(authenticated_server: ProjectServer, authorization: str | None, path: str) -> None: + # unauthorized requests are rejected even before query payload validation + headers = {} if authorization is None else {"Authorization": authorization} + with authenticated_server._app.test_client() as client: + response = client.open(path, method="GET" if path == "/heartbeat" else "POST", headers=headers) + assert response.status_code == 401 + + +@pytest.mark.parametrize("use_wrong_password", [True, False]) +def test_project_server_client_authenticates_requests(authenticated_server: ProjectServer, use_wrong_password: bool) -> None: + # run the authenticated endpoints on an ephemeral local port + http_server = make_server("127.0.0.1", 0, authenticated_server._app) + thread = threading.Thread(target=http_server.serve_forever, daemon=True) + thread.start() + try: + + def check_client(): + serena_config = SerenaConfig() + serena_config.auth_secret = "wrong-secret" if use_wrong_password else authenticated_server.get_auth_secret() + client = ProjectServerClient(serena_config, port=http_server.server_port) + assert client.query_project("other", "find_symbol", "{}") == "other" + + # construction authenticates the heartbeat, raising a Connection error if using the wrong password + if use_wrong_password: + with pytest.raises(expected_exception=ConnectionError, match="401"): + check_client() + else: + check_client() + + finally: + http_server.shutdown() + thread.join(timeout=5) + http_server.server_close() + + def test_cached_project_lookup_is_not_blocked_by_unrelated_cold_load(project_server: ProjectServer) -> None: cached_root = Path("/cached") cold_root = Path("/cold") diff --git a/test/serena/test_repl_tool.py b/test/serena/test_repl_tool.py new file mode 100644 index 00000000..9cf5dd38 --- /dev/null +++ b/test/serena/test_repl_tool.py @@ -0,0 +1,382 @@ +""" +Tests for the REPL tool, which executes Python code against the facade entrypoint `s`. +""" + +import os +from unittest.mock import MagicMock + +import pytest + +from serena.config.serena_config import ApiInclusionDefinition +from serena.language_backend import BuiltinLanguageBackend +from serena.repl.api.edit_api import EditApi +from serena.repl.api.lsp_api import LspApi +from serena.repl.external_project import ExternalProjectExecution +from serena.repl.facade import ApiScope, Facade, FacadeApi, FacadeMethodInfo, facade_method +from serena.repl.repl import SerenaRepl +from serena.session import SerenaSession +from serena.tools import FindSymbolTool, SerenaReplTool +from solidlsp.ls_config import LanguageServerId +from test.conftest import agent_for_project_context + + +class TestReplExecution: + """Tests the code execution mechanics of the REPL, which do not require a project.""" + + @pytest.fixture + def repl(self) -> SerenaRepl: + return SerenaRepl([Facade.from_api(LspApi(MagicMock()), ApiScope())], ApiScope()) + + def test_last_expression_defines_result(self, repl: SerenaRepl) -> None: + assert repl.execute("x = 20\ny = 22\nx + y") == "42" + assert repl.execute("x = 20\ny = 22") == "None" # no trailing expression + assert repl.execute("") == "None" + + def test_return_yields_a_hint(self, repl: SerenaRepl) -> None: + result = repl.execute("x = 1\nreturn x") + assert result.startswith("SyntaxError") and "line 2" in result + assert "last expression is the result" in result + + def test_single_expression_is_evaluated(self, repl: SerenaRepl) -> None: + assert repl.execute("1 + 2") == "3" + + def test_list_is_rendered_element_wise(self, repl: SerenaRepl) -> None: + assert repl.execute('["a", "b"]') == "a\nb" + + def test_error_reports_type_message_and_line(self, repl: SerenaRepl) -> None: + result = repl.execute("x = 1\nraise ValueError('boom')") + assert result.startswith("ValueError: boom") + assert "line 2" in result + + def test_syntax_error_reports_line(self, repl: SerenaRepl) -> None: + result = repl.execute("x = 1\ny = (2") + assert result.startswith("SyntaxError") + assert "line 2" in result + + def test_top_level_names_persist_within_session(self, repl: SerenaRepl) -> None: + session = SerenaSession("a") + repl.execute("x = 20\ndef double(v):\n return 2 * v\nfor i in range(3):\n pass\nimport os as os_module", session) + assert repl.execute("double(x) + i", session) == "42" + assert repl.execute("os_module.sep is not None", session) == "True" + + # local variables of nested scopes do not persist; persisted items can be listed and cleared + assert "v" not in repl.execute("s.vars()", session) + listing = repl.execute("s.vars()", session) + assert "x: int = 20" in listing and "double: function" in listing + assert repl.execute("s.clear()", session) == "Removed 4 persisted item(s)." + assert repl.execute("s.vars()", session) == "No persisted variables." + assert "NameError" in repl.execute("x", session) + + def test_sessions_have_separate_namespaces(self, repl: SerenaRepl) -> None: + session_a = SerenaSession("a") + repl.execute("x = 1", session_a) + assert "NameError" in repl.execute("x", SerenaSession("b")) + assert repl.execute("x", session_a) == "1" + # ... and execution without a session persists nothing + repl.execute("y = 1") + assert "NameError" in repl.execute("y") + + def test_persisted_functions_use_the_current_entrypoint(self) -> None: + # a function defined against one REPL instance uses the entrypoint of the REPL that later calls it + session = SerenaSession("a") + SerenaRepl([Facade.from_api(LspApi(MagicMock()), ApiScope())], ApiScope()).execute("def facades():\n return s.info()", session) + rebuilt_repl = SerenaRepl([Facade.from_api(EditApi(MagicMock()), ApiScope())], ApiScope()) + overview = rebuilt_repl.execute("facades()", session) + assert "s.edit" in overview and "s.lsp" not in overview + + @pytest.mark.parametrize("builtin_backend", [BuiltinLanguageBackend.LSP, BuiltinLanguageBackend.JETBRAINS]) + @pytest.mark.parametrize("read_only", [True, False]) + def test_external_project_dispatch(self, builtin_backend: BuiltinLanguageBackend, read_only: bool) -> None: + agent = MagicMock() + backend = builtin_backend.get_instance() + agent.get_language_backend.return_value = backend + + class FakeExternalProject(ExternalProjectExecution): + def __init__(self) -> None: + super().__init__("other", read_only=read_only, agent=agent) + self.calls: list[tuple[str, str, tuple, dict]] = [] + + def call_remotely(self, facade_name: str, method_name: str, args: tuple, kwargs: dict) -> str: + self.calls.append((facade_name, method_name, args, kwargs)) + return "remote result" + + class LocalApi(FacadeApi): + @facade_method(can_edit=True) + def write(self, content: str) -> str: + return f"local result: {content}" + + # expose a server-backed read and a backend-dependent write + facades = [Facade.from_api(LspApi(agent), ApiScope()), Facade.from_api(LocalApi(agent, "local", "local operations"), ApiScope())] + repl = SerenaRepl(facades, ApiScope()) + external_project = FakeExternalProject() + repl.entrypoint.set_external_project_(external_project) + + # methods explicitly requiring the project server are executed remotely + assert repl.execute('s.lsp.find_symbol("Foo", depth=1)') == "remote result" + assert external_project.calls == [("lsp", "find_symbol", ("Foo",), {"depth": 1})] + + # writes obey the context's access mode and use the selected backend + result = repl.execute('s.local.write("content")') + if read_only: + assert "PermissionError" in result and "read-only" in result + assert len(external_project.calls) == 1 + elif backend.is_lsp(): + assert result == "remote result" + assert external_project.calls[-1] == ("local", "write", ("content",), {}) + else: + assert result == "local result: content" + assert len(external_project.calls) == 1 + + # leaving the external execution context restores local execution + repl.entrypoint.set_external_project_(None) + assert repl.execute('s.local.write("restored")') == "local result: restored" + + def test_facade_discovery(self, repl: SerenaRepl) -> None: + overview = repl.execute("s.info()") + assert "s.lsp" in overview + assert "find_symbol" in overview # method names are listed, but not signatures + assert "name_path_pattern" not in overview + facade_info = repl.execute('s.info("lsp")') + assert "find_symbol(" in facade_info + method_info = repl.execute('s.info("lsp.find_symbol")') + assert "name_path_pattern" in method_info + + def test_type_discovery(self, repl: SerenaRepl) -> None: + # the overview names the result types of methods returning objects that can be processed in code + assert "find_symbol -> LspSymbolCollection" in repl.execute("s.info()") + + # signatures render type names without module paths, and point to the documentation of referenced return types + method_info = repl.execute('s.info("lsp.find_symbol")') + assert "-> LspSymbolCollection" in method_info and "lsp_api." not in method_info + assert 's.info("LspSymbolCollection")' in method_info + + # the facade description documents the operations only and lists the result types by name + facade_info = repl.execute('s.info("lsp")') + assert "type LspSymbolCollection" not in facade_info + assert "Result types: " in facade_info and "LspSymbolCollection" in facade_info + + # types can be requested via the facade or by bare name, and their curated members are documented + type_info = repl.execute('s.info("lsp.LanguageServerSymbol")') + assert type_info == repl.execute('s.info("LanguageServerSymbol")') + assert "get_name_path() -> str" in type_info and "iter_children()" in type_info + assert "to_dict" not in type_info # not among the curated members + + # types reachable through annotations are documented without being declared: SymbolKind (a parameter type of + # iter_ancestors) is documented as an enum with its members, both transitively and on request + assert "enum SymbolKind" in type_info and "SymbolKind.Class = 5" in type_info + assert "enum SymbolKind" in repl.execute('s.info("SymbolKind")') + + # TypedDicts reachable through annotations are documented with their keys + assert "keys:" in repl.execute('s.info("Diagnostic")') and "severity" in repl.execute('s.info("Diagnostic")') + assert "represent()" not in repl.execute('s.info("LspSymbolCollection")') # the representation mechanism is not exposed + + def test_contained_types_are_documented_once_per_session(self, repl: SerenaRepl) -> None: + session = SerenaSession("test") + + # a type's documentation includes the types it contains (transitively) + first = repl.execute('s.info("LspReferenceCollection")', session) + assert "type LspReferenceCollection" in first + assert "type ReferenceInLanguageServerSymbol" in first and "type LanguageServerSymbol" in first + + # a contained type documented earlier in the session is only pointed to; an explicit request yields it again + second = repl.execute('s.info("LspSymbolCollection")', session) + assert "type LspSymbolCollection" in second + assert "type LanguageServerSymbol: documented earlier" in second and "get_name_path()" not in second + assert "get_name_path()" in repl.execute('s.info("LanguageServerSymbol")', session) + + # another session is unaffected + assert "get_name_path()" in repl.execute('s.info("LspSymbolCollection")', SerenaSession("other")) + + def test_info_documents_several_items(self, repl: SerenaRepl) -> None: + info = repl.execute('s.info("lsp.find_symbol", "nope", "lsp.LspSymbolCollection")') + assert "lsp.find_symbol(" in info and "type LspSymbolCollection" in info + assert "Unknown item 'nope'" in info # an unknown item does not prevent the documentation of the others + + +class TestFacade: + """Tests the indirection between facades and their implementations.""" + + class DummyApi(FacadeApi): + def __init__(self, agent: MagicMock) -> None: + super().__init__(agent, name="dummy", description="a dummy facade") + + @facade_method() + def add(self, a: int, b: int) -> int: + """Adds two numbers.""" + return a + b + + @facade_method(can_edit=True) + def secret(self) -> str: + return "hidden" + + @facade_method(optional=True, beta=True) + def extra(self) -> str: + return "extra" + + @facade_method(niche=True) + def rarely(self, x: int) -> str: + """Rarely needed operation. + + :param x: some parameter + """ + return str(x) + + def undecorated(self) -> str: + """Public within Serena, but not exposed, since it is not decorated.""" + return "internal" + + def _internal(self) -> None: + pass + + @staticmethod + def _scope(*definitions: ApiInclusionDefinition, **kwargs: list[str]) -> ApiScope: + """ + :param definitions: definitions to apply in order + :param kwargs: an additional definition (`included_apis`/`excluded_apis`) to apply last + """ + scope = ApiScope() + for definition in definitions: + scope.process(definition) + if kwargs: + scope.process(ApiInclusionDefinition(**kwargs)) + return scope + + def test_optional_facade_is_opt_in(self) -> None: + # an optional facade is disabled unless it is included explicitly + facade = Facade.from_api(self.DummyApi(MagicMock()), ApiScope(), is_optional=True) + assert not facade.is_enabled() + assert facade.enabled_method_names == [] + + # including the facade enables its non-optional methods + facade = Facade.from_api(self.DummyApi(MagicMock()), self._scope(included_apis=["dummy"]), is_optional=True) + assert facade.is_enabled() + assert set(facade.enabled_method_names) == {"add", "secret", "rarely"} + + # including a single method enables the facade with just that method + facade = Facade.from_api(self.DummyApi(MagicMock()), self._scope(included_apis=["dummy.add"]), is_optional=True) + assert facade.is_enabled() + assert facade.enabled_method_names == ["add"] + + def test_api_scope_facade_exclusion_and_method_inclusion(self) -> None: + # excluding the facade disables everything + facade = Facade.from_api(self.DummyApi(MagicMock()), self._scope(excluded_apis=["dummy"])) + assert facade.enabled_method_names == [] + + # an excluded facade is opt-in: a method inclusion enables exactly that method + facade = Facade.from_api(self.DummyApi(MagicMock()), self._scope(excluded_apis=["dummy"], included_apis=["dummy.add"])) + assert facade.enabled_method_names == ["add"] + + def test_entrypoint_omits_excluded_facades(self) -> None: + def create_repl(scope: ApiScope) -> SerenaRepl: + return SerenaRepl([Facade.from_api(self.DummyApi(MagicMock()), scope)], scope) + + assert "s.dummy" in create_repl(ApiScope()).execute("s.info()") + assert "s.dummy" not in create_repl(self._scope(excluded_apis=["dummy"])).execute("s.info()") + # a method inclusion keeps the facade available (with just that method) + overview = create_repl(self._scope(excluded_apis=["dummy"], included_apis=["dummy.add"])).execute("s.info()") + assert "s.dummy" in overview and "methods: add" in overview + + def test_api_scope_later_definitions_take_precedence(self) -> None: + scope = self._scope( + ApiInclusionDefinition(included_apis=["dummy.extra"]), + ApiInclusionDefinition(excluded_apis=["dummy.extra", "dummy.add"]), + ApiInclusionDefinition(included_apis=["dummy.add"]), + ) + facade = Facade.from_api(self.DummyApi(MagicMock()), scope) + assert set(facade.enabled_method_names) == {"add", "secret", "rarely"} + + def test_api_scope_read_only_excludes_editing_methods(self) -> None: + scope = self._scope(included_apis=["dummy.secret"]) + scope.exclude_editing() + facade = Facade.from_api(self.DummyApi(MagicMock()), scope) + assert set(facade.enabled_method_names) == {"add", "rarely"} + + def test_enabled_methods_delegate_to_implementation(self) -> None: + facade = Facade.from_api(self.DummyApi(MagicMock()), ApiScope()) + assert facade.add(1, 2) == 3 + assert "dummy.add(a: int, b: int) -> int" in facade.describe() + assert "Adds two numbers." in facade.describe_member("add") + + def test_disabled_methods_are_inaccessible_and_undocumented(self) -> None: + facade = Facade.from_api(self.DummyApi(MagicMock()), self._scope(excluded_apis=["dummy.secret"])) + assert facade.add(1, 2) == 3 + with pytest.raises(AttributeError): + facade.secret() + with pytest.raises(ValueError): + facade.describe_member("secret") + assert "secret" not in facade.describe() + assert "_internal" not in facade.describe() + + def test_undecorated_methods_are_not_exposed(self) -> None: + facade = Facade.from_api(self.DummyApi(MagicMock()), ApiScope()) + assert self.DummyApi(MagicMock()).undecorated() == "internal" # usable from within Serena + with pytest.raises(AttributeError): + facade.undecorated() + with pytest.raises(ValueError): + facade.get_method("undecorated") + assert "undecorated" not in facade.describe() + assert "undecorated" not in facade.enabled_method_names + + def test_niche_methods_are_summarised_in_facade_description(self) -> None: + facade = Facade.from_api(self.DummyApi(MagicMock()), ApiScope()) + description = facade.describe() + assert "dummy.rarely: Rarely needed operation." in description + assert ":param x:" not in description # only the summary, no signature or full documentation + assert ":param x:" in facade.describe_member("rarely") # full documentation on request + + def test_optional_methods_are_disabled_by_default(self) -> None: + facade = Facade.from_api(self.DummyApi(MagicMock()), ApiScope()) + assert "extra" not in facade.enabled_method_names + with pytest.raises(AttributeError): + facade.extra() + facade.get_method("extra").enabled = True + assert facade.extra() == "extra" + + # an explicit inclusion enables it + facade = Facade.from_api(self.DummyApi(MagicMock()), self._scope(included_apis=["dummy.extra"])) + assert "extra" in facade.enabled_method_names + + def test_corresponding_tool(self) -> None: + facade = Facade.from_api(self.DummyApi(MagicMock()), ApiScope()) + assert facade.get_method("add").info.get_corresponding_tool_name() is None + + lsp_facade = Facade.from_api(LspApi(MagicMock()), ApiScope()) + info = lsp_facade.get_method("find_symbol").info + assert info.corresponding_tool is FindSymbolTool + assert info.get_corresponding_tool_name() == "find_symbol" + + def test_method_info_mirrors_decorator(self) -> None: + facade = Facade.from_api(self.DummyApi(MagicMock()), ApiScope()) + assert facade.get_method("add").info == FacadeMethodInfo(name="add") + assert facade.get_method("secret").info.can_edit + assert facade.get_method("extra").info == FacadeMethodInfo(name="extra", optional=True, beta=True) + + def test_enablement_can_be_changed(self) -> None: + facade = Facade.from_api(self.DummyApi(MagicMock()), ApiScope()) + facade.get_method("secret").enabled = False + with pytest.raises(AttributeError): + facade.secret() + facade.get_method("secret").enabled = True + assert facade.secret() == "hidden" + + +@pytest.mark.python +class TestLspFacade: + _SERVICES_FILE = os.path.join("test_repo", "services.py") + + def test_find_symbol_via_repl(self) -> None: + with agent_for_project_context(LanguageServerId.PYTHON) as agent: + tool = agent.get_tool(SerenaReplTool) + session_id = agent.create_session().session_id + + # a returned collection is rendered, identifying the symbol and its file + rendered = tool.apply(session_id, 's.lsp.find_symbol("create_user")') + assert "create_user" in rendered + assert "services.py" in rendered + + # the underlying symbols are accessible from code, e.g. to retrieve a body without rendering the collection + body = tool.apply( + session_id, + f'result = s.lsp.find_symbol("create_user", relative_path={self._SERVICES_FILE!r})\nresult.symbols[0].body', + ) + assert body.startswith("def create_user") diff --git a/test/serena/test_serena_agent.py b/test/serena/test_serena_agent.py index cd7cbb6d..448206d2 100644 --- a/test/serena/test_serena_agent.py +++ b/test/serena/test_serena_agent.py @@ -14,12 +14,13 @@ from _pytest.mark import Mark, MarkDecorator, ParameterSet from serena.agent import SerenaAgent from serena.config.context_mode import SerenaAgentContext -from serena.config.serena_config import ProjectConfig, RegisteredProject, SerenaConfig +from serena.config.serena_config import AgentInterface, ProjectConfig, RegisteredProject, SerenaConfig +from serena.lsp.lsp_diagnostics import DiagnosticsContext from serena.project import Project +from serena.session import SessionRegistry from serena.tools import ( SUCCESS_RESULT, ActivateProjectTool, - EditingToolWithDiagnostics, FindDeclarationTool, FindImplementationsTool, FindReferencingSymbolsTool, @@ -30,6 +31,7 @@ from serena.tools import ( ReplaceInFilesTool, ReplaceSymbolBodyTool, SafeDeleteSymbol, + SerenaReplTool, Tool, ) from solidlsp.ls_config import LanguageServerId @@ -824,15 +826,15 @@ def read_project_file(project: Project, relative_path: str) -> str: def parse_edit_diagnostics_result(result: str) -> dict: """Utility function to parse the diagnostic payload returned by edit tools.""" - assert EditingToolWithDiagnostics.DIAGNOSTICS_KEY in result + assert DiagnosticsContext.DIAGNOSTICS_KEY in result d = json.loads(result) - return d[EditingToolWithDiagnostics.DIAGNOSTICS_KEY] + return d[DiagnosticsContext.DIAGNOSTICS_KEY] @contextmanager def project_file_modification_context(serena_agent: SerenaAgent, relative_path: str) -> Iterator[None]: """Context manager to modify a project file and revert the changes after use.""" - project = serena_agent.get_active_project() + project = serena_agent.get_active_project_or_raise() file_path = os.path.join(project.project_root, relative_path) # Read the original content @@ -901,6 +903,43 @@ class TestSerenaAgent: finally: agent.on_shutdown(timeout=5) + @pytest.mark.python + @pytest.mark.skipif(not language_server_tests_enabled(LanguageServerId.PYTHON), reason="python tests are disabled in this environment") + @pytest.mark.parametrize("context_name", ["desktop-app", "grok"], ids=["multi_project", "single_project"]) + def test_repl_interface_exposes_fixed_toolset(self, serena_config, context_name: str): + # the toolset is fixed regardless of tool inclusions/exclusions (e.g. the context's or the configuration's); + # only the single-project property of the context matters (no project activation in that case) + serena_config.agent_interface = AgentInterface.REPL + serena_config.included_optional_tools = ["get_diagnostics_for_symbol"] + context = SerenaAgentContext.from_name(context_name) + agent = SerenaAgent(project="test_repo_python", serena_config=serena_config, context=context) + agent.execute_task(lambda: None) + try: + exposed = {tool.get_name() for tool in agent.get_exposed_tool_instances()} + expected = {"serena_repl", "initial_instructions"} | (set() if context.single_project else {"activate_project"}) + assert exposed == expected + assert "s.lsp" in agent.get_tool(SerenaReplTool).apply(agent.create_session().session_id, "s.info()") + + # the facade listing is part of the (fixed) tool description in single-project sessions, + # and of the activation message otherwise (where the facades depend on the activated project) + tool_description = agent.get_tool(SerenaReplTool).get_apply_docstring() + activation_message = agent.get_project_activation_message("test_session") + assert ("s.lsp:" in tool_description) == context.single_project + assert ("s.lsp:" in activation_message) == (not context.single_project) + + # prompts refer to operations by their qualified REPL names, e.g. `lsp.find_symbol` instead of the tool name + system_prompt = agent.create_system_prompt() + assert "`lsp.find_symbol`" in system_prompt + assert "`find_symbol`" not in system_prompt + + # the instructions establish a session, whose id can be used with session-aware tools + session_id_match = re.search(r"session id is `(\w+)`", system_prompt) + assert session_id_match is not None + session_id = session_id_match.group(1) + assert "s.lsp" in agent.get_tool(SerenaReplTool).apply(session_id, "s.info()") + finally: + agent.on_shutdown(timeout=5) + def _symbol_matches_expected_name(self, symbol: dict, expected_name: str) -> bool: return ( symbol.get("name") == expected_name @@ -1362,13 +1401,17 @@ class TestSerenaAgent: class TestPromptProvision: - class MockContext: - def __init__(self, session_id: str): - self.session = session_id - @classmethod def _call_tool(cls, agent: SerenaAgent, tool_class: type[Tool], session_id: str = "global", **kwargs) -> str: - result = agent.get_tool(tool_class).apply_ex(mcp_ctx=cls.MockContext(session_id), catch_exceptions=False, **kwargs) + old_method = SessionRegistry._next_session_id + if tool_class == InitialInstructionsTool: + SessionRegistry._next_session_id = lambda x: session_id # type: ignore + else: + kwargs["session_id"] = session_id + try: + result = agent.get_tool(tool_class).apply_ex(catch_exceptions=False, **kwargs) + finally: + SessionRegistry._next_session_id = old_method return result @staticmethod @@ -1417,6 +1460,7 @@ class TestPromptProvision: # now activate another project which dynamically enables a new mode (no-onboarding) reg_project = serena_agent.serena_config.get_registered_project(project_name2) + assert reg_project is not None reg_project.project_config.default_modes = ["no-onboarding"] expected_new_mode_message = "The onboarding process is not applied." result2 = self._call_tool(serena_agent, ActivateProjectTool, project=project_name2, session_id=session1) diff --git a/test/serena/test_shell_api.py b/test/serena/test_shell_api.py new file mode 100644 index 00000000..e10ce019 --- /dev/null +++ b/test/serena/test_shell_api.py @@ -0,0 +1,45 @@ +""" +Tests for the shell facade API. +""" + +import sys +from pathlib import Path +from unittest.mock import MagicMock + +import pytest + +from serena.config.serena_config import SerenaConfig +from serena.project import Project +from serena.repl.api.shell_api import ShellApi +from serena.repl.facade import ApiScope, Facade + + +@pytest.fixture +def api(tmp_path: Path) -> ShellApi: + (tmp_path / "sub").mkdir() + project = Project.load(str(tmp_path), serena_config=SerenaConfig(gui_log_window=False, web_dashboard=False)) + agent = MagicMock() + agent.get_active_project_or_raise.return_value = project + agent.serena_config.default_max_tool_answer_chars = 10000 + return ShellApi(agent) + + +def test_facade_exposes_shell_command_as_editing_operation(api: ShellApi) -> None: + facade = Facade.from_api(api, ApiScope()) + assert facade.name == "shell" + assert facade.enabled_method_names == ["execute_shell_command"] + assert facade.get_method("execute_shell_command").info.can_edit + + +def test_execute_shell_command(api: ShellApi, tmp_path: Path) -> None: + print_cwd = "cd" if sys.platform == "win32" else "pwd" + + output = api.execute_shell_command(f"{print_cwd}") + assert output.return_code == 0 + assert Path(output.stdout.strip()).resolve() == tmp_path.resolve() + assert '"stdout"' in output.represent() and '"return_code"' in output.represent() + + # a relative working directory is resolved against the project root and must exist + assert Path(api.execute_shell_command(print_cwd, cwd="sub").stdout.strip()).resolve() == (tmp_path / "sub").resolve() + with pytest.raises(FileNotFoundError): + api.execute_shell_command(print_cwd, cwd="missing") diff --git a/test/serena/test_symbol.py b/test/serena/test_symbol.py index db315270..5dfc2029 100644 --- a/test/serena/test_symbol.py +++ b/test/serena/test_symbol.py @@ -251,10 +251,10 @@ class TestSymbolDictTypes: :param key_type: the corresponding key type (Literal[...]) that the dict should have for keys """ dict_type_keys = dict_type.__annotations__.keys() - assert len(dict_type_keys) == len(key_type.__args__), ( - f"Expected {len(key_type.__args__)} keys in {dict_type}, but got {len(dict_type_keys)}" + assert len(dict_type_keys) == len(key_type.__args__), ( # type: ignore + f"Expected {len(key_type.__args__)} keys in {dict_type}, but got {len(dict_type_keys)}" # type: ignore ) - for expected_key in key_type.__args__: + for expected_key in key_type.__args__: # type: ignore assert expected_key in dict_type_keys, f"Expected key '{expected_key}' not found in {dict_type}" def test_ls_symbol_dict_type(self): diff --git a/test/serena/test_text_utils.py b/test/serena/test_text_utils.py index 991aca9c..61cb3e85 100644 --- a/test/serena/test_text_utils.py +++ b/test/serena/test_text_utils.py @@ -3,7 +3,14 @@ from collections.abc import Callable import pytest from serena.util.file_proxy import FileCollection, FileProxy -from serena.util.text_utils import GlobMatcher, LineType, MultiFileContentReplacer, search_files, search_text +from serena.util.text_utils import ( + ContentReplacer, + GlobMatcher, + LineType, + MultiFileContentReplacer, + search_files, + search_text, +) class TestSearchText: @@ -691,3 +698,51 @@ class TestMultiFileContentReplacer: occ = replacer.find_occurrences([(path, content)], "old_pkg", "new_pkg")[0] with pytest.raises(AssertionError): replacer.apply_to_content("completely different content", [occ]) + + +class TestBackreferenceExpansion: + """$!N backreferences in regex-mode replacements refer to matched groups. A group that + exists but did not participate in the match (e.g. inside an optional construct that was + skipped) must expand to the empty string; a reference to a group that the search + expression does not define at all must fail with an error naming the problem instead of + a raw IndexError. Literal mode has no backreference expansion at all: the replacement is + used verbatim (observed in practice when an agent tried to document the $!N convention + itself and the literal-mode replacement crashed instead of writing the text). + """ + + def test_unmatched_group_expands_to_empty_string(self): + replacer = ContentReplacer(mode="regex", allow_multiple_occurrences=False) + needle = r"EA_INPUT(?:\((\w*)\))?" + + # the group participated and captured an empty string (empty parentheses) + assert replacer.replace("EA_INPUT()\n", needle, r"EA_INPUT$!1(...)") == "EA_INPUT(...)\n" + # the group did not participate at all (no parentheses) + assert replacer.replace("EA_INPUT\n", needle, r"EA_INPUT$!1(...)") == "EA_INPUT(...)\n" + + def test_matched_group_expands_to_its_value(self): + replacer = ContentReplacer(mode="regex", allow_multiple_occurrences=False) + assert replacer.replace("id=alpha", r"id=(\w+)", r"[$!1]") == "[alpha]" + + def test_nonexistent_group_reference_raises_clear_error(self): + replacer = ContentReplacer(mode="regex", allow_multiple_occurrences=False) + with pytest.raises(ValueError, match="does not exist"): + replacer.replace("id=alpha", r"id=(\w+)", r"[$!2]") + + def test_literal_mode_repl_is_verbatim(self): + """Literal mode has no groups at all and no backreference expansion: a replacement + containing $!N sequences is written as-is instead of failing with a backreference error. + """ + replacer = ContentReplacer(mode="literal", allow_multiple_occurrences=False) + assert replacer.replace("literal needle", "literal needle", "$!1 stuff $!2") == "$!1 stuff $!2" + + def test_multi_file_replacer_expands_unmatched_group_to_empty_string(self): + replacer = MultiFileContentReplacer(mode="regex") + files = [("f.txt", "EA_INPUT\n")] + occurrences = replacer.find_occurrences(files, r"EA_INPUT(?:\((\w*)\))?", r"EA_INPUT$!1(...)") + assert [o.replacement for o in occurrences] == ["EA_INPUT(...)"] + + def test_multi_file_replacer_nonexistent_group_reference_raises_clear_error(self): + replacer = MultiFileContentReplacer(mode="regex") + files = [("f.txt", "id=alpha\n")] + with pytest.raises(ValueError, match="does not exist"): + replacer.find_occurrences(files, r"id=(\w+)", r"[$!2]") diff --git a/test/serena/util/test_file_system.py b/test/serena/util/test_file_system.py index 6e366f3b..4e1e7bcb 100644 --- a/test/serena/util/test_file_system.py +++ b/test/serena/util/test_file_system.py @@ -920,3 +920,101 @@ class TestGitignoreParserPermissionError: finally: # Restore permissions so teardown can clean up os.chmod(unreadable, old_mode) + + +class TestWriteFileAtomicSymlinks: + """``write_file_atomic`` replaces ``open(path, "w")`` at its call sites, so it has to agree + with it about symlinks: a plain write follows the link and updates its target, whereas a bare + ``os.replace`` onto the link path would swap the link itself out for a regular file and leave + the target holding stale content (issue #1958 asks for symlink behaviour to be preserved + before source files use this). + """ + + @staticmethod + def _symlink_or_skip(link: Path, target: Path) -> None: + """Windows needs developer mode or admin rights to create a symlink; skip there rather + than fail, matching how ``test_memories_manager.py`` handles the same limitation. + """ + try: + link.symlink_to(target) + except OSError as e: + pytest.skip(f"cannot create symlinks on this platform/permissions: {e}") + + def test_writes_through_a_symlink_instead_of_replacing_it(self, tmp_path): + target = tmp_path / "real.txt" + target.write_text("old", encoding="utf-8") + link = tmp_path / "link.txt" + self._symlink_or_skip(link, target) + + write_file_atomic(str(link), "new", encoding="utf-8") + + assert link.is_symlink(), "the symlink must survive the write, not be replaced by a regular file" + assert target.read_text(encoding="utf-8") == "new", "the content must reach the link's target" + + def test_writes_through_a_symlink_pointing_outside_its_directory(self, tmp_path): + outside = tmp_path / "outside" + outside.mkdir() + target = outside / "real.txt" + target.write_text("old", encoding="utf-8") + inside = tmp_path / "inside" + inside.mkdir() + link = inside / "link.txt" + self._symlink_or_skip(link, target) + + write_file_atomic(str(link), "new", encoding="utf-8") + + assert link.is_symlink() + assert target.read_text(encoding="utf-8") == "new" + assert list(inside.iterdir()) == [link], "no temp file may be left beside the link" + + def test_broken_symlink_creates_its_target(self, tmp_path): + """``open(path, "w")`` on a dangling link creates the target; this must do the same.""" + target = tmp_path / "missing.txt" + link = tmp_path / "link.txt" + self._symlink_or_skip(link, target) + + write_file_atomic(str(link), "new", encoding="utf-8") + + assert link.is_symlink() + assert target.read_text(encoding="utf-8") == "new" + + def test_writes_through_a_symlinked_parent_directory(self, tmp_path): + """The path is resolved in full, so a symlinked *directory* on the way to the file is + followed too, and the temporary file is created in the destination's real directory (it has + to be on the same filesystem as the destination for the rename to be atomic). + """ + real_dir = tmp_path / "real_dir" + real_dir.mkdir() + target = real_dir / "file.txt" + target.write_text("old", encoding="utf-8") + link_dir = tmp_path / "link_dir" + self._symlink_or_skip(link_dir, real_dir) + + write_file_atomic(str(link_dir / "file.txt"), "new", encoding="utf-8") + + assert link_dir.is_symlink(), "the directory symlink must survive" + assert target.read_text(encoding="utf-8") == "new" + assert list(real_dir.iterdir()) == [target], "no temp file may be left in the real directory" + + def test_non_ascii_filename_round_trips(self, tmp_path): + target = tmp_path / "測試檔案.txt" + try: + target.write_text("old", encoding="utf-8") + except (OSError, UnicodeError) as e: + pytest.skip(f"cannot create non-ASCII filenames on this filesystem: {e}") + + write_file_atomic(str(target), "new", encoding="utf-8") + + assert target.read_text(encoding="utf-8") == "new" + assert list(tmp_path.iterdir()) == [target] + + def test_regular_file_is_written_in_place(self, tmp_path): + """Control: the symlink handling must not change the ordinary case.""" + target = tmp_path / "plain.txt" + target.write_text("old", encoding="utf-8") + + write_file_atomic(str(target), "new", encoding="utf-8") + + assert not target.is_symlink() + assert target.read_text(encoding="utf-8") == "new" + assert list(tmp_path.iterdir()) == [target] diff --git a/test/solidlsp/al/test_al_basic.py b/test/solidlsp/al/test_al_basic.py index 62e129e5..e6310c36 100644 --- a/test/solidlsp/al/test_al_basic.py +++ b/test/solidlsp/al/test_al_basic.py @@ -264,7 +264,7 @@ class TestALHoverInjection: char = start.get("character", 0) hover = language_server.request_hover(file_path, line, char) if hover and "contents" in hover: - return hover, hover["contents"].get("value", "") + return hover, hover["contents"].get("value", "") # type: ignore return hover, None return None, None @@ -286,7 +286,7 @@ class TestALHoverInjection: char = start.get("character", 0) hover = language_server.request_hover(file_path, line, char) if hover and "contents" in hover: - return hover, hover["contents"].get("value", "") + return hover, hover["contents"].get("value", "") # type: ignore return hover, None return None, None @@ -373,7 +373,7 @@ class TestALHoverInjection: hover = language_server.request_hover(file_path, line, char) assert hover is not None, "Hover should return a result for field" - value = hover.get("contents", {}).get("value", "") + value = hover.get("contents", {}).get("value", "") # type: ignore # Field hover should NOT start with ** (no injection) assert not value.startswith("**"), f"Field hover should not have injected name. Got: {value[:200]}" return @@ -445,7 +445,7 @@ class TestALPathNormalization: hover = language_server.request_hover(file_path, line, char) assert hover is not None, "Hover should return a result" - value = hover.get("contents", {}).get("value", "") + value = hover.get("contents", {}).get("value", "") # type: ignore assert '**Table 50000 "TEST Customer"**' in value, f"Hover should have injection. Got: {value[:200]}" return @@ -466,7 +466,7 @@ class TestALPathNormalization: hover = language_server.request_hover(file_path, line, char) assert hover is not None, "Hover should return a result" - value = hover.get("contents", {}).get("value", "") + value = hover.get("contents", {}).get("value", "") # type: ignore assert '**Table 50000 "TEST Customer"**' in value, f"Hover should have injection. Got: {value[:200]}" return @@ -491,7 +491,7 @@ class TestALPathNormalization: # Request hover with forward slash path (different format) hover = language_server.request_hover(file_path_forward, line, char) assert hover is not None, "Hover should return a result" - value = hover.get("contents", {}).get("value", "") + value = hover.get("contents", {}).get("value", "") # type: ignore assert '**Table 50000 "TEST Customer"**' in value, ( f"Hover injection should work with mixed path formats. Got: {value[:200]}" ) @@ -518,7 +518,7 @@ class TestALPathNormalization: # Request hover with backslash path (different format) hover = language_server.request_hover(file_path_backslash, line, char) assert hover is not None, "Hover should return a result" - value = hover.get("contents", {}).get("value", "") + value = hover.get("contents", {}).get("value", "") # type: ignore assert '**Table 50000 "TEST Customer"**' in value, ( f"Hover injection should work with mixed path formats. Got: {value[:200]}" ) @@ -553,7 +553,7 @@ class TestALPathNormalization: # Request hover with different path format hover = language_server.request_hover(hover_path, line, char) assert hover is not None, f"Hover should return a result for {symbol_name}" - value = hover.get("contents", {}).get("value", "") + value = hover.get("contents", {}).get("value", "") # type: ignore assert f"**{expected_injection}**" in value, ( f"Hover for {symbol_name} should have injection with mixed paths. Got: {value[:200]}" ) diff --git a/test/solidlsp/angular/test_angular_basic.py b/test/solidlsp/angular/test_angular_basic.py index 8b377d5c..c1d6bbd3 100644 --- a/test/solidlsp/angular/test_angular_basic.py +++ b/test/solidlsp/angular/test_angular_basic.py @@ -125,6 +125,7 @@ class TestAngularLanguageServerBasics: # +1 puts the cursor inside the identifier rather than on its leading boundary. refs = language_server.request_references(src_path, coords.line, coords.col + 1) ref_paths = {r.get("relativePath", "") for r in refs} + ref_paths = {p for p in ref_paths if p} # filter out any empty relativePath entries assert any(p.endswith("app.component.html") for p in ref_paths), ( f"Expected references for setName to include its template callsite in app.component.html, got: {ref_paths}" ) diff --git a/test/solidlsp/angular/test_angular_error_cases.py b/test/solidlsp/angular/test_angular_error_cases.py index 6ce05e8c..53a1a30b 100644 --- a/test/solidlsp/angular/test_angular_error_cases.py +++ b/test/solidlsp/angular/test_angular_error_cases.py @@ -21,6 +21,7 @@ import time import pytest from solidlsp import SolidLanguageServer +from solidlsp.language_servers.angular_language_server import AngularLanguageServer from solidlsp.ls_config import LanguageServerId from solidlsp.ls_exceptions import SolidLSPException from test.conftest import _create_ls @@ -344,6 +345,7 @@ class TestAngularStartupCleanup: with pytest.raises(RuntimeError, match="simulated ngserver init"): ls.start() + assert isinstance(ls, AngularLanguageServer) assert ls._ts_server is None, "TS companion was not cleared after startup failure" assert ls._html_server is None, "HTML companion was not cleared after startup failure" diff --git a/test/solidlsp/clojure/test_clojure_indexing.py b/test/solidlsp/clojure/test_clojure_indexing.py index da5121b0..e1048e79 100644 --- a/test/solidlsp/clojure/test_clojure_indexing.py +++ b/test/solidlsp/clojure/test_clojure_indexing.py @@ -54,7 +54,7 @@ class TestClojureProjectIndexing: # extra.clj contains two real call sites (in double-product and triple-product); # they must be returned regardless of whether the file was opened beforehand - extra_refs = [r for r in refs if r.get("relativePath", "").endswith("extra.clj")] + extra_refs = [r for r in refs if r.get("relativePath", "").endswith("extra.clj")] # type: ignore assert extra_refs, ( "Expected references to 'multiply' to include call sites from extra.clj, " f"but got files: {sorted(ref_paths)}. " @@ -82,7 +82,9 @@ class TestClojureProjectIndexing: ref_paths = {r.get("relativePath", "") for r in refs} consumer_refs = [ - r for r in refs if r.get("relativePath", "").replace("\\", "/").endswith("sub_module/src/sub_module_app/consumer.clj") + r + for r in refs + if r.get("relativePath", "").replace("\\", "/").endswith("sub_module/src/sub_module_app/consumer.clj") # type: ignore ] assert consumer_refs, ( "Expected references to 'multiply' to include call sites from the sibling module " diff --git a/test/solidlsp/crystal/test_crystal_basic.py b/test/solidlsp/crystal/test_crystal_basic.py index ae5d618d..b1a1fd83 100644 --- a/test/solidlsp/crystal/test_crystal_basic.py +++ b/test/solidlsp/crystal/test_crystal_basic.py @@ -14,6 +14,7 @@ import os import pytest from solidlsp import SolidLanguageServer +from solidlsp.language_servers.crystal_language_server import CrystalLanguageServer from solidlsp.ls_config import LanguageServerId from test.conftest import language_server_tests_enabled from test.solidlsp.conftest import format_symbol_for_assert, has_malformed_name, request_all_symbols @@ -94,7 +95,8 @@ class TestCrystalDefinition: file_path = os.path.join("src", "main.cr") # wait for Crystalline to compile the project - language_server.language_server._wait_for_compilation() + assert isinstance(language_server, CrystalLanguageServer) + language_server._wait_for_compilation() # Calculator.new on line 35 (0-indexed: 34), col 13 -> Calculator class on line 3 (0-indexed: 2) definitions = language_server.request_definition(file_path, 34, 13) diff --git a/test/solidlsp/csharp/test_csharp_basic.py b/test/solidlsp/csharp/test_csharp_basic.py index 5f0a0107..1babaf2e 100644 --- a/test/solidlsp/csharp/test_csharp_basic.py +++ b/test/solidlsp/csharp/test_csharp_basic.py @@ -7,6 +7,7 @@ import pytest from solidlsp import SolidLanguageServer from solidlsp.language_servers.csharp_language_server import ( + CSharpLanguageServer, breadth_first_file_scan, find_solution_or_project_file, ) @@ -200,6 +201,28 @@ class TestCSharpLanguageServer: ), f"Expected ConsoleGreeter.FormatGreeting symbol, got: {implementing_symbols}" +class TestCSharpExtractBaseNameAndType: + """Regression tests for _extract_base_name_and_type, no running language server needed.""" + + @pytest.mark.parametrize( + ("roslyn_name", "expected"), + [ + # Property whose type is a tuple: the literal '(' in the type must not be + # mistaken for a method's parameter list. + ("Position : (int X, string Y)", ("Position", ": (int X, string Y)")), + ("Name : string", ("Name", ": string")), + ("Add(int, int) : int", ("Add", "(int, int) : int")), + ("ToString()", ("ToString", "()")), + ("SimpleMethod", ("SimpleMethod", "")), + # Both still have a '(' before the first " : ", so they keep the method branch. + ("GetPair() : (int, int)", ("GetPair", "() : (int, int)")), + ("Merge((int, int) a, (int, int) b) : void", ("Merge", "((int, int) a, (int, int) b) : void")), + ], + ) + def test_extract_base_name_and_type(self, roslyn_name: str, expected: tuple[str, str]) -> None: + assert CSharpLanguageServer._extract_base_name_and_type(roslyn_name) == expected + + @pytest.mark.csharp class TestCSharpSolutionProjectOpening: """Test C# language server solution and project opening functionality.""" diff --git a/test/solidlsp/dart/test_dart_basic.py b/test/solidlsp/dart/test_dart_basic.py index 96f93b7e..598fd879 100644 --- a/test/solidlsp/dart/test_dart_basic.py +++ b/test/solidlsp/dart/test_dart_basic.py @@ -249,6 +249,7 @@ class TestDartLanguageServer: # Find coordinates of 'final result = a + b;' - test position on 'result' with language_server.open_file(file_path, open_in_ls=False) as f: pos = find_text_coordinates(f.contents, r"final (result) = a \+ b;") + assert pos is not None defining_symbol = language_server.request_defining_symbol(file_path, pos.line, pos.col) diff --git a/test/solidlsp/erlang/test_erlang_ignored_dirs.py b/test/solidlsp/erlang/test_erlang_ignored_dirs.py index 38013a6d..5aaa6310 100644 --- a/test/solidlsp/erlang/test_erlang_ignored_dirs.py +++ b/test/solidlsp/erlang/test_erlang_ignored_dirs.py @@ -4,6 +4,7 @@ from pathlib import Path import pytest from solidlsp import SolidLanguageServer +from solidlsp.language_servers.erlang_language_server import ErlangLanguageServer from solidlsp.ls_config import LanguageServerId from test.conftest import language_server_tests_enabled, start_ls_context @@ -146,6 +147,8 @@ def test_symbol_tree_excludes_build_dirs(language_server: SolidLanguageServer): @pytest.mark.parametrize("language_server", [LanguageServerId.ERLANG], indirect=True) def test_ignore_compiled_files(language_server: SolidLanguageServer): """Test that compiled Erlang files are ignored.""" + assert isinstance(language_server, ErlangLanguageServer) + # Test that beam files are ignored assert language_server.is_ignored_filename("module.beam"), "BEAM files should be ignored" assert language_server.is_ignored_filename("app.beam"), "BEAM files should be ignored" @@ -164,6 +167,7 @@ def test_rebar_directories_ignored(language_server: SolidLanguageServer): assert language_server.is_ignored_dirname(".rebar3"), "rebar3 cache should be ignored" # Test that rebar.lock and rebar.config are not ignored (they are configuration files) + assert isinstance(language_server, ErlangLanguageServer) assert not language_server.is_ignored_filename("rebar.config"), "rebar.config should not be ignored" assert not language_server.is_ignored_filename("rebar.lock"), "rebar.lock should not be ignored" diff --git a/test/solidlsp/julia/test_fatou.py b/test/solidlsp/julia/test_fatou.py index d6f7a6d5..daa6bdcc 100644 --- a/test/solidlsp/julia/test_fatou.py +++ b/test/solidlsp/julia/test_fatou.py @@ -25,7 +25,7 @@ class TestFatouLanguageServer: def test_cross_file_references(self, language_server: SolidLanguageServer) -> None: references = language_server.request_references("src/fatou_a.jl", line=0, column=2) - locations = {(reference["relativePath"].replace("\\", "/"), reference["range"]["start"]["line"]) for reference in references} + locations = {(reference["relativePath"].replace("\\", "/"), reference["range"]["start"]["line"]) for reference in references} # type: ignore assert locations >= {("src/fatou_a.jl", 1), ("src/fatou_b.jl", 0)} def test_file_matching(self) -> None: diff --git a/test/solidlsp/kotlin/test_kotlin_dependency_provider.py b/test/solidlsp/kotlin/test_kotlin_dependency_provider.py index a2c06283..16d84bad 100644 --- a/test/solidlsp/kotlin/test_kotlin_dependency_provider.py +++ b/test/solidlsp/kotlin/test_kotlin_dependency_provider.py @@ -38,42 +38,42 @@ class TestKotlinDependencyProvider: ".win.zip", "zip", ("bin", "intellij-server.exe"), - "f2daaa476f26d99301b406f76de6d87c437d04dc72f06845154619d8f991c51f", + "a9b471b16025b1bfb3b0a097862580abb40e3c35406c44242c18b1d70f5d0e44", ), ( PlatformId.WIN_arm64, "-aarch64.win.zip", "zip", ("bin", "intellij-server.exe"), - "73a552a6a420158622e5ad8d96b53da8aa8ced3f88a24fded01575927a2fd8e7", + "3bf008d8c94fa70eb13fc998eaa42f29b9d13f368984d4cec46277808f94e1de", ), ( PlatformId.LINUX_x64, ".tar.gz", "gztar", (f"kotlin-server-{DEFAULT_KOTLIN_LSP_VERSION}", "bin", "intellij-server"), - "2d99d8e198fbe4aa8f4481e37799724ce94803b4ea12a60b416040e3fcd7cc5e", + "1e11d2e5fefbf9ea215ad8dd6be95f2222897cd086e8cb7a661a52084a590405", ), ( PlatformId.LINUX_arm64, "-aarch64.tar.gz", "gztar", (f"kotlin-server-{DEFAULT_KOTLIN_LSP_VERSION}", "bin", "intellij-server"), - "2317831c6e5607d05b7ebc1da655330125ce0e3d66fbf24517dfce442debc14e", + "ec7cb254a6662a07fff9f10e4365226afab6c40008f8a974c10ac5e785d6510f", ), ( PlatformId.OSX_x64, ".sit", "zip", (f"kotlin-server-{DEFAULT_KOTLIN_LSP_VERSION}", "bin", "intellij-server"), - "17369fda97c85418ac24ab38a9df56b21522a3468dfe193832fe455c13920745", + "62ab735947b1c855b505f64f5db8fbd7ff0b52a35ab1897938c6dbfc7b24c8a3", ), ( PlatformId.OSX_arm64, "-aarch64.sit", "zip", (f"kotlin-server-{DEFAULT_KOTLIN_LSP_VERSION}", "bin", "intellij-server"), - "6ba6021a706b21e64cef33f7e2b79f187c0910320722bb2d3ed05ad1115ec43f", + "95da3fc6d3b9092c7616345044a05edb85e5408dc648d081e4e433595c892bec", ), ], ) @@ -176,12 +176,12 @@ class TestKotlinDependencyProvider: "kotlin-lsp-261.13587.0-linux-aarch64.zip", "kotlin-lsp-261.13587.0-mac-x64.zip", "kotlin-lsp-261.13587.0-mac-aarch64.zip", - "kotlin-server-262.9593.0.win.zip", - "kotlin-server-262.9593.0-aarch64.win.zip", - "kotlin-server-262.9593.0.tar.gz", - "kotlin-server-262.9593.0-aarch64.tar.gz", - "kotlin-server-262.9593.0.sit", - "kotlin-server-262.9593.0-aarch64.sit", + "kotlin-server-263.4702.0.win.zip", + "kotlin-server-263.4702.0-aarch64.win.zip", + "kotlin-server-263.4702.0.tar.gz", + "kotlin-server-263.4702.0-aarch64.tar.gz", + "kotlin-server-263.4702.0.sit", + "kotlin-server-263.4702.0-aarch64.sit", } @pytest.mark.parametrize( diff --git a/test/solidlsp/pascal/test_pascal_basic.py b/test/solidlsp/pascal/test_pascal_basic.py index 74576459..bfa14728 100644 --- a/test/solidlsp/pascal/test_pascal_basic.py +++ b/test/solidlsp/pascal/test_pascal_basic.py @@ -192,7 +192,7 @@ class TestPascalLanguageServerBasics: contents = hover.get("contents", {}) value = contents.get("value", "") if isinstance(contents, dict) else str(contents) else: - value = hover.contents.value if hasattr(hover.contents, "value") else str(hover.contents) + value = hover.contents.value if hasattr(hover.contents, "value") else str(hover.contents) # type: ignore # Should contain the function signature assert "CalculateSum" in value, f"Hover should show function name. Got: {value[:500]}" diff --git a/test/solidlsp/python/test_symbol_retrieval.py b/test/solidlsp/python/test_symbol_retrieval.py index 38185f31..45460d73 100644 --- a/test/solidlsp/python/test_symbol_retrieval.py +++ b/test/solidlsp/python/test_symbol_retrieval.py @@ -46,6 +46,7 @@ class TestLanguageServerSymbols: with language_server.open_file(file_path, open_in_ls=False) as f: file_content = f.contents coords = find_text_coordinates(file_content, r"(status): str") + assert coords is not None ref_symbols = [ref.symbol for ref in language_server.request_referencing_symbols(file_path, coords.line, coords.col)] assert len(ref_symbols) > 0 diff --git a/test/solidlsp/r/test_r_basic.py b/test/solidlsp/r/test_r_basic.py index 84a9ab51..e44b1864 100644 --- a/test/solidlsp/r/test_r_basic.py +++ b/test/solidlsp/r/test_r_basic.py @@ -9,7 +9,7 @@ import pytest from solidlsp import SolidLanguageServer from solidlsp.ls_config import LanguageServerId -from test.conftest import language_server_tests_enabled +from test.conftest import is_ci, language_server_tests_enabled from test.solidlsp.conftest import format_symbol_for_assert, has_malformed_name, request_all_symbols @@ -41,6 +41,7 @@ class TestRLanguageServer: expected_functions = {"calculate_mean", "process_data", "create_data_frame"} assert expected_functions.issubset(function_names), f"Expected functions {expected_functions} but found {function_names}" + @pytest.mark.xfail(is_ci, reason="Test is flaky") # See #1040 @pytest.mark.parametrize("language_server", [LanguageServerId.R], indirect=True) def test_find_definition_across_files(self, language_server: SolidLanguageServer): """Test finding function definitions across files.""" @@ -58,6 +59,7 @@ class TestRLanguageServer: # Definition should be around line 37 (0-indexed: 36) where create_data_frame is defined assert definition_location["range"]["start"]["line"] >= 35 + @pytest.mark.xfail(is_ci, reason="Test is flaky") # See #1040 @pytest.mark.parametrize("language_server", [LanguageServerId.R], indirect=True) def test_find_references_across_files(self, language_server: SolidLanguageServer): """Test finding function references across files.""" diff --git a/test/solidlsp/ruby/test_ruby_symbol_retrieval.py b/test/solidlsp/ruby/test_ruby_symbol_retrieval.py index 67a29eb2..1f4e9aa0 100644 --- a/test/solidlsp/ruby/test_ruby_symbol_retrieval.py +++ b/test/solidlsp/ruby/test_ruby_symbol_retrieval.py @@ -590,7 +590,9 @@ class TestRubyLanguageServerSymbols: pos = find_text_coordinates(fb.contents, r"user = @service\.(create_user)") # Verify that we can find the method definition + assert pos is not None defining_symbol = language_server.request_defining_symbol(file_path, pos.line, pos.col) + assert defining_symbol is not None assert "name" in defining_symbol assert "kind" in defining_symbol assert defining_symbol.get("name") == "create_user" diff --git a/test/solidlsp/rust/test_rust_basic.py b/test/solidlsp/rust/test_rust_basic.py index 446cdfbe..ac7c8d06 100644 --- a/test/solidlsp/rust/test_rust_basic.py +++ b/test/solidlsp/rust/test_rust_basic.py @@ -68,7 +68,7 @@ class TestRustLanguageServer: implementations = language_server.request_implementation(os.path.join("src", "lib.rs"), *pos) assert implementations, "Expected at least one implementation of Greeter.format_greeting" - assert any("src/lib.rs" in implementation.get("relativePath", "").replace("\\", "/") for implementation in implementations), ( + assert any("src/lib.rs" in implementation.get("relativePath", "").replace("\\", "/") for implementation in implementations), ( # type: ignore f"Expected ConsoleGreeter.format_greeting in implementations, got: {implementations}" ) @@ -81,7 +81,7 @@ class TestRustLanguageServer: implementing_symbols = language_server.request_implementing_symbols(os.path.join("src", "lib.rs"), *pos) assert implementing_symbols, "Expected implementing symbols for Greeter.format_greeting" assert any( - symbol.get("name") == "format_greeting" and "src/lib.rs" in symbol["location"].get("relativePath", "").replace("\\", "/") + symbol.get("name") == "format_greeting" and "src/lib.rs" in symbol["location"].get("relativePath", "").replace("\\", "/") # type: ignore for symbol in implementing_symbols ), f"Expected ConsoleGreeter.format_greeting symbol, got: {implementing_symbols}" diff --git a/test/solidlsp/scss/test_scss_basic.py b/test/solidlsp/scss/test_scss_basic.py index 7266f85d..9ae51a34 100644 --- a/test/solidlsp/scss/test_scss_basic.py +++ b/test/solidlsp/scss/test_scss_basic.py @@ -130,6 +130,7 @@ class TestScssReferences: line, col = coords.line, coords.col refs = language_server.request_references(path, line, col + 2) ref_paths = {r.get("relativePath", "") for r in refs} + ref_paths = {r for r in ref_paths if r} # filter out empty strings assert any(p.endswith("buttons.scss") for p in ref_paths), ( f"Expected card-surface references to include buttons.scss, got: {ref_paths}" ) @@ -148,6 +149,7 @@ class TestScssReferences: line, col = coords.line, coords.col refs = language_server.request_references(path, line, col + 2) ref_paths = {r.get("relativePath", "") for r in refs} + ref_paths = {r for r in ref_paths if r} # filter out empty strings assert any(p.endswith("buttons.scss") for p in ref_paths), ( f"Expected $color-primary references to include buttons.scss, got: {ref_paths}" ) diff --git a/test/solidlsp/svelte/test_svelte_basic.py b/test/solidlsp/svelte/test_svelte_basic.py index 3757ff83..e799c139 100644 --- a/test/solidlsp/svelte/test_svelte_basic.py +++ b/test/solidlsp/svelte/test_svelte_basic.py @@ -20,7 +20,7 @@ class TestSvelteLanguageServer: def test_svelte_language_server_root_matches_repo_path(self, language_server: SolidLanguageServer, repo_path: Path) -> None: assert language_server.is_running() assert repo_path.resolve() == svelte_test_conftest.repo_path.resolve() - assert Path(language_server.language_server.repo_path).resolve() == repo_path.resolve() + assert Path(language_server.repository_root_path).resolve() == repo_path.resolve() @pytest.mark.parametrize("language_server", [LanguageServerId.SVELTE], indirect=True) def test_svelte_and_typescript_files_in_symbol_tree(self, language_server: SolidLanguageServer) -> None: @@ -74,12 +74,13 @@ class TestSvelteLanguageServer: def test_definition_from_component_import_to_svelte_file(self, language_server: SolidLanguageServer) -> None: file_path = os.path.join("src", "lib", "components", "Header.svelte") coords = find_text_coordinates(read_repo_file(language_server, file_path), r"(count)") + assert coords is not None definitions = language_server.request_definition(file_path, coords.line, coords.col) - definition_paths = sorted(definition["relativePath"].replace("\\", "/") for definition in definitions) + definition_paths = sorted(definition["relativePath"].replace("\\", "/") for definition in definitions) # type: ignore assert len(definitions) == 1, definition_paths - assert definitions[0]["relativePath"].replace("\\", "/") == "src/lib/components/Counter.svelte", definition_paths + assert definitions[0]["relativePath"].replace("\\", "/") == "src/lib/components/Counter.svelte", definition_paths # type: ignore @pytest.mark.parametrize("language_server", [LanguageServerId.SVELTE], indirect=True) def test_diagnostics_in_typescript_file(self, language_server: SolidLanguageServer) -> None: diff --git a/test/solidlsp/svelte/test_svelte_references.py b/test/solidlsp/svelte/test_svelte_references.py index 2fc7bdbf..c7777799 100644 --- a/test/solidlsp/svelte/test_svelte_references.py +++ b/test/solidlsp/svelte/test_svelte_references.py @@ -12,7 +12,7 @@ class TestSvelteReferences: @pytest.mark.parametrize("language_server", [LanguageServerId.SVELTE], indirect=True) def test_references_across_svelte_and_typescript(self, language_server: SolidLanguageServer) -> None: refs = language_server.request_references(os.path.join("src", "lib", "components", "Words.svelte"), 1, 17) - ref_paths = {ref["relativePath"].replace("\\", "/") for ref in refs} + ref_paths = {ref["relativePath"].replace("\\", "/") for ref in refs} # type: ignore assert "src/routes/(sverdle)/words.server.ts" in ref_paths, sorted(ref_paths) assert "src/lib/game.ts" in ref_paths, sorted(ref_paths) @@ -21,6 +21,6 @@ class TestSvelteReferences: @pytest.mark.parametrize("language_server", [LanguageServerId.SVELTE], indirect=True) def test_references_from_typescript_file(self, language_server: SolidLanguageServer) -> None: refs = language_server.request_references(os.path.join("src", "lib", "game.ts"), 3, 13) - ref_paths = {ref["relativePath"].replace("\\", "/") for ref in refs} + ref_paths = {ref["relativePath"].replace("\\", "/") for ref in refs} # type: ignore assert "src/routes/(sverdle)/+page.server.ts" in ref_paths, sorted(ref_paths) diff --git a/test/solidlsp/svelte/test_svelte_rename.py b/test/solidlsp/svelte/test_svelte_rename.py index ec550b82..c40e5d78 100644 --- a/test/solidlsp/svelte/test_svelte_rename.py +++ b/test/solidlsp/svelte/test_svelte_rename.py @@ -56,6 +56,7 @@ class TestSvelteRename: def test_rename_svelte_export_updates_svelte_importers(self, language_server: SolidLanguageServer) -> None: file_path = os.path.join("src", "lib", "components", "Counter.svelte") coords = find_text_coordinates(read_repo_file(language_server, file_path), r"(count)") + assert coords is not None workspace_edit = language_server.request_rename_symbol_edit(file_path, coords.line, coords.col, "score") @@ -69,6 +70,7 @@ class TestSvelteRename: def test_rename_svelte_export_updates_ts_and_svelte_files(self, language_server: SolidLanguageServer) -> None: file_path = os.path.join("src", "lib", "components", "Words.svelte") coords = find_text_coordinates(read_repo_file(language_server, file_path), r"(words)") + assert coords is not None workspace_edit = language_server.request_rename_symbol_edit(file_path, coords.line, coords.col, "vocabulary") diff --git a/test/solidlsp/test_dart_root_uri.py b/test/solidlsp/test_dart_root_uri.py new file mode 100644 index 00000000..46d6f4d6 --- /dev/null +++ b/test/solidlsp/test_dart_root_uri.py @@ -0,0 +1,38 @@ +# SPDX-License-Identifier: MIT + +import tempfile +from pathlib import Path + +from solidlsp.language_servers.dart_language_server import DartLanguageServer + + +def _make_dart_ls() -> DartLanguageServer: + ls = object.__new__(DartLanguageServer) + ls._custom_settings = {} + # Windows: as_uri() rejects drive-less paths like "/tmp/..." + project_dir = Path(tempfile.mkdtemp(prefix="fake-dart-project-")) / "project" + project_dir.mkdir(parents=True, exist_ok=True) + ls.repository_root_path = str(project_dir) + + class _Cfg: + @staticmethod + def get_absolute_workspace_folders(root): + return [root] + + @staticmethod + def get_absolute_additional_workspace_folders(root): + return [] + + ls.config = _Cfg() + # custom_settings property reads from _custom_settings on SolidLanguageServer + return ls + + +def test_dart_omits_root_uri(): + builder = _make_dart_ls()._create_initialize_params_builder() + params = builder.build() + # rootUri must be present (some servers reject an undefined key) but null, so that only + # workspaceFolders determine the analysis roots (oraios/serena#2045). + assert params["rootUri"] is None + assert params["rootPath"] is None + assert params["workspaceFolders"] diff --git a/test/solidlsp/test_initialize_params_root_uri.py b/test/solidlsp/test_initialize_params_root_uri.py new file mode 100644 index 00000000..7ab80c98 --- /dev/null +++ b/test/solidlsp/test_initialize_params_root_uri.py @@ -0,0 +1,39 @@ +# SPDX-License-Identifier: MIT + +import tempfile +from pathlib import Path + +from solidlsp.initialize_params import DefaultInitializeParamsBuilder + + +class _FakeLS: + # Windows: as_uri() rejects drive-less paths like "/tmp/..." + repository_root_path = str(Path(tempfile.mkdtemp(prefix="fake-project-")) / "root") + + class config: + @staticmethod + def get_absolute_workspace_folders(root): + return [root] + + @staticmethod + def get_absolute_additional_workspace_folders(root): + return [] + + custom_settings: dict = {} + + +def test_default_builder_sets_root_uri(): + builder = DefaultInitializeParamsBuilder(_FakeLS()) + params = builder.build() + assert "rootUri" in params + assert "rootPath" in params + + +def test_builder_sends_null_root_uri_when_disabled(): + builder = DefaultInitializeParamsBuilder(_FakeLS(), set_root_uri=False) + params = builder.build() + # keys must be present (some servers reject an undefined rootUri); values are null + assert params["rootUri"] is None + assert params["rootPath"] is None + assert params["processId"] is not None + assert params["clientInfo"] == {"name": "Serena"} diff --git a/test/solidlsp/test_ls_start_cleanup.py b/test/solidlsp/test_ls_start_cleanup.py index 7e574311..5246ff43 100644 --- a/test/solidlsp/test_ls_start_cleanup.py +++ b/test/solidlsp/test_ls_start_cleanup.py @@ -38,7 +38,7 @@ def test_start_stops_process_when_start_server_raises_after_spawning(): with pytest.raises(RuntimeError, match="capability assertion"): server.start() - server.server.stop.assert_called_once() + server.server.stop.assert_called_once() # type: ignore assert server.server_started is False @@ -51,5 +51,5 @@ def test_start_does_not_call_stop_when_start_server_raises_before_spawning(): with pytest.raises(RuntimeError, match="capability assertion"): server.start() - server.server.stop.assert_not_called() + server.server.stop.assert_not_called() # type: ignore assert server.server_started is False diff --git a/test/solidlsp/test_pdeathsig.py b/test/solidlsp/test_pdeathsig.py index 662143c9..9be4be2a 100644 --- a/test/solidlsp/test_pdeathsig.py +++ b/test/solidlsp/test_pdeathsig.py @@ -115,8 +115,8 @@ def test_language_server_process_survives_a_short_lived_calling_thread() -> None text=True, ) try: - ready_line = driver.stdout.readline() - assert ready_line.strip() == "READY", f"driver failed to start the language server: {driver.stderr.read()}" + ready_line = driver.stdout.readline() # type: ignore + assert ready_line.strip() == "READY", f"driver failed to start the language server: {driver.stderr.read()}" # type: ignore time.sleep(2) # The driver's own argv also contains `marker` (it's passed as sys.argv[1]), so exclude @@ -143,8 +143,10 @@ def test_language_server_process_dies_with_a_sigkilled_serena() -> None: text=True, ) try: - ready_line = driver.stdout.readline() - assert ready_line.strip() == "READY", f"driver failed to start the language server: {driver.stderr.read()}" + ready_line = driver.stdout.readline() # type: ignore + stderr = driver.stderr + assert stderr is not None, "stderr should be captured" + assert ready_line.strip() == "READY", f"driver failed to start the language server: {stderr.read()}" assert _find_marked_processes(marker), "language server process never started" driver.kill() # SIGKILL: simulates Serena being killed without a chance to clean up diff --git a/test/solidlsp/test_process_group_cleanup.py b/test/solidlsp/test_process_group_cleanup.py index 9bd4098c..473a803b 100644 --- a/test/solidlsp/test_process_group_cleanup.py +++ b/test/solidlsp/test_process_group_cleanup.py @@ -80,7 +80,9 @@ def _spawn_ready(src: str) -> subprocess.Popen: test_pdeathsig.py's driver pattern (deterministic sync instead of a blind sleep). """ proc = subprocess.Popen([sys.executable, "-c", src], start_new_session=True, stdout=subprocess.PIPE, text=True) - ready_line = proc.stdout.readline() + stdout = proc.stdout + assert stdout is not None + ready_line = stdout.readline() assert ready_line.strip() == "READY", f"helper process failed to start: {ready_line!r}" return proc @@ -341,8 +343,10 @@ class TestPsutilDenialConsequences: """ ) proc = subprocess.Popen([sys.executable, "-c", src], start_new_session=True, stdout=subprocess.PIPE, text=True) - child_pid = int(proc.stdout.readline().strip()) - ready_line = proc.stdout.readline() + stdout = proc.stdout + assert stdout is not None + child_pid = int(stdout.readline().strip()) + ready_line = stdout.readline() assert ready_line.strip() == "READY", f"helper process failed to start: {ready_line!r}" return proc, child_pid diff --git a/test/solidlsp/util/test_godot_symbol_fix.py b/test/solidlsp/util/test_godot_symbol_fix.py new file mode 100644 index 00000000..0b3559a2 --- /dev/null +++ b/test/solidlsp/util/test_godot_symbol_fix.py @@ -0,0 +1,80 @@ +"""Unit tests for GodotLanguageServer's document-symbol range correction. + +See oraios/serena#1974: Godot's GDScript parser can report a symbol's end column one +column past the line-end convention every other language server follows, which silently +rolls over into the following line's content (or, at the last line of a body, into the +separating blank line) when the position is later turned into a text index. +""" + +from __future__ import annotations + +from solidlsp.language_servers.godot_language_server import GodotLanguageServer + +_fix_range_end = GodotLanguageServer._fix_range_end +_fix_symbol_ranges = GodotLanguageServer._fix_symbol_ranges + + +def _range(end_line: int, end_char: int, start_line: int = 0, start_char: int = 0) -> dict: + return {"start": {"line": start_line, "character": start_char}, "end": {"line": end_line, "character": end_char}} + + +# The exact fixture from the reported issue: a 13-char line whose reported end column is 14. +_ISSUE_1952_LINES = ["extends Node", "", "func first():", '\tprint("old")', "", "func second():", '\tprint("keep")'] + + +def test_fix_range_end_corrects_the_measured_off_by_one() -> None: + rng = _range(end_line=3, end_char=14) + _fix_range_end(rng, _ISSUE_1952_LINES) + assert rng["end"]["character"] == 13 + + +def test_fix_range_end_leaves_a_correct_one_past_end_column_alone() -> None: + rng = _range(end_line=3, end_char=13) + _fix_range_end(rng, _ISSUE_1952_LINES) + assert rng["end"]["character"] == 13 + + +def test_fix_range_end_does_not_guess_at_a_larger_overshoot() -> None: + # Two past the line's end is not the measured Godot mechanism; leave it as reported + # rather than assume the same +1 correction applies. + rng = _range(end_line=3, end_char=15) + _fix_range_end(rng, _ISSUE_1952_LINES) + assert rng["end"]["character"] == 15 + + +def test_fix_range_end_ignores_an_out_of_range_line() -> None: + rng = _range(end_line=99, end_char=5) + _fix_range_end(rng, _ISSUE_1952_LINES) + assert rng["end"]["character"] == 5 + + +def test_fix_range_end_tolerates_a_missing_end() -> None: + _fix_range_end({"start": {"line": 0, "character": 0}}, _ISSUE_1952_LINES) # must not raise + + +def test_fix_symbol_ranges_corrects_range_selection_range_and_location_recursively() -> None: + # line 2, "func first():", is 13 chars: end_char=14 is the off-by-one, end_char=5 is not. + child = { + "name": "inner", + "range": _range(end_line=3, end_char=14), + "selectionRange": _range(end_line=2, end_char=14, start_line=2), + } + root = { + "name": "first", + "range": _range(end_line=3, end_char=14), + "selectionRange": _range(end_line=2, end_char=5, start_line=2), + "children": [child], + } + _fix_symbol_ranges(root, _ISSUE_1952_LINES) + + assert root["range"]["end"]["character"] == 13 + assert root["selectionRange"]["end"]["character"] == 5 # unaffected, no overshoot + assert child["range"]["end"]["character"] == 13 + assert child["selectionRange"]["end"]["character"] == 13 # line 2 is 13 chars, 14 was the overshoot + + +def test_fix_symbol_ranges_corrects_symbol_information_style_location() -> None: + # SymbolInformation (the flat, non-hierarchical shape) nests its range under "location". + symbol = {"name": "first", "location": {"uri": "file:///script.gd", "range": _range(end_line=3, end_char=14)}} + _fix_symbol_ranges(symbol, _ISSUE_1952_LINES) + assert symbol["location"]["range"]["end"]["character"] == 13 diff --git a/uv.lock b/uv.lock index cb77e62c..ea17f0cc 100644 --- a/uv.lock +++ b/uv.lock @@ -821,6 +821,19 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/7e/f5/f66802a942d491edb555dd61e3a9961140fd64c90bce1eafd741609d334d/httpcore-1.0.9-py3-none-any.whl", hash = "sha256:2d400746a40668fc9dec9810239072b40b4484b640a8c38fd654a024c7a1bf55", size = 78784, upload-time = "2025-04-24T22:06:20.566Z" }, ] +[[package]] +name = "httpcore2" +version = "2.9.1" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "h11" }, + { name = "truststore" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/39/a8/20ed1ed79cbc2ecdf5301c0968ab7c85547212e2a7bd126ddd2d986e206e/httpcore2-2.9.1.tar.gz", hash = "sha256:4d8acbf8b306f48c9d6046591fd5ba4037d1b1b1000d140fc2c3eab1e9a0c0e2", size = 67089, upload-time = "2026-07-24T09:21:03.867Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/9f/fb/46c52b781975c335a2bcf1072c7bbc007cbdc8d674217f5ee1daba2c848b/httpcore2-2.9.1-py3-none-any.whl", hash = "sha256:6182472379e855fe4221246a2bb7ecede403bc61c6798062ae1787d051ccde26", size = 82809, upload-time = "2026-07-24T09:21:01.178Z" }, +] + [[package]] name = "httpx" version = "0.28.1" @@ -842,12 +855,19 @@ http2 = [ ] [[package]] -name = "httpx-sse" -version = "0.4.3" +name = "httpx2" +version = "2.9.1" source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/0f/4c/751061ffa58615a32c31b2d82e8482be8dd4a89154f003147acee90f2be9/httpx_sse-0.4.3.tar.gz", hash = "sha256:9b1ed0127459a66014aec3c56bebd93da3c1bc8bb6618c8082039a44889a755d", size = 15943, upload-time = "2025-10-10T21:48:22.271Z" } +dependencies = [ + { name = "anyio" }, + { name = "httpcore2" }, + { name = "idna" }, + { name = "truststore" }, + { name = "typing-extensions", marker = "python_full_version < '3.13'" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/21/14/38128fbafd7e0ed41d874df6c9a653d47c2d111cfe59e2b4ac95161b4abd/httpx2-2.9.1.tar.gz", hash = "sha256:1932a768737e3666291582833da748cc4e563c337cf96706fccc04fa6e58764a", size = 95458, upload-time = "2026-07-24T09:21:04.972Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/d2/fd/6668e5aec43ab844de6fc74927e155a3b37bf40d7c3790e49fc0406b6578/httpx_sse-0.4.3-py3-none-any.whl", hash = "sha256:0ac1c9fe3c0afad2e0ebb25a934a59f4c7823b60792691f779fad2c5568830fc", size = 8960, upload-time = "2025-10-10T21:48:21.158Z" }, + { url = "https://files.pythonhosted.org/packages/13/b8/cfd91c4ab9134d386d48f0b6ac662ff3d4be6efdee59ee1c67ebc3c0487c/httpx2-2.9.1-py3-none-any.whl", hash = "sha256:1820fe14a9ab1107bfeff39259987429450b070ec0ff38cc87eb0d8c97fdc71a", size = 91191, upload-time = "2026-07-24T09:21:02.6Z" }, ] [[package]] @@ -861,11 +881,11 @@ wheels = [ [[package]] name = "idna" -version = "3.15" +version = "3.18" source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/82/77/7b3966d0b9d1d31a36ddf1746926a11dface89a83409bf1483f0237aa758/idna-3.15.tar.gz", hash = "sha256:ca962446ea538f7092a95e057da437618e886f4d349216d2b1e294abfdb65fdc", size = 199245, upload-time = "2026-05-12T22:45:57.011Z" } +sdist = { url = "https://files.pythonhosted.org/packages/cd/63/9496c57188a2ee585e0f1db071d75089a11e98aa86eb99d9d7618fc1edce/idna-3.18.tar.gz", hash = "sha256:ffb385a7e039654cef1ab9ef32c6fafe283c0c0467bba1d9029738ce4a14a848", size = 196711, upload-time = "2026-06-02T14:34:07.794Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/d2/23/408243171aa9aaba178d3e2559159c24c1171a641aa83b67bdd3394ead8e/idna-3.15-py3-none-any.whl", hash = "sha256:048adeaf8c2d788c40fee287673ccaa74c24ffd8dcf09ffa555a2fbb59f10ac8", size = 72340, upload-time = "2026-05-12T22:45:55.733Z" }, + { url = "https://files.pythonhosted.org/packages/1e/5e/d4e9f1a599fb8e573b7b87160658329fbf28d19eac2718f51fc3def3aa5a/idna-3.18-py3-none-any.whl", hash = "sha256:7f952cbe720b688055e3f87de14f5c3e5fdaa8bc3928985c4077ca689de849a2", size = 65455, upload-time = "2026-06-02T14:34:06.319Z" }, ] [[package]] @@ -1353,15 +1373,15 @@ wheels = [ [[package]] name = "mcp" -version = "1.28.1" +version = "2.2.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "anyio" }, - { name = "httpx" }, - { name = "httpx-sse" }, + { name = "httpx2" }, { name = "jsonschema" }, + { name = "mcp-types" }, + { name = "opentelemetry-api" }, { name = "pydantic" }, - { name = "pydantic-settings" }, { name = "pyjwt", extra = ["crypto"] }, { name = "python-multipart" }, { name = "pywin32", marker = "sys_platform == 'win32'" }, @@ -1371,9 +1391,22 @@ dependencies = [ { name = "typing-inspection" }, { name = "uvicorn", marker = "sys_platform != 'emscripten'" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/6e/77/9450b8f251a13affb6281997d0523c4615f8a8b35d0b21ff30db3a5aac9d/mcp-1.28.1.tar.gz", hash = "sha256:d51e36a5f5644faea4f85ea649bfffa6bc6c26770d42798ad6a3de3d2ba69683", size = 638501, upload-time = "2026-06-26T12:57:29.093Z" } +sdist = { url = "https://files.pythonhosted.org/packages/76/31/ac54fb0fdd5b37de704486e288bba4fbbb463f24cfcfedbede407b854513/mcp-2.2.0.tar.gz", hash = "sha256:2dc37ecb1974becdcebdbf7561e7c15a07dbbf20ba21ba16c3593b3038b3afbd", size = 4084129, upload-time = "2026-09-07T16:06:23.439Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/e2/5e/d118fce19f87a2e7d8101c35c8ae0ec289098a4df0ff244cec23e415aca0/mcp-1.28.1-py3-none-any.whl", hash = "sha256:2726bca5e7193f61c5dde8b12500a6de2d9acf6d1a1c0be9e8c2e706437991df", size = 222620, upload-time = "2026-06-26T12:57:27.218Z" }, + { url = "https://files.pythonhosted.org/packages/1b/ff/8e7eade68b8a28f7da0ed1085544341b51f9c935dbf6b95c76b7edfea6a0/mcp-2.2.0-py3-none-any.whl", hash = "sha256:bde982589473a060ae145e3406e9a5333fe538c97229ba841f5a7f92be004f81", size = 365656, upload-time = "2026-09-07T16:06:19.711Z" }, +] + +[[package]] +name = "mcp-types" +version = "2.2.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "pydantic" }, + { name = "typing-extensions" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/ae/91/762d7755d971aff8a28d75f7961656148edf27875c8026e6385aaab08ae7/mcp_types-2.2.0.tar.gz", hash = "sha256:d3ed53703ddd10d9c6399f29d322bb66f3f67ab41348ac8556ba23e07fedefad", size = 65892, upload-time = "2026-09-07T16:06:25.187Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/8f/d7/6ffba5d8cd5dd9b8a19478875c50e04945314ba5074e84d749283f27f62d/mcp_types-2.2.0-py3-none-any.whl", hash = "sha256:ea476b73ee86709ab5abc9452385ed36cc05907e582355622e294595c9a04f13", size = 69106, upload-time = "2026-09-07T16:06:21.461Z" }, ] [[package]] @@ -1565,6 +1598,18 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/a0/c4/c2971a3ba4c6103a3d10c4b0f24f461ddc027f0f09763220cf35ca1401b3/nest_asyncio-1.6.0-py3-none-any.whl", hash = "sha256:87af6efd6b5e897c81050477ef65c62e2b2f35d51703cae01aff2905b1852e1c", size = 5195, upload-time = "2024-01-21T14:25:17.223Z" }, ] +[[package]] +name = "opentelemetry-api" +version = "1.44.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "typing-extensions" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/ee/8b/aa9e2d8b8dfa7c946f7dec5d1f8f6ba8eca062f43509a06bdb5ce93d26c0/opentelemetry_api-1.44.0.tar.gz", hash = "sha256:67647e5e9566edcf421166fdf022b3537f818635daa852b289e34604dc6fb33a", size = 72406, upload-time = "2026-07-16T15:25:32.678Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/ca/6f/a04e900f465ff3221ccc395522503e2d10e79fa21f2723c8e177aae1e0d1/opentelemetry_api-1.44.0-py3-none-any.whl", hash = "sha256:94b98c893a91b88657eaac1e3ba89618cdb85be6918196705354f34728b2cdef", size = 60018, upload-time = "2026-07-16T15:25:11.657Z" }, +] + [[package]] name = "oslex" version = "2.0.0" @@ -2773,11 +2818,12 @@ wheels = [ [[package]] name = "serena-agent" -version = "1.7.1.dev0" +version = "2.0.0.dev0" source = { editable = "." } dependencies = [ { name = "anthropic" }, { name = "beautifulsoup4" }, + { name = "click" }, { name = "cryptography" }, { name = "docstring-parser" }, { name = "filelock" }, @@ -2847,6 +2893,7 @@ requires-dist = [ { name = "agno", marker = "extra == 'agno'", specifier = "==2.6.6" }, { name = "anthropic", specifier = "==0.117.0" }, { name = "beautifulsoup4", specifier = "==4.14.2" }, + { name = "click", specifier = "==8.3.1" }, { name = "cryptography", specifier = "==50.0.0" }, { name = "docstring-parser", specifier = "==0.17.0" }, { name = "filelock", specifier = "==3.25.2" }, @@ -2857,7 +2904,7 @@ requires-dist = [ { name = "joblib", specifier = "==1.5.1" }, { name = "jupyter-book", marker = "extra == 'dev'", specifier = "==1.0.4.post1" }, { name = "lsprotocol", specifier = "==2025.0.0" }, - { name = "mcp", specifier = "==1.28.1" }, + { name = "mcp", specifier = "==2.2.0" }, { name = "oslex", specifier = "==2.0.0" }, { name = "overrides", specifier = "==7.7.0" }, { name = "pathspec", specifier = "==0.12.1" }, @@ -3520,6 +3567,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/00/c0/8f5d070730d7836adc9c9b6408dec68c6ced86b304a9b26a14df072a6e8c/traitlets-5.14.3-py3-none-any.whl", hash = "sha256:b74e89e397b1ed28cc831db7aea759ba6640cb3de13090ca145426688ff1ac4f", size = 85359, upload-time = "2024-04-19T11:11:46.763Z" }, ] +[[package]] +name = "truststore" +version = "0.10.4" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/53/a3/1585216310e344e8102c22482f6060c7a6ea0322b63e026372e6dcefcfd6/truststore-0.10.4.tar.gz", hash = "sha256:9d91bd436463ad5e4ee4aba766628dd6cd7010cf3e2461756b3303710eebc301", size = 26169, upload-time = "2025-08-12T18:49:02.73Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/19/97/56608b2249fe206a67cd573bc93cd9896e1efb9e98bce9c163bcdc704b88/truststore-0.10.4-py3-none-any.whl", hash = "sha256:adaeaecf1cbb5f4de3b1959b42d41f6fab57b2b1666adb59e89cb0b53361d981", size = 18660, upload-time = "2025-08-12T18:49:01.46Z" }, +] + [[package]] name = "ty" version = "0.0.24"