Skip to content
Open
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
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -298,7 +298,7 @@ oci = [
"dstack[server]",
]
nebius = [
"nebius>=0.3.4,<0.4",
"nebius>=0.6.13,<0.7",
"dstack[server]",
]
fluentbit = [
Expand Down
83 changes: 81 additions & 2 deletions src/dstack/_internal/core/backends/nebius/compute.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
import shlex
import time
from collections.abc import Iterable
from decimal import ROUND_CEILING, Decimal
from functools import cached_property
from typing import List, Optional

Expand Down Expand Up @@ -66,6 +67,7 @@
max=Memory.parse("8192GB"), # max for the NETWORK_SSD disk type
)
WAIT_FOR_DISK_TIMEOUT = 20
WAIT_FOR_PRICING_POLICY_TIMEOUT = 10
WAIT_FOR_INSTANCE_TIMEOUT = 30
WAIT_FOR_INSTANCE_UPDATE_INTERVAL = 2.5
DELETE_INSTANCE_TIMEOUT = 25
Expand Down Expand Up @@ -136,7 +138,11 @@ def get_all_offers_with_availability(
extra_filter=_supported_instances,
)
return [
offer.with_availability(availability=InstanceAvailability.UNKNOWN) for offer in offers
offer.with_availability(
availability=InstanceAvailability.UNKNOWN,
price=_get_price_adjusted_for_pricing_policy(offer),
)
for offer in offers
]

def get_offers_modifiers(
Expand All @@ -156,6 +162,9 @@ def create_instance(
# instance.
instance_name = generate_unique_instance_name(instance_config)
platform, preset = instance_offer.instance.name.split()
offer_backend_data = validate_extra_ignore(
NebiusOfferBackendData, instance_offer.backend_data
)
cluster_id = None
if placement_group:
assert placement_group.provisioning_data is not None
Expand Down Expand Up @@ -186,6 +195,7 @@ def create_instance(
labels=labels,
)
create_instance_op = None
create_pricing_policy_op = None
try:
logger.debug("Blocking until disk %s is created", create_disk_op.resource_id)
resources.wait_for_operation(create_disk_op, timeout=WAIT_FOR_DISK_TIMEOUT)
Expand All @@ -195,6 +205,27 @@ def create_instance(
f"Create disk operation failed. Message: {raw_op.status.message}."
f" Details: {raw_op.status.details}"
)
if (
instance_offer.instance.resources.spot
and not offer_backend_data.is_preemptible_flat_rate
):
create_pricing_policy_op = resources.create_pricing_policy(
sdk=self._sdk,
name=instance_name,
project_id=self._region_to_project_id[instance_offer.region],
platform=platform,
# Prevent the spot price from growing above what dstack shows
max_price=str(_get_pricing_policy_max_price(instance_offer)),
)
resources.wait_for_operation(
create_pricing_policy_op, timeout=WAIT_FOR_PRICING_POLICY_TIMEOUT
)
if not create_pricing_policy_op.successful():
raw_op = create_pricing_policy_op.raw()
raise ProvisioningError(
f"Create pricing policy operation failed. Message: {raw_op.status.message}."
f" Details: {raw_op.status.details}"
)
create_instance_op = resources.create_instance(
sdk=self._sdk,
name=instance_name,
Expand All @@ -209,6 +240,11 @@ def create_instance(
disk_id=create_disk_op.resource_id,
subnet_id=self._get_subnet_id(instance_offer.region),
preemptible=instance_offer.instance.resources.spot,
pricing_policy_id=(
create_pricing_policy_op.resource_id
if create_pricing_policy_op is not None
else None
),
labels=labels,
)
_wait_for_instance(self._sdk, create_instance_op)
Expand All @@ -226,6 +262,18 @@ def create_instance(
logger.exception(
"Could not delete instance %s: %s", create_instance_op.resource_id, e
)
if create_pricing_policy_op is not None:
try:
with resources.ignore_errors([StatusCode.NOT_FOUND]):
resources.delete_pricing_policy(
self._sdk, create_pricing_policy_op.resource_id
)
except Exception as e:
logger.exception(
"Could not delete pricing policy %s: %s",
create_pricing_policy_op.resource_id,
e,
)
try:
with resources.ignore_errors([StatusCode.NOT_FOUND]):
resources.delete_disk(self._sdk, create_disk_op.resource_id)
Expand All @@ -245,7 +293,12 @@ def create_instance(
username="ubuntu",
dockerized=True,
backend_data=NebiusInstanceBackendData(
boot_disk_id=create_disk_op.resource_id
boot_disk_id=create_disk_op.resource_id,
pricing_policy_id=(
create_pricing_policy_op.resource_id
if create_pricing_policy_op is not None
else None
),
).model_dump_json(),
)

Expand Down Expand Up @@ -284,6 +337,9 @@ def terminate_instance(
)
with resources.ignore_errors([StatusCode.NOT_FOUND]):
resources.delete_disk(self._sdk, backend_data_parsed.boot_disk_id)
if backend_data_parsed.pricing_policy_id is not None:
with resources.ignore_errors([StatusCode.NOT_FOUND]):
resources.delete_pricing_policy(self._sdk, backend_data_parsed.pricing_policy_id)

def create_placement_group(
self,
Expand Down Expand Up @@ -348,6 +404,7 @@ def is_suitable_placement_group(

class NebiusInstanceBackendData(CoreModel):
boot_disk_id: str
pricing_policy_id: str | None = None

@classmethod
def load(cls, raw: Optional[str]) -> "NebiusInstanceBackendData":
Expand Down Expand Up @@ -406,3 +463,25 @@ def _wait_for_instance(sdk: SDK, op: SDKOperation[Operation]) -> None:
def _supported_instances(offer: InstanceOffer) -> bool:
platform, _ = offer.instance.name.split()
return platform in SUPPORTED_PLATFORMS


def _get_price_adjusted_for_pricing_policy(offer: InstanceOffer) -> float:
if (
not offer.instance.resources.spot
or validate_extra_ignore(
NebiusOfferBackendData, offer.backend_data
).is_preemptible_flat_rate
):
return offer.price
max_price = _get_pricing_policy_max_price(offer)
return float(max_price * _get_pricing_policy_price_units(offer))


def _get_pricing_policy_max_price(offer: InstanceOffer) -> Decimal:
# Pricing policies allow at most 3 decimal places. Round upward.
price = Decimal(str(offer.price)) / _get_pricing_policy_price_units(offer)
return price.quantize(Decimal("0.001"), rounding=ROUND_CEILING)


def _get_pricing_policy_price_units(offer: InstanceOffer) -> int:
return len(offer.instance.resources.gpus) or 1
4 changes: 4 additions & 0 deletions src/dstack/_internal/core/backends/nebius/models.py
Original file line number Diff line number Diff line change
Expand Up @@ -182,6 +182,10 @@ class NebiusConfig(NebiusStoredConfig):

class NebiusOfferBackendData(CoreModel):
fabrics: set[str] = set()
is_preemptible_flat_rate: bool = False
"""
True if the price of this spot instance does not change and pricing policies are not supported.
"""

@field_serializer("fabrics")
def _serialize_fabrics(self, value: set[str]) -> list[str]:
Expand Down
45 changes: 45 additions & 0 deletions src/dstack/_internal/core/backends/nebius/resources.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,16 @@
from nebius.aio.operation import Operation as SDKOperation
from nebius.aio.service_error import RequestError, StatusCode
from nebius.aio.token.renewable import OPTION_RENEW_REQUEST_TIMEOUT, OPTION_RENEW_SYNCHRONOUS
from nebius.api.nebius.billing.v1 import (
ComputeInstanceSpec,
ComputeInstanceSpecV1,
CreatePricingPolicyRequest,
DeletePricingPolicyRequest,
MaxPriceV1,
PricingMethod,
PricingPolicyServiceClient,
PricingPolicySpec,
)
from nebius.api.nebius.common.v1 import Operation, ResourceMetadata
from nebius.api.nebius.compute.v1 import (
AttachedDiskSpec,
Expand All @@ -37,6 +47,7 @@
PublicIPAddress,
ResourcesSpec,
SourceImageFamily,
SpotPricingPolicySpec,
)
from nebius.api.nebius.iam.v1 import (
Container,
Expand Down Expand Up @@ -299,6 +310,34 @@ def delete_disk(sdk: SDK, disk_id: str) -> None:
)


def create_pricing_policy(
sdk: SDK, name: str, project_id: str, platform: str, max_price: str
) -> SDKOperation[Operation]:
client = PricingPolicyServiceClient(sdk)
request = CreatePricingPolicyRequest(
metadata=ResourceMetadata(name=name, parent_id=project_id),
spec=PricingPolicySpec(
compute_instance_spec=ComputeInstanceSpec(v1=ComputeInstanceSpecV1(platform=platform)),
pricing=PricingMethod(max_price_v1=MaxPriceV1(max_price=max_price)),
),
)
return LOOP.await_(
client.create(
request, per_retry_timeout=REQUEST_TIMEOUT, auth_options=REQUEST_AUTH_OPTIONS
)
)


def delete_pricing_policy(sdk: SDK, pricing_policy_id: str) -> None:
LOOP.await_(
PricingPolicyServiceClient(sdk).delete(
DeletePricingPolicyRequest(id=pricing_policy_id),
per_retry_timeout=REQUEST_TIMEOUT,
auth_options=REQUEST_AUTH_OPTIONS,
)
)


def create_instance(
sdk: SDK,
name: str,
Expand All @@ -310,6 +349,7 @@ def create_instance(
disk_id: str,
subnet_id: str,
preemptible: bool,
pricing_policy_id: Optional[str],
labels: Dict[str, str],
) -> SDKOperation[Operation]:
client = InstanceServiceClient(sdk)
Expand Down Expand Up @@ -341,6 +381,11 @@ def create_instance(
if preemptible
else None,
recovery_policy=InstanceRecoveryPolicy.FAIL if preemptible else None,
spot_pricing_policy=(
SpotPricingPolicySpec(id=pricing_policy_id)
if pricing_policy_id is not None
else None
),
),
)
with wrap_capacity_errors():
Expand Down
2 changes: 1 addition & 1 deletion src/dstack/_internal/core/models/instances.py
Original file line number Diff line number Diff line change
Expand Up @@ -205,7 +205,7 @@ def with_availability(self, **kwargs) -> "InstanceOfferWithAvailability":
"""Convert to InstanceOfferWithAvailability without re-serializing/re-validating fields.
The result shares nested objects with self. This is generally safe because callers
discard the original InstanceOffer after conversion."""
return InstanceOfferWithAvailability.model_construct(**self.__dict__, **kwargs)
return InstanceOfferWithAvailability.model_construct(**{**self.__dict__, **kwargs})


class InstanceOfferWithAvailability(InstanceOffer):
Expand Down
111 changes: 111 additions & 0 deletions src/tests/_internal/core/backends/nebius/test_compute.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,111 @@
import pytest

from dstack._internal.core.backends.nebius import compute as compute_module
from dstack._internal.core.backends.nebius.compute import NebiusCompute
from dstack._internal.core.backends.nebius.models import (
NebiusConfig,
NebiusServiceAccountCreds,
)
from dstack._internal.core.models.backends.base import BackendType
from dstack._internal.core.models.instances import (
Gpu,
InstanceAvailability,
InstanceOffer,
InstanceType,
Resources,
)


def make_compute() -> NebiusCompute:
return NebiusCompute(
NebiusConfig(
creds=NebiusServiceAccountCreds(
service_account_id="service-account-id",
public_key_id="public-key-id",
private_key_content="private-key",
)
)
)


def make_offer(
price: float, spot: bool, is_preemptible_flat_rate: bool, gpu_count: int = 8
) -> InstanceOffer:
return InstanceOffer(
backend=BackendType.NEBIUS,
instance=InstanceType(
name=f"gpu-h100-sxm {gpu_count}gpu-128vcpu-1600gb",
resources=Resources(
cpus=128,
memory_mib=1600 * 1024,
gpus=[Gpu(name="H100", memory_mib=80 * 1024)] * gpu_count,
spot=spot,
),
),
region="eu-north1",
price=price,
backend_data={
"fabrics": ["fabric-2"],
"is_preemptible_flat_rate": is_preemptible_flat_rate,
},
)


class TestGetAllOffersWithAvailability:
@pytest.fixture(autouse=True)
def _mock_region_to_project_id(self, mocker):
mocker.patch.object(NebiusCompute, "_region_to_project_id", {"eu-north1": "project-id"})

@pytest.mark.parametrize(
("price", "gpu_count", "expected_price"),
[
(1.23, 1, 1.23),
(1.2, 1, 1.2),
(1.2341, 1, 1.235),
(1.2349, 1, 1.235),
(1.2350001, 1, 1.236),
(0.0001, 1, 0.001),
(2.007, 1, 2.007),
(9.84, 8, 9.84),
(1.23, 8, 1.232),
(1.2341, 0, 1.235),
],
)
def test_rounds_spot_price_up_to_3_decimal_places_per_gpu(
self, mocker, price, gpu_count, expected_price
):
mocker.patch.object(
compute_module,
"get_catalog_offers",
return_value=[
make_offer(price, spot=True, is_preemptible_flat_rate=False, gpu_count=gpu_count)
],
)

offers = make_compute().get_all_offers_with_availability(unallocated_resources=False)

assert len(offers) == 1
assert offers[0].price == expected_price
assert offers[0].availability == InstanceAvailability.UNKNOWN

def test_keeps_on_demand_price_as_is(self, mocker):
mocker.patch.object(
compute_module,
"get_catalog_offers",
return_value=[make_offer(1.2341, spot=False, is_preemptible_flat_rate=False)],
)

offers = make_compute().get_all_offers_with_availability(unallocated_resources=False)

assert offers[0].price == 1.2341

def test_keeps_flat_rate_spot_price_as_is(self, mocker):
mocker.patch.object(
compute_module,
"get_catalog_offers",
return_value=[make_offer(1.2341, spot=True, is_preemptible_flat_rate=True)],
)

offers = make_compute().get_all_offers_with_availability(unallocated_resources=False)

assert offers[0].price == 1.2341
Loading