mirror of
https://github.com/langgenius/dify.git
synced 2026-09-24 23:22:26 +08:00
refactor: pass session into hit testing service (#37785)
This commit is contained in:
@@ -181,7 +181,9 @@ class TestHitTestingService:
|
||||
# ── Response formatting ────────────────────────────────────────────
|
||||
|
||||
@patch("core.rag.datasource.retrieval_service.RetrievalService.format_retrieval_documents")
|
||||
def test_compact_retrieve_response_should_format_correctly(self, mock_format: MagicMock) -> None:
|
||||
def test_compact_retrieve_response_should_format_correctly(
|
||||
self, mock_format: MagicMock, db_session_with_containers: Session
|
||||
) -> None:
|
||||
query = "test query"
|
||||
mock_doc = MagicMock(spec=Document)
|
||||
|
||||
@@ -189,7 +191,9 @@ class TestHitTestingService:
|
||||
mock_record.model_dump.return_value = {"content": "formatted content"}
|
||||
mock_format.return_value = [mock_record]
|
||||
|
||||
response = _RetrieveResponse.model_validate(HitTestingService.compact_retrieve_response(query, [mock_doc]))
|
||||
response = _RetrieveResponse.model_validate(
|
||||
HitTestingService.compact_retrieve_response(db_session_with_containers, query, [mock_doc])
|
||||
)
|
||||
|
||||
assert response.query.content == query
|
||||
assert len(response.records) == 1
|
||||
@@ -242,6 +246,7 @@ class TestHitTestingService:
|
||||
|
||||
response = _RetrieveResponse.model_validate(
|
||||
HitTestingService.external_retrieve(
|
||||
db_session_with_containers,
|
||||
dataset=dataset,
|
||||
query='test "query"',
|
||||
account=account,
|
||||
@@ -269,7 +274,9 @@ class TestHitTestingService:
|
||||
dataset = _create_dataset(db_session_with_containers, provider="vendor")
|
||||
account = MagicMock()
|
||||
|
||||
response = _RetrieveResponse.model_validate(HitTestingService.external_retrieve(dataset, "test query", account))
|
||||
response = _RetrieveResponse.model_validate(
|
||||
HitTestingService.external_retrieve(db_session_with_containers, dataset, "test query", account)
|
||||
)
|
||||
|
||||
assert response.query.content == "test query"
|
||||
assert response.records == []
|
||||
@@ -292,6 +299,7 @@ class TestHitTestingService:
|
||||
|
||||
response = _RetrieveResponse.model_validate(
|
||||
HitTestingService.retrieve(
|
||||
db_session_with_containers,
|
||||
dataset=dataset,
|
||||
query="test query",
|
||||
account=account,
|
||||
@@ -320,7 +328,11 @@ class TestHitTestingService:
|
||||
|
||||
retrieval_model = {
|
||||
"search_method": "semantic_search",
|
||||
"metadata_filtering_conditions": {"some": "condition"},
|
||||
"metadata_filtering_conditions": {
|
||||
"conditions": [
|
||||
{"name": "category", "comparison_operator": "is", "value": "test"},
|
||||
],
|
||||
},
|
||||
"top_k": 5,
|
||||
"reranking_enable": False,
|
||||
"score_threshold_enabled": False,
|
||||
@@ -330,6 +342,7 @@ class TestHitTestingService:
|
||||
mock_retrieve.return_value = retrieved_documents
|
||||
|
||||
HitTestingService.retrieve(
|
||||
db_session_with_containers,
|
||||
dataset=dataset,
|
||||
query="test query",
|
||||
account=account,
|
||||
@@ -352,7 +365,11 @@ class TestHitTestingService:
|
||||
|
||||
retrieval_model = {
|
||||
"search_method": "semantic_search",
|
||||
"metadata_filtering_conditions": {"some": "condition"},
|
||||
"metadata_filtering_conditions": {
|
||||
"conditions": [
|
||||
{"name": "category", "comparison_operator": "is", "value": "test"},
|
||||
],
|
||||
},
|
||||
"top_k": 5,
|
||||
"reranking_enable": False,
|
||||
"score_threshold_enabled": False,
|
||||
@@ -362,6 +379,7 @@ class TestHitTestingService:
|
||||
|
||||
response = _RetrieveResponse.model_validate(
|
||||
HitTestingService.retrieve(
|
||||
db_session_with_containers,
|
||||
dataset=dataset,
|
||||
query="test query",
|
||||
account=account,
|
||||
@@ -393,6 +411,7 @@ class TestHitTestingService:
|
||||
mock_retrieve.return_value = retrieved_documents
|
||||
|
||||
HitTestingService.retrieve(
|
||||
db_session_with_containers,
|
||||
dataset=dataset,
|
||||
query="test query",
|
||||
account=account,
|
||||
@@ -452,6 +471,7 @@ class TestHitTestingService:
|
||||
mock_retrieve.return_value = retrieved_documents
|
||||
|
||||
HitTestingService.retrieve(
|
||||
db_session_with_containers,
|
||||
dataset=dataset,
|
||||
query="test query",
|
||||
account=account,
|
||||
@@ -477,11 +497,15 @@ class TestHitTestingService:
|
||||
"doc_metadata": {"source": "manual"},
|
||||
}
|
||||
|
||||
def test_dump_retrieval_records_returns_dumped_records_without_document_ids(self) -> None:
|
||||
def test_dump_retrieval_records_returns_dumped_records_without_document_ids(
|
||||
self, db_session_with_containers: Session
|
||||
) -> None:
|
||||
segment = _build_segment(document_id="")
|
||||
record = RetrievalSegments.model_validate({"segment": segment, "score": 0.95})
|
||||
|
||||
records = _DUMPED_RETRIEVAL_RECORDS.validate_python(HitTestingService._dump_retrieval_records([record]))
|
||||
records = _DUMPED_RETRIEVAL_RECORDS.validate_python(
|
||||
HitTestingService._dump_retrieval_records(db_session_with_containers, [record])
|
||||
)
|
||||
|
||||
assert len(records) == 1
|
||||
assert records[0].segment.id == segment.id
|
||||
@@ -493,7 +517,9 @@ class TestHitTestingService:
|
||||
segment = _create_segment(db_session_with_containers, document=document)
|
||||
record = RetrievalSegments.model_validate({"segment": segment, "score": 0.9})
|
||||
|
||||
records = _DUMPED_RETRIEVAL_RECORDS.validate_python(HitTestingService._dump_retrieval_records([record]))
|
||||
records = _DUMPED_RETRIEVAL_RECORDS.validate_python(
|
||||
HitTestingService._dump_retrieval_records(db_session_with_containers, [record])
|
||||
)
|
||||
|
||||
assert len(records) == 1
|
||||
dumped_segment = records[0].segment
|
||||
@@ -515,7 +541,7 @@ class TestHitTestingService:
|
||||
segment = _create_segment(db_session_with_containers)
|
||||
record = RetrievalSegments.model_validate({"segment": segment, "score": 0.95})
|
||||
|
||||
result = HitTestingService._dump_retrieval_records([record])
|
||||
result = HitTestingService._dump_retrieval_records(db_session_with_containers, [record])
|
||||
|
||||
assert result == []
|
||||
assert "Skipping hit-testing records with missing documents" in caplog.text
|
||||
|
||||
@@ -147,8 +147,7 @@ class TestHitTestingServiceRetrieve:
|
||||
Provides a mocked database session for testing database operations
|
||||
like adding and committing DatasetQuery records.
|
||||
"""
|
||||
with patch("services.hit_testing_service.db.session", autospec=True) as mock_db:
|
||||
yield mock_db
|
||||
return MagicMock()
|
||||
|
||||
def test_retrieve_success_with_default_retrieval_model(self, mock_db_session):
|
||||
"""
|
||||
@@ -186,7 +185,9 @@ class TestHitTestingServiceRetrieve:
|
||||
mock_format.return_value = mock_records
|
||||
|
||||
# Act
|
||||
result = HitTestingService.retrieve(dataset, query, account, retrieval_model, external_retrieval_model)
|
||||
result = HitTestingService.retrieve(
|
||||
mock_db_session, dataset, query, account, retrieval_model, external_retrieval_model
|
||||
)
|
||||
|
||||
# Assert
|
||||
assert result["query"]["content"] == query
|
||||
@@ -232,7 +233,9 @@ class TestHitTestingServiceRetrieve:
|
||||
mock_format.return_value = mock_records
|
||||
|
||||
# Act
|
||||
result = HitTestingService.retrieve(dataset, query, account, retrieval_model, external_retrieval_model)
|
||||
result = HitTestingService.retrieve(
|
||||
mock_db_session, dataset, query, account, retrieval_model, external_retrieval_model
|
||||
)
|
||||
|
||||
# Assert
|
||||
assert result["query"]["content"] == query
|
||||
@@ -257,9 +260,11 @@ class TestHitTestingServiceRetrieve:
|
||||
retrieval_model = {
|
||||
"metadata_filtering_conditions": {
|
||||
"conditions": [
|
||||
{"field": "category", "operator": "is", "value": "test"},
|
||||
{"name": "category", "comparison_operator": "is", "value": "test"},
|
||||
],
|
||||
},
|
||||
"reranking_enable": False,
|
||||
"score_threshold_enabled": False,
|
||||
}
|
||||
external_retrieval_model = {}
|
||||
|
||||
@@ -286,7 +291,9 @@ class TestHitTestingServiceRetrieve:
|
||||
mock_format.return_value = mock_records
|
||||
|
||||
# Act
|
||||
result = HitTestingService.retrieve(dataset, query, account, retrieval_model, external_retrieval_model)
|
||||
result = HitTestingService.retrieve(
|
||||
mock_db_session, dataset, query, account, retrieval_model, external_retrieval_model
|
||||
)
|
||||
|
||||
# Assert
|
||||
assert result["query"]["content"] == query
|
||||
@@ -308,9 +315,11 @@ class TestHitTestingServiceRetrieve:
|
||||
retrieval_model = {
|
||||
"metadata_filtering_conditions": {
|
||||
"conditions": [
|
||||
{"field": "category", "operator": "is", "value": "test"},
|
||||
{"name": "category", "comparison_operator": "is", "value": "test"},
|
||||
],
|
||||
},
|
||||
"reranking_enable": False,
|
||||
"score_threshold_enabled": False,
|
||||
}
|
||||
external_retrieval_model = {}
|
||||
|
||||
@@ -327,7 +336,9 @@ class TestHitTestingServiceRetrieve:
|
||||
mock_format.return_value = []
|
||||
|
||||
# Act
|
||||
result = HitTestingService.retrieve(dataset, query, account, retrieval_model, external_retrieval_model)
|
||||
result = HitTestingService.retrieve(
|
||||
mock_db_session, dataset, query, account, retrieval_model, external_retrieval_model
|
||||
)
|
||||
|
||||
# Assert
|
||||
assert result["query"]["content"] == query
|
||||
@@ -344,6 +355,8 @@ class TestHitTestingServiceRetrieve:
|
||||
dataset_retrieval_model = {
|
||||
"search_method": RetrievalMethod.HYBRID_SEARCH,
|
||||
"top_k": 3,
|
||||
"reranking_enable": False,
|
||||
"score_threshold_enabled": False,
|
||||
}
|
||||
dataset = HitTestingTestDataFactory.create_dataset_mock(retrieval_model=dataset_retrieval_model)
|
||||
account = HitTestingTestDataFactory.create_user_mock()
|
||||
@@ -366,7 +379,9 @@ class TestHitTestingServiceRetrieve:
|
||||
mock_format.return_value = mock_records
|
||||
|
||||
# Act
|
||||
result = HitTestingService.retrieve(dataset, query, account, retrieval_model, external_retrieval_model)
|
||||
result = HitTestingService.retrieve(
|
||||
mock_db_session, dataset, query, account, retrieval_model, external_retrieval_model
|
||||
)
|
||||
|
||||
# Assert
|
||||
assert result["query"]["content"] == query
|
||||
@@ -391,8 +406,7 @@ class TestHitTestingServiceExternalRetrieve:
|
||||
Provides a mocked database session for testing database operations
|
||||
like adding and committing DatasetQuery records.
|
||||
"""
|
||||
with patch("services.hit_testing_service.db.session", autospec=True) as mock_db:
|
||||
yield mock_db
|
||||
return MagicMock()
|
||||
|
||||
def test_external_retrieve_success(self, mock_db_session):
|
||||
"""
|
||||
@@ -424,7 +438,7 @@ class TestHitTestingServiceExternalRetrieve:
|
||||
|
||||
# Act
|
||||
result = HitTestingService.external_retrieve(
|
||||
dataset, query, account, external_retrieval_model, metadata_filtering_conditions
|
||||
mock_db_session, dataset, query, account, external_retrieval_model, metadata_filtering_conditions
|
||||
)
|
||||
|
||||
# Assert
|
||||
@@ -455,7 +469,7 @@ class TestHitTestingServiceExternalRetrieve:
|
||||
|
||||
# Act
|
||||
result = HitTestingService.external_retrieve(
|
||||
dataset, query, account, external_retrieval_model, metadata_filtering_conditions
|
||||
mock_db_session, dataset, query, account, external_retrieval_model, metadata_filtering_conditions
|
||||
)
|
||||
|
||||
# Assert
|
||||
@@ -490,7 +504,7 @@ class TestHitTestingServiceExternalRetrieve:
|
||||
|
||||
# Act
|
||||
result = HitTestingService.external_retrieve(
|
||||
dataset, query, account, external_retrieval_model, metadata_filtering_conditions
|
||||
mock_db_session, dataset, query, account, external_retrieval_model, metadata_filtering_conditions
|
||||
)
|
||||
|
||||
# Assert
|
||||
@@ -524,7 +538,7 @@ class TestHitTestingServiceExternalRetrieve:
|
||||
|
||||
# Act
|
||||
result = HitTestingService.external_retrieve(
|
||||
dataset, query, account, external_retrieval_model, metadata_filtering_conditions
|
||||
mock_db_session, dataset, query, account, external_retrieval_model, metadata_filtering_conditions
|
||||
)
|
||||
|
||||
# Assert
|
||||
@@ -565,7 +579,7 @@ class TestHitTestingServiceCompactRetrieveResponse:
|
||||
mock_format.return_value = mock_records
|
||||
|
||||
# Act
|
||||
result = HitTestingService.compact_retrieve_response(query, documents)
|
||||
result = HitTestingService.compact_retrieve_response(MagicMock(), query, documents)
|
||||
|
||||
# Assert
|
||||
assert result["query"]["content"] == query
|
||||
@@ -591,7 +605,7 @@ class TestHitTestingServiceCompactRetrieveResponse:
|
||||
mock_format.return_value = []
|
||||
|
||||
# Act
|
||||
result = HitTestingService.compact_retrieve_response(query, documents)
|
||||
result = HitTestingService.compact_retrieve_response(MagicMock(), query, documents)
|
||||
|
||||
# Assert
|
||||
assert result["query"]["content"] == query
|
||||
@@ -708,7 +722,7 @@ class TestHitTestingServiceHitTestingArgsCheck:
|
||||
args = {"query": ""}
|
||||
|
||||
# Act & Assert
|
||||
with pytest.raises(ValueError, match="Query is required and cannot exceed 250 characters"):
|
||||
with pytest.raises(ValueError, match="Query or attachment_ids is required"):
|
||||
HitTestingService.hit_testing_args_check(args)
|
||||
|
||||
def test_hit_testing_args_check_none_query(self):
|
||||
@@ -721,7 +735,7 @@ class TestHitTestingServiceHitTestingArgsCheck:
|
||||
args = {"query": None}
|
||||
|
||||
# Act & Assert
|
||||
with pytest.raises(ValueError, match="Query is required and cannot exceed 250 characters"):
|
||||
with pytest.raises(ValueError, match="Query or attachment_ids is required"):
|
||||
HitTestingService.hit_testing_args_check(args)
|
||||
|
||||
def test_hit_testing_args_check_too_long_query(self):
|
||||
@@ -734,7 +748,7 @@ class TestHitTestingServiceHitTestingArgsCheck:
|
||||
args = {"query": "a" * 251}
|
||||
|
||||
# Act & Assert
|
||||
with pytest.raises(ValueError, match="Query is required and cannot exceed 250 characters"):
|
||||
with pytest.raises(ValueError, match="Query cannot exceed 250 characters"):
|
||||
HitTestingService.hit_testing_args_check(args)
|
||||
|
||||
def test_hit_testing_args_check_exactly_250_characters(self):
|
||||
|
||||
Reference in New Issue
Block a user