111 lines
3.9 KiB
Python
111 lines
3.9 KiB
Python
|
|
# Copyright (c) "Neo4j"
|
||
|
|
# Neo4j Sweden AB [https://neo4j.com]
|
||
|
|
# #
|
||
|
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||
|
|
# you may not use this file except in compliance with the License.
|
||
|
|
# You may obtain a copy of the License at
|
||
|
|
# #
|
||
|
|
# https://www.apache.org/licenses/LICENSE-2.0
|
||
|
|
# #
|
||
|
|
# Unless required by applicable law or agreed to in writing, software
|
||
|
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||
|
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||
|
|
# See the License for the specific language governing permissions and
|
||
|
|
# limitations under the License.
|
||
|
|
from __future__ import annotations # Reminder: May be removed after Python 3.9 is EOL.
|
||
|
|
|
||
|
|
import inspect
|
||
|
|
from typing import Any, Optional
|
||
|
|
from unittest.mock import MagicMock, patch
|
||
|
|
|
||
|
|
import pytest
|
||
|
|
from neo4j_graphrag.exceptions import Neo4jVersionError
|
||
|
|
from neo4j_graphrag.retrievers.base import Retriever
|
||
|
|
from neo4j_graphrag.types import RawSearchResult, RetrieverResult
|
||
|
|
|
||
|
|
|
||
|
|
@pytest.mark.parametrize(
|
||
|
|
"db_version,expected_exception",
|
||
|
|
[
|
||
|
|
(((5, 18, 0), True, True), None),
|
||
|
|
(((5, 3, 0), True, True), Neo4jVersionError),
|
||
|
|
(((5, 19, 0), False, True), None),
|
||
|
|
(((4, 3, 5), False, True), Neo4jVersionError),
|
||
|
|
(((5, 23, 0), False, True), None),
|
||
|
|
(((2025, 1, 0), False, True), None),
|
||
|
|
(((2025, 1, 0), True, True), None),
|
||
|
|
],
|
||
|
|
)
|
||
|
|
@patch("neo4j_graphrag.retrievers.base.get_version")
|
||
|
|
def test_retriever_version_support(
|
||
|
|
mock_get_version: MagicMock,
|
||
|
|
driver: MagicMock,
|
||
|
|
db_version: tuple[tuple[int, ...], bool],
|
||
|
|
expected_exception: Optional[type[ValueError]],
|
||
|
|
) -> None:
|
||
|
|
mock_get_version.return_value = db_version
|
||
|
|
|
||
|
|
class MockRetriever(Retriever):
|
||
|
|
def get_search_results(self, *args: Any, **kwargs: Any) -> RawSearchResult:
|
||
|
|
return RawSearchResult(records=[])
|
||
|
|
|
||
|
|
if expected_exception:
|
||
|
|
with pytest.raises(expected_exception):
|
||
|
|
MockRetriever(driver=driver)
|
||
|
|
else:
|
||
|
|
MockRetriever(driver=driver)
|
||
|
|
|
||
|
|
|
||
|
|
@patch("neo4j_graphrag.retrievers.base.get_version")
|
||
|
|
def test_retriever_search_docstring_copied(
|
||
|
|
mock_get_version: MagicMock,
|
||
|
|
driver: MagicMock,
|
||
|
|
) -> None:
|
||
|
|
mock_get_version.return_value = ((5, 23, 0), False, False)
|
||
|
|
|
||
|
|
class MockRetriever(Retriever):
|
||
|
|
def get_search_results(self, query: str, top_k: int = 10) -> RawSearchResult:
|
||
|
|
"""My fabulous docstring"""
|
||
|
|
return RawSearchResult(records=[])
|
||
|
|
|
||
|
|
retriever = MockRetriever(driver=driver)
|
||
|
|
assert retriever.search.__doc__ == "My fabulous docstring"
|
||
|
|
signature = inspect.signature(retriever.search)
|
||
|
|
assert "query" in signature.parameters
|
||
|
|
query_param = signature.parameters["query"]
|
||
|
|
assert query_param.default == query_param.empty
|
||
|
|
assert query_param.annotation == "str"
|
||
|
|
assert "top_k" in signature.parameters
|
||
|
|
top_k_param = signature.parameters["top_k"]
|
||
|
|
assert top_k_param.default == 10
|
||
|
|
assert top_k_param.annotation == "int"
|
||
|
|
|
||
|
|
|
||
|
|
@patch("neo4j_graphrag.retrievers.base.get_version")
|
||
|
|
def test_retriever_search_docstring_unchanged(
|
||
|
|
mock_get_version: MagicMock,
|
||
|
|
driver: MagicMock,
|
||
|
|
) -> None:
|
||
|
|
mock_get_version.return_value = ((5, 23, 0), False, False)
|
||
|
|
|
||
|
|
class MockRetrieverForNoise(Retriever):
|
||
|
|
def get_search_results(self, query: str, top_k: int = 10) -> RawSearchResult:
|
||
|
|
"""My fabulous docstring"""
|
||
|
|
return RawSearchResult(records=[])
|
||
|
|
|
||
|
|
class MockRetriever(Retriever):
|
||
|
|
def get_search_results(self, *args: Any, **kwargs: Any) -> RawSearchResult:
|
||
|
|
return RawSearchResult(records=[])
|
||
|
|
|
||
|
|
def search(self, query: str, top_k: int = 10) -> RetrieverResult:
|
||
|
|
"""My fabulous docstring that I do not want to be updated"""
|
||
|
|
return RetrieverResult(items=[])
|
||
|
|
|
||
|
|
assert MockRetrieverForNoise.search is not MockRetriever.search
|
||
|
|
|
||
|
|
retriever = MockRetriever(driver=driver)
|
||
|
|
assert (
|
||
|
|
retriever.search.__doc__
|
||
|
|
== "My fabulous docstring that I do not want to be updated"
|
||
|
|
)
|