diff --git a/api/tests/unit_tests/core/app/apps/test_base_app_runner.py b/api/tests/unit_tests/core/app/apps/test_base_app_runner.py index deb9ab4d2af..dcd9c2b76af 100644 --- a/api/tests/unit_tests/core/app/apps/test_base_app_runner.py +++ b/api/tests/unit_tests/core/app/apps/test_base_app_runner.py @@ -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=[""]) + model_config = ModelConfigWithCredentialsEntity.model_construct(mode="completion", stop=[""]) 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=[""]) + model_config = ModelConfigWithCredentialsEntity.model_construct(mode="chat", stop=[""]) 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"