test: migrate application generation sessions and ORM models to SQLite (#40509)

This commit is contained in:
Asuka Minato 2026-09-04 12:16:31 +00:00 committed by GitHub
parent 1b4ae3bcd2
commit 08c7a2c5f3
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
4 changed files with 448 additions and 275 deletions

View File

@ -12,9 +12,13 @@ from __future__ import annotations
import contextlib import contextlib
import inspect import inspect
import json import json
from decimal import Decimal
from types import SimpleNamespace
import pytest import pytest
from pytest_mock import MockerFixture from pytest_mock import MockerFixture
from sqlalchemy import event
from sqlalchemy.orm import Session
import core.app.apps.agent_app.app_generator as module import core.app.apps.agent_app.app_generator as module
from core.app.apps.agent_app.app_generator import ( from core.app.apps.agent_app.app_generator import (
@ -27,14 +31,107 @@ from core.app.entities.app_invoke_entities import InvokeFrom, UserFrom
from core.app.entities.queue_entities import QueueAnnotationReplyEvent from core.app.entities.queue_entities import QueueAnnotationReplyEvent
from core.workflow.file_reference import build_file_reference from core.workflow.file_reference import build_file_reference
from models import Account, AppModelConfig from models import Account, AppModelConfig
from models.agent import Agent, AgentConfigSnapshot, AgentScope, AgentSource, AgentStatus
from models.agent_config_entities import AgentSoulConfig
from models.enums import AppStatus, ConversationFromSource
from models.model import App, AppMode, Conversation, Message, MessageAnnotation
MODULE = "core.app.apps.agent_app.app_generator" MODULE = "core.app.apps.agent_app.app_generator"
class DummyAccount: def _account(user_id: str = "user") -> Account:
def __init__(self, user_id: str) -> None: account = Account(name="User", email=f"{user_id}@example.com")
self.id = user_id account.id = user_id
self.session_id = f"session-{user_id}" return account
def _app(*, app_model_config_id: str | None = None) -> App:
return App(
id="app1",
tenant_id="tenant",
name="Agent App",
description="",
mode=AppMode.AGENT,
app_model_config_id=app_model_config_id,
status=AppStatus.NORMAL,
enable_site=False,
enable_api=False,
api_rpm=0,
api_rph=0,
)
def _agent(*, agent_id: str = "agent1") -> Agent:
return Agent(
id=agent_id,
tenant_id="tenant",
name="Agent",
scope=AgentScope.ROSTER,
source=AgentSource.AGENT_APP,
status=AgentStatus.ACTIVE,
app_id="app1",
)
def _snapshot(*, snapshot_id: str = "snap1", agent_id: str = "agent1") -> AgentConfigSnapshot:
return AgentConfigSnapshot(
id=snapshot_id,
tenant_id="tenant",
agent_id=agent_id,
version=1,
config_snapshot=AgentSoulConfig(),
home_snapshot_id="home-1",
)
def _conversation(*, invoke_from: InvokeFrom = InvokeFrom.WEB_APP) -> Conversation:
conversation = Conversation(
id="conv",
app_id="app1",
mode=AppMode.AGENT,
name="Conversation",
invoke_from=invoke_from,
from_source=ConversationFromSource.CONSOLE,
from_account_id="user",
is_deleted=False,
)
conversation._inputs = {}
return conversation
def _message(*, query: str = "query") -> Message:
message = Message(
id="msg",
app_id="app1",
conversation_id="conv",
query=query,
message={"role": "user", "content": query},
answer="",
message_unit_price=Decimal(0),
answer_unit_price=Decimal(0),
currency="USD",
from_source=ConversationFromSource.CONSOLE,
from_account_id="user",
)
message._inputs = {}
return message
_CURRENT_SESSION: Session | None = None
@pytest.fixture(autouse=True)
def _bind_real_session(sqlite_session: Session, monkeypatch: pytest.MonkeyPatch):
global _CURRENT_SESSION
_CURRENT_SESSION = sqlite_session
monkeypatch.setattr(module.db, "session", sqlite_session)
yield
_CURRENT_SESSION = None
def _session() -> Session:
assert _CURRENT_SESSION is not None
return _CURRENT_SESSION
@pytest.fixture @pytest.fixture
@ -46,35 +143,35 @@ def generator(mocker: MockerFixture) -> AgentAppGenerator:
class TestGenerateGuards: class TestGenerateGuards:
def test_rejects_blocking_mode(self, generator: AgentAppGenerator, mocker: MockerFixture): def test_rejects_blocking_mode(self, generator: AgentAppGenerator):
with pytest.raises(AgentAppGeneratorError, match="only supports streaming"): with pytest.raises(AgentAppGeneratorError, match="only supports streaming"):
generator.generate( generator.generate(
app_model=mocker.MagicMock(), app_model=_app(),
user=DummyAccount("u"), user=_account("u"),
args={}, args={},
invoke_from=InvokeFrom.WEB_APP, invoke_from=InvokeFrom.WEB_APP,
session=mocker.MagicMock(), session=_session(),
streaming=False, streaming=False,
) )
def test_requires_query(self, generator: AgentAppGenerator, mocker: MockerFixture): def test_requires_query(self, generator: AgentAppGenerator):
with pytest.raises(AgentAppGeneratorError, match="query is required"): with pytest.raises(AgentAppGeneratorError, match="query is required"):
generator.generate( generator.generate(
app_model=mocker.MagicMock(), app_model=_app(),
user=DummyAccount("u"), user=_account("u"),
args={"inputs": {}}, args={"inputs": {}},
invoke_from=InvokeFrom.WEB_APP, invoke_from=InvokeFrom.WEB_APP,
session=mocker.MagicMock(), session=_session(),
) )
def test_rejects_blank_query(self, generator: AgentAppGenerator, mocker: MockerFixture): def test_rejects_blank_query(self, generator: AgentAppGenerator):
with pytest.raises(AgentAppGeneratorError, match="query is required"): with pytest.raises(AgentAppGeneratorError, match="query is required"):
generator.generate( generator.generate(
app_model=mocker.MagicMock(), app_model=_app(),
user=DummyAccount("u"), user=_account("u"),
args={"query": " ", "inputs": {}}, args={"query": " ", "inputs": {}},
invoke_from=InvokeFrom.WEB_APP, invoke_from=InvokeFrom.WEB_APP,
session=mocker.MagicMock(), session=_session(),
) )
@ -94,18 +191,17 @@ class TestGenerateSuccess:
) )
def test_generate_orchestrates_and_starts_worker(self, generator, mocker: MockerFixture): def test_generate_orchestrates_and_starts_worker(self, generator, mocker: MockerFixture):
app_model = mocker.MagicMock(id="app1", tenant_id="tenant", mode="agent") config = AppModelConfig(app_id="app1")
app_model.app_model_config_id = "config-1" config.id = "config-1"
user = DummyAccount("user") session = _session()
session = mocker.MagicMock() session.add(config)
session.commit()
app_model = _app(app_model_config_id=config.id)
user = _account()
generator._resolve_agent = mocker.MagicMock( generator._resolve_agent = mocker.MagicMock(return_value=(_agent(), "snap1", "snapshot", AgentSoulConfig()))
return_value=(mocker.MagicMock(id="agent1"), "snap1", "snapshot", mocker.MagicMock())
)
generator._prepare_user_inputs = mocker.MagicMock(return_value={"x": 1}) generator._prepare_user_inputs = mocker.MagicMock(return_value={"x": 1})
generator._init_generate_records = mocker.MagicMock( generator._init_generate_records = mocker.MagicMock(return_value=(_conversation(), _message()))
return_value=(mocker.MagicMock(id="conv", mode="agent"), mocker.MagicMock(id="msg"))
)
generator._handle_response = mocker.MagicMock(return_value="raw-response") generator._handle_response = mocker.MagicMock(return_value="raw-response")
mocker.patch( mocker.patch(
@ -151,23 +247,19 @@ class TestGenerateSuccess:
session=session, session=session,
conversation=None, conversation=None,
) )
session.get.assert_called_once_with(AppModelConfig, "config-1") assert session.get(AppModelConfig, "config-1") is config
assert generate_entity.call_args.kwargs["prompt_file_mappings"] == file_mappings assert generate_entity.call_args.kwargs["prompt_file_mappings"] == file_mappings
assert "agent_runtime_exit_intent" not in generate_entity.call_args.kwargs assert "agent_runtime_exit_intent" not in generate_entity.call_args.kwargs
def test_generate_loads_existing_conversation(self, generator: AgentAppGenerator, mocker: MockerFixture): def test_generate_loads_existing_conversation(self, generator: AgentAppGenerator, mocker: MockerFixture):
app_model = mocker.MagicMock(id="app1", tenant_id="tenant", mode="agent") app_model = _app()
generator._resolve_agent = mocker.MagicMock( generator._resolve_agent = mocker.MagicMock(
return_value=(mocker.MagicMock(id="a"), "snap1", "snapshot", mocker.MagicMock()) return_value=(_agent(agent_id="a"), "snap1", "snapshot", AgentSoulConfig())
) )
generator._prepare_user_inputs = mocker.MagicMock(return_value={}) generator._prepare_user_inputs = mocker.MagicMock(return_value={})
generator._init_generate_records = mocker.MagicMock( generator._init_generate_records = mocker.MagicMock(return_value=(_conversation(), _message()))
return_value=(mocker.MagicMock(id="conv", mode="agent"), mocker.MagicMock(id="msg"))
)
generator._handle_response = mocker.MagicMock(return_value="raw") generator._handle_response = mocker.MagicMock(return_value="raw")
get_conv = mocker.patch( get_conv = mocker.patch(f"{MODULE}.ConversationService.get_conversation", return_value=_conversation())
f"{MODULE}.ConversationService.get_conversation", return_value=mocker.MagicMock(id="conv")
)
mocker.patch(f"{MODULE}.AgentAppConfigManager.get_app_config", return_value=mocker.MagicMock(variables=[])) mocker.patch(f"{MODULE}.AgentAppConfigManager.get_app_config", return_value=mocker.MagicMock(variables=[]))
mocker.patch(f"{MODULE}.load_annotation_reply_config", return_value={"enabled": False}) mocker.patch(f"{MODULE}.load_annotation_reply_config", return_value={"enabled": False})
mocker.patch(f"{MODULE}.ModelConfigConverter.convert", return_value=mocker.MagicMock()) mocker.patch(f"{MODULE}.ModelConfigConverter.convert", return_value=mocker.MagicMock())
@ -176,8 +268,8 @@ class TestGenerateSuccess:
mocker.patch(f"{MODULE}.MessageBasedAppQueueManager", return_value=mocker.MagicMock()) mocker.patch(f"{MODULE}.MessageBasedAppQueueManager", return_value=mocker.MagicMock())
mocker.patch(f"{MODULE}.threading.Thread", return_value=mocker.MagicMock()) mocker.patch(f"{MODULE}.threading.Thread", return_value=mocker.MagicMock())
mocker.patch(f"{MODULE}.AgentAppGenerateResponseConverter.convert", return_value={"result": "ok"}) mocker.patch(f"{MODULE}.AgentAppGenerateResponseConverter.convert", return_value={"result": "ok"})
session = mocker.MagicMock() session = _session()
user = DummyAccount("user") user = _account()
generator.generate( generator.generate(
app_model=app_model, app_model=app_model,
@ -200,16 +292,12 @@ class TestGenerateSuccess:
def test_generate_does_not_include_trace_session_id_in_extras( def test_generate_does_not_include_trace_session_id_in_extras(
self, generator: AgentAppGenerator, mocker: MockerFixture self, generator: AgentAppGenerator, mocker: MockerFixture
): ):
app_model = mocker.MagicMock(id="app1", tenant_id="tenant", mode="agent") app_model = _app()
user = DummyAccount("user") user = _account()
generator._resolve_agent = mocker.MagicMock( generator._resolve_agent = mocker.MagicMock(return_value=(_agent(), "snap1", "snapshot", AgentSoulConfig()))
return_value=(mocker.MagicMock(id="agent1"), "snap1", "snapshot", mocker.MagicMock())
)
generator._prepare_user_inputs = mocker.MagicMock(return_value={}) generator._prepare_user_inputs = mocker.MagicMock(return_value={})
generator._init_generate_records = mocker.MagicMock( generator._init_generate_records = mocker.MagicMock(return_value=(_conversation(), _message()))
return_value=(mocker.MagicMock(id="conv", mode="agent"), mocker.MagicMock(id="msg"))
)
generator._handle_response = mocker.MagicMock(return_value="raw-response") generator._handle_response = mocker.MagicMock(return_value="raw-response")
mocker.patch( mocker.patch(
@ -230,7 +318,7 @@ class TestGenerateSuccess:
user=user, user=user,
args={"query": "hello", "inputs": {}, "trace_session_id": "session-1"}, args={"query": "hello", "inputs": {}, "trace_session_id": "session-1"},
invoke_from=InvokeFrom.WEB_APP, invoke_from=InvokeFrom.WEB_APP,
session=mocker.MagicMock(), session=_session(),
streaming=True, streaming=True,
) )
@ -255,25 +343,22 @@ class TestGenerateWorker:
handled=False, handled=False,
guard_query="query", guard_query="query",
): ):
generator._get_conversation = mocker.MagicMock(return_value=mocker.MagicMock(id="conv")) generator._get_conversation = mocker.MagicMock(return_value=_conversation())
generator._get_message = mocker.MagicMock(return_value=mocker.MagicMock(id="msg")) generator._get_message = mocker.MagicMock(return_value=_message())
generator._run_input_guards = mocker.MagicMock(return_value=(handled, guard_query, None)) generator._run_input_guards = mocker.MagicMock(return_value=(handled, guard_query, None))
resolved_agent = mocker.MagicMock(id="a") resolved_agent = _agent(agent_id="a")
resolved_config = mocker.MagicMock(id="s", home_snapshot_id="home-1") resolved_config = _snapshot(snapshot_id="s", agent_id="a")
generator._resolve_agent_by_id = mocker.MagicMock( resolver_sessions: list[Session] = []
return_value=(resolved_agent, resolved_config, mocker.MagicMock())
) def resolve_agent_by_id(**kwargs):
session = mocker.MagicMock() resolver_sessions.append(kwargs["session"])
session.get.return_value = mocker.MagicMock(id="app1") return resolved_agent, resolved_config, AgentSoulConfig()
session_context = mocker.MagicMock()
session_context.__enter__.return_value = session generator._resolve_agent_by_id = mocker.MagicMock(side_effect=resolve_agent_by_id)
session_maker = mocker.patch(f"{MODULE}.session_factory.get_session_maker").return_value session = _session()
session_maker.begin.return_value = session_context if session.get(App, "app1") is None:
resolver_session = mocker.MagicMock() session.add(_app())
resolver_context = mocker.MagicMock() session.commit()
resolver_context.__enter__.return_value = resolver_session
mocker.patch(f"{MODULE}.session_factory.create_session", return_value=resolver_context)
mocker.patch(f"{MODULE}.db.session.close")
mocker.patch(f"{MODULE}.DifyRunContext", return_value=mocker.MagicMock()) mocker.patch(f"{MODULE}.DifyRunContext", return_value=mocker.MagicMock())
mocker.patch(f"{MODULE}.AgentAppRuntimeRequestBuilder", return_value=mocker.MagicMock()) mocker.patch(f"{MODULE}.AgentAppRuntimeRequestBuilder", return_value=mocker.MagicMock())
mocker.patch(f"{MODULE}.create_agent_backend_run_client", return_value=mocker.MagicMock()) mocker.patch(f"{MODULE}.create_agent_backend_run_client", return_value=mocker.MagicMock())
@ -283,7 +368,7 @@ class TestGenerateWorker:
if run_side_effect is not None: if run_side_effect is not None:
runner.run.side_effect = run_side_effect runner.run.side_effect = run_side_effect
mocker.patch(f"{MODULE}.AgentAppRunner", return_value=runner) mocker.patch(f"{MODULE}.AgentAppRunner", return_value=runner)
return runner, resolver_session return runner, resolver_sessions
def _call( def _call(
self, self,
@ -300,6 +385,7 @@ class TestGenerateWorker:
flask_app=mocker.MagicMock(), flask_app=mocker.MagicMock(),
context=mocker.MagicMock(), context=mocker.MagicMock(),
application_generate_entity=mocker.MagicMock( application_generate_entity=mocker.MagicMock(
app_config=SimpleNamespace(app_id="app1", tenant_id="tenant"),
agent_id="a", agent_id="a",
agent_config_snapshot_id="s", agent_config_snapshot_id="s",
agent_session_scope_config_version_id=session_scope_config_version_id, agent_session_scope_config_version_id=session_scope_config_version_id,
@ -315,11 +401,12 @@ class TestGenerateWorker:
) )
def test_happy_path_runs_backend(self, generator: AgentAppGenerator, mocker: MockerFixture): def test_happy_path_runs_backend(self, generator: AgentAppGenerator, mocker: MockerFixture):
runner, resolver_session = self._wire(generator, mocker) runner, resolver_sessions = self._wire(generator, mocker)
queue_manager = mocker.MagicMock() queue_manager = mocker.MagicMock()
self._call(generator, mocker, queue_manager) self._call(generator, mocker, queue_manager)
runner.run.assert_called_once() runner.run.assert_called_once()
assert generator._resolve_agent_by_id.call_args.kwargs["session"] is resolver_session assert resolver_sessions == [generator._resolve_agent_by_id.call_args.kwargs["session"]]
assert resolver_sessions[0].get_bind() is not None
assert runner.run.call_args.kwargs["home_snapshot_id"] == "home-1" assert runner.run.call_args.kwargs["home_snapshot_id"] == "home-1"
assert "home_snapshot_ref" not in runner.run.call_args.kwargs assert "home_snapshot_ref" not in runner.run.call_args.kwargs
queue_manager.publish_error.assert_not_called() queue_manager.publish_error.assert_not_called()
@ -384,20 +471,30 @@ class TestGenerateWorker:
def test_annotation_reply_publishes_after_guard_transaction_commits(self, generator, mocker: MockerFixture): def test_annotation_reply_publishes_after_guard_transaction_commits(self, generator, mocker: MockerFixture):
runner, _ = self._wire(generator, mocker, handled=True) runner, _ = self._wire(generator, mocker, handled=True)
annotation_reply = mocker.MagicMock(id="annotation-1", content="annotated answer") annotation_reply = MessageAnnotation(
app_id="app1",
question="query",
content="annotated answer",
account_id="user",
)
generator._run_input_guards.return_value = (True, "query", annotation_reply) generator._run_input_guards.return_value = (True, "query", annotation_reply)
events: list[str] = [] events: list[str] = []
guard_context = module.session_factory.get_session_maker.return_value.begin.return_value
guard_context.__exit__.side_effect = lambda *args: events.append("commit") or False
queue_manager = mocker.MagicMock() queue_manager = mocker.MagicMock()
def record_commit(_session: Session) -> None:
events.append("commit")
def publish(event, *_args): def publish(event, *_args):
if isinstance(event, QueueAnnotationReplyEvent): if isinstance(event, QueueAnnotationReplyEvent):
events.append("publish") events.append("publish")
queue_manager.publish.side_effect = publish queue_manager.publish.side_effect = publish
self._call(generator, mocker, queue_manager) event.listen(type(_session()), "after_commit", record_commit)
try:
self._call(generator, mocker, queue_manager)
finally:
event.remove(type(_session()), "after_commit", record_commit)
assert events == ["commit", "publish"] assert events == ["commit", "publish"]
runner.run.assert_not_called() runner.run.assert_not_called()
@ -451,16 +548,12 @@ class TestResumeAfterFormSubmission:
composition's user-prompt layer matches the suspended snapshot (never blank).""" composition's user-prompt layer matches the suspended snapshot (never blank)."""
def _wire(self, generator, mocker: MockerFixture): def _wire(self, generator, mocker: MockerFixture):
generator._resolve_agent = mocker.MagicMock( generator._resolve_agent = mocker.MagicMock(return_value=(_agent(), "snap1", "draft", AgentSoulConfig()))
return_value=(mocker.MagicMock(id="agent1"), "snap1", "draft", mocker.MagicMock()) generator._init_generate_records = mocker.MagicMock(return_value=(_conversation(), _message()))
)
generator._init_generate_records = mocker.MagicMock(
return_value=(mocker.MagicMock(id="conv", mode="agent"), mocker.MagicMock(id="msg"))
)
generator._handle_response = mocker.MagicMock(return_value=None) generator._handle_response = mocker.MagicMock(return_value=None)
get_conversation = mocker.patch( get_conversation = mocker.patch(
f"{MODULE}.ConversationService.get_conversation", f"{MODULE}.ConversationService.get_conversation",
return_value=mocker.MagicMock(id="conv", invoke_from=InvokeFrom.WEB_APP), return_value=_conversation(),
) )
mocker.patch(f"{MODULE}.AgentAppConfigManager.get_app_config", return_value=mocker.MagicMock(variables=[])) mocker.patch(f"{MODULE}.AgentAppConfigManager.get_app_config", return_value=mocker.MagicMock(variables=[]))
mocker.patch(f"{MODULE}.load_annotation_reply_config", return_value={"enabled": False}) mocker.patch(f"{MODULE}.load_annotation_reply_config", return_value={"enabled": False})
@ -478,12 +571,13 @@ class TestResumeAfterFormSubmission:
def test_resume_resends_paused_turn_query(self, generator, mocker: MockerFixture): def test_resume_resends_paused_turn_query(self, generator, mocker: MockerFixture):
entity, get_conversation = self._wire(generator, mocker) entity, get_conversation = self._wire(generator, mocker)
app_model = mocker.MagicMock(id="app1", tenant_id="tenant", mode="agent") session = _session()
app_model.app_model_config_id = "config-1" config = AppModelConfig(app_id="app1")
user = DummyAccount("user") config.id = "config-1"
session = mocker.MagicMock() session.add_all([config, _conversation(), _message(query="original question")])
session.get.return_value = mocker.MagicMock() session.commit()
session.scalar.return_value = mocker.MagicMock(query="original question") app_model = _app(app_model_config_id=config.id)
user = _account()
generator.resume_after_form_submission( generator.resume_after_form_submission(
app_model=app_model, app_model=app_model,
@ -504,17 +598,16 @@ class TestResumeAfterFormSubmission:
session=session, session=session,
) )
assert generator._init_generate_records.call_args.kwargs["session"] is session assert generator._init_generate_records.call_args.kwargs["session"] is session
session.get.assert_called_once_with(AppModelConfig, "config-1") assert session.get(AppModelConfig, "config-1") is config
assert generator._resolve_agent.call_args.kwargs["session"] is session assert generator._resolve_agent.call_args.kwargs["session"] is session
def test_resume_falls_back_to_placeholder_when_no_paused_message(self, generator, mocker: MockerFixture): def test_resume_falls_back_to_placeholder_when_no_paused_message(self, generator, mocker: MockerFixture):
entity, _ = self._wire(generator, mocker) entity, _ = self._wire(generator, mocker)
session = mocker.MagicMock() session = _session()
session.scalar.return_value = None
generator.resume_after_form_submission( generator.resume_after_form_submission(
app_model=mocker.MagicMock(id="app1", tenant_id="tenant", mode="agent"), app_model=_app(),
user=DummyAccount("user"), user=_account(),
conversation_id="conv", conversation_id="conv",
form_id="form-1", form_id="form-1",
invoke_from=InvokeFrom.WEB_APP, invoke_from=InvokeFrom.WEB_APP,
@ -526,15 +619,17 @@ class TestResumeAfterFormSubmission:
def test_resume_uses_build_draft_for_debugger_conversation(self, generator, mocker: MockerFixture): def test_resume_uses_build_draft_for_debugger_conversation(self, generator, mocker: MockerFixture):
self._wire(generator, mocker) self._wire(generator, mocker)
conversation = mocker.MagicMock(id="conv", invoke_from=InvokeFrom.DEBUGGER) conversation = _conversation(invoke_from=InvokeFrom.DEBUGGER)
mocker.patch(f"{MODULE}.ConversationService.get_conversation", return_value=conversation) mocker.patch(f"{MODULE}.ConversationService.get_conversation", return_value=conversation)
generator._resolve_resume_draft.return_value = ("debug_build", "draft-build-1") generator._resolve_resume_draft.return_value = ("debug_build", "draft-build-1")
account_user = Account(name="Test Account", email="test@example.com") account_user = Account(name="Test Account", email="test@example.com")
account_user.id = "user" account_user.id = "user"
app_model = mocker.MagicMock(id="app1", tenant_id="tenant", mode="agent") session = _session()
app_model.app_model_config_id = "config-1" config = AppModelConfig(app_id="app1")
session = mocker.MagicMock() config.id = "config-1"
session.scalar.return_value = mocker.MagicMock(query="original question") session.add_all([config, conversation, _message(query="original question")])
session.commit()
app_model = _app(app_model_config_id=config.id)
generator.resume_after_form_submission( generator.resume_after_form_submission(
app_model=app_model, app_model=app_model,

View File

@ -3,7 +3,7 @@ from datetime import datetime
import pytest import pytest
from pytest_mock import MockerFixture from pytest_mock import MockerFixture
from sqlalchemy import event from sqlalchemy import event
from sqlalchemy.orm import Session, sessionmaker from sqlalchemy.orm import Session
from core.agent.entities import AgentEntity from core.agent.entities import AgentEntity
from core.app.apps.agent_chat.app_runner import AgentChatAppRunner from core.app.apps.agent_chat.app_runner import AgentChatAppRunner
@ -12,7 +12,7 @@ from core.moderation.base import ModerationError
from graphon.model_runtime.entities.llm_entities import LLMMode from graphon.model_runtime.entities.llm_entities import LLMMode
from graphon.model_runtime.entities.model_entities import ModelFeature, ModelPropertyKey from graphon.model_runtime.entities.model_entities import ModelFeature, ModelPropertyKey
from models.enums import ConversationFromSource from models.enums import ConversationFromSource
from models.model import App, AppMode, Conversation, Message from models.model import App, AppMode, Conversation, Message, MessageAnnotation
@pytest.fixture @pytest.fixture
@ -74,9 +74,12 @@ def runner(sqlite_session: Session):
return AgentChatAppRunner() return AgentChatAppRunner()
@pytest.fixture(autouse=True) def _records(session: Session) -> tuple[Conversation, Message]:
def _patch_create_session(mocker: MockerFixture, sqlite_session_factory: sessionmaker[Session]) -> None: conversation = session.get(Conversation, "conv")
mocker.patch("core.app.apps.agent_chat.app_runner.create_session", side_effect=sqlite_session_factory) message = session.get(Message, "msg")
assert conversation is not None
assert message is not None
return conversation, message
class TestAgentChatAppRunnerRun: class TestAgentChatAppRunnerRun:
@ -88,14 +91,14 @@ class TestAgentChatAppRunnerRun:
assert app is not None assert app is not None
sqlite_session.delete(app) sqlite_session.delete(app)
sqlite_session.commit() sqlite_session.commit()
conversation, message = _records(sqlite_session)
with pytest.raises(ValueError): with pytest.raises(ValueError):
runner.run(generate_entity, mocker.MagicMock(), mocker.MagicMock(), mocker.MagicMock(), sqlite_session) runner.run(generate_entity, mocker.MagicMock(), conversation, message, sqlite_session)
def test_run_moderation_error_direct_output( def test_run_moderation_error_direct_output(
self, runner: AgentChatAppRunner, mocker: MockerFixture, sqlite_session: Session self, runner: AgentChatAppRunner, mocker: MockerFixture, sqlite_session: Session
): ):
app_record = mocker.MagicMock(id="app1", tenant_id="tenant")
app_config = mocker.MagicMock(app_id="app1", tenant_id="tenant", prompt_template=mocker.MagicMock()) app_config = mocker.MagicMock(app_id="app1", tenant_id="tenant", prompt_template=mocker.MagicMock())
app_config.agent = mocker.MagicMock() app_config.agent = mocker.MagicMock()
generate_entity = mocker.MagicMock( generate_entity = mocker.MagicMock(
@ -111,15 +114,15 @@ class TestAgentChatAppRunnerRun:
mocker.patch.object(runner, "organize_prompt_messages", return_value=([], None)) mocker.patch.object(runner, "organize_prompt_messages", return_value=([], None))
mocker.patch.object(runner, "moderation_for_inputs", side_effect=ModerationError("bad")) mocker.patch.object(runner, "moderation_for_inputs", side_effect=ModerationError("bad"))
mocker.patch.object(runner, "direct_output") mocker.patch.object(runner, "direct_output")
conversation, message = _records(sqlite_session)
runner.run(generate_entity, mocker.MagicMock(), mocker.MagicMock(), mocker.MagicMock(), sqlite_session) runner.run(generate_entity, mocker.MagicMock(), conversation, message, sqlite_session)
runner.direct_output.assert_called_once() runner.direct_output.assert_called_once()
def test_run_annotation_reply_short_circuits( def test_run_annotation_reply_short_circuits(
self, runner: AgentChatAppRunner, mocker: MockerFixture, sqlite_session: Session self, runner: AgentChatAppRunner, mocker: MockerFixture, sqlite_session: Session
): ):
app_record = mocker.MagicMock(id="app1", tenant_id="tenant")
app_config = mocker.MagicMock(app_id="app1", tenant_id="tenant", prompt_template=mocker.MagicMock()) app_config = mocker.MagicMock(app_id="app1", tenant_id="tenant", prompt_template=mocker.MagicMock())
app_config.agent = mocker.MagicMock() app_config.agent = mocker.MagicMock()
generate_entity = mocker.MagicMock( generate_entity = mocker.MagicMock(
@ -136,13 +139,19 @@ class TestAgentChatAppRunnerRun:
mocker.patch.object(runner, "organize_prompt_messages", return_value=([], None)) mocker.patch.object(runner, "organize_prompt_messages", return_value=([], None))
mocker.patch.object(runner, "moderation_for_inputs", return_value=(None, {}, "q")) mocker.patch.object(runner, "moderation_for_inputs", return_value=(None, {}, "q"))
annotation = mocker.MagicMock(id="anno", content="answer") annotation = MessageAnnotation(
app_id="app1",
question="q",
content="answer",
account_id="user",
)
annotation_query = mocker.patch.object(runner, "query_app_annotations_to_reply", return_value=annotation) annotation_query = mocker.patch.object(runner, "query_app_annotations_to_reply", return_value=annotation)
mocker.patch.object(runner, "direct_output") mocker.patch.object(runner, "direct_output")
queue_manager = mocker.MagicMock() queue_manager = mocker.MagicMock()
write_session = sqlite_session write_session = sqlite_session
runner.run(generate_entity, queue_manager, mocker.MagicMock(), mocker.MagicMock(), write_session) conversation, message = _records(sqlite_session)
runner.run(generate_entity, queue_manager, conversation, message, write_session)
queue_manager.publish.assert_called_once() queue_manager.publish.assert_called_once()
assert annotation_query.call_args.kwargs["session"] is write_session assert annotation_query.call_args.kwargs["session"] is write_session
@ -151,7 +160,6 @@ class TestAgentChatAppRunnerRun:
def test_run_hosting_moderation_short_circuits( def test_run_hosting_moderation_short_circuits(
self, runner: AgentChatAppRunner, mocker: MockerFixture, sqlite_session: Session self, runner: AgentChatAppRunner, mocker: MockerFixture, sqlite_session: Session
): ):
app_record = mocker.MagicMock(id="app1", tenant_id="tenant")
app_config = mocker.MagicMock(app_id="app1", tenant_id="tenant", prompt_template=mocker.MagicMock()) app_config = mocker.MagicMock(app_id="app1", tenant_id="tenant", prompt_template=mocker.MagicMock())
app_config.agent = mocker.MagicMock() app_config.agent = mocker.MagicMock()
generate_entity = mocker.MagicMock( generate_entity = mocker.MagicMock(
@ -170,11 +178,11 @@ class TestAgentChatAppRunnerRun:
mocker.patch.object(runner, "moderation_for_inputs", return_value=(None, {}, "q")) mocker.patch.object(runner, "moderation_for_inputs", return_value=(None, {}, "q"))
mocker.patch.object(runner, "query_app_annotations_to_reply", return_value=None) mocker.patch.object(runner, "query_app_annotations_to_reply", return_value=None)
mocker.patch.object(runner, "check_hosting_moderation", return_value=True) mocker.patch.object(runner, "check_hosting_moderation", return_value=True)
conversation, message = _records(sqlite_session)
runner.run(generate_entity, mocker.MagicMock(), mocker.MagicMock(), mocker.MagicMock(), sqlite_session) runner.run(generate_entity, mocker.MagicMock(), conversation, message, sqlite_session)
def test_run_model_schema_missing(self, runner: AgentChatAppRunner, mocker: MockerFixture, sqlite_session: Session): def test_run_model_schema_missing(self, runner: AgentChatAppRunner, mocker: MockerFixture, sqlite_session: Session):
app_record = mocker.MagicMock(id="app1", tenant_id="tenant")
app_config = mocker.MagicMock(app_id="app1", tenant_id="tenant", prompt_template=mocker.MagicMock()) app_config = mocker.MagicMock(app_id="app1", tenant_id="tenant", prompt_template=mocker.MagicMock())
app_config.agent = AgentEntity(provider="p", model="m", strategy=AgentEntity.Strategy.CHAIN_OF_THOUGHT) app_config.agent = AgentEntity(provider="p", model="m", strategy=AgentEntity.Strategy.CHAIN_OF_THOUGHT)
@ -203,9 +211,10 @@ class TestAgentChatAppRunnerRun:
llm_instance = mocker.MagicMock() llm_instance = mocker.MagicMock()
llm_instance.model_type_instance.get_model_schema.return_value = None llm_instance.model_type_instance.get_model_schema.return_value = None
mocker.patch("core.app.apps.agent_chat.app_runner.ModelInstance", return_value=llm_instance) mocker.patch("core.app.apps.agent_chat.app_runner.ModelInstance", return_value=llm_instance)
conversation, message = _records(sqlite_session)
with pytest.raises(ValueError): with pytest.raises(ValueError):
runner.run(generate_entity, mocker.MagicMock(), mocker.MagicMock(), mocker.MagicMock(), sqlite_session) runner.run(generate_entity, mocker.MagicMock(), conversation, message, sqlite_session)
@pytest.mark.parametrize( @pytest.mark.parametrize(
("mode", "expected_runner"), ("mode", "expected_runner"),
@ -222,7 +231,6 @@ class TestAgentChatAppRunnerRun:
expected_runner, expected_runner,
sqlite_session: Session, sqlite_session: Session,
): ):
app_record = mocker.MagicMock(id="app1", tenant_id="tenant")
app_config = mocker.MagicMock(app_id="app1", tenant_id="tenant", prompt_template=mocker.MagicMock()) app_config = mocker.MagicMock(app_id="app1", tenant_id="tenant", prompt_template=mocker.MagicMock())
app_config.agent = AgentEntity(provider="p", model="m", strategy=AgentEntity.Strategy.CHAIN_OF_THOUGHT) app_config.agent = AgentEntity(provider="p", model="m", strategy=AgentEntity.Strategy.CHAIN_OF_THOUGHT)
@ -256,8 +264,7 @@ class TestAgentChatAppRunnerRun:
llm_instance.model_type_instance.get_model_schema.return_value = model_schema llm_instance.model_type_instance.get_model_schema.return_value = model_schema
mocker.patch("core.app.apps.agent_chat.app_runner.ModelInstance", return_value=llm_instance) mocker.patch("core.app.apps.agent_chat.app_runner.ModelInstance", return_value=llm_instance)
conversation = mocker.MagicMock(id="conv") conversation, message = _records(sqlite_session)
message = mocker.MagicMock(id="msg")
runner_cls = mocker.MagicMock() runner_cls = mocker.MagicMock()
mocker.patch(f"core.app.apps.agent_chat.app_runner.{expected_runner}", runner_cls) mocker.patch(f"core.app.apps.agent_chat.app_runner.{expected_runner}", runner_cls)
@ -268,25 +275,23 @@ class TestAgentChatAppRunnerRun:
runner_instance.run.side_effect = lambda **_kwargs: events.append("agent-run") or [] runner_instance.run.side_effect = lambda **_kwargs: events.append("agent-run") or []
mocker.patch.object(runner, "_handle_invoke_result") mocker.patch.object(runner, "_handle_invoke_result")
session = sqlite_session session = sqlite_session
event.listen(session, "after_commit", lambda _session: events.append("commit"))
original_close = session.close
def close_session() -> None: def record_commit(_session: Session) -> None:
events.append("close") events.append("commit")
original_close()
session.close = close_session event.listen(session, "after_commit", record_commit)
try:
runner.run(generate_entity, mocker.MagicMock(), conversation, message, session)
finally:
event.remove(session, "after_commit", record_commit)
runner.run(generate_entity, mocker.MagicMock(), conversation, message, session) assert events == ["commit", "commit", "agent-run"]
assert events == ["commit", "close", "commit", "close", "agent-run"]
runner_instance.run.assert_called_once() runner_instance.run.assert_called_once()
runner._handle_invoke_result.assert_called_once() runner._handle_invoke_result.assert_called_once()
def test_run_invalid_llm_mode_raises( def test_run_invalid_llm_mode_raises(
self, runner: AgentChatAppRunner, mocker: MockerFixture, sqlite_session: Session self, runner: AgentChatAppRunner, mocker: MockerFixture, sqlite_session: Session
): ):
app_record = mocker.MagicMock(id="app1", tenant_id="tenant")
app_config = mocker.MagicMock(app_id="app1", tenant_id="tenant", prompt_template=mocker.MagicMock()) app_config = mocker.MagicMock(app_id="app1", tenant_id="tenant", prompt_template=mocker.MagicMock())
app_config.agent = AgentEntity(provider="p", model="m", strategy=AgentEntity.Strategy.CHAIN_OF_THOUGHT) app_config.agent = AgentEntity(provider="p", model="m", strategy=AgentEntity.Strategy.CHAIN_OF_THOUGHT)
@ -320,8 +325,7 @@ class TestAgentChatAppRunnerRun:
llm_instance.model_type_instance.get_model_schema.return_value = model_schema llm_instance.model_type_instance.get_model_schema.return_value = model_schema
mocker.patch("core.app.apps.agent_chat.app_runner.ModelInstance", return_value=llm_instance) mocker.patch("core.app.apps.agent_chat.app_runner.ModelInstance", return_value=llm_instance)
conversation = mocker.MagicMock(id="conv") conversation, message = _records(sqlite_session)
message = mocker.MagicMock(id="msg")
with pytest.raises(ValueError): with pytest.raises(ValueError):
runner.run(generate_entity, mocker.MagicMock(), conversation, message, sqlite_session) runner.run(generate_entity, mocker.MagicMock(), conversation, message, sqlite_session)
@ -329,7 +333,6 @@ class TestAgentChatAppRunnerRun:
def test_run_function_calling_strategy_selected_by_features( def test_run_function_calling_strategy_selected_by_features(
self, runner: AgentChatAppRunner, mocker: MockerFixture, sqlite_session: Session self, runner: AgentChatAppRunner, mocker: MockerFixture, sqlite_session: Session
): ):
app_record = mocker.MagicMock(id="app1", tenant_id="tenant")
app_config = mocker.MagicMock(app_id="app1", tenant_id="tenant", prompt_template=mocker.MagicMock()) app_config = mocker.MagicMock(app_id="app1", tenant_id="tenant", prompt_template=mocker.MagicMock())
app_config.agent = AgentEntity(provider="p", model="m", strategy=AgentEntity.Strategy.CHAIN_OF_THOUGHT) app_config.agent = AgentEntity(provider="p", model="m", strategy=AgentEntity.Strategy.CHAIN_OF_THOUGHT)
@ -363,8 +366,7 @@ class TestAgentChatAppRunnerRun:
llm_instance.model_type_instance.get_model_schema.return_value = model_schema llm_instance.model_type_instance.get_model_schema.return_value = model_schema
mocker.patch("core.app.apps.agent_chat.app_runner.ModelInstance", return_value=llm_instance) mocker.patch("core.app.apps.agent_chat.app_runner.ModelInstance", return_value=llm_instance)
conversation = mocker.MagicMock(id="conv") conversation, message = _records(sqlite_session)
message = mocker.MagicMock(id="msg")
runner_cls = mocker.MagicMock() runner_cls = mocker.MagicMock()
mocker.patch("core.app.apps.agent_chat.app_runner.FunctionCallAgentRunner", runner_cls) mocker.patch("core.app.apps.agent_chat.app_runner.FunctionCallAgentRunner", runner_cls)
@ -382,7 +384,6 @@ class TestAgentChatAppRunnerRun:
def test_run_conversation_not_found( def test_run_conversation_not_found(
self, runner: AgentChatAppRunner, mocker: MockerFixture, sqlite_session: Session self, runner: AgentChatAppRunner, mocker: MockerFixture, sqlite_session: Session
): ):
app_record = mocker.MagicMock(id="app1", tenant_id="tenant")
app_config = mocker.MagicMock(app_id="app1", tenant_id="tenant", prompt_template=mocker.MagicMock()) app_config = mocker.MagicMock(app_id="app1", tenant_id="tenant", prompt_template=mocker.MagicMock())
app_config.agent = AgentEntity(provider="p", model="m", strategy=AgentEntity.Strategy.FUNCTION_CALLING) app_config.agent = AgentEntity(provider="p", model="m", strategy=AgentEntity.Strategy.FUNCTION_CALLING)
@ -403,8 +404,7 @@ class TestAgentChatAppRunnerRun:
user_id="user", user_id="user",
) )
conversation_record = sqlite_session.get(Conversation, "conv") conversation_record, message_record = _records(sqlite_session)
assert conversation_record is not None
sqlite_session.delete(conversation_record) sqlite_session.delete(conversation_record)
sqlite_session.commit() sqlite_session.commit()
mocker.patch.object(runner, "organize_prompt_messages", return_value=([], None)) mocker.patch.object(runner, "organize_prompt_messages", return_value=([], None))
@ -416,13 +416,12 @@ class TestAgentChatAppRunnerRun:
runner.run( runner.run(
generate_entity, generate_entity,
mocker.MagicMock(), mocker.MagicMock(),
mocker.MagicMock(id="conv"), conversation_record,
mocker.MagicMock(id="msg"), message_record,
sqlite_session, sqlite_session,
) )
def test_run_message_not_found(self, runner: AgentChatAppRunner, mocker: MockerFixture, sqlite_session: Session): def test_run_message_not_found(self, runner: AgentChatAppRunner, mocker: MockerFixture, sqlite_session: Session):
app_record = mocker.MagicMock(id="app1", tenant_id="tenant")
app_config = mocker.MagicMock(app_id="app1", tenant_id="tenant", prompt_template=mocker.MagicMock()) app_config = mocker.MagicMock(app_id="app1", tenant_id="tenant", prompt_template=mocker.MagicMock())
app_config.agent = AgentEntity(provider="p", model="m", strategy=AgentEntity.Strategy.FUNCTION_CALLING) app_config.agent = AgentEntity(provider="p", model="m", strategy=AgentEntity.Strategy.FUNCTION_CALLING)
@ -443,8 +442,7 @@ class TestAgentChatAppRunnerRun:
user_id="user", user_id="user",
) )
message_record = sqlite_session.get(Message, "msg") conversation_record, message_record = _records(sqlite_session)
assert message_record is not None
sqlite_session.delete(message_record) sqlite_session.delete(message_record)
sqlite_session.commit() sqlite_session.commit()
mocker.patch.object(runner, "organize_prompt_messages", return_value=([], None)) mocker.patch.object(runner, "organize_prompt_messages", return_value=([], None))
@ -456,15 +454,14 @@ class TestAgentChatAppRunnerRun:
runner.run( runner.run(
generate_entity, generate_entity,
mocker.MagicMock(), mocker.MagicMock(),
mocker.MagicMock(id="conv"), conversation_record,
mocker.MagicMock(id="msg"), message_record,
sqlite_session, sqlite_session,
) )
def test_run_invalid_agent_strategy_raises( def test_run_invalid_agent_strategy_raises(
self, runner: AgentChatAppRunner, mocker: MockerFixture, sqlite_session: Session self, runner: AgentChatAppRunner, mocker: MockerFixture, sqlite_session: Session
): ):
app_record = mocker.MagicMock(id="app1", tenant_id="tenant")
app_config = mocker.MagicMock(app_id="app1", tenant_id="tenant", prompt_template=mocker.MagicMock()) app_config = mocker.MagicMock(app_id="app1", tenant_id="tenant", prompt_template=mocker.MagicMock())
app_config.agent = mocker.MagicMock(strategy="invalid", provider="p", model="m") app_config.agent = mocker.MagicMock(strategy="invalid", provider="p", model="m")
@ -498,8 +495,7 @@ class TestAgentChatAppRunnerRun:
llm_instance.model_type_instance.get_model_schema.return_value = model_schema llm_instance.model_type_instance.get_model_schema.return_value = model_schema
mocker.patch("core.app.apps.agent_chat.app_runner.ModelInstance", return_value=llm_instance) mocker.patch("core.app.apps.agent_chat.app_runner.ModelInstance", return_value=llm_instance)
conversation = mocker.MagicMock(id="conv") conversation, message = _records(sqlite_session)
message = mocker.MagicMock(id="msg")
with pytest.raises(ValueError): with pytest.raises(ValueError):
runner.run(generate_entity, mocker.MagicMock(), conversation, message, sqlite_session) runner.run(generate_entity, mocker.MagicMock(), conversation, message, sqlite_session)

View File

@ -1,10 +1,12 @@
from contextlib import contextmanager from decimal import Decimal
from types import SimpleNamespace from types import SimpleNamespace
from unittest.mock import ANY, MagicMock, Mock, patch from unittest.mock import ANY, MagicMock, Mock, patch
import pytest import pytest
from sqlalchemy.orm import Session, sessionmaker from sqlalchemy import event
from sqlalchemy.orm import Session
import core.app.apps.chat.app_generator as generator_module
from core.app.apps.chat.app_generator import ChatAppGenerator from core.app.apps.chat.app_generator import ChatAppGenerator
from core.app.apps.chat.app_runner import ChatAppRunner from core.app.apps.chat.app_runner import ChatAppRunner
from core.app.apps.exc import GenerateTaskStoppedError from core.app.apps.exc import GenerateTaskStoppedError
@ -12,7 +14,9 @@ from core.app.entities.app_invoke_entities import InvokeFrom
from core.app.entities.queue_entities import QueueAnnotationReplyEvent from core.app.entities.queue_entities import QueueAnnotationReplyEvent
from core.moderation.base import ModerationError from core.moderation.base import ModerationError
from graphon.model_runtime.errors.invoke import InvokeAuthorizationError from graphon.model_runtime.errors.invoke import InvokeAuthorizationError
from models.model import App, AppMode, IconType from models import Account
from models.enums import ConversationFromSource
from models.model import App, AppMode, AppModelConfig, Conversation, IconType, Message, MessageAnnotation
class DummyGenerateEntity: class DummyGenerateEntity:
@ -31,19 +35,8 @@ class DummyQueueManager:
self.published.append((event, pub_from)) self.published.append((event, pub_from))
@contextmanager def _app() -> App:
def patched_create_session(session_factory: sessionmaker[Session]): return App(
@contextmanager
def create_session():
with session_factory() as session:
yield session
with patch("core.app.apps.chat.app_runner.create_session", create_session):
yield
def _persist_app(session: Session) -> App:
app = App(
id="app-1", id="app-1",
tenant_id="tenant-1", tenant_id="tenant-1",
name="Chat app", name="Chat app",
@ -54,9 +47,75 @@ def _persist_app(session: Session) -> App:
enable_site=False, enable_site=False,
enable_api=True, enable_api=True,
) )
session.add(app)
def _account() -> Account:
account = Account(name="User", email="user-1@example.com")
account.id = "user-1"
return account
def _conversation() -> Conversation:
conversation = Conversation(
id="c1",
app_id="app-1",
app_model_config_id=None,
model_provider=None,
override_model_configs=None,
model_id=None,
mode=AppMode.CHAT,
name="Conversation",
inputs={},
introduction="",
system_instruction="",
system_instruction_tokens=0,
status="normal",
invoke_from=InvokeFrom.SERVICE_API,
from_source=ConversationFromSource.API,
from_end_user_id=None,
from_account_id="user-1",
)
return conversation
def _message() -> Message:
return Message(
id="m1",
app_id="app-1",
conversation_id="c1",
inputs={},
query="hi",
message={},
message_tokens=0,
message_unit_price=Decimal(0),
message_price_unit=Decimal(0),
answer="",
answer_tokens=0,
answer_unit_price=Decimal(0),
answer_price_unit=Decimal(0),
provider_response_latency=0,
total_price=Decimal(0),
currency="USD",
invoke_from=InvokeFrom.SERVICE_API,
from_source=ConversationFromSource.API,
from_end_user_id=None,
from_account_id="user-1",
app_mode=AppMode.CHAT,
)
def _persist_records(session: Session) -> tuple[App, Conversation, Message]:
app = _app()
conversation = _conversation()
message = _message()
session.add_all([app, conversation, message])
session.commit() session.commit()
return app return app, conversation, message
@pytest.fixture(autouse=True)
def _bind_db_session(sqlite_session: Session, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(generator_module.db, "session", sqlite_session)
class TestChatAppGenerator: class TestChatAppGenerator:
@ -65,8 +124,8 @@ class TestChatAppGenerator:
with pytest.raises(ValueError): with pytest.raises(ValueError):
generator.generate( generator.generate(
session=unbound_session, session=unbound_session,
app_model=SimpleNamespace(), app_model=_app(),
user=SimpleNamespace(), user=_account(),
args={"inputs": {}}, args={"inputs": {}},
invoke_from=InvokeFrom.SERVICE_API, invoke_from=InvokeFrom.SERVICE_API,
streaming=False, streaming=False,
@ -77,8 +136,8 @@ class TestChatAppGenerator:
with pytest.raises(ValueError): with pytest.raises(ValueError):
generator.generate( generator.generate(
session=unbound_session, session=unbound_session,
app_model=SimpleNamespace(), app_model=_app(),
user=SimpleNamespace(), user=_account(),
args={"query": 1, "inputs": {}}, args={"query": 1, "inputs": {}},
invoke_from=InvokeFrom.SERVICE_API, invoke_from=InvokeFrom.SERVICE_API,
streaming=False, streaming=False,
@ -86,8 +145,8 @@ class TestChatAppGenerator:
def test_generate_debugger_overrides_model_config(self, unbound_session: Session): def test_generate_debugger_overrides_model_config(self, unbound_session: Session):
generator = ChatAppGenerator() generator = ChatAppGenerator()
app_model = SimpleNamespace(id="app-1", tenant_id="tenant-1") app_model = _app()
user = SimpleNamespace(id="user-1", session_id="session-1") user = _account()
args = { args = {
"query": "hi", "query": "hi",
"inputs": {}, "inputs": {},
@ -119,12 +178,12 @@ class TestChatAppGenerator:
patch( patch(
"core.app.apps.chat.app_generator.ChatAppGenerateResponseConverter.convert", return_value={"ok": True} "core.app.apps.chat.app_generator.ChatAppGenerateResponseConverter.convert", return_value={"ok": True}
), ),
patch.object(ChatAppGenerator, "_get_app_model_config", return_value=SimpleNamespace(to_dict=lambda: {})), patch.object(ChatAppGenerator, "_get_app_model_config", return_value=AppModelConfig(app_id="app-1")),
patch.object(ChatAppGenerator, "_prepare_user_inputs", return_value={}), patch.object(ChatAppGenerator, "_prepare_user_inputs", return_value={}),
patch.object( patch.object(
ChatAppGenerator, ChatAppGenerator,
"_init_generate_records", "_init_generate_records",
return_value=(SimpleNamespace(id="c1", mode="chat"), SimpleNamespace(id="m1")), return_value=(_conversation(), _message()),
), ),
patch.object(ChatAppGenerator, "_handle_response", return_value={"response": True}), patch.object(ChatAppGenerator, "_handle_response", return_value={"response": True}),
patch("core.app.apps.chat.app_generator.copy_current_request_context", side_effect=lambda f: f), patch("core.app.apps.chat.app_generator.copy_current_request_context", side_effect=lambda f: f),
@ -141,13 +200,14 @@ class TestChatAppGenerator:
def test_generate_uses_session_for_annotation_reply(self, unbound_session: Session): def test_generate_uses_session_for_annotation_reply(self, unbound_session: Session):
generator = ChatAppGenerator() generator = ChatAppGenerator()
app_model = SimpleNamespace(id="app-1", tenant_id="tenant-1") app_model = _app()
app_model_config = MagicMock(id="config-1", app_id="app-1") app_model_config = AppModelConfig(app_id="app-1")
annotation_reply = {"enabled": False} annotation_reply = {"enabled": False}
user = SimpleNamespace(id="user-1", session_id="session-1") user = _account()
with ( with (
patch.object(ChatAppGenerator, "_get_app_model_config", return_value=app_model_config), patch.object(ChatAppGenerator, "_get_app_model_config", return_value=app_model_config),
patch.object(AppModelConfig, "to_dict", return_value={}) as to_dict,
patch( patch(
"core.app.apps.chat.app_generator.load_annotation_reply_config", "core.app.apps.chat.app_generator.load_annotation_reply_config",
return_value=annotation_reply, return_value=annotation_reply,
@ -168,40 +228,34 @@ class TestChatAppGenerator:
) )
load_annotation_reply_config.assert_called_once_with(unbound_session, "app-1") load_annotation_reply_config.assert_called_once_with(unbound_session, "app-1")
app_model_config.to_dict.assert_called_once_with(annotation_reply=annotation_reply) to_dict.assert_called_once_with(annotation_reply=annotation_reply)
assert get_app_config.call_args.kwargs["annotation_reply"] is annotation_reply assert get_app_config.call_args.kwargs["annotation_reply"] is annotation_reply
def test_generate_rejects_model_config_override_for_non_debugger(self, unbound_session: Session): def test_generate_rejects_model_config_override_for_non_debugger(self, unbound_session: Session):
generator = ChatAppGenerator() generator = ChatAppGenerator()
with pytest.raises(ValueError): with pytest.raises(ValueError):
with ( with (
patch.object( patch.object(ChatAppGenerator, "_get_app_model_config", return_value=AppModelConfig(app_id="app-1")),
ChatAppGenerator, "_get_app_model_config", return_value=SimpleNamespace(to_dict=lambda: {}) patch.object(AppModelConfig, "to_dict", return_value={}),
),
): ):
generator.generate( generator.generate(
session=unbound_session, session=unbound_session,
app_model=SimpleNamespace(tenant_id="t1", id="a1", mode=AppMode.CHAT.value), app_model=_app(),
user=SimpleNamespace(id="u1", session_id="s1"), user=_account(),
args={"query": "hi", "inputs": {}, "model_config": {"foo": "bar"}}, args={"query": "hi", "inputs": {}, "model_config": {"foo": "bar"}},
invoke_from=InvokeFrom.SERVICE_API, invoke_from=InvokeFrom.SERVICE_API,
streaming=False, streaming=False,
) )
def test_generate_worker_handles_exceptions(self, unbound_session_factory: sessionmaker[Session]): def test_generate_worker_handles_exceptions(self):
generator = ChatAppGenerator() generator = ChatAppGenerator()
queue_manager = DummyQueueManager() queue_manager = DummyQueueManager()
entity = DummyGenerateEntity(task_id="t1", user_id="u1") entity = DummyGenerateEntity(task_id="t1", user_id="u1")
with ( with (
patch.object(ChatAppGenerator, "_get_conversation", return_value=SimpleNamespace()), patch.object(ChatAppGenerator, "_get_conversation", return_value=_conversation()),
patch.object(ChatAppGenerator, "_get_message", return_value=SimpleNamespace()), patch.object(ChatAppGenerator, "_get_message", return_value=_message()),
patch("core.app.apps.chat.app_generator.ChatAppRunner.run", side_effect=InvokeAuthorizationError()), patch("core.app.apps.chat.app_generator.ChatAppRunner.run", side_effect=InvokeAuthorizationError()),
patch(
"core.app.apps.chat.app_generator.session_factory",
SimpleNamespace(create_session=unbound_session_factory),
),
patch("core.app.apps.chat.app_generator.db.session.close"),
): ):
generator._generate_worker( generator._generate_worker(
flask_app=Mock(app_context=Mock(return_value=Mock(__enter__=Mock(), __exit__=Mock()))), flask_app=Mock(app_context=Mock(return_value=Mock(__enter__=Mock(), __exit__=Mock()))),
@ -214,14 +268,9 @@ class TestChatAppGenerator:
assert queue_manager.published assert queue_manager.published
with ( with (
patch.object(ChatAppGenerator, "_get_conversation", return_value=SimpleNamespace()), patch.object(ChatAppGenerator, "_get_conversation", return_value=_conversation()),
patch.object(ChatAppGenerator, "_get_message", return_value=SimpleNamespace()), patch.object(ChatAppGenerator, "_get_message", return_value=_message()),
patch("core.app.apps.chat.app_generator.ChatAppRunner.run", side_effect=GenerateTaskStoppedError()), patch("core.app.apps.chat.app_generator.ChatAppRunner.run", side_effect=GenerateTaskStoppedError()),
patch(
"core.app.apps.chat.app_generator.session_factory",
SimpleNamespace(create_session=unbound_session_factory),
),
patch("core.app.apps.chat.app_generator.db.session.close"),
): ):
generator._generate_worker( generator._generate_worker(
flask_app=Mock(app_context=Mock(return_value=Mock(__enter__=Mock(), __exit__=Mock()))), flask_app=Mock(app_context=Mock(return_value=Mock(__enter__=Mock(), __exit__=Mock()))),
@ -233,7 +282,7 @@ class TestChatAppGenerator:
class TestChatAppRunner: class TestChatAppRunner:
def test_run_raises_when_app_missing(self, sqlite_session_factory: sessionmaker[Session], unbound_session: Session): def test_run_raises_when_app_missing(self, sqlite_session: Session):
runner = ChatAppRunner() runner = ChatAppRunner()
app_config = SimpleNamespace( app_config = SimpleNamespace(
app_id="app-1", tenant_id="tenant-1", prompt_template=None, external_data_variables=[] app_id="app-1", tenant_id="tenant-1", prompt_template=None, external_data_variables=[]
@ -251,20 +300,17 @@ class TestChatAppRunner:
invoke_from=InvokeFrom.SERVICE_API, invoke_from=InvokeFrom.SERVICE_API,
) )
with patched_create_session(sqlite_session_factory): with pytest.raises(ValueError):
with pytest.raises(ValueError): runner.run(
runner.run( app_generate_entity,
app_generate_entity, DummyQueueManager(),
DummyQueueManager(), _conversation(),
SimpleNamespace(), _message(),
SimpleNamespace(id="m1"), sqlite_session,
unbound_session, )
)
def test_run_moderation_error_direct_output( def test_run_moderation_error_direct_output(self, sqlite_session: Session):
self, sqlite_session: Session, sqlite_session_factory: sessionmaker[Session] _, conversation, message = _persist_records(sqlite_session)
):
_persist_app(sqlite_session)
runner = ChatAppRunner() runner = ChatAppRunner()
app_config = SimpleNamespace( app_config = SimpleNamespace(
app_id="app-1", app_id="app-1",
@ -289,7 +335,6 @@ class TestChatAppRunner:
) )
with ( with (
patched_create_session(sqlite_session_factory),
patch.object(ChatAppRunner, "organize_prompt_messages", return_value=([], [])), patch.object(ChatAppRunner, "organize_prompt_messages", return_value=([], [])),
patch.object(ChatAppRunner, "moderation_for_inputs", side_effect=ModerationError("blocked")), patch.object(ChatAppRunner, "moderation_for_inputs", side_effect=ModerationError("blocked")),
patch.object(ChatAppRunner, "direct_output") as mock_direct, patch.object(ChatAppRunner, "direct_output") as mock_direct,
@ -297,17 +342,15 @@ class TestChatAppRunner:
runner.run( runner.run(
app_generate_entity, app_generate_entity,
DummyQueueManager(), DummyQueueManager(),
SimpleNamespace(), conversation,
SimpleNamespace(id="m1"), message,
sqlite_session, sqlite_session,
) )
mock_direct.assert_called_once() mock_direct.assert_called_once()
def test_run_annotation_reply_short_circuits( def test_run_annotation_reply_short_circuits(self, sqlite_session: Session):
self, sqlite_session: Session, sqlite_session_factory: sessionmaker[Session] _, conversation, message = _persist_records(sqlite_session)
):
_persist_app(sqlite_session)
runner = ChatAppRunner() runner = ChatAppRunner()
app_config = SimpleNamespace( app_config = SimpleNamespace(
app_id="app-1", app_id="app-1",
@ -331,26 +374,28 @@ class TestChatAppRunner:
invoke_from=InvokeFrom.SERVICE_API, invoke_from=InvokeFrom.SERVICE_API,
) )
annotation = SimpleNamespace(id="ann-1", content="answer") annotation = MessageAnnotation(
app_id="app-1",
question="hi",
content="answer",
account_id="user-1",
)
with ( with (
patched_create_session(sqlite_session_factory),
patch.object(ChatAppRunner, "organize_prompt_messages", return_value=([], [])), patch.object(ChatAppRunner, "organize_prompt_messages", return_value=([], [])),
patch.object(ChatAppRunner, "moderation_for_inputs", return_value=(None, {}, "hi")), patch.object(ChatAppRunner, "moderation_for_inputs", return_value=(None, {}, "hi")),
patch.object(ChatAppRunner, "query_app_annotations_to_reply", return_value=annotation) as annotation_query, patch.object(ChatAppRunner, "query_app_annotations_to_reply", return_value=annotation) as annotation_query,
patch.object(ChatAppRunner, "direct_output") as mock_direct, patch.object(ChatAppRunner, "direct_output") as mock_direct,
): ):
queue_manager = DummyQueueManager() queue_manager = DummyQueueManager()
runner.run(app_generate_entity, queue_manager, SimpleNamespace(), SimpleNamespace(id="m1"), sqlite_session) runner.run(app_generate_entity, queue_manager, conversation, message, sqlite_session)
assert any(isinstance(item[0], QueueAnnotationReplyEvent) for item in queue_manager.published) assert any(isinstance(item[0], QueueAnnotationReplyEvent) for item in queue_manager.published)
assert annotation_query.call_args.kwargs["session"] is sqlite_session assert annotation_query.call_args.kwargs["session"] is sqlite_session
mock_direct.assert_called_once() mock_direct.assert_called_once()
def test_run_returns_when_hosting_moderation_blocks( def test_run_returns_when_hosting_moderation_blocks(self, sqlite_session: Session):
self, sqlite_session: Session, sqlite_session_factory: sessionmaker[Session] _, conversation, message = _persist_records(sqlite_session)
):
_persist_app(sqlite_session)
runner = ChatAppRunner() runner = ChatAppRunner()
app_config = SimpleNamespace( app_config = SimpleNamespace(
app_id="app-1", app_id="app-1",
@ -374,7 +419,6 @@ class TestChatAppRunner:
) )
with ( with (
patched_create_session(sqlite_session_factory),
patch.object(ChatAppRunner, "organize_prompt_messages", return_value=([], [])), patch.object(ChatAppRunner, "organize_prompt_messages", return_value=([], [])),
patch.object(ChatAppRunner, "moderation_for_inputs", return_value=(None, {}, "hi")), patch.object(ChatAppRunner, "moderation_for_inputs", return_value=(None, {}, "hi")),
patch.object(ChatAppRunner, "query_app_annotations_to_reply", return_value=None), patch.object(ChatAppRunner, "query_app_annotations_to_reply", return_value=None),
@ -383,15 +427,13 @@ class TestChatAppRunner:
runner.run( runner.run(
app_generate_entity, app_generate_entity,
DummyQueueManager(), DummyQueueManager(),
SimpleNamespace(), conversation,
SimpleNamespace(id="m1"), message,
sqlite_session, sqlite_session,
) )
def test_run_closes_explicit_session_before_stream_consumption( def test_run_closes_explicit_session_before_stream_consumption(self, sqlite_session: Session):
self, sqlite_session: Session, sqlite_session_factory: sessionmaker[Session] _, conversation, message = _persist_records(sqlite_session)
):
_persist_app(sqlite_session)
runner = ChatAppRunner() runner = ChatAppRunner()
app_config = SimpleNamespace( app_config = SimpleNamespace(
app_id="app-1", app_id="app-1",
@ -418,8 +460,9 @@ class TestChatAppRunner:
events = [] events = []
queue_manager = DummyQueueManager() queue_manager = DummyQueueManager()
model_instance = MagicMock() model_instance = MagicMock()
original_commit = sqlite_session.commit
original_close = sqlite_session.close def record_commit(_session: Session) -> None:
events.append("commit")
def invoke_stream(): def invoke_stream():
events.append("first-chunk") events.append("first-chunk")
@ -429,26 +472,27 @@ class TestChatAppRunner:
events.append("invoke") events.append("invoke")
return invoke_stream() return invoke_stream()
with ( event.listen(sqlite_session, "after_commit", record_commit)
patched_create_session(sqlite_session_factory), try:
patch.object(ChatAppRunner, "organize_prompt_messages", return_value=([], [])), with (
patch.object(ChatAppRunner, "moderation_for_inputs", return_value=(None, {}, "hi")), patch.object(ChatAppRunner, "organize_prompt_messages", return_value=([], [])),
patch.object(ChatAppRunner, "query_app_annotations_to_reply", return_value=None), patch.object(ChatAppRunner, "moderation_for_inputs", return_value=(None, {}, "hi")),
patch.object(ChatAppRunner, "check_hosting_moderation", return_value=False), patch.object(ChatAppRunner, "query_app_annotations_to_reply", return_value=None),
patch.object(ChatAppRunner, "recalc_llm_max_tokens"), patch.object(ChatAppRunner, "check_hosting_moderation", return_value=False),
patch.object( patch.object(ChatAppRunner, "recalc_llm_max_tokens"),
ChatAppRunner, patch.object(
"_handle_invoke_result", ChatAppRunner,
side_effect=lambda invoke_result, **kwargs: list(invoke_result), "_handle_invoke_result",
) as mock_handle, side_effect=lambda invoke_result, **kwargs: list(invoke_result),
patch("core.app.apps.chat.app_runner.ModelInstance", return_value=model_instance), ) as mock_handle,
patch.object(sqlite_session, "commit", side_effect=lambda: (events.append("commit"), original_commit())[1]), patch("core.app.apps.chat.app_runner.ModelInstance", return_value=model_instance),
patch.object(sqlite_session, "close", side_effect=lambda: (events.append("close"), original_close())[1]), ):
): model_instance.invoke_llm.side_effect = invoke_llm
model_instance.invoke_llm.side_effect = invoke_llm runner.run(app_generate_entity, queue_manager, conversation, message, sqlite_session)
runner.run(app_generate_entity, queue_manager, SimpleNamespace(), SimpleNamespace(id="m1"), sqlite_session) finally:
event.remove(sqlite_session, "after_commit", record_commit)
assert events == ["commit", "close", "commit", "close", "invoke", "first-chunk"] assert events == ["commit", "commit", "invoke", "first-chunk"]
mock_handle.assert_called_once_with( mock_handle.assert_called_once_with(
invoke_result=ANY, invoke_result=ANY,
queue_manager=queue_manager, queue_manager=queue_manager,

View File

@ -13,6 +13,7 @@ Covers:
- get_response_generator (ended / non-ended workflow run) - get_response_generator (ended / non-ended workflow run)
""" """
import json
import threading import threading
import uuid import uuid
from collections.abc import Callable from collections.abc import Callable
@ -26,7 +27,11 @@ from sqlalchemy.orm import Session
import services.app_generate_service as ags_module import services.app_generate_service as ags_module
from core.app.entities.app_invoke_entities import InvokeFrom from core.app.entities.app_invoke_entities import InvokeFrom
from enums import DeploymentEdition, QuotaType from enums import DeploymentEdition, QuotaType
from models.model import AppMode from graphon.enums import WorkflowExecutionStatus
from models.account import Account
from models.enums import AppStatus, CreatorUserRole
from models.model import App, AppMode
from models.workflow import Workflow, WorkflowRun, WorkflowRunTriggeredFrom, WorkflowType
from services.app_generate_service import AppGenerateService from services.app_generate_service import AppGenerateService
from services.errors.app import ( from services.errors.app import (
TriggerWorkflowServiceModeUnavailableError, TriggerWorkflowServiceModeUnavailableError,
@ -67,19 +72,27 @@ class _DummyRateLimit:
return generator return generator
def _make_app(mode: AppMode | str, *, max_active_requests: int = 0, is_agent: bool = False) -> MagicMock: def _make_app(mode: AppMode | str, *, max_active_requests: int = 0) -> App:
app = MagicMock() app = App(
app.mode = mode id="app-id",
app.id = "app-id" tenant_id="tenant-id",
app.tenant_id = "tenant-id" name="App",
app.max_active_requests = max_active_requests description="",
app.is_agent = is_agent mode=AppMode.CHAT if isinstance(mode, str) and mode == "invalid-mode" else mode,
app.is_agent_with_session.return_value = is_agent status=AppStatus.NORMAL,
enable_site=False,
enable_api=False,
api_rpm=0,
api_rph=0,
max_active_requests=max_active_requests,
)
if mode == "invalid-mode":
app.mode = mode # type: ignore[assignment]
return app return app
def _make_user() -> MagicMock: def _make_user() -> Account:
user = MagicMock() user = Account(name="User", email="user@example.com")
user.id = "user-id" user.id = "user-id"
return user return user
@ -95,14 +108,42 @@ def _make_workflow(
workflow_id: str = "workflow-id", workflow_id: str = "workflow-id",
created_by: str = "owner-id", created_by: str = "owner-id",
node_types: tuple[str, ...] = (), node_types: tuple[str, ...] = (),
) -> MagicMock: ) -> Workflow:
workflow = MagicMock() return Workflow(
workflow.id = workflow_id id=workflow_id,
workflow.created_by = created_by tenant_id="tenant-id",
workflow.walk_nodes.return_value = [ app_id="app-id",
(f"node-{index}", {"type": node_type}) for index, node_type in enumerate(node_types) type=WorkflowType.WORKFLOW,
] version=Workflow.VERSION_DRAFT,
return workflow graph=json.dumps(
{
"nodes": [
{"id": f"node-{index}", "data": {"type": node_type}} for index, node_type in enumerate(node_types)
],
"edges": [],
}
),
features={},
created_by=created_by,
environment_variables=[],
conversation_variables=[],
)
def _make_workflow_run(*, run_id: str, ended: bool) -> WorkflowRun:
run = WorkflowRun(
tenant_id="tenant-id",
app_id="app-id",
workflow_id="workflow-id",
type=WorkflowType.WORKFLOW,
triggered_from=WorkflowRunTriggeredFrom.APP_RUN,
version="published",
status=WorkflowExecutionStatus.SUCCEEDED if ended else WorkflowExecutionStatus.RUNNING,
created_by_role=CreatorUserRole.ACCOUNT,
created_by="user-id",
)
run.id = run_id
return run
@contextmanager @contextmanager
@ -335,7 +376,8 @@ class TestGenerate(_RealSessionTest):
"services.app_generate_service.AgentChatAppGenerator.convert_to_event_stream", "services.app_generate_service.AgentChatAppGenerator.convert_to_event_stream",
side_effect=lambda x: x, side_effect=lambda x: x,
) )
app = _make_app(AppMode.CHAT, is_agent=True) app = _make_app(AppMode.CHAT)
is_agent = mocker.patch.object(App, "is_agent_with_session", return_value=True)
session = self.session session = self.session
result = AppGenerateService.generate( result = AppGenerateService.generate(
app_model=app, app_model=app,
@ -347,7 +389,7 @@ class TestGenerate(_RealSessionTest):
) )
assert result == {"result": "agent-via-flag"} assert result == {"result": "agent-via-flag"}
gen_spy.assert_called_once() gen_spy.assert_called_once()
app.is_agent_with_session.assert_called_once_with(session=session) is_agent.assert_called_once_with(session=session)
# -- AGENT -------------------------------------------------------------- # -- AGENT --------------------------------------------------------------
def test_agent_mode_passes_session(self, mocker: MockerFixture): def test_agent_mode_passes_session(self, mocker: MockerFixture):
@ -383,7 +425,7 @@ class TestGenerate(_RealSessionTest):
"services.app_generate_service.ChatAppGenerator.convert_to_event_stream", "services.app_generate_service.ChatAppGenerator.convert_to_event_stream",
side_effect=lambda x: x, side_effect=lambda x: x,
) )
app = _make_app(AppMode.CHAT, is_agent=False) app = _make_app(AppMode.CHAT)
result = AppGenerateService.generate( result = AppGenerateService.generate(
app_model=app, app_model=app,
user=_make_user(), user=_make_user(),
@ -600,7 +642,7 @@ class TestGenerate(_RealSessionTest):
# -- Invalid mode ------------------------------------------------------- # -- Invalid mode -------------------------------------------------------
def test_invalid_mode_raises(self, mocker: MockerFixture): def test_invalid_mode_raises(self, mocker: MockerFixture):
app = _make_app("invalid-mode", is_agent=False) app = _make_app("invalid-mode")
with pytest.raises(ValueError, match="Invalid app mode"): with pytest.raises(ValueError, match="Invalid app mode"):
AppGenerateService.generate( AppGenerateService.generate(
app_model=app, app_model=app,
@ -1039,9 +1081,7 @@ class TestGenerateMoreLikeThis(_RealSessionTest):
class TestGetResponseGenerator: class TestGetResponseGenerator:
def test_non_ended_workflow_run(self, mocker: MockerFixture): def test_non_ended_workflow_run(self, mocker: MockerFixture):
app = _make_app(AppMode.ADVANCED_CHAT) app = _make_app(AppMode.ADVANCED_CHAT)
workflow_run = MagicMock() workflow_run = _make_workflow_run(run_id="run-1", ended=False)
workflow_run.id = "run-1"
workflow_run.status.is_ended.return_value = False
gen_instance = MagicMock() gen_instance = MagicMock()
gen_instance.retrieve_events.return_value = iter([{"event": "started"}]) gen_instance.retrieve_events.return_value = iter([{"event": "started"}])
@ -1057,9 +1097,7 @@ class TestGetResponseGenerator:
def test_ended_workflow_run_still_returns_generator(self, mocker: MockerFixture): def test_ended_workflow_run_still_returns_generator(self, mocker: MockerFixture):
"""Even when the run is ended, the current code still returns a generator (TODO branch).""" """Even when the run is ended, the current code still returns a generator (TODO branch)."""
app = _make_app(AppMode.WORKFLOW) app = _make_app(AppMode.WORKFLOW)
workflow_run = MagicMock() workflow_run = _make_workflow_run(run_id="run-2", ended=True)
workflow_run.id = "run-2"
workflow_run.status.is_ended.return_value = True
gen_instance = MagicMock() gen_instance = MagicMock()
gen_instance.retrieve_events.return_value = iter([]) gen_instance.retrieve_events.return_value = iter([])