test: use SQLite sessions in services misc (#39114)

This commit is contained in:
Asuka Minato 2026-07-22 11:11:52 +09:00 committed by GitHub
parent d5e5227c6e
commit 2d15743b96
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
2 changed files with 108 additions and 88 deletions

View File

@ -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

View File

@ -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