test: migrate controller sessions to SQLite (#40083)

Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
This commit is contained in:
Asuka Minato
2026-08-11 08:29:40 +00:00
committed by GitHub
co-authored by autofix-ci[bot]
parent e3812cf72b
commit 08008f6305
23 changed files with 637 additions and 496 deletions
@@ -1,77 +1,132 @@
import uuid
from types import SimpleNamespace
from unittest.mock import MagicMock
import pytest
from sqlalchemy.orm import Session
from sqlalchemy import event
from sqlalchemy.engine import Engine
from sqlalchemy.orm import Session, scoped_session, sessionmaker
from commands import system as system_commands
from models.account import Account, Tenant, TenantAccountJoin, TenantAccountRole
from models.enums import CustomizeTokenStrategy
from models.model import App, AppMode, IconType, Site
def test_fix_app_site_missing_passes_loaded_session_to_signal(monkeypatch: pytest.MonkeyPatch) -> None:
account = object()
tenant = MagicMock()
tenant.get_accounts.return_value = [account]
app = SimpleNamespace(id="app-1", tenant_id="tenant-1")
def _persist_missing_site_owner(session: Session) -> tuple[Account, App]:
"""Persist an app without a Site and its complete tenant owner chain."""
tenant = Tenant(name="Command workspace")
account = Account(name="Owner", email=f"owner-{uuid.uuid4()}@example.com")
membership = TenantAccountJoin(
tenant_id=tenant.id,
account_id=account.id,
current=True,
role=TenantAccountRole.OWNER,
)
app = App(
id=str(uuid.uuid4()),
tenant_id=tenant.id,
name="Missing Site App",
mode=AppMode.CHAT,
icon_type=IconType.EMOJI,
icon="chat",
icon_background="#FFFFFF",
enable_site=True,
enable_api=False,
created_by=account.id,
)
session.add_all([tenant, account, membership, app])
session.commit()
return account, app
session = Session()
def _site_for(app: App) -> Site:
return Site(
app_id=app.id,
title=app.name,
default_language="en-US",
customize_token_strategy=CustomizeTokenStrategy.UUID,
code=f"site-{app.id}",
)
def _bind_command_database(
monkeypatch: pytest.MonkeyPatch,
sqlite_engine: Engine,
sqlite_session_factory: sessionmaker[Session],
) -> tuple[scoped_session[Session], Session]:
"""Expose a callable real scoped session through the Flask-SQLAlchemy shape."""
command_sessions = scoped_session(sqlite_session_factory)
command_session = command_sessions()
monkeypatch.setattr(
system_commands,
"db",
SimpleNamespace(engine=sqlite_engine, session=command_sessions),
)
return command_sessions, command_session
def test_fix_app_site_missing_passes_loaded_session_to_signal(
monkeypatch: pytest.MonkeyPatch,
sqlite_engine: Engine,
sqlite_session: Session,
sqlite_session_factory: sessionmaker[Session],
) -> None:
account, app = _persist_missing_site_owner(sqlite_session)
command_sessions, command_session = _bind_command_database(monkeypatch, sqlite_engine, sqlite_session_factory)
phase_events: list[str] = []
scalar = MagicMock(return_value=app)
get = MagicMock(return_value=tenant)
commit = MagicMock(side_effect=lambda: phase_events.append("commit"))
monkeypatch.setattr(session, "scalar", scalar)
monkeypatch.setattr(session, "get", get)
monkeypatch.setattr(session, "commit", commit)
event.listen(command_session, "after_commit", lambda _session: phase_events.append("commit"))
scoped_session = MagicMock(return_value=session)
scoped_session.scalar.return_value = app
def create_site(sender: App, *, account: Account, session: Session) -> None:
phase_events.append("signal")
assert sender.id == app.id
assert account.id == account_id
assert session is command_session
session.add(_site_for(sender))
connection = MagicMock()
connection.execute.side_effect = [[SimpleNamespace(id=app.id)], []]
engine = MagicMock()
engine.begin.return_value.__enter__.return_value = connection
monkeypatch.setattr(system_commands, "db", SimpleNamespace(engine=engine, session=scoped_session))
send = MagicMock(side_effect=lambda *_args, **_kwargs: phase_events.append("signal"))
account_id = account.id
send = MagicMock(side_effect=create_site)
monkeypatch.setattr(system_commands.app_was_created, "send", send)
system_commands.fix_app_site_missing.callback()
try:
system_commands.fix_app_site_missing.callback()
finally:
command_sessions.remove()
scoped_session.assert_called_once_with()
scalar.assert_called_once()
get.assert_called_once_with(system_commands.Tenant, app.tenant_id)
tenant.get_accounts.assert_called_once_with(session=session)
send.assert_called_once_with(app, account=account, session=session)
commit.assert_called_once_with()
send.assert_called_once()
assert phase_events == ["signal", "commit"]
assert isinstance(send.call_args.kwargs["session"], Session)
sqlite_session.expire_all()
persisted_site = sqlite_session.query(Site).filter_by(app_id=app.id).one()
assert persisted_site.title == app.name
def test_fix_app_site_missing_rolls_back_when_signal_fails(monkeypatch: pytest.MonkeyPatch) -> None:
account = object()
tenant = MagicMock()
tenant.get_accounts.return_value = [account]
app = SimpleNamespace(id="app-1", tenant_id="tenant-1")
session = MagicMock()
def test_fix_app_site_missing_rolls_back_when_signal_fails(
monkeypatch: pytest.MonkeyPatch,
sqlite_engine: Engine,
sqlite_session: Session,
sqlite_session_factory: sessionmaker[Session],
) -> None:
_account, app = _persist_missing_site_owner(sqlite_session)
command_sessions, command_session = _bind_command_database(monkeypatch, sqlite_engine, sqlite_session_factory)
phase_events: list[str] = []
session.scalar.return_value = app
session.get.return_value = tenant
session.rollback.side_effect = lambda: phase_events.append("rollback")
event.listen(command_session, "after_rollback", lambda _session: phase_events.append("rollback"))
connection = MagicMock()
connection.execute.side_effect = [[SimpleNamespace(id=app.id)], []]
engine = MagicMock()
engine.begin.return_value.__enter__.return_value = connection
monkeypatch.setattr(system_commands, "db", SimpleNamespace(engine=engine, session=MagicMock(return_value=session)))
def fail_signal(*_args, **_kwargs) -> None:
def fail_signal(sender: App, **_kwargs: object) -> None:
phase_events.append("signal")
# Ensure the command's next raw scan terminates while its own transaction
# still exercises the rollback path.
with sqlite_session_factory() as observer:
observer.add(_site_for(sender))
observer.commit()
raise RuntimeError("failed")
monkeypatch.setattr(system_commands.app_was_created, "send", MagicMock(side_effect=fail_signal))
system_commands.fix_app_site_missing.callback()
try:
system_commands.fix_app_site_missing.callback()
finally:
command_sessions.remove()
session.rollback.assert_called_once_with()
session.commit.assert_not_called()
assert phase_events == ["signal", "rollback"]
sqlite_session.expire_all()
assert sqlite_session.get(App, app.id) is not None
@@ -2,9 +2,8 @@
from __future__ import annotations
from unittest.mock import MagicMock
import pytest
from sqlalchemy.orm import Session
from controllers.common.app_access import (
APP_LIST_PERMISSION_KEYS,
@@ -106,28 +105,32 @@ class TestResolveAppAccessFilter:
lambda tenant_id, account_id: whitelist,
)
def test_default_preview_is_unrestricted(self, monkeypatch: pytest.MonkeyPatch):
def test_default_preview_is_unrestricted(self, monkeypatch: pytest.MonkeyPatch, unbound_session: Session):
self._patch_whitelist(monkeypatch, ResourceWhitelistResources(unrestricted=True))
permissions = _permissions(app_default_keys=["app.preview"])
flt = resolve_app_access_filter("tenant-1", "acc-1", session=MagicMock(), permissions=permissions)
flt = resolve_app_access_filter("tenant-1", "acc-1", session=unbound_session, permissions=permissions)
assert flt.accessible_app_ids is None
assert flt.can_manage_own_apps is False
def test_default_preview_overrides_whitelist_restriction(self, monkeypatch: pytest.MonkeyPatch):
def test_default_preview_overrides_whitelist_restriction(
self, monkeypatch: pytest.MonkeyPatch, unbound_session: Session
):
self._patch_whitelist(monkeypatch, ResourceWhitelistResources(unrestricted=False, resource_ids=["app-9"]))
permissions = _permissions(
workspace_keys=["app.full_access", "app.create_and_management"],
)
flt = resolve_app_access_filter("tenant-1", "acc-1", session=MagicMock(), permissions=permissions)
flt = resolve_app_access_filter("tenant-1", "acc-1", session=unbound_session, permissions=permissions)
# Workspace-level preview grant defeats the whitelist restriction.
assert flt.accessible_app_ids is None
assert flt.can_manage_own_apps is True
def test_override_apps_collected_without_default_preview(self, monkeypatch: pytest.MonkeyPatch):
def test_override_apps_collected_without_default_preview(
self, monkeypatch: pytest.MonkeyPatch, unbound_session: Session
):
self._patch_whitelist(monkeypatch, ResourceWhitelistResources(unrestricted=True))
permissions = _permissions(
app_overrides=[
@@ -136,23 +139,23 @@ class TestResolveAppAccessFilter:
],
)
flt = resolve_app_access_filter("tenant-1", "acc-1", session=MagicMock(), permissions=permissions)
flt = resolve_app_access_filter("tenant-1", "acc-1", session=unbound_session, permissions=permissions)
assert flt.accessible_app_ids == {"app-1"}
def test_whitelist_union_with_override_apps(self, monkeypatch: pytest.MonkeyPatch):
def test_whitelist_union_with_override_apps(self, monkeypatch: pytest.MonkeyPatch, unbound_session: Session):
self._patch_whitelist(monkeypatch, ResourceWhitelistResources(unrestricted=False, resource_ids=["app-5"]))
permissions = _permissions(
app_overrides=[ResourcePermissionKeys(resource_id="app-1", permission_keys=["app.acl.preview"])],
)
flt = resolve_app_access_filter("tenant-1", "acc-1", session=MagicMock(), permissions=permissions)
flt = resolve_app_access_filter("tenant-1", "acc-1", session=unbound_session, permissions=permissions)
assert flt.accessible_app_ids == {"app-1", "app-5"}
def test_fetches_permissions_when_not_supplied(self, monkeypatch: pytest.MonkeyPatch):
def test_fetches_permissions_when_not_supplied(self, monkeypatch: pytest.MonkeyPatch, unbound_session: Session):
self._patch_whitelist(monkeypatch, ResourceWhitelistResources(unrestricted=False, resource_ids=[]))
session = MagicMock()
session = unbound_session
captured: dict[str, object] = {}
def get_permissions(tenant_id: str, account_id: str, *, session: object):
@@ -5,6 +5,7 @@ from unittest.mock import ANY, MagicMock, PropertyMock, patch
import pytest
from flask import Flask
from sqlalchemy.orm import Session
from werkzeug.exceptions import Forbidden, NotFound
import services
@@ -162,13 +163,21 @@ def _dataset_detail_object() -> SimpleNamespace:
)
class TestExternalApiTemplateListApi:
class _UsesSQLiteSession:
session: Session
@pytest.fixture(autouse=True)
def _inject_sqlite_session(self, sqlite_session: Session) -> None:
self.session = sqlite_session
class TestExternalApiTemplateListApi(_UsesSQLiteSession):
def test_get_success(self, app: Flask):
api = ExternalApiTemplateListApi()
method = inspect.unwrap(api.get)
api_item = _external_api_object("api-1")
session = MagicMock()
session = self.session
with (
app.test_request_context("/?page=2&limit=1&keyword=vector"),
@@ -207,7 +216,7 @@ class TestExternalApiTemplateListApi:
},
}
created = _external_api_object("api-created")
session = MagicMock()
session = self.session
with (
app.test_request_context("/", json=payload),
@@ -247,7 +256,7 @@ class TestExternalApiTemplateListApi:
patch.object(ExternalDatasetService, "validate_api_list"),
):
with pytest.raises(Forbidden):
method(api, ExternalKnowledgeApiPayload.model_validate(payload), MagicMock(), "tenant-1", current_user)
method(api, ExternalKnowledgeApiPayload.model_validate(payload), self.session, "tenant-1", current_user)
def test_post_duplicate_name(self, app: Flask, current_user: Account):
api = ExternalApiTemplateListApi()
@@ -266,15 +275,15 @@ class TestExternalApiTemplateListApi:
),
):
with pytest.raises(DatasetNameDuplicateError):
method(api, ExternalKnowledgeApiPayload.model_validate(payload), MagicMock(), "tenant-1", current_user)
method(api, ExternalKnowledgeApiPayload.model_validate(payload), self.session, "tenant-1", current_user)
class TestExternalApiTemplateApi:
class TestExternalApiTemplateApi(_UsesSQLiteSession):
def test_get_success_returns_template_contract(self, app: Flask):
api = ExternalApiTemplateApi()
method = inspect.unwrap(api.get)
template = _external_api_object("api-detail")
session = MagicMock()
session = self.session
with (
app.test_request_context("/"),
@@ -306,7 +315,7 @@ class TestExternalApiTemplateApi:
),
):
with pytest.raises(NotFound):
method(api, MagicMock(), "tenant-1", "api-id")
method(api, self.session, "tenant-1", "api-id")
def test_patch_success_uses_validated_payload_and_returns_template(self, app: Flask, current_user: Account):
api = ExternalApiTemplateApi()
@@ -321,7 +330,7 @@ class TestExternalApiTemplateApi:
},
}
updated = _external_api_object("api-updated")
session = MagicMock()
session = self.session
with (
app.test_request_context("/", json=payload),
@@ -362,15 +371,15 @@ class TestExternalApiTemplateApi:
with app.test_request_context("/"):
with pytest.raises(Forbidden):
method(api, MagicMock(), "tenant-1", current_user, "api-id")
method(api, self.session, "tenant-1", current_user, "api-id")
class TestExternalApiUseCheckApi:
class TestExternalApiUseCheckApi(_UsesSQLiteSession):
def test_get_scopes_usage_check_to_current_tenant(self, app: Flask):
api = ExternalApiUseCheckApi()
method = inspect.unwrap(api.get)
session = MagicMock()
session = self.session
with (
app.test_request_context("/"),
@@ -387,7 +396,7 @@ class TestExternalApiUseCheckApi:
mock_use_check.assert_called_once_with("api-id", "tenant-1", session=ANY)
class TestExternalDatasetCreateApi:
class TestExternalDatasetCreateApi(_UsesSQLiteSession):
def test_create_success(self, app: Flask, current_user: Account):
api = ExternalDatasetCreateApi()
method = inspect.unwrap(api.post)
@@ -450,10 +459,12 @@ class TestExternalDatasetCreateApi:
patch.object(type(console_ns), "payload", new_callable=PropertyMock, return_value=payload),
):
with pytest.raises(Forbidden):
method(api, ExternalDatasetCreatePayload.model_validate(payload), MagicMock(), "tenant-1", current_user)
method(
api, ExternalDatasetCreatePayload.model_validate(payload), self.session, "tenant-1", current_user
)
class TestExternalKnowledgeHitTestingApi:
class TestExternalKnowledgeHitTestingApi(_UsesSQLiteSession):
def test_hit_testing_dataset_not_found(self, app: Flask, current_user: Account):
api = ExternalKnowledgeHitTestingApi()
method = inspect.unwrap(api.post)
@@ -467,7 +478,7 @@ class TestExternalKnowledgeHitTestingApi:
),
):
with pytest.raises(NotFound):
method(api, ExternalHitTestingPayload(query="test"), MagicMock(), current_user, "dataset-id")
method(api, ExternalHitTestingPayload(query="test"), self.session, current_user, "dataset-id")
def test_hit_testing_success(self, app: Flask, current_user: Account):
api = ExternalKnowledgeHitTestingApi()
@@ -498,7 +509,7 @@ class TestExternalKnowledgeHitTestingApi:
}
],
}
session = MagicMock()
session = self.session
with (
app.test_request_context("/", json=payload),
@@ -555,8 +566,6 @@ class TestBedrockRetrievalApi:
]
}
session = MagicMock()
with (
app.test_request_context("/", json=payload),
patch.object(type(console_ns), "payload", new_callable=PropertyMock, return_value=payload),
@@ -576,7 +585,7 @@ class TestBedrockRetrievalApi:
assert knowledge_id == "knowledge-base-1"
class TestExternalApiTemplateListApiAdvanced:
class TestExternalApiTemplateListApiAdvanced(_UsesSQLiteSession):
def test_post_duplicate_name_error(self, app: Flask, current_user: Account):
api = ExternalApiTemplateListApi()
method = inspect.unwrap(api.post)
@@ -593,7 +602,7 @@ class TestExternalApiTemplateListApiAdvanced:
),
):
with pytest.raises(DatasetNameDuplicateError):
method(api, ExternalKnowledgeApiPayload.model_validate(payload), MagicMock(), "tenant-1", current_user)
method(api, ExternalKnowledgeApiPayload.model_validate(payload), self.session, "tenant-1", current_user)
def test_get_with_pagination(self, app: Flask):
api = ExternalApiTemplateListApi()
@@ -608,7 +617,7 @@ class TestExternalApiTemplateListApiAdvanced:
return_value=(templates, 25),
) as get_external_knowledge_apis,
):
resp, status = method(api, ExternalApiTemplateListQuery(page=2, limit=3), MagicMock(), "tenant-1")
resp, status = method(api, ExternalApiTemplateListQuery(page=2, limit=3), self.session, "tenant-1")
assert status == 200
assert resp == {
@@ -621,7 +630,7 @@ class TestExternalApiTemplateListApiAdvanced:
get_external_knowledge_apis.assert_called_once_with(2, 3, "tenant-1", None, session=ANY)
class TestExternalDatasetCreateApiAdvanced:
class TestExternalDatasetCreateApiAdvanced(_UsesSQLiteSession):
def test_create_forbidden(self, app: Flask, current_user: Account):
"""Test creating external dataset without permission"""
api = ExternalDatasetCreateApi()
@@ -638,10 +647,12 @@ class TestExternalDatasetCreateApiAdvanced:
with app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload):
with pytest.raises(Forbidden):
method(api, ExternalDatasetCreatePayload.model_validate(payload), MagicMock(), "tenant-1", current_user)
method(
api, ExternalDatasetCreatePayload.model_validate(payload), self.session, "tenant-1", current_user
)
class TestExternalKnowledgeHitTestingApiAdvanced:
class TestExternalKnowledgeHitTestingApiAdvanced(_UsesSQLiteSession):
def test_hit_testing_dataset_not_found(self, app: Flask, current_user: Account):
"""Test hit testing on non-existent dataset"""
api = ExternalKnowledgeHitTestingApi()
@@ -661,7 +672,7 @@ class TestExternalKnowledgeHitTestingApiAdvanced:
),
):
with pytest.raises(NotFound):
method(api, ExternalHitTestingPayload.model_validate(payload), MagicMock(), current_user, "ds-1")
method(api, ExternalHitTestingPayload.model_validate(payload), self.session, current_user, "ds-1")
def test_hit_testing_with_custom_retrieval_model(self, app: Flask, current_user: Account):
api = ExternalKnowledgeHitTestingApi()
@@ -673,7 +684,7 @@ class TestExternalKnowledgeHitTestingApiAdvanced:
"external_retrieval_model": {"type": "bm25"},
"metadata_filtering_conditions": {"status": "active"},
}
session = MagicMock()
session = self.session
with (
app.test_request_context("/", json=payload),
@@ -65,9 +65,8 @@ class TestGetRagPipeline:
assert result is pipeline
get_pipeline_by_id.assert_called_once_with("pipeline-1", "tenant-1", session=session_factory.return_value)
def test_load_rag_pipeline_uses_provided_session(self, mocker: MockerFixture):
def test_load_rag_pipeline_uses_provided_session(self, mocker: MockerFixture, sqlite_session: Session):
pipeline = Pipeline(tenant_id="tenant-id", name="Test Pipeline")
session = Mock(spec=Session)
mocker.patch(
"controllers.console.datasets.wraps.current_account_with_tenant",
@@ -78,10 +77,10 @@ class TestGetRagPipeline:
return_value=pipeline,
)
result = load_rag_pipeline(session, "pipeline-1")
result = load_rag_pipeline(sqlite_session, "pipeline-1")
assert result is pipeline
get_pipeline_by_id.assert_called_once_with("pipeline-1", "tenant-1", session=session)
get_pipeline_by_id.assert_called_once_with("pipeline-1", "tenant-1", session=sqlite_session)
def test_pipeline_id_removed_from_kwargs(self, mocker: MockerFixture):
pipeline = Pipeline(tenant_id="tenant-id", name="Test Pipeline")
@@ -1,42 +1,31 @@
from collections.abc import Iterator
from datetime import datetime
from types import SimpleNamespace
from unittest.mock import MagicMock, call
from typing import NamedTuple
from uuid import uuid4
import pytest
from flask import Flask
from pydantic import ValidationError
from sqlalchemy.engine import Engine
from sqlalchemy.orm import Session, scoped_session, sessionmaker
from sqlalchemy.orm import Session, sessionmaker
import controllers.console.explore.banner as banner_module
from models.base import TypeBase
from models.enums import BannerStatus
from models.model import ExporleBanner
from repositories.explore_banner_query_repository import ExploreBannerQueryRepository
from services.explore_banner_query_service import ExploreBannerQueryService, ExploreBannerRecord
@pytest.fixture
def banner_session(sqlite_engine: Engine) -> Iterator[scoped_session[Session]]:
"""Create the banner table without its PostgreSQL-cast server defaults."""
table = TypeBase.metadata.tables[ExporleBanner.__tablename__]
status_default = table.c.status.server_default
language_default = table.c.language.server_default
table.c.status.server_default = None
table.c.language.server_default = None
try:
TypeBase.metadata.create_all(sqlite_engine, tables=[table])
finally:
table.c.status.server_default = status_default
table.c.language.server_default = language_default
class FakeExploreBannerQuery:
def __init__(self, responses: dict[str, tuple[ExploreBannerRecord, ...]] | None = None) -> None:
self.responses = responses or {}
self.requested_languages: list[str] = []
session = scoped_session(sessionmaker(bind=sqlite_engine, expire_on_commit=False))
try:
yield session
finally:
session.remove()
def list_enabled(self, language: str) -> tuple[ExploreBannerRecord, ...]:
self.requested_languages.append(language)
return self.responses.get(language, ())
class _ApplicationServicesStub(NamedTuple):
explore_banner_queries: ExploreBannerQueryService
def _content(
@@ -90,46 +79,59 @@ def _record(
)
def _use_sqlite_banner_service(
monkeypatch: pytest.MonkeyPatch,
sqlite_session_factory: sessionmaker[Session],
) -> None:
service = ExploreBannerQueryService(
banners=ExploreBannerQueryRepository(sqlite_session_factory),
is_enabled=lambda: True,
)
monkeypatch.setattr(
banner_module,
"application_services",
lambda: _ApplicationServicesStub(explore_banner_queries=service),
)
class TestExploreBannerQueryService:
def test_returns_empty_without_querying_when_disabled(self) -> None:
banners = MagicMock()
banners = FakeExploreBannerQuery()
service = ExploreBannerQueryService(banners=banners, is_enabled=lambda: False)
assert service.list_for_language("fr-FR") == ()
banners.list_enabled.assert_not_called()
assert banners.requested_languages == []
def test_returns_requested_language(self) -> None:
record = _record()
banners = MagicMock()
banners.list_enabled.return_value = (record,)
banners = FakeExploreBannerQuery({"fr-FR": (record,)})
service = ExploreBannerQueryService(banners=banners, is_enabled=lambda: True)
assert service.list_for_language("fr-FR") == (record,)
banners.list_enabled.assert_called_once_with("fr-FR")
assert banners.requested_languages == ["fr-FR"]
def test_falls_back_to_en_us(self) -> None:
record = _record(title="fallback")
banners = MagicMock()
banners.list_enabled.side_effect = [(), (record,)]
banners = FakeExploreBannerQuery({"en-US": (record,)})
service = ExploreBannerQueryService(banners=banners, is_enabled=lambda: True)
assert service.list_for_language("es-ES") == (record,)
assert banners.list_enabled.call_args_list == [
call("es-ES"),
call("en-US"),
]
assert banners.requested_languages == ["es-ES", "en-US"]
def test_does_not_repeat_default_language_query(self) -> None:
banners = MagicMock()
banners.list_enabled.return_value = ()
banners = FakeExploreBannerQuery()
service = ExploreBannerQueryService(banners=banners, is_enabled=lambda: True)
assert service.list_for_language("en-US") == ()
banners.list_enabled.assert_called_once_with("en-US")
assert banners.requested_languages == ["en-US"]
class TestExploreBannerQueryRepository:
def test_filters_language_and_status_and_orders_by_sort(self, banner_session: scoped_session[Session]) -> None:
def test_filters_language_and_status_and_orders_by_sort(
self,
sqlite_session: Session,
sqlite_session_factory: sessionmaker[Session],
) -> None:
created_at = datetime(2024, 1, 1)
second = _banner(
title="second",
@@ -158,10 +160,10 @@ class TestExploreBannerQueryRepository:
link="https://example.com/english",
created_at=created_at,
)
banner_session.add_all([second, first, disabled, english])
banner_session.commit()
sqlite_session.add_all([second, first, disabled, english])
sqlite_session.commit()
repository = ExploreBannerQueryRepository(banner_session.session_factory)
repository = ExploreBannerQueryRepository(sqlite_session_factory)
result = repository.list_enabled("fr-FR")
assert [banner.id for banner in result] == [first.id, second.id]
@@ -176,21 +178,29 @@ class TestExploreBannerQueryRepository:
class TestBannerApi:
def test_get_serializes_requested_language(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
queries = MagicMock()
queries.list_for_language.return_value = (_record(),)
monkeypatch.setattr(
banner_module,
"application_services",
lambda: SimpleNamespace(explore_banner_queries=queries),
def test_get_serializes_requested_language(
self,
app: Flask,
monkeypatch: pytest.MonkeyPatch,
sqlite_session: Session,
sqlite_session_factory: sessionmaker[Session],
) -> None:
banner = _banner(
title="hello",
language="fr-FR",
link="https://example.com",
created_at=datetime(2024, 1, 1),
)
sqlite_session.add(banner)
sqlite_session.commit()
_use_sqlite_banner_service(monkeypatch, sqlite_session_factory)
with app.test_request_context("/?language=fr-FR"):
result = banner_module.BannerApi().get()
assert result == [
{
"id": "banner-1",
"id": banner.id,
"content": {
"category": "Featured",
"title": "hello",
@@ -203,31 +213,48 @@ class TestBannerApi:
"created_at": "2024-01-01T00:00:00",
}
]
queries.list_for_language.assert_called_once_with("fr-FR")
def test_get_uses_default_language(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
queries = MagicMock()
queries.list_for_language.return_value = ()
monkeypatch.setattr(
banner_module,
"application_services",
lambda: SimpleNamespace(explore_banner_queries=queries),
def test_get_uses_default_language(
self,
app: Flask,
monkeypatch: pytest.MonkeyPatch,
sqlite_session: Session,
sqlite_session_factory: sessionmaker[Session],
) -> None:
banner = _banner(
title="default",
language="en-US",
link="https://example.com/default",
created_at=datetime(2024, 1, 2),
)
sqlite_session.add(banner)
sqlite_session.commit()
_use_sqlite_banner_service(monkeypatch, sqlite_session_factory)
with app.test_request_context("/"):
result = banner_module.BannerApi().get()
assert result == []
queries.list_for_language.assert_called_once_with("en-US")
assert result[0]["id"] == banner.id
assert result[0]["content"]["title"] == "default"
def test_get_allows_empty_supporting_copy(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
queries = MagicMock()
queries.list_for_language.return_value = (_record(category="", description=""),)
monkeypatch.setattr(
banner_module,
"application_services",
lambda: SimpleNamespace(explore_banner_queries=queries),
def test_get_allows_empty_supporting_copy(
self,
app: Flask,
monkeypatch: pytest.MonkeyPatch,
sqlite_session: Session,
sqlite_session_factory: sessionmaker[Session],
) -> None:
banner = _banner(
title="hello",
language="en-US",
link="https://example.com",
created_at=datetime(2024, 1, 3),
)
banner.content["category"] = ""
banner.content["description"] = ""
sqlite_session.add(banner)
sqlite_session.commit()
_use_sqlite_banner_service(monkeypatch, sqlite_session_factory)
with app.test_request_context("/"):
result = banner_module.BannerApi().get()
@@ -239,15 +266,23 @@ class TestBannerApi:
"img-src": "https://example.com/banner.png",
}
def test_get_rejects_invalid_content(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
queries = MagicMock()
queries.list_for_language.return_value = (_record()._replace(content={"title": "invalid"}),)
monkeypatch.setattr(
banner_module,
"application_services",
lambda: SimpleNamespace(explore_banner_queries=queries),
def test_get_rejects_invalid_content(
self,
app: Flask,
monkeypatch: pytest.MonkeyPatch,
sqlite_session: Session,
sqlite_session_factory: sessionmaker[Session],
) -> None:
banner = _banner(
title="invalid",
language="en-US",
link="https://example.com",
created_at=datetime(2024, 1, 4),
)
banner.content = {"title": "invalid"}
sqlite_session.add(banner)
sqlite_session.commit()
_use_sqlite_banner_service(monkeypatch, sqlite_session_factory)
with app.test_request_context("/"):
with pytest.raises(ValidationError):
banner_module.BannerApi().get()
with app.test_request_context("/"), pytest.raises(ValidationError):
banner_module.BannerApi().get()
@@ -4,6 +4,7 @@ from unittest.mock import MagicMock, PropertyMock, patch
import pytest
from flask import Flask
from sqlalchemy.orm import Session
from werkzeug.exceptions import InternalServerError
import controllers.console.explore.completion as completion_module
@@ -392,7 +393,9 @@ class TestChatApi:
api, completion_module.ChatMessagePayload.model_validate(payload_data), MagicMock(), user, chat_app
)
def test_invalid_conversation_id_fails_fast_as_not_found(self, app: Flask, chat_app, user) -> None:
def test_invalid_conversation_id_fails_fast_as_not_found(
self, app: Flask, chat_app, user, unbound_session: Session
) -> None:
# A nonexistent conversation_id must fail fast as 404, before the streaming
# generator is created. Previously the lookup only ran inside the generator,
# so an invalid id surfaced as a hang instead of a clean error.
@@ -407,7 +410,7 @@ class TestChatApi:
get_conversation_mock = MagicMock(
side_effect=completion_module.services.errors.conversation.ConversationNotExistsError()
)
session = MagicMock()
session = unbound_session
api = completion_module.ChatApi()
method = unwrap(api.post)
@@ -131,7 +131,7 @@ def test_trial_workflow_uses_trial_scoped_simple_account_model() -> None:
assert module.simple_account_model.__schema__["properties"].keys() >= {"id", "name", "email"}
def test_trial_dataset_list_preserves_slim_dataset_fields(app: Flask):
def test_trial_dataset_list_preserves_slim_dataset_fields(app: Flask, unbound_session: Session):
class DatasetListItem:
id = "dataset-1"
name = "Dataset"
@@ -150,8 +150,6 @@ def test_trial_dataset_list_preserves_slim_dataset_fields(app: Flask):
api = module.DatasetListApi()
method = unwrap(api.get)
app_model = SimpleNamespace(tenant_id="tenant-1")
session = MagicMock()
with (
app.test_request_context("/?page=1&limit=20&ids=dataset-1"),
patch.object(
@@ -160,9 +158,9 @@ def test_trial_dataset_list_preserves_slim_dataset_fields(app: Flask):
return_value=([DatasetListItem()], 1),
) as get_datasets,
):
result = method(api, session, app_model)
result = method(api, unbound_session, app_model)
get_datasets.assert_called_once_with(["dataset-1"], "tenant-1", session=session)
get_datasets.assert_called_once_with(["dataset-1"], "tenant-1", session=unbound_session)
assert result == {
"data": [
{
@@ -195,8 +193,9 @@ def test_trial_app_handlers_use_explicit_read_session(api_type: type) -> None:
assert tuple(signature(api_type.get).parameters)[:3] == ("self", "session", "app_model")
def test_trial_app_detail_serializes_with_explicit_session(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
session = MagicMock()
def test_trial_app_detail_serializes_with_explicit_session(
app: Flask, monkeypatch: pytest.MonkeyPatch, unbound_session: Session
) -> None:
app_model = MagicMock()
response_view = MagicMock()
get_app = MagicMock(return_value=app_model)
@@ -208,11 +207,11 @@ def test_trial_app_detail_serializes_with_explicit_session(app: Flask, monkeypat
monkeypatch.setattr(module.TrialAppDetailResponse, "model_validate", MagicMock(return_value=validated))
with app.test_request_context("/"):
result = unwrap(module.AppApi.get)(module.AppApi(), session, app_model)
result = unwrap(module.AppApi.get)(module.AppApi(), unbound_session, app_model)
assert result == {"id": "app-1"}
get_app.assert_called_once_with(app_model, session=session)
build_view.assert_called_once_with(app_model, session=session)
get_app.assert_called_once_with(app_model, session=unbound_session)
build_view.assert_called_once_with(app_model, session=unbound_session)
module.TrialAppDetailResponse.model_validate.assert_called_once_with(response_view, from_attributes=True)
@@ -882,14 +881,14 @@ class TestTrialMessageSuggestedQuestionApi:
class TestTrialAppParameterApi:
def test_app_unavailable(self) -> None:
def test_app_unavailable(self, unbound_session: Session) -> None:
api = module.TrialAppParameterApi()
method = unwrap(api.get)
with pytest.raises(AppUnavailableError):
method(api, MagicMock(), None)
method(api, unbound_session, None)
def test_success_non_workflow(self, valid_parameters: dict[str, object]) -> None:
def test_success_non_workflow(self, valid_parameters: dict[str, object], unbound_session: Session) -> None:
api = module.TrialAppParameterApi()
method = unwrap(api.get)
@@ -899,7 +898,6 @@ class TestTrialAppParameterApi:
mode=AppMode.CHAT,
app_model_config_with_session=MagicMock(return_value=app_model_config),
)
session = MagicMock()
annotation_reply = {"enabled": False}
with (
@@ -917,14 +915,14 @@ class TestTrialAppParameterApi:
return_value=MagicMock(model_dump=lambda mode=None: {"ok": True}),
),
):
result = method(api, session, app_model)
result = method(api, unbound_session, app_model)
assert result == {"ok": True}
app_model.app_model_config_with_session.assert_called_once_with(session=session)
load_annotation_reply.assert_called_once_with(session, "app-1")
app_model.app_model_config_with_session.assert_called_once_with(session=unbound_session)
load_annotation_reply.assert_called_once_with(unbound_session, "app-1")
app_model_config.to_dict.assert_called_once_with(annotation_reply=annotation_reply)
def test_success_workflow(self, valid_parameters: dict[str, object]) -> None:
def test_success_workflow(self, valid_parameters: dict[str, object], unbound_session: Session) -> None:
api = module.TrialAppParameterApi()
method = unwrap(api.get)
@@ -934,8 +932,6 @@ class TestTrialAppParameterApi:
mode=AppMode.WORKFLOW,
workflow_with_session=MagicMock(return_value=workflow),
)
session = MagicMock()
with (
patch.object(module, "get_parameters_from_feature_dict", return_value=valid_parameters),
patch.object(
@@ -944,10 +940,10 @@ class TestTrialAppParameterApi:
return_value=MagicMock(model_dump=lambda mode=None: {"ok": True}),
),
):
result = method(api, session, app_model)
result = method(api, unbound_session, app_model)
assert result == {"ok": True}
app_model.workflow_with_session.assert_called_once_with(session=session)
app_model.workflow_with_session.assert_called_once_with(session=unbound_session)
workflow.user_input_form.assert_called_once_with(to_old_structure=True)
@@ -1446,7 +1442,7 @@ class TestTrialSitApi:
class TestAppWorkflowApi:
def test_uses_injected_session(self) -> None:
def test_uses_injected_session(self, unbound_session: Session) -> None:
api = module.AppWorkflowApi()
method = unwrap(api.get)
created_by = SimpleNamespace(id="account-1", name="Creator", email="creator@example.com")
@@ -1489,9 +1485,7 @@ class TestAppWorkflowApi:
workflow_id="workflow-1",
workflow_with_session=MagicMock(return_value=workflow),
)
session = MagicMock()
result = method(api, session, app_model)
result = method(api, unbound_session, app_model)
assert result == {
"id": "workflow-1",
@@ -1535,10 +1529,10 @@ class TestAppWorkflowApi:
],
"rag_pipeline_variables": [],
}
app_model.workflow_with_session.assert_called_once_with(session=session)
workflow.get_created_by_account.assert_called_once_with(session=session)
workflow.get_updated_by_account.assert_called_once_with(session=session)
workflow.get_tool_published.assert_called_once_with(session=session)
app_model.workflow_with_session.assert_called_once_with(session=unbound_session)
workflow.get_created_by_account.assert_called_once_with(session=unbound_session)
workflow.get_updated_by_account.assert_called_once_with(session=unbound_session)
workflow.get_tool_published.assert_called_once_with(session=unbound_session)
class TestTrialChatAudioApiExceptionHandlers:
@@ -2,18 +2,19 @@ from __future__ import annotations
import inspect
from collections.abc import Callable
from types import SimpleNamespace
from typing import cast
from unittest.mock import MagicMock, patch
from unittest.mock import patch
import pytest
from sqlalchemy import event, select
from sqlalchemy.orm import Session
from werkzeug.exceptions import Forbidden
from controllers.console.apikey import BaseApiKeyListResource, BaseApiKeyResource
from models import Account
from models.account import AccountStatus, TenantAccountRole
from models.enums import ApiTokenType
from models.model import ApiToken, App, AppMode
from models.model import ApiToken, App, AppMode, IconType
from services.agent.errors import AgentAccessNotReadyError
@@ -45,77 +46,80 @@ def _make_account(role: TenantAccountRole) -> Account:
return account
def test_list_api_keys_uses_injected_session_and_tenant_id() -> None:
def _persist_app(session: Session, *, mode: AppMode = AppMode.CHAT) -> App:
app = App(
id="app-1",
tenant_id="tenant-1",
name="API key app",
mode=mode,
icon_type=IconType.EMOJI,
icon="chat",
icon_background="#ffffff",
enable_site=False,
enable_api=True,
)
session.add(app)
session.flush()
return app
def test_list_api_keys_uses_injected_session_and_tenant_id(sqlite_session: Session) -> None:
resource = _make_list_resource()
raw_get = cast(
Callable[[BaseApiKeyListResource, object, str, str], dict[str, object]],
inspect.unwrap(BaseApiKeyListResource.get),
)
session = MagicMock()
session.execute.return_value.scalar_one_or_none.return_value = SimpleNamespace()
api_key = SimpleNamespace(
id="key-1",
session = sqlite_session
_persist_app(session)
api_key = ApiToken(
type=ApiTokenType.APP,
token="app-token",
last_used_at=None,
created_at=None,
app_id="app-1",
tenant_id="tenant-1",
)
session.scalars.return_value.all.return_value = [api_key]
api_key.id = "key-1"
session.add(api_key)
session.commit()
result = raw_get(resource, session, "app-1", "tenant-1")
data = cast(list[dict[str, object]], result["data"])
session.execute.assert_called_once()
session.scalars.assert_called_once()
assert result == {
"data": [
{
"id": "key-1",
"type": "app",
"token": "app-token",
"last_used_at": None,
"created_at": None,
}
]
}
assert len(data) == 1
assert data[0]["id"] == "key-1"
assert data[0]["token"] == "app-token"
def test_create_api_key_uses_injected_session_and_tenant_id() -> None:
def test_create_api_key_uses_injected_session_and_tenant_id(sqlite_session: Session) -> None:
resource = _make_list_resource()
raw_post = cast(
Callable[[BaseApiKeyListResource, object, str, str], tuple[dict[str, object], int]],
inspect.unwrap(BaseApiKeyListResource.post),
)
session = MagicMock()
session.execute.return_value.scalar_one_or_none.return_value = SimpleNamespace()
session.scalar.return_value = 0
def add_api_token(api_token: ApiToken) -> None:
api_token.id = "key-1"
session = sqlite_session
_persist_app(session)
commits: list[str] = []
event.listen(session, "after_commit", lambda _session: commits.append("commit"))
with patch(
"controllers.console.apikey.ApiToken.generate_api_key", return_value="app-generated-token"
) as generate_api_key:
session.add.side_effect = add_api_token
result, status = raw_post(resource, session, "app-1", "tenant-1")
assert status == 201
assert result["token"] == "app-generated-token"
api_token = session.add.call_args.args[0]
api_token = session.scalar(select(ApiToken).where(ApiToken.token == "app-generated-token"))
assert api_token is not None
assert api_token.app_id == "app-1"
assert api_token.tenant_id == "tenant-1"
assert api_token.type == ApiTokenType.APP
generate_api_key.assert_called_once_with("app-", 24, session=session)
session.execute.assert_called_once()
session.scalar.assert_called_once()
session.commit.assert_called_once()
assert commits == ["commit"]
def test_create_agent_api_key_requires_published_access() -> None:
def test_create_agent_api_key_requires_published_access(sqlite_session: Session) -> None:
resource = _make_list_resource()
app = App(id="app-1", tenant_id="tenant-1", mode=AppMode.AGENT)
session = MagicMock()
session.execute.return_value.scalar_one_or_none.return_value = app
session = sqlite_session
app = _persist_app(session, mode=AppMode.AGENT)
with patch(
"controllers.console.apikey.AppService.ensure_agent_app_access_ready",
@@ -125,18 +129,17 @@ def test_create_agent_api_key_requires_published_access() -> None:
resource._create_api_key("app-1", "tenant-1", session=session)
ensure_access_ready.assert_called_once_with(app, session=session)
session.scalar.assert_not_called()
session.add.assert_not_called()
assert session.scalar(select(ApiToken)) is None
def test_delete_api_key_rejects_non_admin_account() -> None:
def test_delete_api_key_rejects_non_admin_account(sqlite_session: Session) -> None:
resource = _make_key_resource()
raw_delete = cast(
Callable[[BaseApiKeyResource, object, str, str, str, Account], tuple[str, int]],
inspect.unwrap(BaseApiKeyResource.delete),
)
session = MagicMock()
session.execute.return_value.scalar_one_or_none.return_value = SimpleNamespace()
session = sqlite_session
_persist_app(session)
with pytest.raises(Forbidden):
raw_delete(
@@ -148,20 +151,21 @@ def test_delete_api_key_rejects_non_admin_account() -> None:
_make_account(TenantAccountRole.NORMAL),
)
session.execute.assert_called_once()
session.scalar.assert_not_called()
def test_delete_api_key_uses_injected_session_user_and_tenant() -> None:
def test_delete_api_key_uses_injected_session_user_and_tenant(sqlite_session: Session) -> None:
resource = _make_key_resource()
raw_delete = cast(
Callable[[BaseApiKeyResource, object, str, str, str, Account], tuple[str, int]],
inspect.unwrap(BaseApiKeyResource.delete),
)
api_key = SimpleNamespace(token="app-token", type=ApiTokenType.APP)
session = MagicMock()
session.execute.return_value.scalar_one_or_none.return_value = SimpleNamespace()
session.scalar.return_value = api_key
session = sqlite_session
_persist_app(session)
api_key = ApiToken(type=ApiTokenType.APP, token="app-token", app_id="app-1", tenant_id="tenant-1")
api_key.id = "key-1"
session.add(api_key)
session.commit()
commits: list[str] = []
event.listen(session, "after_commit", lambda _session: commits.append("commit"))
with patch("controllers.console.apikey.ApiTokenCache.delete") as delete_cache:
result, status = raw_delete(
@@ -174,8 +178,7 @@ def test_delete_api_key_uses_injected_session_user_and_tenant() -> None:
)
delete_cache.assert_called_once_with("app-token", ApiTokenType.APP)
assert session.execute.call_count == 2
session.scalar.assert_called_once()
session.commit.assert_called_once()
assert session.get(ApiToken, "key-1") is None
assert commits == ["commit"]
assert result == ""
assert status == 204
@@ -641,9 +641,8 @@ class TestWorkspaceInfoApi:
),
),
):
session = MagicMock()
session.get.return_value = tenant
session.commit.side_effect = lambda: events.append("commit")
session = workspace_session()
event.listen(session, "after_commit", lambda _session: events.append("commit"))
result = method(api, session, "t1")
assert result["result"] == "success"
assert events == ["commit", "get_tenant_info"]
@@ -6,22 +6,47 @@ in test_auth_wraps.py; handler tests use inspect.unwrap() to bypass them.
"""
import inspect
from unittest.mock import ANY, MagicMock, patch
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
from uuid import uuid4
import pytest
from flask import Flask
from pydantic import ValidationError
from sqlalchemy import event, select
from sqlalchemy.engine import Engine
from sqlalchemy.orm import Session, scoped_session, sessionmaker
from controllers.inner_api.app import dsl as dsl_module
from controllers.inner_api.app.dsl import (
EnterpriseAppDSLExport,
EnterpriseAppDSLImport,
InnerAppDSLImportPayload,
_get_active_account,
)
from models import Account, App
from models.account import AccountStatus
from models.model import AppMode, IconType
from services.app_dsl_service import Import, ImportStatus
def _persist_app(session: Session) -> App:
app = App(
id=str(uuid4()),
tenant_id=str(uuid4()),
name="DSL App",
mode=AppMode.WORKFLOW,
icon_type=IconType.EMOJI,
icon="robot",
icon_background="#ffffff",
enable_site=False,
enable_api=False,
)
session.add(app)
session.commit()
return app
class TestInnerAppDSLImportPayload:
"""Test InnerAppDSLImportPayload Pydantic model validation."""
@@ -61,32 +86,29 @@ class TestInnerAppDSLImportPayload:
class TestGetActiveAccount:
"""Test the _get_active_account helper function."""
@patch("controllers.inner_api.app.dsl.db")
def test_returns_active_account(self, mock_db):
mock_account = MagicMock()
mock_account.status = AccountStatus.ACTIVE
mock_db.session.scalar.return_value = mock_account
def test_returns_active_account(self, sqlite_session: Session):
account = Account(name="Active", email="user@example.com", status=AccountStatus.ACTIVE)
sqlite_session.add(account)
sqlite_session.commit()
result = _get_active_account("user@example.com")
with patch.object(dsl_module.db, "session", sqlite_session):
result = _get_active_account("user@example.com")
assert result is mock_account
mock_db.session.scalar.assert_called_once()
assert result is account
@patch("controllers.inner_api.app.dsl.db")
def test_returns_none_for_inactive_account(self, mock_db):
mock_account = MagicMock()
mock_account.status = AccountStatus.BANNED
mock_db.session.scalar.return_value = mock_account
def test_returns_none_for_inactive_account(self, sqlite_session: Session):
account = Account(name="Banned", email="banned@example.com", status=AccountStatus.BANNED)
sqlite_session.add(account)
sqlite_session.commit()
result = _get_active_account("banned@example.com")
with patch.object(dsl_module.db, "session", sqlite_session):
result = _get_active_account("banned@example.com")
assert result is None
@patch("controllers.inner_api.app.dsl.db")
def test_returns_none_for_nonexistent_email(self, mock_db):
mock_db.session.scalar.return_value = None
result = _get_active_account("missing@example.com")
def test_returns_none_for_nonexistent_email(self, sqlite_session: Session):
with patch.object(dsl_module.db, "session", sqlite_session):
result = _get_active_account("missing@example.com")
assert result is None
@@ -102,20 +124,29 @@ class TestEnterpriseAppDSLImport:
return EnterpriseAppDSLImport()
@pytest.fixture
def _mock_import_deps(self):
"""Patch db, Session, and AppDslService for import handler tests."""
mock_session = MagicMock()
mock_session.__enter__ = MagicMock(return_value=mock_session)
mock_session.__exit__ = MagicMock(return_value=False)
def _mock_import_deps(self, sqlite_engine: Engine):
"""Bind the handler Session to SQLite and isolate the DSL service boundary."""
self._transaction_events: list[str] = []
def on_commit(session: Session) -> None:
if session.get_bind() is sqlite_engine:
self._transaction_events.append("commit")
def on_rollback(session: Session) -> None:
if session.get_bind() is sqlite_engine:
self._transaction_events.append("rollback")
event.listen(Session, "after_commit", on_commit)
event.listen(Session, "after_rollback", on_rollback)
with (
patch("controllers.inner_api.app.dsl.db"),
patch("controllers.inner_api.app.dsl.Session", return_value=mock_session),
patch.object(dsl_module, "db", SimpleNamespace(engine=sqlite_engine)),
patch("controllers.inner_api.app.dsl.AppDslService") as mock_dsl_cls,
):
self._mock_session = mock_session
self._mock_dsl = MagicMock()
mock_dsl_cls.return_value = self._mock_dsl
yield
event.remove(Session, "after_commit", on_commit)
event.remove(Session, "after_rollback", on_rollback)
def _make_import_result(self, status: ImportStatus, **kwargs) -> Import:
result = Import(
@@ -145,9 +176,9 @@ class TestEnterpriseAppDSLImport:
body, status_code = result
assert status_code == 200
assert body["status"] == "completed"
mock_account.set_tenant_id_with_session.assert_called_once_with("ws-123", session=self._mock_session)
self._mock_session.commit.assert_called_once_with()
self._mock_session.rollback.assert_not_called()
call_session = mock_account.set_tenant_id_with_session.call_args.kwargs["session"]
assert isinstance(call_session, Session)
assert self._transaction_events == ["commit"]
@pytest.mark.usefixtures("_mock_import_deps")
@patch("controllers.inner_api.app.dsl._get_active_account")
@@ -163,13 +194,14 @@ class TestEnterpriseAppDSLImport:
assert status_code == 202
assert body["status"] == "pending"
self._mock_session.commit.assert_called_once_with()
self._mock_session.rollback.assert_not_called()
assert self._transaction_events == ["commit"]
@pytest.mark.usefixtures("_mock_import_deps")
@patch("controllers.inner_api.app.dsl._get_active_account")
def test_import_failed_returns_400(self, mock_get_account, api_instance, app: Flask):
mock_get_account.return_value = MagicMock()
mock_account = MagicMock()
mock_account.set_tenant_id_with_session.side_effect = lambda _tenant_id, *, session: session.execute(select(1))
mock_get_account.return_value = mock_account
self._mock_dsl.import_app.return_value = self._make_import_result(ImportStatus.FAILED)
unwrapped = inspect.unwrap(api_instance.post)
@@ -180,8 +212,7 @@ class TestEnterpriseAppDSLImport:
assert status_code == 400
assert body["status"] == "failed"
self._mock_session.rollback.assert_called_once_with()
self._mock_session.commit.assert_not_called()
assert self._transaction_events == ["rollback"]
@patch("controllers.inner_api.app.dsl._get_active_account")
def test_import_account_not_found_returns_404(self, mock_get_account, api_instance, app: Flask):
@@ -208,44 +239,65 @@ class TestEnterpriseAppDSLExport:
def api_instance(self):
return EnterpriseAppDSLExport()
@pytest.fixture
def scoped_db(self, sqlite_session_factory: sessionmaker[Session]):
db_session = scoped_session(sqlite_session_factory)
with patch.object(dsl_module, "db", SimpleNamespace(session=db_session)):
yield db_session
db_session.remove()
@patch("controllers.inner_api.app.dsl.AppDslService")
@patch("controllers.inner_api.app.dsl.db")
def test_export_success_returns_200(self, mock_db, mock_dsl_cls, api_instance, app: Flask):
mock_app = MagicMock()
mock_db.session.get.return_value = mock_app
def test_export_success_returns_200(
self,
mock_dsl_cls,
api_instance,
app: Flask,
sqlite_session: Session,
scoped_db,
):
app_model = _persist_app(sqlite_session)
mock_dsl_cls.export_dsl.return_value = "version: 0.6.0\nkind: app\n"
unwrapped = inspect.unwrap(api_instance.get)
with app.test_request_context("?include_secret=false"):
result = unwrapped(api_instance, app_id="app-123")
result = unwrapped(api_instance, app_id=app_model.id)
body, status_code = result
assert status_code == 200
assert body["data"] == "version: 0.6.0\nkind: app\n"
mock_dsl_cls.export_dsl.assert_called_once_with(app_model=mock_app, session=ANY, include_secret=False)
call_kwargs = mock_dsl_cls.export_dsl.call_args.kwargs
assert call_kwargs["app_model"].id == app_model.id
assert call_kwargs["session"] is scoped_db()
assert call_kwargs["include_secret"] is False
@patch("controllers.inner_api.app.dsl.AppDslService")
@patch("controllers.inner_api.app.dsl.db")
def test_export_with_secret(self, mock_db, mock_dsl_cls, api_instance, app: Flask):
mock_app = MagicMock()
mock_db.session.get.return_value = mock_app
def test_export_with_secret(
self,
mock_dsl_cls,
api_instance,
app: Flask,
sqlite_session: Session,
scoped_db,
):
app_model = _persist_app(sqlite_session)
mock_dsl_cls.export_dsl.return_value = "yaml-data"
unwrapped = inspect.unwrap(api_instance.get)
with app.test_request_context("?include_secret=true"):
result = unwrapped(api_instance, app_id="app-123")
result = unwrapped(api_instance, app_id=app_model.id)
body, status_code = result
assert status_code == 200
mock_dsl_cls.export_dsl.assert_called_once_with(app_model=mock_app, session=ANY, include_secret=True)
@patch("controllers.inner_api.app.dsl.db")
def test_export_app_not_found_returns_404(self, mock_db, api_instance, app: Flask):
mock_db.session.get.return_value = None
call_kwargs = mock_dsl_cls.export_dsl.call_args.kwargs
assert call_kwargs["app_model"].id == app_model.id
assert call_kwargs["session"] is scoped_db()
assert call_kwargs["include_secret"] is True
def test_export_app_not_found_returns_404(self, api_instance, app: Flask, scoped_db):
assert scoped_db() is not None
unwrapped = inspect.unwrap(api_instance.get)
with app.test_request_context("?include_secret=false"):
result = unwrapped(api_instance, app_id="nonexistent")
result = unwrapped(api_instance, app_id=str(uuid4()))
body, status_code = result
assert status_code == 404
@@ -6,10 +6,10 @@ from unittest.mock import MagicMock, patch
import pytest
from flask import Flask
from sqlalchemy.orm import Session
from controllers.inner_api.agent.files import (
AgentFileDownloadRequestApi,
AgentFileRequestHttpError,
AgentFileUploadRequestApi,
)
from core.workflow.file_reference import build_file_reference
@@ -22,7 +22,7 @@ def _raw[R](method: Callable[..., R]) -> Callable[..., R]:
return cast(Callable[..., R], inspect.unwrap(method))
def test_upload_request_returns_origin_free_uri(app: Flask) -> None:
def test_upload_request_returns_origin_free_uri(app: Flask, unbound_session: Session) -> None:
payload = {
"tenant_id": "tenant-1",
"user_id": "execution-user-1",
@@ -32,7 +32,7 @@ def test_upload_request_returns_origin_free_uri(app: Flask) -> None:
}
tenant = SimpleNamespace(id="tenant-1")
user = SimpleNamespace(id="canonical-end-user-1")
session = MagicMock()
session = unbound_session
with app.test_request_context("/", method="POST", json=payload):
with (
patch(f"{MODULE}.TenantService") as tenant_service,
@@ -54,71 +54,7 @@ def test_upload_request_returns_origin_free_uri(app: Flask) -> None:
)
def test_upload_request_preserves_tenant_scoped_account_owner(app: Flask) -> None:
payload = {
"tenant_id": "tenant-1",
"user_id": "account-1",
"user_from": "account",
"filename": "report.pdf",
"mimetype": "application/pdf",
"conversation_id": "conversation-1",
}
tenant = SimpleNamespace(id="tenant-1")
session = MagicMock()
with app.test_request_context("/", method="POST", json=payload):
with (
patch(f"{MODULE}.TenantService") as tenant_service,
patch(f"{MODULE}.get_user") as get_user,
patch(f"{MODULE}.get_signed_file_uri_for_plugin", return_value="/files/upload/for-plugin?sign=1") as sign,
):
tenant_service.get_tenant_by_id.return_value = tenant
tenant_service.account_belongs_to_tenant.return_value = True
response = _raw(AgentFileUploadRequestApi.post)(AgentFileUploadRequestApi(), session)
assert response == {"upload_uri": "/files/upload/for-plugin?sign=1"}
get_user.assert_not_called()
tenant_service.account_belongs_to_tenant.assert_called_once_with("account-1", "tenant-1", session=session)
sign.assert_called_once_with(
filename="report.pdf",
mimetype="application/pdf",
tenant_id="tenant-1",
user_id="account-1",
conversation_id="conversation-1",
user_from="account",
)
def test_upload_request_rejects_account_outside_tenant_without_signing(app: Flask) -> None:
payload = {
"tenant_id": "tenant-1",
"user_id": "account-outside-tenant",
"user_from": "account",
"filename": "report.pdf",
"mimetype": "application/pdf",
}
tenant = SimpleNamespace(id="tenant-1")
session = MagicMock()
with app.test_request_context("/", method="POST", json=payload):
with (
patch(f"{MODULE}.TenantService") as tenant_service,
patch(f"{MODULE}.get_user") as get_user,
patch(f"{MODULE}.get_signed_file_uri_for_plugin") as sign,
):
tenant_service.get_tenant_by_id.return_value = tenant
tenant_service.account_belongs_to_tenant.return_value = False
with pytest.raises(AgentFileRequestHttpError) as exc_info:
_raw(AgentFileUploadRequestApi.post)(AgentFileUploadRequestApi(), session)
assert exc_info.value.error_code == "user_not_found"
assert exc_info.value.code == 404
tenant_service.account_belongs_to_tenant.assert_called_once_with(
"account-outside-tenant", "tenant-1", session=session
)
get_user.assert_not_called()
sign.assert_not_called()
def test_download_request_returns_origin_free_uri_for_sandbox(app: Flask) -> None:
def test_download_request_returns_origin_free_uri_for_sandbox(app: Flask, unbound_session: Session) -> None:
reference = build_file_reference(record_id="tool-file-1")
payload = {
"tenant_id": "tenant-1",
@@ -128,7 +64,7 @@ def test_download_request_returns_origin_free_uri_for_sandbox(app: Flask) -> Non
"file": {"transfer_method": "tool_file", "reference": reference},
"for_frontend": False,
}
session = MagicMock()
session = unbound_session
with app.test_request_context("/", method="POST", json=payload):
with (
patch(f"{MODULE}.TenantService") as tenant_service,
@@ -158,7 +94,9 @@ def test_download_request_returns_origin_free_uri_for_sandbox(app: Flask) -> Non
)
def test_download_request_binds_frontend_url(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
def test_download_request_binds_frontend_url(
app: Flask, monkeypatch: pytest.MonkeyPatch, unbound_session: Session
) -> None:
reference = build_file_reference(record_id="tool-file-1")
payload = {
"tenant_id": "tenant-1",
@@ -169,7 +107,7 @@ def test_download_request_binds_frontend_url(app: Flask, monkeypatch: pytest.Mon
"for_frontend": True,
}
monkeypatch.setattr(f"{MODULE}.dify_config.FILES_URL", "https://files.example.com")
session = MagicMock()
session = unbound_session
with app.test_request_context("/", method="POST", json=payload):
with (
patch(f"{MODULE}.TenantService") as tenant_service,
@@ -4,6 +4,7 @@ from unittest.mock import MagicMock, patch
import pytest
from flask import Flask
from sqlalchemy.orm import Session
from werkzeug.exceptions import Forbidden, NotFound, Unauthorized
from controllers.openapi.auth.data import AuthData, ExternalIdentity
@@ -150,10 +151,10 @@ def test_load_account_skips_when_already_set():
assert data.caller is existing_caller
def test_load_account_sets_current_tenant_when_tenant_present():
def test_load_account_sets_current_tenant_when_tenant_present(sqlite_session: Session):
account = MagicMock()
tenant = MagicMock()
session = MagicMock()
session = sqlite_session
data = _make_auth_data(account_id=uuid.uuid4(), tenant=tenant)
with (
patch("controllers.openapi.auth.prepare.AccountService.get_account_by_id", return_value=account),
@@ -1,6 +1,8 @@
from types import SimpleNamespace
from unittest.mock import MagicMock
from sqlalchemy.orm import Session
from controllers.openapi._input_schema import EMPTY_INPUT_SCHEMA
from controllers.openapi.apps import _EMPTY_PARAMETERS, build_app_describe_response
from controllers.service_api.app.error import AppUnavailableError
@@ -25,9 +27,9 @@ def _app() -> _FakeApp:
)
def test_fields_none_returns_all_blocks(monkeypatch):
def test_fields_none_returns_all_blocks(monkeypatch, unbound_session: Session):
app = _app()
session = MagicMock()
session = unbound_session
parameters_payload = MagicMock(return_value={"k": "v"})
input_schema = MagicMock(return_value={"s": 1})
monkeypatch.setattr("controllers.openapi.apps.parameters_payload", parameters_payload)
@@ -41,8 +43,8 @@ def test_fields_none_returns_all_blocks(monkeypatch):
input_schema.assert_called_once_with(app, session=session)
def test_fields_subset_limits_blocks(monkeypatch):
session = MagicMock()
def test_fields_subset_limits_blocks(monkeypatch, unbound_session: Session):
session = unbound_session
monkeypatch.setattr("controllers.openapi.apps.parameters_payload", MagicMock(return_value={"k": "v"}))
monkeypatch.setattr("controllers.openapi.apps.build_input_schema", MagicMock(return_value={"s": 1}))
resp = build_app_describe_response(_app(), ["info"], session=session)
@@ -51,8 +53,8 @@ def test_fields_subset_limits_blocks(monkeypatch):
assert resp.input_schema is None
def test_info_omits_author_and_tags(monkeypatch):
session = MagicMock()
def test_info_omits_author_and_tags(monkeypatch, unbound_session: Session):
session = unbound_session
monkeypatch.setattr("controllers.openapi.apps.parameters_payload", MagicMock(return_value={}))
monkeypatch.setattr("controllers.openapi.apps.build_input_schema", MagicMock(return_value={}))
resp = build_app_describe_response(_app(), ["info"], session=session)
@@ -62,21 +64,21 @@ def test_info_omits_author_and_tags(monkeypatch):
assert not hasattr(resp.info, "tags")
def test_parameters_fallback_on_app_unavailable(monkeypatch):
def test_parameters_fallback_on_app_unavailable(monkeypatch, unbound_session: Session):
def _raise(app, *, session):
raise AppUnavailableError()
monkeypatch.setattr("controllers.openapi.apps.parameters_payload", _raise)
monkeypatch.setattr("controllers.openapi.apps.build_input_schema", MagicMock(return_value={"s": 1}))
resp = build_app_describe_response(_app(), ["parameters"], session=MagicMock())
resp = build_app_describe_response(_app(), ["parameters"], session=unbound_session)
assert resp.parameters == dict(_EMPTY_PARAMETERS)
def test_input_schema_fallback_on_app_unavailable(monkeypatch):
def test_input_schema_fallback_on_app_unavailable(monkeypatch, unbound_session: Session):
def _raise(app, *, session):
raise AppUnavailableError()
monkeypatch.setattr("controllers.openapi.apps.parameters_payload", MagicMock(return_value={"k": "v"}))
monkeypatch.setattr("controllers.openapi.apps.build_input_schema", _raise)
resp = build_app_describe_response(_app(), ["input_schema"], session=MagicMock())
resp = build_app_describe_response(_app(), ["input_schema"], session=unbound_session)
assert resp.input_schema == dict(EMPTY_INPUT_SCHEMA)
@@ -8,6 +8,7 @@ from types import SimpleNamespace
from unittest.mock import MagicMock
import pytest
from sqlalchemy.orm import Session
from controllers.openapi.apps import ( # pyright: ignore[reportPrivateUsage]
_EMPTY_PARAMETERS,
@@ -35,10 +36,10 @@ def _fake_app(**overrides):
return SimpleNamespace(**base)
def test_parameters_payload_raises_app_unavailable_when_no_config():
def test_parameters_payload_raises_app_unavailable_when_no_config(unbound_session: Session):
app = _fake_app(mode="chat")
app.app_model_config_with_session = MagicMock(return_value=None)
session = MagicMock()
session = unbound_session
with pytest.raises(AppUnavailableError):
parameters_payload(app, session=session)
@@ -14,6 +14,7 @@ from unittest.mock import MagicMock, patch
import pytest
from pydantic import ValidationError
from sqlalchemy.orm import Session
from controllers.openapi.apps_permitted_external import (
PermittedExternalAppDescribeApi,
@@ -69,10 +70,10 @@ def test_query_accepts_valid_mode():
assert q.mode.value == "chat"
def test_describe_forwards_request_session_to_response_builder():
def test_describe_forwards_request_session_to_response_builder(unbound_session: Session):
api = PermittedExternalAppDescribeApi()
method = inspect.unwrap(api.get)
session = MagicMock()
session = unbound_session
app = MagicMock()
auth_data = SimpleNamespace(app=app)
query = SimpleNamespace(fields={"info"})
@@ -3,6 +3,7 @@
from __future__ import annotations
import pytest
from sqlalchemy.orm import Session
from controllers.openapi._input_schema import _form_to_jsonschema
@@ -97,6 +98,7 @@ from models.model import AppMode
def _stub_app(mode: AppMode, *, form: list[dict] | None = None, has_workflow: bool | None = None):
"""Returns a MagicMock whose explicit config getters are wired up."""
app = MagicMock()
app.id = "00000000-0000-0000-0000-000000000001"
app.mode = mode
if mode in (AppMode.WORKFLOW, AppMode.ADVANCED_CHAT):
if has_workflow is False:
@@ -111,20 +113,15 @@ def _stub_app(mode: AppMode, *, form: list[dict] | None = None, has_workflow: bo
app.app_model_config_with_session.return_value = None
else:
app_model_config = MagicMock()
app_model_config.app_id = app.id
app_model_config.to_dict.return_value = {"user_input_form": form or []}
app.app_model_config_with_session.return_value = app_model_config
return app
def _session() -> MagicMock:
session = MagicMock()
session.scalar.return_value = None
return session
def test_chat_mode_includes_query() -> None:
def test_chat_mode_includes_query(sqlite_session: Session) -> None:
app = _stub_app(AppMode.CHAT, form=[{"text-input": {"variable": "x", "label": "X", "required": True}}])
session = _session()
session = sqlite_session
schema = build_input_schema(app, session=session)
assert schema["$schema"] == "https://json-schema.org/draft/2020-12/schema"
assert "query" in schema["properties"]
@@ -137,33 +134,33 @@ def test_chat_mode_includes_query() -> None:
app.app_model_config_with_session.return_value.to_dict.assert_called_once_with(annotation_reply={"enabled": False})
def test_agent_chat_mode_includes_query() -> None:
def test_agent_chat_mode_includes_query(sqlite_session: Session) -> None:
app = _stub_app(AppMode.AGENT_CHAT, form=[])
schema = build_input_schema(app, session=_session())
schema = build_input_schema(app, session=sqlite_session)
assert "query" in schema["properties"]
def test_advanced_chat_mode_includes_query() -> None:
def test_advanced_chat_mode_includes_query(sqlite_session: Session) -> None:
app = _stub_app(AppMode.ADVANCED_CHAT, form=[])
schema = build_input_schema(app, session=_session())
schema = build_input_schema(app, session=sqlite_session)
assert "query" in schema["properties"]
def test_workflow_mode_omits_query() -> None:
def test_workflow_mode_omits_query(sqlite_session: Session) -> None:
app = _stub_app(AppMode.WORKFLOW, form=[])
schema = build_input_schema(app, session=_session())
schema = build_input_schema(app, session=sqlite_session)
assert "query" not in schema["properties"]
assert schema["required"] == ["inputs"]
def test_completion_mode_omits_query() -> None:
def test_completion_mode_omits_query(sqlite_session: Session) -> None:
app = _stub_app(AppMode.COMPLETION, form=[])
schema = build_input_schema(app, session=_session())
schema = build_input_schema(app, session=sqlite_session)
assert "query" not in schema["properties"]
assert schema["required"] == ["inputs"]
def test_inputs_required_driven_by_form() -> None:
def test_inputs_required_driven_by_form(sqlite_session: Session) -> None:
app = _stub_app(
AppMode.CHAT,
form=[
@@ -171,20 +168,20 @@ def test_inputs_required_driven_by_form() -> None:
{"text-input": {"variable": "context", "label": "Context", "required": False}},
],
)
schema = build_input_schema(app, session=_session())
schema = build_input_schema(app, session=sqlite_session)
assert schema["properties"]["inputs"]["required"] == ["industry"]
def test_misconfigured_chat_raises_app_unavailable() -> None:
def test_misconfigured_chat_raises_app_unavailable(sqlite_session: Session) -> None:
app = _stub_app(AppMode.CHAT, has_workflow=False)
with pytest.raises(AppUnavailableError):
build_input_schema(app, session=_session())
build_input_schema(app, session=sqlite_session)
def test_misconfigured_workflow_raises_app_unavailable() -> None:
def test_misconfigured_workflow_raises_app_unavailable(sqlite_session: Session) -> None:
app = _stub_app(AppMode.WORKFLOW, has_workflow=False)
with pytest.raises(AppUnavailableError):
build_input_schema(app, session=_session())
build_input_schema(app, session=sqlite_session)
def test_empty_input_schema_sentinel_shape() -> None:
@@ -160,7 +160,7 @@ class TestAudioServiceMockedBehavior:
return mock
@patch.object(AudioService, "transcript_asr")
def test_transcript_asr_returns_response(self, mock_asr, mock_app, mock_file):
def test_transcript_asr_returns_response(self, mock_asr, mock_app, mock_file, sqlite_session: Session):
"""Test ASR transcription returns response dict."""
mock_response = {"text": "Transcribed text"}
mock_asr.return_value = mock_response
@@ -168,14 +168,13 @@ class TestAudioServiceMockedBehavior:
result = AudioService.transcript_asr(
app_model=mock_app,
file=mock_file,
session=Mock(),
session=sqlite_session,
end_user="user_123",
)
assert result["text"] == "Transcribed text"
@patch.object(AudioService, "transcript_tts")
@pytest.mark.parametrize("sqlite_session", [()], indirect=True)
def test_transcript_tts_returns_response(self, mock_tts, mock_app, sqlite_session: Session):
"""Test TTS transcription returns response."""
mock_response = {"audio": "base64_audio_data"}
@@ -494,7 +494,6 @@ class TestConversationService:
assert hasattr(result, "limit")
assert hasattr(result, "has_more")
@pytest.mark.parametrize("sqlite_session", [(Conversation,)], indirect=True)
def test_rename_returns_conversation(self, sqlite_session: Session):
"""Test rename returns updated conversation."""
conversation_id = "00000000-0000-0000-0000-000000000001"
@@ -542,7 +541,6 @@ class TestConversationApiController:
with pytest.raises(NotChatAppError):
handler(api, app_model=app_model, end_user=end_user)
@pytest.mark.parametrize("sqlite_session", [(Conversation,)], indirect=True)
def test_list_last_not_found(
self,
app: Flask,
@@ -615,7 +615,6 @@ class TestHitlServiceApi:
assert response.data.paused_nodes == ["node-1"]
assert response.data.reasons == [{"TYPE": "human_input_required", "form_id": "form-1", "expiration_time": 1}]
@pytest.mark.parametrize("sqlite_session", [(HumanInputForm,)], indirect=True)
def test_service_api_pause_event_serializes_hitl_reason(
self,
monkeypatch: pytest.MonkeyPatch,
@@ -630,7 +629,7 @@ class TestHitlServiceApi:
reason=WorkflowStartReason.INITIAL,
)
expiration_time = datetime(2024, 1, 1)
expiration_time = datetime(2024, 1, 1, tzinfo=UTC)
_persist_human_input_form(sqlite_session, expiration_time=expiration_time)
monkeypatch.setattr(workflow_response_converter, "db", SimpleNamespace(engine=sqlite_engine))
@@ -693,7 +692,6 @@ class TestHitlServiceApi:
assert hi_resp.data.expiration_time == int(expiration_time.timestamp())
# Snapshot payload contract
@pytest.mark.parametrize("sqlite_session", [(HumanInputForm,)], indirect=True)
def test_snapshot_events_include_pause_payload_contract(
self,
monkeypatch: pytest.MonkeyPatch,
@@ -703,7 +701,7 @@ class TestHitlServiceApi:
workflow_run = _build_workflow_run(WorkflowExecutionStatus.PAUSED)
snapshot = _build_snapshot(WorkflowNodeExecutionStatus.PAUSED)
resumption_context = _build_resumption_context("task-ctx")
expiration_time = datetime(2024, 1, 1)
expiration_time = datetime(2024, 1, 1, tzinfo=UTC)
_persist_human_input_form(sqlite_session, expiration_time=expiration_time)
monkeypatch.setattr(
"services.workflow_event_snapshot_service.load_form_dispositions_by_form_id",
@@ -574,7 +574,6 @@ class TestPipelineRunApiPost:
)
@patch("controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.RagPipelineService")
@patch("controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.service_api_ns")
@pytest.mark.parametrize("sqlite_session", [(Dataset,)], indirect=True)
def test_post_success_streaming(
self, mock_ns, mock_svc_cls, mock_current_user, mock_gen_svc, mock_helper, app, sqlite_session: Session
):
@@ -609,7 +608,6 @@ class TestPipelineRunApiPost:
mock_svc_cls.assert_called_once_with(sqlite_session)
mock_gen_svc.generate.assert_called_once()
@pytest.mark.parametrize("sqlite_session", [(Dataset,)], indirect=True)
def test_post_not_found(self, app: Flask, sqlite_session: Session):
"""Test NotFound when dataset check fails."""
with app.test_request_context("/datasets/test/pipeline/run", method="POST"):
@@ -624,7 +622,6 @@ class TestPipelineRunApiPost:
@patch("controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.current_user", new="not_account")
@patch("controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.service_api_ns")
@pytest.mark.parametrize("sqlite_session", [(Dataset,)], indirect=True)
def test_post_forbidden_non_account_user(self, mock_ns, app: Flask, sqlite_session: Session):
"""Test Forbidden when current_user is not an Account."""
tenant_id = str(uuid.uuid4())
@@ -20,10 +20,11 @@ import json
import uuid
from dataclasses import dataclass
from datetime import UTC, datetime
from unittest.mock import MagicMock, Mock, patch
from unittest.mock import Mock, patch
import pytest
from flask import Flask
from sqlalchemy.orm import Session
from werkzeug.exceptions import Forbidden, NotFound
from controllers.common.errors import FileTooLargeError as FileTooLargeHTTPError
@@ -44,8 +45,8 @@ from controllers.service_api.dataset.document import (
)
from controllers.service_api.dataset.error import ArchivedDocumentImmutableError
from core.rag.index_processor.constant.index_type import IndexStructureType
from models.dataset import Dataset, Document
from models.enums import DataSourceType, DocumentCreatedFrom, DocumentDocType, IndexingStatus
from models.dataset import Dataset, Document, DocumentSegment
from models.enums import DataSourceType, DocumentCreatedFrom, DocumentDocType, IndexingStatus, SegmentStatus
from services.dataset_service import DocumentService
from services.entities.knowledge_entities.knowledge_entities import ProcessRule, RetrievalModel
from services.errors.file import FileTooLargeError as FileTooLargeServiceError
@@ -176,6 +177,26 @@ def _expected_document_response(document: Document) -> dict[str, object]:
}
def _persist_segments(session: Session, document: Document, count: int = 5) -> None:
session.add_all(
[
DocumentSegment(
tenant_id=document.tenant_id,
dataset_id=document.dataset_id,
document_id=document.id,
position=index,
content=f"segment {index}",
word_count=20,
tokens=5,
created_by=document.created_by,
status=SegmentStatus.COMPLETED,
)
for index in range(1, count + 1)
]
)
session.flush()
class TestDocumentTextCreatePayload:
"""Test suite for DocumentTextCreatePayload Pydantic model."""
@@ -370,10 +391,10 @@ class TestDocumentService:
assert result.indexing_status == "completed"
@patch.object(DocumentService, "delete_document")
def test_delete_document_called(self, mock_delete):
def test_delete_document_called(self, mock_delete, sqlite_session: Session):
"""Test delete_document is called with document."""
document = make_serializable_document()
session = Mock()
session = sqlite_session
DocumentService.delete_document(document=document, session=session)
mock_delete.assert_called_once_with(document=document, session=session)
@@ -550,24 +571,23 @@ class TestDocumentDisplayStatusLogic:
class TestDocumentServiceBatchMethods:
"""Test DocumentService batch operations."""
def test_get_documents_by_ids(self):
def test_get_documents_by_ids(self, sqlite_session: Session):
"""Test batch retrieval of documents by IDs."""
dataset_id = str(uuid.uuid4())
doc_ids = [str(uuid.uuid4()), str(uuid.uuid4())]
session = Mock()
mock_result = Mock()
mock_result.all.return_value = [Mock(id=doc_ids[0]), Mock(id=doc_ids[1])]
session.scalars.return_value = mock_result
session = sqlite_session
session.add_all([make_serializable_document(id=document_id, dataset_id=dataset_id) for document_id in doc_ids])
session.flush()
documents = DocumentService.get_documents_by_ids(dataset_id, doc_ids, session)
assert len(documents) == 2
session.scalars.assert_called_once()
assert {document.id for document in documents} == set(doc_ids)
def test_get_documents_by_ids_empty(self):
def test_get_documents_by_ids_empty(self, sqlite_session: Session):
"""Test batch retrieval with empty list returns empty."""
assert DocumentService.get_documents_by_ids("ds_id", [], Mock()) == []
assert DocumentService.get_documents_by_ids("ds_id", [], sqlite_session) == []
class TestDocumentServiceFileOperations:
@@ -575,7 +595,7 @@ class TestDocumentServiceFileOperations:
@patch("services.dataset_service.file_helpers.get_signed_file_url")
@patch("services.dataset_service.DocumentService._get_upload_file_for_upload_file_document")
def test_get_document_download_url(self, mock_get_file, mock_signed_url):
def test_get_document_download_url(self, mock_get_file, mock_signed_url, sqlite_session: Session):
"""Test generation of download URL."""
mock_doc = Mock()
mock_file = Mock()
@@ -583,7 +603,7 @@ class TestDocumentServiceFileOperations:
mock_get_file.return_value = mock_file
mock_signed_url.return_value = "https://example.com/download"
session = Mock()
session = sqlite_session
url = DocumentService.get_document_download_url(mock_doc, session)
assert url == "https://example.com/download"
@@ -597,7 +617,7 @@ class TestDocumentServiceSaveValidation:
@patch("services.dataset_service.DatasetService.check_doc_form")
@patch("services.dataset_service.FeatureService.get_features")
@patch("services.dataset_service.current_user")
def test_save_document_validates_doc_form(self, mock_user, mock_features, mock_check_form):
def test_save_document_validates_doc_form(self, mock_user, mock_features, mock_check_form, sqlite_session: Session):
"""Test that doc_form is validated during save."""
mock_user.current_tenant_id = "tenant_id"
dataset = Mock()
@@ -610,7 +630,7 @@ class TestDocumentServiceSaveValidation:
pass
mock_check_form.side_effect = TestStopError()
session = Mock()
session = sqlite_session
# Skip actual logic by mocking dependent calls or raising error to stop early
with pytest.raises(TestStopError):
@@ -663,6 +683,7 @@ class TestDocumentApiGet:
app: Flask,
mock_tenant: str,
mock_doc_detail: Document,
sqlite_session: Session,
) -> None:
"""Test successful document retrieval with metadata='all'."""
# Arrange
@@ -670,9 +691,10 @@ class TestDocumentApiGet:
mock_dataset = make_dataset(id=dataset_id, tenant_id=mock_tenant, summary_index_setting=None)
mock_doc_svc.get_document.return_value = mock_doc_detail
mock_dataset_svc.get_dataset_for_tenant.return_value = mock_dataset
mock_dataset_svc.get_process_rules.return_value = {"mode": "automatic", "rules": {}}
session = MagicMock()
session.scalar.side_effect = [5, 0]
session = sqlite_session
_persist_segments(session, mock_doc_detail)
# Act
with app.test_request_context(
@@ -726,15 +748,16 @@ class TestDocumentApiGet:
assert response["summary_index_status"] is None
@patch("controllers.service_api.dataset.document.DocumentService")
def test_get_document_not_found(self, mock_doc_svc: Mock, app: Flask, mock_tenant: str) -> None:
def test_get_document_not_found(
self, mock_doc_svc: Mock, app: Flask, mock_tenant: str, sqlite_session: Session
) -> None:
"""Test 404 when document is not found."""
# Arrange
dataset_id = str(uuid.uuid4())
mock_dataset = make_dataset(id=dataset_id, tenant_id=mock_tenant)
mock_doc_svc.get_document.return_value = None
session = MagicMock()
session.scalar.return_value = mock_dataset
session = sqlite_session
# Act & Assert
with app.test_request_context(
@@ -754,7 +777,12 @@ class TestDocumentApiGet:
@patch("controllers.service_api.dataset.document.DocumentService")
def test_get_document_forbidden_wrong_tenant(
self, mock_doc_svc: Mock, app: Flask, mock_tenant: str, mock_doc_detail: Document
self,
mock_doc_svc: Mock,
app: Flask,
mock_tenant: str,
mock_doc_detail: Document,
sqlite_session: Session,
) -> None:
"""Test 403 when document tenant doesn't match request tenant."""
# Arrange
@@ -763,8 +791,9 @@ class TestDocumentApiGet:
mock_doc_detail.tenant_id = "different-tenant-id"
mock_doc_svc.get_document.return_value = mock_doc_detail
session = MagicMock()
session.scalar.return_value = mock_dataset
session = sqlite_session
session.add(mock_dataset)
session.flush()
# Act & Assert
with app.test_request_context(
@@ -784,7 +813,12 @@ class TestDocumentApiGet:
@patch("controllers.service_api.dataset.document.DocumentService")
def test_get_document_metadata_only(
self, mock_doc_svc: Mock, app: Flask, mock_tenant: str, mock_doc_detail: Document
self,
mock_doc_svc: Mock,
app: Flask,
mock_tenant: str,
mock_doc_detail: Document,
sqlite_session: Session,
) -> None:
"""Test document retrieval with metadata='only'."""
# Arrange
@@ -792,8 +826,9 @@ class TestDocumentApiGet:
mock_dataset = make_dataset(id=dataset_id, tenant_id=mock_tenant, summary_index_setting=None)
mock_doc_svc.get_document.return_value = mock_doc_detail
session = MagicMock()
session.scalar.return_value = mock_dataset
session = sqlite_session
session.add(mock_dataset)
session.flush()
# Act
with app.test_request_context(
@@ -827,6 +862,7 @@ class TestDocumentApiGet:
app: Flask,
mock_tenant: str,
mock_doc_detail: Document,
sqlite_session: Session,
) -> None:
"""Test document retrieval with metadata='without'."""
# Arrange
@@ -834,9 +870,10 @@ class TestDocumentApiGet:
mock_dataset = make_dataset(id=dataset_id, tenant_id=mock_tenant, summary_index_setting=None)
mock_doc_svc.get_document.return_value = mock_doc_detail
mock_dataset_svc.get_dataset_for_tenant.return_value = mock_dataset
mock_dataset_svc.get_process_rules.return_value = {"mode": "automatic", "rules": {}}
session = MagicMock()
session.scalar.side_effect = [5, 0]
session = sqlite_session
_persist_segments(session, mock_doc_detail)
# Act
with app.test_request_context(
@@ -895,7 +932,12 @@ class TestDocumentApiGet:
@patch("controllers.service_api.dataset.document.DocumentService")
def test_get_document_invalid_metadata_value(
self, mock_doc_svc: Mock, app: Flask, mock_tenant: str, mock_doc_detail: Document
self,
mock_doc_svc: Mock,
app: Flask,
mock_tenant: str,
mock_doc_detail: Document,
sqlite_session: Session,
) -> None:
"""Test error when metadata parameter has invalid value."""
# Arrange
@@ -903,8 +945,9 @@ class TestDocumentApiGet:
mock_dataset = make_dataset(id=dataset_id, tenant_id=mock_tenant, summary_index_setting=None)
mock_doc_svc.get_document.return_value = mock_doc_detail
session = MagicMock()
session.scalar.return_value = mock_dataset
session = sqlite_session
session.add(mock_dataset)
session.flush()
# Act & Assert
with app.test_request_context(
@@ -1464,6 +1507,7 @@ class TestDocumentUpdateByTextApiPost:
app: Flask,
mock_tenant,
mock_dataset,
sqlite_session: Session,
):
"""Test successful document update by text."""
_setup_billing_mocks(mock_validate_token, mock_feature_svc, mock_tenant)
@@ -1799,6 +1843,7 @@ class TestDocumentUpdateByFileApiPatch:
app: Flask,
mock_tenant,
mock_dataset,
sqlite_session: Session,
):
"""Test legacy POST aliases still dispatch while marked deprecated."""
_setup_billing_mocks(mock_validate_token, mock_feature_svc, mock_tenant)
@@ -1808,8 +1853,7 @@ class TestDocumentUpdateByFileApiPatch:
)
doc_id = str(uuid.uuid4())
session = MagicMock()
session.scalar.return_value = 0
session = sqlite_session
with app.test_request_context(
f"/datasets/{mock_dataset.id}/documents/{doc_id}/{route_name}",
method="POST",
@@ -17,10 +17,11 @@ Decorator strategy:
import uuid
from inspect import unwrap
from unittest.mock import MagicMock, Mock, patch
from unittest.mock import Mock, patch
import pytest
from flask import Flask
from sqlalchemy.orm import Session
from werkzeug.exceptions import NotFound
from controllers.service_api.dataset.metadata import (
@@ -56,7 +57,15 @@ def mock_dataset():
# ---------------------------------------------------------------------------
class TestDatasetMetadataCreatePost:
class _UsesSQLiteSession:
session: Session
@pytest.fixture(autouse=True)
def _inject_sqlite_session(self, sqlite_session: Session) -> None:
self.session = sqlite_session
class TestDatasetMetadataCreatePost(_UsesSQLiteSession):
"""Tests for DatasetMetadataCreateServiceApi.post().
``post`` is wrapped by ``@cloud_edition_billing_rate_limit_check``
@@ -64,7 +73,7 @@ class TestDatasetMetadataCreatePost:
"""
@staticmethod
def _call_post(api, session: MagicMock, **kwargs):
def _call_post(api, session: Session, **kwargs):
return unwrap(api.post)(api, session, **kwargs)
@patch("controllers.service_api.dataset.metadata.MetadataService")
@@ -91,7 +100,7 @@ class TestDatasetMetadataCreatePost:
json={"type": "string", "name": "Author"},
):
api = DatasetMetadataCreateServiceApi()
session = MagicMock()
session = self.session
response, status = self._call_post(
api,
session,
@@ -120,7 +129,7 @@ class TestDatasetMetadataCreatePost:
json={"type": "string", "name": "Author"},
):
api = DatasetMetadataCreateServiceApi()
session = MagicMock()
session = self.session
with pytest.raises(NotFound):
self._call_post(
api,
@@ -130,7 +139,7 @@ class TestDatasetMetadataCreatePost:
)
class TestDatasetMetadataCreateGet:
class TestDatasetMetadataCreateGet(_UsesSQLiteSession):
"""Tests for DatasetMetadataCreateServiceApi.get()."""
@patch("controllers.service_api.dataset.metadata.MetadataService")
@@ -191,14 +200,14 @@ class TestDatasetMetadataCreateGet:
# ---------------------------------------------------------------------------
class TestDatasetMetadataServiceApiPatch:
class TestDatasetMetadataServiceApiPatch(_UsesSQLiteSession):
"""Tests for DatasetMetadataServiceApi.patch().
``patch`` is wrapped by ``@cloud_edition_billing_rate_limit_check``.
"""
@staticmethod
def _call_patch(api, session: MagicMock, **kwargs):
def _call_patch(api, session: Session, **kwargs):
return unwrap(api.patch)(api, session, **kwargs)
@patch("controllers.service_api.dataset.metadata.MetadataService")
@@ -225,7 +234,7 @@ class TestDatasetMetadataServiceApiPatch:
json={"name": "New Name"},
):
api = DatasetMetadataServiceApi()
session = MagicMock()
session = self.session
response, status = self._call_patch(
api,
session,
@@ -256,7 +265,7 @@ class TestDatasetMetadataServiceApiPatch:
json={"name": "x"},
):
api = DatasetMetadataServiceApi()
session = MagicMock()
session = self.session
with pytest.raises(NotFound):
self._call_patch(
api,
@@ -267,14 +276,14 @@ class TestDatasetMetadataServiceApiPatch:
)
class TestDatasetMetadataServiceApiDelete:
class TestDatasetMetadataServiceApiDelete(_UsesSQLiteSession):
"""Tests for DatasetMetadataServiceApi.delete().
``delete`` is wrapped by ``@cloud_edition_billing_rate_limit_check``.
"""
@staticmethod
def _call_delete(api, session: MagicMock, **kwargs):
def _call_delete(api, session: Session, **kwargs):
return unwrap(api.delete)(api, session, **kwargs)
@patch("controllers.service_api.dataset.metadata.MetadataService")
@@ -300,7 +309,7 @@ class TestDatasetMetadataServiceApiDelete:
method="DELETE",
):
api = DatasetMetadataServiceApi()
session = MagicMock()
session = self.session
response = self._call_delete(
api,
session,
@@ -329,7 +338,7 @@ class TestDatasetMetadataServiceApiDelete:
method="DELETE",
):
api = DatasetMetadataServiceApi()
session = MagicMock()
session = self.session
with pytest.raises(NotFound):
self._call_delete(
api,
@@ -345,7 +354,7 @@ class TestDatasetMetadataServiceApiDelete:
# ---------------------------------------------------------------------------
class TestDatasetMetadataBuiltInFieldGet:
class TestDatasetMetadataBuiltInFieldGet(_UsesSQLiteSession):
"""Tests for DatasetMetadataBuiltInFieldServiceApi.get()."""
@patch("controllers.service_api.dataset.metadata.MetadataService")
@@ -380,14 +389,14 @@ class TestDatasetMetadataBuiltInFieldGet:
# ---------------------------------------------------------------------------
class TestDatasetMetadataBuiltInFieldAction:
class TestDatasetMetadataBuiltInFieldAction(_UsesSQLiteSession):
"""Tests for DatasetMetadataBuiltInFieldActionServiceApi.post().
``post`` is wrapped by ``@cloud_edition_billing_rate_limit_check``.
"""
@staticmethod
def _call_post(api, session: MagicMock, **kwargs):
def _call_post(api, session: Session, **kwargs):
return unwrap(api.post)(api, session, **kwargs)
@patch("controllers.service_api.dataset.metadata.MetadataService")
@@ -411,7 +420,7 @@ class TestDatasetMetadataBuiltInFieldAction:
method="POST",
):
api = DatasetMetadataBuiltInFieldActionServiceApi()
session = MagicMock()
session = self.session
response, status = self._call_post(
api,
session,
@@ -445,7 +454,7 @@ class TestDatasetMetadataBuiltInFieldAction:
method="POST",
):
api = DatasetMetadataBuiltInFieldActionServiceApi()
session = MagicMock()
session = self.session
response, status = self._call_post(
api,
session,
@@ -473,7 +482,7 @@ class TestDatasetMetadataBuiltInFieldAction:
method="POST",
):
api = DatasetMetadataBuiltInFieldActionServiceApi()
session = MagicMock()
session = self.session
with pytest.raises(NotFound):
self._call_post(
api,
@@ -489,14 +498,14 @@ class TestDatasetMetadataBuiltInFieldAction:
# ---------------------------------------------------------------------------
class TestDocumentMetadataEditPost:
class TestDocumentMetadataEditPost(_UsesSQLiteSession):
"""Tests for DocumentMetadataEditServiceApi.post().
``post`` is wrapped by ``@cloud_edition_billing_rate_limit_check``.
"""
@staticmethod
def _call_post(api, session: MagicMock, **kwargs):
def _call_post(api, session: Session, **kwargs):
return unwrap(api.post)(api, session, **kwargs)
@patch("controllers.service_api.dataset.metadata.MetadataService")
@@ -522,7 +531,7 @@ class TestDocumentMetadataEditPost:
json={"operation_data": []},
):
api = DocumentMetadataEditServiceApi()
session = MagicMock()
session = self.session
response, status = self._call_post(
api,
session,
@@ -550,7 +559,7 @@ class TestDocumentMetadataEditPost:
json={"operation_data": []},
):
api = DocumentMetadataEditServiceApi()
session = MagicMock()
session = self.session
with pytest.raises(NotFound):
self._call_post(
api,
@@ -9,6 +9,7 @@ from unittest.mock import MagicMock, patch
import pytest
from flask import Flask
from sqlalchemy.orm import Session
from controllers.web.completion import ChatApi, ChatStopApi, CompletionApi, CompletionStopApi
from controllers.web.error import (
@@ -168,9 +169,10 @@ class TestChatApi:
mock_get_conversation: MagicMock,
mock_generate: MagicMock,
app: Flask,
unbound_session: Session,
) -> None:
mock_ns.payload = {"inputs": {}, "query": "hi", "conversation_id": str(uuid.uuid4())}
session = MagicMock()
session = unbound_session
with app.test_request_context("/chat-messages", method="POST"):
unwrap(ChatApi.post)(ChatApi(), session, _chat_app(), _end_user())