diff --git a/api/tests/unit_tests/services/hit_service.py b/api/tests/unit_tests/services/hit_service.py index 0257fd43676..2a456dc4b9d 100644 --- a/api/tests/unit_tests/services/hit_service.py +++ b/api/tests/unit_tests/services/hit_service.py @@ -6,17 +6,25 @@ which handles retrieval testing operations for datasets, including internal dataset retrieval and external knowledge base retrieval. """ +import json from typing import Any from unittest.mock import MagicMock, Mock, patch import pytest +from sqlalchemy import select +from sqlalchemy.orm import Session from core.rag.models.document import Document from core.rag.retrieval.retrieval_methods import RetrievalMethod from models import Account -from models.dataset import Dataset +from models.dataset import Dataset, DatasetQuery from services.hit_testing_service import HitTestingService +pytestmark = [ + pytest.mark.usefixtures("sqlite_session"), + pytest.mark.parametrize("sqlite_session", [(DatasetQuery,)], indirect=True), +] + class HitTestingTestDataFactory: """ @@ -139,17 +147,7 @@ class TestHitTestingServiceRetrieve: various retrieval model configurations, metadata filtering, and query logging. """ - @pytest.fixture - def mock_db_session(self): - """ - Mock database session. - - Provides a mocked database session for testing database operations - like adding and committing DatasetQuery records. - """ - return MagicMock() - - def test_retrieve_success_with_default_retrieval_model(self, mock_db_session): + def test_retrieve_success_with_default_retrieval_model(self, sqlite_session: Session): """ Test successful retrieval with default retrieval model. @@ -186,17 +184,20 @@ class TestHitTestingServiceRetrieve: # Act result = HitTestingService.retrieve( - dataset, query, account, retrieval_model, external_retrieval_model, session=mock_db_session + dataset, query, account, retrieval_model, external_retrieval_model, session=sqlite_session ) # Assert assert result["query"]["content"] == query assert len(result["records"]) == 2 mock_retrieve.assert_called_once() - mock_db_session.add.assert_called_once() - mock_db_session.commit.assert_called_once() + query_log = sqlite_session.scalar(select(DatasetQuery)) + assert query_log is not None + assert query_log.dataset_id == dataset.id + assert query_log.created_by == account.id + assert json.loads(query_log.content) == [{"content_type": "text_query", "content": query}] - def test_retrieve_success_with_custom_retrieval_model(self, mock_db_session): + def test_retrieve_success_with_custom_retrieval_model(self, sqlite_session: Session): """ Test successful retrieval with custom retrieval model. @@ -234,7 +235,7 @@ class TestHitTestingServiceRetrieve: # Act result = HitTestingService.retrieve( - dataset, query, account, retrieval_model, external_retrieval_model, session=mock_db_session + dataset, query, account, retrieval_model, external_retrieval_model, session=sqlite_session ) # Assert @@ -246,7 +247,7 @@ class TestHitTestingServiceRetrieve: assert call_kwargs["score_threshold"] == 0.7 assert call_kwargs["reranking_model"] == retrieval_model["reranking_model"] - def test_retrieve_with_metadata_filtering(self, mock_db_session): + def test_retrieve_with_metadata_filtering(self, sqlite_session: Session): """ Test retrieval with metadata filtering conditions. @@ -292,7 +293,7 @@ class TestHitTestingServiceRetrieve: # Act result = HitTestingService.retrieve( - dataset, query, account, retrieval_model, external_retrieval_model, session=mock_db_session + dataset, query, account, retrieval_model, external_retrieval_model, session=sqlite_session ) # Assert @@ -301,7 +302,7 @@ class TestHitTestingServiceRetrieve: call_kwargs = mock_retrieve.call_args[1] assert call_kwargs["document_ids_filter"] == ["doc-1", "doc-2"] - def test_retrieve_with_metadata_filtering_no_documents(self, mock_db_session): + def test_retrieve_with_metadata_filtering_no_documents(self, sqlite_session: Session): """ Test retrieval with metadata filtering that returns no documents. @@ -337,14 +338,14 @@ class TestHitTestingServiceRetrieve: # Act result = HitTestingService.retrieve( - dataset, query, account, retrieval_model, external_retrieval_model, session=mock_db_session + dataset, query, account, retrieval_model, external_retrieval_model, session=sqlite_session ) # Assert assert result["query"]["content"] == query assert result["records"] == [] - def test_retrieve_with_dataset_retrieval_model(self, mock_db_session): + def test_retrieve_with_dataset_retrieval_model(self, sqlite_session: Session): """ Test retrieval using dataset's retrieval model when not provided. @@ -380,7 +381,7 @@ class TestHitTestingServiceRetrieve: # Act result = HitTestingService.retrieve( - dataset, query, account, retrieval_model, external_retrieval_model, session=mock_db_session + dataset, query, account, retrieval_model, external_retrieval_model, session=sqlite_session ) # Assert @@ -398,17 +399,7 @@ class TestHitTestingServiceExternalRetrieve: including query escaping, response formatting, and provider validation. """ - @pytest.fixture - def mock_db_session(self): - """ - Mock database session. - - Provides a mocked database session for testing database operations - like adding and committing DatasetQuery records. - """ - return MagicMock() - - def test_external_retrieve_success(self, mock_db_session): + def test_external_retrieve_success(self, sqlite_session: Session): """ Test successful external retrieval. @@ -443,7 +434,7 @@ class TestHitTestingServiceExternalRetrieve: account, external_retrieval_model, metadata_filtering_conditions, - session=mock_db_session, + session=sqlite_session, ) # Assert @@ -455,10 +446,13 @@ class TestHitTestingServiceExternalRetrieve: mock_external_retrieve.assert_called_once() # Verify query was escaped assert mock_external_retrieve.call_args[1]["query"] == 'test query with \\"quotes\\"' - mock_db_session.add.assert_called_once() - mock_db_session.commit.assert_called_once() + query_log = sqlite_session.scalar(select(DatasetQuery)) + assert query_log is not None + assert query_log.dataset_id == dataset.id + assert query_log.content == query + assert query_log.created_by == account.id - def test_external_retrieve_non_external_provider(self, mock_db_session): + def test_external_retrieve_non_external_provider(self, sqlite_session: Session): """ Test external retrieval with non-external provider (should return empty). @@ -474,15 +468,15 @@ class TestHitTestingServiceExternalRetrieve: # Act result = HitTestingService.external_retrieve( - dataset, query, account, external_retrieval_model, metadata_filtering_conditions, session=mock_db_session + dataset, query, account, external_retrieval_model, metadata_filtering_conditions, session=sqlite_session ) # Assert assert result["query"]["content"] == query assert result["records"] == [] - mock_db_session.add.assert_not_called() + assert sqlite_session.scalar(select(DatasetQuery)) is None - def test_external_retrieve_with_metadata_filtering(self, mock_db_session): + def test_external_retrieve_with_metadata_filtering(self, sqlite_session: Session): """ Test external retrieval with metadata filtering conditions. @@ -514,7 +508,7 @@ class TestHitTestingServiceExternalRetrieve: account, external_retrieval_model, metadata_filtering_conditions, - session=mock_db_session, + session=sqlite_session, ) # Assert @@ -523,7 +517,7 @@ class TestHitTestingServiceExternalRetrieve: call_kwargs = mock_external_retrieve.call_args[1] assert call_kwargs["metadata_filtering_conditions"] == metadata_filtering_conditions - def test_external_retrieve_empty_documents(self, mock_db_session): + def test_external_retrieve_empty_documents(self, sqlite_session: Session): """ Test external retrieval with empty document list. @@ -553,7 +547,7 @@ class TestHitTestingServiceExternalRetrieve: account, external_retrieval_model, metadata_filtering_conditions, - session=mock_db_session, + session=sqlite_session, ) # Assert @@ -569,7 +563,7 @@ class TestHitTestingServiceCompactRetrieveResponse: ensuring documents are properly formatted into retrieval records. """ - def test_compact_retrieve_response_success(self): + def test_compact_retrieve_response_success(self, sqlite_session: Session): """ Test successful response formatting. @@ -587,7 +581,6 @@ class TestHitTestingServiceCompactRetrieveResponse: HitTestingTestDataFactory.create_retrieval_record_mock(content="Doc 1", score=0.95), HitTestingTestDataFactory.create_retrieval_record_mock(content="Doc 2", score=0.85), ] - session = MagicMock() with patch( "services.hit_testing_service.RetrievalService.format_retrieval_documents", autospec=True @@ -595,7 +588,7 @@ class TestHitTestingServiceCompactRetrieveResponse: mock_format.return_value = mock_records # Act - result = HitTestingService.compact_retrieve_response(query, documents, session=session) + result = HitTestingService.compact_retrieve_response(query, documents, session=sqlite_session) # Assert assert result["query"]["content"] == query @@ -603,10 +596,11 @@ class TestHitTestingServiceCompactRetrieveResponse: assert result["records"][0]["content"] == "Doc 1" assert result["records"][0]["score"] == 0.95 mock_format.assert_called_once() - assert mock_format.call_args.args[0] is not session + assert mock_format.call_args.args[0] is not sqlite_session + assert mock_format.call_args.args[0].get_bind() is sqlite_session.get_bind() assert mock_format.call_args.args[1] == documents - def test_compact_retrieve_response_empty_documents(self): + def test_compact_retrieve_response_empty_documents(self, sqlite_session: Session): """ Test response formatting with empty document list. @@ -616,7 +610,6 @@ class TestHitTestingServiceCompactRetrieveResponse: # Arrange query = "test query" documents = [] - session = MagicMock() with patch( "services.hit_testing_service.RetrievalService.format_retrieval_documents", autospec=True @@ -624,13 +617,14 @@ class TestHitTestingServiceCompactRetrieveResponse: mock_format.return_value = [] # Act - result = HitTestingService.compact_retrieve_response(query, documents, session=session) + result = HitTestingService.compact_retrieve_response(query, documents, session=sqlite_session) # Assert assert result["query"]["content"] == query assert result["records"] == [] mock_format.assert_called_once() - assert mock_format.call_args.args[0] is not session + assert mock_format.call_args.args[0] is not sqlite_session + assert mock_format.call_args.args[0].get_bind() is sqlite_session.get_bind() assert mock_format.call_args.args[1] == documents diff --git a/api/tests/unit_tests/services/workflow/test_workflow_event_snapshot_service_additional.py b/api/tests/unit_tests/services/workflow/test_workflow_event_snapshot_service_additional.py index be6f9ff1fc0..8efd7370a73 100644 --- a/api/tests/unit_tests/services/workflow/test_workflow_event_snapshot_service_additional.py +++ b/api/tests/unit_tests/services/workflow/test_workflow_event_snapshot_service_additional.py @@ -10,6 +10,8 @@ from typing import Any, cast from unittest.mock import MagicMock import pytest +from sqlalchemy import event +from sqlalchemy.engine import Engine from sqlalchemy.orm import Session, sessionmaker from core.app.app_config.entities import WorkflowUIBasedAppConfig @@ -21,8 +23,9 @@ from core.app.layers.pause_state_persist_layer import ( ) from graphon.enums import WorkflowExecutionStatus from graphon.runtime import GraphRuntimeState, VariablePool -from models.enums import CreatorUserRole -from models.model import AppMode +from models.base import TypeBase +from models.enums import CreatorUserRole, MessageStatus +from models.model import AppMode, Message from models.workflow import WorkflowRun from repositories.entities.workflow_pause import WorkflowPauseEntity from services import workflow_event_snapshot_service as service_module @@ -79,23 +82,47 @@ def _build_resumption_context(task_id: str) -> WorkflowResumptionContext: ) -class _SessionContext: - def __init__(self, session: Any) -> None: - self._session = session - - def __enter__(self) -> Any: - return self._session - - def __exit__(self, exc_type: Any, exc: Any, tb: Any) -> bool: - return False +@pytest.fixture +def message_session_maker(sqlite_engine: Engine) -> sessionmaker[Session]: + """Create real sessions containing only workflow messages.""" + TypeBase.metadata.create_all(sqlite_engine, tables=[TypeBase.metadata.tables[Message.__tablename__]]) + return sessionmaker(bind=sqlite_engine, expire_on_commit=False) -class _SessionMaker: - def __init__(self, session: Any) -> None: - self._session = session - - def __call__(self) -> _SessionContext: - return _SessionContext(self._session) +def _persist_message(session_maker: sessionmaker[Session]) -> Message: + message = Message( + app_id="app-1", + model_provider="provider", + model_id="model", + override_model_configs=None, + conversation_id="conv-1", + inputs={}, + query="hello", + message="", + message_tokens=0, + message_unit_price=0, + message_price_unit=0, + answer="answer", + answer_tokens=0, + answer_unit_price=0, + answer_price_unit=0, + parent_message_id=None, + provider_response_latency=0, + total_price=0, + currency="USD", + invoke_from=InvokeFrom.WEB_APP, + from_source="api", + from_end_user_id="user-1", + from_account_id=None, + app_mode=AppMode.WORKFLOW, + status=MessageStatus.NORMAL, + workflow_run_id="run-1", + ) + message.id = "msg-1" + with session_maker() as session: + session.add(message) + session.commit() + return message class _SubscriptionContext: @@ -150,12 +177,11 @@ class _PauseEntity(WorkflowPauseEntity): class TestWorkflowEventSnapshotHelpers: - def test_get_message_context_by_conversation_should_return_none_when_no_message(self) -> None: - session = SimpleNamespace(scalar=MagicMock(return_value=None)) - session_maker = _SessionMaker(session) - + def test_get_message_context_by_conversation_should_return_none_when_no_message( + self, message_session_maker: sessionmaker[Session] + ) -> None: result = service_module._get_message_context_by_conversation( - cast(sessionmaker[Session], session_maker), + message_session_maker, conversation_id="conv-1", workflow_run_id="run-1", ) @@ -163,22 +189,22 @@ class TestWorkflowEventSnapshotHelpers: assert result is None def test_get_message_context_by_conversation_should_default_created_at_to_zero_when_message_has_no_timestamp( - self, + self, message_session_maker: sessionmaker[Session] ) -> None: - message = SimpleNamespace( - id="msg-1", - conversation_id="conv-1", - created_at=None, - answer="answer", - ) - session = SimpleNamespace(scalar=MagicMock(return_value=message)) - session_maker = _SessionMaker(session) + _persist_message(message_session_maker) - result = service_module._get_message_context_by_conversation( - cast(sessionmaker[Session], session_maker), - conversation_id="conv-1", - workflow_run_id="run-1", - ) + def clear_created_at(message: Message, _context: Any) -> None: + message.created_at = None # type: ignore[assignment] + + event.listen(Message, "load", clear_created_at) + try: + result = service_module._get_message_context_by_conversation( + message_session_maker, + conversation_id="conv-1", + workflow_run_id="run-1", + ) + finally: + event.remove(Message, "load", clear_created_at) assert result is not None assert result.created_at == 0