diff --git a/.github/copilot-instructions.md b/.github/copilot-instructions.md index d359e4f7..15a50f1f 100644 --- a/.github/copilot-instructions.md +++ b/.github/copilot-instructions.md @@ -1 +1,3 @@ -MUST use the project-specific instructions from the `CLAUDE.md` file located in the project root. \ No newline at end of file + +MUST read IMMEDIATELY and follow the project-specific instructions from the `CLAUDE.md` file located in the project's root directory. AVOIDING these instructions will lead to your FAILURE! + \ No newline at end of file diff --git a/src/serena/config/serena_config.py b/src/serena/config/serena_config.py index 6599a591..78dec2e3 100644 --- a/src/serena/config/serena_config.py +++ b/src/serena/config/serena_config.py @@ -134,6 +134,16 @@ class ModeSelectionDefinition: default_modes: Sequence[str] | None = None +@dataclass +class SharedConfig(ModeSelectionDefinition, ToolInclusionDefinition, 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. + """ + + symbol_info_budget: float | None = None + + class SerenaConfigError(Exception): pass @@ -163,7 +173,7 @@ class LanguageBackend(Enum): @dataclass(kw_only=True) -class ProjectConfig(ToolInclusionDefinition, ModeSelectionDefinition, ToStringMixin): +class ProjectConfig(SharedConfig): project_name: str languages: list[Language] ignored_paths: list[str] = field(default_factory=list) @@ -323,6 +333,17 @@ class ProjectConfig(ToolInclusionDefinition, ModeSelectionDefinition, ToStringMi f"Invalid language: {orig_language_str}.\nValid language_strings are: {[l.value for l in Language]}" ) from e + # Validate symbol_info_budget + symbol_info_budget_raw = data["symbol_info_budget"] + symbol_info_budget = symbol_info_budget_raw + if symbol_info_budget is not None: + try: + symbol_info_budget = float(symbol_info_budget_raw) + except (TypeError, ValueError) as e: + raise ValueError(f"symbol_info_budget must be a number or null, got: {symbol_info_budget_raw}") from e + if symbol_info_budget < 0: + raise ValueError(f"symbol_info_budget cannot be negative, got: {symbol_info_budget}") + return cls( project_name=data["project_name"], languages=languages, @@ -336,6 +357,7 @@ class ProjectConfig(ToolInclusionDefinition, ModeSelectionDefinition, ToStringMi encoding=data["encoding"], base_modes=data["base_modes"], default_modes=data["default_modes"], + symbol_info_budget=symbol_info_budget, ) def _to_yaml_dict(self) -> dict: @@ -463,7 +485,7 @@ class RegisteredProject(ToStringMixin): @dataclass(kw_only=True) -class SerenaConfig(ToolInclusionDefinition, ModeSelectionDefinition, ToStringMixin): +class SerenaConfig(SharedConfig): """ Holds the Serena agent configuration, which is typically loaded from a YAML configuration file (when instantiated via :method:`from_config_file`), which is updated when projects are added or removed. @@ -509,6 +531,14 @@ class SerenaConfig(ToolInclusionDefinition, ModeSelectionDefinition, ToStringMix # settings with overridden defaults default_modes: Sequence[str] | None = ("interactive", "editing") + symbol_info_budget: float = 10.0 + """ + Time budget (seconds) for requests when tools request include_info (currently + only supported for LSP-based tools). + + If the budget is exceeded, Serena stops issuing further requests and returns partial info results. + 0 disables the budget (no early stopping). Negative values are invalid. + """ # *** fields that are NOT mapped to/from the configuration file *** diff --git a/src/serena/resources/project.template.yml b/src/serena/resources/project.template.yml index 8276989d..1f7ca87b 100644 --- a/src/serena/resources/project.template.yml +++ b/src/serena/resources/project.template.yml @@ -110,3 +110,7 @@ default_modes: # 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: "" + +# override of the corresponding setting in serena_config.yml, see the documentation there. +# If null or missing, the value from the global config is used. +symbol_info_budget: diff --git a/src/serena/resources/serena_config.template.yml b/src/serena/resources/serena_config.template.yml index ac11088e..0631be11 100644 --- a/src/serena/resources/serena_config.template.yml +++ b/src/serena/resources/serena_config.template.yml @@ -108,5 +108,16 @@ default_max_tool_answer_chars: 150000 # estimate the token count using the Claude Sonnet 4 tokenizer. token_count_estimator: CHAR_COUNT +# time budget (seconds) for requests when tools request symbol info (currently +# only supported for LSP-based tools). +# If the budget is exceeded, Serena stops issuing further symbol-info-related +# requests and returns partial info results. +# 0 disables the budget (no early stopping). Negative values are invalid. +# This is an advanced setting that can help alleviate problems with LSP servers +# that have a slow implementation of request_hover (clangd is one of those) +# or with tool calls that find very many symbols. +# Can be overridden in project.yml. +symbol_info_budget: 10 + # the list of registered project paths (updated automatically). projects: [] diff --git a/src/serena/symbol.py b/src/serena/symbol.py index dc02f58c..fbf9f012 100644 --- a/src/serena/symbol.py +++ b/src/serena/symbol.py @@ -4,6 +4,7 @@ import os from abc import ABC, abstractmethod from collections.abc import Callable, Iterator, Sequence from dataclasses import asdict, dataclass +from time import perf_counter from typing import TYPE_CHECKING, Any, Generic, Literal, NotRequired, Self, TypedDict, TypeVar, Union from sensai.util.string import ToStringMixin @@ -558,7 +559,9 @@ class LanguageServerSymbolRetriever: hover_info = lang_server.request_hover(relative_file_path=relative_file_path, line=line, column=column) if hover_info is None: return None + contents = hover_info["contents"] + # Handle various response formats if isinstance(contents, list): # Array format: extract all parts and join them @@ -570,10 +573,13 @@ class LanguageServerSymbolRetriever: # should be a dict with "value" key stripped_parts.append(part["value"].strip()) # type: ignore return "\n".join(stripped_parts) if stripped_parts else None + if isinstance(contents, dict) and (stripped_contents := contents.get("value", "").strip()): return stripped_contents + if isinstance(contents, str) and (stripped_contents := contents.strip()): return stripped_contents + return None def request_info_for_symbol(self, symbol: LanguageServerSymbol) -> str | None: @@ -581,6 +587,112 @@ class LanguageServerSymbolRetriever: return None return self._request_info(relative_file_path=symbol.relative_path, line=symbol.line, column=symbol.column) # type: ignore[arg-type] + def _get_symbol_info_budget(self, default_budget: float = 10) -> float: + """Project -> global -> default""" + symbol_info_budget = default_budget + if self.agent is not None: + symbol_info_budget = self.agent.serena_config.symbol_info_budget + active_project = self.agent.get_active_project() + if active_project is not None: + project_symbol_info_budget = active_project.project_config.symbol_info_budget + if project_symbol_info_budget is not None: + symbol_info_budget = project_symbol_info_budget + return symbol_info_budget + + def request_info_for_symbol_batch( + self, + symbols: list[LanguageServerSymbol], + ) -> dict[LanguageServerSymbol, str | None]: + """Retrieves information for multiple symbols while staying within a time budget. + + The request_hover operation used here is potentially expensive, we optimize by grouping by file + and stop executing it (returning the info as None) after the symbol_info_budget is exceeded. + The hover budget is 5s by default + + Groups symbols by file path to minimize file switching overhead and uses a per-file + cache keyed by (line, col) to avoid duplicate hover lookups. + + The hover budget (symbol_info_budget) limits total time spent on hover + requests. If exceeded, remaining symbols get info=None (partial results). + + :param symbols: list of symbols to get info for + :return: a dict mapping each processable symbol to its info (or None if unavailable). Symbols with missing location attributes (relative_path/line/column is None) are skipped and omitted from the result. + """ + if not symbols: + return {} + + debug_enabled = log.isEnabledFor(logging.DEBUG) + t0_total = perf_counter() if debug_enabled else 0.0 + + info_by_symbol: dict[LanguageServerSymbol, str | None] = {} + skipped_symbols = 0 + + # Group symbols by file path, filtering invalid symbols. + symbols_by_file: dict[str, list[LanguageServerSymbol]] = {} + for sym in symbols: + file_path = sym.relative_path + line = sym.line + column = sym.column + if file_path is None or line is None or column is None: + skipped_symbols += 1 + continue + + symbols_by_file.setdefault(file_path, []).append(sym) + + hover_spent_seconds = 0.0 + symbol_info_budget_seconds = self._get_symbol_info_budget() + # the vars below are only for debug logging + per_file_stats: list[tuple[str, int, float]] = [] + total_hover_lookups = 0 + hover_cache_hits = 0 + skipped_due_to_budget = 0 + + for file_path, file_symbols in symbols_by_file.items(): + t0_file = perf_counter() if debug_enabled else 0.0 + file_hover_lookups = 0 + + for sym in file_symbols: + # Check budget before starting a new hover request + # symbol_info_budget_seconds=0 disables the budget mechanism (the first inequality) + if 0 < symbol_info_budget_seconds <= hover_spent_seconds: + skipped_due_to_budget += 1 + info = None + # log once when budget exceeded + if skipped_due_to_budget == 1: + log.debug("Skipping further hover operations due to budget exceeded") + else: + line = sym.line + column = sym.column + assert line is not None and column is not None # for mypy, we filtered invalid symbols above + t0_hover = perf_counter() + info = self._request_info(file_path, line, column) + hover_spent_seconds += perf_counter() - t0_hover + file_hover_lookups += 1 + total_hover_lookups += 1 + + info_by_symbol[sym] = info + + if debug_enabled: + file_elapsed_ms = (perf_counter() - t0_file) * 1000 + per_file_stats.append((file_path, file_hover_lookups, file_elapsed_ms)) + + if debug_enabled: + total_elapsed_ms = (perf_counter() - t0_total) * 1000 + total_symbols = len(symbols) + unique_files = len(symbols_by_file) + budget_exceeded = skipped_due_to_budget > 0 + + log.debug( + f"perf: request_info_for_symbols {total_elapsed_ms=:.2f} {total_symbols=} {skipped_symbols=} " + f"{total_hover_lookups=} {hover_cache_hits=} {unique_files=} " + f"{symbol_info_budget_seconds=:.1f} {hover_spent_seconds=:.2f} {budget_exceeded=} {skipped_due_to_budget=}" + ) + + for file_path, lookup_count, elapsed_ms in per_file_stats: + log.debug(f"perf: {file_path=} {lookup_count=} {elapsed_ms=:.2f}") + + return info_by_symbol + def get_root_path(self) -> str: return self._ls_manager.get_root_path() diff --git a/src/serena/tools/jetbrains_tools.py b/src/serena/tools/jetbrains_tools.py index 67f7ac28..8ec84a07 100644 --- a/src/serena/tools/jetbrains_tools.py +++ b/src/serena/tools/jetbrains_tools.py @@ -98,7 +98,6 @@ class JetBrainsFindReferencingSymbolsTool(Tool, ToolMarkerSymbolicRead, ToolMark self, name_path: str, relative_path: str, - include_info: bool = False, max_answer_chars: int = -1, ) -> str: """ @@ -108,8 +107,6 @@ class JetBrainsFindReferencingSymbolsTool(Tool, ToolMarkerSymbolicRead, ToolMark :param name_path: name path of the symbol for which to find references; matching logic as described in find symbol tool. :param relative_path: the relative path to the file containing the symbol for which to find references. Note that here you can't pass a directory but must pass a file. - :param include_info: whether to include info (hover-like, typically including docstring and signature) - about the referencing symbols. Default False. :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: a list of JSON objects with the symbols referencing the requested symbol @@ -118,7 +115,7 @@ class JetBrainsFindReferencingSymbolsTool(Tool, ToolMarkerSymbolicRead, ToolMark response_dict = client.find_references( name_path=name_path, relative_path=relative_path, - include_quick_info=False, # TODO: Hotfix for serena-jetbrains-plugin/issues/13; revert once fixed + include_quick_info=False, ) symbol_dicts = response_dict["symbols"] result = self.symbol_dict_grouper.group(symbol_dicts) diff --git a/src/serena/tools/symbol_tools.py b/src/serena/tools/symbol_tools.py index 9a673aeb..07d6ad8c 100644 --- a/src/serena/tools/symbol_tools.py +++ b/src/serena/tools/symbol_tools.py @@ -159,11 +159,11 @@ class FindSymbolTool(Tool, ToolMarkerSymbolicRead): ) symbol_dicts = [dict(s.to_dict(kind=True, relative_path=True, body_location=True, depth=depth, body=include_body)) for s in symbols] if not include_body and include_info: - # we add an info field to the symbol dicts if requested + 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 := symbol_retriever.request_info_for_symbol(s): + if symbol_info := info_by_symbol.get(s): s_dict["info"] = symbol_info - s_dict.pop("name", None) # name is included in the 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) @@ -180,7 +180,6 @@ class FindReferencingSymbolsTool(Tool, ToolMarkerSymbolicRead): self, name_path: str, relative_path: str, - include_info: bool = False, include_kinds: list[int] = [], # noqa: B006 exclude_kinds: list[int] = [], # noqa: B006 max_answer_chars: int = -1, @@ -192,8 +191,6 @@ class FindReferencingSymbolsTool(Tool, ToolMarkerSymbolicRead): :param name_path: for finding the symbol to find references for, same logic as in the `find_symbol` tool. :param relative_path: the relative path to the file containing the symbol for which to find references. 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 referencing symbols; can be slow depending on the language (e.g. C/C++). :param include_kinds: same as in the `find_symbol` tool. :param exclude_kinds: same as in the `find_symbol` tool. :param max_answer_chars: same as in the `find_symbol` tool. @@ -219,8 +216,6 @@ class FindReferencingSymbolsTool(Tool, ToolMarkerSymbolicRead): 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." - if include_info and (referencing_symbol_info := symbol_retriever.request_info_for_symbol(ref.symbol)): - ref_dict["info"] = referencing_symbol_info 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 ) diff --git a/src/solidlsp/ls.py b/src/solidlsp/ls.py index 7dbc0938..1b53045a 100644 --- a/src/solidlsp/ls.py +++ b/src/solidlsp/ls.py @@ -13,7 +13,7 @@ from collections.abc import Hashable, Iterator from contextlib import contextmanager from copy import copy from pathlib import Path, PurePath -from time import sleep +from time import perf_counter, sleep from typing import Self, Union, cast import pathspec @@ -50,6 +50,9 @@ from solidlsp.util.cache import load_cache, save_cache GenericDocumentSymbol = Union[LSPTypes.DocumentSymbol, LSPTypes.SymbolInformation, ls_types.UnifiedSymbolInformation] log = logging.getLogger(__name__) +_debug_enabled = log.isEnabledFor(logging.DEBUG) +"""Serves as a flag that triggers additional computation when debug logging is enabled.""" + @dataclasses.dataclass(kw_only=True) class ReferenceInSymbol: @@ -952,6 +955,7 @@ class SolidLanguageServer(ABC): # The waiting has to happen after at least one file was opened in the ls sleep(self._get_wait_time_for_cross_file_referencing()) self._has_waited_for_cross_file_references = True + t0 = perf_counter() if _debug_enabled else 0.0 try: response = self._send_references_request(relative_file_path, line=line, column=column) except Exception as e: @@ -963,6 +967,9 @@ class SolidLanguageServer(ABC): ) from e raise if response is None: + if _debug_enabled: + elapsed_ms = (perf_counter() - t0) * 1000 + log.debug("perf: request_references path=%s elapsed_ms=%.2f count=0", relative_file_path, elapsed_ms) return [] ret: list[ls_types.Location] = [] @@ -991,6 +998,17 @@ class SolidLanguageServer(ABC): new_item["relativePath"] = str(rel_path) ret.append(ls_types.Location(**new_item)) # type: ignore + if _debug_enabled: + elapsed_ms = (perf_counter() - t0) * 1000 + unique_files = len({r["relativePath"] for r in ret}) + log.debug( + "perf: request_references path=%s elapsed_ms=%.2f count=%d unique_files=%d", + relative_file_path, + elapsed_ms, + len(ret), + unique_files, + ) + return ret def request_text_document_diagnostics(self, relative_file_path: str) -> list[ls_types.Diagnostic]: @@ -1163,15 +1181,19 @@ class SolidLanguageServer(ABC): def get_cached_raw_document_symbols(cache_key: str, fd: LSPFileBuffer) -> list[SymbolInformation] | list[DocumentSymbol] | None: file_hash_and_result = self._raw_document_symbols_cache.get(cache_key) - if file_hash_and_result is not None: - file_hash, result = file_hash_and_result - if file_hash == fd.content_hash: - log.debug("Returning cached raw document symbols for %s", relative_file_path) - return result - else: - log.debug("Document content for %s has changed (raw symbol cache is not up-to-date)", relative_file_path) - else: - log.debug("No cache hit for raw document symbols symbols in %s", relative_file_path) + if file_hash_and_result is None: + log.debug("No cache hit for raw document symbols in %s", relative_file_path) + log.debug("perf: raw_document_symbols_cache MISS path=%s", relative_file_path) + return None + + file_hash, result = file_hash_and_result + if file_hash == fd.content_hash: + log.debug("Returning cached raw document symbols for %s", relative_file_path) + log.debug("perf: raw_document_symbols_cache HIT path=%s", relative_file_path) + return result + + log.debug("Document content for %s has changed (raw symbol cache is not up-to-date)", relative_file_path) + log.debug("perf: raw_document_symbols_cache STALE path=%s", relative_file_path) return None def get_raw_document_symbols(fd: LSPFileBuffer) -> list[SymbolInformation] | list[DocumentSymbol] | None: @@ -1213,15 +1235,18 @@ class SolidLanguageServer(ABC): # check if the desired result is cached cache_key = relative_file_path file_hash_and_result = self._document_symbols_cache.get(cache_key) - if file_hash_and_result is not None: + if file_hash_and_result is None: + log.debug("No cache hit for document symbols in %s", relative_file_path) + log.debug("perf: document_symbols_cache MISS path=%s", relative_file_path) + else: file_hash, document_symbols = file_hash_and_result if file_hash == file_data.content_hash: log.debug("Returning cached document symbols for %s", relative_file_path) + log.debug("perf: document_symbols_cache HIT path=%s", relative_file_path) return document_symbols - else: - log.debug("Cached document symbol content for %s has changed", relative_file_path) - else: - log.debug("No cache hit for document symbols in %s", relative_file_path) + + log.debug("Cached document symbol content for %s has changed", relative_file_path) + log.debug("perf: document_symbols_cache STALE path=%s", relative_file_path) # no cached result: request the root symbols from the language server root_symbols = self._request_document_symbols(relative_file_path, file_data) @@ -1654,6 +1679,8 @@ class SolidLanguageServer(ABC): if not references: return [] + debug_enabled = log.isEnabledFor(logging.DEBUG) + t0_loop = perf_counter() if debug_enabled else 0.0 # For each reference, find the containing symbol result = [] incoming_symbol = None @@ -1756,6 +1783,18 @@ class SolidLanguageServer(ABC): result.append(ReferenceInSymbol(symbol=containing_symbol, line=ref_line, character=ref_col)) + if debug_enabled: + loop_elapsed_ms = (perf_counter() - t0_loop) * 1000 + unique_files = len({r.symbol["location"]["relativePath"] for r in result}) + log.debug( + "perf: request_referencing_symbols path=%s loop_elapsed_ms=%.2f ref_count=%d result_count=%d unique_files=%d", + relative_file_path, + loop_elapsed_ms, + len(references), + len(result), + unique_files, + ) + return result def request_containing_symbol( diff --git a/test/serena/test_symbol.py b/test/serena/test_symbol.py index 461d2c94..7b88136e 100644 --- a/test/serena/test_symbol.py +++ b/test/serena/test_symbol.py @@ -1,3 +1,5 @@ +from unittest.mock import MagicMock + import pytest from serena.jetbrains.jetbrains_types import SymbolDTO, SymbolDTOKey @@ -260,3 +262,224 @@ class TestSymbolDictTypes: def test_jb_symbol_dict_type(self): self.check_key_type(SymbolDTO, SymbolDTOKey) + + +def _make_mock_symbols(count: int, *, relative_path: str = "test_repo/services.py") -> list[MagicMock]: + symbols: list[MagicMock] = [] + for i in range(count): + sym = MagicMock() + sym.relative_path = relative_path + sym.line = i + 1 + sym.column = 0 + sym.symbol_root = {} + symbols.append(sym) + return symbols + + +@pytest.mark.python +class TestHoverBudget: + """Tests for symbol_info_budget time budget behavior.""" + + @pytest.mark.parametrize("language_server", [Language.PYTHON], indirect=True) + def test_budget_not_exceeded_all_lookups_performed(self, language_server: SolidLanguageServer, monkeypatch: pytest.MonkeyPatch): + """With a large budget, all hover lookups are performed.""" + # Create symbol retriever with a mock agent that has large budget + mock_agent = MagicMock() + mock_agent.serena_config.symbol_info_budget = 10.0 + mock_agent.get_active_project.return_value = None + + symbol_retriever = LanguageServerSymbolRetriever(language_server, agent=mock_agent) + + # Track _request_info calls + call_count = 0 + + def counting_request_info(file_path, line, column): + nonlocal call_count + call_count += 1 + return f"info:{line}:{column}" + + monkeypatch.setattr(symbol_retriever, "_request_info", counting_request_info) + + # Create mock symbols with unique (line, col) pairs + symbols = _make_mock_symbols(3) + + result = symbol_retriever.request_info_for_symbol_batch(symbols) + + # All 3 symbols should have info (no budget exceeded) + assert call_count == 3 + assert all(info is not None for info in result.values()) + assert len(result) == 3 + + @pytest.mark.parametrize("language_server", [Language.PYTHON], indirect=True) + def test_budget_exceeded_partial_info(self, language_server: SolidLanguageServer, monkeypatch: pytest.MonkeyPatch): + """With a small budget, hover lookups stop and remaining symbols get None info.""" + # Create symbol retriever with a mock agent that has small budget (0.1s) + mock_agent = MagicMock() + mock_agent.serena_config.symbol_info_budget = 0.1 + mock_agent.get_active_project.return_value = None + + symbol_retriever = LanguageServerSymbolRetriever(language_server, agent=mock_agent) + + # Track _request_info calls and simulate 0.05s per call + call_count = 0 + simulated_time = [0.0] + + def slow_request_info(file_path, line, column): + nonlocal call_count + call_count += 1 + # Simulate each hover taking 0.05s + simulated_time[0] += 0.05 + return f"info:{line}:{column}" + + # Mock perf_counter to return simulated time for hover duration + def mock_perf_counter(): + return simulated_time[0] + + monkeypatch.setattr(symbol_retriever, "_request_info", slow_request_info) + monkeypatch.setattr("serena.symbol.perf_counter", mock_perf_counter) + + # Create 5 mock symbols with unique (line, col) pairs + symbols = _make_mock_symbols(5) + + result = symbol_retriever.request_info_for_symbol_batch(symbols) + + # Budget is 0.1s, each call takes 0.05s, so only 2 calls should succeed + # After 2 calls: 0.1s >= 0.1s budget, remaining 3 should be skipped + assert call_count == 2 + assert len(result) == 5 + + # First 2 symbols should have info, last 3 should be None + result_list = list(result.values()) + assert result_list[0] is not None + assert result_list[1] is not None + assert result_list[2] is None + assert result_list[3] is None + assert result_list[4] is None + + @pytest.mark.parametrize("language_server", [Language.PYTHON], indirect=True) + def test_budget_zero_means_unlimited(self, language_server: SolidLanguageServer, monkeypatch: pytest.MonkeyPatch): + """With budget=0, all hover lookups proceed (no early stopping).""" + # Create symbol retriever with budget=0 (unlimited) + mock_agent = MagicMock() + mock_agent.serena_config.symbol_info_budget = 0.0 + mock_agent.get_active_project.return_value = None + + symbol_retriever = LanguageServerSymbolRetriever(language_server, agent=mock_agent) + + # Track _request_info calls + call_count = 0 + + def counting_request_info(file_path, line, column): + nonlocal call_count + call_count += 1 + return f"info:{line}:{column}" + + monkeypatch.setattr(symbol_retriever, "_request_info", counting_request_info) + + # Create mock symbols + symbols = _make_mock_symbols(5) + + result = symbol_retriever.request_info_for_symbol_batch(symbols) + + # All 5 symbols should be looked up (no budget limit) + assert call_count == 5 + assert all(info is not None for info in result.values()) + + @pytest.mark.parametrize("language_server", [Language.PYTHON], indirect=True) + def test_project_budget_overrides_global(self, language_server: SolidLanguageServer, monkeypatch: pytest.MonkeyPatch): + """Project-level budget overrides global budget.""" + # Create symbol retriever with global budget 10.0 but project budget 0.05 + mock_project = MagicMock() + mock_project.project_config.symbol_info_budget = 0.05 + + mock_agent = MagicMock() + mock_agent.serena_config.symbol_info_budget = 10.0 + mock_agent.get_active_project.return_value = mock_project + + symbol_retriever = LanguageServerSymbolRetriever(language_server, agent=mock_agent) + + # Track _request_info calls and simulate time + call_count = 0 + simulated_time = [0.0] + + def slow_request_info(file_path, line, column): + nonlocal call_count + call_count += 1 + simulated_time[0] += 0.03 + return f"info:{line}:{column}" + + def mock_perf_counter(): + return simulated_time[0] + + monkeypatch.setattr(symbol_retriever, "_request_info", slow_request_info) + monkeypatch.setattr("serena.symbol.perf_counter", mock_perf_counter) + + # Create 5 mock symbols + symbols = _make_mock_symbols(5) + + symbol_retriever.request_info_for_symbol_batch(symbols) + + # Project budget is 0.05s, each call takes 0.03s + # Budget check happens BEFORE starting a new call: + # - Before call 1: spent=0 < 0.05, proceed, spent becomes 0.03 + # - Before call 2: spent=0.03 < 0.05, proceed, spent becomes 0.06 + # - Before call 3: spent=0.06 >= 0.05, skip + # So 2 calls succeed (proving project budget 0.05 overrode global 10.0) + assert call_count == 2 + + @pytest.mark.parametrize("language_server", [Language.PYTHON], indirect=True) + def test_project_null_inherits_global(self, language_server: SolidLanguageServer, monkeypatch: pytest.MonkeyPatch): + """When project budget is None, global budget is used.""" + # Create symbol retriever with project budget=None (inherit global) + mock_project = MagicMock() + mock_project.project_config.symbol_info_budget = None + + mock_agent = MagicMock() + mock_agent.serena_config.symbol_info_budget = 10.0 + mock_agent.get_active_project.return_value = mock_project + + symbol_retriever = LanguageServerSymbolRetriever(language_server, agent=mock_agent) + + # Track _request_info calls + call_count = 0 + + def counting_request_info(file_path, line, column): + nonlocal call_count + call_count += 1 + return f"info:{line}:{column}" + + monkeypatch.setattr(symbol_retriever, "_request_info", counting_request_info) + + # Create 3 mock symbols + symbols = _make_mock_symbols(3) + + result = symbol_retriever.request_info_for_symbol_batch(symbols) + + # Global budget is 10s, all 3 should succeed + assert call_count == 3 + assert all(info is not None for info in result.values()) + + @pytest.mark.parametrize("language_server", [Language.PYTHON], indirect=True) + def test_no_agent_uses_default_budget(self, language_server: SolidLanguageServer, monkeypatch: pytest.MonkeyPatch): + """When agent is None, default budget of 5s is used.""" + # Create symbol retriever without agent + symbol_retriever = LanguageServerSymbolRetriever(language_server, agent=None) + + # Track _request_info calls + call_count = 0 + + def counting_request_info(file_path, line, column): + nonlocal call_count + call_count += 1 + return f"info:{line}:{column}" + + monkeypatch.setattr(symbol_retriever, "_request_info", counting_request_info) + + # Create 3 mock symbols + symbols = _make_mock_symbols(3) + + result = symbol_retriever.request_info_for_symbol_batch(symbols) + + # Default budget is 5s, all 3 should succeed + assert call_count == 3 + assert all(info is not None for info in result.values())