참고소스 수정본
This commit is contained in:
282
참고/neo4j-graphrag-python-main/tests/unit/retrievers/external/test_pinecone.py
vendored
Normal file
282
참고/neo4j-graphrag-python-main/tests/unit/retrievers/external/test_pinecone.py
vendored
Normal file
@@ -0,0 +1,282 @@
|
||||
# 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 unittest import mock
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import neo4j
|
||||
import pytest
|
||||
from neo4j_graphrag.exceptions import RetrieverInitializationError
|
||||
from neo4j_graphrag.retrievers import PineconeNeo4jRetriever
|
||||
from neo4j_graphrag.retrievers.external.utils import get_match_query
|
||||
from neo4j_graphrag.types import RetrieverResult, RetrieverResultItem
|
||||
from pinecone import Pinecone
|
||||
|
||||
|
||||
@pytest.fixture(scope="function")
|
||||
def client() -> MagicMock:
|
||||
return MagicMock(spec=Pinecone)
|
||||
|
||||
|
||||
def test_pinecone_retriever_invalid_return_properties(
|
||||
driver: MagicMock, client: MagicMock
|
||||
) -> None:
|
||||
with pytest.raises(RetrieverInitializationError) as exc_info:
|
||||
PineconeNeo4jRetriever(
|
||||
driver=driver,
|
||||
client=client,
|
||||
index_name="dummy-text",
|
||||
id_property_neo4j="dummy-text",
|
||||
return_properties=42, # type: ignore
|
||||
)
|
||||
|
||||
assert "return_properties" in str(exc_info.value)
|
||||
assert "Input should be a valid list" in str(exc_info.value)
|
||||
|
||||
|
||||
def test_pinecone_retriever_invalid_retrieval_query(
|
||||
driver: MagicMock, client: MagicMock
|
||||
) -> None:
|
||||
with pytest.raises(RetrieverInitializationError) as exc_info:
|
||||
PineconeNeo4jRetriever(
|
||||
driver=driver,
|
||||
client=client,
|
||||
index_name="dummy-text",
|
||||
id_property_neo4j="dummy-text",
|
||||
retrieval_query=42, # type: ignore
|
||||
)
|
||||
|
||||
assert "retrieval_query" in str(exc_info.value)
|
||||
assert "Input should be a valid string" in str(exc_info.value)
|
||||
|
||||
|
||||
def test_pinecone_retriever_search_happy_path(
|
||||
driver: MagicMock, client: MagicMock
|
||||
) -> None:
|
||||
retriever = PineconeNeo4jRetriever(
|
||||
driver=driver,
|
||||
client=client,
|
||||
index_name="dummy-text",
|
||||
id_property_neo4j="sync_id",
|
||||
)
|
||||
with mock.patch.object(retriever, "index") as mock_index:
|
||||
top_k = 5
|
||||
mock_index.query.return_value = {
|
||||
"matches": [
|
||||
{"id": f"node_{i}", "score": i / top_k, "values": []}
|
||||
for i in range(top_k)
|
||||
],
|
||||
"namespace": "",
|
||||
"usage": {"read_units": top_k},
|
||||
}
|
||||
driver.execute_query.return_value = (
|
||||
[
|
||||
neo4j.Record({"node": {"sync_id": f"node_{i}"}, "score": i / top_k})
|
||||
for i in range(top_k)
|
||||
],
|
||||
None,
|
||||
None,
|
||||
)
|
||||
query_vector = [1.0 for _ in range(1536)]
|
||||
search_query = get_match_query()
|
||||
records = retriever.search(query_vector=query_vector)
|
||||
|
||||
driver.execute_query.assert_called_once_with(
|
||||
search_query,
|
||||
{
|
||||
"match_params": [(f"node_{i}", i / top_k) for i in range(top_k)],
|
||||
"id_property": "sync_id",
|
||||
},
|
||||
database_=None,
|
||||
routing_=neo4j.RoutingControl.READ,
|
||||
)
|
||||
|
||||
assert records == RetrieverResult(
|
||||
items=[
|
||||
RetrieverResultItem(
|
||||
content="<Record node={'sync_id': "
|
||||
+ f"'node_{i}'"
|
||||
+ "} "
|
||||
+ f"score={i / top_k}>",
|
||||
metadata=None,
|
||||
)
|
||||
for i in range(top_k)
|
||||
],
|
||||
metadata={"__retriever": "PineconeNeo4jRetriever"},
|
||||
)
|
||||
|
||||
|
||||
def test_invalid_neo4j_database_name(driver: MagicMock, client: MagicMock) -> None:
|
||||
with pytest.raises(RetrieverInitializationError) as exc_info:
|
||||
PineconeNeo4jRetriever(
|
||||
driver=driver,
|
||||
client=client,
|
||||
index_name="dummy-text",
|
||||
id_property_neo4j="sync_id",
|
||||
neo4j_database=42, # type: ignore
|
||||
)
|
||||
|
||||
assert "neo4j_database" in str(exc_info.value)
|
||||
assert "Input should be a valid string" in str(exc_info.value)
|
||||
|
||||
|
||||
def test_pinecone_retriever_search_return_properties(
|
||||
driver: MagicMock, client: MagicMock
|
||||
) -> None:
|
||||
retriever = PineconeNeo4jRetriever(
|
||||
driver=driver,
|
||||
client=client,
|
||||
index_name="dummy-text",
|
||||
id_property_neo4j="sync_id",
|
||||
return_properties=["sync_id"],
|
||||
)
|
||||
with mock.patch.object(retriever, "index") as mock_index:
|
||||
top_k = 5
|
||||
mock_index.query.return_value = {
|
||||
"matches": [
|
||||
{"id": f"node_{i}", "score": i / top_k, "values": []}
|
||||
for i in range(top_k)
|
||||
],
|
||||
"namespace": "",
|
||||
"usage": {"read_units": top_k},
|
||||
}
|
||||
driver.execute_query.return_value = (
|
||||
[
|
||||
neo4j.Record({"node": {"sync_id": f"node_{i}"}, "score": i / top_k})
|
||||
for i in range(top_k)
|
||||
],
|
||||
None,
|
||||
None,
|
||||
)
|
||||
query_vector = [1.0 for _ in range(1536)]
|
||||
search_query = get_match_query(return_properties=["sync_id"])
|
||||
records = retriever.search(
|
||||
query_vector=query_vector,
|
||||
)
|
||||
|
||||
driver.execute_query.assert_called_once_with(
|
||||
search_query,
|
||||
{
|
||||
"match_params": [(f"node_{i}", i / top_k) for i in range(top_k)],
|
||||
"id_property": "sync_id",
|
||||
},
|
||||
database_=None,
|
||||
routing_=neo4j.RoutingControl.READ,
|
||||
)
|
||||
|
||||
assert records == RetrieverResult(
|
||||
items=[
|
||||
RetrieverResultItem(
|
||||
content="<Record node={'sync_id': "
|
||||
+ f"'node_{i}'"
|
||||
+ "} "
|
||||
+ f"score={i / top_k}>",
|
||||
metadata=None,
|
||||
)
|
||||
for i in range(top_k)
|
||||
],
|
||||
metadata={"__retriever": "PineconeNeo4jRetriever"},
|
||||
)
|
||||
|
||||
|
||||
def test_pinecone_retriever_search_retrieval_query(
|
||||
driver: MagicMock, client: MagicMock
|
||||
) -> None:
|
||||
retrieval_query = "WITH node MATCH (node)--(m) RETURN n, m LIMIT 10"
|
||||
retriever = PineconeNeo4jRetriever(
|
||||
driver=driver,
|
||||
client=client,
|
||||
index_name="dummy-text",
|
||||
id_property_neo4j="sync_id",
|
||||
retrieval_query=retrieval_query,
|
||||
)
|
||||
with mock.patch.object(retriever, "index") as mock_index:
|
||||
top_k = 5
|
||||
mock_index.query.return_value = {
|
||||
"matches": [
|
||||
{"id": f"node_{i}", "score": i / top_k, "values": []}
|
||||
for i in range(top_k)
|
||||
],
|
||||
"namespace": "",
|
||||
"usage": {"read_units": top_k},
|
||||
}
|
||||
driver.execute_query.return_value = (
|
||||
[
|
||||
neo4j.Record({"node": {"sync_id": f"node_{i}"}, "score": i / top_k})
|
||||
for i in range(top_k)
|
||||
],
|
||||
None,
|
||||
None,
|
||||
)
|
||||
query_vector = [1.0 for _ in range(1536)]
|
||||
search_query = get_match_query(retrieval_query=retrieval_query)
|
||||
records = retriever.search(
|
||||
query_vector=query_vector,
|
||||
)
|
||||
|
||||
driver.execute_query.assert_called_once_with(
|
||||
search_query,
|
||||
{
|
||||
"match_params": [(f"node_{i}", i / top_k) for i in range(top_k)],
|
||||
"id_property": "sync_id",
|
||||
},
|
||||
database_=None,
|
||||
routing_=neo4j.RoutingControl.READ,
|
||||
)
|
||||
|
||||
assert records == RetrieverResult(
|
||||
items=[
|
||||
RetrieverResultItem(
|
||||
content="<Record node={'sync_id': "
|
||||
+ f"'node_{i}'"
|
||||
+ "} "
|
||||
+ f"score={i / top_k}>",
|
||||
metadata=None,
|
||||
)
|
||||
for i in range(top_k)
|
||||
],
|
||||
metadata={"__retriever": "PineconeNeo4jRetriever"},
|
||||
)
|
||||
|
||||
|
||||
def test_pinecone_retriever_with_result_format_function(
|
||||
driver: MagicMock,
|
||||
client: MagicMock,
|
||||
neo4j_record: MagicMock,
|
||||
result_formatter: MagicMock,
|
||||
) -> None:
|
||||
retriever = PineconeNeo4jRetriever(
|
||||
driver=driver,
|
||||
client=client,
|
||||
index_name="dummy-text",
|
||||
id_property_neo4j="sync_id",
|
||||
result_formatter=result_formatter,
|
||||
)
|
||||
with mock.patch.object(retriever, "index"):
|
||||
driver.execute_query.return_value = (
|
||||
[neo4j_record],
|
||||
None,
|
||||
None,
|
||||
)
|
||||
query_vector = [1.0 for _ in range(1536)]
|
||||
records = retriever.search(query_vector=query_vector)
|
||||
|
||||
assert records == RetrieverResult(
|
||||
items=[
|
||||
RetrieverResultItem(
|
||||
content="dummy-node", metadata={"score": 1.0, "node_id": 123}
|
||||
),
|
||||
],
|
||||
metadata={"__retriever": "PineconeNeo4jRetriever"},
|
||||
)
|
||||
337
참고/neo4j-graphrag-python-main/tests/unit/retrievers/external/test_qdrant.py
vendored
Normal file
337
참고/neo4j-graphrag-python-main/tests/unit/retrievers/external/test_qdrant.py
vendored
Normal file
@@ -0,0 +1,337 @@
|
||||
# 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 typing import Any
|
||||
from unittest import mock
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import neo4j
|
||||
import pytest
|
||||
from neo4j_graphrag.exceptions import RetrieverInitializationError
|
||||
from neo4j_graphrag.retrievers import QdrantNeo4jRetriever
|
||||
from neo4j_graphrag.retrievers.external.utils import get_match_query
|
||||
from neo4j_graphrag.types import RetrieverResult, RetrieverResultItem
|
||||
from qdrant_client import QdrantClient
|
||||
from qdrant_client.http.models import QueryResponse, ScoredPoint
|
||||
|
||||
|
||||
@pytest.fixture(scope="function")
|
||||
def client() -> MagicMock:
|
||||
return MagicMock(spec=QdrantClient)
|
||||
|
||||
|
||||
def test_qdrant_retriever_search_happy_path(
|
||||
driver: MagicMock, client: MagicMock
|
||||
) -> None:
|
||||
retriever = QdrantNeo4jRetriever(
|
||||
driver=driver,
|
||||
client=client,
|
||||
collection_name="dummy-text",
|
||||
id_property_neo4j="sync_id",
|
||||
id_property_external="sync_id",
|
||||
)
|
||||
with mock.patch.object(retriever, "client") as mock_client:
|
||||
top_k = 5
|
||||
mock_client.query_points.return_value = QueryResponse(
|
||||
points=[
|
||||
ScoredPoint(
|
||||
id=i,
|
||||
version=0,
|
||||
score=i / top_k,
|
||||
payload={
|
||||
"sync_id": f"node_{i}",
|
||||
},
|
||||
)
|
||||
for i in range(top_k)
|
||||
]
|
||||
)
|
||||
driver.execute_query.return_value = (
|
||||
[
|
||||
neo4j.Record({"node": {"sync_id": f"node_{i}"}, "score": i / top_k})
|
||||
for i in range(top_k)
|
||||
],
|
||||
None,
|
||||
None,
|
||||
)
|
||||
query_vector = [1.0 for _ in range(1536)]
|
||||
search_query = get_match_query()
|
||||
records = retriever.search(query_vector=query_vector)
|
||||
|
||||
driver.execute_query.assert_called_once_with(
|
||||
search_query,
|
||||
{
|
||||
"match_params": [(f"node_{i}", i / top_k) for i in range(top_k)],
|
||||
"id_property": "sync_id",
|
||||
},
|
||||
database_=None,
|
||||
routing_=neo4j.RoutingControl.READ,
|
||||
)
|
||||
|
||||
assert records == RetrieverResult(
|
||||
items=[
|
||||
RetrieverResultItem(
|
||||
content="<Record node={'sync_id': "
|
||||
+ f"'node_{i}'"
|
||||
+ "} "
|
||||
+ f"score={i / top_k}>",
|
||||
metadata=None,
|
||||
)
|
||||
for i in range(top_k)
|
||||
],
|
||||
metadata={"__retriever": "QdrantNeo4jRetriever"},
|
||||
)
|
||||
|
||||
|
||||
def test_invalid_neo4j_database_name(driver: MagicMock, client: MagicMock) -> None:
|
||||
with pytest.raises(RetrieverInitializationError) as exc_info:
|
||||
QdrantNeo4jRetriever(
|
||||
driver=driver,
|
||||
client=client,
|
||||
collection_name="dummy-text",
|
||||
id_property_neo4j="sync_id",
|
||||
neo4j_database=42, # type: ignore
|
||||
)
|
||||
|
||||
assert "neo4j_database" in str(exc_info.value)
|
||||
assert "Input should be a valid string" in str(exc_info.value)
|
||||
|
||||
|
||||
def test_qdrant_retriever_search_return_properties(
|
||||
driver: MagicMock, client: MagicMock
|
||||
) -> None:
|
||||
retriever = QdrantNeo4jRetriever(
|
||||
driver=driver,
|
||||
client=client,
|
||||
collection_name="dummy-text",
|
||||
id_property_neo4j="sync_id",
|
||||
id_property_external="sync_id",
|
||||
return_properties=["sync_id"],
|
||||
)
|
||||
with mock.patch.object(retriever, "client") as mock_client:
|
||||
top_k = 5
|
||||
mock_client.query_points.return_value = QueryResponse(
|
||||
points=[
|
||||
ScoredPoint(
|
||||
id=i,
|
||||
version=0,
|
||||
score=i / top_k,
|
||||
payload={
|
||||
"sync_id": f"node_{i}",
|
||||
},
|
||||
)
|
||||
for i in range(top_k)
|
||||
]
|
||||
)
|
||||
driver.execute_query.return_value = (
|
||||
[
|
||||
neo4j.Record({"node": {"sync_id": f"node_{i}"}, "score": i / top_k})
|
||||
for i in range(top_k)
|
||||
],
|
||||
None,
|
||||
None,
|
||||
)
|
||||
query_vector = [1.0 for _ in range(1536)]
|
||||
search_query = get_match_query(return_properties=["sync_id"])
|
||||
records = retriever.search(
|
||||
query_vector=query_vector,
|
||||
)
|
||||
|
||||
driver.execute_query.assert_called_once_with(
|
||||
search_query,
|
||||
{
|
||||
"match_params": [(f"node_{i}", i / top_k) for i in range(top_k)],
|
||||
"id_property": "sync_id",
|
||||
},
|
||||
database_=None,
|
||||
routing_=neo4j.RoutingControl.READ,
|
||||
)
|
||||
|
||||
assert records == RetrieverResult(
|
||||
items=[
|
||||
RetrieverResultItem(
|
||||
content="<Record node={'sync_id': "
|
||||
+ f"'node_{i}'"
|
||||
+ "} "
|
||||
+ f"score={i / top_k}>",
|
||||
metadata=None,
|
||||
)
|
||||
for i in range(top_k)
|
||||
],
|
||||
metadata={"__retriever": "QdrantNeo4jRetriever"},
|
||||
)
|
||||
|
||||
|
||||
def test_qdrant_retriever_search_retrieval_query(
|
||||
driver: MagicMock, client: MagicMock
|
||||
) -> None:
|
||||
retrieval_query = "WITH node MATCH (node)--(m) RETURN n, m LIMIT 10"
|
||||
retriever = QdrantNeo4jRetriever(
|
||||
driver=driver,
|
||||
client=client,
|
||||
collection_name="dummy-text",
|
||||
id_property_neo4j="sync_id",
|
||||
id_property_external="sync_id",
|
||||
retrieval_query=retrieval_query,
|
||||
)
|
||||
with mock.patch.object(retriever, "client") as mock_client:
|
||||
top_k = 5
|
||||
mock_client.query_points.return_value = QueryResponse(
|
||||
points=[
|
||||
ScoredPoint(
|
||||
id=i,
|
||||
version=0,
|
||||
score=i / top_k,
|
||||
payload={
|
||||
"sync_id": f"node_{i}",
|
||||
},
|
||||
)
|
||||
for i in range(top_k)
|
||||
]
|
||||
)
|
||||
driver.execute_query.return_value = (
|
||||
[
|
||||
neo4j.Record({"node": {"sync_id": f"node_{i}"}, "score": i / top_k})
|
||||
for i in range(top_k)
|
||||
],
|
||||
None,
|
||||
None,
|
||||
)
|
||||
query_vector = [1.0 for _ in range(1536)]
|
||||
search_query = get_match_query(retrieval_query=retrieval_query)
|
||||
records = retriever.search(
|
||||
query_vector=query_vector,
|
||||
)
|
||||
|
||||
driver.execute_query.assert_called_once_with(
|
||||
search_query,
|
||||
{
|
||||
"match_params": [(f"node_{i}", i / top_k) for i in range(top_k)],
|
||||
"id_property": "sync_id",
|
||||
},
|
||||
database_=None,
|
||||
routing_=neo4j.RoutingControl.READ,
|
||||
)
|
||||
|
||||
assert records == RetrieverResult(
|
||||
items=[
|
||||
RetrieverResultItem(
|
||||
content="<Record node={'sync_id': "
|
||||
+ f"'node_{i}'"
|
||||
+ "} "
|
||||
+ f"score={i / top_k}>",
|
||||
metadata=None,
|
||||
)
|
||||
for i in range(top_k)
|
||||
],
|
||||
metadata={"__retriever": "QdrantNeo4jRetriever"},
|
||||
)
|
||||
|
||||
|
||||
def test_qdrant_retriever_invalid_return_properties(
|
||||
driver: MagicMock, client: MagicMock
|
||||
) -> None:
|
||||
with pytest.raises(RetrieverInitializationError) as exc_info:
|
||||
QdrantNeo4jRetriever(
|
||||
driver=driver,
|
||||
client=client,
|
||||
collection_name="dummy-text",
|
||||
id_property_neo4j="dummy-text",
|
||||
return_properties=42, # type: ignore
|
||||
)
|
||||
|
||||
assert "return_properties" in str(exc_info.value)
|
||||
assert "Input should be a valid list" in str(exc_info.value)
|
||||
|
||||
|
||||
def test_qdrant_retriever_invalid_retrieval_query(
|
||||
driver: MagicMock, client: MagicMock
|
||||
) -> None:
|
||||
with pytest.raises(RetrieverInitializationError) as exc_info:
|
||||
QdrantNeo4jRetriever(
|
||||
driver=driver,
|
||||
client=client,
|
||||
collection_name="dummy-text",
|
||||
id_property_neo4j="dummy-text",
|
||||
retrieval_query=42, # type: ignore
|
||||
)
|
||||
|
||||
assert "retrieval_query" in str(exc_info.value)
|
||||
assert "Input should be a valid string" in str(exc_info.value)
|
||||
|
||||
|
||||
def test_qdrant_retriever_search_custom_match_id_getter(
|
||||
driver: MagicMock, client: MagicMock
|
||||
) -> None:
|
||||
def my_id_getter(point: ScoredPoint) -> Any:
|
||||
if point.payload is None:
|
||||
raise Exception("Payload is None")
|
||||
return point.payload["data"]["id"]
|
||||
|
||||
retriever = QdrantNeo4jRetriever(
|
||||
driver=driver,
|
||||
client=client,
|
||||
collection_name="dummy-text",
|
||||
id_property_neo4j="sync_id",
|
||||
id_property_getter=my_id_getter,
|
||||
)
|
||||
with mock.patch.object(retriever, "client") as mock_client:
|
||||
top_k = 5
|
||||
mock_client.query_points.return_value = QueryResponse(
|
||||
points=[
|
||||
ScoredPoint(
|
||||
id=i,
|
||||
version=0,
|
||||
score=i / top_k,
|
||||
payload={
|
||||
"data": {"id": f"node_{i}"},
|
||||
},
|
||||
)
|
||||
for i in range(top_k)
|
||||
]
|
||||
)
|
||||
driver.execute_query.return_value = (
|
||||
[
|
||||
neo4j.Record({"node": {"sync_id": f"node_{i}"}, "score": i / top_k})
|
||||
for i in range(top_k)
|
||||
],
|
||||
None,
|
||||
None,
|
||||
)
|
||||
query_vector = [1.0 for _ in range(1536)]
|
||||
search_query = get_match_query()
|
||||
records = retriever.search(query_vector=query_vector)
|
||||
|
||||
driver.execute_query.assert_called_once_with(
|
||||
search_query,
|
||||
{
|
||||
"match_params": [(f"node_{i}", i / top_k) for i in range(top_k)],
|
||||
"id_property": "sync_id",
|
||||
},
|
||||
database_=None,
|
||||
routing_=neo4j.RoutingControl.READ,
|
||||
)
|
||||
|
||||
assert records == RetrieverResult(
|
||||
items=[
|
||||
RetrieverResultItem(
|
||||
content="<Record node={'sync_id': "
|
||||
+ f"'node_{i}'"
|
||||
+ "} "
|
||||
+ f"score={i / top_k}>",
|
||||
metadata=None,
|
||||
)
|
||||
for i in range(top_k)
|
||||
],
|
||||
metadata={"__retriever": "QdrantNeo4jRetriever"},
|
||||
)
|
||||
310
참고/neo4j-graphrag-python-main/tests/unit/retrievers/external/test_weaviate.py
vendored
Normal file
310
참고/neo4j-graphrag-python-main/tests/unit/retrievers/external/test_weaviate.py
vendored
Normal file
@@ -0,0 +1,310 @@
|
||||
# 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 types import SimpleNamespace
|
||||
from typing import Optional
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import neo4j
|
||||
import pytest
|
||||
from neo4j_graphrag.exceptions import RetrieverInitializationError
|
||||
from neo4j_graphrag.retrievers import WeaviateNeo4jRetriever
|
||||
from neo4j_graphrag.retrievers.external.utils import get_match_query
|
||||
from neo4j_graphrag.types import RetrieverResult, RetrieverResultItem
|
||||
from weaviate.client import WeaviateClient
|
||||
|
||||
|
||||
# Weaviate class with fake methods
|
||||
class WClient(WeaviateClient):
|
||||
def __init__(
|
||||
self,
|
||||
node_id_value: Optional[str] = None,
|
||||
node_match_score: Optional[float] = None,
|
||||
) -> None:
|
||||
self.collections = MagicMock()
|
||||
self.collections.get = MagicMock()
|
||||
query = MagicMock()
|
||||
self.collections.get.return_value = SimpleNamespace(query=query)
|
||||
query.near_text.return_value = SimpleNamespace(
|
||||
objects=[
|
||||
SimpleNamespace(
|
||||
properties={"neo4j_id": node_id_value},
|
||||
metadata=SimpleNamespace(certainty=node_match_score),
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
def test_text_search_remote_vector_store_happy_path(driver: MagicMock) -> None:
|
||||
query_text = "may thy knife chip and shatter"
|
||||
top_k = 5
|
||||
node_id_value = "node-test-id"
|
||||
node_match_score = 0.9
|
||||
|
||||
wc = WClient(node_id_value=node_id_value, node_match_score=node_match_score)
|
||||
|
||||
retriever = WeaviateNeo4jRetriever(
|
||||
driver=driver,
|
||||
client=wc,
|
||||
collection="dummy-collection",
|
||||
id_property_neo4j="sync_id",
|
||||
id_property_external="neo4j_id",
|
||||
)
|
||||
driver.execute_query.return_value = [
|
||||
[neo4j.Record({"node": {"sync_id": node_id_value}, "score": node_match_score})],
|
||||
None,
|
||||
None,
|
||||
]
|
||||
search_query = get_match_query()
|
||||
|
||||
records = retriever.search(query_text=query_text, top_k=top_k)
|
||||
|
||||
driver.execute_query.assert_called_once_with(
|
||||
search_query,
|
||||
{
|
||||
"match_params": [
|
||||
(node_id_value, node_match_score),
|
||||
],
|
||||
"id_property": "sync_id",
|
||||
},
|
||||
database_=None,
|
||||
routing_=neo4j.RoutingControl.READ,
|
||||
)
|
||||
assert records == RetrieverResult(
|
||||
items=[
|
||||
RetrieverResultItem(
|
||||
content="<Record node={'sync_id': 'node-test-id'} score=0.9>",
|
||||
metadata=None,
|
||||
),
|
||||
],
|
||||
metadata={"__retriever": "WeaviateNeo4jRetriever"},
|
||||
)
|
||||
|
||||
|
||||
def test_invalid_neo4j_database_name(driver: MagicMock) -> None:
|
||||
node_id_value = "node-test-id"
|
||||
node_match_score = 0.9
|
||||
|
||||
wc = WClient(node_id_value=node_id_value, node_match_score=node_match_score)
|
||||
|
||||
with pytest.raises(RetrieverInitializationError) as exc_info:
|
||||
WeaviateNeo4jRetriever(
|
||||
driver=driver,
|
||||
client=wc,
|
||||
collection="dummy-collection",
|
||||
id_property_neo4j="sync_id",
|
||||
id_property_external="neo4j_id",
|
||||
neo4j_database=42, # type: ignore
|
||||
)
|
||||
|
||||
assert "neo4j_database" in str(exc_info.value)
|
||||
assert "Input should be a valid string" in str(exc_info.value)
|
||||
|
||||
|
||||
def test_text_search_remote_vector_store_return_properties(driver: MagicMock) -> None:
|
||||
query_text = "may thy knife chip and shatter"
|
||||
top_k = 5
|
||||
node_id_value = "node-test-id"
|
||||
node_match_score = 0.9
|
||||
|
||||
wc = WClient(node_id_value=node_id_value, node_match_score=node_match_score)
|
||||
|
||||
retriever = WeaviateNeo4jRetriever(
|
||||
driver=driver,
|
||||
client=wc,
|
||||
collection="dummy-collection",
|
||||
id_property_neo4j="sync_id",
|
||||
id_property_external="neo4j_id",
|
||||
return_properties=["sync_id"],
|
||||
)
|
||||
driver.execute_query.return_value = [
|
||||
[neo4j.Record({"node": {"sync_id": node_id_value}, "score": node_match_score})],
|
||||
None,
|
||||
None,
|
||||
]
|
||||
search_query = get_match_query(return_properties=["sync_id"])
|
||||
|
||||
records = retriever.search(query_text=query_text, top_k=top_k)
|
||||
|
||||
driver.execute_query.assert_called_once_with(
|
||||
search_query,
|
||||
{
|
||||
"match_params": [
|
||||
(node_id_value, node_match_score),
|
||||
],
|
||||
"id_property": "sync_id",
|
||||
},
|
||||
database_=None,
|
||||
routing_=neo4j.RoutingControl.READ,
|
||||
)
|
||||
assert records == RetrieverResult(
|
||||
items=[
|
||||
RetrieverResultItem(
|
||||
content="<Record node={'sync_id': 'node-test-id'} score=0.9>",
|
||||
metadata=None,
|
||||
),
|
||||
],
|
||||
metadata={"__retriever": "WeaviateNeo4jRetriever"},
|
||||
)
|
||||
|
||||
|
||||
def test_text_search_remote_vector_store_retrieval_query(driver: MagicMock) -> None:
|
||||
query_text = "may thy knife chip and shatter"
|
||||
top_k = 5
|
||||
node_id_value = "node-test-id"
|
||||
node_match_score = 0.9
|
||||
retrieval_query = "WITH node MATCH (node)--(m) RETURN n, m LIMIT 10"
|
||||
|
||||
wc = WClient(node_id_value=node_id_value, node_match_score=node_match_score)
|
||||
|
||||
retriever = WeaviateNeo4jRetriever(
|
||||
driver=driver,
|
||||
client=wc,
|
||||
collection="dummy-collection",
|
||||
id_property_neo4j="sync_id",
|
||||
id_property_external="neo4j_id",
|
||||
retrieval_query=retrieval_query,
|
||||
)
|
||||
driver.execute_query.return_value = [
|
||||
[neo4j.Record({"node": {"sync_id": node_id_value}, "score": node_match_score})],
|
||||
None,
|
||||
None,
|
||||
]
|
||||
search_query = get_match_query(retrieval_query=retrieval_query)
|
||||
|
||||
records = retriever.search(query_text=query_text, top_k=top_k)
|
||||
|
||||
driver.execute_query.assert_called_once_with(
|
||||
search_query,
|
||||
{
|
||||
"match_params": [
|
||||
(node_id_value, node_match_score),
|
||||
],
|
||||
"id_property": "sync_id",
|
||||
},
|
||||
database_=None,
|
||||
routing_=neo4j.RoutingControl.READ,
|
||||
)
|
||||
|
||||
assert records == RetrieverResult(
|
||||
items=[
|
||||
RetrieverResultItem(
|
||||
content="<Record node={'sync_id': 'node-test-id'} score=0.9>",
|
||||
metadata=None,
|
||||
),
|
||||
],
|
||||
metadata={"__retriever": "WeaviateNeo4jRetriever"},
|
||||
)
|
||||
|
||||
|
||||
def test_match_query() -> None:
|
||||
match_query = get_match_query()
|
||||
expected = (
|
||||
"UNWIND $match_params AS match_param "
|
||||
"WITH match_param[0] AS match_id_value, match_param[1] AS score "
|
||||
"MATCH (node) "
|
||||
"WHERE node[$id_property] = match_id_value "
|
||||
"RETURN node, score"
|
||||
)
|
||||
assert match_query.strip() == expected.strip()
|
||||
|
||||
|
||||
def test_match_query_with_return_properties() -> None:
|
||||
match_query = get_match_query(return_properties=["name", "age"])
|
||||
expected = (
|
||||
"UNWIND $match_params AS match_param "
|
||||
"WITH match_param[0] AS match_id_value, match_param[1] AS score "
|
||||
"MATCH (node) "
|
||||
"WHERE node[$id_property] = match_id_value "
|
||||
"RETURN node {.name, .age} AS node, labels(node) AS nodeLabels, elementId(node) AS elementId, elementId(node) AS id, score"
|
||||
)
|
||||
assert match_query.strip() == expected.strip()
|
||||
|
||||
|
||||
def test_match_query_with_retrieval_query() -> None:
|
||||
retrieval_query = "WITH node MATCH (node)--(m) RETURN node, m LIMIT 10"
|
||||
match_query = get_match_query(retrieval_query=retrieval_query)
|
||||
expected = (
|
||||
"UNWIND $match_params AS match_param "
|
||||
"WITH match_param[0] AS match_id_value, match_param[1] AS score "
|
||||
"MATCH (node) "
|
||||
"WHERE node[$id_property] = match_id_value " + retrieval_query
|
||||
)
|
||||
assert match_query.strip() == expected.strip()
|
||||
|
||||
|
||||
def test_match_query_with_both_return_properties_and_retrieval_query() -> None:
|
||||
# Should ignore return_properties
|
||||
retrieval_query = "WITH node MATCH (node)--(m) RETURN node, m LIMIT 10"
|
||||
match_query = get_match_query(
|
||||
return_properties=["name", "age"], retrieval_query=retrieval_query
|
||||
)
|
||||
expected = (
|
||||
"UNWIND $match_params AS match_param "
|
||||
"WITH match_param[0] AS match_id_value, match_param[1] AS score "
|
||||
"MATCH (node) "
|
||||
"WHERE node[$id_property] = match_id_value " + retrieval_query
|
||||
)
|
||||
assert match_query.strip() == expected.strip()
|
||||
|
||||
|
||||
def test_match_query_with_custom_node_label() -> None:
|
||||
match_query = get_match_query(
|
||||
return_properties=["name", "age"], node_label="`MyNodeLabel`"
|
||||
)
|
||||
expected = (
|
||||
"UNWIND $match_params AS match_param "
|
||||
"WITH match_param[0] AS match_id_value, match_param[1] AS score "
|
||||
"MATCH (node:`MyNodeLabel`) "
|
||||
"WHERE node[$id_property] = match_id_value "
|
||||
"RETURN node {.name, .age} AS node, labels(node) AS nodeLabels, elementId(node) AS elementId, elementId(node) AS id, score "
|
||||
)
|
||||
assert match_query.strip() == expected.strip()
|
||||
|
||||
|
||||
def test_weaviate_retriever_with_result_format_function(
|
||||
driver: MagicMock, neo4j_record: MagicMock, result_formatter: MagicMock
|
||||
) -> None:
|
||||
query_text = "may thy knife chip and shatter"
|
||||
top_k = 5
|
||||
node_id_value = "node-test-id"
|
||||
node_match_score = 0.9
|
||||
|
||||
wc = WClient(node_id_value=node_id_value, node_match_score=node_match_score)
|
||||
|
||||
retriever = WeaviateNeo4jRetriever(
|
||||
driver=driver,
|
||||
client=wc,
|
||||
collection="dummy-collection",
|
||||
id_property_neo4j="sync_id",
|
||||
id_property_external="neo4j_id",
|
||||
result_formatter=result_formatter,
|
||||
)
|
||||
driver.execute_query.return_value = [
|
||||
[neo4j_record],
|
||||
None,
|
||||
None,
|
||||
]
|
||||
|
||||
records = retriever.search(query_text=query_text, top_k=top_k)
|
||||
|
||||
assert records == RetrieverResult(
|
||||
items=[
|
||||
RetrieverResultItem(
|
||||
content="dummy-node", metadata={"score": 1.0, "node_id": 123}
|
||||
),
|
||||
],
|
||||
metadata={"__retriever": "WeaviateNeo4jRetriever"},
|
||||
)
|
||||
Reference in New Issue
Block a user