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

This commit is contained in:
Byron.wang
2026-08-08 11:52:15 +00:00
committed by GitHub
parent 2a742486c5
commit 4e18ab0ad4
18 changed files with 1214 additions and 102 deletions
@@ -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