Merge pull request #956 from n-r-w/hover_budget

This commit is contained in:
Michael Panchenko authored and GitHub committed 2026-02-20 17:44:23 +01:00
commit dcbf08520d
9 files changed
+443 -30

No files matched your search

+3 -1
View File
@@ -1 +1,3 @@
MUST use the project-specific instructions from the `CLAUDE.md` file located in the project root.
<non_negotiable critical="true">
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!
</non_negotiable>
+32 -2
View File
@@ -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 ***
@@ -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:
@@ -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: []
+112
View File
@@ -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()
+1 -4
View File
@@ -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)
+3 -8
View File
@@ -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
)
+54 -15
View File
@@ -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(
+223
View File
@@ -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())