mirror of
https://github.com/tiennm99/serena.git
synced 2026-10-11 03:13:51 +00:00
Merge pull request #956 from n-r-w/hover_budget
This commit is contained in:
9 files changed
+443
-30
No files matched your search
@@ -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>
|
||||
@@ -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: []
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
@@ -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(
|
||||
|
||||
@@ -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())
|
||||
Reference in new issue
Block a user