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
@@ -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