Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
19 changes: 19 additions & 0 deletions src/dstack/_internal/core/backends/aws/compute.py
Original file line number Diff line number Diff line change
Expand Up @@ -528,6 +528,25 @@ def is_suitable_placement_group(
return False
return placement_group.configuration.region == instance_offer.region

def are_placement_groups_compatible_with_reservation(
self,
instance_offer: InstanceOffer,
reservation: str,
) -> bool:
# AWS rejects launches into Capacity Blocks that specify a placement group.
# Capacity Block instances are already placed close together in EC2 UltraClusters.
try:
capacity_block = aws_resources.get_reservation(
ec2_client=self.session.client("ec2", region_name=instance_offer.region),
reservation_id=reservation,
is_capacity_block=True,
active_only=False,
)
except botocore.exceptions.ClientError as e:
logger.warning("Failed to get reservation %s: %s", reservation, e)
return True
return capacity_block is None

def create_gateway_replica(
self,
configuration: GatewayReplicaConfiguration,
Expand Down
7 changes: 5 additions & 2 deletions src/dstack/_internal/core/backends/aws/resources.py
Original file line number Diff line number Diff line change
Expand Up @@ -678,8 +678,11 @@ def get_reservation(
instance_count: int = 0,
instance_types: Optional[List[str]] = None,
is_capacity_block: bool = False,
active_only: bool = True,
) -> Optional[Dict[str, Any]]:
filters = [{"Name": "state", "Values": ["active"]}]
filters = []
if active_only:
filters.append({"Name": "state", "Values": ["active"]})
if instance_types:
filters.append({"Name": "instance-type", "Values": instance_types})
try:
Expand Down Expand Up @@ -707,7 +710,7 @@ def get_reservation(
if instance_count > 0 and reservation["AvailableInstanceCount"] < instance_count:
return None

if is_capacity_block and reservation["ReservationType"] != "capacity-block":
if is_capacity_block and reservation.get("ReservationType") != "capacity-block":
return None

return reservation
18 changes: 12 additions & 6 deletions src/dstack/_internal/core/backends/base/compute.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,7 +26,6 @@
DSTACK_RUNNER_SSH_PORT,
DSTACK_SHIM_HTTP_PORT,
)
from dstack._internal.core.models.backends.base import BackendType
from dstack._internal.core.models.compute_groups import ComputeGroup, ComputeGroupProvisioningData
from dstack._internal.core.models.gateways import (
GatewayLoadBalancerConfiguration,
Expand Down Expand Up @@ -415,7 +414,7 @@ def run_job(
user=run.user,
ssh_keys=[SSHKey(public=project_ssh_public_key.strip())],
volumes=volumes,
reservation=job.job_spec.requirements.reservation,
reservation=requirements.reservation,
tags=run.run_spec.merged_profile.tags,
)
instance_offer = instance_offer.model_copy()
Expand Down Expand Up @@ -559,13 +558,20 @@ def is_suitable_placement_group(
"""
pass

def are_placement_groups_compatible_with_reservations(self, backend_type: BackendType) -> bool:
def are_placement_groups_compatible_with_reservation(
self,
instance_offer: InstanceOffer,
reservation: str,
) -> bool:
"""
Whether placement groups can be used for instances provisioned in reservations.
Whether a placement group can be used for an instance provisioned in the reservation.

May perform API calls.

Arguments:
backend_type: matches the backend type of this compute, unless this compute is a proxy
for other backends (dstack Sky)
instance_offer: the offer to provision. Its backend matches the backend type of this
compute, unless this compute is a proxy for other backends (dstack Sky)
reservation: the reservation to provision the instance in
"""
return True

Expand Down
6 changes: 5 additions & 1 deletion src/dstack/_internal/core/backends/gcp/compute.py
Original file line number Diff line number Diff line change
Expand Up @@ -594,7 +594,11 @@ def is_suitable_placement_group(
) -> bool:
return placement_group.configuration.region == instance_offer.region

def are_placement_groups_compatible_with_reservations(self, backend_type: BackendType) -> bool:
def are_placement_groups_compatible_with_reservation(
self,
instance_offer: InstanceOffer,
reservation: str,
) -> bool:
# Cannot use our own placement policies when provisioning in a reservation.
# Instead, we use the placement policy defined in reservation settings.
return False
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -47,6 +47,7 @@
from dstack._internal.server.services.logging import fmt
from dstack._internal.server.services.offers import get_instance_offer_with_restricted_az
from dstack._internal.server.services.placement import (
can_use_placement_groups,
get_fleet_placement_group_models,
placement_group_model_to_placement_group,
placement_group_model_to_placement_group_optional,
Expand Down Expand Up @@ -140,9 +141,10 @@ async def create_cloud_instance(instance_model: InstanceModel) -> ProcessResult:
and cluster_context.is_current_instance_master
and instance_offer.backend in BACKENDS_WITH_PLACEMENT_GROUPS_SUPPORT
and isinstance(compute, ComputeWithPlacementGroupSupport)
and (
compute.are_placement_groups_compatible_with_reservations(instance_offer.backend)
or instance_configuration.reservation is None
and await can_use_placement_groups(
compute=compute,
instance_offer=instance_offer,
reservation=instance_configuration.reservation,
)
):
(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -122,6 +122,7 @@
)
from dstack._internal.server.services.pipelines import PipelineHinterProtocol
from dstack._internal.server.services.placement import (
can_use_placement_groups,
find_or_create_suitable_placement_group,
get_placement_group_model_for_job,
placement_group_model_to_placement_group_optional,
Expand Down Expand Up @@ -2457,9 +2458,10 @@ async def _provision_new_capacity(
and is_cloud_cluster(fleet_model)
and offer.backend in BACKENDS_WITH_PLACEMENT_GROUPS_SUPPORT
and isinstance(compute, ComputeWithPlacementGroupSupport)
and (
compute.are_placement_groups_compatible_with_reservations(offer.backend)
or job.job_spec.requirements.reservation is None
and await can_use_placement_groups(
compute=compute,
instance_offer=offer,
reservation=requirements.reservation,
)
):
placement_group_model = await find_or_create_suitable_placement_group(
Expand Down
14 changes: 14 additions & 0 deletions src/dstack/_internal/server/services/placement.py
Original file line number Diff line number Diff line change
Expand Up @@ -121,6 +121,20 @@ def get_placement_group_model_for_job(
return placement_group_model


async def can_use_placement_groups(
compute: ComputeWithPlacementGroupSupport,
instance_offer: InstanceOffer,
reservation: Optional[str],
) -> bool:
if reservation is None:
return True
return await run_async(
compute.are_placement_groups_compatible_with_reservation,
instance_offer,
reservation,
)


async def find_or_create_suitable_placement_group(
fleet_model: FleetModel,
placement_groups: list[PlacementGroupModel],
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,7 @@
create_project,
get_fleet_configuration,
get_fleet_spec,
get_instance_configuration,
get_instance_offer_with_availability,
get_job_provisioning_data,
get_placement_group_provisioning_data,
Expand Down Expand Up @@ -749,6 +750,66 @@ async def test_create_placement_group_if_placement_cluster(
assert backend_mock.compute.return_value.create_placement_group.call_count == 0
assert len(placement_groups) == 0

@pytest.mark.parametrize("compatible", [True, False])
async def test_creates_placement_group_in_reservation_only_if_compatible(
self,
test_db,
session: AsyncSession,
worker: InstanceWorker,
compatible: bool,
) -> None:
project = await create_project(session=session)
fleet = await create_fleet(
session,
project,
spec=get_fleet_spec(
conf=get_fleet_configuration(
placement=InstanceGroupPlacement.CLUSTER,
nodes=FleetNodesSpec(min=1, target=1, max=1),
)
),
)
instance_configuration = get_instance_configuration()
instance_configuration.reservation = "test-reservation"
instance = await create_instance(
session=session,
project=project,
fleet=fleet,
status=InstanceStatus.PENDING,
offer=None,
job_provisioning_data=None,
instance_configuration=instance_configuration,
)
await _set_current_master_instance(session, fleet, instance)
offer = get_instance_offer_with_availability()
backend_mock = Mock()
backend_mock.TYPE = BackendType.AWS
compute_mock = Mock(spec=ComputeMockSpec)
backend_mock.compute.return_value = compute_mock
compute_mock.get_offers.return_value = [offer]
compute_mock.are_placement_groups_compatible_with_reservation.return_value = compatible
compute_mock.create_instance.return_value = get_job_provisioning_data()
compute_mock.create_placement_group.return_value = get_placement_group_provisioning_data()
with patch("dstack._internal.server.services.backends.get_project_backends") as m:
m.return_value = [backend_mock]
await process_instance(session, worker, instance)

await session.refresh(instance)
assert instance.status == InstanceStatus.PROVISIONING
compute_mock.are_placement_groups_compatible_with_reservation.assert_called_once_with(
offer, "test-reservation"
)
placement_groups = (await session.execute(select(PlacementGroupModel))).scalars().all()
created_placement_group = compute_mock.create_instance.call_args[0][2]
if compatible:
assert compute_mock.create_placement_group.call_count == 1
assert len(placement_groups) == 1
assert isinstance(created_placement_group, PlacementGroup)
else:
assert compute_mock.create_placement_group.call_count == 0
assert len(placement_groups) == 0
assert created_placement_group is None

@pytest.mark.parametrize("can_reuse", [True, False])
async def test_reuses_placement_group_between_offers_if_the_group_is_suitable(
self,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -803,6 +803,66 @@ async def test_creates_placement_group_for_cluster_fleet(
placement_group = (await session.execute(select(PlacementGroupModel))).scalar()
assert placement_group is not None

@pytest.mark.parametrize("compatible", [True, False])
async def test_creates_placement_group_in_reservation_only_if_compatible(
self, test_db, session: AsyncSession, worker: JobSubmittedWorker, compatible: bool
):
project = await create_project(session=session)
user = await create_user(session=session)
repo = await create_repo(session=session, project_id=project.id)
fleet_spec = get_fleet_spec()
fleet_spec.configuration.placement = InstanceGroupPlacement.CLUSTER
fleet_spec.configuration.nodes = FleetNodesSpec(min=0, target=0, max=None)
fleet = await create_fleet(session=session, project=project, spec=fleet_spec)
run_spec = get_run_spec(run_name="test-run", repo_id=repo.name)
run_spec.configuration.reservation = "test-reservation"
run = await create_run(
session=session,
project=project,
repo=repo,
user=user,
fleet=fleet,
run_name="test-run",
run_spec=run_spec,
)
job = await create_job(session=session, run=run, instance_assigned=True)
offer = get_instance_offer_with_availability(backend=BackendType.AWS)

with patch("dstack._internal.server.services.backends.get_project_backends") as m:
backend_mock = Mock()
compute_mock = Mock(spec=ComputeMockSpec)
backend_mock.TYPE = BackendType.AWS
backend_mock.compute.return_value = compute_mock
m.return_value = [backend_mock]
compute_mock.get_offers.return_value = [offer]
compute_mock.are_placement_groups_compatible_with_reservation.return_value = compatible
compute_mock.run_job.return_value = get_job_provisioning_data(
backend=BackendType.AWS,
)
compute_mock.create_placement_group.return_value = (
get_placement_group_provisioning_data()
)

await _process_job(session=session, worker=worker, job_model=job)

await session.refresh(job)
assert job.status == JobStatus.PROVISIONING
compute_mock.are_placement_groups_compatible_with_reservation.assert_called_once()
assert (
compute_mock.are_placement_groups_compatible_with_reservation.call_args[0][1]
== "test-reservation"
)
compute_mock.run_job.assert_called_once()
placement_group = (await session.execute(select(PlacementGroupModel))).scalar()
if compatible:
compute_mock.create_placement_group.assert_called_once()
assert isinstance(compute_mock.run_job.call_args[0][6], PlacementGroup)
assert placement_group is not None
else:
compute_mock.create_placement_group.assert_not_called()
assert compute_mock.run_job.call_args[0][6] is None
assert placement_group is None

async def test_marks_unused_existing_placement_groups_for_cleanup(
self, test_db, session: AsyncSession, worker: JobSubmittedWorker
):
Expand Down
Loading