참고소스 수정본
This commit is contained in:
337
참고/neo4j-graphrag-python-main/tests/unit/tool/test_tool.py
Normal file
337
참고/neo4j-graphrag-python-main/tests/unit/tool/test_tool.py
Normal file
@@ -0,0 +1,337 @@
|
||||
import pytest
|
||||
from typing import Any
|
||||
from neo4j_graphrag.tool import (
|
||||
StringParameter,
|
||||
IntegerParameter,
|
||||
NumberParameter,
|
||||
BooleanParameter,
|
||||
ArrayParameter,
|
||||
ObjectParameter,
|
||||
Tool,
|
||||
ToolParameter,
|
||||
ParameterType,
|
||||
)
|
||||
|
||||
|
||||
def test_string_parameter() -> None:
|
||||
param = StringParameter(description="A string", required=True, enum=["a", "b"])
|
||||
assert param.description == "A string"
|
||||
assert param.required is True
|
||||
assert param.enum == ["a", "b"]
|
||||
d = param.model_dump_tool()
|
||||
assert d["type"] == ParameterType.STRING
|
||||
assert d["enum"] == ["a", "b"]
|
||||
# Note: 'required' is handled at the object level, not individual parameter level
|
||||
assert "required" not in d
|
||||
|
||||
|
||||
def test_integer_parameter() -> None:
|
||||
param = IntegerParameter(description="An int", minimum=0, maximum=10)
|
||||
d = param.model_dump_tool()
|
||||
assert d["type"] == ParameterType.INTEGER
|
||||
assert d["minimum"] == 0
|
||||
assert d["maximum"] == 10
|
||||
|
||||
|
||||
def test_number_parameter() -> None:
|
||||
param = NumberParameter(description="A number", minimum=1.5, maximum=3.5)
|
||||
d = param.model_dump_tool()
|
||||
assert d["type"] == ParameterType.NUMBER
|
||||
assert d["minimum"] == 1.5
|
||||
assert d["maximum"] == 3.5
|
||||
|
||||
|
||||
def test_boolean_parameter() -> None:
|
||||
param = BooleanParameter(description="A bool")
|
||||
d = param.model_dump_tool()
|
||||
assert d["type"] == ParameterType.BOOLEAN
|
||||
assert d["description"] == "A bool"
|
||||
|
||||
|
||||
def test_array_parameter_and_validation() -> None:
|
||||
arr_param = ArrayParameter(
|
||||
description="An array",
|
||||
items=StringParameter(description="str"),
|
||||
min_items=1,
|
||||
max_items=5,
|
||||
)
|
||||
d = arr_param.model_dump_tool()
|
||||
assert d["type"] == ParameterType.ARRAY
|
||||
assert d["items"]["type"] == ParameterType.STRING
|
||||
assert d["minItems"] == 1
|
||||
assert d["maxItems"] == 5
|
||||
|
||||
# Test items as dict
|
||||
arr_param2 = ArrayParameter(
|
||||
description="Arr with dict",
|
||||
items={"type": "string", "description": "str"}, # type: ignore
|
||||
)
|
||||
assert isinstance(arr_param2.items, StringParameter)
|
||||
|
||||
# Test error on invalid items
|
||||
with pytest.raises(ValueError):
|
||||
# Use type: ignore to bypass type checking for this intentional error case
|
||||
ArrayParameter(description="bad", items=123).validate_items() # type: ignore
|
||||
|
||||
|
||||
def test_object_parameter_and_validation() -> None:
|
||||
obj_param = ObjectParameter(
|
||||
description="Obj",
|
||||
properties={
|
||||
"foo": StringParameter(description="foo"),
|
||||
"bar": IntegerParameter(description="bar"),
|
||||
},
|
||||
required_properties=["foo"],
|
||||
additional_properties=False,
|
||||
)
|
||||
d = obj_param.model_dump_tool()
|
||||
assert d["type"] == ParameterType.OBJECT
|
||||
assert d["properties"]["foo"]["type"] == ParameterType.STRING
|
||||
assert d["required"] == ["foo"]
|
||||
assert d["additionalProperties"] is False
|
||||
|
||||
# Test properties as dicts
|
||||
obj_param2 = ObjectParameter(
|
||||
description="Obj2",
|
||||
properties={
|
||||
"foo": {"type": "string", "description": "foo"}, # type: ignore
|
||||
},
|
||||
)
|
||||
assert isinstance(obj_param2.properties["foo"], StringParameter)
|
||||
|
||||
# Test error on invalid property
|
||||
with pytest.raises(ValueError):
|
||||
# Use type: ignore to bypass type checking for this intentional error case
|
||||
ObjectParameter(
|
||||
description="bad",
|
||||
properties={"foo": 123}, # type: ignore
|
||||
).validate_properties()
|
||||
|
||||
|
||||
def test_from_dict() -> None:
|
||||
d = {"type": ParameterType.STRING, "description": "desc"}
|
||||
param = ToolParameter.from_dict(d)
|
||||
assert isinstance(param, StringParameter)
|
||||
assert param.description == "desc"
|
||||
|
||||
obj_dict = {
|
||||
"type": "object",
|
||||
"description": "obj",
|
||||
"properties": {"foo": {"type": "string", "description": "foo"}},
|
||||
}
|
||||
obj_param = ToolParameter.from_dict(obj_dict)
|
||||
assert isinstance(obj_param, ObjectParameter)
|
||||
assert isinstance(obj_param.properties["foo"], StringParameter)
|
||||
|
||||
arr_dict = {
|
||||
"type": "array",
|
||||
"description": "arr",
|
||||
"items": {"type": "integer", "description": "int"},
|
||||
}
|
||||
arr_param = ToolParameter.from_dict(arr_dict)
|
||||
assert isinstance(arr_param, ArrayParameter)
|
||||
assert isinstance(arr_param.items, IntegerParameter)
|
||||
|
||||
# Test unknown type
|
||||
with pytest.raises(ValueError):
|
||||
ToolParameter.from_dict({"type": "unknown", "description": "bad"})
|
||||
|
||||
# Test missing type
|
||||
with pytest.raises(ValueError):
|
||||
ToolParameter.from_dict({"description": "no type"})
|
||||
|
||||
|
||||
def test_required_parameter() -> None:
|
||||
# Test that individual parameters don't include 'required' field (it's handled at object level)
|
||||
string_param = StringParameter(description="Required string", required=True)
|
||||
assert "required" not in string_param.model_dump_tool()
|
||||
|
||||
integer_param = IntegerParameter(description="Required integer", required=True)
|
||||
assert "required" not in integer_param.model_dump_tool()
|
||||
|
||||
number_param = NumberParameter(description="Required number", required=True)
|
||||
assert "required" not in number_param.model_dump_tool()
|
||||
|
||||
boolean_param = BooleanParameter(description="Required boolean", required=True)
|
||||
assert "required" not in boolean_param.model_dump_tool()
|
||||
|
||||
array_param = ArrayParameter(
|
||||
description="Required array",
|
||||
items=StringParameter(description="item"),
|
||||
required=True,
|
||||
)
|
||||
assert "required" not in array_param.model_dump_tool()
|
||||
|
||||
object_param = ObjectParameter(
|
||||
description="Required object",
|
||||
properties={"prop": StringParameter(description="property")},
|
||||
required=True,
|
||||
)
|
||||
assert "required" not in object_param.model_dump_tool()
|
||||
|
||||
# Test that optional parameters also don't include the required field
|
||||
optional_param = StringParameter(description="Optional string", required=False)
|
||||
assert "required" not in optional_param.model_dump_tool()
|
||||
|
||||
|
||||
def test_object_parameter_additional_properties_always_present() -> None:
|
||||
"""Test that additionalProperties is always present in ObjectParameter schema, fixing OpenAI compatibility."""
|
||||
|
||||
# Test additionalProperties=True (default)
|
||||
obj_param_true = ObjectParameter(
|
||||
description="Object with additional properties",
|
||||
properties={"prop": StringParameter(description="A property")},
|
||||
additional_properties=True,
|
||||
)
|
||||
schema_true = obj_param_true.model_dump_tool()
|
||||
assert "additionalProperties" in schema_true
|
||||
assert schema_true["additionalProperties"] is True
|
||||
|
||||
# Test additionalProperties=False
|
||||
obj_param_false = ObjectParameter(
|
||||
description="Object without additional properties",
|
||||
properties={"prop": StringParameter(description="A property")},
|
||||
additional_properties=False,
|
||||
)
|
||||
schema_false = obj_param_false.model_dump_tool()
|
||||
assert "additionalProperties" in schema_false
|
||||
assert schema_false["additionalProperties"] is False
|
||||
|
||||
|
||||
def test_json_schema_compatibility() -> None:
|
||||
"""Test that the generated schema is compatible with JSON Schema specification."""
|
||||
|
||||
# Create a complex object with nested properties and required fields
|
||||
nested_obj = ObjectParameter(
|
||||
description="Nested object",
|
||||
properties={
|
||||
"nested_prop": StringParameter(description="Nested string"),
|
||||
},
|
||||
additional_properties=True,
|
||||
)
|
||||
|
||||
main_obj = ObjectParameter(
|
||||
description="Main object",
|
||||
properties={
|
||||
"required_string": StringParameter(description="Required string"),
|
||||
"optional_number": NumberParameter(description="Optional number"),
|
||||
"nested_object": nested_obj,
|
||||
},
|
||||
required_properties=["required_string"],
|
||||
additional_properties=False,
|
||||
)
|
||||
|
||||
schema = main_obj.model_dump_tool()
|
||||
|
||||
# Verify JSON Schema structure
|
||||
assert schema["type"] == "object"
|
||||
assert "properties" in schema
|
||||
assert "required" in schema
|
||||
assert "additionalProperties" in schema
|
||||
|
||||
# Check required is an array (not boolean on individual properties)
|
||||
assert isinstance(schema["required"], list)
|
||||
assert "required_string" in schema["required"]
|
||||
assert len(schema["required"]) == 1
|
||||
|
||||
# Check individual properties don't have 'required' field
|
||||
for prop_name, prop_schema in schema["properties"].items():
|
||||
assert "required" not in prop_schema
|
||||
|
||||
# Check additionalProperties is properly set at all levels
|
||||
assert schema["additionalProperties"] is False
|
||||
assert schema["properties"]["nested_object"]["additionalProperties"] is True
|
||||
|
||||
|
||||
def test_text2cypher_retriever_schema_compatibility() -> None:
|
||||
"""Test the specific schema structure that caused the OpenAI API error."""
|
||||
|
||||
# Simulate the Text2CypherRetriever parameter structure
|
||||
prompt_params = ObjectParameter(
|
||||
description="Parameter prompt_params",
|
||||
properties={},
|
||||
additional_properties=True, # This was missing in the original bug
|
||||
)
|
||||
|
||||
t2c_params = ObjectParameter(
|
||||
description="Parameters for Text2CypherRetriever",
|
||||
properties={
|
||||
"query_text": StringParameter(description="Parameter query_text"),
|
||||
"prompt_params": prompt_params,
|
||||
},
|
||||
required_properties=["query_text"],
|
||||
additional_properties=False,
|
||||
)
|
||||
|
||||
schema = t2c_params.model_dump_tool()
|
||||
|
||||
# Verify the fix: prompt_params should have additionalProperties
|
||||
prompt_params_schema = schema["properties"]["prompt_params"]
|
||||
assert "additionalProperties" in prompt_params_schema
|
||||
assert prompt_params_schema["additionalProperties"] is True
|
||||
|
||||
# Verify query_text doesn't have individual 'required' field
|
||||
query_text_schema = schema["properties"]["query_text"]
|
||||
assert "required" not in query_text_schema
|
||||
|
||||
# Verify required array at object level
|
||||
assert schema["required"] == ["query_text"]
|
||||
|
||||
|
||||
def test_exclude_parameter_in_object_schema() -> None:
|
||||
"""Test that exclude parameter works correctly in ObjectParameter.model_dump_tool()."""
|
||||
|
||||
obj_param = ObjectParameter(
|
||||
description="Test object",
|
||||
properties={
|
||||
"prop1": StringParameter(description="Property 1"),
|
||||
"prop2": IntegerParameter(description="Property 2"),
|
||||
},
|
||||
required_properties=["prop1"],
|
||||
additional_properties=True,
|
||||
)
|
||||
|
||||
# Test excluding required field
|
||||
schema_no_required = obj_param.model_dump_tool(exclude=["required"])
|
||||
assert "required" not in schema_no_required
|
||||
assert "additionalProperties" in schema_no_required # Should still be present
|
||||
|
||||
# Test excluding additionalProperties field
|
||||
schema_no_additional = obj_param.model_dump_tool(exclude=["additional_properties"])
|
||||
assert "additionalProperties" not in schema_no_additional
|
||||
assert "required" in schema_no_additional # Should still be present
|
||||
|
||||
|
||||
def test_tool_class() -> None:
|
||||
def dummy_func(**kwargs: Any) -> dict[str, Any]:
|
||||
return kwargs
|
||||
|
||||
params = ObjectParameter(
|
||||
description="params",
|
||||
properties={"a": StringParameter(description="a")},
|
||||
)
|
||||
tool = Tool(
|
||||
name="mytool",
|
||||
description="desc",
|
||||
parameters=params,
|
||||
execute_func=dummy_func,
|
||||
)
|
||||
assert tool.get_name() == "mytool"
|
||||
assert tool.get_description() == "desc"
|
||||
assert tool.get_parameters()["type"] == ParameterType.OBJECT
|
||||
assert tool.execute(query="query", a="b") == {"query": "query", "a": "b"}
|
||||
|
||||
# Test parameters as dict
|
||||
params_dict = {
|
||||
"type": "object",
|
||||
"description": "params",
|
||||
"properties": {"a": {"type": "string", "description": "a"}},
|
||||
}
|
||||
tool2 = Tool(
|
||||
name="mytool2",
|
||||
description="desc2",
|
||||
parameters=params_dict,
|
||||
execute_func=dummy_func,
|
||||
)
|
||||
assert tool2.get_parameters()["type"] == ParameterType.OBJECT
|
||||
assert tool2.execute(a="b") == {"a": "b"}
|
||||
@@ -0,0 +1,483 @@
|
||||
# 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.mock import MagicMock, patch
|
||||
import neo4j
|
||||
from neo4j_graphrag.embeddings.base import Embedder
|
||||
from neo4j_graphrag.llm.base import LLMInterface
|
||||
from neo4j_graphrag.retrievers import (
|
||||
HybridCypherRetriever,
|
||||
HybridRetriever,
|
||||
Text2CypherRetriever,
|
||||
VectorCypherRetriever,
|
||||
VectorRetriever,
|
||||
)
|
||||
from neo4j_graphrag.tool import Tool
|
||||
|
||||
|
||||
# Mock dependencies for retriever instances
|
||||
def create_mock_driver() -> neo4j.Driver:
|
||||
driver = MagicMock(spec=neo4j.Driver)
|
||||
# Create a mock result object with a records attribute
|
||||
mock_result = MagicMock()
|
||||
mock_result.records = [MagicMock()]
|
||||
driver.execute_query.return_value = mock_result
|
||||
return driver
|
||||
|
||||
|
||||
def create_mock_embedder() -> Embedder:
|
||||
embedder = MagicMock(spec=Embedder)
|
||||
embedder.embed_query.return_value = [0.1, 0.2, 0.3]
|
||||
return embedder
|
||||
|
||||
|
||||
def create_mock_llm() -> LLMInterface:
|
||||
llm = MagicMock()
|
||||
llm.invoke.return_value = "MATCH (n) RETURN n"
|
||||
return llm
|
||||
|
||||
|
||||
# Test conversion with VectorRetriever
|
||||
@patch("neo4j_graphrag.retrievers.base.get_version")
|
||||
def test_convert_vector_retriever_to_tool(mock_get_version: MagicMock) -> None:
|
||||
"""Test conversion of VectorRetriever to a Tool instance with correct attributes."""
|
||||
mock_get_version.return_value = ((5, 20, 0), False, False)
|
||||
driver = create_mock_driver()
|
||||
embedder = create_mock_embedder()
|
||||
retriever = VectorRetriever(
|
||||
driver=driver,
|
||||
index_name="test_index",
|
||||
embedder=embedder,
|
||||
return_properties=["name", "description"],
|
||||
)
|
||||
tool = retriever.convert_to_tool(
|
||||
name="VectorRetriever",
|
||||
description="A tool for vector-based retrieval from Neo4j.",
|
||||
parameter_descriptions={
|
||||
"query_text": "The query text for vector search.",
|
||||
"top_k": "Number of results to return.",
|
||||
},
|
||||
)
|
||||
assert isinstance(tool, Tool)
|
||||
assert tool.get_name() == "VectorRetriever"
|
||||
assert tool.get_description() == "A tool for vector-based retrieval from Neo4j."
|
||||
# Check that the parameters object has the expected properties
|
||||
params = tool.get_parameters()
|
||||
assert "properties" in params
|
||||
assert len(params["properties"]) == 5 # VectorRetriever has 5 parameters
|
||||
assert "query_text" in params["properties"]
|
||||
assert "top_k" in params["properties"]
|
||||
assert "query_vector" in params["properties"]
|
||||
assert "effective_search_ratio" in params["properties"]
|
||||
assert "filters" in params["properties"]
|
||||
|
||||
|
||||
# Test conversion with VectorCypherRetriever
|
||||
@patch("neo4j_graphrag.retrievers.base.get_version")
|
||||
def test_convert_vector_cypher_retriever_to_tool(mock_get_version: MagicMock) -> None:
|
||||
"""Test conversion of VectorCypherRetriever to a Tool instance with correct attributes."""
|
||||
mock_get_version.return_value = ((5, 20, 0), False, False)
|
||||
driver = create_mock_driver()
|
||||
embedder = create_mock_embedder()
|
||||
retriever = VectorCypherRetriever(
|
||||
driver=driver,
|
||||
index_name="test_index",
|
||||
embedder=embedder,
|
||||
retrieval_query="RETURN n",
|
||||
)
|
||||
tool = retriever.convert_to_tool(
|
||||
name="VectorCypherRetriever",
|
||||
description="A tool for vector-cypher retrieval from Neo4j.",
|
||||
parameter_descriptions={
|
||||
"query_text": "The query text for vector-cypher search.",
|
||||
"top_k": "Number of results to return.",
|
||||
},
|
||||
)
|
||||
assert isinstance(tool, Tool)
|
||||
assert tool.get_name() == "VectorCypherRetriever"
|
||||
assert tool.get_description() == "A tool for vector-cypher retrieval from Neo4j."
|
||||
# Check that the parameters object has the expected properties
|
||||
params = tool.get_parameters()
|
||||
assert "properties" in params
|
||||
assert len(params["properties"]) == 6 # VectorCypherRetriever has 6 parameters
|
||||
assert "query_text" in params["properties"]
|
||||
assert "top_k" in params["properties"]
|
||||
assert "query_vector" in params["properties"]
|
||||
assert "effective_search_ratio" in params["properties"]
|
||||
assert "query_params" in params["properties"]
|
||||
assert "filters" in params["properties"]
|
||||
|
||||
|
||||
# Test conversion with HybridRetriever
|
||||
@patch("neo4j_graphrag.retrievers.base.get_version")
|
||||
def test_convert_hybrid_retriever_to_tool(mock_get_version: MagicMock) -> None:
|
||||
"""Test conversion of HybridRetriever to a Tool instance with correct attributes."""
|
||||
mock_get_version.return_value = ((5, 20, 0), False, False)
|
||||
driver = create_mock_driver()
|
||||
embedder = create_mock_embedder()
|
||||
retriever = HybridRetriever(
|
||||
driver=driver,
|
||||
vector_index_name="test_vector_index",
|
||||
fulltext_index_name="test_fulltext_index",
|
||||
embedder=embedder,
|
||||
return_properties=["name", "description"],
|
||||
)
|
||||
tool = retriever.convert_to_tool(
|
||||
name="HybridRetriever",
|
||||
description="A tool for hybrid retrieval from Neo4j.",
|
||||
parameter_descriptions={
|
||||
"query_text": "The query text for hybrid search.",
|
||||
"top_k": "Number of results to return.",
|
||||
},
|
||||
)
|
||||
assert isinstance(tool, Tool)
|
||||
assert tool.get_name() == "HybridRetriever"
|
||||
assert tool.get_description() == "A tool for hybrid retrieval from Neo4j."
|
||||
# Check that the parameters object has the expected properties
|
||||
params = tool.get_parameters()
|
||||
assert "properties" in params
|
||||
assert len(params["properties"]) == 6 # HybridRetriever has 6 parameters
|
||||
assert "query_text" in params["properties"]
|
||||
assert "top_k" in params["properties"]
|
||||
assert "query_vector" in params["properties"]
|
||||
assert "effective_search_ratio" in params["properties"]
|
||||
assert "ranker" in params["properties"]
|
||||
assert "alpha" in params["properties"]
|
||||
|
||||
|
||||
# Test conversion with HybridCypherRetriever
|
||||
@patch("neo4j_graphrag.retrievers.base.get_version")
|
||||
def test_convert_hybrid_cypher_retriever_to_tool(mock_get_version: MagicMock) -> None:
|
||||
"""Test conversion of HybridCypherRetriever to a Tool instance with correct attributes."""
|
||||
mock_get_version.return_value = ((5, 20, 0), False, False)
|
||||
driver = create_mock_driver()
|
||||
embedder = create_mock_embedder()
|
||||
retriever = HybridCypherRetriever(
|
||||
driver=driver,
|
||||
vector_index_name="test_vector_index",
|
||||
fulltext_index_name="test_fulltext_index",
|
||||
embedder=embedder,
|
||||
retrieval_query="RETURN n",
|
||||
)
|
||||
tool = retriever.convert_to_tool(
|
||||
name="HybridCypherRetriever",
|
||||
description="A tool for hybrid-cypher retrieval from Neo4j.",
|
||||
parameter_descriptions={
|
||||
"query_text": "The query text for hybrid-cypher search.",
|
||||
"top_k": "Number of results to return.",
|
||||
},
|
||||
)
|
||||
assert isinstance(tool, Tool)
|
||||
assert tool.get_name() == "HybridCypherRetriever"
|
||||
assert tool.get_description() == "A tool for hybrid-cypher retrieval from Neo4j."
|
||||
# Check that the parameters object has the expected properties
|
||||
params = tool.get_parameters()
|
||||
assert "properties" in params
|
||||
assert len(params["properties"]) == 7 # HybridCypherRetriever has 7 parameters
|
||||
assert "query_text" in params["properties"]
|
||||
assert "query_vector" in params["properties"]
|
||||
assert "top_k" in params["properties"]
|
||||
assert "effective_search_ratio" in params["properties"]
|
||||
assert "query_params" in params["properties"]
|
||||
assert "ranker" in params["properties"]
|
||||
assert "alpha" in params["properties"]
|
||||
|
||||
|
||||
# Test conversion with Text2CypherRetriever
|
||||
@patch("neo4j_graphrag.retrievers.base.get_version")
|
||||
def test_convert_text2cypher_retriever_to_tool(mock_get_version: MagicMock) -> None:
|
||||
"""Test conversion of Text2CypherRetriever to a Tool instance with correct attributes."""
|
||||
mock_get_version.return_value = ((5, 20, 0), False, False)
|
||||
driver = create_mock_driver()
|
||||
llm = create_mock_llm()
|
||||
retriever = Text2CypherRetriever(driver=driver, llm=llm)
|
||||
tool = retriever.convert_to_tool(
|
||||
name="Text2CypherRetriever",
|
||||
description="A tool for text to Cypher retrieval from Neo4j.",
|
||||
parameter_descriptions={
|
||||
"query_text": "The query text for text to Cypher conversion.",
|
||||
},
|
||||
)
|
||||
assert isinstance(tool, Tool)
|
||||
assert tool.get_name() == "Text2CypherRetriever"
|
||||
assert tool.get_description() == "A tool for text to Cypher retrieval from Neo4j."
|
||||
# Check that the parameters object has the expected properties
|
||||
params = tool.get_parameters()
|
||||
assert "properties" in params
|
||||
assert len(params["properties"]) == 2 # Text2CypherRetriever has 2 parameters
|
||||
assert "query_text" in params["properties"]
|
||||
assert "prompt_params" in params["properties"]
|
||||
|
||||
|
||||
# Test conversion with custom name provided
|
||||
@patch("neo4j_graphrag.retrievers.base.get_version")
|
||||
def test_convert_retriever_with_custom_name(
|
||||
mock_get_version: MagicMock,
|
||||
) -> None:
|
||||
"""Test conversion of a retriever to a Tool instance with a custom name."""
|
||||
mock_get_version.return_value = ((5, 20, 0), False, False)
|
||||
driver = create_mock_driver()
|
||||
embedder = create_mock_embedder()
|
||||
retriever = VectorRetriever(
|
||||
driver=driver,
|
||||
index_name="test_index",
|
||||
embedder=embedder,
|
||||
return_properties=["name", "description"],
|
||||
)
|
||||
|
||||
custom_name = "CustomNamedTool"
|
||||
|
||||
tool = retriever.convert_to_tool(
|
||||
name=custom_name,
|
||||
description="A tool with a custom name",
|
||||
parameter_descriptions={
|
||||
"query_text": "The query text for vector search.",
|
||||
},
|
||||
)
|
||||
|
||||
# Verify that the custom name is used instead of the retriever class name
|
||||
assert tool.get_name() == custom_name
|
||||
|
||||
|
||||
# Test conversion with no parameters provided
|
||||
@patch("neo4j_graphrag.retrievers.base.get_version")
|
||||
def test_convert_vector_retriever_to_tool_no_parameters(
|
||||
mock_get_version: MagicMock,
|
||||
) -> None:
|
||||
"""Test conversion of VectorRetriever to a Tool instance when no parameters are provided."""
|
||||
mock_get_version.return_value = ((5, 20, 0), False, False)
|
||||
driver = create_mock_driver()
|
||||
embedder = create_mock_embedder()
|
||||
retriever = VectorRetriever(
|
||||
driver=driver,
|
||||
index_name="test_index",
|
||||
embedder=embedder,
|
||||
return_properties=["name", "description"],
|
||||
)
|
||||
tool = retriever.convert_to_tool(
|
||||
name="VectorRetriever",
|
||||
description="A tool for vector-based retrieval from Neo4j.",
|
||||
)
|
||||
assert isinstance(tool, Tool)
|
||||
assert tool.get_name() == "VectorRetriever"
|
||||
assert tool.get_description() == "A tool for vector-based retrieval from Neo4j."
|
||||
# With the new API, parameters are always auto-inferred from method signature
|
||||
params = tool.get_parameters()
|
||||
assert params is not None
|
||||
assert "properties" in params
|
||||
assert len(params["properties"]) == 5 # VectorRetriever has 5 parameters
|
||||
|
||||
|
||||
# Test tool execution for VectorRetriever
|
||||
@patch("neo4j_graphrag.retrievers.base.get_version")
|
||||
def test_vector_retriever_tool_execution(mock_get_version: MagicMock) -> None:
|
||||
"""Test execution of VectorRetriever tool calls the search method with correct arguments."""
|
||||
mock_get_version.return_value = ((5, 20, 0), False, False)
|
||||
driver = create_mock_driver()
|
||||
embedder = create_mock_embedder()
|
||||
retriever = VectorRetriever(
|
||||
driver=driver,
|
||||
index_name="test_index",
|
||||
embedder=embedder,
|
||||
return_properties=["name", "description"],
|
||||
)
|
||||
# Create the tool first, before mocking
|
||||
with patch.object(VectorRetriever, "_fetch_index_infos"):
|
||||
tool = retriever.convert_to_tool(
|
||||
name="VectorRetriever",
|
||||
description="A tool for vector-based retrieval from Neo4j.",
|
||||
parameter_descriptions={
|
||||
"query_text": "The query text for vector search.",
|
||||
"top_k": "Number of results to return.",
|
||||
},
|
||||
)
|
||||
|
||||
# Now mock the get_search_results method to track calls
|
||||
from neo4j_graphrag.types import RawSearchResult
|
||||
|
||||
get_search_results_mock = MagicMock(
|
||||
return_value=RawSearchResult(records=[], metadata={})
|
||||
)
|
||||
# Use patch to mock the method
|
||||
with patch.object(retriever, "get_search_results", get_search_results_mock):
|
||||
tools = {tool.get_name(): tool}
|
||||
# Simulate indirect invocation as would happen in real usage
|
||||
tool_call_arguments = {"query_text": "test query", "top_k": 5}
|
||||
# Pass the arguments as kwargs
|
||||
result = tools[tool.get_name()].execute(**tool_call_arguments)
|
||||
|
||||
# Since we're using a context manager for patching, we need to verify the call inside the context
|
||||
# We can only check the result, not the method call itself
|
||||
assert result is not None
|
||||
assert hasattr(result, "items") # Should return RetrieverResult now
|
||||
assert isinstance(result.items, list)
|
||||
assert hasattr(result, "metadata")
|
||||
|
||||
|
||||
# Test tool execution for HybridRetriever
|
||||
@patch("neo4j_graphrag.retrievers.base.get_version")
|
||||
def test_hybrid_retriever_tool_execution(mock_get_version: MagicMock) -> None:
|
||||
"""Test execution of HybridRetriever tool calls the search method with correct arguments."""
|
||||
mock_get_version.return_value = ((5, 20, 0), False, False)
|
||||
driver = create_mock_driver()
|
||||
embedder = create_mock_embedder()
|
||||
retriever = HybridRetriever(
|
||||
driver=driver,
|
||||
vector_index_name="test_vector_index",
|
||||
fulltext_index_name="test_fulltext_index",
|
||||
embedder=embedder,
|
||||
return_properties=["name", "description"],
|
||||
)
|
||||
# Create the tool first, before mocking
|
||||
with patch.object(HybridRetriever, "_fetch_index_infos"):
|
||||
tool = retriever.convert_to_tool(
|
||||
name="HybridRetriever",
|
||||
description="A tool for hybrid retrieval from Neo4j.",
|
||||
parameter_descriptions={
|
||||
"query_text": "The query text for hybrid search.",
|
||||
"top_k": "Number of results to return.",
|
||||
},
|
||||
)
|
||||
|
||||
# Now mock the get_search_results method to track calls
|
||||
from neo4j_graphrag.types import RawSearchResult
|
||||
|
||||
get_search_results_mock = MagicMock(
|
||||
return_value=RawSearchResult(records=[], metadata={})
|
||||
)
|
||||
# Use patch to mock the method
|
||||
with patch.object(retriever, "get_search_results", get_search_results_mock):
|
||||
tools = {tool.get_name(): tool}
|
||||
# Simulate indirect invocation as would happen in real usage
|
||||
tool_call_arguments = {"query_text": "test query", "top_k": 5}
|
||||
# Pass the arguments as kwargs
|
||||
result = tools[tool.get_name()].execute(**tool_call_arguments)
|
||||
|
||||
# Since we're using a context manager for patching, we need to verify the call inside the context
|
||||
# We can only check the result, not the method call itself
|
||||
assert result is not None
|
||||
assert hasattr(result, "items") # Should return RetrieverResult now
|
||||
assert isinstance(result.items, list)
|
||||
assert hasattr(result, "metadata")
|
||||
|
||||
|
||||
# Test tool execution for Text2CypherRetriever
|
||||
@patch("neo4j_graphrag.retrievers.base.get_version")
|
||||
def test_text2cypher_retriever_tool_execution(mock_get_version: MagicMock) -> None:
|
||||
"""Test execution of Text2CypherRetriever tool calls the search method with correct arguments."""
|
||||
mock_get_version.return_value = ((5, 20, 0), False, False)
|
||||
driver = create_mock_driver()
|
||||
llm = create_mock_llm()
|
||||
retriever = Text2CypherRetriever(driver=driver, llm=llm)
|
||||
# Create the tool first, before mocking
|
||||
tool = retriever.convert_to_tool(
|
||||
name="Text2CypherRetriever",
|
||||
description="A tool for text to Cypher retrieval from Neo4j.",
|
||||
parameter_descriptions={
|
||||
"query_text": "The query text for text to Cypher conversion.",
|
||||
},
|
||||
)
|
||||
|
||||
# Now mock the get_search_results method to track calls
|
||||
from neo4j_graphrag.types import RawSearchResult
|
||||
|
||||
get_search_results_mock = MagicMock(
|
||||
return_value=RawSearchResult(records=[], metadata={})
|
||||
)
|
||||
# Use patch to mock the method
|
||||
with patch.object(retriever, "get_search_results", get_search_results_mock):
|
||||
tools = {tool.get_name(): tool}
|
||||
# Simulate indirect invocation as would happen in real usage
|
||||
tool_call_arguments = {"query_text": "test query"}
|
||||
# Pass the arguments as kwargs
|
||||
result = tools[tool.get_name()].execute(**tool_call_arguments)
|
||||
|
||||
# Since we're using a context manager for patching, we need to verify the call inside the context
|
||||
# We can only check the result, not the method call itself
|
||||
assert result is not None
|
||||
assert hasattr(result, "items") # Should return RetrieverResult now
|
||||
assert isinstance(result.items, list)
|
||||
assert hasattr(result, "metadata")
|
||||
|
||||
|
||||
# Test tool serialization to JSON format
|
||||
@patch("neo4j_graphrag.retrievers.base.get_version")
|
||||
def test_tool_serialization(mock_get_version: MagicMock) -> None:
|
||||
"""Test that a Tool instance can be serialized to the required JSON format."""
|
||||
mock_get_version.return_value = ((5, 20, 0), False, False)
|
||||
driver = create_mock_driver()
|
||||
embedder = create_mock_embedder()
|
||||
retriever = VectorRetriever(
|
||||
driver=driver,
|
||||
index_name="test_index",
|
||||
embedder=embedder,
|
||||
return_properties=["name", "description"],
|
||||
)
|
||||
tool = retriever.convert_to_tool(
|
||||
name="VectorRetriever",
|
||||
description="A tool for vector-based retrieval from Neo4j.",
|
||||
parameter_descriptions={
|
||||
"query_text": "The query text for vector search.",
|
||||
"top_k": "Number of results to return.",
|
||||
},
|
||||
)
|
||||
# Create a dictionary representation of the tool
|
||||
tool_dict = {
|
||||
"type": "function",
|
||||
"name": tool.get_name(),
|
||||
"description": tool.get_description(),
|
||||
"parameters": tool.get_parameters(),
|
||||
}
|
||||
|
||||
assert tool_dict["type"] == "function"
|
||||
assert tool_dict["name"] == tool.get_name()
|
||||
assert tool_dict["description"] == tool.get_description()
|
||||
assert "parameters" in tool_dict
|
||||
|
||||
# Get parameters and convert to dictionary
|
||||
parameters_any = tool_dict["parameters"]
|
||||
# With the new API, parameters should be a dictionary
|
||||
if isinstance(parameters_any, dict):
|
||||
parameters_dict = parameters_any
|
||||
else:
|
||||
# Handle unexpected parameter format
|
||||
parameters_dict = {
|
||||
str(k): v for k, v in enumerate(parameters_any) if v is not None
|
||||
}
|
||||
|
||||
# Check the parameters structure
|
||||
assert parameters_dict.get("type") == "object"
|
||||
assert "properties" in parameters_dict
|
||||
|
||||
# Check that we have the expected parameter properties
|
||||
# VectorRetriever has all optional parameters (query_vector and query_text are both optional)
|
||||
expected_properties = {
|
||||
"query_vector",
|
||||
"query_text",
|
||||
"top_k",
|
||||
"effective_search_ratio",
|
||||
"filters",
|
||||
}
|
||||
actual_properties = set(parameters_dict.get("properties", {}).keys())
|
||||
assert (
|
||||
expected_properties == actual_properties
|
||||
), f"Expected {expected_properties}, got {actual_properties}"
|
||||
|
||||
# Check additionalProperties if it exists
|
||||
if "additionalProperties" in parameters_dict and not parameters_dict.get(
|
||||
"additionalProperties"
|
||||
):
|
||||
pass # This line is just to satisfy the test, actual check is visual
|
||||
Reference in New Issue
Block a user