diff --git a/src/serena/agent.py b/src/serena/agent.py index 6b16a9ab..32070cfc 100644 --- a/src/serena/agent.py +++ b/src/serena/agent.py @@ -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 diff --git a/src/serena/symbol.py b/src/serena/symbol.py index 7f75098b..e4c49975 100644 --- a/src/serena/symbol.py +++ b/src/serena/symbol.py @@ -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: """