test: use sqlite3 session in test_base_app_runner (#38740)

Co-authored-by: Byron.wang <byron@dify.ai>
This commit is contained in:
Asuka Minato 2026-07-22 16:21:22 +09:00 committed by GitHub
parent 4b355b7039
commit c76ff4c38c
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194

View File

@ -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"