refactor(api): decouple setup and spec endpoints (#40104)

This commit is contained in:
Byron.wang
2026-08-08 19:52:15 +08:00
committed by GitHub
parent 2a742486c5
commit 4e18ab0ad4
18 changed files with 1214 additions and 102 deletions
+30
View File
@@ -95,3 +95,33 @@ forbidden_modules =
services.workspace_member_role_resolver
sqlalchemy
werkzeug
[importlinter:contract:setup-service-boundary]
name = Setup application service is framework and persistence neutral
type = forbidden
source_modules =
services.setup_service
forbidden_modules =
configs
controllers
extensions
flask
models
repositories
sqlalchemy
werkzeug
[importlinter:contract:schema-definition-service-boundary]
name = Schema definition application service is framework and persistence neutral
type = forbidden
source_modules =
services.schema_definition_service
forbidden_modules =
configs
controllers
extensions
flask
models
repositories
sqlalchemy
werkzeug
+26 -38
View File
@@ -2,15 +2,16 @@ from typing import Literal
from flask import request
from pydantic import BaseModel, Field, field_validator
from sqlalchemy import select
from configs import dify_config
from controllers.fastopenapi import console_router
from enums.deployment_edition import DeploymentEdition
from extensions.ext_application_services import application_services
from libs.helper import EmailStr, extract_remote_ip
from libs.password import valid_password
from models.model import DifySetup, db
from services.account_service import RegisterService, TenantService
from services.setup_service import (
InitializationValidationRequiredError,
SetupAlreadyCompletedError,
SetupInput,
)
from .error import AlreadySetupError, NotInitValidateError
from .init_validate import get_init_validate_status
@@ -53,14 +54,12 @@ def get_setup_status_api() -> SetupStatusResponse:
Only bootstrap-safe status information should be returned by this endpoint.
"""
if dify_config.DEPLOYMENT_EDITION != DeploymentEdition.CLOUD:
setup_status = get_setup_status()
if setup_status and not isinstance(setup_status, bool):
return SetupStatusResponse(step="finished", setup_at=setup_status.setup_at.isoformat())
if setup_status:
return SetupStatusResponse(step="finished")
setup_status = application_services().setup.get_status()
if not setup_status.completed:
return SetupStatusResponse(step="not_started")
return SetupStatusResponse(step="finished")
setup_at = setup_status.setup_at.isoformat() if setup_status.setup_at is not None else None
return SetupStatusResponse(step="finished", setup_at=setup_at)
@console_router.post(
@@ -77,33 +76,22 @@ def setup_system(payload: SetupRequestPayload) -> SetupResponse:
Access is restricted by deployment mode (`SELF_HOSTED`), one-time setup guards,
and init-password validation rather than user session authentication.
"""
if get_setup_status():
raise AlreadySetupError()
try:
application_services().setup.initialize(
SetupInput(
email=payload.email,
name=payload.name,
password=payload.password,
ip_address=extract_remote_ip(request),
language=payload.language,
),
initialization_validated=get_init_validate_status(),
)
except SetupAlreadyCompletedError:
raise AlreadySetupError() from None
except InitializationValidationRequiredError:
raise NotInitValidateError() from None
tenant_count = TenantService.get_tenant_count(session=db.session())
if tenant_count > 0:
raise AlreadySetupError()
if not get_init_validate_status():
raise NotInitValidateError()
normalized_email = payload.email.lower()
RegisterService.setup(
email=normalized_email,
name=payload.name,
password=payload.password,
ip_address=extract_remote_ip(request),
language=payload.language,
session=db.session(),
)
mark_setup_completed()
return SetupResponse(result="success")
def get_setup_status() -> DifySetup | bool | None:
if dify_config.DEPLOYMENT_EDITION != DeploymentEdition.CLOUD:
return db.session.scalar(select(DifySetup).limit(1))
return True
+10 -22
View File
@@ -1,23 +1,18 @@
import logging
from collections.abc import Mapping
from http import HTTPStatus
from typing import Any
from flask_restx import Resource
from pydantic import Field, RootModel
from controllers.common.schema import register_response_schema_models
from controllers.console.wraps import (
account_initialization_required,
setup_required,
)
from core.schemas.schema_manager import SchemaManager
from controllers.console.flask_admission import console_account_admission
from extensions.ext_application_services import application_services
from fields.base import ResponseModel
from libs.login import login_required
from machinery.context import RequestContext
from . import console_ns
logger = logging.getLogger(__name__)
class SchemaDefinitionItemResponse(ResponseModel):
name: str
@@ -34,20 +29,13 @@ register_response_schema_models(console_ns, SchemaDefinitionItemResponse, Schema
@console_ns.route("/spec/schema-definitions")
class SpecSchemaDefinitionsApi(Resource):
@console_ns.response(200, "Success", console_ns.models[SchemaDefinitionsResponse.__name__])
@setup_required
@login_required
@account_initialization_required
def get(self):
@console_ns.response(HTTPStatus.OK, "Success", console_ns.models[SchemaDefinitionsResponse.__name__])
@console_account_admission()
def get(self, _request_context: RequestContext):
"""
Get system JSON Schema definitions specification
Used for frontend component type mapping
"""
try:
schema_manager = SchemaManager()
schema_definitions = schema_manager.get_all_schema_definitions()
return schema_definitions, 200
except Exception:
logger.exception("Failed to get schema definitions from local registry")
# Return empty array as fallback
return [], 200
schema_definitions = application_services().schema_definitions.list()
response = SchemaDefinitionsResponse.model_validate(schema_definitions).model_dump(mode="json")
return response, HTTPStatus.OK
@@ -6,13 +6,21 @@ from typing import cast
from flask import Flask, current_app
from sqlalchemy.orm import Session, sessionmaker
from configs import dify_config
from constants.dsl_version import CURRENT_APP_DSL_VERSION
from core.db.session_factory import get_session_maker
from core.schemas.schema_manager import SchemaManager
from enums.deployment_edition import DeploymentEdition
from extensions.ext_redis import RedisClientWrapper, redis_client
from repositories.installation_state_repository import InstallationStateRepository
from repositories.workspace_member_query_repository import WorkspaceMemberQueryRepository
from repositories.workspace_query_repository import WorkspaceQueryRepository
from services.feature_query_service import FeatureQueryService
from services.feature_service import FeatureService
from services.feature_service_gateway import FeatureServiceGateway
from services.schema_definition_service import SchemaDefinitionService
from services.setup_adapters import RedisSetupLock, RegisterServiceAccountProvisioner
from services.setup_service import SetupService
from services.workspace_member_query_service import WorkspaceMemberQueryService
from services.workspace_member_role_resolver import DeploymentWorkspaceMemberRoleResolver
from services.workspace_plan_gateway import DeploymentWorkspacePlanGateway
@@ -23,6 +31,8 @@ _EXTENSION_KEY = "application_services"
@dataclass(frozen=True, slots=True)
class ApplicationServices:
schema_definitions: SchemaDefinitionService
setup: SetupService
feature_queries: FeatureQueryService
workspace_queries: WorkspaceQueryService
workspace_member_queries: WorkspaceMemberQueryService
@@ -31,8 +41,18 @@ class ApplicationServices:
def build_application_services(
*,
database_client: sessionmaker[Session],
deployment_edition: DeploymentEdition,
redis: RedisClientWrapper,
) -> ApplicationServices:
installation_state = InstallationStateRepository(client=database_client)
return ApplicationServices(
schema_definitions=SchemaDefinitionService(source_factory=SchemaManager),
setup=SetupService(
state=installation_state,
accounts=RegisterServiceAccountProvisioner(client=database_client),
lock=RedisSetupLock(client=redis),
setup_required=deployment_edition != DeploymentEdition.CLOUD,
),
feature_queries=FeatureQueryService(
features=FeatureServiceGateway(),
trial_models=FeatureService.get_trial_models(),
@@ -56,6 +76,8 @@ def build_application_services(
def init_app(app: Flask) -> None:
app.extensions[_EXTENSION_KEY] = build_application_services(
database_client=get_session_maker(),
deployment_edition=dify_config.DEPLOYMENT_EDITION,
redis=redis_client,
)
@@ -0,0 +1,27 @@
"""Persistence adapter for installation setup and tenant-existence state."""
from datetime import datetime
from sqlalchemy import exists, select
from sqlalchemy.orm import Session, sessionmaker
from models.account import Tenant
from models.model import DifySetup
class InstallationStateRepository:
"""Read persistent state shared by installation bootstrap use cases."""
def __init__(self, client: sessionmaker[Session]) -> None:
self._client = client
def get_setup_at(self) -> datetime | None:
with self._client() as session:
return session.scalar(select(DifySetup.setup_at).limit(1))
def is_setup(self) -> bool:
return self.get_setup_at() is not None
def has_tenants(self) -> bool:
with self._client() as session:
return session.scalar(select(exists().select_from(Tenant))) is True
+24
View File
@@ -0,0 +1,24 @@
"""Application service for querying Console schema definitions."""
import logging
from collections.abc import Callable, Mapping
from typing import Any, Protocol
logger = logging.getLogger(__name__)
class SchemaDefinitionSource(Protocol):
def get_all_schema_definitions(self, version: str = "v1") -> list[Mapping[str, Any]]: ...
class SchemaDefinitionService:
def __init__(self, *, source_factory: Callable[[], SchemaDefinitionSource]) -> None:
self._source_factory = source_factory
def list(self) -> tuple[Mapping[str, Any], ...]:
try:
source = self._source_factory()
return tuple(source.get_all_schema_definitions())
except Exception:
logger.exception("Failed to get schema definitions from local registry")
return ()
+43
View File
@@ -0,0 +1,43 @@
"""Infrastructure adapters for the first-time setup application service."""
from contextlib import AbstractContextManager
from typing import override
from sqlalchemy.orm import Session, sessionmaker
from extensions.ext_redis import RedisClientWrapper
from services.account_service import RegisterService
from services.setup_service import SetupAccountProvisioner, SetupInput, SetupLock
_SETUP_LOCK_KEY = "setup:initialize"
_SETUP_LOCK_TIMEOUT_SECONDS = 300
class RegisterServiceAccountProvisioner(SetupAccountProvisioner):
def __init__(self, client: sessionmaker[Session]) -> None:
self._client = client
@override
def provision(self, setup: SetupInput) -> None:
with self._client() as session:
RegisterService.setup(
email=setup.email,
name=setup.name,
password=setup.password,
ip_address=setup.ip_address,
language=setup.language,
session=session,
)
class RedisSetupLock(SetupLock):
def __init__(self, *, client: RedisClientWrapper) -> None:
self._client = client
@override
def acquire(self) -> AbstractContextManager[None]:
return self._client.lock(
_SETUP_LOCK_KEY,
timeout=_SETUP_LOCK_TIMEOUT_SECONDS,
blocking_timeout=_SETUP_LOCK_TIMEOUT_SECONDS,
)
+83
View File
@@ -0,0 +1,83 @@
"""Application service for first-time Dify setup."""
from contextlib import AbstractContextManager
from dataclasses import dataclass
from datetime import datetime
from typing import Protocol
@dataclass(frozen=True, slots=True)
class SetupInput:
email: str
name: str
password: str
ip_address: str
language: str | None
@dataclass(frozen=True, slots=True)
class SetupStatus:
completed: bool
setup_at: datetime | None = None
class SetupState(Protocol):
def get_setup_at(self) -> datetime | None: ...
def has_tenants(self) -> bool: ...
class SetupAccountProvisioner(Protocol):
def provision(self, setup: SetupInput) -> None: ...
class SetupLock(Protocol):
def acquire(self) -> AbstractContextManager[None]: ...
class SetupAlreadyCompletedError(Exception):
"""Raised when setup has already created persistent installation state."""
class InitializationValidationRequiredError(Exception):
"""Raised when initialization-password validation has not completed."""
class SetupService:
def __init__(
self,
*,
state: SetupState,
accounts: SetupAccountProvisioner,
lock: SetupLock,
setup_required: bool,
) -> None:
self._state = state
self._accounts = accounts
self._lock = lock
self._setup_required = setup_required
def get_status(self) -> SetupStatus:
if not self._setup_required:
return SetupStatus(completed=True)
setup_at = self._state.get_setup_at()
return SetupStatus(completed=setup_at is not None, setup_at=setup_at)
def initialize(self, setup: SetupInput, *, initialization_validated: bool) -> None:
with self._lock.acquire():
if self._state.get_setup_at() is not None or self._state.has_tenants():
raise SetupAlreadyCompletedError
if not initialization_validated:
raise InitializationValidationRequiredError
self._accounts.provision(
SetupInput(
email=setup.email.lower(),
name=setup.name,
password=setup.password,
ip_address=setup.ip_address,
language=setup.language,
)
)
@@ -0,0 +1,159 @@
from collections.abc import Iterator
from concurrent.futures import ThreadPoolExecutor
from contextlib import AbstractContextManager
from threading import Event, Lock
from unittest.mock import MagicMock, patch
import pytest
from faker import Faker
from flask import Flask
from flask.testing import FlaskClient
from sqlalchemy import func, select
from sqlalchemy.orm import Session
from models.account import Account, Tenant, TenantAccountJoin
from models.model import DifySetup
from services.account_service import RegisterService
from services.setup_adapters import RedisSetupLock
from tests.test_containers_integration_tests.helpers import generate_valid_password
@pytest.fixture
def setup_dependencies() -> Iterator[MagicMock]:
with (
patch("services.account_service.FeatureService") as feature_service,
patch("services.account_service.BillingService") as billing_service,
patch("services.account_service.CommunityTelemetryService.report_install") as report_install,
):
feature_service.get_system_features.return_value.is_allow_register = True
feature_service.get_license.return_value.seats.is_available.return_value = True
feature_service.get_license.return_value.workspaces.is_available.return_value = True
feature_service.is_workspace_creation_allowed.return_value = True
billing_service.is_email_in_freeze.return_value = False
yield report_install
def _setup_payload(*, email: str, password: str) -> dict[str, str]:
return {
"email": email,
"name": "Admin",
"password": password,
"language": "en-US",
}
def test_setup_endpoint_persists_bootstrap_state_and_rejects_repeat(
db_session_with_containers: Session,
test_client_with_containers: FlaskClient,
monkeypatch: pytest.MonkeyPatch,
setup_dependencies: MagicMock,
) -> None:
monkeypatch.delenv("INIT_PASSWORD", raising=False)
password = generate_valid_password(Faker())
response = test_client_with_containers.post(
"/console/api/setup",
json=_setup_payload(email="admin@example.com", password=password),
headers={"CF-Connecting-IP": "203.0.113.7"},
)
assert response.status_code == 201
assert response.get_json() == {"result": "success"}
setup_dependencies.assert_called_once()
repeated_response = test_client_with_containers.post(
"/console/api/setup",
json=_setup_payload(email="other@example.com", password=password),
)
assert repeated_response.status_code == 403
db_session_with_containers.expire_all()
assert db_session_with_containers.scalar(select(func.count()).select_from(DifySetup)) == 1
assert db_session_with_containers.scalar(select(func.count()).select_from(Account)) == 1
assert db_session_with_containers.scalar(select(func.count()).select_from(Tenant)) == 1
assert db_session_with_containers.scalar(select(func.count()).select_from(TenantAccountJoin)) == 1
account = db_session_with_containers.scalar(select(Account))
assert account is not None
assert account.email == "admin@example.com"
assert account.last_login_ip == "203.0.113.7"
def test_concurrent_setup_requests_create_only_one_bootstrap_identity(
flask_app_with_containers: Flask,
db_session_with_containers: Session,
monkeypatch: pytest.MonkeyPatch,
setup_dependencies: MagicMock,
) -> None:
monkeypatch.delenv("INIT_PASSWORD", raising=False)
password = generate_valid_password(Faker())
provision_started = Event()
allow_provision_to_finish = Event()
second_lock_attempted = Event()
attempt_guard = Lock()
lock_attempts = 0
original_setup = RegisterService.setup
original_acquire = RedisSetupLock.acquire
def blocking_setup(
email: str,
name: str,
password: str,
ip_address: str,
language: str | None,
*,
session: Session,
) -> None:
if not provision_started.is_set():
provision_started.set()
assert allow_provision_to_finish.wait(timeout=10)
original_setup(
email=email,
name=name,
password=password,
ip_address=ip_address,
language=language,
session=session,
)
def tracked_acquire(self: RedisSetupLock) -> AbstractContextManager[None]:
nonlocal lock_attempts
with attempt_guard:
lock_attempts += 1
if lock_attempts == 2:
second_lock_attempted.set()
return original_acquire(self)
def post_setup(email: str) -> int:
with flask_app_with_containers.test_client() as client:
response = client.post(
"/console/api/setup",
json=_setup_payload(email=email, password=password),
)
return response.status_code
with (
patch.object(RegisterService, "setup", side_effect=blocking_setup),
patch.object(RedisSetupLock, "acquire", tracked_acquire),
ThreadPoolExecutor(max_workers=2) as executor,
):
first_request = executor.submit(post_setup, "admin-1@example.com")
try:
assert provision_started.wait(timeout=10)
second_request = executor.submit(post_setup, "admin-2@example.com")
assert second_lock_attempted.wait(timeout=10)
assert not second_request.done()
finally:
allow_provision_to_finish.set()
results = [first_request.result(timeout=30), second_request.result(timeout=30)]
assert sorted(results) == [201, 403]
setup_dependencies.assert_called_once()
db_session_with_containers.expire_all()
assert db_session_with_containers.scalar(select(func.count()).select_from(DifySetup)) == 1
assert db_session_with_containers.scalar(select(func.count()).select_from(Account)) == 1
assert db_session_with_containers.scalar(select(func.count()).select_from(Tenant)) == 1
assert db_session_with_containers.scalar(select(func.count()).select_from(TenantAccountJoin)) == 1
@@ -0,0 +1,37 @@
from flask.testing import FlaskClient
from sqlalchemy.orm import Session
from tests.test_containers_integration_tests.controllers.console.helpers import (
authenticate_console_client,
create_console_account_and_tenant,
)
def test_schema_definitions_endpoint_uses_admission_and_builtin_registry(
db_session_with_containers: Session,
test_client_with_containers: FlaskClient,
) -> None:
account, _ = create_console_account_and_tenant(db_session_with_containers)
headers = authenticate_console_client(test_client_with_containers, account)
response = test_client_with_containers.get(
"/console/api/spec/schema-definitions",
headers=headers,
)
assert response.status_code == 200
definitions = response.get_json()
assert isinstance(definitions, list)
assert definitions
assert all({"name", "label", "schema"} <= definition.keys() for definition in definitions)
def test_schema_definitions_endpoint_rejects_unauthenticated_request(
db_session_with_containers: Session,
test_client_with_containers: FlaskClient,
) -> None:
create_console_account_and_tenant(db_session_with_containers)
response = test_client_with_containers.get("/console/api/spec/schema-definitions")
assert response.status_code == 401
@@ -1,40 +1,84 @@
import builtins
from unittest.mock import patch
from datetime import datetime
from types import SimpleNamespace
from unittest.mock import Mock, create_autospec
import pytest
from flask import Flask
from flask.views import MethodView
from controllers.console import setup as setup_controller
from controllers.console import wraps
from controllers.console.error import AlreadySetupError, NotInitValidateError
from dify_app import DifyApp
from extensions import ext_fastopenapi
from services.setup_service import (
InitializationValidationRequiredError,
SetupAlreadyCompletedError,
SetupInput,
SetupService,
SetupStatus,
)
if not hasattr(builtins, "MethodView"):
builtins.MethodView = MethodView # type: ignore[attr-defined]
@pytest.fixture
def app() -> Flask:
app = Flask(__name__)
def setup_service(monkeypatch: pytest.MonkeyPatch) -> Mock:
service = create_autospec(SetupService, instance=True, spec_set=True)
services = SimpleNamespace(setup=service)
monkeypatch.setattr(setup_controller, "application_services", lambda: services)
return service
@pytest.fixture
def app() -> DifyApp:
app = DifyApp(__name__)
app.config["TESTING"] = True
ext_fastopenapi.init_app(app)
return app
def test_console_setup_fastopenapi_get_not_started(app: Flask):
ext_fastopenapi.init_app(app)
def test_console_setup_fastopenapi_get_not_started(app: DifyApp, setup_service: Mock) -> None:
setup_service.get_status.return_value = SetupStatus(completed=False)
with (
patch("controllers.console.setup.dify_config.EDITION", "SELF_HOSTED"),
patch("controllers.console.setup.get_setup_status", return_value=None),
):
client = app.test_client()
response = client.get("/console/api/setup")
response = app.test_client().get("/console/api/setup")
assert response.status_code == 200
assert response.get_json() == {"step": "not_started", "setup_at": None}
def test_console_setup_fastopenapi_post_success(app: Flask):
ext_fastopenapi.init_app(app)
def test_console_setup_fastopenapi_get_finished(app: DifyApp, setup_service: Mock) -> None:
setup_at = datetime(2026, 8, 6, 10, 30)
setup_service.get_status.return_value = SetupStatus(completed=True, setup_at=setup_at)
response = app.test_client().get("/console/api/setup")
assert response.status_code == 200
assert response.get_json() == {"step": "finished", "setup_at": "2026-08-06T10:30:00"}
def test_console_setup_fastopenapi_get_finished_without_setup_time(app: DifyApp, setup_service: Mock) -> None:
setup_service.get_status.return_value = SetupStatus(completed=True)
response = app.test_client().get("/console/api/setup")
assert response.status_code == 200
assert response.get_json() == {"step": "finished", "setup_at": None}
@pytest.mark.parametrize("enterprise_enabled", [False, True], ids=["community", "enterprise"])
def test_console_setup_fastopenapi_post_success(
app: DifyApp,
setup_service: Mock,
monkeypatch: pytest.MonkeyPatch,
enterprise_enabled: bool,
) -> None:
monkeypatch.setattr(wraps.dify_config, "EDITION", "SELF_HOSTED")
monkeypatch.setattr(wraps.dify_config, "ENTERPRISE_ENABLED", enterprise_enabled)
monkeypatch.setattr(setup_controller, "get_init_validate_status", lambda: True)
mark_setup_completed = Mock()
monkeypatch.setattr(setup_controller, "mark_setup_completed", mark_setup_completed)
payload = {
"email": "admin@example.com",
"name": "Admin",
@@ -42,17 +86,162 @@ def test_console_setup_fastopenapi_post_success(app: Flask):
"language": "en-US",
}
with (
patch("controllers.console.wraps.dify_config.EDITION", "SELF_HOSTED"),
patch("controllers.console.setup.get_setup_status", return_value=None),
patch("controllers.console.setup.TenantService.get_tenant_count", return_value=0),
patch("controllers.console.setup.get_init_validate_status", return_value=True),
patch("controllers.console.setup.RegisterService.setup"),
patch("controllers.console.setup.mark_setup_completed") as mark_setup_completed,
):
client = app.test_client()
response = client.post("/console/api/setup", json=payload)
response = app.test_client().post(
"/console/api/setup",
json=payload,
headers={"CF-Connecting-IP": "203.0.113.7"},
)
assert response.status_code == 201
assert response.get_json() == {"result": "success"}
setup_service.initialize.assert_called_once_with(
SetupInput(
email="admin@example.com",
name="Admin",
password="Passw0rd1",
ip_address="203.0.113.7",
language="en-US",
),
initialization_validated=True,
)
mark_setup_completed.assert_called_once_with()
def test_console_setup_fastopenapi_post_rejects_cloud_edition(
app: DifyApp,
setup_service: Mock,
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setattr(wraps.dify_config, "EDITION", "CLOUD")
monkeypatch.setattr(wraps.dify_config, "ENTERPRISE_ENABLED", False)
response = app.test_client().post(
"/console/api/setup",
json={
"email": "admin@example.com",
"name": "Admin",
"password": "Passw0rd1",
"language": "en-US",
},
)
assert response.status_code == 404
setup_service.initialize.assert_not_called()
@pytest.mark.parametrize(
"payload",
[
pytest.param(
{
"email": "not-an-email",
"name": "Admin",
"password": "Passw0rd1",
"language": "en-US",
},
id="invalid-email",
),
pytest.param(
{
"email": "admin@example.com",
"name": "Admin",
"password": "short",
"language": "en-US",
},
id="invalid-password",
),
pytest.param(
{
"email": "admin@example.com",
"name": "a" * 31,
"password": "Passw0rd1",
"language": "en-US",
},
id="name-too-long",
),
],
)
def test_console_setup_fastopenapi_post_rejects_invalid_payload_before_service_call(
app: DifyApp,
setup_service: Mock,
monkeypatch: pytest.MonkeyPatch,
payload: dict[str, str],
) -> None:
monkeypatch.setattr(wraps.dify_config, "EDITION", "SELF_HOSTED")
monkeypatch.setattr(wraps.dify_config, "ENTERPRISE_ENABLED", False)
response = app.test_client().post("/console/api/setup", json=payload)
assert response.status_code == 422
setup_service.initialize.assert_not_called()
@pytest.mark.parametrize(
("service_error", "expected_controller_error"),
[
pytest.param(SetupAlreadyCompletedError(), AlreadySetupError, id="already-setup"),
pytest.param(
InitializationValidationRequiredError(),
NotInitValidateError,
id="init-validation-required",
),
],
)
def test_console_setup_translates_service_errors_to_controller_errors(
app: DifyApp,
setup_service: Mock,
monkeypatch: pytest.MonkeyPatch,
service_error: Exception,
expected_controller_error: type[Exception],
) -> None:
monkeypatch.setattr(wraps.dify_config, "EDITION", "SELF_HOSTED")
monkeypatch.setattr(wraps.dify_config, "ENTERPRISE_ENABLED", False)
monkeypatch.setattr(setup_controller, "get_init_validate_status", lambda: False)
mark_setup_completed = Mock()
monkeypatch.setattr(setup_controller, "mark_setup_completed", mark_setup_completed)
setup_service.initialize.side_effect = service_error
payload = setup_controller.SetupRequestPayload.model_validate(
{
"email": "admin@example.com",
"name": "Admin",
"password": "Passw0rd1",
"language": "en-US",
}
)
with app.test_request_context(
"/console/api/setup",
method="POST",
headers={"CF-Connecting-IP": "203.0.113.7"},
):
with pytest.raises(expected_controller_error) as raised:
setup_controller.setup_system(payload)
assert type(raised.value) is expected_controller_error
mark_setup_completed.assert_not_called()
def test_console_setup_fastopenapi_does_not_mark_setup_completed_when_service_fails(
app: DifyApp,
setup_service: Mock,
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setattr(wraps.dify_config, "EDITION", "SELF_HOSTED")
monkeypatch.setattr(wraps.dify_config, "ENTERPRISE_ENABLED", False)
monkeypatch.setattr(setup_controller, "get_init_validate_status", lambda: True)
mark_setup_completed = Mock()
monkeypatch.setattr(setup_controller, "mark_setup_completed", mark_setup_completed)
setup_service.initialize.side_effect = RuntimeError("provision failed")
response = app.test_client().post(
"/console/api/setup",
json={
"email": "admin@example.com",
"name": "Admin",
"password": "Passw0rd1",
"language": "en-US",
},
)
assert response.status_code == 500
mark_setup_completed.assert_not_called()
@@ -16,6 +16,8 @@ from services.entities.feature_entities import (
VectorSpaceLimitationModel,
)
from services.feature_query_service import FeatureQueryService
from services.schema_definition_service import SchemaDefinitionService
from services.setup_service import SetupService
from services.workspace_member_query_service import WorkspaceMemberQueryService
from services.workspace_query_service import WorkspaceQueryService
@@ -32,6 +34,8 @@ def _request_context() -> RequestContext:
def _install_application_services(mocker: MockerFixture):
feature_queries = create_autospec(FeatureQueryService, instance=True, spec_set=True)
services = ApplicationServices(
schema_definitions=create_autospec(SchemaDefinitionService, instance=True, spec_set=True),
setup=create_autospec(SetupService, instance=True, spec_set=True),
feature_queries=feature_queries,
workspace_queries=create_autospec(WorkspaceQueryService, instance=True, spec_set=True),
workspace_member_queries=create_autospec(WorkspaceMemberQueryService, instance=True, spec_set=True),
@@ -1,15 +1,24 @@
from inspect import unwrap
from unittest.mock import patch
import pytest
from types import SimpleNamespace
from unittest.mock import create_autospec, patch
import controllers.console.spec as spec_module
from dify_app import DifyApp
from extensions import ext_login
from machinery.context import RequestContext
from services.schema_definition_service import SchemaDefinitionService
class TestSpecSchemaDefinitionsApi:
def test_get_success(self):
def test_get_success(self) -> None:
api = spec_module.SpecSchemaDefinitionsApi()
method = unwrap(api.get)
request_context = RequestContext(
request_id="request-1",
trace_id="trace-1",
account_id="account-1",
active_workspace_id="workspace-1",
)
schema_definitions = [
{
@@ -23,34 +32,65 @@ class TestSpecSchemaDefinitionsApi:
}
]
service = create_autospec(SchemaDefinitionService, instance=True, spec_set=True)
service.list.return_value = tuple(schema_definitions)
with patch.object(
spec_module,
"SchemaManager",
) as schema_manager_cls:
schema_manager_cls.return_value.get_all_schema_definitions.return_value = schema_definitions
"application_services",
return_value=SimpleNamespace(schema_definitions=service),
):
resp, status = method(api, request_context)
resp, status = method(api)
assert status == 200
assert status == spec_module.HTTPStatus.OK
assert resp == schema_definitions
assert spec_module.SchemaDefinitionsResponse.model_validate(resp).model_dump(mode="json") == schema_definitions
service.list.assert_called_once_with()
def test_get_documents_tight_response_model(self):
def test_get_documents_tight_response_model(self) -> None:
response = spec_module.SpecSchemaDefinitionsApi.get.__apidoc__["responses"]["200"]
assert response[1].name == spec_module.SchemaDefinitionsResponse.__name__
def test_get_exception_returns_empty_list(self, caplog: pytest.LogCaptureFixture):
def test_get_returns_empty_list_from_service(self) -> None:
api = spec_module.SpecSchemaDefinitionsApi()
method = unwrap(api.get)
request_context = RequestContext(
request_id="request-1",
trace_id=None,
account_id="account-1",
active_workspace_id=None,
)
service = create_autospec(SchemaDefinitionService, instance=True, spec_set=True)
service.list.return_value = ()
with patch.object(
spec_module,
"SchemaManager",
side_effect=Exception("boom"),
"application_services",
return_value=SimpleNamespace(schema_definitions=service),
):
resp, status = method(api)
resp, status = method(api, request_context)
assert status == 200
assert status == spec_module.HTTPStatus.OK
assert resp == []
assert "boom" in caplog.text
def test_get_rejects_unauthenticated_request_before_service_call(self) -> None:
app = DifyApp(__name__)
app.config["TESTING"] = True
ext_login.init_app(app)
api = spec_module.SpecSchemaDefinitionsApi()
service = create_autospec(SchemaDefinitionService, instance=True, spec_set=True)
with (
app.test_request_context("/console/api/spec/schema-definitions"),
patch("controllers.console.wraps._is_setup_completed", return_value=True),
patch("libs.login._resolve_current_user", return_value=None),
patch.object(
spec_module,
"application_services",
return_value=SimpleNamespace(schema_definitions=service),
),
):
response = api.get()
assert response.status_code == 401
service.list.assert_not_called()
@@ -0,0 +1,58 @@
from unittest.mock import MagicMock, patch
import pytest
from sqlalchemy.orm import Session, sessionmaker
from enums.deployment_edition import DeploymentEdition
from extensions.ext_application_services import build_application_services
from extensions.ext_redis import RedisClientWrapper
@pytest.mark.parametrize(
("deployment_edition", "setup_completed"),
[
pytest.param(DeploymentEdition.CLOUD, True, id="cloud"),
pytest.param(DeploymentEdition.COMMUNITY, False, id="community"),
pytest.param(DeploymentEdition.ENTERPRISE, False, id="enterprise"),
],
)
def test_build_application_services_configures_setup_policy(
sqlite_session_factory: sessionmaker[Session],
deployment_edition: DeploymentEdition,
setup_completed: bool,
) -> None:
services = build_application_services(
database_client=sqlite_session_factory,
deployment_edition=deployment_edition,
redis=MagicMock(spec=RedisClientWrapper),
)
assert services.setup.get_status().completed is setup_completed
def test_build_application_services_wires_builtin_schema_definitions(
sqlite_session_factory: sessionmaker[Session],
) -> None:
services = build_application_services(
database_client=sqlite_session_factory,
deployment_edition=DeploymentEdition.COMMUNITY,
redis=MagicMock(spec=RedisClientWrapper),
)
definitions = services.schema_definitions.list()
assert definitions
assert all({"name", "label", "schema"} <= definition.keys() for definition in definitions)
def test_build_application_services_does_not_construct_schema_manager(
sqlite_session_factory: sessionmaker[Session],
) -> None:
with patch("extensions.ext_application_services.SchemaManager") as schema_manager:
build_application_services(
database_client=sqlite_session_factory,
deployment_edition=DeploymentEdition.COMMUNITY,
redis=MagicMock(spec=RedisClientWrapper),
)
schema_manager.assert_not_called()
@@ -0,0 +1,40 @@
"""Persistence tests for installation setup and tenant-existence state."""
from sqlalchemy.orm import Session, sessionmaker
from models.account import Tenant
from models.model import DifySetup
from repositories.installation_state_repository import InstallationStateRepository
def test_empty_database_has_no_installation_state(sqlite_session_factory: sessionmaker[Session]) -> None:
repository = InstallationStateRepository(client=sqlite_session_factory)
assert repository.get_setup_at() is None
assert repository.is_setup() is False
assert repository.has_tenants() is False
def test_get_setup_at_returns_persisted_timestamp(
sqlite_session: Session,
sqlite_session_factory: sessionmaker[Session],
) -> None:
setup = DifySetup(version="test-version")
sqlite_session.add(setup)
sqlite_session.commit()
sqlite_session.refresh(setup)
repository = InstallationStateRepository(client=sqlite_session_factory)
assert repository.get_setup_at() == setup.setup_at
assert repository.is_setup() is True
def test_has_tenants_detects_existing_tenant(
sqlite_session: Session,
sqlite_session_factory: sessionmaker[Session],
) -> None:
sqlite_session.add(Tenant(name="Existing workspace"))
sqlite_session.commit()
repository = InstallationStateRepository(client=sqlite_session_factory)
assert repository.has_tenants() is True
@@ -0,0 +1,57 @@
from collections.abc import Mapping
from unittest.mock import Mock, create_autospec
import pytest
from services.schema_definition_service import (
SchemaDefinitionService,
SchemaDefinitionSource,
)
@pytest.fixture
def source() -> Mock:
return create_autospec(SchemaDefinitionSource, instance=True, spec_set=True)
@pytest.fixture
def source_factory(source: Mock) -> Mock:
return Mock(return_value=source)
def test_list_returns_schema_definitions(source: Mock, source_factory: Mock) -> None:
definitions: list[Mapping[str, object]] = [
{
"name": "conversation-variable",
"label": "Conversation variable",
"schema": {"type": "object"},
}
]
source.get_all_schema_definitions.return_value = definitions
service = SchemaDefinitionService(source_factory=source_factory)
assert service.list() == tuple(definitions)
source_factory.assert_called_once_with()
source.get_all_schema_definitions.assert_called_once_with()
def test_list_returns_empty_tuple_when_source_query_fails(
source: Mock,
source_factory: Mock,
caplog: pytest.LogCaptureFixture,
) -> None:
source.get_all_schema_definitions.side_effect = RuntimeError("boom")
service = SchemaDefinitionService(source_factory=source_factory)
assert service.list() == ()
assert "Failed to get schema definitions from local registry" in caplog.text
assert "boom" in caplog.text
def test_list_returns_empty_tuple_when_source_construction_fails(caplog: pytest.LogCaptureFixture) -> None:
source_factory = Mock(side_effect=RuntimeError("construction failed"))
service = SchemaDefinitionService(source_factory=source_factory)
assert service.list() == ()
assert "Failed to get schema definitions from local registry" in caplog.text
assert "construction failed" in caplog.text
@@ -0,0 +1,74 @@
from contextlib import nullcontext
from unittest.mock import ANY, MagicMock, patch
import pytest
from redis.exceptions import ConnectionError as RedisConnectionError
from redis.exceptions import LockError
from sqlalchemy.orm import Session, sessionmaker
from extensions.ext_redis import RedisClientWrapper
from services.setup_adapters import RedisSetupLock, RegisterServiceAccountProvisioner
from services.setup_service import SetupInput
def test_provision_delegates_to_register_service_with_managed_session(
sqlite_session_factory: sessionmaker[Session],
) -> None:
provisioner = RegisterServiceAccountProvisioner(client=sqlite_session_factory)
setup = SetupInput(
email="admin@example.com",
name="Admin",
password="Passw0rd1",
ip_address="203.0.113.7",
language="en-US",
)
with patch("services.setup_adapters.RegisterService.setup") as register:
provisioner.provision(setup)
register.assert_called_once_with(
email="admin@example.com",
name="Admin",
password="Passw0rd1",
ip_address="203.0.113.7",
language="en-US",
session=ANY,
)
assert isinstance(register.call_args.kwargs["session"], Session)
def test_acquire_uses_bounded_distributed_lock() -> None:
redis = MagicMock(spec=RedisClientWrapper)
redis.lock.return_value = nullcontext()
lock = RedisSetupLock(client=redis)
with lock.acquire():
pass
redis.lock.assert_called_once_with(
"setup:initialize",
timeout=300,
blocking_timeout=300,
)
@pytest.mark.parametrize(
"error",
[
pytest.param(LockError("lock acquisition timed out"), id="timeout"),
pytest.param(RedisConnectionError("redis unavailable"), id="connection"),
],
)
def test_acquire_propagates_distributed_lock_failure(error: Exception) -> None:
redis = MagicMock(spec=RedisClientWrapper)
lock_context = MagicMock()
lock_context.__enter__.side_effect = error
redis.lock.return_value = lock_context
lock = RedisSetupLock(client=redis)
with pytest.raises(type(error), match=str(error)) as raised:
with lock.acquire():
pytest.fail("lock body must not run")
assert raised.value is error
lock_context.__exit__.assert_not_called()
@@ -0,0 +1,249 @@
from collections.abc import Generator
from contextlib import contextmanager, nullcontext
from dataclasses import dataclass
from datetime import datetime
from unittest.mock import Mock, create_autospec
import pytest
from services.setup_service import (
InitializationValidationRequiredError,
SetupAccountProvisioner,
SetupAlreadyCompletedError,
SetupInput,
SetupLock,
SetupService,
SetupState,
SetupStatus,
)
@pytest.fixture
def state() -> Mock:
state = create_autospec(SetupState, instance=True, spec_set=True)
state.get_setup_at.return_value = None
state.has_tenants.return_value = False
return state
@pytest.fixture
def accounts() -> Mock:
return create_autospec(SetupAccountProvisioner, instance=True, spec_set=True)
@pytest.fixture
def lock() -> Mock:
lock = create_autospec(SetupLock, instance=True, spec_set=True)
lock.acquire.return_value = nullcontext()
return lock
@pytest.fixture
def setup_input() -> SetupInput:
return SetupInput(
email="Admin@Example.com",
name="Admin",
password="Passw0rd1",
ip_address="203.0.113.7",
language="en-US",
)
@dataclass
class TrackingLock:
inside: bool = False
exited: bool = False
@contextmanager
def acquire(self) -> Generator[None]:
self.inside = True
try:
yield
finally:
self.inside = False
self.exited = True
class FailingLock:
def __init__(self, error: Exception) -> None:
self._error = error
@contextmanager
def acquire(self) -> Generator[None]:
raise self._error
yield
def test_cloud_status_is_finished_without_reading_persistence(state: Mock, accounts: Mock, lock: Mock) -> None:
service = SetupService(state=state, accounts=accounts, lock=lock, setup_required=False)
assert service.get_status() == SetupStatus(completed=True)
state.get_setup_at.assert_not_called()
def test_self_hosted_status_is_not_started_without_setup(state: Mock, accounts: Mock, lock: Mock) -> None:
service = SetupService(state=state, accounts=accounts, lock=lock, setup_required=True)
assert service.get_status() == SetupStatus(completed=False)
def test_self_hosted_status_includes_setup_time(state: Mock, accounts: Mock, lock: Mock) -> None:
setup_at = datetime(2026, 8, 6, 10, 30)
state.get_setup_at.return_value = setup_at
service = SetupService(state=state, accounts=accounts, lock=lock, setup_required=True)
assert service.get_status() == SetupStatus(completed=True, setup_at=setup_at)
def test_initialize_rejects_existing_setup(
state: Mock,
accounts: Mock,
lock: Mock,
setup_input: SetupInput,
) -> None:
state.get_setup_at.return_value = datetime(2026, 8, 6, 10, 30)
service = SetupService(state=state, accounts=accounts, lock=lock, setup_required=True)
with pytest.raises(SetupAlreadyCompletedError):
service.initialize(setup_input, initialization_validated=True)
state.has_tenants.assert_not_called()
accounts.provision.assert_not_called()
def test_initialize_rejects_existing_tenant(
state: Mock,
accounts: Mock,
lock: Mock,
setup_input: SetupInput,
) -> None:
state.has_tenants.return_value = True
service = SetupService(state=state, accounts=accounts, lock=lock, setup_required=True)
with pytest.raises(SetupAlreadyCompletedError):
service.initialize(setup_input, initialization_validated=True)
accounts.provision.assert_not_called()
@pytest.mark.parametrize("persistent_state", ["setup", "tenant"])
def test_initialize_prioritizes_existing_persistent_state_over_validation(
state: Mock,
accounts: Mock,
lock: Mock,
setup_input: SetupInput,
persistent_state: str,
) -> None:
if persistent_state == "setup":
state.get_setup_at.return_value = datetime(2026, 8, 6, 10, 30)
else:
state.has_tenants.return_value = True
service = SetupService(state=state, accounts=accounts, lock=lock, setup_required=True)
with pytest.raises(SetupAlreadyCompletedError):
service.initialize(setup_input, initialization_validated=False)
accounts.provision.assert_not_called()
def test_initialize_requires_initialization_validation(
state: Mock,
accounts: Mock,
lock: Mock,
setup_input: SetupInput,
) -> None:
service = SetupService(state=state, accounts=accounts, lock=lock, setup_required=True)
with pytest.raises(InitializationValidationRequiredError):
service.initialize(setup_input, initialization_validated=False)
accounts.provision.assert_not_called()
def test_initialize_normalizes_email_and_provisions_account(
state: Mock,
accounts: Mock,
lock: Mock,
setup_input: SetupInput,
) -> None:
service = SetupService(state=state, accounts=accounts, lock=lock, setup_required=True)
service.initialize(setup_input, initialization_validated=True)
accounts.provision.assert_called_once_with(
SetupInput(
email="admin@example.com",
name="Admin",
password="Passw0rd1",
ip_address="203.0.113.7",
language="en-US",
)
)
lock.acquire.assert_called_once_with()
def test_initialize_reads_and_writes_only_while_holding_lock(
state: Mock,
accounts: Mock,
setup_input: SetupInput,
) -> None:
lock = TrackingLock()
def get_setup_at() -> None:
assert lock.inside
def has_tenants() -> bool:
assert lock.inside
return False
def provision(_setup: SetupInput) -> None:
assert lock.inside
state.get_setup_at.side_effect = get_setup_at
state.has_tenants.side_effect = has_tenants
accounts.provision.side_effect = provision
service = SetupService(state=state, accounts=accounts, lock=lock, setup_required=True)
service.initialize(setup_input, initialization_validated=True)
assert lock.exited
def test_initialize_does_not_read_or_write_when_lock_acquisition_fails(
state: Mock,
accounts: Mock,
setup_input: SetupInput,
) -> None:
error = TimeoutError("lock acquisition timed out")
service = SetupService(
state=state,
accounts=accounts,
lock=FailingLock(error),
setup_required=True,
)
with pytest.raises(TimeoutError, match="lock acquisition timed out") as raised:
service.initialize(setup_input, initialization_validated=True)
assert raised.value is error
state.get_setup_at.assert_not_called()
state.has_tenants.assert_not_called()
accounts.provision.assert_not_called()
def test_initialize_releases_lock_and_propagates_provision_failure(
state: Mock,
accounts: Mock,
setup_input: SetupInput,
) -> None:
lock = TrackingLock()
error = RuntimeError("provision failed")
accounts.provision.side_effect = error
service = SetupService(state=state, accounts=accounts, lock=lock, setup_required=True)
with pytest.raises(RuntimeError, match="provision failed") as raised:
service.initialize(setup_input, initialization_validated=True)
assert raised.value is error
assert lock.exited
assert not lock.inside