mirror of
https://github.com/langgenius/dify.git
synced 2026-08-29 03:45:08 +08:00
refactor(api): decouple setup and spec endpoints (#40104)
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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 ()
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user