Added retrieval of line and full document

This commit is contained in:
Michael Panchenko committed 2025-03-24 23:07:09 +01:00
1 parent 6510329388
commit 31639b6e6e
7 files changed
+308 -139

No files matched your search

+67 -1
View File
@@ -11,12 +11,13 @@ import dataclasses
import hashlib
import json
import pickle
import time
import logging
import os
import pathlib
import threading
from contextlib import asynccontextmanager, contextmanager
from serena.text_utils import MatchedConsecutiveLines, LineType, TextLine
from .lsp_protocol_handler.lsp_constants import LSPConstants
from .lsp_protocol_handler import lsp_types as LSPTypes
@@ -34,6 +35,15 @@ from typing import AsyncIterator, Iterator, List, Dict, Optional, Union, Tuple
from .type_helpers import ensure_all_methods_implemented
# Serena dependencies
# We will need to watch out for circular imports, but it's probably better to not
# move all generic util code from serena into multilspy.
# It does however make sense to integrate many text-related utils into the language server
# since it caches (in-memory) file contents, so we can avoid reading from disk.
# Moreover, the way we want to use the language server (for retrieving actual content),
# it makes sense to have more content-related utils directly in it.
@dataclasses.dataclass
class LSPFileBuffer:
"""
@@ -514,7 +524,45 @@ class LanguageServer:
ret.append(multilspy_types.Location(**new_item))
return ret
def retrieve_full_file_content(self, relative_file_path: str) -> str:
"""
Retrieve the full content of the given file.
"""
with self.open_file(relative_file_path) as file_data:
return file_data.contents
def retrieve_content_around_line(self, relative_file_path: str, line: int, context_lines_before: int = 0, context_lines_after: int = 0) -> MatchedConsecutiveLines:
"""
Retrieve the content of the given file around the given line.
:param relative_file_path: The relative path of the file to retrieve the content from
:param line: The line number to retrieve the content around
:param context_lines_before: The number of lines to retrieve before the given line
:param context_lines_after: The number of lines to retrieve after the given line
:return MatchedConsecutiveLines: A container with the desired lines.
"""
with self.open_file(relative_file_path) as file_data:
file_contents = file_data.contents
line_contents = file_contents.split("\n")
start_lineno = max(0, line - context_lines_before)
end_lineno = min(len(line_contents) - 1, line + context_lines_after)
# instantiate TextLines with the write LineType
text_lines: list[TextLine] = []
# before the line
for lineno in range(start_lineno, line):
text_lines.append(TextLine(line_number=lineno, line_content=line_contents[lineno], match_type=LineType.BEFORE_MATCH))
# the line
text_lines.append(TextLine(line_number=line, line_content=line_contents[line], match_type=LineType.MATCH))
# after the line
for lineno in range(line + 1, end_lineno + 1):
text_lines.append(TextLine(line_number=lineno, line_content=line_contents[lineno], match_type=LineType.AFTER_MATCH))
return MatchedConsecutiveLines(lines=text_lines, source_file_path=relative_file_path)
async def request_completions(
self, relative_file_path: str, line: int, column: int, allow_incomplete: bool = False
) -> List[multilspy_types.CompletionItem]:
@@ -1296,6 +1344,24 @@ class SyncLanguageServer:
).result()
return result
def retrieve_full_file_content(self, relative_file_path: str) -> str:
"""
Retrieve the full content of the given file.
"""
return self.language_server.retrieve_full_file_content(relative_file_path)
def retrieve_content_around_line(self, relative_file_path: str, line: int, context_lines_before: int = 0, context_lines_after: int = 0) -> MatchedConsecutiveLines:
"""
Retrieve the content of the given file around the given line.
:param relative_file_path: The relative path of the file to retrieve the content from
:param line: The line number to retrieve the content around
:param context_lines_before: The number of lines to retrieve before the given line
:param context_lines_after: The number of lines to retrieve after the given line
:return MatchedConsecutiveLines: A container with the desired lines.
"""
return self.language_server.retrieve_content_around_line(relative_file_path, line, context_lines_before, context_lines_after)
def start(self) -> "SyncLanguageServer":
"""
Starts the language server process and connects to it. Call shutdown when ready.
+31 -12
View File
@@ -1,5 +1,5 @@
import re
from dataclasses import dataclass
from dataclasses import dataclass, field
from enum import StrEnum
@@ -37,27 +37,46 @@ class TextLine:
@dataclass(kw_only=True)
class TextSearchMatch:
"""Represents a match found in a text file or a string."""
class MatchedConsecutiveLines:
"""Represents a collection of consecutive lines found through some criterion in a text file or a string.
May include lines before, after, and matched.
"""
context_lines: list[TextLine]
lines: list[TextLine]
"""All lines in the context of the match. At least one of them should be of match_type MATCH."""
source_file_path: str | None = None
"""Path to the file where the match was found."""
"""Path to the file where the match was found (Metadata)."""
# set in post-init
lines_before_matched: list[TextLine] = field(default_factory=list)
matched_lines: list[TextLine] = field(default_factory=list)
lines_after_matched: list[TextLine] = field(default_factory=list)
def __post_init__(self):
for line in self.lines:
if line.match_type == LineType.BEFORE_MATCH:
self.lines_before_matched.append(line)
elif line.match_type == LineType.MATCH:
self.matched_lines.append(line)
elif line.match_type == LineType.AFTER_MATCH:
self.lines_after_matched.append(line)
assert len(self.matched_lines) > 0, "At least one matched line is required"
@property
def start_line(self) -> int:
return self.context_lines[0].line_number
return self.lines[0].line_number
@property
def end_line(self) -> int:
return self.context_lines[-1].line_number
return self.lines[-1].line_number
@property
def num_matched_lines(self) -> int:
return len([line for line in self.context_lines if line.match_type == LineType.MATCH])
return len(self.matched_lines)
def to_display_string(self) -> str:
return "\n".join([line.format_line() for line in self.context_lines])
return "\n".join([line.format_line() for line in self.lines])
def search_text(
@@ -68,7 +87,7 @@ def search_text(
context_lines_before: int = 0,
context_lines_after: int = 0,
is_glob: bool = False,
) -> list[TextSearchMatch]:
) -> list[MatchedConsecutiveLines]:
"""
Search for a pattern in text content. Supports both regex and glob-like patterns.
@@ -160,7 +179,7 @@ def search_text(
context_lines.append(TextLine(line_number=line_num, line_content=lines[i], match_type=match_type))
matches.append(TextSearchMatch(context_lines=context_lines, source_file_path=source_file_path))
matches.append(MatchedConsecutiveLines(lines=context_lines, source_file_path=source_file_path))
else:
# Search line by line
for i, line in enumerate(lines):
@@ -183,6 +202,6 @@ def search_text(
context_lines.append(TextLine(line_number=context_line_num, line_content=lines[j], match_type=match_type))
matches.append(TextSearchMatch(context_lines=context_lines, source_file_path=source_file_path))
matches.append(MatchedConsecutiveLines(lines=context_lines, source_file_path=source_file_path))
return matches
+23
View File
@@ -2,6 +2,10 @@ from pathlib import Path
import pytest
from multilspy.language_server import SyncLanguageServer
from multilspy.multilspy_config import Language, MultilspyConfig
from multilspy.multilspy_logger import MultilspyLogger
@pytest.fixture(scope="session")
def resources_dir() -> Path:
@@ -13,3 +17,22 @@ def resources_dir() -> Path:
@pytest.fixture(scope="session")
def repo_path() -> Path:
return Path(__file__).parent / "resources" / "test_repo"
@pytest.fixture(scope="session")
def language_server(repo_path: Path):
"""Create a SyncLanguageServer instance configured to use the test repository."""
config = MultilspyConfig(code_language=Language.PYTHON)
logger = MultilspyLogger()
# Create a language server instance
server = SyncLanguageServer.create(config, logger, str(repo_path))
# Start the server
server.start()
try:
yield server
finally:
# Ensure server is shut down
server.stop()
+167
View File
@@ -0,0 +1,167 @@
"""
Basic integration tests for the language server functionality.
These tests validate the functionality of the language server APIs
like request_references using the test repository.
"""
from pathlib import Path
from multilspy.language_server import SyncLanguageServer
from serena.text_utils import LineType
class TestLanguageServerBasics:
"""Test basic functionality of the language server."""
def test_request_references_user_class(self, language_server: SyncLanguageServer, repo_path: Path):
"""Test request_references on the User class."""
# Get references to the User class in models.py
file_path = str(repo_path / "test_repo" / "models.py")
# Line 31 contains the User class definition
references = language_server.request_references(file_path, 31, 6)
# User class should be referenced in multiple files
assert len(references) > 0
# At least two references should be found (one for the class definition itself)
assert len(references) > 1
def test_request_references_item_class(self, language_server: SyncLanguageServer, repo_path: Path):
"""Test request_references on the Item class."""
# Get references to the Item class in models.py
file_path = str(repo_path / "test_repo" / "models.py")
# Line 56 contains the Item class definition
references = language_server.request_references(file_path, 56, 6)
# Item class should be referenced in multiple places
assert len(references) > 0
# At least one reference should be in services.py (ItemService class)
services_references = [ref for ref in references if "services.py" in ref["uri"]]
assert len(services_references) > 0
def test_request_references_function_parameter(self, language_server: SyncLanguageServer, repo_path: Path):
"""Test request_references on a function parameter."""
# Get references to the id parameter in get_user method
file_path = str(repo_path / "test_repo" / "services.py")
# Line 24 contains the get_user method with id parameter
references = language_server.request_references(file_path, 24, 16)
# id parameter should be referenced within the method
assert len(references) > 0
def test_request_references_create_user_method(self, language_server: SyncLanguageServer, repo_path: Path):
# Get references to the create_user method in UserService
file_path = str(repo_path / "test_repo" / "services.py")
# Line 15 contains the create_user method definition
references = language_server.request_references(file_path, 15, 9)
# Verify that we get valid references
assert len(references) > 1
def test_retrieve_content_around_line(self, language_server: SyncLanguageServer, repo_path: Path):
"""Test retrieve_content_around_line functionality with various scenarios."""
file_path = str(repo_path / "test_repo" / "models.py")
# Scenario 1: Just a single line (User class definition)
line_31 = language_server.retrieve_content_around_line(file_path, 31)
assert len(line_31.lines) == 1
assert "class User(BaseModel):" in line_31.lines[0].line_content
assert line_31.lines[0].line_number == 31
assert line_31.lines[0].match_type == LineType.MATCH
# Scenario 2: Context above and below
with_context_around_user = language_server.retrieve_content_around_line(file_path, 31, 2, 2)
assert len(with_context_around_user.lines) == 5
# Check line content
assert "class User(BaseModel):" in with_context_around_user.matched_lines[0].line_content
assert with_context_around_user.num_matched_lines == 1
assert " User model representing a system user." in with_context_around_user.lines[4].line_content
# Check line numbers
assert with_context_around_user.lines[0].line_number == 29
assert with_context_around_user.lines[1].line_number == 30
assert with_context_around_user.lines[2].line_number == 31
assert with_context_around_user.lines[3].line_number == 32
assert with_context_around_user.lines[4].line_number == 33
# Check match types
assert with_context_around_user.lines[0].match_type == LineType.BEFORE_MATCH
assert with_context_around_user.lines[1].match_type == LineType.BEFORE_MATCH
assert with_context_around_user.lines[2].match_type == LineType.MATCH
assert with_context_around_user.lines[3].match_type == LineType.AFTER_MATCH
assert with_context_around_user.lines[4].match_type == LineType.AFTER_MATCH
# Scenario 3a: Only context above
with_context_above = language_server.retrieve_content_around_line(file_path, 31, 3, 0)
assert len(with_context_above.lines) == 4
assert "return cls(id=id, name=name)" in with_context_above.lines[0].line_content
assert "class User(BaseModel):" in with_context_above.matched_lines[0].line_content
assert with_context_above.num_matched_lines == 1
# Check line numbers
assert with_context_above.lines[0].line_number == 28
assert with_context_above.lines[1].line_number == 29
assert with_context_above.lines[2].line_number == 30
assert with_context_above.lines[3].line_number == 31
# Check match types
assert with_context_above.lines[0].match_type == LineType.BEFORE_MATCH
assert with_context_above.lines[1].match_type == LineType.BEFORE_MATCH
assert with_context_above.lines[2].match_type == LineType.BEFORE_MATCH
assert with_context_above.lines[3].match_type == LineType.MATCH
# Scenario 3b: Only context below
with_context_below = language_server.retrieve_content_around_line(file_path, 31, 0, 3)
assert len(with_context_below.lines) == 4
assert "class User(BaseModel):" in with_context_below.matched_lines[0].line_content
assert with_context_below.num_matched_lines == 1
assert with_context_below.lines[0].line_number == 31
assert with_context_below.lines[1].line_number == 32
assert with_context_below.lines[2].line_number == 33
assert with_context_below.lines[3].line_number == 34
# Check match types
assert with_context_below.lines[0].match_type == LineType.MATCH
assert with_context_below.lines[1].match_type == LineType.AFTER_MATCH
assert with_context_below.lines[2].match_type == LineType.AFTER_MATCH
assert with_context_below.lines[3].match_type == LineType.AFTER_MATCH
# Scenario 4a: Edge case - context above but line is at 0
first_line_with_context_around = language_server.retrieve_content_around_line(file_path, 0, 2, 1)
assert len(first_line_with_context_around.lines) <= 4 # Should have at most 4 lines (line 0 + 1 below + up to 2 above)
assert first_line_with_context_around.lines[0].line_number <= 2 # First line should be at most line 2
# Check match type for the target line
for line in first_line_with_context_around.lines:
if line.line_number == 0:
assert line.match_type == LineType.MATCH
elif line.line_number < 0:
assert line.match_type == LineType.BEFORE_MATCH
else:
assert line.match_type == LineType.AFTER_MATCH
# Scenario 4b: Edge case - context above but line is at 1
second_line_with_context_above = language_server.retrieve_content_around_line(file_path, 1, 3, 1)
assert len(second_line_with_context_above.lines) <= 5 # Should have at most 5 lines (line 1 + 1 below + up to 3 above)
assert second_line_with_context_above.lines[0].line_number <= 1 # First line should be at most line 1
# Check match type for the target line
for line in second_line_with_context_above.lines:
if line.line_number == 1:
assert line.match_type == LineType.MATCH
elif line.line_number < 1:
assert line.match_type == LineType.BEFORE_MATCH
else:
assert line.match_type == LineType.AFTER_MATCH
# Scenario 4c: Edge case - context below but line is at the end of file
# First get the total number of lines in the file
all_content = language_server.retrieve_full_file_content(file_path)
total_lines = len(all_content.split("\n"))
last_line_with_context_around = language_server.retrieve_content_around_line(file_path, total_lines - 1, 1, 3)
assert len(last_line_with_context_around.lines) <= 5 # Should have at most 5 lines (last line + 1 above + up to 3 below)
assert last_line_with_context_around.lines[-1].line_number >= total_lines - 4 # Last line should be at least total_lines - 4
# Check match type for the target line
for line in last_line_with_context_around.lines:
if line.line_number == total_lines - 1:
assert line.match_type == LineType.MATCH
elif line.line_number < total_lines - 1:
assert line.match_type == LineType.BEFORE_MATCH
else:
assert line.match_type == LineType.AFTER_MATCH
@@ -1,83 +0,0 @@
"""
Basic integration tests for the language server functionality.
These tests validate the functionality of the language server APIs
like request_references using the test repository.
"""
from pathlib import Path
import pytest
from multilspy.language_server import SyncLanguageServer
from multilspy.multilspy_config import Language, MultilspyConfig
from multilspy.multilspy_logger import MultilspyLogger
@pytest.fixture(scope="session")
def language_server(repo_path: Path):
"""Create a SyncLanguageServer instance configured to use the test repository."""
config = MultilspyConfig(code_language=Language.PYTHON)
logger = MultilspyLogger()
# Create a language server instance
server = SyncLanguageServer.create(config, logger, str(repo_path))
# Start the server
server.start()
try:
yield server
finally:
# Ensure server is shut down
server.stop()
class TestLanguageServerBasics:
"""Test basic functionality of the language server."""
def test_request_references_user_class(self, language_server: SyncLanguageServer, repo_path: Path):
"""Test request_references on the User class."""
# Get references to the User class in models.py
file_path = str(repo_path / "test_repo" / "models.py")
# Line 31 contains the User class definition
references = language_server.request_references(file_path, 31, 6)
# User class should be referenced in multiple files
assert len(references) > 0
# At least two references should be found (one for the class definition itself)
assert len(references) > 1
def test_request_references_item_class(self, language_server: SyncLanguageServer, repo_path: Path):
"""Test request_references on the Item class."""
# Get references to the Item class in models.py
file_path = str(repo_path / "test_repo" / "models.py")
# Line 56 contains the Item class definition
references = language_server.request_references(file_path, 56, 6)
# Item class should be referenced in multiple places
assert len(references) > 0
# At least one reference should be in services.py (ItemService class)
services_references = [ref for ref in references if "services.py" in ref["uri"]]
assert len(services_references) > 0
def test_request_references_function_parameter(self, language_server: SyncLanguageServer, repo_path: Path):
"""Test request_references on a function parameter."""
# Get references to the id parameter in get_user method
file_path = str(repo_path / "test_repo" / "services.py")
# Line 24 contains the get_user method with id parameter
references = language_server.request_references(file_path, 24, 16)
# id parameter should be referenced within the method
assert len(references) > 0
def test_request_references_create_user_method(self, language_server: SyncLanguageServer, repo_path: Path):
# Get references to the create_user method in UserService
file_path = str(repo_path / "test_repo" / "services.py")
# Line 15 contains the create_user method definition
references = language_server.request_references(file_path, 15, 9)
# Verify that we get valid references
assert len(references) > 1
@@ -8,33 +8,10 @@ These tests focus on the following methods:
from pathlib import Path
import pytest
from multilspy.language_server import SyncLanguageServer
from multilspy.multilspy_config import Language, MultilspyConfig
from multilspy.multilspy_logger import MultilspyLogger
from multilspy.multilspy_types import SymbolKind
@pytest.fixture(scope="session")
def language_server(repo_path: Path):
"""Create a SyncLanguageServer instance configured to use the test repository."""
config = MultilspyConfig(code_language=Language.PYTHON)
logger = MultilspyLogger()
# Create a language server instance
server = SyncLanguageServer.create(config, logger, str(repo_path))
# Start the server
server.start()
try:
yield server
finally:
# Ensure server is shut down
server.stop()
class TestLanguageServerSymbols:
"""Test the language server's symbol-related functionality."""
+20 -20
View File
@@ -21,7 +21,7 @@ class TestTextUtils:
assert matches[0].num_matched_lines == 1
assert matches[0].start_line == 3
assert matches[0].end_line == 3
assert matches[0].context_lines[0].line_content.strip() == 'print("Hello, World!")'
assert matches[0].lines[0].line_content.strip() == 'print("Hello, World!")'
def test_search_text_with_regex_pattern(self):
"""Test searching with a regex pattern."""
@@ -42,10 +42,10 @@ class TestTextUtils:
matches = search_text(pattern, content=content)
assert len(matches) == 3
assert matches[0].context_lines[0].match_type == LineType.MATCH
assert "def __init__" in matches[0].context_lines[0].line_content
assert "def process" in matches[1].context_lines[0].line_content
assert "def filter" in matches[2].context_lines[0].line_content
assert matches[0].lines[0].match_type == LineType.MATCH
assert "def __init__" in matches[0].lines[0].line_content
assert "def process" in matches[1].lines[0].line_content
assert "def filter" in matches[2].lines[0].line_content
def test_search_text_with_compiled_regex(self):
"""Test searching with a pre-compiled regex pattern."""
@@ -68,8 +68,8 @@ class TestTextUtils:
matches = search_text(pattern, content=content)
assert len(matches) == 2
assert "DEBUG = True" in matches[0].context_lines[0].line_content
assert "MAX_RETRIES = 3" in matches[1].context_lines[0].line_content
assert "DEBUG = True" in matches[0].lines[0].line_content
assert "MAX_RETRIES = 3" in matches[1].lines[0].line_content
def test_search_text_with_context_lines(self):
"""Test searching with context lines before and after the match."""
@@ -91,15 +91,15 @@ class TestTextUtils:
# Check the first match with context
first_match = matches[0]
assert len(first_match.context_lines) == 3
assert first_match.context_lines[0].match_type == LineType.BEFORE_MATCH
assert first_match.context_lines[1].match_type == LineType.MATCH
assert first_match.context_lines[2].match_type == LineType.AFTER_MATCH
assert len(first_match.lines) == 3
assert first_match.lines[0].match_type == LineType.BEFORE_MATCH
assert first_match.lines[1].match_type == LineType.MATCH
assert first_match.lines[2].match_type == LineType.AFTER_MATCH
# Verify the content of lines
assert "if a > b:" in first_match.context_lines[0].line_content
assert "return a * c" in first_match.context_lines[1].line_content
assert "elif b > a:" in first_match.context_lines[2].line_content
assert "if a > b:" in first_match.lines[0].line_content
assert "return a * c" in first_match.lines[1].line_content
assert "elif b > a:" in first_match.lines[2].line_content
def test_search_text_with_multiline_match(self):
"""Test searching with multiline pattern matching."""
@@ -120,10 +120,10 @@ class TestTextUtils:
assert len(matches) == 1
multiline_match = matches[0]
assert multiline_match.num_matched_lines >= 3
assert "if n <= 1:" in multiline_match.context_lines[0].line_content
assert "if n <= 1:" in multiline_match.lines[0].line_content
# All matched lines should have match_type == LineType.MATCH
match_lines = [line for line in multiline_match.context_lines if line.match_type == LineType.MATCH]
match_lines = [line for line in multiline_match.lines if line.match_type == LineType.MATCH]
assert len(match_lines) >= 3
def test_search_text_with_glob_pattern(self):
@@ -146,9 +146,9 @@ class TestTextUtils:
matches = search_text("*_user*", content=content, is_glob=True)
assert len(matches) == 3
assert "get_user" in matches[0].context_lines[0].line_content
assert "create_user" in matches[1].context_lines[0].line_content
assert "update_user" in matches[2].context_lines[0].line_content
assert "get_user" in matches[0].lines[0].line_content
assert "create_user" in matches[1].lines[0].line_content
assert "update_user" in matches[2].lines[0].line_content
def test_search_text_with_complex_glob_pattern(self):
"""Test searching with more complex glob patterns."""
@@ -175,7 +175,7 @@ class TestTextUtils:
instance_matches = [
line.line_content
for match in matches
for line in match.context_lines
for line in match.lines
if line.match_type == LineType.MATCH and "isinstance(item," in line.line_content
]
assert len(instance_matches) >= 2