mirror of
https://github.com/langgenius/dify.git
synced 2026-09-01 15:09:21 +08:00
test: migrate app service sessions to SQLite (#40088)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
This commit is contained in:
@@ -3,32 +3,95 @@ from __future__ import annotations
|
||||
from collections.abc import Callable
|
||||
from datetime import datetime
|
||||
from types import SimpleNamespace
|
||||
from typing import cast
|
||||
from unittest.mock import MagicMock, patch
|
||||
from uuid import uuid4
|
||||
|
||||
import pytest
|
||||
from sqlalchemy import event
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from graphon.model_runtime.entities.model_entities import ModelType
|
||||
from models import Account, Tenant
|
||||
from models.account import TenantAccountJoin, TenantAccountRole
|
||||
from models.agent import Agent, AgentIconType, AgentScope, AgentSource, AgentStatus
|
||||
from models.model import App, AppMode, AppModelConfig, IconType
|
||||
from models.workflow import Workflow
|
||||
from services.agent.errors import AgentAccessNotReadyError, AgentNameConflictError
|
||||
from services.app_service import AppListParams, AppService, CreateAppParams
|
||||
|
||||
|
||||
def _persist_account(session: Session) -> Account:
|
||||
tenant = Tenant(name="App Service Workspace")
|
||||
account = Account(name="Test Account", email=f"app-service-{uuid4()}@example.com")
|
||||
membership = TenantAccountJoin(
|
||||
tenant_id=tenant.id,
|
||||
account_id=account.id,
|
||||
current=True,
|
||||
role=TenantAccountRole.OWNER,
|
||||
)
|
||||
account._current_tenant = tenant
|
||||
session.add_all([tenant, account, membership])
|
||||
session.commit()
|
||||
return account
|
||||
|
||||
|
||||
def _persist_app(session: Session, *, tenant_id: str, name: str = "Visible App") -> App:
|
||||
app = App(
|
||||
id=str(uuid4()),
|
||||
tenant_id=tenant_id,
|
||||
name=name,
|
||||
mode=AppMode.CHAT,
|
||||
icon_type=IconType.EMOJI,
|
||||
icon="chat",
|
||||
icon_background="#FFFFFF",
|
||||
enable_site=False,
|
||||
enable_api=False,
|
||||
)
|
||||
session.add(app)
|
||||
session.commit()
|
||||
return app
|
||||
|
||||
|
||||
def _persist_agent_app(session: Session, *, app_name: str = "Old", agent_name: str = "Old") -> tuple[App, Agent]:
|
||||
tenant_id = str(uuid4())
|
||||
creator_id = str(uuid4())
|
||||
app = App(
|
||||
id=str(uuid4()),
|
||||
tenant_id=tenant_id,
|
||||
name=app_name,
|
||||
description="old",
|
||||
mode=AppMode.AGENT,
|
||||
icon_type=IconType.EMOJI,
|
||||
icon="robot",
|
||||
icon_background="#fff",
|
||||
enable_site=False,
|
||||
enable_api=False,
|
||||
created_by=creator_id,
|
||||
)
|
||||
agent = Agent(
|
||||
tenant_id=tenant_id,
|
||||
name=agent_name,
|
||||
description="old",
|
||||
role="research assistant",
|
||||
scope=AgentScope.ROSTER,
|
||||
source=AgentSource.AGENT_APP,
|
||||
status=AgentStatus.ACTIVE,
|
||||
icon_type=AgentIconType.EMOJI,
|
||||
icon="robot",
|
||||
icon_background="#fff",
|
||||
app_id=app.id,
|
||||
created_by=creator_id,
|
||||
)
|
||||
session.add_all([app, agent])
|
||||
session.commit()
|
||||
return app, agent
|
||||
|
||||
|
||||
class TestCreateAppTransactionBoundary:
|
||||
def test_commits_database_state_before_external_side_effects(self) -> None:
|
||||
session = MagicMock()
|
||||
account = Account(name="Test Account", email="test@example.com")
|
||||
account.id = "account-1"
|
||||
account._current_tenant = Tenant(name="Test Tenant")
|
||||
account._current_tenant.id = "tenant-1"
|
||||
def test_commits_database_state_before_external_side_effects(self, sqlite_session: Session) -> None:
|
||||
account = _persist_account(sqlite_session)
|
||||
phase_events: list[str] = []
|
||||
session.commit.side_effect = lambda: phase_events.append("commit")
|
||||
event.listen(sqlite_session, "after_commit", lambda _session: phase_events.append("commit"))
|
||||
|
||||
with (
|
||||
patch(
|
||||
@@ -45,21 +108,18 @@ class TestCreateAppTransactionBoundary:
|
||||
),
|
||||
patch("services.app_service.dify_config.BILLING_ENABLED", False),
|
||||
):
|
||||
AppService().create_app(
|
||||
"tenant-1",
|
||||
app = AppService().create_app(
|
||||
account.current_tenant_id,
|
||||
CreateAppParams(name="Workflow", mode=AppMode.WORKFLOW.value),
|
||||
account,
|
||||
session=session,
|
||||
session=sqlite_session,
|
||||
)
|
||||
|
||||
assert phase_events == ["commit", "signal", "commit", "external"]
|
||||
assert sqlite_session.get(App, app.id) is app
|
||||
|
||||
def test_falls_back_when_default_model_schema_is_unavailable(self) -> None:
|
||||
session = MagicMock()
|
||||
account = Account(name="Test Account", email="test@example.com")
|
||||
account.id = "account-1"
|
||||
account._current_tenant = Tenant(name="Test Tenant")
|
||||
account._current_tenant.id = "tenant-1"
|
||||
def test_falls_back_when_default_model_schema_is_unavailable(self, sqlite_session: Session) -> None:
|
||||
account = _persist_account(sqlite_session)
|
||||
model_type_instance = MagicMock()
|
||||
model_type_instance.get_model_schema.side_effect = ValueError("Base model unknown-model not found")
|
||||
model_instance = SimpleNamespace(
|
||||
@@ -71,9 +131,6 @@ class TestCreateAppTransactionBoundary:
|
||||
model_manager = MagicMock()
|
||||
model_manager.get_default_model_instance.return_value = model_instance
|
||||
model_manager.get_default_provider_model_name.return_value = ("openai", "gpt-4o")
|
||||
added_objects: list[object] = []
|
||||
session.add.side_effect = added_objects.append
|
||||
|
||||
with (
|
||||
patch("services.app_service.ModelManager.for_tenant", return_value=model_manager),
|
||||
patch("services.app_service.app_was_created.send"),
|
||||
@@ -85,13 +142,14 @@ class TestCreateAppTransactionBoundary:
|
||||
patch("services.app_service.dify_config.BILLING_ENABLED", False),
|
||||
):
|
||||
app = AppService().create_app(
|
||||
"tenant-1",
|
||||
account.current_tenant_id,
|
||||
CreateAppParams(name="Chat", mode=AppMode.CHAT.value),
|
||||
account,
|
||||
session=session,
|
||||
session=sqlite_session,
|
||||
)
|
||||
|
||||
app_model_config = next(obj for obj in added_objects if isinstance(obj, AppModelConfig))
|
||||
app_model_config = sqlite_session.get(AppModelConfig, app.app_model_config_id)
|
||||
assert app_model_config is not None
|
||||
assert app.mode == AppMode.CHAT
|
||||
assert app_model_config.model_dict == {
|
||||
"provider": "openai",
|
||||
@@ -100,7 +158,7 @@ class TestCreateAppTransactionBoundary:
|
||||
"completion_params": {},
|
||||
}
|
||||
model_manager.get_default_provider_model_name.assert_called_once_with(
|
||||
tenant_id="tenant-1", model_type=ModelType.LLM
|
||||
tenant_id=account.current_tenant_id, model_type=ModelType.LLM
|
||||
)
|
||||
|
||||
|
||||
@@ -108,17 +166,17 @@ class TestCreateAppTransactionBoundary:
|
||||
"update_status",
|
||||
[AppService.update_app_site_status, AppService.update_app_api_status],
|
||||
)
|
||||
def test_app_status_updates_commit_before_signal(update_status: Callable[..., App]) -> None:
|
||||
app = cast(App, SimpleNamespace(enable_site=False, enable_api=False, mode=AppMode.CHAT))
|
||||
session = MagicMock()
|
||||
def test_app_status_updates_commit_before_signal(update_status: Callable[..., App], sqlite_session: Session) -> None:
|
||||
account = _persist_account(sqlite_session)
|
||||
app = _persist_app(sqlite_session, tenant_id=account.current_tenant_id or "")
|
||||
phase_events: list[str] = []
|
||||
session.commit.side_effect = lambda: phase_events.append("commit")
|
||||
event.listen(sqlite_session, "after_commit", lambda _session: phase_events.append("commit"))
|
||||
|
||||
with (
|
||||
patch("services.app_service.current_user", SimpleNamespace(id="account-1")),
|
||||
patch("services.app_service.current_user", account),
|
||||
patch("services.app_service.app_was_updated.send", side_effect=lambda *_args: phase_events.append("signal")),
|
||||
):
|
||||
update_status(AppService(), app, True, session=session)
|
||||
update_status(AppService(), app, True, session=sqlite_session)
|
||||
|
||||
assert phase_events == ["commit", "signal"]
|
||||
|
||||
@@ -130,27 +188,20 @@ def test_app_status_updates_commit_before_signal(update_status: Callable[..., Ap
|
||||
AppService.update_app_api_status,
|
||||
],
|
||||
)
|
||||
def test_unpublished_agent_app_access_cannot_be_enabled(update_status: Callable[..., App]) -> None:
|
||||
app = cast(
|
||||
App,
|
||||
SimpleNamespace(
|
||||
id="app-1",
|
||||
tenant_id="tenant-1",
|
||||
mode=AppMode.AGENT,
|
||||
enable_site=False,
|
||||
enable_api=False,
|
||||
),
|
||||
)
|
||||
session = MagicMock()
|
||||
session.scalar.return_value = SimpleNamespace(id="agent-1")
|
||||
def test_unpublished_agent_app_access_cannot_be_enabled(
|
||||
update_status: Callable[..., App], sqlite_session: Session
|
||||
) -> None:
|
||||
app, _ = _persist_agent_app(sqlite_session)
|
||||
commits: list[str] = []
|
||||
event.listen(sqlite_session, "after_commit", lambda _session: commits.append("commit"))
|
||||
|
||||
with patch("services.app_service.agent_has_workflow_callable_active_snapshot", return_value=False):
|
||||
with pytest.raises(AgentAccessNotReadyError):
|
||||
update_status(AppService(), app, True, session=session)
|
||||
update_status(AppService(), app, True, session=sqlite_session)
|
||||
|
||||
assert app.enable_site is False
|
||||
assert app.enable_api is False
|
||||
session.commit.assert_not_called()
|
||||
assert commits == []
|
||||
|
||||
|
||||
class TestOpenapiVisibilityHelpers:
|
||||
@@ -160,127 +211,96 @@ class TestOpenapiVisibilityHelpers:
|
||||
gate passes" check so the controller can stay free of SQL.
|
||||
"""
|
||||
|
||||
def test_get_app_by_id_is_plain_session_get(self):
|
||||
def test_get_app_by_id_is_plain_session_get(self, sqlite_session: Session):
|
||||
"""``get_app_by_id`` must NOT apply status / visibility filters
|
||||
— callers (e.g. the openapi auth pipeline) need to differentiate
|
||||
404 (missing) from 403 (``enable_api`` off) and would lose that
|
||||
signal if the helper coalesced both into ``None``.
|
||||
"""
|
||||
mock_session = MagicMock()
|
||||
sentinel_app = App(
|
||||
status="archived",
|
||||
) # explicitly NOT "normal"
|
||||
mock_session.get.return_value = sentinel_app
|
||||
sentinel_app = _persist_app(sqlite_session, tenant_id=str(uuid4()))
|
||||
sentinel_app.status = "archived" # type: ignore[assignment]
|
||||
|
||||
assert AppService.get_app_by_id("app-uuid", mock_session) is sentinel_app
|
||||
mock_session.get.assert_called_once_with(App, "app-uuid")
|
||||
assert AppService.get_app_by_id(sentinel_app.id, sqlite_session) is sentinel_app
|
||||
|
||||
def test_get_app_by_id_returns_none_when_missing(self):
|
||||
mock_session = MagicMock()
|
||||
mock_session.get.return_value = None
|
||||
def test_get_app_by_id_returns_none_when_missing(self, sqlite_session: Session):
|
||||
assert AppService.get_app_by_id(str(uuid4()), sqlite_session) is None
|
||||
|
||||
assert AppService.get_app_by_id("missing", mock_session) is None
|
||||
|
||||
def test_get_visible_app_by_id_returns_app_when_visible(self):
|
||||
mock_session = MagicMock()
|
||||
app = App(
|
||||
status="normal",
|
||||
)
|
||||
mock_session.get.return_value = app
|
||||
def test_get_visible_app_by_id_returns_app_when_visible(self, sqlite_session: Session):
|
||||
app = _persist_app(sqlite_session, tenant_id=str(uuid4()))
|
||||
|
||||
with patch("services.app_service.is_openapi_visible", return_value=True):
|
||||
assert AppService.get_visible_app_by_id("app-uuid", mock_session) is app
|
||||
assert AppService.get_visible_app_by_id(app.id, sqlite_session) is app
|
||||
|
||||
mock_session.get.assert_called_once_with(App, "app-uuid")
|
||||
def test_get_visible_app_by_id_returns_none_when_row_missing(self, sqlite_session: Session):
|
||||
assert AppService.get_visible_app_by_id(str(uuid4()), sqlite_session) is None
|
||||
|
||||
def test_get_visible_app_by_id_returns_none_when_row_missing(self):
|
||||
mock_session = MagicMock()
|
||||
mock_session.get.return_value = None
|
||||
|
||||
assert AppService.get_visible_app_by_id("missing", mock_session) is None
|
||||
|
||||
def test_get_visible_app_by_id_returns_none_when_status_not_normal(self):
|
||||
def test_get_visible_app_by_id_returns_none_when_status_not_normal(self, sqlite_session: Session):
|
||||
"""Soft-deleted/archived rows must not surface on the openapi
|
||||
surface — the helper hides them by returning ``None``.
|
||||
"""
|
||||
mock_session = MagicMock()
|
||||
app = App(
|
||||
status="archived",
|
||||
)
|
||||
mock_session.get.return_value = app
|
||||
app = _persist_app(sqlite_session, tenant_id=str(uuid4()))
|
||||
app.status = "archived" # type: ignore[assignment]
|
||||
|
||||
with patch("services.app_service.is_openapi_visible", return_value=True):
|
||||
assert AppService.get_visible_app_by_id("app-uuid", mock_session) is None
|
||||
assert AppService.get_visible_app_by_id(app.id, sqlite_session) is None
|
||||
|
||||
def test_get_visible_app_by_id_returns_none_when_visibility_gate_rejects(self):
|
||||
def test_get_visible_app_by_id_returns_none_when_visibility_gate_rejects(self, sqlite_session: Session):
|
||||
"""``is_openapi_visible`` is the per-row counterpart to
|
||||
``apply_openapi_gate`` — when it returns False the helper must
|
||||
treat the row as invisible (not "found but unauthorized").
|
||||
"""
|
||||
mock_session = MagicMock()
|
||||
app = App(
|
||||
status="normal",
|
||||
)
|
||||
mock_session.get.return_value = app
|
||||
app = _persist_app(sqlite_session, tenant_id=str(uuid4()))
|
||||
|
||||
with patch("services.app_service.is_openapi_visible", return_value=False):
|
||||
assert AppService.get_visible_app_by_id("app-uuid", mock_session) is None
|
||||
assert AppService.get_visible_app_by_id(app.id, sqlite_session) is None
|
||||
|
||||
def test_find_visible_apps_by_name_returns_scalars_through_visibility_gate(self):
|
||||
def test_find_visible_apps_by_name_returns_scalars_through_visibility_gate(self, sqlite_session: Session):
|
||||
"""Tenant-scoped name lookup. The helper passes the SELECT through
|
||||
``apply_openapi_gate`` and materialises ``.scalars()`` into a list
|
||||
so the controller can branch on length (404 / single / 409).
|
||||
"""
|
||||
mock_session = MagicMock()
|
||||
rows = [App(), App()]
|
||||
mock_session.execute.return_value.scalars.return_value = iter(rows)
|
||||
tenant_id = str(uuid4())
|
||||
rows = [
|
||||
_persist_app(sqlite_session, tenant_id=tenant_id, name="my-app"),
|
||||
_persist_app(sqlite_session, tenant_id=tenant_id, name="my-app"),
|
||||
]
|
||||
_persist_app(sqlite_session, tenant_id=str(uuid4()), name="my-app")
|
||||
|
||||
with patch("services.app_service.apply_openapi_gate", side_effect=lambda q: q) as gate:
|
||||
out = AppService.find_visible_apps_by_name(name="my-app", tenant_id="tenant-1", session=mock_session)
|
||||
out = AppService.find_visible_apps_by_name(name="my-app", tenant_id=tenant_id, session=sqlite_session)
|
||||
|
||||
assert out == rows
|
||||
assert {app.id for app in out} == {app.id for app in rows}
|
||||
# Visibility gate must wrap the SELECT exactly once.
|
||||
gate.assert_called_once()
|
||||
mock_session.execute.assert_called_once()
|
||||
|
||||
def test_find_visible_apps_by_name_returns_empty_list_on_no_match(self):
|
||||
mock_session = MagicMock()
|
||||
mock_session.execute.return_value.scalars.return_value = iter([])
|
||||
|
||||
def test_find_visible_apps_by_name_returns_empty_list_on_no_match(self, sqlite_session: Session):
|
||||
with patch("services.app_service.apply_openapi_gate", side_effect=lambda q: q):
|
||||
out = AppService.find_visible_apps_by_name(name="nope", tenant_id="tenant-1", session=mock_session)
|
||||
out = AppService.find_visible_apps_by_name(name="nope", tenant_id=str(uuid4()), session=sqlite_session)
|
||||
|
||||
assert out == []
|
||||
|
||||
def test_find_visible_apps_by_ids_short_circuits_on_empty_input(self):
|
||||
def test_find_visible_apps_by_ids_short_circuits_on_empty_input(self, unbound_session: Session):
|
||||
"""Empty id list must not emit ``WHERE id IN ()`` — Postgres
|
||||
rejects empty IN lists and the call is a guaranteed no-op
|
||||
anyway. The helper returns ``[]`` without touching the session.
|
||||
"""
|
||||
mock_session = MagicMock()
|
||||
assert AppService.find_visible_apps_by_ids([], unbound_session) == []
|
||||
|
||||
assert AppService.find_visible_apps_by_ids([], mock_session) == []
|
||||
mock_session.execute.assert_not_called()
|
||||
|
||||
def test_find_visible_apps_by_ids_passes_through_visibility_gate(self):
|
||||
def test_find_visible_apps_by_ids_passes_through_visibility_gate(self, sqlite_session: Session):
|
||||
"""Bulk fetch routes through ``apply_openapi_gate`` exactly once
|
||||
and materialises the scalar rows. **No** status filter is
|
||||
applied here — the EE permitted-external pipeline filters
|
||||
non-normal hits in Python so its page count stays anchored.
|
||||
"""
|
||||
mock_session = MagicMock()
|
||||
rows = [App(), App()]
|
||||
mock_session.execute.return_value.scalars.return_value.all.return_value = rows
|
||||
rows = [_persist_app(sqlite_session, tenant_id=str(uuid4())) for _ in range(2)]
|
||||
|
||||
with patch("services.app_service.apply_openapi_gate", side_effect=lambda q: q) as gate:
|
||||
out = AppService.find_visible_apps_by_ids(["a", "b"], mock_session)
|
||||
out = AppService.find_visible_apps_by_ids([app.id for app in rows], sqlite_session)
|
||||
|
||||
assert out == rows
|
||||
assert {app.id for app in out} == {app.id for app in rows}
|
||||
gate.assert_called_once()
|
||||
mock_session.execute.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("sqlite_session", [(Account, App, AppModelConfig)], indirect=True)
|
||||
def test_get_recent_apps_uses_one_tenant_scoped_projection_query(sqlite_session: Session) -> None:
|
||||
tenant_id = str(uuid4())
|
||||
other_tenant_id = str(uuid4())
|
||||
@@ -370,28 +390,39 @@ def test_get_recent_apps_uses_one_tenant_scoped_projection_query(sqlite_session:
|
||||
|
||||
|
||||
class TestAppMeta:
|
||||
def test_loads_workflow_with_caller_session(self):
|
||||
session = MagicMock()
|
||||
session.get.return_value = SimpleNamespace(graph_dict={"nodes": []})
|
||||
app = cast(App, SimpleNamespace(mode=AppMode.WORKFLOW, workflow_id="workflow-1"))
|
||||
def test_loads_workflow_with_caller_session(self, sqlite_session: Session):
|
||||
tenant_id = str(uuid4())
|
||||
app = _persist_app(sqlite_session, tenant_id=tenant_id)
|
||||
app.mode = AppMode.WORKFLOW
|
||||
workflow = Workflow(
|
||||
id=str(uuid4()),
|
||||
tenant_id=tenant_id,
|
||||
app_id=app.id,
|
||||
type="workflow",
|
||||
version="draft",
|
||||
graph='{"nodes": []}',
|
||||
features="{}",
|
||||
created_by=str(uuid4()),
|
||||
)
|
||||
app.workflow_id = workflow.id
|
||||
sqlite_session.add(workflow)
|
||||
sqlite_session.commit()
|
||||
|
||||
assert AppService().get_app_meta(app, session=session) == {"tool_icons": {}}
|
||||
assert AppService().get_app_meta(app, session=sqlite_session) == {"tool_icons": {}}
|
||||
|
||||
session.get.assert_called_once_with(Workflow, "workflow-1")
|
||||
def test_loads_app_model_config_with_caller_session(self, sqlite_session: Session):
|
||||
app = _persist_app(sqlite_session, tenant_id=str(uuid4()))
|
||||
config = AppModelConfig(app_id=app.id, agent_mode='{"tools": []}')
|
||||
sqlite_session.add(config)
|
||||
sqlite_session.flush()
|
||||
app.app_model_config_id = config.id
|
||||
sqlite_session.commit()
|
||||
|
||||
def test_loads_app_model_config_with_caller_session(self):
|
||||
session = MagicMock()
|
||||
session.get.return_value = SimpleNamespace(agent_mode_dict={"tools": []})
|
||||
app = cast(App, SimpleNamespace(mode=AppMode.CHAT, app_model_config_id="config-1"))
|
||||
|
||||
assert AppService().get_app_meta(app, session=session) == {"tool_icons": {}}
|
||||
|
||||
session.get.assert_called_once_with(AppModelConfig, "config-1")
|
||||
assert AppService().get_app_meta(app, session=sqlite_session) == {"tool_icons": {}}
|
||||
|
||||
|
||||
class TestGetApp:
|
||||
def test_legacy_agent_detection_uses_caller_session(self):
|
||||
session = MagicMock()
|
||||
def test_legacy_agent_detection_uses_caller_session(self, unbound_session: Session):
|
||||
app = App(
|
||||
mode=AppMode.CHAT,
|
||||
)
|
||||
@@ -404,13 +435,12 @@ class TestGetApp:
|
||||
patch.object(App, "app_model_config_with_session") as get_model_config,
|
||||
patch("services.app_service.current_user", account),
|
||||
):
|
||||
assert AppService().get_app(app, session=session) is app
|
||||
assert AppService().get_app(app, session=unbound_session) is app
|
||||
|
||||
is_agent.assert_called_once_with(session=session)
|
||||
is_agent.assert_called_once_with(session=unbound_session)
|
||||
get_model_config.assert_not_called()
|
||||
|
||||
def test_agent_model_config_uses_caller_session(self):
|
||||
session = MagicMock()
|
||||
def test_agent_model_config_uses_caller_session(self, unbound_session: Session):
|
||||
app = App(
|
||||
mode=AppMode.AGENT_CHAT,
|
||||
)
|
||||
@@ -423,10 +453,10 @@ class TestGetApp:
|
||||
patch.object(App, "app_model_config_with_session", return_value=None) as get_model_config,
|
||||
patch("services.app_service.current_user", account),
|
||||
):
|
||||
assert AppService().get_app(app, session=session) is app
|
||||
assert AppService().get_app(app, session=unbound_session) is app
|
||||
|
||||
is_agent.assert_not_called()
|
||||
get_model_config.assert_called_once_with(session=session)
|
||||
get_model_config.assert_called_once_with(session=unbound_session)
|
||||
|
||||
|
||||
class TestAgentAppType:
|
||||
@@ -459,43 +489,16 @@ class TestAgentAppType:
|
||||
)
|
||||
assert app.bound_agent_id is None
|
||||
|
||||
def test_update_agent_app_syncs_backing_agent_identity(self):
|
||||
from models.agent import AgentIconType
|
||||
from models.model import AppMode, IconType
|
||||
from services.app_service import AppService
|
||||
|
||||
app = SimpleNamespace(
|
||||
id="app-1",
|
||||
tenant_id="tenant-1",
|
||||
mode=AppMode.AGENT,
|
||||
name="Old",
|
||||
description="old",
|
||||
role="draft",
|
||||
icon_type=IconType.EMOJI,
|
||||
icon="robot",
|
||||
icon_background="#fff",
|
||||
use_icon_as_answer_icon=False,
|
||||
max_active_requests=None,
|
||||
created_by="account-1",
|
||||
)
|
||||
backing_agent = SimpleNamespace(
|
||||
name="Old",
|
||||
description="old",
|
||||
role="draft",
|
||||
icon_type=AgentIconType.EMOJI,
|
||||
icon="robot",
|
||||
icon_background="#fff",
|
||||
updated_by=None,
|
||||
updated_at=None,
|
||||
)
|
||||
def test_update_agent_app_syncs_backing_agent_identity(self, sqlite_session: Session):
|
||||
app, backing_agent = _persist_agent_app(sqlite_session)
|
||||
account_id = str(uuid4())
|
||||
|
||||
with (
|
||||
patch("services.app_service.db") as mock_db,
|
||||
patch("services.app_service.current_user", SimpleNamespace(id="account-2")),
|
||||
patch("services.app_service.current_user", SimpleNamespace(id=account_id)),
|
||||
patch("services.app_service.app_was_updated.send"),
|
||||
):
|
||||
mock_db.session.scalar.return_value = backing_agent
|
||||
updated_app = AppService().update_app(
|
||||
app, # type: ignore[arg-type]
|
||||
app,
|
||||
{
|
||||
"name": "Iris",
|
||||
"description": "agent app",
|
||||
@@ -506,7 +509,7 @@ class TestAgentAppType:
|
||||
"use_icon_as_answer_icon": False,
|
||||
"max_active_requests": 0,
|
||||
},
|
||||
session=mock_db.session,
|
||||
session=sqlite_session,
|
||||
)
|
||||
|
||||
assert updated_app.name == "Iris"
|
||||
@@ -516,46 +519,18 @@ class TestAgentAppType:
|
||||
assert backing_agent.icon_type == AgentIconType.IMAGE
|
||||
assert backing_agent.icon == "file-id"
|
||||
assert backing_agent.icon_background == "#123456"
|
||||
assert backing_agent.updated_by == "account-2"
|
||||
assert backing_agent.updated_by == account_id
|
||||
assert backing_agent.updated_at == updated_app.updated_at
|
||||
|
||||
def test_update_agent_app_preserves_role_when_args_omit_it(self):
|
||||
from models.agent import AgentIconType
|
||||
from models.model import AppMode, IconType
|
||||
from services.app_service import AppService
|
||||
|
||||
app = SimpleNamespace(
|
||||
id="app-1",
|
||||
tenant_id="tenant-1",
|
||||
mode=AppMode.AGENT,
|
||||
name="Old",
|
||||
description="old",
|
||||
role="draft",
|
||||
icon_type=IconType.EMOJI,
|
||||
icon="robot",
|
||||
icon_background="#fff",
|
||||
use_icon_as_answer_icon=False,
|
||||
max_active_requests=None,
|
||||
created_by="account-1",
|
||||
)
|
||||
backing_agent = SimpleNamespace(
|
||||
name="Old",
|
||||
description="old",
|
||||
role="research assistant",
|
||||
icon_type=AgentIconType.EMOJI,
|
||||
icon="robot",
|
||||
icon_background="#fff",
|
||||
updated_by=None,
|
||||
updated_at=None,
|
||||
)
|
||||
def test_update_agent_app_preserves_role_when_args_omit_it(self, sqlite_session: Session):
|
||||
app, backing_agent = _persist_agent_app(sqlite_session)
|
||||
|
||||
with (
|
||||
patch("services.app_service.db") as mock_db,
|
||||
patch("services.app_service.current_user", SimpleNamespace(id="account-2")),
|
||||
patch("services.app_service.current_user", SimpleNamespace(id=str(uuid4()))),
|
||||
patch("services.app_service.app_was_updated.send"),
|
||||
):
|
||||
mock_db.session.scalar.return_value = backing_agent
|
||||
AppService().update_app(
|
||||
app, # type: ignore[arg-type]
|
||||
app,
|
||||
{
|
||||
"name": "Iris",
|
||||
"description": "agent app",
|
||||
@@ -565,48 +540,20 @@ class TestAgentAppType:
|
||||
"use_icon_as_answer_icon": False,
|
||||
"max_active_requests": 0,
|
||||
},
|
||||
session=mock_db.session,
|
||||
session=sqlite_session,
|
||||
)
|
||||
|
||||
assert backing_agent.role == "research assistant"
|
||||
|
||||
def test_update_agent_app_clears_role_when_args_set_empty_string(self):
|
||||
from models.agent import AgentIconType
|
||||
from models.model import AppMode, IconType
|
||||
from services.app_service import AppService
|
||||
|
||||
app = SimpleNamespace(
|
||||
id="app-1",
|
||||
tenant_id="tenant-1",
|
||||
mode=AppMode.AGENT,
|
||||
name="Old",
|
||||
description="old",
|
||||
role="draft",
|
||||
icon_type=IconType.EMOJI,
|
||||
icon="robot",
|
||||
icon_background="#fff",
|
||||
use_icon_as_answer_icon=False,
|
||||
max_active_requests=None,
|
||||
created_by="account-1",
|
||||
)
|
||||
backing_agent = SimpleNamespace(
|
||||
name="Old",
|
||||
description="old",
|
||||
role="research assistant",
|
||||
icon_type=AgentIconType.EMOJI,
|
||||
icon="robot",
|
||||
icon_background="#fff",
|
||||
updated_by=None,
|
||||
updated_at=None,
|
||||
)
|
||||
def test_update_agent_app_clears_role_when_args_set_empty_string(self, sqlite_session: Session):
|
||||
app, backing_agent = _persist_agent_app(sqlite_session)
|
||||
|
||||
with (
|
||||
patch("services.app_service.db") as mock_db,
|
||||
patch("services.app_service.current_user", SimpleNamespace(id="account-2")),
|
||||
patch("services.app_service.current_user", SimpleNamespace(id=str(uuid4()))),
|
||||
patch("services.app_service.app_was_updated.send"),
|
||||
):
|
||||
mock_db.session.scalar.return_value = backing_agent
|
||||
AppService().update_app(
|
||||
app, # type: ignore[arg-type]
|
||||
app,
|
||||
{
|
||||
"name": "Iris",
|
||||
"description": "agent app",
|
||||
@@ -617,50 +564,34 @@ class TestAgentAppType:
|
||||
"use_icon_as_answer_icon": False,
|
||||
"max_active_requests": 0,
|
||||
},
|
||||
session=mock_db.session,
|
||||
session=sqlite_session,
|
||||
)
|
||||
|
||||
assert backing_agent.role == ""
|
||||
|
||||
def test_update_agent_app_duplicate_name_rolls_back_and_raises_conflict(self):
|
||||
from models.agent import AgentIconType
|
||||
from models.model import AppMode, IconType
|
||||
from services.app_service import AppService
|
||||
|
||||
app = SimpleNamespace(
|
||||
id="app-1",
|
||||
tenant_id="tenant-1",
|
||||
mode=AppMode.AGENT,
|
||||
name="Old",
|
||||
description="old",
|
||||
role="draft",
|
||||
icon_type=IconType.EMOJI,
|
||||
icon="robot",
|
||||
icon_background="#fff",
|
||||
use_icon_as_answer_icon=False,
|
||||
max_active_requests=None,
|
||||
created_by="account-1",
|
||||
)
|
||||
backing_agent = SimpleNamespace(
|
||||
name="Old",
|
||||
description="old",
|
||||
role="research assistant",
|
||||
icon_type=AgentIconType.EMOJI,
|
||||
icon="robot",
|
||||
icon_background="#fff",
|
||||
updated_by=None,
|
||||
updated_at=None,
|
||||
def test_update_agent_app_duplicate_name_rolls_back_and_raises_conflict(self, sqlite_session: Session):
|
||||
app, backing_agent = _persist_agent_app(sqlite_session)
|
||||
existing = Agent(
|
||||
tenant_id=app.tenant_id,
|
||||
name="Existing Agent",
|
||||
description="existing",
|
||||
role="",
|
||||
scope=AgentScope.ROSTER,
|
||||
source=AgentSource.ROSTER,
|
||||
status=AgentStatus.ACTIVE,
|
||||
)
|
||||
sqlite_session.add(existing)
|
||||
sqlite_session.commit()
|
||||
rollback_events: list[str] = []
|
||||
event.listen(sqlite_session, "after_rollback", lambda _session: rollback_events.append("rollback"))
|
||||
|
||||
with (
|
||||
patch("services.app_service.db") as mock_db,
|
||||
patch("services.app_service.current_user", SimpleNamespace(id="account-2")),
|
||||
patch("services.app_service.current_user", SimpleNamespace(id=str(uuid4()))),
|
||||
patch("services.app_service.app_was_updated.send"),
|
||||
):
|
||||
mock_db.session.scalar.return_value = backing_agent
|
||||
mock_db.session.commit.side_effect = IntegrityError("duplicate", None, None)
|
||||
with pytest.raises(AgentNameConflictError):
|
||||
AppService().update_app(
|
||||
app, # type: ignore[arg-type]
|
||||
app,
|
||||
{
|
||||
"name": "Existing Agent",
|
||||
"description": "agent app",
|
||||
@@ -671,23 +602,39 @@ class TestAgentAppType:
|
||||
"use_icon_as_answer_icon": False,
|
||||
"max_active_requests": 0,
|
||||
},
|
||||
session=mock_db.session,
|
||||
session=sqlite_session,
|
||||
)
|
||||
|
||||
mock_db.session.rollback.assert_called_once()
|
||||
assert rollback_events == ["rollback"]
|
||||
sqlite_session.expire_all()
|
||||
assert sqlite_session.get(Agent, backing_agent.id).name == "Old" # type: ignore[union-attr]
|
||||
|
||||
def test_delete_agent_app_archives_backing_agent(self):
|
||||
from models.agent import AgentStatus
|
||||
from models.model import AppMode
|
||||
from services.app_service import AppService
|
||||
|
||||
app = SimpleNamespace(id="app-1", tenant_id="tenant-1", mode=AppMode.AGENT)
|
||||
backing_agent = SimpleNamespace(id="agent-1", status=AgentStatus.ACTIVE, archived_by=None, archived_at=None)
|
||||
def test_delete_agent_app_archives_backing_agent(self, sqlite_session: Session):
|
||||
app, backing_agent = _persist_agent_app(sqlite_session)
|
||||
workflow_agents = [
|
||||
Agent(
|
||||
tenant_id=app.tenant_id,
|
||||
name=f"Workflow Agent {index}",
|
||||
description="",
|
||||
role="",
|
||||
scope=AgentScope.WORKFLOW_ONLY,
|
||||
source=AgentSource.WORKFLOW,
|
||||
status=AgentStatus.ACTIVE,
|
||||
app_id=app.id,
|
||||
workflow_id=str(uuid4()),
|
||||
workflow_node_id=f"node-{index}",
|
||||
)
|
||||
for index in range(2)
|
||||
]
|
||||
sqlite_session.add_all(workflow_agents)
|
||||
sqlite_session.commit()
|
||||
account_id = str(uuid4())
|
||||
events: list[str] = []
|
||||
event.listen(sqlite_session, "after_commit", lambda _session: events.append("commit"))
|
||||
|
||||
with (
|
||||
patch("services.app_service.db") as mock_db,
|
||||
patch("services.app_service.current_user", SimpleNamespace(id="account-2")),
|
||||
patch("services.app_service.current_user", SimpleNamespace(id=account_id)),
|
||||
patch("services.app_service.app_was_deleted.send"),
|
||||
patch("services.app_service.BillingService"),
|
||||
patch("services.app_service.EnterpriseService"),
|
||||
patch("services.app_service.FeatureService"),
|
||||
@@ -712,62 +659,67 @@ class TestAgentAppType:
|
||||
side_effect=lambda **_kwargs: events.append("enqueue"),
|
||||
) as mock_enqueue_collection,
|
||||
):
|
||||
mock_db.session.scalar.return_value = backing_agent
|
||||
mock_db.session.commit.side_effect = lambda: events.append("commit")
|
||||
workflow_agents = MagicMock()
|
||||
workflow_agents.all.return_value = ["workflow-agent-1", "workflow-agent-2"]
|
||||
bindings = MagicMock()
|
||||
bindings.all.return_value = []
|
||||
mock_db.session.scalars.side_effect = [workflow_agents, bindings]
|
||||
AppService().delete_app(app, session=mock_db.session) # type: ignore[arg-type]
|
||||
AppService().delete_app(app, session=sqlite_session)
|
||||
|
||||
assert events == ["retire-app-workspaces", "commit", "retire-workflow-agents", "enqueue"]
|
||||
assert backing_agent.status == AgentStatus.ARCHIVED
|
||||
assert backing_agent.archived_by == "account-2"
|
||||
assert backing_agent.archived_at is not None
|
||||
mock_db.session.delete.assert_called_once_with(app)
|
||||
sqlite_session.expire_all()
|
||||
persisted_agent = sqlite_session.get(Agent, backing_agent.id)
|
||||
assert persisted_agent is not None
|
||||
assert sqlite_session.get(App, app.id) is None
|
||||
assert persisted_agent.status == AgentStatus.ARCHIVED
|
||||
assert persisted_agent.archived_by == account_id
|
||||
assert persisted_agent.archived_at is not None
|
||||
mock_workflow_retirement.assert_called_once_with(
|
||||
tenant_id="tenant-1",
|
||||
agent_ids=["workflow-agent-1", "workflow-agent-2"],
|
||||
account_id="account-2",
|
||||
tenant_id=app.tenant_id,
|
||||
agent_ids=[agent.id for agent in workflow_agents],
|
||||
account_id=account_id,
|
||||
)
|
||||
mock_retire_workspaces.assert_called_once_with(
|
||||
session=mock_db.session,
|
||||
tenant_id="tenant-1",
|
||||
app_id="app-1",
|
||||
session=sqlite_session,
|
||||
tenant_id=app.tenant_id,
|
||||
app_id=app.id,
|
||||
)
|
||||
mock_retire_homes.assert_called_once_with(
|
||||
session=mock_db.session,
|
||||
tenant_id="tenant-1",
|
||||
agent_id="agent-1",
|
||||
session=sqlite_session,
|
||||
tenant_id=app.tenant_id,
|
||||
agent_id=backing_agent.id,
|
||||
)
|
||||
mock_enqueue_collection.assert_called_once_with(
|
||||
tenant_id="tenant-1",
|
||||
tenant_id=app.tenant_id,
|
||||
workspace_ids=["workspace-1"],
|
||||
binding_ids=["workflow-binding-1"],
|
||||
home_snapshot_ids=["home-1", "workflow-home-1"],
|
||||
)
|
||||
|
||||
def test_delete_app_commit_failure_does_not_retire_workflow_agents_or_enqueue(self):
|
||||
from models.model import AppMode
|
||||
from services.app_service import AppService
|
||||
|
||||
app = SimpleNamespace(id="app-1", tenant_id="tenant-1", mode=AppMode.WORKFLOW)
|
||||
def test_delete_app_commit_failure_does_not_retire_workflow_agents_or_enqueue(
|
||||
self, sqlite_session: Session, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
app = _persist_app(sqlite_session, tenant_id=str(uuid4()))
|
||||
app.mode = AppMode.WORKFLOW
|
||||
workflow_agent = Agent(
|
||||
tenant_id=app.tenant_id,
|
||||
name="Workflow Agent",
|
||||
description="",
|
||||
role="",
|
||||
scope=AgentScope.WORKFLOW_ONLY,
|
||||
source=AgentSource.WORKFLOW,
|
||||
status=AgentStatus.ACTIVE,
|
||||
app_id=app.id,
|
||||
workflow_id=str(uuid4()),
|
||||
workflow_node_id="node-1",
|
||||
)
|
||||
sqlite_session.add(workflow_agent)
|
||||
sqlite_session.commit()
|
||||
monkeypatch.setattr(sqlite_session, "commit", MagicMock(side_effect=RuntimeError("commit failed")))
|
||||
with (
|
||||
patch("services.app_service.db") as mock_db,
|
||||
patch("services.app_service.current_user", SimpleNamespace(id="account-2")),
|
||||
patch("services.app_service.current_user", SimpleNamespace(id=str(uuid4()))),
|
||||
patch("services.app_service.app_was_deleted.send"),
|
||||
patch("services.app_service.AgentWorkspaceService.retire_all_for_app", return_value=["workspace-1"]),
|
||||
patch("services.app_service.WorkflowAgentRetirementService.retire_unowned") as retire_unowned,
|
||||
patch("services.app_service.enqueue_agent_resource_collection") as enqueue_collection,
|
||||
):
|
||||
mock_db.session.scalar.return_value = None
|
||||
workflow_agents = MagicMock()
|
||||
workflow_agents.all.return_value = ["workflow-agent-1"]
|
||||
mock_db.session.scalars.return_value = workflow_agents
|
||||
mock_db.session.commit.side_effect = RuntimeError("commit failed")
|
||||
|
||||
with pytest.raises(RuntimeError, match="commit failed"):
|
||||
AppService().delete_app(app, session=mock_db.session) # type: ignore[arg-type]
|
||||
AppService().delete_app(app, session=sqlite_session)
|
||||
|
||||
retire_unowned.assert_not_called()
|
||||
enqueue_collection.assert_not_called()
|
||||
|
||||
@@ -224,6 +224,10 @@ def factory():
|
||||
class TestAudioServiceASR:
|
||||
"""Test speech-to-text (ASR) operations."""
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _bind_sqlite_session(self, sqlite_session: Session) -> None:
|
||||
self.session = sqlite_session
|
||||
|
||||
@patch("services.audio_service.ModelManager.for_tenant", autospec=True)
|
||||
def test_transcript_asr_success_chat_mode(self, mock_model_manager_class, factory: AudioServiceTestDataFactory):
|
||||
"""Test successful ASR transcription in CHAT mode."""
|
||||
@@ -242,7 +246,7 @@ class TestAudioServiceASR:
|
||||
mock_model_manager.get_default_model_instance.return_value = mock_model_instance
|
||||
|
||||
# Act
|
||||
result = AudioService.transcript_asr(app_model=app, file=file, session=MagicMock(), end_user="user-123")
|
||||
result = AudioService.transcript_asr(app_model=app, file=file, session=self.session, end_user="user-123")
|
||||
|
||||
# Assert
|
||||
assert result == {"text": "Transcribed text"}
|
||||
@@ -269,7 +273,7 @@ class TestAudioServiceASR:
|
||||
mock_model_manager.get_default_model_instance.return_value = mock_model_instance
|
||||
|
||||
# Act
|
||||
result = AudioService.transcript_asr(app_model=app, file=file, session=MagicMock())
|
||||
result = AudioService.transcript_asr(app_model=app, file=file, session=self.session)
|
||||
|
||||
# Assert
|
||||
assert result == {"text": "Workflow transcribed text"}
|
||||
@@ -290,7 +294,7 @@ class TestAudioServiceASR:
|
||||
mock_model_instance.invoke_speech2text.return_value = "Published Agent transcript"
|
||||
mock_model_manager_class.return_value.get_default_model_instance.return_value = mock_model_instance
|
||||
|
||||
result = AudioService.transcript_asr(app_model=app, file=file, session=MagicMock(), end_user="end-user-1")
|
||||
result = AudioService.transcript_asr(app_model=app, file=file, session=self.session, end_user="end-user-1")
|
||||
|
||||
assert result == {"text": "Published Agent transcript"}
|
||||
mock_roster_service_class.return_value.get_published_agent_soul_for_app.assert_called_once_with(
|
||||
@@ -314,7 +318,7 @@ class TestAudioServiceASR:
|
||||
mock_model_instance.invoke_speech2text.return_value = "Legacy Agent transcript"
|
||||
mock_model_manager_class.return_value.get_default_model_instance.return_value = mock_model_instance
|
||||
|
||||
result = AudioService.transcript_asr(app_model=app, file=file, session=MagicMock())
|
||||
result = AudioService.transcript_asr(app_model=app, file=file, session=self.session)
|
||||
|
||||
assert result == {"text": "Legacy Agent transcript"}
|
||||
|
||||
@@ -333,7 +337,7 @@ class TestAudioServiceASR:
|
||||
app_model=app,
|
||||
agent_soul=agent_soul,
|
||||
file=file,
|
||||
session=MagicMock(),
|
||||
session=self.session,
|
||||
end_user="account-1",
|
||||
)
|
||||
|
||||
@@ -354,7 +358,7 @@ class TestAudioServiceASR:
|
||||
file = factory.create_file_storage_mock()
|
||||
|
||||
with pytest.raises(SpeechToTextDisabledServiceError):
|
||||
AudioService.transcript_agent_asr(app_model=app, agent_soul=agent_soul, file=file, session=MagicMock())
|
||||
AudioService.transcript_agent_asr(app_model=app, agent_soul=agent_soul, file=file, session=self.session)
|
||||
|
||||
@patch("services.audio_service.ModelManager.for_tenant", autospec=True)
|
||||
def test_transcript_agent_asr_preserves_legacy_feature_fallback(
|
||||
@@ -372,7 +376,7 @@ class TestAudioServiceASR:
|
||||
app_model=app,
|
||||
agent_soul=AgentSoulConfig(),
|
||||
file=file,
|
||||
session=MagicMock(),
|
||||
session=self.session,
|
||||
)
|
||||
|
||||
assert result == {"text": "Legacy feature transcript"}
|
||||
@@ -385,7 +389,7 @@ class TestAudioServiceASR:
|
||||
agent_soul = AgentSoulConfig.model_validate({"app_features": {"speech_to_text": {"enabled": False}}})
|
||||
|
||||
with pytest.raises(SpeechToTextDisabledServiceError):
|
||||
AudioService.transcript_agent_asr(app_model=app, agent_soul=agent_soul, file=file, session=MagicMock())
|
||||
AudioService.transcript_agent_asr(app_model=app, agent_soul=agent_soul, file=file, session=self.session)
|
||||
|
||||
def test_transcript_asr_raises_error_when_feature_disabled_chat_mode(self, factory: AudioServiceTestDataFactory):
|
||||
"""Test that ASR raises error when speech-to-text is disabled in CHAT mode."""
|
||||
@@ -399,7 +403,7 @@ class TestAudioServiceASR:
|
||||
|
||||
# Act & Assert
|
||||
with pytest.raises(SpeechToTextDisabledServiceError):
|
||||
AudioService.transcript_asr(app_model=app, file=file, session=MagicMock())
|
||||
AudioService.transcript_asr(app_model=app, file=file, session=self.session)
|
||||
|
||||
def test_transcript_asr_raises_error_when_feature_disabled_workflow_mode(
|
||||
self, factory: AudioServiceTestDataFactory
|
||||
@@ -415,7 +419,7 @@ class TestAudioServiceASR:
|
||||
|
||||
# Act & Assert
|
||||
with pytest.raises(SpeechToTextDisabledServiceError):
|
||||
AudioService.transcript_asr(app_model=app, file=file, session=MagicMock())
|
||||
AudioService.transcript_asr(app_model=app, file=file, session=self.session)
|
||||
|
||||
def test_transcript_asr_raises_error_when_workflow_missing(self, factory: AudioServiceTestDataFactory):
|
||||
"""Test that ASR raises error when workflow is missing in WORKFLOW mode."""
|
||||
@@ -428,7 +432,7 @@ class TestAudioServiceASR:
|
||||
|
||||
# Act & Assert
|
||||
with pytest.raises(SpeechToTextDisabledServiceError):
|
||||
AudioService.transcript_asr(app_model=app, file=file, session=MagicMock())
|
||||
AudioService.transcript_asr(app_model=app, file=file, session=self.session)
|
||||
|
||||
def test_transcript_asr_raises_error_when_no_file_uploaded(self, factory: AudioServiceTestDataFactory):
|
||||
"""Test that ASR raises error when no file is uploaded."""
|
||||
@@ -441,7 +445,7 @@ class TestAudioServiceASR:
|
||||
|
||||
# Act & Assert
|
||||
with pytest.raises(NoAudioUploadedServiceError):
|
||||
AudioService.transcript_asr(app_model=app, file=None, session=MagicMock())
|
||||
AudioService.transcript_asr(app_model=app, file=None, session=self.session)
|
||||
|
||||
def test_transcript_asr_raises_error_for_unsupported_audio_type(self, factory: AudioServiceTestDataFactory):
|
||||
"""Test that ASR raises error for unsupported audio file types."""
|
||||
@@ -455,7 +459,7 @@ class TestAudioServiceASR:
|
||||
|
||||
# Act & Assert
|
||||
with pytest.raises(UnsupportedAudioTypeServiceError):
|
||||
AudioService.transcript_asr(app_model=app, file=file, session=MagicMock())
|
||||
AudioService.transcript_asr(app_model=app, file=file, session=self.session)
|
||||
|
||||
def test_transcript_asr_raises_error_for_large_file(self, factory: AudioServiceTestDataFactory):
|
||||
"""Test that ASR raises error when file exceeds size limit (30MB)."""
|
||||
@@ -471,7 +475,7 @@ class TestAudioServiceASR:
|
||||
|
||||
# Act & Assert
|
||||
with pytest.raises(AudioTooLargeServiceError, match="Audio size larger than 30 mb"):
|
||||
AudioService.transcript_asr(app_model=app, file=file, session=MagicMock())
|
||||
AudioService.transcript_asr(app_model=app, file=file, session=self.session)
|
||||
|
||||
@patch("services.audio_service.ModelManager.for_tenant", autospec=True)
|
||||
def test_transcript_asr_raises_error_when_no_model_instance(
|
||||
@@ -492,10 +496,9 @@ class TestAudioServiceASR:
|
||||
|
||||
# Act & Assert
|
||||
with pytest.raises(ProviderNotSupportSpeechToTextServiceError):
|
||||
AudioService.transcript_asr(app_model=app, file=file, session=MagicMock())
|
||||
AudioService.transcript_asr(app_model=app, file=file, session=self.session)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("sqlite_session", [(Message,)], indirect=True)
|
||||
class TestAudioServiceTTS:
|
||||
"""Test text-to-speech (TTS) operations."""
|
||||
|
||||
|
||||
@@ -11,7 +11,7 @@ import json
|
||||
from unittest.mock import MagicMock, Mock, patch
|
||||
|
||||
import pytest
|
||||
from sqlalchemy import asc, desc
|
||||
from sqlalchemy import asc, desc, event
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from core.app.entities.app_invoke_entities import InvokeFrom
|
||||
@@ -150,18 +150,19 @@ class ConversationServiceTestDataFactory:
|
||||
return conversation
|
||||
|
||||
|
||||
def test_delete_retires_then_commits_before_enqueue(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
def test_delete_retires_then_commits_before_enqueue(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session) -> None:
|
||||
app = ConversationServiceTestDataFactory.create_app()
|
||||
conversation = ConversationServiceTestDataFactory.create_conversation()
|
||||
conversation.agent_workspace_binding_id = "conversation-binding-1"
|
||||
session = MagicMock()
|
||||
sqlite_session.add(conversation)
|
||||
sqlite_session.flush()
|
||||
events: list[str] = []
|
||||
get_binding = MagicMock(return_value=Mock(id="conversation-binding-1"))
|
||||
retire_binding = MagicMock(side_effect=lambda **_kwargs: events.append("retire") or "conversation-binding-1")
|
||||
monkeypatch.setattr(ConversationService, "get_conversation", MagicMock(return_value=conversation))
|
||||
monkeypatch.setattr(AgentWorkspaceService, "get_active_binding", get_binding)
|
||||
monkeypatch.setattr(AgentWorkspaceService, "retire_binding", retire_binding)
|
||||
session.commit.side_effect = lambda: events.append("commit")
|
||||
event.listen(sqlite_session, "after_commit", lambda _session: events.append("commit"))
|
||||
monkeypatch.setattr(
|
||||
conversation_service,
|
||||
"enqueue_agent_resource_collection",
|
||||
@@ -169,19 +170,22 @@ def test_delete_retires_then_commits_before_enqueue(monkeypatch: pytest.MonkeyPa
|
||||
)
|
||||
monkeypatch.setattr(conversation_service.delete_conversation_related_data, "delay", MagicMock())
|
||||
|
||||
ConversationService.delete(app, conversation.id, None, session=session)
|
||||
ConversationService.delete(app, conversation.id, None, session=sqlite_session)
|
||||
|
||||
assert events == ["retire", "commit", "enqueue"]
|
||||
assert get_binding.call_args.kwargs["binding_id"] == "conversation-binding-1"
|
||||
assert retire_binding.call_args.kwargs["binding_id"] == "conversation-binding-1"
|
||||
|
||||
|
||||
def test_delete_commit_failure_does_not_enqueue(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
def test_delete_commit_failure_does_not_enqueue(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session) -> None:
|
||||
app = ConversationServiceTestDataFactory.create_app()
|
||||
conversation = ConversationServiceTestDataFactory.create_conversation()
|
||||
conversation.agent_workspace_binding_id = "binding-1"
|
||||
session = MagicMock()
|
||||
session.commit.side_effect = RuntimeError("commit failed")
|
||||
sqlite_session.add(conversation)
|
||||
sqlite_session.flush()
|
||||
rollback_events: list[str] = []
|
||||
event.listen(sqlite_session, "after_rollback", lambda _session: rollback_events.append("rollback"))
|
||||
monkeypatch.setattr(sqlite_session, "commit", MagicMock(side_effect=RuntimeError("commit failed")))
|
||||
monkeypatch.setattr(ConversationService, "get_conversation", MagicMock(return_value=conversation))
|
||||
monkeypatch.setattr(
|
||||
AgentWorkspaceService,
|
||||
@@ -195,14 +199,13 @@ def test_delete_commit_failure_does_not_enqueue(monkeypatch: pytest.MonkeyPatch)
|
||||
monkeypatch.setattr(conversation_service.delete_conversation_related_data, "delay", delete_related)
|
||||
|
||||
with pytest.raises(RuntimeError, match="commit failed"):
|
||||
ConversationService.delete(app, conversation.id, None, session=session)
|
||||
ConversationService.delete(app, conversation.id, None, session=sqlite_session)
|
||||
|
||||
session.rollback.assert_called_once_with()
|
||||
assert rollback_events == ["rollback"]
|
||||
enqueue_collection.assert_not_called()
|
||||
delete_related.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("sqlite_session", [(Conversation,)], indirect=True)
|
||||
class TestConversationServicePagination:
|
||||
"""Test conversation pagination operations."""
|
||||
|
||||
|
||||
@@ -1,8 +1,13 @@
|
||||
from datetime import datetime
|
||||
from unittest.mock import MagicMock, patch
|
||||
from unittest.mock import patch
|
||||
|
||||
from sqlalchemy import event, select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from core.rag.index_processor.constant.built_in_field import BuiltInField
|
||||
from models import Account
|
||||
from models.dataset import Dataset, DatasetMetadata, DatasetMetadataBinding, Document
|
||||
from models.enums import DataSourceType, DocumentCreatedFrom
|
||||
from services.dataset_service import DocumentService
|
||||
from services.entities.knowledge_entities.knowledge_entities import (
|
||||
DocumentMetadataOperation,
|
||||
@@ -18,57 +23,74 @@ def _account() -> Account:
|
||||
return account
|
||||
|
||||
|
||||
def test_create_metadata_flushes_without_committing_caller_session() -> None:
|
||||
session = MagicMock()
|
||||
session.scalar.return_value = None
|
||||
def test_create_metadata_flushes_without_committing_caller_session(sqlite_session: Session) -> None:
|
||||
transaction_events: list[str] = []
|
||||
event.listen(sqlite_session, "after_commit", lambda _session: transaction_events.append("commit"))
|
||||
event.listen(sqlite_session, "after_rollback", lambda _session: transaction_events.append("rollback"))
|
||||
|
||||
metadata = MetadataService.create_metadata(
|
||||
"dataset-1",
|
||||
MetadataArgs(type="string", name="author"),
|
||||
_account(),
|
||||
"tenant-1",
|
||||
session=session,
|
||||
session=sqlite_session,
|
||||
)
|
||||
|
||||
assert metadata.name == "author"
|
||||
session.flush.assert_called_once_with()
|
||||
session.commit.assert_not_called()
|
||||
session.rollback.assert_not_called()
|
||||
assert sqlite_session.get(DatasetMetadata, metadata.id) is metadata
|
||||
assert transaction_events == []
|
||||
|
||||
|
||||
def _document() -> MagicMock:
|
||||
document = MagicMock()
|
||||
document.id = "document-1"
|
||||
document.name = "Document"
|
||||
document.doc_metadata = {}
|
||||
document.data_source_type = "upload_file"
|
||||
document.upload_date = datetime(2026, 1, 1)
|
||||
document.last_update_date = datetime(2026, 1, 2)
|
||||
document.uploader = "global-session uploader"
|
||||
document.get_uploader.return_value = "caller-session uploader"
|
||||
return document
|
||||
def _dataset(*, built_in_field_enabled: bool) -> Dataset:
|
||||
return Dataset(
|
||||
id="dataset-1",
|
||||
tenant_id="tenant-1",
|
||||
name="Dataset",
|
||||
description="",
|
||||
provider="vendor",
|
||||
created_by="account-1",
|
||||
built_in_field_enabled=built_in_field_enabled,
|
||||
)
|
||||
|
||||
|
||||
def test_enable_built_in_field_uses_caller_session_for_uploader() -> None:
|
||||
session = MagicMock()
|
||||
dataset = MagicMock(id="dataset-1", built_in_field_enabled=False)
|
||||
def _document() -> Document:
|
||||
return Document(
|
||||
id="document-1",
|
||||
tenant_id="tenant-1",
|
||||
dataset_id="dataset-1",
|
||||
position=1,
|
||||
data_source_type=DataSourceType.UPLOAD_FILE,
|
||||
batch="batch-1",
|
||||
name="Document",
|
||||
created_from=DocumentCreatedFrom.API,
|
||||
created_by="account-1",
|
||||
created_at=datetime(2026, 1, 1),
|
||||
updated_at=datetime(2026, 1, 2),
|
||||
doc_metadata={},
|
||||
)
|
||||
|
||||
|
||||
def test_enable_built_in_field_uses_caller_session_for_uploader(sqlite_session: Session) -> None:
|
||||
dataset = _dataset(built_in_field_enabled=False)
|
||||
document = _document()
|
||||
sqlite_session.add_all([_account(), dataset, document])
|
||||
sqlite_session.commit()
|
||||
|
||||
with (
|
||||
patch.object(MetadataService, "knowledge_base_metadata_lock_check"),
|
||||
patch.object(DocumentService, "get_working_documents_by_dataset_id", return_value=[document]),
|
||||
patch("services.metadata_service.redis_client.delete"),
|
||||
):
|
||||
MetadataService.enable_built_in_field(dataset, session)
|
||||
MetadataService.enable_built_in_field(dataset, sqlite_session)
|
||||
|
||||
assert document.doc_metadata[BuiltInField.uploader] == "caller-session uploader"
|
||||
document.get_uploader.assert_called_once_with(session=session)
|
||||
assert document.doc_metadata[BuiltInField.uploader] == "User"
|
||||
|
||||
|
||||
def test_update_documents_metadata_uses_caller_session_for_uploader() -> None:
|
||||
session = MagicMock()
|
||||
dataset = MagicMock(id="dataset-1", tenant_id="tenant-1", built_in_field_enabled=True)
|
||||
def test_update_documents_metadata_uses_caller_session_for_uploader(sqlite_session: Session) -> None:
|
||||
dataset = _dataset(built_in_field_enabled=True)
|
||||
document = _document()
|
||||
sqlite_session.add_all([_account(), dataset, document])
|
||||
sqlite_session.commit()
|
||||
metadata_args = MetadataOperationData(
|
||||
operation_data=[
|
||||
DocumentMetadataOperation(document_id=document.id, metadata_list=[], partial_update=False),
|
||||
@@ -85,23 +107,38 @@ def test_update_documents_metadata_uses_caller_session_for_uploader() -> None:
|
||||
metadata_args,
|
||||
_account(),
|
||||
"tenant-1",
|
||||
session=session,
|
||||
session=sqlite_session,
|
||||
)
|
||||
|
||||
assert document.doc_metadata[BuiltInField.uploader] == "caller-session uploader"
|
||||
document.get_uploader.assert_called_once_with(session=session)
|
||||
assert document.doc_metadata[BuiltInField.uploader] == "User"
|
||||
|
||||
|
||||
def test_get_dataset_metadatas_uses_caller_session() -> None:
|
||||
session = MagicMock()
|
||||
session.scalar.return_value = 2
|
||||
dataset = MagicMock(id="dataset-1", built_in_field_enabled=False)
|
||||
dataset.get_doc_metadata.return_value = [{"id": "metadata-1", "name": "author", "type": "string"}]
|
||||
def test_get_dataset_metadatas_uses_caller_session(monkeypatch, sqlite_session: Session) -> None:
|
||||
dataset = _dataset(built_in_field_enabled=False)
|
||||
sqlite_session.add_all(
|
||||
[
|
||||
DatasetMetadataBinding(
|
||||
tenant_id="tenant-1",
|
||||
dataset_id="dataset-1",
|
||||
document_id=f"document-{index}",
|
||||
metadata_id="metadata-1",
|
||||
created_by="account-1",
|
||||
)
|
||||
for index in range(2)
|
||||
]
|
||||
)
|
||||
sqlite_session.commit()
|
||||
|
||||
result = MetadataService.get_dataset_metadatas(dataset, session)
|
||||
def get_doc_metadata(_dataset: Dataset, *, session: Session) -> list[dict[str, str]]:
|
||||
assert session is sqlite_session
|
||||
return [{"id": "metadata-1", "name": "author", "type": "string"}]
|
||||
|
||||
monkeypatch.setattr(Dataset, "get_doc_metadata", get_doc_metadata)
|
||||
|
||||
result = MetadataService.get_dataset_metadatas(dataset, sqlite_session)
|
||||
|
||||
assert result == {
|
||||
"doc_metadata": [{"id": "metadata-1", "name": "author", "type": "string", "count": 2}],
|
||||
"built_in_field_enabled": False,
|
||||
}
|
||||
dataset.get_doc_metadata.assert_called_once_with(session=session)
|
||||
assert sqlite_session.scalar(select(DatasetMetadataBinding).limit(1)) is not None
|
||||
|
||||
@@ -16,7 +16,7 @@ from typing import Any, cast
|
||||
from unittest.mock import ANY, MagicMock, patch, sentinel
|
||||
|
||||
import pytest
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy import event, select
|
||||
from sqlalchemy.dialects import postgresql
|
||||
from sqlalchemy.engine import Engine
|
||||
from sqlalchemy.orm import Session, sessionmaker
|
||||
@@ -37,7 +37,6 @@ from graphon.variables import StringVariable
|
||||
from graphon.variables.input_entities import VariableEntityType
|
||||
from libs.datetime_utils import naive_utc_now
|
||||
from models.account import Account
|
||||
from models.agent import WorkflowAgentNodeBinding
|
||||
from models.human_input import HumanInputFormRecipient, RecipientType
|
||||
from models.model import App, AppMode
|
||||
from models.tools import BuiltinToolProvider, WorkflowToolProvider
|
||||
@@ -201,11 +200,6 @@ class TestWorkflowAssociatedDataFactory:
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("sqlite_session")
|
||||
@pytest.mark.parametrize(
|
||||
"sqlite_session",
|
||||
[(Workflow, App, WorkflowToolProvider, HumanInputFormRecipient, WorkflowAgentNodeBinding)],
|
||||
indirect=True,
|
||||
)
|
||||
class TestWorkflowService:
|
||||
"""
|
||||
Comprehensive unit tests for WorkflowService methods.
|
||||
@@ -452,7 +446,7 @@ class TestWorkflowService:
|
||||
assert workflow.updated_by == account.id
|
||||
|
||||
def test_sync_draft_workflow_collaborative_save_preserves_environment_variables_and_locks_row(
|
||||
self, workflow_service: WorkflowService
|
||||
self, workflow_service: WorkflowService, sqlite_session: Session
|
||||
) -> None:
|
||||
"""A collaborative graph-only save locks the draft and keeps server environment values."""
|
||||
app = TestWorkflowAssociatedDataFactory.create_app()
|
||||
@@ -466,9 +460,19 @@ class TestWorkflowService:
|
||||
)
|
||||
stale_variable = remote_variable.model_copy(update={"value": "stale-client-value"})
|
||||
workflow.environment_variables = [remote_variable]
|
||||
sqlite_session.add(workflow)
|
||||
sqlite_session.commit()
|
||||
unique_hash = workflow.unique_hash
|
||||
session = MagicMock(spec=Session)
|
||||
session.scalar.return_value = workflow
|
||||
statements = []
|
||||
commits = []
|
||||
|
||||
@event.listens_for(sqlite_session, "do_orm_execute")
|
||||
def capture_statement(execute_state):
|
||||
statements.append(execute_state.statement)
|
||||
|
||||
@event.listens_for(sqlite_session, "before_commit")
|
||||
def capture_commit(_session):
|
||||
commits.append(True)
|
||||
|
||||
result = workflow_service.sync_draft_workflow(
|
||||
app_model=app,
|
||||
@@ -478,16 +482,15 @@ class TestWorkflowService:
|
||||
account=account,
|
||||
environment_variables=[stale_variable],
|
||||
conversation_variables=[],
|
||||
session=session,
|
||||
session=sqlite_session,
|
||||
preserve_environment_variables=True,
|
||||
commit=False,
|
||||
sync_agent_bindings=False,
|
||||
)
|
||||
|
||||
statement = session.scalar.call_args.args[0]
|
||||
assert "FOR UPDATE" in str(statement.compile(dialect=postgresql.dialect()))
|
||||
assert "FOR UPDATE" in str(statements[0].compile(dialect=postgresql.dialect()))
|
||||
assert result.environment_variables == [remote_variable]
|
||||
session.commit.assert_not_called()
|
||||
assert commits == []
|
||||
|
||||
def test_sync_draft_workflow_merges_environment_patch_with_graph_update(
|
||||
self, workflow_service: WorkflowService, sqlite_session: Session
|
||||
@@ -910,27 +913,31 @@ class TestWorkflowService:
|
||||
assert persisted_workflow.updated_by == account.id
|
||||
|
||||
def test_patch_draft_workflow_environment_variables_locks_draft_row(
|
||||
self, workflow_service: WorkflowService
|
||||
self, workflow_service: WorkflowService, sqlite_session: Session
|
||||
) -> None:
|
||||
"""The merge reads the draft with a row lock before applying a partial update."""
|
||||
app = TestWorkflowAssociatedDataFactory.create_app()
|
||||
account = TestWorkflowAssociatedDataFactory.create_account()
|
||||
workflow = TestWorkflowAssociatedDataFactory.create_workflow()
|
||||
session = MagicMock(spec=Session)
|
||||
session.scalar.return_value = workflow
|
||||
sqlite_session.add(workflow)
|
||||
sqlite_session.commit()
|
||||
statements = []
|
||||
|
||||
@event.listens_for(sqlite_session, "do_orm_execute")
|
||||
def capture_statement(execute_state):
|
||||
statements.append(execute_state.statement)
|
||||
|
||||
workflow_service.patch_draft_workflow_environment_variables(
|
||||
app_model=app,
|
||||
environment_variables=[],
|
||||
deleted_environment_variable_ids=[],
|
||||
account=account,
|
||||
session=session,
|
||||
session=sqlite_session,
|
||||
)
|
||||
|
||||
statement = session.scalar.call_args.args[0]
|
||||
compiled_statement = str(statement.compile(dialect=postgresql.dialect()))
|
||||
compiled_statement = str(statements[0].compile(dialect=postgresql.dialect()))
|
||||
assert "FOR UPDATE" in compiled_statement
|
||||
session.commit.assert_called_once_with()
|
||||
assert sqlite_session.get(Workflow, workflow.id) is workflow
|
||||
|
||||
def test_patch_draft_workflow_environment_variables_rejects_conflicting_ids(
|
||||
self, workflow_service: WorkflowService, sqlite_session: Session
|
||||
@@ -1670,7 +1677,6 @@ class TestWorkflowService:
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("sqlite_session")
|
||||
@pytest.mark.parametrize("sqlite_session", [(BuiltinToolProvider,)], indirect=True)
|
||||
class TestWorkflowServiceCredentialValidation:
|
||||
"""
|
||||
Tests for the private credential-validation helpers on WorkflowService.
|
||||
@@ -1779,7 +1785,7 @@ class TestWorkflowServiceCredentialValidation:
|
||||
mock_llm.assert_called_once_with("tenant-1", "openai", "gpt-4")
|
||||
|
||||
def test_validate_workflow_credentials_should_use_llm_environment_variable_model(
|
||||
self, service: WorkflowService
|
||||
self, service: WorkflowService, sqlite_session: Session
|
||||
) -> None:
|
||||
workflow = self._make_workflow(
|
||||
[
|
||||
@@ -1814,7 +1820,7 @@ class TestWorkflowServiceCredentialValidation:
|
||||
patch.object(service, "_validate_llm_model_config") as validate_model,
|
||||
patch.object(service, "_validate_load_balancing_credentials") as validate_load_balancing,
|
||||
):
|
||||
service._validate_workflow_credentials(workflow, session=MagicMock())
|
||||
service._validate_workflow_credentials(workflow, session=sqlite_session)
|
||||
|
||||
validate_model.assert_called_once_with("tenant-1", "new-provider", "new-model")
|
||||
validated_node_data = validate_load_balancing.call_args.args[1]
|
||||
@@ -1826,7 +1832,7 @@ class TestWorkflowServiceCredentialValidation:
|
||||
}
|
||||
|
||||
def test_validate_workflow_credentials_should_reject_llm_environment_variable_mode_mismatch(
|
||||
self, service: WorkflowService
|
||||
self, service: WorkflowService, sqlite_session: Session
|
||||
) -> None:
|
||||
workflow = self._make_workflow(
|
||||
[
|
||||
@@ -1848,7 +1854,7 @@ class TestWorkflowServiceCredentialValidation:
|
||||
]
|
||||
|
||||
with pytest.raises(ValueError, match="uses mode 'completion'.*uses mode 'chat'"):
|
||||
service._validate_workflow_credentials(workflow, session=MagicMock())
|
||||
service._validate_workflow_credentials(workflow, session=sqlite_session)
|
||||
|
||||
def test_validate_workflow_credentials_should_raise_for_llm_node_missing_model(
|
||||
self, service: WorkflowService, sqlite_session: Session
|
||||
@@ -2964,7 +2970,6 @@ class TestWorkflowServiceHumanInputOperations:
|
||||
},
|
||||
)
|
||||
|
||||
@pytest.mark.parametrize("sqlite_session", [(Workflow,)], indirect=True)
|
||||
def test_get_human_input_form_preview_should_raise_if_workflow_not_init(
|
||||
self, service: WorkflowService, sqlite_session: Session
|
||||
) -> None:
|
||||
@@ -2977,7 +2982,6 @@ class TestWorkflowServiceHumanInputOperations:
|
||||
session=sqlite_session,
|
||||
)
|
||||
|
||||
@pytest.mark.parametrize("sqlite_session", [(Workflow,)], indirect=True)
|
||||
def test_get_human_input_form_preview_should_raise_if_wrong_node_type(
|
||||
self, service: WorkflowService, sqlite_session: Session
|
||||
) -> None:
|
||||
@@ -2993,7 +2997,6 @@ class TestWorkflowServiceHumanInputOperations:
|
||||
session=sqlite_session,
|
||||
)
|
||||
|
||||
@pytest.mark.parametrize("sqlite_session", [(Workflow,)], indirect=True)
|
||||
def test_get_human_input_form_preview_success(self, service: WorkflowService, sqlite_session: Session) -> None:
|
||||
app_model = TestWorkflowAssociatedDataFactory.create_app(app_id="app-1", tenant_id="tenant-1")
|
||||
account = TestWorkflowAssociatedDataFactory.create_account(account_id="user-1")
|
||||
@@ -3017,7 +3020,6 @@ class TestWorkflowServiceHumanInputOperations:
|
||||
mock_node.render_form_content_before_submission.assert_called_once()
|
||||
mock_required_cls.return_value.model_dump.assert_called_once()
|
||||
|
||||
@pytest.mark.parametrize("sqlite_session", [(Workflow,)], indirect=True)
|
||||
def test_submit_human_input_form_preview_success(self, service: WorkflowService, sqlite_session: Session) -> None:
|
||||
app_model = TestWorkflowAssociatedDataFactory.create_app(app_id="app-1", tenant_id="tenant-1")
|
||||
account = TestWorkflowAssociatedDataFactory.create_account(account_id="user-1")
|
||||
@@ -3058,7 +3060,6 @@ class TestWorkflowServiceHumanInputOperations:
|
||||
assert result["__rendered_content"] == "Ticket: val1"
|
||||
mock_saver_cls.return_value.save.assert_called_once()
|
||||
|
||||
@pytest.mark.parametrize("sqlite_session", [(Workflow,)], indirect=True)
|
||||
def test_test_human_input_delivery_success(self, service: WorkflowService, sqlite_session: Session) -> None:
|
||||
draft = self._create_human_input_workflow()
|
||||
service.get_draft_workflow = MagicMock(return_value=draft)
|
||||
@@ -3082,7 +3083,6 @@ class TestWorkflowServiceHumanInputOperations:
|
||||
)
|
||||
mock_test_srv.return_value.send_test.assert_called_once()
|
||||
|
||||
@pytest.mark.parametrize("sqlite_session", [(Workflow,)], indirect=True)
|
||||
def test_test_human_input_delivery_failure_cases(self, service: WorkflowService, sqlite_session: Session) -> None:
|
||||
draft = self._create_human_input_workflow()
|
||||
service.get_draft_workflow = MagicMock(return_value=draft)
|
||||
@@ -3100,7 +3100,6 @@ class TestWorkflowServiceHumanInputOperations:
|
||||
session=sqlite_session,
|
||||
)
|
||||
|
||||
@pytest.mark.parametrize("sqlite_session", [(HumanInputFormRecipient,)], indirect=True)
|
||||
def test_load_email_recipients_parsing_failure(self, service: WorkflowService, sqlite_session: Session) -> None:
|
||||
"""Malformed persisted recipient payloads are skipped instead of aborting delivery tests."""
|
||||
recipient = HumanInputFormRecipient(
|
||||
|
||||
@@ -6,6 +6,8 @@ from typing import Any, cast
|
||||
from unittest.mock import MagicMock, call
|
||||
|
||||
import pytest
|
||||
from sqlalchemy import event, select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from core.app.app_config.entities import (
|
||||
AdvancedChatMessageEntity,
|
||||
@@ -20,7 +22,8 @@ from core.app.app_config.entities import (
|
||||
from core.helper import encrypter
|
||||
from core.prompt.utils.prompt_template_parser import PromptTemplateParser
|
||||
from models.api_based_extension import APIBasedExtension, APIBasedExtensionPoint
|
||||
from models.model import Account, App, AppMode, AppModelConfig
|
||||
from models.model import Account, App, AppMode, AppModelConfig, IconType
|
||||
from models.workflow import Workflow, WorkflowType
|
||||
from services.workflow import workflow_converter as converter_module
|
||||
from services.workflow.workflow_converter import WorkflowConverter
|
||||
|
||||
@@ -88,7 +91,9 @@ def test__convert_to_start_node(default_variables: list[VariableEntity]) -> None
|
||||
assert result["data"]["variables"][0]["variable"] == "text_input"
|
||||
|
||||
|
||||
def test__convert_to_http_request_node_for_chatbot(default_variables: list[VariableEntity]) -> None:
|
||||
def test__convert_to_http_request_node_for_chatbot(
|
||||
default_variables: list[VariableEntity], unbound_session: Session
|
||||
) -> None:
|
||||
app_model = MagicMock()
|
||||
app_model.id = "app_id"
|
||||
app_model.tenant_id = "tenant_id"
|
||||
@@ -118,7 +123,7 @@ def test__convert_to_http_request_node_for_chatbot(default_variables: list[Varia
|
||||
app_model=app_model,
|
||||
variables=default_variables,
|
||||
external_data_variables=external_data_variables,
|
||||
session=MagicMock(),
|
||||
session=unbound_session,
|
||||
)
|
||||
|
||||
assert len(nodes) == 2
|
||||
@@ -131,7 +136,9 @@ def test__convert_to_http_request_node_for_chatbot(default_variables: list[Varia
|
||||
assert mapping == {"external_variable": "code_1"}
|
||||
|
||||
|
||||
def test__convert_to_http_request_node_for_workflow_app(default_variables: list[VariableEntity]) -> None:
|
||||
def test__convert_to_http_request_node_for_workflow_app(
|
||||
default_variables: list[VariableEntity], unbound_session: Session
|
||||
) -> None:
|
||||
app_model = MagicMock()
|
||||
app_model.id = "app_id"
|
||||
app_model.tenant_id = "tenant_id"
|
||||
@@ -161,7 +168,7 @@ def test__convert_to_http_request_node_for_workflow_app(default_variables: list[
|
||||
app_model=app_model,
|
||||
variables=default_variables,
|
||||
external_data_variables=external_data_variables,
|
||||
session=MagicMock(),
|
||||
session=unbound_session,
|
||||
)
|
||||
|
||||
body = json.loads(nodes[0]["data"]["body"]["data"])
|
||||
@@ -355,9 +362,10 @@ def test__convert_to_answer_node() -> None:
|
||||
assert node["data"]["type"] == BuiltinNodeTypes.ANSWER
|
||||
|
||||
|
||||
def test_convert_to_workflow_should_raise_when_app_model_config_is_missing(converter: WorkflowConverter) -> None:
|
||||
def test_convert_to_workflow_should_raise_when_app_model_config_is_missing(
|
||||
converter: WorkflowConverter, unbound_session: Session
|
||||
) -> None:
|
||||
app_model = _app_model(app_model_config_id=None)
|
||||
session = MagicMock()
|
||||
|
||||
with pytest.raises(ValueError, match="App model config is required"):
|
||||
converter.convert_to_workflow(
|
||||
@@ -367,10 +375,10 @@ def test_convert_to_workflow_should_raise_when_app_model_config_is_missing(conve
|
||||
icon_type="emoji",
|
||||
icon="robot",
|
||||
icon_background="#fff",
|
||||
session=session,
|
||||
session=unbound_session,
|
||||
)
|
||||
|
||||
session.get.assert_not_called()
|
||||
assert not unbound_session.in_transaction()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
@@ -385,23 +393,26 @@ def test_convert_to_workflow_should_create_new_app_with_fallback_fields(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
source_mode: AppMode,
|
||||
expected_mode: AppMode,
|
||||
sqlite_session: Session,
|
||||
) -> None:
|
||||
class FakeApp:
|
||||
def __init__(self) -> None:
|
||||
self.id = "new-app-id"
|
||||
|
||||
workflow = SimpleNamespace(app_id=None)
|
||||
monkeypatch.setattr(converter, "convert_app_model_config_to_workflow", MagicMock(return_value=workflow))
|
||||
monkeypatch.setattr(converter_module, "App", FakeApp)
|
||||
|
||||
app_model_config = _app_model_config(id="config-1")
|
||||
phase_events: list[str] = []
|
||||
db_session = SimpleNamespace(
|
||||
add=MagicMock(),
|
||||
flush=MagicMock(),
|
||||
commit=MagicMock(side_effect=lambda: phase_events.append("commit")),
|
||||
get=MagicMock(return_value=app_model_config),
|
||||
app_model_config = AppModelConfig(app_id="source-app")
|
||||
app_model_config.id = "config-1"
|
||||
sqlite_session.add(app_model_config)
|
||||
sqlite_session.flush()
|
||||
workflow = Workflow(
|
||||
tenant_id="tenant-1",
|
||||
app_id="source-app",
|
||||
type=WorkflowType.WORKFLOW,
|
||||
version=Workflow.VERSION_DRAFT,
|
||||
graph="{}",
|
||||
features="{}",
|
||||
created_by="account-1",
|
||||
environment_variables=[],
|
||||
conversation_variables=[],
|
||||
)
|
||||
monkeypatch.setattr(converter, "convert_app_model_config_to_workflow", MagicMock(return_value=workflow))
|
||||
phase_events: list[str] = []
|
||||
event.listen(sqlite_session, "after_commit", lambda _session: phase_events.append("commit"))
|
||||
|
||||
send_mock = MagicMock(side_effect=lambda *_args, **_kwargs: phase_events.append("signal"))
|
||||
monkeypatch.setattr(converter_module.app_was_created, "send", send_mock)
|
||||
@@ -409,9 +420,10 @@ def test_convert_to_workflow_should_create_new_app_with_fallback_fields(
|
||||
account = _account(id="account-1")
|
||||
app_model = _app_model(
|
||||
tenant_id="tenant-1",
|
||||
id="source-app",
|
||||
name="Source App",
|
||||
mode=source_mode,
|
||||
icon_type="emoji",
|
||||
icon_type=IconType.EMOJI,
|
||||
icon="sparkles",
|
||||
icon_background="#123456",
|
||||
enable_site=True,
|
||||
@@ -429,27 +441,25 @@ def test_convert_to_workflow_should_create_new_app_with_fallback_fields(
|
||||
icon_type="",
|
||||
icon="",
|
||||
icon_background="",
|
||||
session=db_session,
|
||||
session=sqlite_session,
|
||||
)
|
||||
|
||||
assert new_app.name == "Source App(workflow)"
|
||||
assert new_app.mode == expected_mode
|
||||
assert new_app.icon_type == "emoji"
|
||||
assert new_app.icon_type == IconType.EMOJI
|
||||
assert new_app.icon == "sparkles"
|
||||
assert new_app.icon_background == "#123456"
|
||||
assert new_app.created_by == "account-1"
|
||||
assert workflow.app_id == "new-app-id"
|
||||
db_session.add.assert_called_once()
|
||||
db_session.flush.assert_called_once()
|
||||
assert workflow.app_id == new_app.id
|
||||
assert sqlite_session.get(App, new_app.id) is new_app
|
||||
assert phase_events == ["commit", "signal", "commit"]
|
||||
assert db_session.commit.call_count == 2
|
||||
db_session.get.assert_called_once_with(AppModelConfig, "config-1")
|
||||
send_mock.assert_called_once_with(new_app, account=account, session=db_session)
|
||||
send_mock.assert_called_once_with(new_app, account=account, session=sqlite_session)
|
||||
|
||||
|
||||
def test_convert_app_model_config_to_workflow_should_build_advanced_chat_graph_and_features(
|
||||
converter: WorkflowConverter,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
sqlite_session: Session,
|
||||
) -> None:
|
||||
app_model = _app_model(id="app-1", tenant_id="tenant-1", mode=AppMode.CHAT)
|
||||
app_config = SimpleNamespace(
|
||||
@@ -471,12 +481,6 @@ def test_convert_app_model_config_to_workflow_should_build_advanced_chat_graph_a
|
||||
},
|
||||
)
|
||||
|
||||
class FakeWorkflow:
|
||||
VERSION_DRAFT = "draft"
|
||||
|
||||
def __init__(self, **kwargs: Any) -> None:
|
||||
self.__dict__.update(kwargs)
|
||||
|
||||
monkeypatch.setattr(converter, "_get_new_app_mode", MagicMock(return_value=AppMode.ADVANCED_CHAT))
|
||||
monkeypatch.setattr(converter, "_convert_to_app_config", MagicMock(return_value=app_config))
|
||||
monkeypatch.setattr(
|
||||
@@ -513,15 +517,11 @@ def test_convert_app_model_config_to_workflow_should_build_advanced_chat_graph_a
|
||||
"_convert_to_answer_node",
|
||||
MagicMock(return_value={"id": "answer", "position": None, "data": {"type": BuiltinNodeTypes.ANSWER}}),
|
||||
)
|
||||
monkeypatch.setattr(converter_module, "Workflow", FakeWorkflow)
|
||||
|
||||
db_session = SimpleNamespace(add=MagicMock(), commit=MagicMock())
|
||||
|
||||
workflow = converter.convert_app_model_config_to_workflow(
|
||||
app_model=app_model,
|
||||
app_model_config=_app_model_config(id="cfg"),
|
||||
account_id="account-1",
|
||||
session=db_session,
|
||||
session=sqlite_session,
|
||||
)
|
||||
|
||||
graph = json.loads(workflow.graph)
|
||||
@@ -531,13 +531,13 @@ def test_convert_app_model_config_to_workflow_should_build_advanced_chat_graph_a
|
||||
features = json.loads(workflow.features)
|
||||
assert "opening_statement" in features
|
||||
assert "retriever_resource" in features
|
||||
db_session.add.assert_called_once()
|
||||
db_session.commit.assert_called_once()
|
||||
assert sqlite_session.scalar(select(Workflow).where(Workflow.id == workflow.id)) is workflow
|
||||
|
||||
|
||||
def test_convert_app_model_config_to_workflow_should_build_workflow_mode_with_end_node(
|
||||
converter: WorkflowConverter,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
sqlite_session: Session,
|
||||
) -> None:
|
||||
app_model = _app_model(id="app-1", tenant_id="tenant-1", mode=AppMode.COMPLETION)
|
||||
app_config = SimpleNamespace(
|
||||
@@ -554,12 +554,6 @@ def test_convert_app_model_config_to_workflow_should_build_workflow_mode_with_en
|
||||
},
|
||||
)
|
||||
|
||||
class FakeWorkflow:
|
||||
VERSION_DRAFT = "draft"
|
||||
|
||||
def __init__(self, **kwargs: Any) -> None:
|
||||
self.__dict__.update(kwargs)
|
||||
|
||||
monkeypatch.setattr(converter, "_get_new_app_mode", MagicMock(return_value=AppMode.WORKFLOW))
|
||||
monkeypatch.setattr(converter, "_convert_to_app_config", MagicMock(return_value=app_config))
|
||||
monkeypatch.setattr(
|
||||
@@ -580,15 +574,11 @@ def test_convert_app_model_config_to_workflow_should_build_workflow_mode_with_en
|
||||
"_convert_to_end_node",
|
||||
MagicMock(return_value={"id": "end", "position": None, "data": {"type": BuiltinNodeTypes.END}}),
|
||||
)
|
||||
monkeypatch.setattr(converter_module, "Workflow", FakeWorkflow)
|
||||
|
||||
db_session = SimpleNamespace(add=MagicMock(), commit=MagicMock())
|
||||
|
||||
workflow = converter.convert_app_model_config_to_workflow(
|
||||
app_model=app_model,
|
||||
app_model_config=_app_model_config(id="cfg"),
|
||||
account_id="account-1",
|
||||
session=db_session,
|
||||
session=sqlite_session,
|
||||
)
|
||||
|
||||
graph = json.loads(workflow.graph)
|
||||
@@ -597,11 +587,13 @@ def test_convert_app_model_config_to_workflow_should_build_workflow_mode_with_en
|
||||
|
||||
features = json.loads(workflow.features)
|
||||
assert set(features.keys()) == {"text_to_speech", "file_upload", "sensitive_word_avoidance"}
|
||||
assert sqlite_session.scalar(select(Workflow).where(Workflow.id == workflow.id)) is workflow
|
||||
|
||||
|
||||
def test_convert_to_app_config_should_route_to_correct_manager(
|
||||
converter: WorkflowConverter,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
unbound_session: Session,
|
||||
) -> None:
|
||||
agent_result = SimpleNamespace(kind="agent")
|
||||
chat_result = SimpleNamespace(kind="chat")
|
||||
@@ -614,7 +606,6 @@ def test_convert_to_app_config_should_route_to_correct_manager(
|
||||
monkeypatch.setattr(converter_module.ChatAppConfigManager, "get_app_config", chat_get_app_config)
|
||||
monkeypatch.setattr(converter_module.CompletionAppConfigManager, "get_app_config", completion_get_app_config)
|
||||
monkeypatch.setattr(converter_module, "load_annotation_reply_config", load_annotation_reply)
|
||||
session = MagicMock()
|
||||
agent_mode_app = _app_model(mode=AppMode.AGENT_CHAT, is_agent_with_session=MagicMock(return_value=False))
|
||||
agent_flag_app = _app_model(mode=AppMode.CHAT, is_agent_with_session=MagicMock(return_value=True))
|
||||
chat_app = _app_model(mode=AppMode.CHAT, is_agent_with_session=MagicMock(return_value=False))
|
||||
@@ -627,31 +618,36 @@ def test_convert_to_app_config_should_route_to_correct_manager(
|
||||
from_agent_mode = converter._convert_to_app_config(
|
||||
app_model=agent_mode_app,
|
||||
app_model_config=agent_mode_config,
|
||||
session=session,
|
||||
session=unbound_session,
|
||||
)
|
||||
from_agent_flag = converter._convert_to_app_config(
|
||||
app_model=agent_flag_app,
|
||||
app_model_config=agent_flag_config,
|
||||
session=session,
|
||||
session=unbound_session,
|
||||
)
|
||||
from_chat_mode = converter._convert_to_app_config(
|
||||
app_model=chat_app,
|
||||
app_model_config=chat_config,
|
||||
session=session,
|
||||
session=unbound_session,
|
||||
)
|
||||
from_completion_mode = converter._convert_to_app_config(
|
||||
app_model=completion_app,
|
||||
app_model_config=completion_config,
|
||||
session=session,
|
||||
session=unbound_session,
|
||||
)
|
||||
|
||||
assert from_agent_mode is agent_result
|
||||
assert from_agent_flag is agent_result
|
||||
assert from_chat_mode is chat_result
|
||||
assert from_completion_mode is completion_result
|
||||
agent_flag_app.is_agent_with_session.assert_called_once_with(session=session)
|
||||
agent_flag_app.is_agent_with_session.assert_called_once_with(session=unbound_session)
|
||||
load_annotation_reply.assert_has_calls(
|
||||
[call(session, "app-1"), call(session, "app-2"), call(session, "app-3"), call(session, "app-4")]
|
||||
[
|
||||
call(unbound_session, "app-1"),
|
||||
call(unbound_session, "app-2"),
|
||||
call(unbound_session, "app-3"),
|
||||
call(unbound_session, "app-4"),
|
||||
]
|
||||
)
|
||||
assert all(
|
||||
manager_call.kwargs["annotation_reply"] == {"enabled": False}
|
||||
@@ -660,18 +656,20 @@ def test_convert_to_app_config_should_route_to_correct_manager(
|
||||
)
|
||||
|
||||
|
||||
def test_convert_to_app_config_should_raise_for_invalid_app_mode(converter: WorkflowConverter) -> None:
|
||||
def test_convert_to_app_config_should_raise_for_invalid_app_mode(
|
||||
converter: WorkflowConverter, unbound_session: Session
|
||||
) -> None:
|
||||
app_model = _app_model(mode=AppMode.WORKFLOW, is_agent_with_session=MagicMock(return_value=False))
|
||||
session = MagicMock()
|
||||
|
||||
with pytest.raises(ValueError, match="Invalid app mode"):
|
||||
converter._convert_to_app_config(
|
||||
app_model=app_model, app_model_config=_app_model_config(id="cfg"), session=session
|
||||
app_model=app_model, app_model_config=_app_model_config(id="cfg"), session=unbound_session
|
||||
)
|
||||
|
||||
|
||||
def test_convert_to_http_request_node_should_skip_non_api_and_missing_extension_id(
|
||||
converter: WorkflowConverter,
|
||||
unbound_session: Session,
|
||||
) -> None:
|
||||
app_model = _app_model(id="app-1", tenant_id="tenant-1", mode=AppMode.CHAT)
|
||||
external_data_variables = [
|
||||
@@ -683,7 +681,7 @@ def test_convert_to_http_request_node_should_skip_non_api_and_missing_extension_
|
||||
app_model=app_model,
|
||||
variables=[],
|
||||
external_data_variables=external_data_variables,
|
||||
session=MagicMock(),
|
||||
session=unbound_session,
|
||||
)
|
||||
|
||||
assert nodes == []
|
||||
|
||||
@@ -720,6 +720,7 @@ def test_start_buffering_should_set_done_event_when_subscription_raises() -> Non
|
||||
|
||||
def test_build_workflow_event_stream_should_emit_ping_and_terminal_snapshot_event(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
unbound_session_factory: sessionmaker[Session],
|
||||
) -> None:
|
||||
# Arrange
|
||||
workflow_run = _build_workflow_run_additional(status=WorkflowExecutionStatus.PAUSED)
|
||||
@@ -764,7 +765,7 @@ def test_build_workflow_event_stream_should_emit_ping_and_terminal_snapshot_even
|
||||
"_build_snapshot_events",
|
||||
MagicMock(return_value=[{"event": StreamEvent.WORKFLOW_FINISHED, "task_id": "task-1"}]),
|
||||
)
|
||||
session_maker = MagicMock()
|
||||
session_maker = unbound_session_factory
|
||||
|
||||
# Act
|
||||
events = list(
|
||||
@@ -817,6 +818,7 @@ def test_build_workflow_event_stream_should_emit_ping_and_terminal_snapshot_even
|
||||
)
|
||||
def test_build_advanced_chat_snapshot_requires_conversation_context(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
unbound_session_factory: sessionmaker[Session],
|
||||
resumption_context: WorkflowResumptionContext | None,
|
||||
expected_error: str,
|
||||
) -> None:
|
||||
@@ -848,7 +850,7 @@ def test_build_advanced_chat_snapshot_requires_conversation_context(
|
||||
workflow_run=workflow_run,
|
||||
tenant_id="tenant-1",
|
||||
app_id="app-1",
|
||||
session_maker=MagicMock(),
|
||||
session_maker=unbound_session_factory,
|
||||
)
|
||||
conversation_lookup.assert_not_called()
|
||||
app_lookup.assert_not_called()
|
||||
@@ -856,6 +858,7 @@ def test_build_advanced_chat_snapshot_requires_conversation_context(
|
||||
|
||||
def test_build_non_suspended_advanced_chat_snapshot_uses_app_scoped_fallback(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
unbound_session_factory: sessionmaker[Session],
|
||||
) -> None:
|
||||
# Arrange
|
||||
workflow_run = _build_workflow_run_additional(status=WorkflowExecutionStatus.RUNNING)
|
||||
@@ -877,7 +880,7 @@ def test_build_non_suspended_advanced_chat_snapshot_uses_app_scoped_fallback(
|
||||
conversation_lookup,
|
||||
)
|
||||
monkeypatch.setattr(service_module, "_get_message_context_by_app", app_lookup)
|
||||
session_maker = MagicMock()
|
||||
session_maker = unbound_session_factory
|
||||
|
||||
# Act
|
||||
event_stream = build_workflow_event_stream(
|
||||
|
||||
Reference in New Issue
Block a user