test: use SQLite sessions in controllers service api (#39098)

Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
Co-authored-by: Byron.wang <byron@dify.ai>
This commit is contained in:
Asuka Minato 2026-07-28 12:19:17 +09:00 committed by GitHub
parent 1e61078e93
commit 81fb57639e
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
2 changed files with 65 additions and 34 deletions

View File

@ -15,12 +15,15 @@ Focus on:
"""
import uuid
from collections.abc import Iterator
from inspect import unwrap
from types import SimpleNamespace
from unittest.mock import Mock, patch
import pytest
from flask import Flask
from sqlalchemy import Engine
from sqlalchemy.orm import Session
from werkzeug.exceptions import BadRequest, InternalServerError, NotFound
from controllers.service_api.app.error import NotChatAppError
@ -44,6 +47,14 @@ from services.errors.message import (
from services.message_service import MessageService
@pytest.fixture
def orm_session(sqlite_engine: Engine) -> Iterator[Session]:
"""Provide a real caller-owned session for MessageService interface tests."""
with Session(sqlite_engine, expire_on_commit=False) as session:
yield session
class TestMessageListQuery:
"""Test suite for MessageListQuery Pydantic model."""
@ -253,7 +264,7 @@ class TestMessageService:
assert callable(MessageService.get_suggested_questions_after_answer)
@patch.object(MessageService, "pagination_by_first_id")
def test_pagination_by_first_id_returns_pagination_result(self, mock_pagination):
def test_pagination_by_first_id_returns_pagination_result(self, mock_pagination, orm_session: Session):
"""Test pagination_by_first_id returns expected format."""
mock_result = Mock()
mock_result.data = []
@ -267,7 +278,7 @@ class TestMessageService:
conversation_id=str(uuid.uuid4()),
first_id=None,
limit=20,
session=Mock(),
session=orm_session,
)
assert hasattr(result, "data")
@ -275,7 +286,7 @@ class TestMessageService:
assert hasattr(result, "has_more")
@patch.object(MessageService, "pagination_by_first_id")
def test_pagination_raises_conversation_not_exists_error(self, mock_pagination):
def test_pagination_raises_conversation_not_exists_error(self, mock_pagination, orm_session: Session):
"""Test pagination raises ConversationNotExistsError."""
import services.errors.conversation
@ -288,11 +299,11 @@ class TestMessageService:
conversation_id="invalid_id",
first_id=None,
limit=20,
session=Mock(),
session=orm_session,
)
@patch.object(MessageService, "pagination_by_first_id")
def test_pagination_raises_first_message_not_exists_error(self, mock_pagination):
def test_pagination_raises_first_message_not_exists_error(self, mock_pagination, orm_session: Session):
"""Test pagination raises FirstMessageNotExistsError."""
mock_pagination.side_effect = FirstMessageNotExistsError()
@ -303,11 +314,11 @@ class TestMessageService:
conversation_id=str(uuid.uuid4()),
first_id="invalid_first_id",
limit=20,
session=Mock(),
session=orm_session,
)
@patch.object(MessageService, "create_feedback")
def test_create_feedback_with_rating_and_content(self, mock_create_feedback):
def test_create_feedback_with_rating_and_content(self, mock_create_feedback, orm_session: Session):
"""Test create_feedback with rating and content."""
mock_create_feedback.return_value = None
@ -317,13 +328,13 @@ class TestMessageService:
user=Mock(spec=EndUser),
rating=FeedbackRating.LIKE,
content="Great response!",
session=Mock(),
session=orm_session,
)
mock_create_feedback.assert_called_once()
@patch.object(MessageService, "create_feedback")
def test_create_feedback_raises_message_not_exists_error(self, mock_create_feedback):
def test_create_feedback_raises_message_not_exists_error(self, mock_create_feedback, orm_session: Session):
"""Test create_feedback raises MessageNotExistsError."""
mock_create_feedback.side_effect = MessageNotExistsError()
@ -334,11 +345,11 @@ class TestMessageService:
user=Mock(spec=EndUser),
rating=FeedbackRating.LIKE,
content=None,
session=Mock(),
session=orm_session,
)
@patch.object(MessageService, "get_all_messages_feedbacks")
def test_get_all_messages_feedbacks_returns_list(self, mock_get_feedbacks):
def test_get_all_messages_feedbacks_returns_list(self, mock_get_feedbacks, orm_session: Session):
"""Test get_all_messages_feedbacks returns list of feedbacks."""
mock_feedbacks = [
{"message_id": str(uuid.uuid4()), "rating": "like"},
@ -346,13 +357,15 @@ class TestMessageService:
]
mock_get_feedbacks.return_value = mock_feedbacks
result = MessageService.get_all_messages_feedbacks(app_model=Mock(spec=App), page=1, limit=20, session=Mock())
result = MessageService.get_all_messages_feedbacks(
app_model=Mock(spec=App), page=1, limit=20, session=orm_session
)
assert len(result) == 2
assert result[0]["rating"] == "like"
@patch.object(MessageService, "get_suggested_questions_after_answer")
def test_get_suggested_questions_returns_questions_list(self, mock_get_questions):
def test_get_suggested_questions_returns_questions_list(self, mock_get_questions, orm_session: Session):
"""Test get_suggested_questions_after_answer returns list of questions."""
mock_questions = ["What about this aspect?", "Can you elaborate on that?", "How does this relate to...?"]
mock_get_questions.return_value = mock_questions
@ -362,14 +375,14 @@ class TestMessageService:
user=Mock(spec=EndUser),
message_id=str(uuid.uuid4()),
invoke_from=Mock(),
session=Mock(),
session=orm_session,
)
assert len(result) == 3
assert isinstance(result[0], str)
@patch.object(MessageService, "get_suggested_questions_after_answer")
def test_get_suggested_questions_raises_disabled_error(self, mock_get_questions):
def test_get_suggested_questions_raises_disabled_error(self, mock_get_questions, orm_session: Session):
"""Test get_suggested_questions_after_answer raises SuggestedQuestionsAfterAnswerDisabledError."""
mock_get_questions.side_effect = SuggestedQuestionsAfterAnswerDisabledError()
@ -379,11 +392,11 @@ class TestMessageService:
user=Mock(spec=EndUser),
message_id=str(uuid.uuid4()),
invoke_from=Mock(),
session=Mock(),
session=orm_session,
)
@patch.object(MessageService, "get_suggested_questions_after_answer")
def test_get_suggested_questions_raises_message_not_exists_error(self, mock_get_questions):
def test_get_suggested_questions_raises_message_not_exists_error(self, mock_get_questions, orm_session: Session):
"""Test get_suggested_questions_after_answer raises MessageNotExistsError."""
mock_get_questions.side_effect = MessageNotExistsError()
@ -393,7 +406,7 @@ class TestMessageService:
user=Mock(spec=EndUser),
message_id="invalid_message_id",
invoke_from=Mock(),
session=Mock(),
session=orm_session,
)

View File

@ -24,6 +24,7 @@ from unittest.mock import Mock, patch
import pytest
from flask import Flask
from sqlalchemy.orm import Session
from werkzeug.datastructures import FileStorage
from werkzeug.exceptions import Forbidden, NotFound
@ -38,6 +39,7 @@ from controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow import (
)
from core.app.entities.app_invoke_entities import InvokeFrom
from models.account import Account
from models.dataset import Dataset
from services.errors.file import FileTooLargeError, UnsupportedFileTypeError
from services.rag_pipeline.entity.pipeline_service_api_entities import (
DatasourceNodeRunApiEntity,
@ -46,6 +48,20 @@ from services.rag_pipeline.entity.pipeline_service_api_entities import (
from services.rag_pipeline.rag_pipeline import RagPipelineService
def _persist_dataset(session: Session, *, tenant_id: str, dataset_id: str) -> Dataset:
dataset = Dataset(
id=dataset_id,
tenant_id=tenant_id,
name="Pipeline dataset",
created_by="account-1",
data_source_type=None,
indexing_technique=None,
)
session.add(dataset)
session.commit()
return dataset
class TestDatasourceNodeRunPayload:
"""Test suite for DatasourceNodeRunPayload Pydantic model."""
@ -550,13 +566,15 @@ class TestPipelineRunApiPost:
)
@patch("controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.RagPipelineService")
@patch("controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.service_api_ns")
def test_post_success_streaming(self, mock_ns, mock_svc_cls, mock_current_user, mock_gen_svc, mock_helper, app):
@pytest.mark.parametrize("sqlite_session", [(Dataset,)], indirect=True)
def test_post_success_streaming(
self, mock_ns, mock_svc_cls, mock_current_user, mock_gen_svc, mock_helper, app, sqlite_session: Session
):
"""Test successful pipeline run with streaming response."""
tenant_id = str(uuid.uuid4())
dataset_id = str(uuid.uuid4())
session = Mock()
session.scalar.return_value = Mock()
_persist_dataset(sqlite_session, tenant_id=tenant_id, dataset_id=dataset_id)
mock_ns.payload = {
"inputs": {"key": "val"},
@ -577,33 +595,33 @@ class TestPipelineRunApiPost:
with app.test_request_context("/datasets/test/pipeline/run", method="POST"):
api = PipelineRunApi()
response = api.post.__wrapped__(api, session, tenant_id=tenant_id, dataset_id=dataset_id)
response = api.post.__wrapped__(api, sqlite_session, tenant_id=tenant_id, dataset_id=dataset_id)
assert response == {"result": "ok"}
mock_svc_cls.assert_called_once_with(session)
mock_svc_cls.assert_called_once_with(sqlite_session)
mock_gen_svc.generate.assert_called_once()
def test_post_not_found(self, app: Flask):
@pytest.mark.parametrize("sqlite_session", [(Dataset,)], indirect=True)
def test_post_not_found(self, app: Flask, sqlite_session: Session):
"""Test NotFound when dataset check fails."""
session = Mock()
session.scalar.return_value = None
with app.test_request_context("/datasets/test/pipeline/run", method="POST"):
api = PipelineRunApi()
with pytest.raises(NotFound):
api.post.__wrapped__(
api,
session,
sqlite_session,
tenant_id=str(uuid.uuid4()),
dataset_id=str(uuid.uuid4()),
)
@patch("controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.current_user", new="not_account")
@patch("controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.service_api_ns")
def test_post_forbidden_non_account_user(self, mock_ns, app: Flask):
@pytest.mark.parametrize("sqlite_session", [(Dataset,)], indirect=True)
def test_post_forbidden_non_account_user(self, mock_ns, app: Flask, sqlite_session: Session):
"""Test Forbidden when current_user is not an Account."""
session = Mock()
session.scalar.return_value = Mock()
tenant_id = str(uuid.uuid4())
dataset_id = str(uuid.uuid4())
_persist_dataset(sqlite_session, tenant_id=tenant_id, dataset_id=dataset_id)
mock_ns.payload = {
"inputs": {},
"datasource_type": "online_document",
@ -618,9 +636,9 @@ class TestPipelineRunApiPost:
with pytest.raises(Forbidden):
api.post.__wrapped__(
api,
session,
tenant_id=str(uuid.uuid4()),
dataset_id=str(uuid.uuid4()),
sqlite_session,
tenant_id=tenant_id,
dataset_id=dataset_id,
)