mirror of
https://github.com/langgenius/dify.git
synced 2026-09-05 08:48:10 +08:00
test: migrate telemetry and core entity ORM models to SQLite (#40574)
This commit is contained in:
parent
eeda983b17
commit
beecf5730a
@ -1,5 +1,4 @@
|
||||
from datetime import UTC, datetime
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
import pytest
|
||||
@ -11,6 +10,7 @@ from core.entities.mcp_provider import (
|
||||
MCPProviderEntity,
|
||||
)
|
||||
from core.mcp.types import OAuthTokens
|
||||
from models.tools import MCPToolProvider
|
||||
|
||||
|
||||
def _build_mcp_provider_entity() -> MCPProviderEntity:
|
||||
@ -37,24 +37,25 @@ def _build_mcp_provider_entity() -> MCPProviderEntity:
|
||||
def test_from_db_model_maps_fields() -> None:
|
||||
# Arrange
|
||||
now = datetime(2025, 1, 1, tzinfo=UTC)
|
||||
db_provider = SimpleNamespace(
|
||||
id="provider-1",
|
||||
db_provider = MCPToolProvider(
|
||||
server_identifier="server-1",
|
||||
name="Example MCP",
|
||||
tenant_id="tenant-1",
|
||||
user_id="user-1",
|
||||
server_url="encrypted-server-url",
|
||||
headers={"Authorization": "enc"},
|
||||
server_url_hash="server-url-hash",
|
||||
encrypted_headers='{"Authorization": "enc"}',
|
||||
timeout=15,
|
||||
sse_read_timeout=120,
|
||||
authed=True,
|
||||
credentials={"access_token": "enc-token"},
|
||||
tool_dict=[{"name": "search"}],
|
||||
encrypted_credentials='{"access_token": "enc-token"}',
|
||||
tools='[{"name": "search"}]',
|
||||
icon=None,
|
||||
created_at=now,
|
||||
updated_at=now,
|
||||
identity_mode="off",
|
||||
)
|
||||
db_provider.id = "provider-1"
|
||||
db_provider.created_at = now
|
||||
db_provider.updated_at = now
|
||||
|
||||
# Act
|
||||
entity = MCPProviderEntity.from_db_model(db_provider)
|
||||
|
||||
@ -11,6 +11,7 @@ import pytest
|
||||
from sqlalchemy import Engine, event
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
import models.model as model_module
|
||||
from core.memory import token_buffer_memory as memory_module
|
||||
from core.memory.token_buffer_memory import TokenBufferMemory
|
||||
from graphon.file import FileTransferMethod, FileType
|
||||
@ -23,8 +24,14 @@ from graphon.model_runtime.entities import (
|
||||
)
|
||||
from models.base import TypeBase
|
||||
from models.enums import ConversationFromSource, CreatorUserRole, MessageFileBelongsTo
|
||||
from models.model import AppMode, Message, MessageFile
|
||||
from models.workflow import Workflow, WorkflowType
|
||||
from models.model import App, AppMode, Conversation, Message, MessageFile
|
||||
from models.workflow import (
|
||||
Workflow,
|
||||
WorkflowExecutionStatus,
|
||||
WorkflowRun,
|
||||
WorkflowRunTriggeredFrom,
|
||||
WorkflowType,
|
||||
)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers / shared fixtures
|
||||
@ -44,7 +51,7 @@ class Database:
|
||||
def database(sqlite_engine: Engine, monkeypatch: pytest.MonkeyPatch) -> Iterator[Database]:
|
||||
TypeBase.metadata.create_all(
|
||||
sqlite_engine,
|
||||
tables=[Message.__table__, MessageFile.__table__, Workflow.__table__],
|
||||
tables=[App.__table__, Conversation.__table__, Message.__table__, MessageFile.__table__, Workflow.__table__],
|
||||
)
|
||||
statements: list[tuple[str, object]] = []
|
||||
|
||||
@ -55,17 +62,42 @@ def database(sqlite_engine: Engine, monkeypatch: pytest.MonkeyPatch) -> Iterator
|
||||
with Session(sqlite_engine, expire_on_commit=False) as session:
|
||||
database = Database(engine=sqlite_engine, session=session, statements=statements)
|
||||
monkeypatch.setattr(memory_module, "db", database)
|
||||
monkeypatch.setattr(model_module, "db", database)
|
||||
yield database
|
||||
event.remove(sqlite_engine, "before_cursor_execute", record_statement)
|
||||
|
||||
|
||||
def _make_conversation(mode: AppMode = AppMode.CHAT) -> MagicMock:
|
||||
"""Return a minimal Conversation mock."""
|
||||
conv = MagicMock()
|
||||
conv.id = str(uuid4())
|
||||
conv.mode = mode
|
||||
conv.model_config = {}
|
||||
return conv
|
||||
def _make_app(*, app_id: str | None = None, mode: AppMode = AppMode.CHAT) -> App:
|
||||
"""Return a real transient app with the ownership fields used by memory."""
|
||||
return App(
|
||||
id=app_id or str(uuid4()),
|
||||
tenant_id=str(uuid4()),
|
||||
name="Memory test app",
|
||||
mode=mode,
|
||||
enable_site=False,
|
||||
enable_api=False,
|
||||
)
|
||||
|
||||
|
||||
def _make_conversation(mode: AppMode = AppMode.CHAT, *, app_id: str | None = None) -> Conversation:
|
||||
"""Return a real transient conversation configured without database-backed model settings."""
|
||||
return Conversation(
|
||||
id=str(uuid4()),
|
||||
app_id=app_id or str(uuid4()),
|
||||
mode=mode,
|
||||
name="Memory test conversation",
|
||||
override_model_configs="{}",
|
||||
_inputs={},
|
||||
from_source=ConversationFromSource.API,
|
||||
)
|
||||
|
||||
|
||||
def _persist_conversation(database: Database, mode: AppMode = AppMode.CHAT) -> Conversation:
|
||||
app = _make_app(mode=mode)
|
||||
conversation = _make_conversation(mode, app_id=app.id)
|
||||
database.session.add_all([app, conversation])
|
||||
database.session.commit()
|
||||
return conversation
|
||||
|
||||
|
||||
def _make_model_instance() -> MagicMock:
|
||||
@ -75,15 +107,51 @@ def _make_model_instance() -> MagicMock:
|
||||
return mi
|
||||
|
||||
|
||||
def _make_message(answer: str = "hello", answer_tokens: int = 5) -> MagicMock:
|
||||
msg = MagicMock()
|
||||
msg.id = str(uuid4())
|
||||
msg.query = "user query"
|
||||
msg.answer = answer
|
||||
msg.answer_tokens = answer_tokens
|
||||
msg.workflow_run_id = str(uuid4())
|
||||
msg.created_at = MagicMock()
|
||||
return msg
|
||||
def _make_message(answer: str = "hello", answer_tokens: int = 5) -> Message:
|
||||
return Message(
|
||||
id=str(uuid4()),
|
||||
app_id=str(uuid4()),
|
||||
conversation_id=str(uuid4()),
|
||||
_inputs={},
|
||||
query="user query",
|
||||
message={},
|
||||
message_unit_price=Decimal(0),
|
||||
answer=answer,
|
||||
answer_tokens=answer_tokens,
|
||||
answer_unit_price=Decimal(0),
|
||||
currency="USD",
|
||||
from_source=ConversationFromSource.API,
|
||||
workflow_run_id=str(uuid4()),
|
||||
created_at=datetime.now(UTC).replace(tzinfo=None),
|
||||
)
|
||||
|
||||
|
||||
def _make_message_file() -> MessageFile:
|
||||
return MessageFile(
|
||||
message_id=str(uuid4()),
|
||||
type=FileType.IMAGE,
|
||||
transfer_method=FileTransferMethod.REMOTE_URL,
|
||||
created_by_role=CreatorUserRole.ACCOUNT,
|
||||
created_by=str(uuid4()),
|
||||
belongs_to=MessageFileBelongsTo.USER,
|
||||
url="https://example.com/image.png",
|
||||
)
|
||||
|
||||
|
||||
def _make_workflow_run(*, workflow_id: str | None = None) -> WorkflowRun:
|
||||
return WorkflowRun(
|
||||
tenant_id=str(uuid4()),
|
||||
app_id=str(uuid4()),
|
||||
workflow_id=workflow_id or str(uuid4()),
|
||||
type=WorkflowType.CHAT,
|
||||
triggered_from=WorkflowRunTriggeredFrom.APP_RUN,
|
||||
version="1",
|
||||
graph="{}",
|
||||
inputs="{}",
|
||||
status=WorkflowExecutionStatus.SUCCEEDED,
|
||||
created_by_role=CreatorUserRole.ACCOUNT,
|
||||
created_by=str(uuid4()),
|
||||
)
|
||||
|
||||
|
||||
def _persist_message(
|
||||
@ -229,7 +297,7 @@ class TestBuildPromptMessageWithFiles:
|
||||
message_files=[],
|
||||
text_content="hello",
|
||||
message=_make_message(),
|
||||
app_record=MagicMock(),
|
||||
app_record=_make_app(),
|
||||
is_user_message=True,
|
||||
)
|
||||
|
||||
@ -272,9 +340,8 @@ class TestBuildPromptMessageWithFiles:
|
||||
url="http://example.com/img.png", format="png", mime_type="image/png"
|
||||
)
|
||||
|
||||
mock_message_file = MagicMock()
|
||||
mock_app_record = MagicMock()
|
||||
mock_app_record.tenant_id = "tenant-1"
|
||||
message_file = _make_message_file()
|
||||
app_record = _make_app()
|
||||
|
||||
with (
|
||||
patch(
|
||||
@ -291,10 +358,10 @@ class TestBuildPromptMessageWithFiles:
|
||||
),
|
||||
):
|
||||
result = mem._build_prompt_message_with_files(
|
||||
message_files=[mock_message_file],
|
||||
message_files=[message_file],
|
||||
text_content="user text",
|
||||
message=_make_message(),
|
||||
app_record=mock_app_record,
|
||||
app_record=app_record,
|
||||
is_user_message=True,
|
||||
)
|
||||
|
||||
@ -318,8 +385,7 @@ class TestBuildPromptMessageWithFiles:
|
||||
real_image_content = ImagePromptMessageContent(
|
||||
url="http://example.com/img.png", format="png", mime_type="image/png"
|
||||
)
|
||||
mock_app_record = MagicMock()
|
||||
mock_app_record.tenant_id = "tenant-1"
|
||||
app_record = _make_app()
|
||||
|
||||
with (
|
||||
patch(
|
||||
@ -336,10 +402,10 @@ class TestBuildPromptMessageWithFiles:
|
||||
),
|
||||
):
|
||||
mem._build_prompt_message_with_files(
|
||||
message_files=[MagicMock()],
|
||||
message_files=[_make_message_file()],
|
||||
text_content="user text",
|
||||
message=_make_message(),
|
||||
app_record=mock_app_record,
|
||||
app_record=app_record,
|
||||
is_user_message=True,
|
||||
)
|
||||
|
||||
@ -359,8 +425,7 @@ class TestBuildPromptMessageWithFiles:
|
||||
real_image_content = ImagePromptMessageContent(
|
||||
url="http://example.com/img.png", format="png", mime_type="image/png"
|
||||
)
|
||||
mock_app_record = MagicMock()
|
||||
mock_app_record.tenant_id = "tenant-1"
|
||||
app_record = _make_app()
|
||||
|
||||
with (
|
||||
patch(
|
||||
@ -377,10 +442,10 @@ class TestBuildPromptMessageWithFiles:
|
||||
),
|
||||
):
|
||||
result = mem._build_prompt_message_with_files(
|
||||
message_files=[MagicMock()],
|
||||
message_files=[_make_message_file()],
|
||||
text_content="ai text",
|
||||
message=_make_message(),
|
||||
app_record=mock_app_record,
|
||||
app_record=app_record,
|
||||
is_user_message=False,
|
||||
)
|
||||
|
||||
@ -399,8 +464,7 @@ class TestBuildPromptMessageWithFiles:
|
||||
mock_file_extra_config = MagicMock()
|
||||
mock_file_extra_config.image_config = mock_image_config
|
||||
|
||||
mock_app_record = MagicMock()
|
||||
mock_app_record.tenant_id = "tenant-1"
|
||||
app_record = _make_app()
|
||||
|
||||
real_image_content = ImagePromptMessageContent(
|
||||
url="http://example.com/img.png", format="png", mime_type="image/png"
|
||||
@ -421,10 +485,10 @@ class TestBuildPromptMessageWithFiles:
|
||||
) as mock_to_prompt,
|
||||
):
|
||||
mem._build_prompt_message_with_files(
|
||||
message_files=[MagicMock()],
|
||||
message_files=[_make_message_file()],
|
||||
text_content="user text",
|
||||
message=_make_message(),
|
||||
app_record=mock_app_record,
|
||||
app_record=app_record,
|
||||
is_user_message=True,
|
||||
)
|
||||
# Ensure the LOW detail was passed through
|
||||
@ -445,7 +509,7 @@ class TestBuildPromptMessageWithFiles:
|
||||
return_value=mock_file_extra_config,
|
||||
):
|
||||
result = mem._build_prompt_message_with_files(
|
||||
message_files=[MagicMock()],
|
||||
message_files=[_make_message_file()],
|
||||
text_content="hello",
|
||||
message=_make_message(),
|
||||
app_record=None, # <-- forces the else branch → file_objs = []
|
||||
@ -460,10 +524,9 @@ class TestBuildPromptMessageWithFiles:
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
@pytest.mark.parametrize("mode", [AppMode.ADVANCED_CHAT, AppMode.WORKFLOW])
|
||||
def test_workflow_mode_no_app_raises(self, mode):
|
||||
def test_workflow_mode_no_app_raises(self, mode, database: Database):
|
||||
"""Raises ValueError when conversation.app is falsy."""
|
||||
conv = _make_conversation(mode)
|
||||
conv.app = None
|
||||
mem = TokenBufferMemory(conversation=conv, model_instance=_make_model_instance())
|
||||
|
||||
with pytest.raises(ValueError, match="App not found for conversation"):
|
||||
@ -471,15 +534,14 @@ class TestBuildPromptMessageWithFiles:
|
||||
message_files=[],
|
||||
text_content="text",
|
||||
message=_make_message(),
|
||||
app_record=MagicMock(),
|
||||
app_record=_make_app(),
|
||||
is_user_message=True,
|
||||
)
|
||||
|
||||
@pytest.mark.parametrize("mode", [AppMode.ADVANCED_CHAT, AppMode.WORKFLOW])
|
||||
def test_workflow_mode_no_workflow_run_id_raises(self, mode):
|
||||
def test_workflow_mode_no_workflow_run_id_raises(self, mode, database: Database):
|
||||
"""Raises ValueError when message.workflow_run_id is falsy."""
|
||||
conv = _make_conversation(mode)
|
||||
conv.app = MagicMock()
|
||||
conv = _persist_conversation(database, mode)
|
||||
|
||||
message = _make_message()
|
||||
message.workflow_run_id = None # force missing
|
||||
@ -491,16 +553,14 @@ class TestBuildPromptMessageWithFiles:
|
||||
message_files=[],
|
||||
text_content="text",
|
||||
message=message,
|
||||
app_record=MagicMock(),
|
||||
app_record=_make_app(),
|
||||
is_user_message=True,
|
||||
)
|
||||
|
||||
@pytest.mark.parametrize("mode", [AppMode.ADVANCED_CHAT, AppMode.WORKFLOW])
|
||||
def test_workflow_mode_workflow_run_not_found_raises(self, mode):
|
||||
def test_workflow_mode_workflow_run_not_found_raises(self, mode, database: Database):
|
||||
"""Raises ValueError when workflow_run_repo returns None."""
|
||||
conv = _make_conversation(mode)
|
||||
mock_app = MagicMock()
|
||||
conv.app = mock_app
|
||||
conv = _persist_conversation(database, mode)
|
||||
|
||||
mem = TokenBufferMemory(conversation=conv, model_instance=_make_model_instance())
|
||||
mem._workflow_run_repo = MagicMock()
|
||||
@ -511,46 +571,39 @@ class TestBuildPromptMessageWithFiles:
|
||||
message_files=[],
|
||||
text_content="text",
|
||||
message=_make_message(),
|
||||
app_record=MagicMock(),
|
||||
app_record=_make_app(),
|
||||
is_user_message=True,
|
||||
)
|
||||
|
||||
@pytest.mark.parametrize("mode", [AppMode.ADVANCED_CHAT, AppMode.WORKFLOW])
|
||||
def test_workflow_mode_workflow_not_found_raises(self, mode, database: Database):
|
||||
"""Raises ValueError when Workflow lookup returns None."""
|
||||
conv = _make_conversation(mode)
|
||||
conv.app = MagicMock()
|
||||
|
||||
mock_workflow_run = MagicMock()
|
||||
mock_workflow_run.workflow_id = str(uuid4())
|
||||
conv = _persist_conversation(database, mode)
|
||||
workflow_run = _make_workflow_run()
|
||||
|
||||
mem = TokenBufferMemory(conversation=conv, model_instance=_make_model_instance())
|
||||
mem._workflow_run_repo = MagicMock()
|
||||
mem._workflow_run_repo.get_workflow_run_by_id.return_value = mock_workflow_run
|
||||
mem._workflow_run_repo.get_workflow_run_by_id.return_value = workflow_run
|
||||
|
||||
with pytest.raises(ValueError, match="Workflow not found"):
|
||||
mem._build_prompt_message_with_files(
|
||||
message_files=[],
|
||||
text_content="text",
|
||||
message=_make_message(),
|
||||
app_record=MagicMock(),
|
||||
app_record=_make_app(),
|
||||
is_user_message=True,
|
||||
)
|
||||
|
||||
@pytest.mark.parametrize("mode", [AppMode.ADVANCED_CHAT, AppMode.WORKFLOW])
|
||||
def test_workflow_mode_success_no_files_user(self, mode, database: Database):
|
||||
"""Happy path: workflow mode, no message files → plain UserPromptMessage."""
|
||||
conv = _make_conversation(mode)
|
||||
conv.app = MagicMock()
|
||||
|
||||
mock_workflow_run = MagicMock()
|
||||
mock_workflow_run.workflow_id = str(uuid4())
|
||||
|
||||
workflow = _persist_workflow(database, workflow_id=mock_workflow_run.workflow_id)
|
||||
conv = _persist_conversation(database, mode)
|
||||
workflow_run = _make_workflow_run()
|
||||
workflow = _persist_workflow(database, workflow_id=workflow_run.workflow_id)
|
||||
|
||||
mem = TokenBufferMemory(conversation=conv, model_instance=_make_model_instance())
|
||||
mem._workflow_run_repo = MagicMock()
|
||||
mem._workflow_run_repo.get_workflow_run_by_id.return_value = mock_workflow_run
|
||||
mem._workflow_run_repo.get_workflow_run_by_id.return_value = workflow_run
|
||||
|
||||
with patch(
|
||||
"core.memory.token_buffer_memory.FileUploadConfigManager.convert",
|
||||
@ -560,7 +613,7 @@ class TestBuildPromptMessageWithFiles:
|
||||
message_files=[],
|
||||
text_content="wf text",
|
||||
message=_make_message(),
|
||||
app_record=MagicMock(),
|
||||
app_record=_make_app(),
|
||||
is_user_message=True,
|
||||
)
|
||||
|
||||
@ -583,7 +636,7 @@ class TestBuildPromptMessageWithFiles:
|
||||
message_files=[],
|
||||
text_content="text",
|
||||
message=_make_message(),
|
||||
app_record=MagicMock(),
|
||||
app_record=_make_app(),
|
||||
is_user_message=True,
|
||||
)
|
||||
|
||||
@ -596,23 +649,22 @@ class TestBuildPromptMessageWithFiles:
|
||||
class TestGetHistoryPromptMessages:
|
||||
"""Tests for persisted history retrieval, file batching, and pruning."""
|
||||
|
||||
def _make_memory(self, mode: AppMode = AppMode.CHAT) -> TokenBufferMemory:
|
||||
conv = _make_conversation(mode)
|
||||
conv.app = MagicMock()
|
||||
def _make_memory(self, database: Database, mode: AppMode = AppMode.CHAT) -> TokenBufferMemory:
|
||||
conv = _persist_conversation(database, mode)
|
||||
return TokenBufferMemory(conversation=conv, model_instance=_make_model_instance())
|
||||
|
||||
def test_returns_empty_when_no_messages(self, database: Database) -> None:
|
||||
assert self._make_memory().get_history_prompt_messages() == []
|
||||
assert self._make_memory(database).get_history_prompt_messages() == []
|
||||
|
||||
def test_skips_newest_message_without_answer(self, database: Database) -> None:
|
||||
mem = self._make_memory()
|
||||
mem = self._make_memory(database)
|
||||
message = _persist_message(database, mem.conversation.id, answer="", answer_tokens=0)
|
||||
|
||||
assert mem.get_history_prompt_messages() == []
|
||||
assert database.session.get(Message, message.id) is message
|
||||
|
||||
def test_message_with_answer_returns_user_and_assistant_prompts(self, database: Database) -> None:
|
||||
mem = self._make_memory()
|
||||
mem = self._make_memory(database)
|
||||
_persist_message(database, mem.conversation.id, query="My query", answer="My answer", answer_tokens=10)
|
||||
|
||||
result = mem.get_history_prompt_messages()
|
||||
@ -624,7 +676,7 @@ class TestGetHistoryPromptMessages:
|
||||
assert result[1].content == "My answer"
|
||||
|
||||
def test_history_is_conversation_scoped(self, database: Database) -> None:
|
||||
mem = self._make_memory()
|
||||
mem = self._make_memory(database)
|
||||
_persist_message(database, mem.conversation.id, answer="visible")
|
||||
_persist_message(database, "other-conversation", answer="hidden")
|
||||
|
||||
@ -642,12 +694,12 @@ class TestGetHistoryPromptMessages:
|
||||
message_limit: int | None,
|
||||
expected_limit: int,
|
||||
) -> None:
|
||||
mem = self._make_memory()
|
||||
mem = self._make_memory(database)
|
||||
before = len(database.statements)
|
||||
|
||||
mem.get_history_prompt_messages(message_limit=message_limit)
|
||||
|
||||
statements = database.statements[before:]
|
||||
statements = [entry for entry in database.statements[before:] if "FROM messages" in entry[0]]
|
||||
assert len(statements) == 1
|
||||
sql, parameters = statements[0]
|
||||
assert "LIMIT" in sql
|
||||
@ -667,7 +719,7 @@ class TestGetHistoryPromptMessages:
|
||||
belongs_to: MessageFileBelongsTo | None,
|
||||
is_user_message: bool,
|
||||
) -> None:
|
||||
mem = self._make_memory()
|
||||
mem = self._make_memory(database)
|
||||
message = _persist_message(database, mem.conversation.id)
|
||||
message_file = _persist_message_file(database, message, belongs_to=belongs_to)
|
||||
built_prompt = (
|
||||
@ -685,7 +737,7 @@ class TestGetHistoryPromptMessages:
|
||||
assert built_prompt in result
|
||||
|
||||
def test_message_files_are_batch_loaded_with_constant_query_count(self, database: Database) -> None:
|
||||
mem = self._make_memory()
|
||||
mem = self._make_memory(database)
|
||||
base_time = datetime.now(UTC).replace(tzinfo=None)
|
||||
messages = [
|
||||
_persist_message(
|
||||
@ -703,7 +755,8 @@ class TestGetHistoryPromptMessages:
|
||||
result = mem.get_history_prompt_messages()
|
||||
|
||||
selects = [sql for sql, _ in database.statements[before:] if sql.lstrip().upper().startswith("SELECT")]
|
||||
assert len(selects) == 3
|
||||
assert len(selects) == 4
|
||||
assert sum("FROM apps" in sql for sql in selects) == 1
|
||||
assert len(result) == 10
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
@ -721,7 +774,7 @@ class TestGetHistoryPromptMessages:
|
||||
max_token_limit: int,
|
||||
expected_length: int,
|
||||
) -> None:
|
||||
mem = self._make_memory()
|
||||
mem = self._make_memory(database)
|
||||
mem.model_instance.get_llm_num_tokens.side_effect = token_values
|
||||
_persist_message(database, mem.conversation.id)
|
||||
|
||||
@ -740,7 +793,6 @@ class TestGetHistoryPromptText:
|
||||
|
||||
def _make_memory(self) -> TokenBufferMemory:
|
||||
conv = _make_conversation()
|
||||
conv.app = MagicMock()
|
||||
return TokenBufferMemory(conversation=conv, model_instance=_make_model_instance())
|
||||
|
||||
def test_empty_messages_returns_empty_string(self):
|
||||
|
||||
@ -1,9 +1,9 @@
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import patch
|
||||
|
||||
from core.plugin.backwards_invocation.model import PluginModelBackwardsInvocation
|
||||
from core.plugin.entities.request import RequestInvokeSummary
|
||||
from graphon.model_runtime.entities.message_entities import UserPromptMessage
|
||||
from models.account import Tenant
|
||||
|
||||
|
||||
def test_system_model_helpers_forward_user_id():
|
||||
@ -36,7 +36,8 @@ def test_system_model_helpers_forward_user_id():
|
||||
|
||||
|
||||
def test_invoke_summary_uses_same_user_scope_for_token_helpers():
|
||||
tenant = SimpleNamespace(id="tenant-1")
|
||||
tenant = Tenant(name="Test Workspace")
|
||||
tenant.id = "tenant-1"
|
||||
payload = RequestInvokeSummary(text="short", instruction="keep it concise")
|
||||
|
||||
with (
|
||||
|
||||
@ -1,6 +1,5 @@
|
||||
import sys
|
||||
import types
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
@ -9,6 +8,7 @@ from sqlalchemy.orm import Session
|
||||
from core.rag.datasource.keyword.keyword_factory import Keyword
|
||||
from core.rag.datasource.keyword.keyword_type import KeyWordType
|
||||
from core.rag.models.document import Document
|
||||
from models.dataset import Dataset
|
||||
|
||||
|
||||
def test_get_keyword_factory_returns_jieba_factory(monkeypatch: pytest.MonkeyPatch):
|
||||
@ -29,7 +29,13 @@ def test_get_keyword_factory_raises_for_unsupported_type():
|
||||
|
||||
|
||||
def test_keyword_initialization_uses_configured_factory(monkeypatch: pytest.MonkeyPatch):
|
||||
dataset = SimpleNamespace(id="dataset-1")
|
||||
dataset = Dataset(
|
||||
id="dataset-1",
|
||||
tenant_id="tenant-1",
|
||||
name="Test Dataset",
|
||||
description="",
|
||||
created_by="account-1",
|
||||
)
|
||||
fake_processor = MagicMock()
|
||||
|
||||
monkeypatch.setattr("core.rag.datasource.keyword.keyword_factory.dify_config.KEYWORD_STORE", KeyWordType.JIEBA)
|
||||
|
||||
@ -2,39 +2,60 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from datetime import UTC, datetime
|
||||
from typing import cast
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from graphon.enums import WorkflowNodeExecutionMetadataKey
|
||||
from models.enums import CreatorUserRole
|
||||
from models.workflow import (
|
||||
WorkflowNodeExecutionModel,
|
||||
WorkflowNodeExecutionStatus,
|
||||
WorkflowNodeExecutionTriggeredFrom,
|
||||
)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _make_execution(**overrides) -> MagicMock:
|
||||
"""Return a minimal WorkflowNodeExecutionModel mock."""
|
||||
execution = MagicMock()
|
||||
execution.tenant_id = overrides.get("tenant_id", "tenant-1")
|
||||
execution.app_id = overrides.get("app_id", "app-1")
|
||||
execution.workflow_id = overrides.get("workflow_id", "wf-1")
|
||||
execution.id = overrides.get("id", "exec-1")
|
||||
execution.node_id = overrides.get("node_id", "node-1")
|
||||
execution.node_type = overrides.get("node_type", "llm")
|
||||
execution.title = overrides.get("title", "My LLM Node")
|
||||
execution.status = overrides.get("status", "succeeded")
|
||||
execution.error = overrides.get("error")
|
||||
execution.elapsed_time = overrides.get("elapsed_time", 1.5)
|
||||
execution.index = overrides.get("index", 1)
|
||||
execution.predecessor_node_id = overrides.get("predecessor_node_id")
|
||||
execution.created_at = overrides.get("created_at", datetime(2024, 1, 1, tzinfo=UTC))
|
||||
execution.finished_at = overrides.get("finished_at", datetime(2024, 1, 1, 0, 0, 5, tzinfo=UTC))
|
||||
execution.workflow_run_id = overrides.get("workflow_run_id", "run-1")
|
||||
execution.inputs_dict = overrides.get("inputs_dict", {"prompt": "hello"})
|
||||
execution.outputs_dict = overrides.get("outputs_dict", {"answer": "world"})
|
||||
execution.process_data_dict = overrides.get("process_data_dict", {})
|
||||
execution.execution_metadata_dict = overrides.get("execution_metadata_dict", {})
|
||||
return execution
|
||||
def _make_execution(**overrides: object) -> WorkflowNodeExecutionModel:
|
||||
"""Return a real transient execution with JSON-backed model properties."""
|
||||
|
||||
inputs = overrides.get("inputs_dict", {"prompt": "hello"})
|
||||
outputs = overrides.get("outputs_dict", {"answer": "world"})
|
||||
process_data = overrides.get("process_data_dict", {})
|
||||
metadata = overrides.get("execution_metadata_dict", {})
|
||||
return WorkflowNodeExecutionModel(
|
||||
id=cast(str, overrides.get("id", "exec-1")),
|
||||
tenant_id=cast(str, overrides.get("tenant_id", "tenant-1")),
|
||||
app_id=cast(str, overrides.get("app_id", "app-1")),
|
||||
workflow_id=cast(str, overrides.get("workflow_id", "wf-1")),
|
||||
triggered_from=WorkflowNodeExecutionTriggeredFrom.WORKFLOW_RUN,
|
||||
workflow_run_id=cast(str | None, overrides.get("workflow_run_id", "run-1")),
|
||||
index=cast(int, overrides.get("index", 1)),
|
||||
predecessor_node_id=cast(str | None, overrides.get("predecessor_node_id")),
|
||||
node_execution_id=None,
|
||||
node_id=cast(str, overrides.get("node_id", "node-1")),
|
||||
node_type=cast(str, overrides.get("node_type", "llm")),
|
||||
title=cast(str, overrides.get("title", "My LLM Node")),
|
||||
agent_workspace_binding_id=None,
|
||||
inputs=json.dumps(inputs) if inputs is not None else None,
|
||||
process_data=json.dumps(process_data) if process_data is not None else None,
|
||||
outputs=json.dumps(outputs) if outputs is not None else None,
|
||||
status=WorkflowNodeExecutionStatus(cast(str, overrides.get("status", "succeeded"))),
|
||||
error=cast(str | None, overrides.get("error")),
|
||||
elapsed_time=cast(float, overrides.get("elapsed_time", 1.5)),
|
||||
execution_metadata=json.dumps(metadata),
|
||||
created_at=cast(datetime, overrides.get("created_at", datetime(2024, 1, 1, tzinfo=UTC))),
|
||||
created_by_role=CreatorUserRole.ACCOUNT,
|
||||
created_by="user-1",
|
||||
finished_at=cast(
|
||||
datetime | None,
|
||||
overrides.get("finished_at", datetime(2024, 1, 1, 0, 0, 5, tzinfo=UTC)),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@ -289,8 +310,8 @@ class TestEnqueueDraftNodeExecutionTrace:
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _make_llm_execution() -> MagicMock:
|
||||
"""Return a WorkflowNodeExecutionModel mock that mimics a real LLM node.
|
||||
def _make_llm_execution() -> WorkflowNodeExecutionModel:
|
||||
"""Return a real WorkflowNodeExecutionModel that mimics an LLM node.
|
||||
|
||||
The field values match what graphon/nodes/llm/node.py produces:
|
||||
- process_data_dict contains model_provider, model_name, and usage
|
||||
|
||||
@ -15,6 +15,7 @@ from pydantic import ValidationError
|
||||
|
||||
import libs.helper as helper_module
|
||||
from libs.helper import TokenManager
|
||||
from models.account import Account
|
||||
|
||||
|
||||
def _build_fake_redis(storage: dict[str, str]):
|
||||
@ -70,7 +71,8 @@ def test_token_manager_roundtrip_uses_explicit_email_with_account(monkeypatch: p
|
||||
storage: dict[str, str] = {}
|
||||
monkeypatch.setattr(helper_module, "redis_client", _build_fake_redis(storage))
|
||||
|
||||
account = SimpleNamespace(id="acc-1", email="old@example.com")
|
||||
account = Account(name="Test User", email="old@example.com")
|
||||
account.id = "acc-1"
|
||||
|
||||
token = TokenManager.generate_token(
|
||||
account=account,
|
||||
|
||||
Loading…
Reference in New Issue
Block a user