Skip to content

Commit 699ddb0

Browse files
committed
Fix Vast.ai offer order in dstack offer --fleet
Order by score rather than by price, the same way offers are already ordered in apply plans and `dstack offer` without `--fleet`.
1 parent 97dd535 commit 699ddb0

5 files changed

Lines changed: 184 additions & 19 deletions

File tree

src/dstack/_internal/server/services/backends/__init__.py

Lines changed: 4 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,4 @@
11
import asyncio
2-
import heapq
32
import json
43
import time
54
from collections.abc import Iterable, Iterator
@@ -43,6 +42,7 @@
4342
from dstack._internal.core.models.runs import Requirements
4443
from dstack._internal.server import settings
4544
from dstack._internal.server.models import BackendModel, DecryptedString, ProjectModel
45+
from dstack._internal.server.services.offers import merge_offer_iterables
4646
from dstack._internal.settings import LOCAL_BACKEND_ENABLED
4747
from dstack._internal.utils.common import run_async
4848
from dstack._internal.utils.logging import get_logger
@@ -459,7 +459,7 @@ async def get_backend_offers(
459459
backends: List[Backend],
460460
requirements: Requirements,
461461
exclude_not_available: bool = False,
462-
) -> Iterator[Tuple[Backend, InstanceOfferWithAvailability]]:
462+
) -> Iterable[Tuple[Backend, InstanceOfferWithAvailability]]:
463463
"""
464464
Yields backend offers satisfying `requirements` sorted by price.
465465
"""
@@ -474,7 +474,7 @@ def get_filtered_offers_with_backends(
474474

475475
logger.debug("Requesting instance offers from backends: %s", [b.TYPE.value for b in backends])
476476
tasks = [run_async(get_offers_tracked, backend, requirements) for backend in backends]
477-
offers_by_backend = []
477+
offers_by_backend: list[Iterable[tuple[Backend, InstanceOfferWithAvailability]]] = []
478478
for backend, result in zip(backends, await asyncio.gather(*tasks, return_exceptions=True)):
479479
if isinstance(result, BackendError):
480480
logger.warning(
@@ -491,9 +491,7 @@ def get_filtered_offers_with_backends(
491491
)
492492
continue
493493
offers_by_backend.append(get_filtered_offers_with_backends(backend, result))
494-
# Merge preserving order for every backend.
495-
offers = heapq.merge(*offers_by_backend, key=lambda i: i[1].price)
496-
return offers
494+
return merge_offer_iterables(*offers_by_backend)
497495

498496

499497
def check_backend_type_available(backend_type: BackendType):

src/dstack/_internal/server/services/offers.py

Lines changed: 17 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
1+
import heapq
12
import itertools
23
from collections.abc import Container, Iterable, Iterator
3-
from typing import List, Literal, Optional, Tuple, Union
4+
from typing import List, Literal, Optional, Tuple, TypeVar, Union
45

56
import gpuhunt
67

@@ -116,6 +117,21 @@ async def get_offers_by_requirements(
116117
return sorted(offers, key=lambda i: not i[1].availability.is_available())
117118

118119

120+
T = TypeVar("T")
121+
122+
123+
def merge_offer_iterables(
124+
*iterables: Iterable[tuple[T, InstanceOfferWithAvailability]],
125+
) -> Iterable[tuple[T, InstanceOfferWithAvailability]]:
126+
"""
127+
Merge offers from different sources (e.g., different backends, different fleets).
128+
129+
Some backends produce offers that are not sorted by price (e.g., `vastai` sorts by pod score).
130+
That backend-specific order is preserved.
131+
"""
132+
return heapq.merge(*iterables, key=lambda i: i[1].price)
133+
134+
119135
def is_divisible_into_blocks(
120136
cpu_count: int, gpu_count: int, blocks: Union[int, Literal["auto"]]
121137
) -> tuple[bool, int]:

src/dstack/_internal/server/services/runs/plan.py

Lines changed: 13 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -52,7 +52,10 @@
5252
is_multinode_job,
5353
remove_job_spec_sensitive_info,
5454
)
55-
from dstack._internal.server.services.offers import get_offers_by_requirements
55+
from dstack._internal.server.services.offers import (
56+
get_offers_by_requirements,
57+
merge_offer_iterables,
58+
)
5659
from dstack._internal.server.services.requirements.combine import (
5760
combine_fleet_and_run_profiles,
5861
combine_fleet_and_run_requirements,
@@ -711,11 +714,10 @@ async def get_backend_offers_in_run_candidate_fleets(
711714
run_model=None,
712715
run_spec=run_spec,
713716
)
714-
deduplicated_backend_offers: dict[
715-
Hashable,
716-
tuple[Backend, InstanceOfferWithAvailability],
717-
] = {}
717+
seen_offer_identities = set()
718+
offers: list[tuple[Backend, InstanceOfferWithAvailability]] = []
718719
for candidate_fleet_model in candidate_fleet_models:
720+
offers_from_fleet = []
719721
for backend, offer in await _get_backend_offers_in_fleet(
720722
project=project,
721723
fleet_model=candidate_fleet_model,
@@ -724,13 +726,12 @@ async def get_backend_offers_in_run_candidate_fleets(
724726
volumes=volumes,
725727
max_offers=max_offers_per_fleet,
726728
):
727-
deduplicated_backend_offers.setdefault(
728-
_get_backend_offer_identity(offer),
729-
(backend, offer),
730-
)
731-
backend_offers = list(deduplicated_backend_offers.values())
732-
backend_offers.sort(key=lambda offer: offer[1].price)
733-
return backend_offers
729+
offer_identity = _get_backend_offer_identity(offer)
730+
if offer_identity not in seen_offer_identities:
731+
offers_from_fleet.append((backend, offer))
732+
seen_offer_identities.add(offer_identity)
733+
offers = list(merge_offer_iterables(offers, offers_from_fleet))
734+
return offers
734735

735736

736737
async def _get_offers_in_run_candidate_fleets(

src/dstack/_internal/server/testing/common.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -751,11 +751,13 @@ def get_fleet_configuration(
751751
name: str = "test-fleet",
752752
nodes: FleetNodesSpec = FleetNodesSpec(min=1, target=1, max=1),
753753
placement: Optional[InstanceGroupPlacement] = None,
754+
backends: Optional[list[BackendType]] = None,
754755
) -> FleetConfiguration:
755756
return FleetConfiguration(
756757
name=name,
757758
nodes=nodes,
758759
placement=placement,
760+
backends=backends,
759761
)
760762

761763

src/tests/_internal/server/routers/test_runs.py

Lines changed: 148 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -21,6 +21,7 @@
2121
ScalingSpec,
2222
ServiceConfiguration,
2323
TaskConfiguration,
24+
parse_run_configuration,
2425
)
2526
from dstack._internal.core.models.fleets import FleetNodesSpec
2627
from dstack._internal.core.models.gateways import GatewayStatus
@@ -66,6 +67,7 @@
6667
create_run,
6768
create_user,
6869
get_auth_headers,
70+
get_fleet_configuration,
6971
get_fleet_spec,
7072
get_instance_offer_with_availability,
7173
get_job_provisioning_data,
@@ -1916,6 +1918,152 @@ async def test_returns_no_offers_if_imported_fleet_specified_without_project_pre
19161918
assert response_json["project_name"] == "importer"
19171919
assert len(response_json["job_plans"][0]["offers"]) == 0
19181920

1921+
@pytest.mark.asyncio
1922+
@pytest.mark.parametrize("test_db", ["sqlite", "postgres"], indirect=True)
1923+
@pytest.mark.parametrize(
1924+
"configuration",
1925+
[
1926+
pytest.param({"type": "dev-environment"}, id="regular-configuration"),
1927+
pytest.param(
1928+
{"type": "task", "commands": [":"], "image": "scratch"},
1929+
id="special-configuration-used-by-dstack-offer-cli-command",
1930+
),
1931+
pytest.param(
1932+
{"type": "task", "commands": [":"], "image": "scratch", "fleets": ["test-fleet"]},
1933+
id="special-configuration-used-by-dstack-offer-cli-command-with-fleets", # --fleet
1934+
),
1935+
],
1936+
)
1937+
async def test_preserves_backend_specific_offer_order(
1938+
self,
1939+
test_db,
1940+
session: AsyncSession,
1941+
client: AsyncClient,
1942+
configuration: dict,
1943+
) -> None:
1944+
user = await create_user(session=session, global_role=GlobalRole.USER)
1945+
project = await create_project(session=session, owner=user)
1946+
await add_project_member(
1947+
session=session,
1948+
project=project,
1949+
user=user,
1950+
project_role=ProjectRole.USER,
1951+
)
1952+
repo = await create_repo(session=session, project_id=project.id)
1953+
await create_fleet(
1954+
session=session,
1955+
project=project,
1956+
spec=get_fleet_spec(conf=get_fleet_configuration(name="test-fleet")),
1957+
)
1958+
1959+
run_spec = get_run_spec(
1960+
repo_id=repo.name, configuration=parse_run_configuration(configuration)
1961+
)
1962+
body = {"run_spec": run_spec.dict()}
1963+
1964+
backend_mock_aws = Mock()
1965+
backend_mock_aws.TYPE = BackendType.AWS
1966+
backend_mock_aws.compute.return_value.get_offers.return_value = [
1967+
get_instance_offer_with_availability(backend=BackendType.AWS, price=1.0),
1968+
get_instance_offer_with_availability(backend=BackendType.AWS, price=4.0),
1969+
]
1970+
backend_mock_vastai = Mock()
1971+
backend_mock_vastai.TYPE = BackendType.VASTAI
1972+
backend_mock_vastai.compute.return_value.get_offers.return_value = [
1973+
# not ordered by price - custom order should be preserved
1974+
get_instance_offer_with_availability(backend=BackendType.VASTAI, price=3.0),
1975+
get_instance_offer_with_availability(backend=BackendType.VASTAI, price=2.0),
1976+
]
1977+
1978+
with patch("dstack._internal.server.services.backends.get_project_backends") as m:
1979+
m.return_value = [backend_mock_aws, backend_mock_vastai]
1980+
response = await client.post(
1981+
f"/api/project/{project.name}/runs/get_plan",
1982+
headers=get_auth_headers(user.token),
1983+
json=body,
1984+
)
1985+
1986+
assert response.status_code == 200, response.json()
1987+
offers = [(o["backend"], o["price"]) for o in response.json()["job_plans"][0]["offers"]]
1988+
expected_offers = [
1989+
(BackendType.AWS.value, 1.0),
1990+
(BackendType.VASTAI.value, 3.0),
1991+
(BackendType.VASTAI.value, 2.0),
1992+
(BackendType.AWS.value, 4.0),
1993+
]
1994+
assert offers == expected_offers
1995+
1996+
@pytest.mark.asyncio
1997+
@pytest.mark.parametrize("test_db", ["sqlite", "postgres"], indirect=True)
1998+
async def test_offer_cli_preserves_backend_specific_offer_order_across_fleets(
1999+
self, test_db, session: AsyncSession, client: AsyncClient
2000+
) -> None:
2001+
user = await create_user(session=session, global_role=GlobalRole.USER)
2002+
project = await create_project(session=session, owner=user)
2003+
await add_project_member(
2004+
session=session,
2005+
project=project,
2006+
user=user,
2007+
project_role=ProjectRole.USER,
2008+
)
2009+
repo = await create_repo(session=session, project_id=project.id)
2010+
await create_fleet(
2011+
session=session,
2012+
project=project,
2013+
spec=get_fleet_spec(
2014+
conf=get_fleet_configuration(name="fleet-aws", backends=[BackendType.AWS])
2015+
),
2016+
)
2017+
await create_fleet(
2018+
session=session,
2019+
project=project,
2020+
spec=get_fleet_spec(
2021+
conf=get_fleet_configuration(name="fleet-vastai", backends=[BackendType.VASTAI])
2022+
),
2023+
)
2024+
2025+
run_spec = get_run_spec(
2026+
repo_id=repo.name,
2027+
configuration=TaskConfiguration(
2028+
commands=[":"],
2029+
image="scratch",
2030+
fleets=["fleet-aws", "fleet-vastai"],
2031+
),
2032+
)
2033+
body = {"run_spec": run_spec.dict()}
2034+
2035+
backend_mock_aws = Mock()
2036+
backend_mock_aws.TYPE = BackendType.AWS
2037+
backend_mock_aws.compute.return_value.get_offers.return_value = [
2038+
get_instance_offer_with_availability(backend=BackendType.AWS, price=1.0),
2039+
get_instance_offer_with_availability(backend=BackendType.AWS, price=4.0),
2040+
]
2041+
backend_mock_vastai = Mock()
2042+
backend_mock_vastai.TYPE = BackendType.VASTAI
2043+
backend_mock_vastai.compute.return_value.get_offers.return_value = [
2044+
# not ordered by price - custom order should be preserved
2045+
get_instance_offer_with_availability(backend=BackendType.VASTAI, price=3.0),
2046+
get_instance_offer_with_availability(backend=BackendType.VASTAI, price=2.0),
2047+
]
2048+
2049+
with patch("dstack._internal.server.services.backends.get_project_backends") as m:
2050+
m.return_value = [backend_mock_aws, backend_mock_vastai]
2051+
response = await client.post(
2052+
f"/api/project/{project.name}/runs/get_plan",
2053+
headers=get_auth_headers(user.token),
2054+
json=body,
2055+
)
2056+
2057+
assert response.status_code == 200, response.json()
2058+
offers = [(o["backend"], o["price"]) for o in response.json()["job_plans"][0]["offers"]]
2059+
expected_offers = [
2060+
(BackendType.AWS.value, 1.0),
2061+
(BackendType.VASTAI.value, 3.0),
2062+
(BackendType.VASTAI.value, 2.0),
2063+
(BackendType.AWS.value, 4.0),
2064+
]
2065+
assert offers == expected_offers
2066+
19192067
@pytest.mark.asyncio
19202068
@pytest.mark.parametrize("test_db", ["sqlite", "postgres"], indirect=True)
19212069
async def test_offer_cli_returns_offers_from_all_specified_fleets(

0 commit comments

Comments
 (0)