mirror of
https://github.com/langgenius/dify.git
synced 2026-09-24 23:22:26 +08:00
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:
co-authored by
autofix-ci[bot]
parent
e3812cf72b
commit
08008f6305
@@ -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",
|
||||
|
||||
-3
@@ -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())
|
||||
|
||||
Reference in New Issue
Block a user