diff --git a/api/tests/unit_tests/controllers/web/test_human_input_form.py b/api/tests/unit_tests/controllers/web/test_human_input_form.py index bcafd269dfd..3408e3049d1 100644 --- a/api/tests/unit_tests/controllers/web/test_human_input_form.py +++ b/api/tests/unit_tests/controllers/web/test_human_input_form.py @@ -4,11 +4,13 @@ from __future__ import annotations from datetime import UTC, datetime from types import SimpleNamespace -from typing import Any from unittest.mock import MagicMock +from uuid import uuid4 import pytest from flask import Flask +from sqlalchemy import Engine +from sqlalchemy.orm import Session, sessionmaker from werkzeug.exceptions import Forbidden import controllers.web.human_input_form as human_input_module @@ -16,13 +18,21 @@ import controllers.web.site as site_module from controllers.web.error import WebFormRateLimitExceededError from core.workflow.nodes.human_input.entities import ParagraphInputConfig, SelectInputConfig, StringListSource from core.workflow.nodes.human_input.enums import ValueSourceType +from models import Tenant +from models.enums import CustomizeTokenStrategy from models.human_input import RecipientType +from models.model import App, AppMode, IconType, Site from services.feature_service import FeatureModel from services.human_input_service import FormExpiredError HumanInputFormApi = human_input_module.HumanInputFormApi HumanInputFormUploadTokenApi = human_input_module.HumanInputFormUploadTokenApi -TenantStatus = human_input_module.TenantStatus + +SQLITE_MODELS = (Tenant, App, Site) +pytestmark = [ + pytest.mark.usefixtures("sqlite_session"), + pytest.mark.parametrize("sqlite_session", [SQLITE_MODELS], indirect=True), +] @pytest.fixture @@ -34,36 +44,65 @@ def app() -> Flask: return app -class _FakeSession: - """Simple stand-in for db.session that returns pre-seeded objects.""" +@pytest.fixture +def database_session( + sqlite_session: Session, + sqlite_engine: Engine, + monkeypatch: pytest.MonkeyPatch, +) -> Session: + """Bind model/controller database access to the shared SQLite session.""" - def __init__(self, mapping: dict[str, Any]): - self._mapping = mapping - - def get(self, model, ident): - return self._mapping.get(model.__name__) - - def scalar(self, stmt): - # Extract the model name from the select statement's column_descriptions - try: - name = stmt.column_descriptions[0]["entity"].__name__ - except (AttributeError, IndexError, KeyError): - return None - return self._mapping.get(name) + database = SimpleNamespace(engine=sqlite_engine, session=sqlite_session) + monkeypatch.setattr(human_input_module, "db", database) + monkeypatch.setattr("models.model.db", database) + return sqlite_session -class _FakeDB: - """Minimal db stub exposing engine and session.""" - - def __init__(self, session: _FakeSession): - self.session = session - self.engine = object() +def _persist_app_site(session: Session, *, include_site: bool = True) -> tuple[Tenant, App, Site | None]: + tenant = Tenant(name="Tenant", plan="basic") + tenant.custom_config_dict = {"remove_webapp_brand": True, "replace_webapp_logo": None} + app_model = App( + id=str(uuid4()), + tenant_id=tenant.id, + name="Human Input App", + mode=AppMode.CHAT, + icon_type=IconType.EMOJI, + icon="robot", + icon_background="#fff", + enable_site=True, + enable_api=False, + ) + site = None + models: list[object] = [tenant, app_model] + if include_site: + site = Site( + app_id=app_model.id, + title="My Site", + default_language="en", + customize_token_strategy=CustomizeTokenStrategy.UUID, + icon_type=IconType.EMOJI, + icon="robot", + icon_background="#fff", + description="desc", + input_placeholder="Ask the app", + chat_color_theme="light", + chat_color_theme_inverted=False, + custom_disclaimer="", + prompt_public=False, + show_workflow_steps=True, + use_icon_as_answer_icon=False, + ) + models.append(site) + session.add_all(models) + session.commit() + return tenant, app_model, site -def test_get_form_includes_site(monkeypatch: pytest.MonkeyPatch, app: Flask): +def test_get_form_includes_site(monkeypatch: pytest.MonkeyPatch, app: Flask, database_session: Session): """GET returns form definition merged with site payload.""" expiration_time = datetime(2099, 1, 1, tzinfo=UTC) + _, app_model, _ = _persist_app_site(database_session) class _FakeDefinition: def model_dump(self, mode: str | None = None): @@ -77,9 +116,9 @@ def test_get_form_includes_site(monkeypatch: pytest.MonkeyPatch, app: Flask): class _FakeForm: def __init__(self, expiration: datetime): - self.workflow_run_id = "workflow-1" - self.app_id = "app-1" - self.tenant_id = "tenant-1" + self.workflow_run_id = None + self.app_id = app_model.id + self.tenant_id = app_model.tenant_id self.expiration_time = expiration self.recipient_type = RecipientType.BACKSTAGE @@ -92,32 +131,6 @@ def test_get_form_includes_site(monkeypatch: pytest.MonkeyPatch, app: Flask): monkeypatch.setattr(human_input_module, "_FORM_ACCESS_RATE_LIMITER", limiter_mock) monkeypatch.setattr(human_input_module, "extract_remote_ip", lambda req: "203.0.113.10") - tenant = SimpleNamespace( - id="tenant-1", - status=TenantStatus.NORMAL, - plan="basic", - custom_config_dict={"remove_webapp_brand": True, "replace_webapp_logo": False}, - ) - app_model = SimpleNamespace(id="app-1", tenant_id="tenant-1", tenant=tenant, enable_site=True) - workflow_run = SimpleNamespace(app_id="app-1") - site_model = SimpleNamespace( - title="My Site", - icon_type="emoji", - icon="robot", - icon_background="#fff", - description="desc", - input_placeholder="Ask the app", - default_language="en", - chat_color_theme="light", - chat_color_theme_inverted=False, - copyright=None, - privacy_policy=None, - custom_disclaimer="", - prompt_public=False, - show_workflow_steps=True, - use_icon_as_answer_icon=False, - ) - # Patch service to return fake form. service_mock = MagicMock() service_mock.get_form_by_token.return_value = form @@ -125,10 +138,6 @@ def test_get_form_includes_site(monkeypatch: pytest.MonkeyPatch, app: Flask): service_mock.resolve_form_inputs.return_value = [resolved_input] monkeypatch.setattr(human_input_module, "HumanInputService", lambda engine: service_mock) - # Patch db session. - db_stub = _FakeDB(_FakeSession({"WorkflowRun": workflow_run, "App": app_model, "Site": site_model})) - monkeypatch.setattr(human_input_module, "db", db_stub) - monkeypatch.setattr( site_module.FeatureService, "get_features", @@ -153,7 +162,7 @@ def test_get_form_includes_site(monkeypatch: pytest.MonkeyPatch, app: Flask): assert body["user_actions"] == [{"id": "approve", "title": "Approve", "button_style": "default"}] assert body["expiration_time"] == int(expiration_time.timestamp()) assert body["site"] == { - "app_id": "app-1", + "app_id": app_model.id, "end_user_id": None, "enable_site": True, "site": { @@ -187,10 +196,11 @@ def test_get_form_includes_site(monkeypatch: pytest.MonkeyPatch, app: Flask): limiter_mock.increment_rate_limit.assert_called_once_with("203.0.113.10") -def test_get_form_uses_runtime_select_options(monkeypatch: pytest.MonkeyPatch, app: Flask): +def test_get_form_uses_runtime_select_options(monkeypatch: pytest.MonkeyPatch, app: Flask, database_session: Session): """GET returns variable-backed select options resolved from runtime state.""" expiration_time = datetime(2099, 1, 1, tzinfo=UTC) + _, app_model, _ = _persist_app_site(database_session) configured_inputs = [ { "type": "select", @@ -225,9 +235,9 @@ def test_get_form_uses_runtime_select_options(monkeypatch: pytest.MonkeyPatch, a class _FakeForm: def __init__(self, expiration: datetime): - self.workflow_run_id = "workflow-1" - self.app_id = "app-1" - self.tenant_id = "tenant-1" + self.workflow_run_id = None + self.app_id = app_model.id + self.tenant_id = app_model.tenant_id self.recipient_type = RecipientType.STANDALONE_WEB_APP self.expiration_time = expiration @@ -239,37 +249,11 @@ def test_get_form_uses_runtime_select_options(monkeypatch: pytest.MonkeyPatch, a monkeypatch.setattr(human_input_module, "_FORM_ACCESS_RATE_LIMITER", limiter_mock) monkeypatch.setattr(human_input_module, "extract_remote_ip", lambda req: "203.0.113.10") - tenant = SimpleNamespace( - id="tenant-1", - status=TenantStatus.NORMAL, - plan="basic", - custom_config_dict={}, - ) - app_model = SimpleNamespace(id="app-1", tenant_id="tenant-1", tenant=tenant, enable_site=True) - site_model = SimpleNamespace( - title="My Site", - icon_type="emoji", - icon="robot", - icon_background="#fff", - description="desc", - input_placeholder="Ask the app", - default_language="en", - chat_color_theme="light", - chat_color_theme_inverted=False, - copyright=None, - privacy_policy=None, - custom_disclaimer="", - prompt_public=False, - show_workflow_steps=True, - use_icon_as_answer_icon=False, - ) - form = _FakeForm(expiration_time) service_mock = MagicMock() service_mock.get_form_by_token.return_value = form service_mock.resolve_form_inputs.return_value = runtime_inputs monkeypatch.setattr(human_input_module, "HumanInputService", lambda engine: service_mock) - monkeypatch.setattr(human_input_module, "db", _FakeDB(_FakeSession({"App": app_model, "Site": site_model}))) def mock_get_features(tenant_id: str, exclude_vector_space: bool = False): return FeatureModel(can_replace_logo=True) @@ -284,7 +268,9 @@ def test_get_form_uses_runtime_select_options(monkeypatch: pytest.MonkeyPatch, a service_mock.resolve_form_inputs.assert_called_once_with(form) -def test_create_upload_token_returns_token_and_form_expiration(monkeypatch: pytest.MonkeyPatch, app: Flask): +def test_create_upload_token_returns_token_and_form_expiration( + monkeypatch: pytest.MonkeyPatch, app: Flask, sqlite_engine: Engine +): """POST returns a HITL upload token for an active form token.""" expiration_time = datetime(2099, 1, 1, tzinfo=UTC) @@ -312,7 +298,7 @@ def test_create_upload_token_returns_token_and_form_expiration(monkeypatch: pyte "HumanInputFileUploadService", _service_factory, ) - monkeypatch.setattr(human_input_module, "db", SimpleNamespace(engine=object())) + monkeypatch.setattr(human_input_module, "db", SimpleNamespace(engine=sqlite_engine)) limiter_mock = MagicMock() limiter_mock.is_rate_limited.return_value = False @@ -329,14 +315,18 @@ def test_create_upload_token_returns_token_and_form_expiration(monkeypatch: pyte } repo_factory.assert_called_once() assert captured["workflow_run_repository"] is workflow_run_repository + session_factory = captured["session_factory"] + assert isinstance(session_factory, sessionmaker) + assert session_factory.kw["bind"] is sqlite_engine service_mock.issue_upload_token.assert_called_once_with("token-1") limiter_mock.increment_rate_limit.assert_called_once_with("203.0.113.10") -def test_get_form_allows_backstage_token(monkeypatch: pytest.MonkeyPatch, app: Flask): +def test_get_form_allows_backstage_token(monkeypatch: pytest.MonkeyPatch, app: Flask, database_session: Session): """GET returns form payload for backstage token.""" expiration_time = datetime(2099, 1, 2, tzinfo=UTC) + _, app_model, _ = _persist_app_site(database_session) class _FakeDefinition: def model_dump(self, mode: str | None = None): @@ -350,9 +340,9 @@ def test_get_form_allows_backstage_token(monkeypatch: pytest.MonkeyPatch, app: F class _FakeForm: def __init__(self, expiration: datetime): - self.workflow_run_id = "workflow-1" - self.app_id = "app-1" - self.tenant_id = "tenant-1" + self.workflow_run_id = None + self.app_id = app_model.id + self.tenant_id = app_model.tenant_id self.expiration_time = expiration def get_definition(self): @@ -363,40 +353,11 @@ def test_get_form_allows_backstage_token(monkeypatch: pytest.MonkeyPatch, app: F limiter_mock.is_rate_limited.return_value = False monkeypatch.setattr(human_input_module, "_FORM_ACCESS_RATE_LIMITER", limiter_mock) monkeypatch.setattr(human_input_module, "extract_remote_ip", lambda req: "203.0.113.10") - tenant = SimpleNamespace( - id="tenant-1", - status=TenantStatus.NORMAL, - plan="basic", - custom_config_dict={"remove_webapp_brand": True, "replace_webapp_logo": False}, - ) - app_model = SimpleNamespace(id="app-1", tenant_id="tenant-1", tenant=tenant, enable_site=True) - workflow_run = SimpleNamespace(app_id="app-1") - site_model = SimpleNamespace( - title="My Site", - icon_type="emoji", - icon="robot", - icon_background="#fff", - description="desc", - input_placeholder="Ask the app", - default_language="en", - chat_color_theme="light", - chat_color_theme_inverted=False, - copyright=None, - privacy_policy=None, - custom_disclaimer="", - prompt_public=False, - show_workflow_steps=True, - use_icon_as_answer_icon=False, - ) - service_mock = MagicMock() service_mock.get_form_by_token.return_value = form service_mock.resolve_form_inputs.return_value = [] monkeypatch.setattr(human_input_module, "HumanInputService", lambda engine: service_mock) - db_stub = _FakeDB(_FakeSession({"WorkflowRun": workflow_run, "App": app_model, "Site": site_model})) - monkeypatch.setattr(human_input_module, "db", db_stub) - monkeypatch.setattr( site_module.FeatureService, "get_features", @@ -421,7 +382,7 @@ def test_get_form_allows_backstage_token(monkeypatch: pytest.MonkeyPatch, app: F assert body["user_actions"] == [] assert body["expiration_time"] == int(expiration_time.timestamp()) assert body["site"] == { - "app_id": "app-1", + "app_id": app_model.id, "end_user_id": None, "enable_site": True, "site": { @@ -455,10 +416,13 @@ def test_get_form_allows_backstage_token(monkeypatch: pytest.MonkeyPatch, app: F limiter_mock.increment_rate_limit.assert_called_once_with("203.0.113.10") -def test_get_form_raises_forbidden_when_site_missing(monkeypatch: pytest.MonkeyPatch, app: Flask): +def test_get_form_raises_forbidden_when_site_missing( + monkeypatch: pytest.MonkeyPatch, app: Flask, database_session: Session +): """GET raises Forbidden if site cannot be resolved.""" expiration_time = datetime(2099, 1, 3, tzinfo=UTC) + _, app_model, _ = _persist_app_site(database_session, include_site=False) class _FakeDefinition: def model_dump(self, mode: str | None = None): @@ -472,9 +436,9 @@ def test_get_form_raises_forbidden_when_site_missing(monkeypatch: pytest.MonkeyP class _FakeForm: def __init__(self, expiration: datetime): - self.workflow_run_id = "workflow-1" - self.app_id = "app-1" - self.tenant_id = "tenant-1" + self.workflow_run_id = None + self.app_id = app_model.id + self.tenant_id = app_model.tenant_id self.expiration_time = expiration def get_definition(self): @@ -485,17 +449,10 @@ def test_get_form_raises_forbidden_when_site_missing(monkeypatch: pytest.MonkeyP limiter_mock.is_rate_limited.return_value = False monkeypatch.setattr(human_input_module, "_FORM_ACCESS_RATE_LIMITER", limiter_mock) monkeypatch.setattr(human_input_module, "extract_remote_ip", lambda req: "203.0.113.10") - tenant = SimpleNamespace(status=TenantStatus.NORMAL) - app_model = SimpleNamespace(id="app-1", tenant_id="tenant-1", tenant=tenant) - workflow_run = SimpleNamespace(app_id="app-1") - service_mock = MagicMock() service_mock.get_form_by_token.return_value = form monkeypatch.setattr(human_input_module, "HumanInputService", lambda engine: service_mock) - db_stub = _FakeDB(_FakeSession({"WorkflowRun": workflow_run, "App": app_model, "Site": None})) - monkeypatch.setattr(human_input_module, "db", db_stub) - with app.test_request_context("/api/form/human_input/token-1", method="GET"): with pytest.raises(Forbidden): HumanInputFormApi().get("token-1") @@ -503,7 +460,7 @@ def test_get_form_raises_forbidden_when_site_missing(monkeypatch: pytest.MonkeyP limiter_mock.increment_rate_limit.assert_called_once_with("203.0.113.10") -def test_submit_form_accepts_backstage_token(monkeypatch: pytest.MonkeyPatch, app: Flask): +def test_submit_form_accepts_backstage_token(monkeypatch: pytest.MonkeyPatch, app: Flask, sqlite_engine: Engine): """POST forwards backstage submissions to the service.""" class _FakeForm: @@ -517,7 +474,7 @@ def test_submit_form_accepts_backstage_token(monkeypatch: pytest.MonkeyPatch, ap service_mock = MagicMock() service_mock.get_form_by_token.return_value = form monkeypatch.setattr(human_input_module, "HumanInputService", lambda engine: service_mock) - monkeypatch.setattr(human_input_module, "db", _FakeDB(_FakeSession({}))) + monkeypatch.setattr(human_input_module, "db", SimpleNamespace(engine=sqlite_engine)) with app.test_request_context( "/api/form/human_input/token-1", @@ -550,7 +507,6 @@ def test_submit_form_rate_limited(monkeypatch: pytest.MonkeyPatch, app: Flask): service_mock = MagicMock() service_mock.get_form_by_token.return_value = None monkeypatch.setattr(human_input_module, "HumanInputService", lambda engine: service_mock) - monkeypatch.setattr(human_input_module, "db", _FakeDB(_FakeSession({}))) with app.test_request_context( "/api/form/human_input/token-1", @@ -576,7 +532,6 @@ def test_get_form_rate_limited(monkeypatch: pytest.MonkeyPatch, app: Flask): service_mock = MagicMock() service_mock.get_form_by_token.return_value = None monkeypatch.setattr(human_input_module, "HumanInputService", lambda engine: service_mock) - monkeypatch.setattr(human_input_module, "db", _FakeDB(_FakeSession({}))) with app.test_request_context("/api/form/human_input/token-1", method="GET"): with pytest.raises(WebFormRateLimitExceededError): @@ -587,7 +542,7 @@ def test_get_form_rate_limited(monkeypatch: pytest.MonkeyPatch, app: Flask): service_mock.get_form_by_token.assert_not_called() -def test_get_form_raises_expired(monkeypatch: pytest.MonkeyPatch, app: Flask): +def test_get_form_raises_expired(monkeypatch: pytest.MonkeyPatch, app: Flask, sqlite_engine: Engine): class _FakeForm: pass @@ -600,7 +555,7 @@ def test_get_form_raises_expired(monkeypatch: pytest.MonkeyPatch, app: Flask): service_mock.get_form_by_token.return_value = form service_mock.ensure_form_active.side_effect = FormExpiredError("form-id") monkeypatch.setattr(human_input_module, "HumanInputService", lambda engine: service_mock) - monkeypatch.setattr(human_input_module, "db", _FakeDB(_FakeSession({}))) + monkeypatch.setattr(human_input_module, "db", SimpleNamespace(engine=sqlite_engine)) with app.test_request_context("/api/form/human_input/token-1", method="GET"): with pytest.raises(FormExpiredError):