Bugfix: reset of lang server also replaces the reference in SymbolManager

This commit is contained in:
Michael Panchenko committed 2025-06-16 23:45:38 +02:00
1 parent 47b50198ec
commit 6d54024966
2 files changed
+26 -15

No files matched your search

+4
View File
@@ -1056,11 +1056,15 @@ class SerenaAgent:
ls_timeout=ls_timeout,
trace_lsp_communication=self.serena_config.trace_lsp_communication,
)
log.info(f"Starting the language server for {self._active_project.project_name}")
self.language_server.start()
if not self.language_server.is_running():
raise RuntimeError(
f"Failed to start the language server for {self._active_project.project_name} at {self._active_project.project_root}"
)
assert self.symbol_manager is not None, "Should never be None with an active project"
log.debug("Setting the language server in the agent's symbol manager")
self.symbol_manager.set_language_server(self.language_server)
def get_tool(self, tool_class: type[TTool]) -> TTool:
return self._all_tools[tool_class] # type: ignore
+22 -15
View File
@@ -494,9 +494,16 @@ class SymbolManager:
:param agent: the agent to use (only needed for marking files as modified). You can pass None if you don't
need an agent to be avare of file modifications performed by the symbol manager.
"""
self.lang_server = lang_server
self._lang_server = lang_server
self.agent = agent
def set_language_server(self, lang_server: SyncLanguageServer) -> None:
"""
Set the language server to use for symbol retrieval and editing operations.
This is useful if you want to change the language server after initializing the SymbolManager.
"""
self._lang_server = lang_server
def find_by_name(
self,
name_path: str,
@@ -512,7 +519,7 @@ class SymbolManager:
to symbols within a specific file or directory.
"""
symbols: list[Symbol] = []
symbol_roots = self.lang_server.request_full_symbol_tree(within_relative_path=within_relative_path, include_body=include_body)
symbol_roots = self._lang_server.request_full_symbol_tree(within_relative_path=within_relative_path, include_body=include_body)
for root in symbol_roots:
symbols.extend(
Symbol(root).find(
@@ -522,14 +529,14 @@ class SymbolManager:
return symbols
def get_document_symbols(self, relative_path: str) -> list[Symbol]:
symbol_dicts, roots = self.lang_server.request_document_symbols(relative_path, include_body=False)
symbol_dicts, roots = self._lang_server.request_document_symbols(relative_path, include_body=False)
symbols = [Symbol(s) for s in symbol_dicts]
return symbols
def find_by_location(self, location: SymbolLocation) -> Symbol | None:
if location.relative_path is None:
return None
symbol_dicts, roots = self.lang_server.request_document_symbols(location.relative_path, include_body=False)
symbol_dicts, roots = self._lang_server.request_document_symbols(location.relative_path, include_body=False)
for symbol_dict in symbol_dicts:
symbol = Symbol(symbol_dict)
if symbol.location == location:
@@ -598,7 +605,7 @@ class SymbolManager:
assert symbol_location.relative_path is not None
assert symbol_location.line is not None
assert symbol_location.column is not None
references = self.lang_server.request_referencing_symbols(
references = self._lang_server.request_referencing_symbols(
relative_file_path=symbol_location.relative_path,
line=symbol_location.line,
column=symbol_location.column,
@@ -618,9 +625,9 @@ class SymbolManager:
@contextmanager
def _edited_file(self, relative_path: str) -> Iterator[None]:
with self.lang_server.open_file(relative_path) as file_buffer:
with self._lang_server.open_file(relative_path) as file_buffer:
yield
root_path = self.lang_server.language_server.repository_root_path
root_path = self._lang_server.language_server.repository_root_path
abs_path = os.path.join(root_path, relative_path)
with open(abs_path, "w", encoding="utf-8") as f:
f.write(file_buffer.contents)
@@ -641,7 +648,7 @@ class SymbolManager:
def _get_code_file_content(self, relative_path: str) -> str:
"""Get the content of a file using the language server."""
return self.lang_server.language_server.retrieve_full_file_content(relative_path)
return self._lang_server.language_server.retrieve_full_file_content(relative_path)
def replace_body(self, name_path: str, relative_file_path: str, body: str, *, use_same_indentation: bool = True) -> None:
"""
@@ -690,8 +697,8 @@ class SymbolManager:
# make sure body always ends with at least one newline
if not body.endswith("\n"):
body += "\n"
self.lang_server.delete_text_between_positions(location.relative_path, start_pos, end_pos)
self.lang_server.insert_text_at_position(location.relative_path, start_line, start_col, body)
self._lang_server.delete_text_between_positions(location.relative_path, start_pos, end_pos)
self._lang_server.insert_text_at_position(location.relative_path, start_line, start_col, body)
def insert_after_symbol(
self,
@@ -774,7 +781,7 @@ class SymbolManager:
col = 0
with self._edited_symbol_location(location):
self.lang_server.insert_text_at_position(location.relative_path, line=line, column=col, text_to_be_inserted=body)
self._lang_server.insert_text_at_position(location.relative_path, line=line, column=col, text_to_be_inserted=body)
def insert_before_symbol(
self,
@@ -827,7 +834,7 @@ class SymbolManager:
body += "\n"
assert location.relative_path is not None
self.lang_server.insert_text_at_position(location.relative_path, line=line, column=col, text_to_be_inserted=body)
self._lang_server.insert_text_at_position(location.relative_path, line=line, column=col, text_to_be_inserted=body)
def insert_at_line(self, relative_path: str, line: int, content: str) -> None:
"""
@@ -837,7 +844,7 @@ class SymbolManager:
:param content: the content to insert
"""
with self._edited_file(relative_path):
self.lang_server.insert_text_at_position(relative_path, line, 0, content)
self._lang_server.insert_text_at_position(relative_path, line, 0, content)
def delete_lines(self, relative_path: str, start_line: int, end_line: int) -> None:
"""
@@ -852,7 +859,7 @@ class SymbolManager:
with self._edited_file(relative_path):
start_pos = Position(line=start_line, character=start_col)
end_pos = Position(line=end_line_for_delete, character=end_col)
self.lang_server.delete_text_between_positions(relative_path, start_pos, end_pos)
self._lang_server.delete_text_between_positions(relative_path, start_pos, end_pos)
def delete_symbol_at_location(self, location: SymbolLocation) -> None:
"""
@@ -862,7 +869,7 @@ class SymbolManager:
assert location.relative_path is not None
assert symbol.body_start_position is not None
assert symbol.body_end_position is not None
self.lang_server.delete_text_between_positions(location.relative_path, symbol.body_start_position, symbol.body_end_position)
self._lang_server.delete_text_between_positions(location.relative_path, symbol.body_start_position, symbol.body_end_position)
def delete_symbol(self, name_path: str, relative_file_path: str) -> None:
"""