diff --git a/src/dstack/_internal/cli/services/configurators/gateway.py b/src/dstack/_internal/cli/services/configurators/gateway.py index d13697866..f7e216f0d 100644 --- a/src/dstack/_internal/cli/services/configurators/gateway.py +++ b/src/dstack/_internal/cli/services/configurators/gateway.py @@ -24,7 +24,7 @@ GatewaySpec, GatewayStatus, ) -from dstack._internal.core.services.diff import diff_models +from dstack._internal.core.services.gateways import diff_gateway_configurations from dstack._internal.utils.common import local_time from dstack._internal.utils.logging import get_logger from dstack._internal.utils.nested_list import NestedList, NestedListItem @@ -60,15 +60,11 @@ def apply_configuration( confirm_message += "Create the gateway?" else: action_message += f"Found gateway [code]{plan.effective_spec.configuration.name}[/]." - diff = diff_models( + diff = diff_gateway_configurations( plan.current_resource.configuration, plan.effective_spec.configuration, ) - changed_fields = list(diff.keys()) - if ( - plan.current_resource.configuration == plan.effective_spec.configuration - or changed_fields == ["default"] - ): + if not diff: if command_args.yes and not command_args.force: # --force is required only with --yes, # otherwise we may ask for force apply interactively. diff --git a/src/dstack/_internal/core/compatibility/gateways.py b/src/dstack/_internal/core/compatibility/gateways.py index 6e72546f0..d76c6cece 100644 --- a/src/dstack/_internal/core/compatibility/gateways.py +++ b/src/dstack/_internal/core/compatibility/gateways.py @@ -4,15 +4,30 @@ GatewayConfiguration, GatewaySpec, ) -from dstack._internal.server.schemas.gateways import SetDefaultGatewayRequest +from dstack._internal.server.schemas.gateways import ( + GetGatewayPlanRequest, + SetDefaultGatewayRequest, +) + + +def get_get_plan_excludes(body: GetGatewayPlanRequest) -> IncludeExcludeDictType: + return {"spec": get_gateway_spec_excludes(body.spec)} def get_apply_plan_excludes(plan_input: ApplyGatewayPlanInput) -> IncludeExcludeDictType: - apply_plan_excludes: IncludeExcludeDictType = {} + apply_plan_excludes: IncludeExcludeDictType = { + "spec": get_gateway_spec_excludes(plan_input.spec) + } if plan_input.current_resource is not None: # `Gateway.backend` and `Gateway.region` are deprecated and never set since 0.21. # Not sending them lets 0.22 drop the fields without breaking 0.21 clients. - apply_plan_excludes["current_resource"] = {"backend": True, "region": True} + apply_plan_excludes["current_resource"] = { + "backend": True, + "region": True, + "configuration": _get_gateway_configuration_excludes( + plan_input.current_resource.configuration + ), + } return {"plan": apply_plan_excludes} @@ -49,4 +64,8 @@ def _get_gateway_configuration_excludes( configuration: GatewayConfiguration, ) -> IncludeExcludeDictType: configuration_excludes: IncludeExcludeDictType = {} + + if configuration.default is None: + configuration_excludes["default"] = True + return configuration_excludes diff --git a/src/dstack/_internal/core/models/gateways.py b/src/dstack/_internal/core/models/gateways.py index 02b3e1c32..a070c2a92 100644 --- a/src/dstack/_internal/core/models/gateways.py +++ b/src/dstack/_internal/core/models/gateways.py @@ -55,7 +55,18 @@ class GatewayCertificate(RootModel[Annotated[AnyGatewayCertificate, Field(discri class GatewayConfiguration(CoreModel): type: Literal["gateway"] = "gateway" name: Annotated[Optional[str], Field(description="The gateway name")] = None - default: Annotated[bool, Field(description="Make the gateway default")] = False + default: Annotated[ + Optional[bool], + Field( + description=( + "Whether the gateway is the project's default. Can be updated in-place." + " If unset when creating a new gateway," + " the gateway will become the default unless there is already a default gateway." + " If unset when updating the gateway in-place," + " the gateway's default status will not change" + ) + ), + ] = None backend: Annotated[BackendType, Field(description="The gateway backend")] region: Annotated[str, Field(description="The gateway region")] instance_type: Annotated[ diff --git a/src/dstack/_internal/core/services/gateways.py b/src/dstack/_internal/core/services/gateways.py new file mode 100644 index 000000000..446c7b7fb --- /dev/null +++ b/src/dstack/_internal/core/services/gateways.py @@ -0,0 +1,11 @@ +from dstack._internal.core.models.gateways import GatewayConfiguration +from dstack._internal.core.services.diff import ModelDiff, diff_models + + +def diff_gateway_configurations(old: GatewayConfiguration, new: GatewayConfiguration) -> ModelDiff: + return diff_models( + old, + new, + # default=None => default should stay unchanged => shouldn't be in the diff + reset={"default"} if new.default is None else {}, + ) diff --git a/src/dstack/_internal/server/compatibility/gateways.py b/src/dstack/_internal/server/compatibility/gateways.py index a21aea6c4..f58e7a3bd 100644 --- a/src/dstack/_internal/server/compatibility/gateways.py +++ b/src/dstack/_internal/server/compatibility/gateways.py @@ -2,12 +2,33 @@ from packaging.version import Version -from dstack._internal.core.models.gateways import Gateway, GatewayPlan +from dstack._internal.core.models.gateways import ( + Gateway, + GatewayConfiguration, + GatewayPlan, + GatewaySpec, +) + + +def patch_gateway_configuration_in_request( + configuration: GatewayConfiguration, client_version: Optional[Version] +) -> None: + if client_version is None: + return + if client_version < Version("0.21.1") and configuration.default is False: + # Pre-0.21.1 clients send `default=false` both when `default` was omitted and when it was + # set to `false` explicitly. Assume it was omitted, which is more common and more useful. + configuration.default = None + + +def patch_gateway_spec_in_request(spec: GatewaySpec, client_version: Optional[Version]) -> None: + patch_gateway_configuration_in_request(spec.configuration, client_version) def patch_gateway(gateway: Gateway, client_version: Optional[Version]) -> None: if client_version is None: return + _patch_gateway_configuration(gateway.configuration, client_version) if client_version < Version("0.20.25"): gateway.instance_id = "" gateway.ip_address = "\n".join(r.hostname for r in gateway.replicas if r.hostname) @@ -26,5 +47,16 @@ def patch_gateway(gateway: Gateway, client_version: Optional[Version]) -> None: def patch_gateway_plan(plan: GatewayPlan, client_version: Optional[Version]) -> None: if client_version is None: return + _patch_gateway_configuration(plan.spec.configuration, client_version) + _patch_gateway_configuration(plan.effective_spec.configuration, client_version) if plan.current_resource is not None: patch_gateway(plan.current_resource, client_version) + + +def _patch_gateway_configuration( + configuration: GatewayConfiguration, client_version: Optional[Version] +): + if client_version is None: + return + if client_version < Version("0.21.1") and configuration.default is None: + configuration.default = False diff --git a/src/dstack/_internal/server/routers/gateways.py b/src/dstack/_internal/server/routers/gateways.py index 26764a81b..03c57677f 100644 --- a/src/dstack/_internal/server/routers/gateways.py +++ b/src/dstack/_internal/server/routers/gateways.py @@ -9,7 +9,12 @@ import dstack._internal.server.services.gateways as gateways from dstack._internal.core.errors import ResourceNotExistsError from dstack._internal.core.models.common import EntityReference -from dstack._internal.server.compatibility.gateways import patch_gateway, patch_gateway_plan +from dstack._internal.server.compatibility.gateways import ( + patch_gateway, + patch_gateway_configuration_in_request, + patch_gateway_plan, + patch_gateway_spec_in_request, +) from dstack._internal.server.db import get_session from dstack._internal.server.deps import Project from dstack._internal.server.models import ProjectModel, UserModel @@ -83,6 +88,7 @@ async def get_plan( This is an optional step before calling `/apply`. """ user, project = user_project + patch_gateway_spec_in_request(body.spec, client_version) plan = await gateways.get_plan( session=session, project=project, @@ -105,6 +111,7 @@ async def apply_plan( Creates a new gateway or updates an existing gateway in-place. """ user, project = user_project + patch_gateway_spec_in_request(body.plan.spec, client_version) gateway = await gateways.apply_plan( session=session, user=user, @@ -129,6 +136,7 @@ async def create_gateway( Deprecated in favor of `/apply`. """ user, project = user_project + patch_gateway_configuration_in_request(body.configuration, client_version) gateway = await gateways.create_gateway( session=session, user=user, diff --git a/src/dstack/_internal/server/services/gateways/__init__.py b/src/dstack/_internal/server/services/gateways/__init__.py index dddf5f919..3cebb3675 100644 --- a/src/dstack/_internal/server/services/gateways/__init__.py +++ b/src/dstack/_internal/server/services/gateways/__init__.py @@ -49,11 +49,8 @@ LetsEncryptGatewayCertificate, ) from dstack._internal.core.services import validate_dstack_resource_name -from dstack._internal.core.services.diff import ( - ModelDiff, - diff_models, - format_diff_fields_for_event, -) +from dstack._internal.core.services.diff import ModelDiff, format_diff_fields_for_event +from dstack._internal.core.services.gateways import diff_gateway_configurations from dstack._internal.proxy.gateway.const import SERVICE_SCALING_WINDOWS from dstack._internal.proxy.gateway.schemas.stats import PerWindowStats, Stat from dstack._internal.server import settings @@ -92,7 +89,7 @@ from dstack._internal.utils.logging import get_logger logger = get_logger(__name__) -_CONF_UPDATABLE_FIELDS = frozenset({"domain"}) +_CONF_UPDATABLE_FIELDS = frozenset({"domain", "default"}) if FeatureFlags.GATEWAY_SCALING: _CONF_UPDATABLE_FIELDS |= {"replicas"} @@ -292,7 +289,7 @@ async def create_gateway( await session.commit() default_gateway = await get_project_default_gateway_model(session=session, project=project) - if default_gateway is None or configuration.default: + if default_gateway is None and configuration.default is None or configuration.default: await set_default_gateway( session=session, project=project, @@ -309,7 +306,9 @@ async def create_gateway( load_backend_type=True, ) assert gateway is not None - return gateway_model_to_gateway(gateway, default_gateway_id=default_gateway.id) + return gateway_model_to_gateway( + gateway, default_gateway_id=default_gateway.id if default_gateway is not None else None + ) async def connect_to_gateway_with_retry( @@ -430,7 +429,11 @@ async def set_gateway_wildcard_domain( async def set_default_gateway( - session: AsyncSession, project: ProjectModel, ref: EntityReference, user: Optional[UserModel] + session: AsyncSession, + project: ProjectModel, + ref: EntityReference, + user: Optional[UserModel], + commit: bool = True, ): gateway = await get_project_gateway_model_by_reference( session=session, project=project, ref=ref @@ -470,7 +473,28 @@ async def set_default_gateway( events.Target.from_model(project), ], ) - await session.commit() + if commit: + await session.commit() + + +async def unset_default_gateway( + session: AsyncSession, project: ProjectModel, expect_gateway_id: uuid.UUID, user: UserModel +) -> None: + gateway = await get_project_default_gateway_model(session, project) + if gateway is None or gateway.id != expect_gateway_id: + return + await session.execute( + update(ProjectModel).where(ProjectModel.id == project.id).values(default_gateway_id=None) + ) + events.emit( + session, + "Gateway unset as project default", + actor=events.UserActor.from_user(user), + targets=[ + events.Target.from_model(gateway), + events.Target.from_model(project), + ], + ) async def list_project_gateway_models( @@ -849,7 +873,6 @@ def get_gateway_configuration(gateway_model: GatewayModel) -> GatewayConfigurati # Handle gateways created before GatewayConfiguration was introduced return GatewayConfiguration( name=gateway_model.name, - default=False, backend=gateway_model.backend.type, region=gateway_model.region, domain=gateway_model.wildcard_domain, @@ -979,7 +1002,10 @@ async def get_plan( current_gateway_model, default_gateway_id=project.default_gateway_id ) if _can_update_gateway_in_place( - diff_models(current_gateway.configuration, effective_spec.configuration) + diff_gateway_configurations( + current_gateway.configuration, + effective_spec.configuration, + ) ): action = ApplyAction.UPDATE @@ -1055,7 +1081,10 @@ async def apply_plan( "Failed to apply plan. Resource has been changed. Try again or use force apply." ) - diff = diff_models(current_configuration, new_configuration) + diff = diff_gateway_configurations( + current_configuration, + new_configuration, + ) if not _can_update_gateway_in_place(diff): raise ServerClientError( f"Gateway {new_configuration.name!r} cannot be updated in-place." @@ -1069,6 +1098,21 @@ async def apply_plan( if new_configuration.replicas is not None else GATEWAY_REPLICAS_DEFAULT ) + if new_configuration.default is True: + await set_default_gateway( + session=session, + project=project, + ref=EntityReference(name=gateway_model.name, project=None), + user=user, + commit=False, + ) + elif new_configuration.default is False: + await unset_default_gateway( + session=session, + project=project, + expect_gateway_id=gateway_model.id, + user=user, + ) gateway_model.configuration = new_configuration.model_dump_json() gateway_model.last_update_at = get_current_datetime() events.emit( diff --git a/src/dstack/api/server/_gateways.py b/src/dstack/api/server/_gateways.py index 31c351459..e1fc1bd77 100644 --- a/src/dstack/api/server/_gateways.py +++ b/src/dstack/api/server/_gateways.py @@ -3,6 +3,7 @@ from dstack._internal.core.compatibility.gateways import ( get_apply_plan_excludes, get_create_gateway_excludes, + get_get_plan_excludes, get_set_default_gateway_excludes, ) from dstack._internal.core.models.common import validate_extra_ignore @@ -46,7 +47,8 @@ def get(self, project_name: str, gateway_name: str) -> Gateway: def get_plan(self, project_name: str, spec: GatewaySpec) -> GatewayPlan: body = GetGatewayPlanRequest(spec=spec) resp = self._request( - f"/api/project/{project_name}/gateways/get_plan", body=body.model_dump_json() + f"/api/project/{project_name}/gateways/get_plan", + body=body.model_dump_json(exclude=get_get_plan_excludes(body)), ) return validate_extra_ignore(GatewayPlan, resp.json()) diff --git a/src/tests/_internal/server/routers/test_gateways.py b/src/tests/_internal/server/routers/test_gateways.py index e60e25bdc..7753ccc6a 100644 --- a/src/tests/_internal/server/routers/test_gateways.py +++ b/src/tests/_internal/server/routers/test_gateways.py @@ -1,4 +1,4 @@ -from typing import Any +from typing import Any, Optional from unittest.mock import patch import pytest @@ -1795,13 +1795,60 @@ async def test_importer_member_cannot_apply_on_exporter_project( @pytest.mark.asyncio @pytest.mark.parametrize("test_db", ["sqlite", "postgres"], indirect=True) - async def test_creates_new_gateway(self, test_db, session: AsyncSession, client: AsyncClient): + @pytest.mark.parametrize("default", [None, True]) + async def test_creates_new_gateway( + self, test_db, session: AsyncSession, client: AsyncClient, default: Optional[bool] + ): user = await create_user(session, global_role=GlobalRole.USER) project = await create_project(session) await add_project_member( session=session, project=project, user=user, project_role=ProjectRole.ADMIN ) await create_backend(session, project.id, backend_type=BackendType.AWS) + configuration: dict[str, Any] = { + "type": "gateway", + "name": "my-gateway", + "backend": "aws", + "region": "us-east-1", + } + if default is not None: + configuration["default"] = default + response = await client.post( + f"/api/project/{project.name}/gateways/apply", + json={ + "plan": { + "spec": {"configuration": configuration}, + "current_resource": None, + }, + "force": False, + }, + headers=get_auth_headers(user.token), + ) + assert response.status_code == 200 + data = response.json() + assert data["name"] == "my-gateway" + assert data["status"] == "submitted" + # There is no other gateway in the project, so this one becomes the default + # regardless of whether `default` is omitted or set to `true`. + assert data["default"] is True + events = await list_events(session) + assert events[0].message == "Gateway created. Status: SUBMITTED" + + await session.refresh(project) + assert str(project.default_gateway_id) == data["id"] + + @pytest.mark.asyncio + @pytest.mark.parametrize("test_db", ["sqlite", "postgres"], indirect=True) + async def test_creates_new_gateway_with_default_false_as_not_default( + self, test_db, session: AsyncSession, client: AsyncClient + ): + user = await create_user(session, global_role=GlobalRole.USER) + project = await create_project(session) + await add_project_member( + session=session, project=project, user=user, project_role=ProjectRole.ADMIN + ) + await create_backend(session, project.id, backend_type=BackendType.AWS) + response = await client.post( f"/api/project/{project.name}/gateways/apply", json={ @@ -1812,6 +1859,7 @@ async def test_creates_new_gateway(self, test_db, session: AsyncSession, client: "name": "my-gateway", "backend": "aws", "region": "us-east-1", + "default": False, } }, "current_resource": None, @@ -1822,10 +1870,77 @@ async def test_creates_new_gateway(self, test_db, session: AsyncSession, client: ) assert response.status_code == 200 data = response.json() - assert data["name"] == "my-gateway" - assert data["status"] == "submitted" + assert data["default"] is False + assert data["configuration"]["default"] is False + + await session.refresh(project) + assert project.default_gateway_id is None + + @pytest.mark.asyncio + @pytest.mark.parametrize("test_db", ["sqlite", "postgres"], indirect=True) + async def test_creates_new_gateway_with_default_true_supersedes_existing_default( + self, test_db, session: AsyncSession, client: AsyncClient + ): + user = await create_user(session, global_role=GlobalRole.USER) + project = await create_project(session) + await add_project_member( + session=session, project=project, user=user, project_role=ProjectRole.ADMIN + ) + await create_backend(session, project.id, backend_type=BackendType.AWS) + + response = await client.post( + f"/api/project/{project.name}/gateways/apply", + json={ + "plan": { + "spec": { + "configuration": { + "type": "gateway", + "name": "first-gateway", + "backend": "aws", + "region": "us-east-1", + } + }, + "current_resource": None, + }, + "force": False, + }, + headers=get_auth_headers(user.token), + ) + assert response.status_code == 200 + first_gateway_id = response.json()["id"] + + await session.refresh(project) + assert str(project.default_gateway_id) == first_gateway_id + + await clear_events(session) + response = await client.post( + f"/api/project/{project.name}/gateways/apply", + json={ + "plan": { + "spec": { + "configuration": { + "type": "gateway", + "name": "second-gateway", + "backend": "aws", + "region": "us-east-1", + "default": True, + } + }, + "current_resource": None, + }, + "force": False, + }, + headers=get_auth_headers(user.token), + ) + assert response.status_code == 200 + data = response.json() + assert data["default"] is True + + await session.refresh(project) + assert str(project.default_gateway_id) == data["id"] events = await list_events(session) - assert events[0].message == "Gateway created. Status: SUBMITTED" + assert any(e.message == "Gateway set as project default" for e in events) + assert any(e.message == "Gateway unset as project default" for e in events) @pytest.mark.asyncio @pytest.mark.parametrize("test_db", ["sqlite", "postgres"], indirect=True) @@ -2318,3 +2433,463 @@ async def test_rejects_apply_on_failed_gateway( ) assert response.status_code == 400 assert "FAILED status" in response.json()["detail"][0]["msg"] + + +class TestApplyGatewayPlanDefault: + @pytest.mark.asyncio + @pytest.mark.parametrize("test_db", ["sqlite", "postgres"], indirect=True) + @pytest.mark.parametrize("populate_configuration", [True, False]) + async def test_sets_default_in_place( + self, test_db, session: AsyncSession, client: AsyncClient, populate_configuration: bool + ): + user = await create_user(session, global_role=GlobalRole.USER) + project = await create_project(session) + await add_project_member( + session=session, project=project, user=user, project_role=ProjectRole.ADMIN + ) + backend = await create_backend(session, project.id, backend_type=BackendType.AWS) + first_gateway = await create_gateway( + session=session, + project_id=project.id, + backend_id=backend.id, + name="first-gateway", + region="us-east-1", + populate_configuration=populate_configuration, + ) + await create_gateway_compute( + session=session, backend_id=backend.id, gateway_id=first_gateway.id + ) + second_gateway = await create_gateway( + session=session, + project_id=project.id, + backend_id=backend.id, + name="second-gateway", + region="us-east-1", + populate_configuration=populate_configuration, + ) + await create_gateway_compute( + session=session, backend_id=backend.id, gateway_id=second_gateway.id + ) + response = await client.post( + f"/api/project/{project.name}/gateways/set_default", + json={"name": first_gateway.name}, + headers=get_auth_headers(user.token), + ) + assert response.status_code == 200 + + get_response = await client.post( + f"/api/project/{project.name}/gateways/get", + json={"name": "second-gateway"}, + headers=get_auth_headers(user.token), + ) + assert get_response.status_code == 200 + current_resource = get_response.json() + assert current_resource["default"] is False + + await clear_events(session) + response = await client.post( + f"/api/project/{project.name}/gateways/apply", + json={ + "plan": { + "spec": { + "configuration": { + "type": "gateway", + "name": "second-gateway", + "backend": "aws", + "region": "us-east-1", + "default": True, + } + }, + "current_resource": current_resource, + }, + "force": False, + }, + headers=get_auth_headers(user.token), + ) + assert response.status_code == 200 + data = response.json() + assert data["default"] is True + assert data["configuration"]["default"] is True + + await session.refresh(project) + assert project.default_gateway_id == second_gateway.id + + events = await list_events(session) + assert any(e.message == "Gateway set as project default" for e in events) + assert any(e.message == "Gateway unset as project default" for e in events) + assert any("Gateway updated." in e.message and "default" in e.message for e in events) + + @pytest.mark.asyncio + @pytest.mark.parametrize("test_db", ["sqlite", "postgres"], indirect=True) + @pytest.mark.parametrize("populate_configuration", [True, False]) + async def test_unsets_default_in_place( + self, test_db, session: AsyncSession, client: AsyncClient, populate_configuration: bool + ): + user = await create_user(session, global_role=GlobalRole.USER) + project = await create_project(session) + await add_project_member( + session=session, project=project, user=user, project_role=ProjectRole.ADMIN + ) + backend = await create_backend(session, project.id, backend_type=BackendType.AWS) + gateway = await create_gateway( + session=session, + project_id=project.id, + backend_id=backend.id, + name="my-gateway", + region="us-east-1", + populate_configuration=populate_configuration, + ) + await create_gateway_compute(session=session, backend_id=backend.id, gateway_id=gateway.id) + response = await client.post( + f"/api/project/{project.name}/gateways/set_default", + json={"name": gateway.name}, + headers=get_auth_headers(user.token), + ) + assert response.status_code == 200 + + get_response = await client.post( + f"/api/project/{project.name}/gateways/get", + json={"name": "my-gateway"}, + headers=get_auth_headers(user.token), + ) + current_resource = get_response.json() + assert current_resource["default"] is True + + await clear_events(session) + response = await client.post( + f"/api/project/{project.name}/gateways/apply", + json={ + "plan": { + "spec": { + "configuration": { + "type": "gateway", + "name": "my-gateway", + "backend": "aws", + "region": "us-east-1", + "default": False, + } + }, + "current_resource": current_resource, + }, + "force": False, + }, + headers=get_auth_headers(user.token), + ) + assert response.status_code == 200 + data = response.json() + assert data["default"] is False + assert data["configuration"]["default"] is False + events = await list_events(session) + assert any(e.message == "Gateway unset as project default" for e in events) + + await session.refresh(project) + assert project.default_gateway_id is None + + @pytest.mark.asyncio + @pytest.mark.parametrize("test_db", ["sqlite", "postgres"], indirect=True) + @pytest.mark.parametrize("initial_default", [True, False]) + async def test_omitted_default_leaves_current_status_unchanged( + self, test_db, session: AsyncSession, client: AsyncClient, initial_default: bool + ): + user = await create_user(session, global_role=GlobalRole.USER) + project = await create_project(session) + await add_project_member( + session=session, project=project, user=user, project_role=ProjectRole.ADMIN + ) + backend = await create_backend(session, project.id, backend_type=BackendType.AWS) + gateway = await create_gateway( + session=session, + project_id=project.id, + backend_id=backend.id, + name="my-gateway", + region="us-east-1", + ) + await create_gateway_compute(session=session, backend_id=backend.id, gateway_id=gateway.id) + if initial_default: + response = await client.post( + f"/api/project/{project.name}/gateways/set_default", + json={"name": gateway.name}, + headers=get_auth_headers(user.token), + ) + assert response.status_code == 200 + + get_response = await client.post( + f"/api/project/{project.name}/gateways/get", + json={"name": "my-gateway"}, + headers=get_auth_headers(user.token), + ) + current_resource = get_response.json() + assert current_resource["default"] is initial_default + + await clear_events(session) + response = await client.post( + f"/api/project/{project.name}/gateways/apply", + json={ + "plan": { + "spec": { + "configuration": { + "type": "gateway", + "name": "my-gateway", + "backend": "aws", + "region": "us-east-1", + "domain": "new.example.com", + } + }, + "current_resource": current_resource, + }, + "force": False, + }, + headers=get_auth_headers(user.token), + ) + assert response.status_code == 200 + data = response.json() + assert data["default"] is initial_default + assert data["configuration"]["domain"] == "new.example.com" + events = await list_events(session) + assert not any("default" in e.message for e in events) + + @pytest.mark.asyncio + @pytest.mark.parametrize("test_db", ["sqlite", "postgres"], indirect=True) + @pytest.mark.parametrize("initial_default", [True, False]) + async def test_legacy_client_default_false_is_treated_as_omitted( + self, test_db, session: AsyncSession, client: AsyncClient, initial_default: bool + ): + """Pre-0.21.1 clients always send `default: false` when the user does not request + a default status change, since they predate `default: null`.""" + user = await create_user(session, global_role=GlobalRole.USER) + project = await create_project(session) + await add_project_member( + session=session, project=project, user=user, project_role=ProjectRole.ADMIN + ) + backend = await create_backend(session, project.id, backend_type=BackendType.AWS) + gateway = await create_gateway( + session=session, + project_id=project.id, + backend_id=backend.id, + name="my-gateway", + region="us-east-1", + ) + await create_gateway_compute(session=session, backend_id=backend.id, gateway_id=gateway.id) + if initial_default: + response = await client.post( + f"/api/project/{project.name}/gateways/set_default", + json={"name": gateway.name}, + headers=get_auth_headers(user.token), + ) + assert response.status_code == 200 + + get_response = await client.post( + f"/api/project/{project.name}/gateways/get", + json={"name": "my-gateway"}, + headers=get_auth_headers(user.token), + ) + current_resource = get_response.json() + assert current_resource["default"] is initial_default + + await clear_events(session) + response = await client.post( + f"/api/project/{project.name}/gateways/apply", + json={ + "plan": { + "spec": { + "configuration": { + "type": "gateway", + "name": "my-gateway", + "backend": "aws", + "region": "us-east-1", + "domain": "new.example.com", + "default": False, + } + }, + "current_resource": current_resource, + }, + "force": False, + }, + headers={**get_auth_headers(user.token), "x-api-version": "0.21.0"}, + ) + assert response.status_code == 200 + data = response.json() + assert data["default"] is initial_default + assert data["configuration"]["domain"] == "new.example.com" + events = await list_events(session) + assert not any("default" in e.message for e in events) + + @pytest.mark.asyncio + @pytest.mark.parametrize("test_db", ["sqlite", "postgres"], indirect=True) + async def test_force_apply_default_true_restores_default( + self, test_db, session: AsyncSession, client: AsyncClient + ): + user = await create_user(session, global_role=GlobalRole.USER) + project = await create_project(session) + await add_project_member( + session=session, project=project, user=user, project_role=ProjectRole.ADMIN + ) + await create_backend(session, project.id, backend_type=BackendType.AWS) + + response = await client.post( + f"/api/project/{project.name}/gateways/apply", + json={ + "plan": { + "spec": { + "configuration": { + "type": "gateway", + "name": "first-gateway", + "backend": "aws", + "region": "us-east-1", + "default": True, + } + }, + "current_resource": None, + }, + "force": False, + }, + headers=get_auth_headers(user.token), + ) + assert response.status_code == 200 + + await session.refresh(project) + assert project.default_gateway_id is not None + assert str(project.default_gateway_id) == response.json()["id"] + + response = await client.post( + f"/api/project/{project.name}/gateways/apply", + json={ + "plan": { + "spec": { + "configuration": { + "type": "gateway", + "name": "second-gateway", + "backend": "aws", + "region": "us-east-1", + "default": True, + } + }, + "current_resource": None, + }, + "force": False, + }, + headers=get_auth_headers(user.token), + ) + assert response.status_code == 200 + + await session.refresh(project) + assert str(project.default_gateway_id) == response.json()["id"] + + get_response = await client.post( + f"/api/project/{project.name}/gateways/get", + json={"name": "first-gateway"}, + headers=get_auth_headers(user.token), + ) + assert get_response.status_code == 200 + current_resource = get_response.json() + assert current_resource["default"] is False + + response = await client.post( + f"/api/project/{project.name}/gateways/apply", + json={ + "plan": { + "spec": { + "configuration": { + "type": "gateway", + "name": "first-gateway", + "backend": "aws", + "region": "us-east-1", + "default": True, + } + }, + "current_resource": current_resource, + }, + "force": True, + }, + headers=get_auth_headers(user.token), + ) + assert response.status_code == 200 + data = response.json() + assert data["default"] is True + + await session.refresh(project) + assert data["id"] == str(project.default_gateway_id) + + @pytest.mark.asyncio + @pytest.mark.parametrize("test_db", ["sqlite", "postgres"], indirect=True) + async def test_force_apply_default_false_unsets_default( + self, test_db, session: AsyncSession, client: AsyncClient + ): + user = await create_user(session, global_role=GlobalRole.USER) + project = await create_project(session) + await add_project_member( + session=session, project=project, user=user, project_role=ProjectRole.ADMIN + ) + await create_backend(session, project.id, backend_type=BackendType.AWS) + + response = await client.post( + f"/api/project/{project.name}/gateways/apply", + json={ + "plan": { + "spec": { + "configuration": { + "type": "gateway", + "name": "my-gateway", + "backend": "aws", + "region": "us-east-1", + "default": False, + } + }, + "current_resource": None, + }, + "force": False, + }, + headers=get_auth_headers(user.token), + ) + assert response.status_code == 200 + data = response.json() + gateway_id = data["id"] + assert data["default"] is False + + await session.refresh(project) + assert project.default_gateway_id is None + + response = await client.post( + f"/api/project/{project.name}/gateways/set_default", + json={"name": "my-gateway"}, + headers=get_auth_headers(user.token), + ) + assert response.status_code == 200 + + await session.refresh(project) + assert str(project.default_gateway_id) == gateway_id + + get_response = await client.post( + f"/api/project/{project.name}/gateways/get", + json={"name": "my-gateway"}, + headers=get_auth_headers(user.token), + ) + assert get_response.status_code == 200 + current_resource = get_response.json() + assert current_resource["default"] is True + + response = await client.post( + f"/api/project/{project.name}/gateways/apply", + json={ + "plan": { + "spec": { + "configuration": { + "type": "gateway", + "name": "my-gateway", + "backend": "aws", + "region": "us-east-1", + "default": False, + } + }, + "current_resource": current_resource, + }, + "force": True, + }, + headers=get_auth_headers(user.token), + ) + assert response.status_code == 200 + data = response.json() + assert data["default"] is False + + await session.refresh(project) + assert project.default_gateway_id is None