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
84 changes: 51 additions & 33 deletions src/dstack/_internal/server/services/runs/plan.py
Original file line number Diff line number Diff line change
Expand Up @@ -127,6 +127,11 @@ async def get_job_plans(
else:
candidate_fleet_models = None

skip_backend_offers = (
run_spec.merged_profile.creation_policy == CreationPolicy.REUSE
or run_spec.merged_profile.instances is not None
)

if run_spec.configuration.type == "service":
replica_group_names = [g.name for g in run_spec.configuration.replica_groups]
else:
Expand All @@ -149,6 +154,7 @@ async def get_job_plans(
master_job_provisioning_data=None,
volumes=volumes,
exclude_not_available=False,
skip_backend_offers=skip_backend_offers,
)
elif run_spec.merged_profile.instances is not None:
instance_offers = await get_targeted_instance_offers(
Expand All @@ -166,6 +172,7 @@ async def get_job_plans(
run_spec=run_spec,
job=jobs[0],
volumes=volumes,
skip_backend_offers=skip_backend_offers,
)
else:
instance_offers, backend_offers = await _get_non_fleet_offers(
Expand All @@ -174,13 +181,13 @@ async def get_job_plans(
run_spec=run_spec,
job=jobs[0],
volumes=volumes,
skip_backend_offers=skip_backend_offers,
)

for job in jobs:
job_plan = _get_job_plan(
instance_offers=instance_offers,
backend_offers=backend_offers,
profile=run_spec.merged_profile,
job=job,
max_offers=max_offers,
)
Expand Down Expand Up @@ -315,6 +322,7 @@ async def find_optimal_fleet_with_offers(
master_job_provisioning_data: Optional[JobProvisioningData],
volumes: Optional[list[list[Volume]]],
exclude_not_available: bool,
skip_backend_offers: bool = False,
skip_backend_offers_on_pool_capacity: bool = False,
) -> tuple[
Optional[FleetModel],
Expand Down Expand Up @@ -397,17 +405,18 @@ async def find_optimal_fleet_with_offers(
)
)

# If any candidate fleet has pool capacity, the optimal fleet will be one of
# those, so backend offers from any fleet won't affect selection — skip them entirely when allowed.
skip_backend_offers = skip_backend_offers_on_pool_capacity and any(
candidate.has_pool_capacity for candidate in candidates
_skip_backend_offers = skip_backend_offers or (
# If any candidate fleet has pool capacity, the optimal fleet will be one of
# those, so backend offers from any fleet won't affect selection — skip them entirely when allowed.
skip_backend_offers_on_pool_capacity
and any(candidate.has_pool_capacity for candidate in candidates)
)

# Second step: gather backend offers unless skipped.
candidates_with_backend_offers: list[_FleetCandidateWithBackendOffers] = []
for candidate in candidates:
backend_offers: list[tuple[Backend, InstanceOfferWithAvailability]]
if skip_backend_offers:
if _skip_backend_offers:
backend_offers = []
else:
backend_offers = await _get_backend_offers_in_fleet(
Expand Down Expand Up @@ -439,7 +448,7 @@ async def find_optimal_fleet_with_offers(
optimal = min(candidates_with_backend_offers, key=lambda c: c.sort_key)
optimal_fleet_model = optimal.candidate.fleet_model
instance_offers = optimal.candidate.instance_offers
if skip_backend_offers:
if _skip_backend_offers:
backend_offers = []
else:
# Refetch backend offers without limit to return all offers for the optimal fleet.
Expand Down Expand Up @@ -783,6 +792,7 @@ async def _get_non_fleet_offers(
run_spec: RunSpec,
job: Job,
volumes: list[list[Volume]],
skip_backend_offers: bool = False,
) -> tuple[
list[tuple[InstanceModel, InstanceOfferWithAvailability]],
list[tuple[Backend, InstanceOfferWithAvailability]],
Expand All @@ -798,16 +808,20 @@ async def _get_non_fleet_offers(
job=job,
volumes=volumes,
)
backend_offers = await get_offers_by_requirements(
project=project,
profile=run_spec.merged_profile,
requirements=job.job_spec.requirements,
exclude_not_available=False,
multinode=is_multinode_job(job),
volumes=volumes,
privileged=job.job_spec.privileged,
instance_mounts=check_run_spec_requires_instance_mounts(run_spec),
)
backend_offers: list[tuple[Backend, InstanceOfferWithAvailability]]
if skip_backend_offers:
backend_offers = []
else:
backend_offers = await get_offers_by_requirements(
project=project,
profile=run_spec.merged_profile,
requirements=job.job_spec.requirements,
exclude_not_available=False,
multinode=is_multinode_job(job),
volumes=volumes,
privileged=job.job_spec.privileged,
instance_mounts=check_run_spec_requires_instance_mounts(run_spec),
)
return instance_offers, backend_offers


Expand Down Expand Up @@ -861,6 +875,7 @@ async def _get_offers_in_run_candidate_fleets(
run_spec: RunSpec,
job: Job,
volumes: list[list[Volume]],
skip_backend_offers: bool = False,
) -> tuple[
list[tuple[InstanceModel, InstanceOfferWithAvailability]],
list[tuple[Backend, InstanceOfferWithAvailability]],
Expand Down Expand Up @@ -891,19 +906,24 @@ async def _get_offers_in_run_candidate_fleets(
)
)
instance_offers.sort(key=lambda offer: offer[1].price or 0)
# TODO: Intentionally pass `max_offers_per_fleet=None` here. `dstack offer --fleet ...`
# is expected to return the exact `total_offers`, so capping backend offers per selected
# fleet would make that total approximate. We already deduplicate identical backend offers
# while merging selected fleets via `_get_backend_offer_identity()`. Revisit adding a cap
# only if this path causes real performance or memory problems.
backend_offers = await get_backend_offers_in_run_candidate_fleets(
session=session,
project=project,
run_spec=run_spec,
job=job,
volumes=volumes,
max_offers_per_fleet=None,
)

backend_offers: list[tuple[Backend, InstanceOfferWithAvailability]]
if skip_backend_offers:
backend_offers = []
else:
# TODO: Intentionally pass `max_offers_per_fleet=None` here. `dstack offer --fleet ...`
# is expected to return the exact `total_offers`, so capping backend offers per selected
# fleet would make that total approximate. We already deduplicate identical backend offers
# while merging selected fleets via `_get_backend_offer_identity()`. Revisit adding a cap
# only if this path causes real performance or memory problems.
backend_offers = await get_backend_offers_in_run_candidate_fleets(
session=session,
project=project,
run_spec=run_spec,
job=job,
volumes=volumes,
max_offers_per_fleet=None,
)
return instance_offers, backend_offers


Expand Down Expand Up @@ -946,14 +966,12 @@ def _freeze_offer_identity_value(value: object) -> Hashable:
def _get_job_plan(
instance_offers: list[tuple[InstanceModel, InstanceOfferWithAvailability]],
backend_offers: list[tuple[Backend, InstanceOfferWithAvailability]],
profile: Profile,
job: Job,
max_offers: Optional[int],
) -> JobPlan:
job_offers: list[InstanceOfferWithAvailability] = []
job_offers.extend(offer for _, offer in instance_offers)
if profile.creation_policy == CreationPolicy.REUSE_OR_CREATE and profile.instances is None:
job_offers.extend(offer for _, offer in backend_offers)
job_offers.extend(offer for _, offer in backend_offers)
job_offers.sort(key=lambda offer: not offer.availability.is_available())
remove_job_spec_sensitive_info(job.job_spec)
return JobPlan(
Expand Down
103 changes: 85 additions & 18 deletions src/tests/_internal/server/services/runs/test_plan.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
import copy
from unittest.mock import AsyncMock
from unittest.mock import AsyncMock, Mock

import pytest
from sqlalchemy.ext.asyncio import AsyncSession
Expand Down Expand Up @@ -27,8 +27,8 @@
_freeze_offer_identity_value,
_get_backend_offer_identity,
_get_backend_offers_in_fleet,
_get_job_plan,
get_backend_offers_in_run_candidate_fleets,
get_job_plans,
get_targeted_instance_offers,
)
from dstack._internal.server.testing.common import (
Expand Down Expand Up @@ -86,31 +86,98 @@ def test_get_backend_offer_identity_uses_full_offer_payload(self) -> None:
assert _get_backend_offer_identity(offer) != _get_backend_offer_identity(different_offer)


class TestGetJobPlan:
class TestGetJobPlansBackendOffers:
"""
Backend offers are requested only for `creation_policy: reuse-or-create` runs without
an explicit `instances` selector. `get_job_plans` decides this once via `skip_backend_offers`
and forwards it to the offer collectors.
"""

@pytest.mark.asyncio
async def test_excludes_backend_offers_when_instances_specified(self) -> None:
@pytest.mark.parametrize("test_db", ["sqlite", "postgres"], indirect=True)
@pytest.mark.parametrize(
("creation_policy", "expected_skip_backend_offers"),
[
(CreationPolicy.REUSE, True),
(CreationPolicy.REUSE_OR_CREATE, False),
],
)
async def test_skips_backend_offers_by_creation_policy(
self,
test_db,
session: AsyncSession,
monkeypatch: pytest.MonkeyPatch,
creation_policy: CreationPolicy,
expected_skip_backend_offers: bool,
) -> None:
user = await create_user(session=session)
project = await create_project(session=session, owner=user)
repo = await create_repo(session=session, project_id=project.id)
run_spec = get_run_spec(
repo_id="test-repo",
configuration=TaskConfiguration(image="debian", commands=["echo"]),
repo_id=repo.name,
configuration=TaskConfiguration(
image="debian", commands=["echo"], creation_policy=creation_policy
),
)
monkeypatch.setattr(
"dstack._internal.server.services.runs.plan._select_candidate_fleet_models",
AsyncMock(return_value=[Mock()]),
)
find_optimal_fleet_with_offers_mock = AsyncMock(return_value=(Mock(), [], []))
monkeypatch.setattr(
"dstack._internal.server.services.runs.plan.find_optimal_fleet_with_offers",
find_optimal_fleet_with_offers_mock,
)
jobs = await get_jobs_from_run_spec(run_spec=run_spec, secrets={}, replica_num=0)
instance_offer = get_instance_offer_with_availability()
backend_offer = get_instance_offer_with_availability()

job_plan = _get_job_plan(
instance_offers=[(None, instance_offer)], # type: ignore[list-item]
backend_offers=[(None, backend_offer)], # type: ignore[list-item]
profile=Profile(
name="default",
creation_policy=CreationPolicy.REUSE_OR_CREATE,
await get_job_plans(
session=session,
project=project,
run_spec=run_spec,
max_offers=None,
)

find_optimal_fleet_with_offers_mock.assert_awaited_once()
await_args = find_optimal_fleet_with_offers_mock.await_args
assert await_args is not None
assert await_args.kwargs["skip_backend_offers"] is expected_skip_backend_offers

@pytest.mark.asyncio
@pytest.mark.parametrize("test_db", ["sqlite", "postgres"], indirect=True)
async def test_excludes_backend_offers_when_instances_specified(
self,
test_db,
session: AsyncSession,
monkeypatch: pytest.MonkeyPatch,
) -> None:
user = await create_user(session=session)
project = await create_project(session=session, owner=user)
repo = await create_repo(session=session, project_id=project.id)
run_spec = get_run_spec(
repo_id=repo.name,
configuration=TaskConfiguration(
image="debian",
commands=["echo"],
instances=[InstanceNameSelector(name="my-fleet-0")],
),
job=jobs[0],
)
instance_offer = get_instance_offer_with_availability(price=1.0)
get_targeted_instance_offers_mock = AsyncMock(return_value=[(Mock(), instance_offer)])
monkeypatch.setattr(
"dstack._internal.server.services.runs.plan.get_targeted_instance_offers",
get_targeted_instance_offers_mock,
)

job_plans = await get_job_plans(
session=session,
project=project,
run_spec=run_spec,
max_offers=None,
)

assert job_plan.total_offers == 1
assert job_plan.offers == [instance_offer]
get_targeted_instance_offers_mock.assert_awaited_once()
assert len(job_plans) == 1
assert job_plans[0].total_offers == 1
assert job_plans[0].offers == [instance_offer]


class TestGetPlan:
Expand Down
Loading