mirror of
https://github.com/langgenius/dify.git
synced 2026-08-30 17:11:50 +08:00
test: migrate residual controller sessions and ORM models to SQLite (#40606)
Co-authored-by: Byron Wang <byron@dify.ai>
This commit is contained in:
+12
-5
@@ -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("/"),
|
||||
|
||||
+12
-5
@@ -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):
|
||||
|
||||
+12
-5
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user