test: migrate residual controller sessions and ORM models to SQLite (#40606)

Co-authored-by: Byron Wang <byron@dify.ai>
This commit is contained in:
Asuka Minato
2026-08-15 15:15:11 +00:00
committed by GitHub
parent c2e68335d4
commit 1651bcc774
6 changed files with 196 additions and 106 deletions
@@ -27,12 +27,19 @@ from controllers.console.datasets.rag_pipeline.datasource_auth import (
)
from core.plugin.impl.oauth import OAuthHandler
from graphon.model_runtime.errors.validate import CredentialsValidateFailedError
from models.account import Account
from services.datasource_provider_service import DatasourceProviderService
from services.plugin.oauth_service import OAuthProxyService
_PROVIDER_ID = "langgenius/notion_datasource/notion"
def _account() -> Account:
account = Account(name="Datasource Auth Tester", email="datasource-auth@example.com")
account.id = "user-1"
return account
def _i18n(text: str) -> dict[str, str]:
return {"en_US": text, "zh_Hans": text, "pt_BR": text, "ja_JP": text}
@@ -106,7 +113,7 @@ class TestDatasourcePluginOAuthAuthorizationUrl:
api = DatasourcePluginOAuthAuthorizationUrl()
method = inspect.unwrap(api.get)
user = MagicMock(id="user-1")
user = _account()
oauth_client = {"client_id": "abc", "client_secret": "shh", "scopes": ["read", "write"]}
auth_url_payload = {
"authorization_url": "https://auth.example.com/oauth?client_id=abc&state=xyz",
@@ -155,7 +162,7 @@ class TestDatasourcePluginOAuthAuthorizationUrl:
def test_get_no_oauth_config(self, app: Flask):
api = DatasourcePluginOAuthAuthorizationUrl()
method = inspect.unwrap(api.get)
user = MagicMock(id="user-1")
user = _account()
with (
app.test_request_context("/"),
@@ -172,7 +179,7 @@ class TestDatasourcePluginOAuthAuthorizationUrl:
api = DatasourcePluginOAuthAuthorizationUrl()
method = inspect.unwrap(api.get)
user = MagicMock(id="user-1")
user = _account()
with (
app.test_request_context("/"),
@@ -443,7 +450,7 @@ class TestDatasourceAuth:
def test_get_success(self, app: Flask):
api = DatasourceAuth()
method = inspect.unwrap(api.get)
user = MagicMock(id="user-1")
user = _account()
with (
app.test_request_context("/"),
@@ -474,7 +481,7 @@ class TestDatasourceAuth:
def test_get_empty_list(self, app: Flask):
api = DatasourceAuth()
method = inspect.unwrap(api.get)
user = MagicMock(id="user-1")
user = _account()
with (
app.test_request_context("/"),
@@ -17,9 +17,16 @@ from controllers.console.datasets.rag_pipeline.rag_pipeline_datasets import (
CreateRagPipelineDatasetApi,
RagPipelineDatasetImportPayload,
)
from models.account import Account, TenantAccountRole
from services.entities.dsl_entities import ImportStatus
def _account(*, editor: bool) -> Account:
account = Account(name="RAG Pipeline Tester", email="rag-pipeline@example.com")
account.role = TenantAccountRole.EDITOR if editor else TenantAccountRole.NORMAL
return account
class TestCreateRagPipelineDatasetApi:
def _valid_payload(self) -> dict[str, str]:
return {"yaml_content": "name: test"}
@@ -29,7 +36,7 @@ class TestCreateRagPipelineDatasetApi:
method = unwrap(api.post)
payload = self._valid_payload()
user = MagicMock(is_dataset_editor=True)
user = _account(editor=True)
import_info = {
"id": "import-1",
"status": ImportStatus.COMPLETED,
@@ -69,7 +76,7 @@ class TestCreateRagPipelineDatasetApi:
method = unwrap(api.post)
payload = self._valid_payload()
user = MagicMock(is_dataset_editor=False)
user = _account(editor=False)
with (
app.test_request_context("/", json=payload),
@@ -83,7 +90,7 @@ class TestCreateRagPipelineDatasetApi:
method = unwrap(api.post)
payload = self._valid_payload()
user = MagicMock(is_dataset_editor=True)
user = _account(editor=True)
mock_service = MagicMock()
mock_service.create_rag_pipeline_dataset.side_effect = services.errors.dataset.DatasetNameDuplicateError()
@@ -104,7 +111,7 @@ class TestCreateRagPipelineDatasetApi:
method = unwrap(api.post)
payload: dict[str, str] = {}
user = MagicMock(is_dataset_editor=True)
user = _account(editor=True)
with (
app.test_request_context("/", json=payload),
@@ -119,7 +126,7 @@ class TestCreateEmptyRagPipelineDatasetApi:
api = CreateEmptyRagPipelineDatasetApi()
method = unwrap(api.post)
user = MagicMock(is_dataset_editor=False)
user = _account(editor=False)
with app.test_request_context("/"):
with pytest.raises(Forbidden):
@@ -19,12 +19,19 @@ from controllers.console.datasets.rag_pipeline.rag_pipeline_import import (
RagPipelineImportPayload,
)
from core.plugin.entities.plugin import PluginDependency, PluginDependencyType
from models.account import Account
from models.dataset import Pipeline
from models.engine import db
from services.entities.dsl_entities import CheckDependenciesResult, ImportStatus
from services.rag_pipeline.rag_pipeline_dsl_service import RagPipelineImportInfo
def _account() -> Account:
account = Account(name="RAG Import Tester", email="rag-import@example.com")
account.id = "account-1"
return account
@pytest.fixture
def app() -> Iterator[Flask]:
app = Flask(__name__)
@@ -48,7 +55,7 @@ class TestRagPipelineImportApi:
method = unwrap(api.post)
payload = self._payload()
user = MagicMock()
user = _account()
result = RagPipelineImportInfo(
id="import-1",
status=ImportStatus.COMPLETED,
@@ -87,7 +94,7 @@ class TestRagPipelineImportApi:
method = unwrap(api.post)
payload = self._payload()
user = MagicMock()
user = _account()
result = RagPipelineImportInfo(
id="import-1",
status=ImportStatus.FAILED,
@@ -120,7 +127,7 @@ class TestRagPipelineImportApi:
method = unwrap(api.post)
payload = self._payload()
user = MagicMock()
user = _account()
result = RagPipelineImportInfo(
id="import-1",
status=ImportStatus.PENDING,
@@ -154,7 +161,7 @@ class TestRagPipelineImportConfirmApi:
api = RagPipelineImportConfirmApi()
method = unwrap(api.post)
user = MagicMock()
user = _account()
result = RagPipelineImportInfo(
id="import-1",
status=ImportStatus.COMPLETED,
@@ -184,7 +191,7 @@ class TestRagPipelineImportConfirmApi:
api = RagPipelineImportConfirmApi()
method = unwrap(api.post)
user = MagicMock()
user = _account()
result = RagPipelineImportInfo(
id="import-1",
status=ImportStatus.FAILED,
@@ -1,5 +1,5 @@
from inspect import unwrap
from unittest.mock import MagicMock, patch
from unittest.mock import patch
from flask import Flask
@@ -7,6 +7,13 @@ from controllers.console.workspace.agent_providers import (
AgentProviderApi,
AgentProviderListApi,
)
from models.account import Account
def _account() -> Account:
account = Account(name="Agent Provider Tester", email="agent-provider@example.com")
account.id = "user1"
return account
class TestAgentProviderListApi:
@@ -14,7 +21,7 @@ class TestAgentProviderListApi:
api = AgentProviderListApi()
method = unwrap(api.get)
user = MagicMock(id="user1")
user = _account()
tenant_id = "tenant1"
providers = [{"name": "openai"}, {"name": "anthropic"}]
@@ -33,7 +40,7 @@ class TestAgentProviderListApi:
api = AgentProviderListApi()
method = unwrap(api.get)
user = MagicMock(id="user1")
user = _account()
tenant_id = "tenant1"
with (
@@ -53,7 +60,7 @@ class TestAgentProviderApi:
api = AgentProviderApi()
method = unwrap(api.get)
user = MagicMock(id="user1")
user = _account()
tenant_id = "tenant1"
provider_name = "openai"
provider_data = {"name": "openai", "models": ["gpt-4"]}
@@ -73,7 +80,7 @@ class TestAgentProviderApi:
api = AgentProviderApi()
method = unwrap(api.get)
user = MagicMock(id="user1")
user = _account()
tenant_id = "tenant1"
provider_name = "unknown"
@@ -9,7 +9,6 @@ Service API controller tests.
import uuid
from collections.abc import Iterator
from dataclasses import dataclass
from unittest.mock import Mock
import pytest
from flask import Flask
@@ -19,7 +18,7 @@ from sqlalchemy.orm import Session
from core.rag.index_processor.constant.index_type import IndexStructureType
from models.account import Account, Tenant, TenantAccountJoin, TenantAccountRole, TenantStatus
from models.base import TypeBase
from models.model import App, AppMode, EndUser
from models.model import ApiToken, App, AppMode, EndUser, EndUserType
@dataclass(frozen=True)
@@ -82,11 +81,15 @@ def mock_app_id():
@pytest.fixture
def mock_end_user(mock_tenant_id):
"""Create a mock EndUser model with required attributes."""
"""Create a real EndUser model with required attributes."""
user = EndUser(
id=str(uuid.uuid4()),
external_user_id=f"external_{uuid.uuid4().hex[:8]}",
tenant_id=mock_tenant_id,
app_id=None,
type=EndUserType.SERVICE_API,
name="Service API User",
session_id=str(uuid.uuid4()),
)
return user
@@ -103,52 +106,44 @@ def mock_app_model(mock_app_id, mock_tenant_id):
status="normal",
enable_api=True,
)
app.author_name = "Test Author"
app.tags = []
# Mock workflow for workflow apps
app.workflow = None
app.app_model_config = None
return app
@pytest.fixture
def mock_tenant(mock_tenant_id):
"""Create a Tenant model."""
tenant = Mock()
tenant = Tenant(name="Service API Tenant", status=TenantStatus.NORMAL)
tenant.id = mock_tenant_id
tenant.status = TenantStatus.NORMAL
return tenant
@pytest.fixture
def mock_account():
"""Create an Account model."""
account = Mock()
account = Account(name="Service API Account", email=f"service-{uuid.uuid4()}@example.com")
account.id = str(uuid.uuid4())
return account
@pytest.fixture
def mock_api_token(mock_app_id, mock_tenant_id):
"""Create a mock API token for authentication tests."""
token = Mock()
token.app_id = mock_app_id
token.tenant_id = mock_tenant_id
token.token = f"test_token_{uuid.uuid4().hex[:8]}"
token.type = "app"
return token
"""Create a real API token for authentication tests."""
return ApiToken(
app_id=mock_app_id,
tenant_id=mock_tenant_id,
token=f"test_token_{uuid.uuid4().hex[:8]}",
type="app",
)
@pytest.fixture
def mock_dataset_api_token(mock_tenant_id):
"""Create a mock API token for dataset endpoints."""
token = Mock()
token.tenant_id = mock_tenant_id
token.token = f"dataset_token_{uuid.uuid4().hex[:8]}"
token.type = "dataset"
return token
"""Create a real API token for dataset endpoints."""
return ApiToken(
tenant_id=mock_tenant_id,
token=f"dataset_token_{uuid.uuid4().hex[:8]}",
type="dataset",
)
@pytest.fixture
@@ -4,8 +4,10 @@ from __future__ import annotations
import builtins
import inspect
import json
import uuid
from datetime import datetime
from pathlib import Path
from types import ModuleType, SimpleNamespace
from unittest.mock import ANY, patch
from uuid import uuid4
@@ -15,7 +17,18 @@ from flask import Flask
from flask.views import MethodView
from controllers.common.controller_schemas import MessageListQuery
from core.app.entities.app_invoke_entities import InvokeFrom
from core.entities.execution_extra_content import HumanInputContent
from models.enums import ConversationFromSource, EndUserType, FeedbackFromSource, FeedbackRating
from models.model import (
App,
AppMode,
Conversation,
EndUser,
Message,
MessageAgentThought,
MessageFeedback,
)
# Ensure flask_restx.api finds MethodView during import.
if not hasattr(builtins, "MethodView"):
@@ -36,8 +49,9 @@ def _load_controller_module():
from flask_restx import Namespace
stub = ModuleType(parent_module_name)
stub.__file__ = "controllers/web/__init__.py"
stub.__path__ = ["controllers/web"]
web_controller_dir = Path(__file__).resolve().parents[4] / "controllers" / "web"
stub.__file__ = str(web_controller_dir / "__init__.py")
stub.__path__ = [str(web_controller_dir)]
stub.__package__ = "controllers"
stub.__spec__ = importlib.util.spec_from_loader(parent_module_name, loader=None, is_package=True)
stub.web_ns = Namespace("web", description="Web API", path="/")
@@ -67,7 +81,12 @@ def app() -> Flask:
return app
def test_message_list_mapping(app: Flask) -> None:
@pytest.mark.parametrize(
"sqlite_session",
[(App, Conversation, EndUser, Message, MessageAgentThought, MessageFeedback)],
indirect=True,
)
def test_message_list_mapping(app: Flask, monkeypatch: pytest.MonkeyPatch, sqlite_session) -> None:
conversation_id = str(uuid4())
message_id = str(uuid4())
@@ -75,87 +94,135 @@ def test_message_list_mapping(app: Flask) -> None:
resource_created_at = datetime(2024, 1, 1, 13, 0, 0)
thought_created_at = datetime(2024, 1, 1, 14, 0, 0)
retriever_resource_obj = SimpleNamespace(
id="res-obj",
message_id=message_id,
position=2,
dataset_id="ds-1",
dataset_name="dataset",
document_id="doc-1",
document_name="document",
data_source_type="file",
segment_id="seg-1",
score=0.9,
hit_count=1,
word_count=10,
segment_position=0,
index_node_hash="hash",
content="content",
created_at=resource_created_at,
)
retriever_resource = {
"id": "res-obj",
"message_id": message_id,
"position": 2,
"dataset_id": "ds-1",
"dataset_name": "dataset",
"document_id": "doc-1",
"document_name": "document",
"data_source_type": "file",
"segment_id": "seg-1",
"score": 0.9,
"hit_count": 1,
"word_count": 10,
"segment_position": 0,
"index_node_hash": "hash",
"content": "content",
"created_at": int(resource_created_at.timestamp()),
}
agent_thought = SimpleNamespace(
id="thought-1",
chain_id=None,
agent_thought = MessageAgentThought(
message_chain_id="chain-1",
message_id=message_id,
position=1,
created_by_role="end_user",
created_by="end-user-1",
thought="thinking",
tool="tool",
tool_labels={"label": "value"},
tool_labels_str=json.dumps({"label": "value"}),
tool_input="{}",
created_at=thought_created_at,
observation="observed",
files=["file-a"],
message_files=json.dumps(["file-a"]),
)
agent_thought.id = "thought-1"
agent_thought.created_at = thought_created_at
message_file_obj = SimpleNamespace(
id="file-obj",
filename="b.txt",
type="file",
url=None,
mime_type=None,
size=None,
transfer_method="local",
belongs_to=None,
upload_file_id=None,
message_files = [
{"id": "file-dict", "filename": "a.txt", "type": "file", "transfer_method": "local"},
{"id": "file-obj", "filename": "b.txt", "type": "file", "transfer_method": "local"},
]
app_model = App(
id="app-1",
tenant_id="tenant-1",
name="Chat App",
mode=AppMode.CHAT,
enable_site=False,
enable_api=False,
)
message = SimpleNamespace(
end_user = EndUser(
id="end-user-1",
tenant_id="tenant-1",
app_id=app_model.id,
type=EndUserType.BROWSER,
name="Web User",
session_id="session-1",
)
conversation = Conversation(
id=conversation_id,
app_id=app_model.id,
mode=AppMode.CHAT,
name="Conversation",
_inputs={},
status="normal",
from_source=ConversationFromSource.API,
from_end_user_id=end_user.id,
)
message = Message(
id=message_id,
app_id=app_model.id,
conversation_id=conversation_id,
parent_message_id=None,
inputs={"foo": "bar"},
_inputs={"foo": "bar"},
query="hello",
re_sign_file_url_answer="answer",
user_feedback=SimpleNamespace(rating="like"),
retriever_resources=[
{"id": "res-dict", "message_id": message_id, "position": 1},
retriever_resource_obj,
],
message={},
answer="answer",
message_unit_price=0,
message_price_unit=0,
answer_unit_price=0,
answer_price_unit=0,
provider_response_latency=0,
total_price=0,
currency="USD",
invoke_from=InvokeFrom.SERVICE_API,
from_source=ConversationFromSource.API,
from_end_user_id=end_user.id,
app_mode=AppMode.CHAT,
message_metadata=json.dumps(
{
"meta": "value",
"retriever_resources": [
{"id": "res-dict", "message_id": message_id, "position": 1},
retriever_resource,
],
}
),
created_at=created_at,
agent_thoughts=[agent_thought],
message_files=[
{"id": "file-dict", "filename": "a.txt", "type": "file", "transfer_method": "local"},
message_file_obj,
],
status="normal",
error=None,
message_metadata_dict={"meta": "value"},
extra_contents=[
)
message.set_extra_contents(
[
HumanInputContent(
workflow_run_id=str(uuid.uuid4()),
submitted=True,
)
],
).model_dump(mode="json")
]
)
feedback = MessageFeedback(
app_id=app_model.id,
conversation_id=conversation_id,
message_id=message_id,
rating=FeedbackRating.LIKE,
from_source=FeedbackFromSource.USER,
from_end_user_id=end_user.id,
)
sqlite_session.add_all([app_model, end_user, conversation, message, feedback, agent_thought])
sqlite_session.commit()
pagination = SimpleNamespace(limit=20, has_more=False, data=[message])
app_model = SimpleNamespace(mode="chat")
end_user = SimpleNamespace()
def message_files_with_session(_self, *, session):
del session
return message_files
monkeypatch.setattr(Message, "message_files_with_session", message_files_with_session)
with (
patch.object(message_module.MessageService, "pagination_by_first_id", return_value=pagination) as mock_page,
patch.object(message_module.db, "session", return_value=sqlite_session),
app.test_request_context(f"/messages?conversation_id={conversation_id}&limit=20"),
):
query = MessageListQuery.model_validate({"conversation_id": conversation_id, "limit": 20})
@@ -172,7 +239,7 @@ def test_message_list_mapping(app: Flask) -> None:
assert item["inputs"] == {"foo": "bar"}
assert item["answer"] == "answer"
assert item["feedback"]["rating"] == "like"
assert item["metadata"] == {"meta": "value"}
assert item["metadata"]["meta"] == "value"
assert item["created_at"] == int(created_at.timestamp())
assert item["retriever_resources"][0]["id"] == "res-dict"
@@ -181,8 +248,8 @@ def test_message_list_mapping(app: Flask) -> None:
assert item["agent_thoughts"][0]["chain_id"] == "chain-1"
assert item["agent_thoughts"][0]["created_at"] == int(thought_created_at.timestamp())
assert item["extra_contents"][0]["workflow_run_id"] == message.extra_contents[0].workflow_run_id
assert item["extra_contents"][0]["submitted"] == message.extra_contents[0].submitted
assert item["extra_contents"][0]["workflow_run_id"] == message.extra_contents[0]["workflow_run_id"]
assert item["extra_contents"][0]["submitted"] == message.extra_contents[0]["submitted"]
assert item["message_files"][0]["id"] == "file-dict"
assert item["message_files"][1]["id"] == "file-obj"