diff --git a/src/srm/api/databus/dependencies.py b/src/srm/api/databus/dependencies.py index e72f9fea3a311d5c611097a0bf9bf4ebfcd2e99c..a3234b9aa9ad1d82aca8a634b7fbe405acb775e8 100644 --- a/src/srm/api/databus/dependencies.py +++ b/src/srm/api/databus/dependencies.py @@ -10,7 +10,7 @@ from srm.adapters.database.repos.runtime_inventory import ( SqlServiceInstanceRepository, SqlServiceOrderRepository, ) -from srm.adapters.database.repos.topology import SqlDomainRepository, SqlZoneRepository +from srm.adapters.database.repos.topology import SqlZoneRepository from srm.adapters.databus.nats_connection_manager import NatsConnectionManager from srm.adapters.databus.nats_publisher import NatsPublisher from srm.application.command_handlers.deploy_service import DeployServiceCommandCoordinator @@ -37,7 +37,6 @@ def build_deploy_service_use_case( service_instances=SqlServiceInstanceRepository(session), capability_instances=SqlCapabilityInstanceRepository(session), zones=SqlZoneRepository(session), - domains=SqlDomainRepository(session), publisher=publisher, ) diff --git a/src/srm/application/services/capability_placement.py b/src/srm/application/services/capability_placement.py index 5beaf08c111b89ef1d8b18fc06eff0eb32edd1fb..a8b6249b09b3f7dbfe54db7805c6c8fb72fb5b1f 100644 --- a/src/srm/application/services/capability_placement.py +++ b/src/srm/application/services/capability_placement.py @@ -16,7 +16,7 @@ from srm.domain.models.topology import ( Zone, ZoneState, ) -from srm.domain.ports.database.topology import DomainRepository, ZoneRepository +from srm.domain.ports.database.topology import ZoneRepository @dataclass(frozen=True, slots=True) @@ -47,40 +47,17 @@ class CapabilityPlacementPlanner: def __init__( self, zones: ZoneRepository, - domains: DomainRepository, ) -> None: self._zones = zones - self._domains = domains async def place(self, request: CapabilityPlacementRequest) -> CapabilityPlacement | None: if request.pin.domain_id is not None and request.pin.zone_id is None: return None - if request.pin.domain_id is not None: - return await self._place_pinned_domain(request) if request.pin.zone_id is not None: return await self._place_pinned_zone(request) return await self._place_unpinned(request) - async def _place_pinned_domain( - self, - request: CapabilityPlacementRequest, - ) -> CapabilityPlacement | None: - assert request.pin.zone_id is not None - assert request.pin.domain_id is not None - - zone = await self._zones.get_by_id(request.pin.zone_id) - if zone is None or zone.state != ZoneState.ACTIVE: - return None - - domain = await self._domains.get_by_id(request.pin.domain_id) - if domain is None: - return None - if domain.zone_id != zone.id or domain.kind != request.domain_kind: - return None - - return self._select_from_domain(zone.id, domain, request) - async def _place_pinned_zone( self, request: CapabilityPlacementRequest, diff --git a/src/srm/application/use_cases/deploy_service.py b/src/srm/application/use_cases/deploy_service.py index 615cb2a09bfb5fdd8caa3f7db65b5c7f90aecb62..6099003d0a7f5f1799a10c2b1f2d5820a96d5898 100644 --- a/src/srm/application/use_cases/deploy_service.py +++ b/src/srm/application/use_cases/deploy_service.py @@ -38,7 +38,7 @@ from srm.domain.ports.database.runtime_inventory import ( ServiceInstanceRepository, ServiceOrderRepository, ) -from srm.domain.ports.database.topology import DomainRepository, ZoneRepository +from srm.domain.ports.database.topology import ZoneRepository from srm.domain.ports.databus.events import ( OperationCompletedStatus, OperationStatusState, @@ -109,7 +109,6 @@ class DeployServiceUseCase: service_instances: ServiceInstanceRepository, capability_instances: CapabilityInstanceRepository, zones: ZoneRepository, - domains: DomainRepository, publisher: DataBusPublisher, placement_planner: CapabilityPlacementPlanner | None = None, ) -> None: @@ -120,9 +119,8 @@ class DeployServiceUseCase: self._service_instances = service_instances self._capability_instances = capability_instances self._zones = zones - self._domains = domains self._publisher = publisher - self._placement_planner = placement_planner or CapabilityPlacementPlanner(zones, domains) + self._placement_planner = placement_planner or CapabilityPlacementPlanner(zones) async def accept(self, command: DeployServiceCommand) -> DeployServiceResult: existing_order = await self._service_orders.get_by_operation_id(command.operation_id) @@ -491,12 +489,10 @@ class DeployServiceUseCase: candidate_zones: list[Zone] | None, ) -> DeployTargetPlacement | None: if target.zone_id is not None: - if target.domain_id is None: - zone = await self._zones.get_by_id(target.zone_id) - if zone is None: - return None - return self._place_target_in_loaded_zone(target, deploy_material, zone) - return await self._place_target_in_pin(target, deploy_material) + zone = await self._zones.get_by_id(target.zone_id) + if zone is None: + return None + return self._place_target_in_loaded_zone(target, deploy_material, zone) if target.domain_id is not None: return None @@ -514,35 +510,6 @@ class DeployServiceUseCase: return placement return None - async def _place_target_in_pin( - self, - target: DeployTargetCommand, - deploy_material: list[tuple[ServiceCapabilityRequirement, RuntimeKind]], - ) -> DeployTargetPlacement | None: - requirement_placements: list[DeployRequirementPlacement] = [] - for requirement, runtime_kind in deploy_material: - placement = await self._placement_planner.place( - self._placement_request( - requirement, - runtime_kind, - PlacementPin(zone_id=target.zone_id, domain_id=target.domain_id), - ) - ) - if placement is None: - return None - requirement_placements.append( - DeployRequirementPlacement( - requirement_id=requirement.id, - placement=placement, - ) - ) - - return DeployTargetPlacement( - app_instance_id=target.app_instance_id, - zone_id=requirement_placements[0].placement.zone_id, - requirements=requirement_placements, - ) - def _place_target_in_loaded_zone( self, target: DeployTargetCommand, @@ -556,7 +523,7 @@ class DeployServiceUseCase: self._placement_request( requirement, runtime_kind, - PlacementPin(zone_id=zone.id), + PlacementPin(zone_id=zone.id, domain_id=target.domain_id), ), ) if placement is None: diff --git a/tests/api/databus/test_dependencies.py b/tests/api/databus/test_dependencies.py index 7b9908a50fa751bba5d54e536c431b232b766079..6ea44be5b6eee0c69171dae1c52d102c39b26909 100644 --- a/tests/api/databus/test_dependencies.py +++ b/tests/api/databus/test_dependencies.py @@ -17,7 +17,6 @@ def test_get_deploy_service_use_case_wires_required_repositories() -> None: patch("srm.api.databus.dependencies.SqlServiceInstanceRepository") as instances, patch("srm.api.databus.dependencies.SqlCapabilityInstanceRepository") as capabilities, patch("srm.api.databus.dependencies.SqlZoneRepository") as zones, - patch("srm.api.databus.dependencies.SqlDomainRepository") as domains, patch("srm.api.databus.dependencies.NatsPublisher") as publisher, patch("srm.api.databus.dependencies.DeployServiceUseCase") as use_case, ): @@ -30,7 +29,6 @@ def test_get_deploy_service_use_case_wires_required_repositories() -> None: instances.assert_called_once_with(session) capabilities.assert_called_once_with(session) zones.assert_called_once_with(session) - domains.assert_called_once_with(session) publisher.assert_called_once_with(connection_manager) use_case.assert_called_once_with( service_specifications=specs.return_value, @@ -40,7 +38,6 @@ def test_get_deploy_service_use_case_wires_required_repositories() -> None: service_instances=instances.return_value, capability_instances=capabilities.return_value, zones=zones.return_value, - domains=domains.return_value, publisher=publisher.return_value, ) assert result is use_case.return_value diff --git a/tests/unit/test_deploy_service_use_case.py b/tests/unit/test_deploy_service_use_case.py index 7d0068a1e7d68f2c482768428c537345c941b07d..a374f253adf40ccd4f9a530f475ac9aa9ace0f1b 100644 --- a/tests/unit/test_deploy_service_use_case.py +++ b/tests/unit/test_deploy_service_use_case.py @@ -252,7 +252,6 @@ def _use_case( service_deployment_units: AsyncMock | None = None, service_capability_requirements: AsyncMock | None = None, zones: AsyncMock | None = None, - domains: AsyncMock | None = None, publisher: AsyncMock, ) -> DeployServiceUseCase: deployment_units = service_deployment_units @@ -268,14 +267,11 @@ def _use_case( _deploy_requirement(service_specification_id=uuid4(), deployment_unit_id=unit_id) ] zone_repo = zones or AsyncMock() - domain_repo = domains or AsyncMock() if zones is None: zone_id = uuid4() domain = _domain(uuid4(), zone_id=zone_id) zone_repo.get_by_id.return_value = _zone(zone_id, domains=[domain]) zone_repo.list_active_resource_zones.return_value = [_zone(zone_id, domains=[domain])] - if domains is None: - domain_repo.get_by_id.return_value = None return DeployServiceUseCase( service_specifications=service_specifications or AsyncMock(), @@ -285,7 +281,6 @@ def _use_case( service_instances=service_instances or AsyncMock(), capability_instances=capability_instances or AsyncMock(), zones=zone_repo, - domains=domain_repo, publisher=publisher, ) @@ -624,7 +619,6 @@ async def test_execute_publishes_failed_before_start_when_target_zone_pin_is_mis service_orders=service_orders, service_specifications=service_specifications, zones=zones, - domains=domains, publisher=publisher, ).accept(command) @@ -667,19 +661,17 @@ async def test_execute_publishes_failed_before_start_when_target_domain_pin_is_m zones = AsyncMock() zones.get_by_id.return_value = _zone(zone_id) domains = AsyncMock() - domains.get_by_id.return_value = None publisher = AsyncMock() await _use_case( service_orders=service_orders, service_specifications=service_specifications, zones=zones, - domains=domains, publisher=publisher, ).accept(command) zones.get_by_id.assert_awaited_once_with(zone_id) - domains.get_by_id.assert_awaited_once_with(domain_id) + domains.get_by_id.assert_not_awaited() assert publisher.publish.await_count == 2 subject, payload = publisher.publish.await_args_list[0].args assert subject == "event.srm.operation.status" @@ -714,7 +706,6 @@ async def test_execute_publishes_failed_before_start_when_target_domain_pin_has_ service_orders=service_orders, service_specifications=service_specifications, zones=zones, - domains=domains, publisher=publisher, ).accept(command) @@ -757,19 +748,17 @@ async def test_execute_publishes_failed_before_start_when_target_domain_is_outsi zones = AsyncMock() zones.get_by_id.return_value = _zone(zone_id) domains = AsyncMock() - domains.get_by_id.return_value = _domain(domain_id, zone_id=uuid4()) publisher = AsyncMock() await _use_case( service_orders=service_orders, service_specifications=service_specifications, zones=zones, - domains=domains, publisher=publisher, ).accept(command) zones.get_by_id.assert_awaited_once_with(zone_id) - domains.get_by_id.assert_awaited_once_with(domain_id) + domains.get_by_id.assert_not_awaited() assert publisher.publish.await_count == 2 subject, payload = publisher.publish.await_args_list[0].args assert subject == "event.srm.operation.status" @@ -1224,6 +1213,92 @@ async def test_execute_loads_pinned_zone_once_for_multiple_requirements() -> Non publisher.publish.assert_not_awaited() +async def test_execute_loads_pinned_zone_once_for_domain_pin_multiple_requirements() -> None: + operation_id = uuid4() + zone_id = uuid4() + domain_id = uuid4() + helm_unit_id = uuid4() + container_unit_id = uuid4() + command = _command( + operation_id, + targets=[ + DeployTargetCommand( + app_instance_id=uuid4(), + zone_id=zone_id, + domain_id=domain_id, + ) + ], + ) + service_orders = AsyncMock() + service_orders.get_by_operation_id.return_value = None + service_orders.create.return_value = _service_order(operation_id) + service_instances = AsyncMock() + service_specifications = AsyncMock() + service_specifications.get_by_id.return_value = _service_specification( + id=command.service_specification_id, + state=ServiceSpecificationState.ACTIVE, + ) + service_deployment_units = AsyncMock() + service_deployment_units.list_by_service_specification_id.return_value = [ + _deployment_unit( + id=helm_unit_id, + service_specification_id=command.service_specification_id, + runtime_kind=RuntimeKind.HELM, + ), + _deployment_unit( + id=container_unit_id, + service_specification_id=command.service_specification_id, + runtime_kind=RuntimeKind.CONTAINER, + ), + ] + service_capability_requirements = AsyncMock() + service_capability_requirements.list_by_service_specification_id.return_value = [ + _deploy_requirement( + service_specification_id=command.service_specification_id, + deployment_unit_id=helm_unit_id, + ref="deploy-helm", + ), + _deploy_requirement( + service_specification_id=command.service_specification_id, + deployment_unit_id=container_unit_id, + ref="deploy-container", + ), + ] + zones = AsyncMock() + zones.get_by_id.return_value = _zone( + zone_id, + domains=[ + _domain( + domain_id, + zone_id=zone_id, + runtime_kinds=[RuntimeKind.HELM, RuntimeKind.CONTAINER], + ) + ], + ) + domains = AsyncMock() + publisher = AsyncMock() + + use_case = _use_case( + service_orders=service_orders, + service_instances=service_instances, + service_specifications=service_specifications, + service_deployment_units=service_deployment_units, + service_capability_requirements=service_capability_requirements, + zones=zones, + publisher=publisher, + ) + result = await use_case.accept(command) + + zones.get_by_id.assert_awaited_once_with(zone_id) + domains.get_by_id.assert_not_awaited() + service_instances.create.assert_awaited_once() + assert result.accepted_order_id is not None + assert len(result.target_placements) == 1 + assert result.target_placements[0].zone_id == zone_id + assert len(result.target_placements[0].requirements) == 2 + publisher.publish.assert_not_awaited() + + async def test_execute_accepts_valid_target_zone_and_domain_pins() -> None: operation_id = uuid4() zone_id = uuid4() @@ -1250,9 +1325,11 @@ async def test_execute_accepts_valid_target_zone_and_domain_pins() -> None: state=ServiceSpecificationState.ACTIVE, ) zones = AsyncMock() - zones.get_by_id.return_value = _zone(zone_id) + zones.get_by_id.return_value = _zone( + zone_id, + domains=[_domain(domain_id, zone_id=zone_id)], + ) domains = AsyncMock() - domains.get_by_id.return_value = _domain(domain_id, zone_id=zone_id) publisher = AsyncMock() use_case = _use_case( @@ -1261,13 +1338,12 @@ async def test_execute_accepts_valid_target_zone_and_domain_pins() -> None: capability_instances=capability_instances, service_specifications=service_specifications, zones=zones, - domains=domains, publisher=publisher, ) result = await use_case.accept(command) zones.get_by_id.assert_awaited_once_with(zone_id) - domains.get_by_id.assert_awaited_once_with(domain_id) + domains.get_by_id.assert_not_awaited() service_orders.create.assert_awaited_once() (created_order,) = service_orders.create.await_args.args assert created_order.operation_id == operation_id