mirror of
https://github.com/langgenius/dify.git
synced 2026-07-25 05:28:35 +08:00
test: use sqlite3 session in test_base_app_runner (#38740)
Co-authored-by: Byron.wang <byron@dify.ai>
This commit is contained in:
parent
4b355b7039
commit
c76ff4c38c
@ -2,10 +2,10 @@ from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from contextlib import nullcontext
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
from sqlalchemy import func, select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from core.app.app_config.entities import (
|
||||
AdvancedChatMessageEntity,
|
||||
@ -15,7 +15,13 @@ from core.app.app_config.entities import (
|
||||
)
|
||||
from core.app.apps.base_app_runner import AppRunner
|
||||
from core.app.apps.exc import GenerateTaskStoppedError
|
||||
from core.app.entities.app_invoke_entities import InvokeFrom
|
||||
from core.app.apps.message_based_app_queue_manager import MessageBasedAppQueueManager
|
||||
from core.app.entities.app_invoke_entities import (
|
||||
AppGenerateEntity,
|
||||
EasyUIBasedAppGenerateEntity,
|
||||
InvokeFrom,
|
||||
ModelConfigWithCredentialsEntity,
|
||||
)
|
||||
from core.app.entities.queue_entities import (
|
||||
QueueAgentMessageEvent,
|
||||
QueueLLMChunkEvent,
|
||||
@ -29,9 +35,9 @@ from graphon.model_runtime.entities.message_entities import (
|
||||
PromptMessageRole,
|
||||
TextPromptMessageContent,
|
||||
)
|
||||
from graphon.model_runtime.entities.model_entities import ModelPropertyKey
|
||||
from graphon.model_runtime.entities.model_entities import AIModelEntity, ModelPropertyKey
|
||||
from graphon.model_runtime.errors.invoke import InvokeBadRequestError
|
||||
from models.model import AppMode
|
||||
from models.model import App, AppMode, Message, MessageFile
|
||||
|
||||
|
||||
class _DummyParameterRule:
|
||||
@ -40,13 +46,29 @@ class _DummyParameterRule:
|
||||
self.use_template = use_template
|
||||
|
||||
|
||||
class _QueueRecorder:
|
||||
def __init__(self) -> None:
|
||||
self.events: list[object] = []
|
||||
class _TokenCountingModel:
|
||||
token_count: int
|
||||
|
||||
def publish(self, event, pub_from):
|
||||
_ = pub_from
|
||||
self.events.append(event)
|
||||
def __init__(self, token_count: int) -> None:
|
||||
self.token_count = token_count
|
||||
|
||||
def get_llm_num_tokens(self, messages: list[AssistantPromptMessage]) -> int:
|
||||
return self.token_count
|
||||
|
||||
|
||||
def _queue_manager() -> MessageBasedAppQueueManager:
|
||||
return MessageBasedAppQueueManager(
|
||||
task_id="task-id",
|
||||
user_id="user-id",
|
||||
invoke_from=InvokeFrom.SERVICE_API,
|
||||
conversation_id="conversation-id",
|
||||
app_mode=AppMode.CHAT.value,
|
||||
message_id="message-id",
|
||||
)
|
||||
|
||||
|
||||
def _published_events(queue_manager: MessageBasedAppQueueManager) -> list[object]:
|
||||
return [message.event for message in queue_manager.listen()]
|
||||
|
||||
|
||||
class _ClosableStream:
|
||||
@ -70,11 +92,11 @@ class TestAppRunner:
|
||||
def test_recalc_llm_max_tokens_updates_parameters(self, monkeypatch: pytest.MonkeyPatch):
|
||||
runner = AppRunner()
|
||||
|
||||
model_schema = SimpleNamespace(
|
||||
model_schema = AIModelEntity.model_construct(
|
||||
model_properties={ModelPropertyKey.CONTEXT_SIZE: 100},
|
||||
parameter_rules=[_DummyParameterRule("max_tokens")],
|
||||
)
|
||||
model_config = SimpleNamespace(
|
||||
model_config = ModelConfigWithCredentialsEntity.model_construct(
|
||||
provider_model_bundle=object(),
|
||||
model="mock",
|
||||
model_schema=model_schema,
|
||||
@ -83,7 +105,7 @@ class TestAppRunner:
|
||||
|
||||
monkeypatch.setattr(
|
||||
"core.app.apps.base_app_runner.ModelInstance",
|
||||
lambda provider_model_bundle, model: SimpleNamespace(get_llm_num_tokens=lambda messages: 80),
|
||||
lambda provider_model_bundle, model: _TokenCountingModel(80),
|
||||
)
|
||||
|
||||
runner.recalc_llm_max_tokens(model_config, prompt_messages=[AssistantPromptMessage(content="hi")])
|
||||
@ -93,11 +115,11 @@ class TestAppRunner:
|
||||
def test_recalc_llm_max_tokens_returns_minus_one_when_no_context(self, monkeypatch: pytest.MonkeyPatch):
|
||||
runner = AppRunner()
|
||||
|
||||
model_schema = SimpleNamespace(
|
||||
model_schema = AIModelEntity.model_construct(
|
||||
model_properties={},
|
||||
parameter_rules=[_DummyParameterRule("max_tokens")],
|
||||
)
|
||||
model_config = SimpleNamespace(
|
||||
model_config = ModelConfigWithCredentialsEntity.model_construct(
|
||||
provider_model_bundle=object(),
|
||||
model="mock",
|
||||
model_schema=model_schema,
|
||||
@ -106,17 +128,16 @@ class TestAppRunner:
|
||||
|
||||
monkeypatch.setattr(
|
||||
"core.app.apps.base_app_runner.ModelInstance",
|
||||
lambda provider_model_bundle, model: SimpleNamespace(get_llm_num_tokens=lambda messages: 10),
|
||||
lambda provider_model_bundle, model: _TokenCountingModel(10),
|
||||
)
|
||||
|
||||
assert runner.recalc_llm_max_tokens(model_config, prompt_messages=[]) == -1
|
||||
|
||||
def test_direct_output_streaming_publishes_chunks_and_end(self, monkeypatch: pytest.MonkeyPatch):
|
||||
def test_direct_output_streaming_publishes_chunks_and_end(self):
|
||||
runner = AppRunner()
|
||||
queue = _QueueRecorder()
|
||||
app_generate_entity = SimpleNamespace(model_conf=SimpleNamespace(model="mock"), stream=True)
|
||||
|
||||
monkeypatch.setattr("core.app.apps.base_app_runner.time.sleep", lambda _: None)
|
||||
queue = _queue_manager()
|
||||
model_config = ModelConfigWithCredentialsEntity.model_construct(model="mock")
|
||||
app_generate_entity = EasyUIBasedAppGenerateEntity.model_construct(model_conf=model_config, stream=True)
|
||||
|
||||
runner.direct_output(
|
||||
queue_manager=queue,
|
||||
@ -126,12 +147,13 @@ class TestAppRunner:
|
||||
stream=True,
|
||||
)
|
||||
|
||||
assert any(isinstance(event, QueueLLMChunkEvent) for event in queue.events)
|
||||
assert isinstance(queue.events[-1], QueueMessageEndEvent)
|
||||
events = _published_events(queue)
|
||||
assert any(isinstance(event, QueueLLMChunkEvent) for event in events)
|
||||
assert isinstance(events[-1], QueueMessageEndEvent)
|
||||
|
||||
def test_handle_invoke_result_direct_publishes_end_event(self):
|
||||
runner = AppRunner()
|
||||
queue = _QueueRecorder()
|
||||
queue = _queue_manager()
|
||||
llm_result = LLMResult(
|
||||
model="mock",
|
||||
prompt_messages=[],
|
||||
@ -145,11 +167,11 @@ class TestAppRunner:
|
||||
stream=False,
|
||||
)
|
||||
|
||||
assert isinstance(queue.events[-1], QueueMessageEndEvent)
|
||||
assert isinstance(_published_events(queue)[-1], QueueMessageEndEvent)
|
||||
|
||||
def test_handle_invoke_result_invalid_type_raises(self):
|
||||
runner = AppRunner()
|
||||
queue = _QueueRecorder()
|
||||
queue = _queue_manager()
|
||||
|
||||
with pytest.raises(NotImplementedError):
|
||||
runner._handle_invoke_result(
|
||||
@ -160,7 +182,7 @@ class TestAppRunner:
|
||||
|
||||
def test_organize_prompt_messages_simple_template(self, monkeypatch: pytest.MonkeyPatch):
|
||||
runner = AppRunner()
|
||||
model_config = SimpleNamespace(mode="chat", stop=["STOP"])
|
||||
model_config = ModelConfigWithCredentialsEntity.model_construct(mode="chat", stop=["STOP"])
|
||||
prompt_template_entity = PromptTemplateEntity(
|
||||
prompt_type=PromptTemplateEntity.PromptType.SIMPLE,
|
||||
simple_prompt_template="hello",
|
||||
@ -172,7 +194,7 @@ class TestAppRunner:
|
||||
)
|
||||
|
||||
prompt_messages, stop = runner.organize_prompt_messages(
|
||||
app_record=SimpleNamespace(mode=AppMode.CHAT.value),
|
||||
app_record=App(mode=AppMode.CHAT.value),
|
||||
model_config=model_config,
|
||||
prompt_template_entity=prompt_template_entity,
|
||||
inputs={},
|
||||
@ -185,7 +207,7 @@ class TestAppRunner:
|
||||
|
||||
def test_organize_prompt_messages_advanced_completion_template(self, monkeypatch: pytest.MonkeyPatch):
|
||||
runner = AppRunner()
|
||||
model_config = SimpleNamespace(mode="completion", stop=["<END>"])
|
||||
model_config = ModelConfigWithCredentialsEntity.model_construct(mode="completion", stop=["<END>"])
|
||||
captured: dict[str, object] = {}
|
||||
prompt_template_entity = PromptTemplateEntity(
|
||||
prompt_type=PromptTemplateEntity.PromptType.ADVANCED,
|
||||
@ -202,7 +224,7 @@ class TestAppRunner:
|
||||
monkeypatch.setattr("core.app.apps.base_app_runner.AdvancedPromptTransform.get_prompt", _fake_advanced_prompt)
|
||||
|
||||
prompt_messages, stop = runner.organize_prompt_messages(
|
||||
app_record=SimpleNamespace(mode=AppMode.CHAT.value),
|
||||
app_record=App(mode=AppMode.CHAT.value),
|
||||
model_config=model_config,
|
||||
prompt_template_entity=prompt_template_entity,
|
||||
inputs={},
|
||||
@ -218,7 +240,7 @@ class TestAppRunner:
|
||||
|
||||
def test_organize_prompt_messages_advanced_chat_template(self, monkeypatch: pytest.MonkeyPatch):
|
||||
runner = AppRunner()
|
||||
model_config = SimpleNamespace(mode="chat", stop=["<END>"])
|
||||
model_config = ModelConfigWithCredentialsEntity.model_construct(mode="chat", stop=["<END>"])
|
||||
captured: dict[str, object] = {}
|
||||
prompt_template_entity = PromptTemplateEntity(
|
||||
prompt_type=PromptTemplateEntity.PromptType.ADVANCED,
|
||||
@ -237,7 +259,7 @@ class TestAppRunner:
|
||||
monkeypatch.setattr("core.app.apps.base_app_runner.AdvancedPromptTransform.get_prompt", _fake_advanced_prompt)
|
||||
|
||||
prompt_messages, stop = runner.organize_prompt_messages(
|
||||
app_record=SimpleNamespace(mode=AppMode.CHAT.value),
|
||||
app_record=App(mode=AppMode.CHAT.value),
|
||||
model_config=model_config,
|
||||
prompt_template_entity=prompt_template_entity,
|
||||
inputs={},
|
||||
@ -254,8 +276,8 @@ class TestAppRunner:
|
||||
|
||||
with pytest.raises(InvokeBadRequestError, match="Advanced completion prompt template is required"):
|
||||
runner.organize_prompt_messages(
|
||||
app_record=SimpleNamespace(mode=AppMode.CHAT.value),
|
||||
model_config=SimpleNamespace(mode="completion", stop=[]),
|
||||
app_record=App(mode=AppMode.CHAT.value),
|
||||
model_config=ModelConfigWithCredentialsEntity.model_construct(mode="completion", stop=[]),
|
||||
prompt_template_entity=PromptTemplateEntity(prompt_type=PromptTemplateEntity.PromptType.ADVANCED),
|
||||
inputs={},
|
||||
files=[],
|
||||
@ -263,18 +285,16 @@ class TestAppRunner:
|
||||
|
||||
with pytest.raises(InvokeBadRequestError, match="Advanced chat prompt template is required"):
|
||||
runner.organize_prompt_messages(
|
||||
app_record=SimpleNamespace(mode=AppMode.CHAT.value),
|
||||
model_config=SimpleNamespace(mode="chat", stop=[]),
|
||||
app_record=App(mode=AppMode.CHAT.value),
|
||||
model_config=ModelConfigWithCredentialsEntity.model_construct(mode="chat", stop=[]),
|
||||
prompt_template_entity=PromptTemplateEntity(prompt_type=PromptTemplateEntity.PromptType.ADVANCED),
|
||||
inputs={},
|
||||
files=[],
|
||||
)
|
||||
|
||||
def test_handle_invoke_result_stream_routes_chunks_and_builds_message(
|
||||
self, monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture
|
||||
):
|
||||
def test_handle_invoke_result_stream_routes_chunks_and_builds_message(self, caplog: pytest.LogCaptureFixture):
|
||||
runner = AppRunner()
|
||||
queue = _QueueRecorder()
|
||||
queue = _queue_manager()
|
||||
|
||||
image_content = ImagePromptMessageContent(
|
||||
url="https://example.com/image.png", format="png", mime_type="image/png"
|
||||
@ -286,11 +306,9 @@ class TestAppRunner:
|
||||
prompt_messages=[AssistantPromptMessage(content="prompt")],
|
||||
delta=LLMResultChunkDelta(
|
||||
index=0,
|
||||
message=AssistantPromptMessage.model_construct(
|
||||
message=AssistantPromptMessage(
|
||||
content=[
|
||||
"a",
|
||||
TextPromptMessageContent(data="b"),
|
||||
SimpleNamespace(data="c"),
|
||||
TextPromptMessageContent(data="abc"),
|
||||
image_content,
|
||||
]
|
||||
),
|
||||
@ -305,21 +323,25 @@ class TestAppRunner:
|
||||
agent=False,
|
||||
)
|
||||
|
||||
assert isinstance(queue.events[0], QueueLLMChunkEvent)
|
||||
assert isinstance(queue.events[-1], QueueMessageEndEvent)
|
||||
assert queue.events[-1].llm_result.message.content == "abc"
|
||||
events = _published_events(queue)
|
||||
assert isinstance(events[0], QueueLLMChunkEvent)
|
||||
assert isinstance(events[-1], QueueMessageEndEvent)
|
||||
assert events[-1].llm_result.message.content == "abc"
|
||||
assert "Received multimodal output but missing required parameters" in caplog.messages
|
||||
|
||||
def test_handle_invoke_result_stream_agent_mode_handles_multimodal_errors(
|
||||
self, monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture
|
||||
):
|
||||
runner = AppRunner()
|
||||
queue = _QueueRecorder()
|
||||
queue = _queue_manager()
|
||||
|
||||
def raise_multimodal_error(**kwargs):
|
||||
raise RuntimeError("failed to save image")
|
||||
|
||||
monkeypatch.setattr(
|
||||
runner,
|
||||
"_handle_multimodal_image_content",
|
||||
MagicMock(side_effect=RuntimeError("failed to save image")),
|
||||
raise_multimodal_error,
|
||||
)
|
||||
usage = LLMUsage.empty_usage()
|
||||
|
||||
@ -353,22 +375,37 @@ class TestAppRunner:
|
||||
tenant_id="tenant-id",
|
||||
)
|
||||
|
||||
assert isinstance(queue.events[0], QueueAgentMessageEvent)
|
||||
assert isinstance(queue.events[-1], QueueMessageEndEvent)
|
||||
assert queue.events[-1].llm_result.usage == usage
|
||||
events = _published_events(queue)
|
||||
assert isinstance(events[0], QueueAgentMessageEvent)
|
||||
assert isinstance(events[-1], QueueMessageEndEvent)
|
||||
assert events[-1].llm_result.usage == usage
|
||||
assert "Failed to handle multimodal image output" in caplog.messages
|
||||
|
||||
def test_handle_invoke_result_stream_commits_message_file_before_publish(self, monkeypatch: pytest.MonkeyPatch):
|
||||
@pytest.mark.parametrize("sqlite_session", [()], indirect=True)
|
||||
def test_handle_invoke_result_stream_commits_message_file_before_publish(
|
||||
self,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
sqlite_session: Session,
|
||||
):
|
||||
runner = AppRunner()
|
||||
runner._handle_multimodal_image_content = MagicMock(return_value="message-file-1")
|
||||
session = MagicMock()
|
||||
monkeypatch.setattr(
|
||||
runner,
|
||||
"_handle_multimodal_image_content",
|
||||
lambda **kwargs: "message-file-1",
|
||||
)
|
||||
events: list[str] = []
|
||||
session.commit.side_effect = lambda: events.append("commit")
|
||||
original_commit = sqlite_session.commit
|
||||
|
||||
def commit():
|
||||
events.append("commit")
|
||||
original_commit()
|
||||
|
||||
monkeypatch.setattr(sqlite_session, "commit", commit)
|
||||
monkeypatch.setattr(
|
||||
"core.app.apps.base_app_runner.session_factory.create_session",
|
||||
lambda: nullcontext(session),
|
||||
lambda: nullcontext(sqlite_session),
|
||||
)
|
||||
queue = _QueueRecorder()
|
||||
queue = _queue_manager()
|
||||
original_publish = queue.publish
|
||||
|
||||
def publish(event, pub_from):
|
||||
@ -407,7 +444,7 @@ class TestAppRunner:
|
||||
|
||||
assert events == ["commit", "publish"]
|
||||
|
||||
def test_handle_invoke_result_stream_closes_generator_when_stopped(self):
|
||||
def test_handle_invoke_result_stream_closes_generator_when_stopped(self, monkeypatch: pytest.MonkeyPatch):
|
||||
runner = AppRunner()
|
||||
chunk = LLMResultChunk(
|
||||
model="stream-model",
|
||||
@ -416,9 +453,8 @@ class TestAppRunner:
|
||||
)
|
||||
stream = _ClosableStream([chunk])
|
||||
|
||||
queue_manager = SimpleNamespace(
|
||||
publish=MagicMock(side_effect=GenerateTaskStoppedError("stopped")),
|
||||
)
|
||||
queue_manager = _queue_manager()
|
||||
monkeypatch.setattr(queue_manager, "_is_stopped", lambda: True)
|
||||
|
||||
with pytest.raises(GenerateTaskStoppedError):
|
||||
runner._handle_invoke_result_stream(
|
||||
@ -429,7 +465,11 @@ class TestAppRunner:
|
||||
|
||||
assert stream.closed is True
|
||||
|
||||
def test_handle_multimodal_image_content_fallback_return_branch(self, monkeypatch: pytest.MonkeyPatch):
|
||||
@pytest.mark.parametrize("sqlite_session", [(MessageFile,)], indirect=True)
|
||||
def test_handle_multimodal_image_content_fallback_return_branch(
|
||||
self,
|
||||
sqlite_session: Session,
|
||||
):
|
||||
runner = AppRunner()
|
||||
|
||||
class _ToggleBool:
|
||||
@ -442,19 +482,17 @@ class TestAppRunner:
|
||||
self._index += 1
|
||||
return value
|
||||
|
||||
content = SimpleNamespace(
|
||||
# The fallback is reachable only when the fields change truthiness between the guard and branch checks.
|
||||
content = ImagePromptMessageContent.model_construct(
|
||||
url=_ToggleBool([False, False]),
|
||||
base64_data=_ToggleBool([True, False]),
|
||||
mime_type="image/png",
|
||||
)
|
||||
|
||||
db_session = SimpleNamespace(add=MagicMock(), flush=MagicMock(), refresh=MagicMock())
|
||||
monkeypatch.setattr("core.app.apps.base_app_runner.ToolFileManager", lambda: MagicMock())
|
||||
|
||||
queue_manager = SimpleNamespace(invoke_from=InvokeFrom.SERVICE_API, publish=MagicMock())
|
||||
queue_manager = _queue_manager()
|
||||
|
||||
runner._handle_multimodal_image_content(
|
||||
session=db_session,
|
||||
session=sqlite_session,
|
||||
content=content,
|
||||
message_id="message-id",
|
||||
user_id="user-id",
|
||||
@ -462,20 +500,20 @@ class TestAppRunner:
|
||||
queue_manager=queue_manager,
|
||||
)
|
||||
|
||||
db_session.add.assert_not_called()
|
||||
queue_manager.publish.assert_not_called()
|
||||
message_file_count = sqlite_session.scalar(select(func.count()).select_from(MessageFile))
|
||||
assert message_file_count == 0
|
||||
|
||||
def test_check_hosting_moderation_direct_output_called(self, monkeypatch: pytest.MonkeyPatch):
|
||||
runner = AppRunner()
|
||||
queue = _QueueRecorder()
|
||||
app_generate_entity = SimpleNamespace(stream=False)
|
||||
queue = _queue_manager()
|
||||
app_generate_entity = EasyUIBasedAppGenerateEntity.model_construct(stream=False)
|
||||
direct_output_calls: list[dict[str, object]] = []
|
||||
|
||||
monkeypatch.setattr(
|
||||
"core.app.apps.base_app_runner.HostingModerationFeature.check",
|
||||
lambda self, application_generate_entity, prompt_messages: True,
|
||||
)
|
||||
direct_output = MagicMock()
|
||||
monkeypatch.setattr(runner, "direct_output", direct_output)
|
||||
monkeypatch.setattr(runner, "direct_output", lambda **kwargs: direct_output_calls.append(kwargs))
|
||||
|
||||
result = runner.check_hosting_moderation(
|
||||
application_generate_entity=app_generate_entity,
|
||||
@ -484,7 +522,7 @@ class TestAppRunner:
|
||||
)
|
||||
|
||||
assert result is True
|
||||
assert direct_output.called
|
||||
assert len(direct_output_calls) == 1
|
||||
|
||||
def test_fill_in_inputs_from_external_data_tools(self, monkeypatch: pytest.MonkeyPatch):
|
||||
runner = AppRunner()
|
||||
@ -509,7 +547,7 @@ class TestAppRunner:
|
||||
"core.app.apps.base_app_runner.InputModeration.check",
|
||||
lambda self, app_id, tenant_id, app_config, inputs, query, message_id, trace_manager: (True, {}, ""),
|
||||
)
|
||||
app_generate_entity = SimpleNamespace(app_config=SimpleNamespace(), trace_manager=None)
|
||||
app_generate_entity = AppGenerateEntity.model_construct(app_config=None, trace_manager=None)
|
||||
|
||||
result = runner.moderation_for_inputs(
|
||||
app_id="app",
|
||||
@ -522,7 +560,12 @@ class TestAppRunner:
|
||||
|
||||
assert result == (True, {}, "")
|
||||
|
||||
def test_query_app_annotations_to_reply(self, monkeypatch: pytest.MonkeyPatch):
|
||||
@pytest.mark.parametrize("sqlite_session", [()], indirect=True)
|
||||
def test_query_app_annotations_to_reply(
|
||||
self,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
sqlite_session: Session,
|
||||
):
|
||||
runner = AppRunner()
|
||||
monkeypatch.setattr(
|
||||
"core.app.apps.base_app_runner.AnnotationReplyFeature.query",
|
||||
@ -530,12 +573,12 @@ class TestAppRunner:
|
||||
)
|
||||
|
||||
response = runner.query_app_annotations_to_reply(
|
||||
app_record=SimpleNamespace(),
|
||||
message=SimpleNamespace(),
|
||||
app_record=App(),
|
||||
message=Message(),
|
||||
query="hello",
|
||||
user_id="user",
|
||||
invoke_from=InvokeFrom.WEB_APP,
|
||||
session=MagicMock(),
|
||||
session=sqlite_session,
|
||||
)
|
||||
|
||||
assert response == "reply"
|
||||
|
||||
Loading…
Reference in New Issue
Block a user