diff --git a/src/dstack/_internal/server/services/runs/plan.py b/src/dstack/_internal/server/services/runs/plan.py index 29aef4ee4..a6c6c1270 100644 --- a/src/dstack/_internal/server/services/runs/plan.py +++ b/src/dstack/_internal/server/services/runs/plan.py @@ -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: @@ -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( @@ -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( @@ -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, ) @@ -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], @@ -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( @@ -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. @@ -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]], @@ -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 @@ -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]], @@ -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 @@ -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( diff --git a/src/tests/_internal/server/services/runs/test_plan.py b/src/tests/_internal/server/services/runs/test_plan.py index ce586171c..01ff0c4e3 100644 --- a/src/tests/_internal/server/services/runs/test_plan.py +++ b/src/tests/_internal/server/services/runs/test_plan.py @@ -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 @@ -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 ( @@ -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: