Skip to content

Commit ad15e90

Browse files
authored
Allow configuring subnet_ids in Azure settings (#3955)
Similarly to `vpc_ids`, allow selecting specific subnets to be attached to dstack VMs. ```yaml projects: - name: main backends: - type: azure subscription_id: ... tenant_id: ... creds: type: client client_id: ... client_secret: ... subnet_ids: westeurope: my-resource-group/my-vpc/my-subnet regions: [westeurope] ```
1 parent fc16f72 commit ad15e90

6 files changed

Lines changed: 164 additions & 18 deletions

File tree

‎mkdocs/docs/concepts/backends.md‎

Lines changed: 18 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -403,16 +403,30 @@ There are two ways to configure Azure: using a client secret or using the defaul
403403
- type: azure
404404
creds:
405405
type: default
406-
regions: [westeurope]
407-
vpc_ids:
408-
westeurope: myNetworkResourceGroup/myNetworkName
406+
regions: [westeurope]
407+
vpc_ids:
408+
westeurope: myNetworkResourceGroup/myNetworkName
409+
```
410+
411+
Alternatively, specify `subnet_ids` to target specific subnets:
412+
413+
```yaml
414+
projects:
415+
- name: main
416+
backends:
417+
- type: azure
418+
creds:
419+
type: default
420+
regions: [westeurope]
421+
subnet_ids:
422+
westeurope: myNetworkResourceGroup/myNetworkName/mySubnetName
409423
```
410424

411425

412426
??? info "Private subnets"
413427
By default, `dstack` provisions instances with public IPs and permits inbound SSH traffic.
414428
If you want `dstack` to use private subnets and provision instances without public IPs,
415-
specify custom networks using `vpc_ids` and set `public_ips` to `false`.
429+
specify custom networks using `vpc_ids` or `subnet_ids`, and set `public_ips` to `false`.
416430

417431
```yaml
418432
projects:

‎src/dstack/_internal/core/backends/azure/compute.py‎

Lines changed: 36 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -143,6 +143,7 @@ def create_instance(
143143
network_client=self._network_client,
144144
resource_group=self.config.resource_group,
145145
vpc_ids=self.config.vpc_ids,
146+
subnet_ids=self.config.subnet_ids,
146147
location=location,
147148
allocate_public_ip=allocate_public_ip,
148149
)
@@ -252,6 +253,7 @@ def create_gateway(
252253
network_client=self._network_client,
253254
resource_group=self.config.resource_group,
254255
vpc_ids=self.config.vpc_ids,
256+
subnet_ids=self.config.subnet_ids,
255257
location=configuration.region,
256258
allocate_public_ip=True,
257259
)
@@ -326,9 +328,38 @@ def get_resource_group_network_subnet_or_error(
326328
network_client: network_mgmt.NetworkManagementClient,
327329
resource_group: Optional[str],
328330
vpc_ids: Optional[Dict[str, str]],
331+
subnet_ids: Optional[Dict[str, str]],
329332
location: str,
330333
allocate_public_ip: bool,
331334
) -> Tuple[str, str, str]:
335+
if subnet_ids is not None and location in subnet_ids:
336+
subnet_id = subnet_ids[location]
337+
try:
338+
net_resource_group, network_name, subnet_name = _parse_config_subnet_id(subnet_id)
339+
except Exception:
340+
raise ComputeError(
341+
"Subnet specified in incorrect format."
342+
" Supported format for `subnet_ids` values: 'networkResourceGroupName/networkName/subnetName'"
343+
)
344+
try:
345+
subnet = network_client.subnets.get(net_resource_group, network_name, subnet_name)
346+
except ResourceNotFoundError:
347+
raise ComputeError(
348+
f"Subnet {subnet_name} not found in network {network_name}"
349+
f" in resource group {net_resource_group}"
350+
)
351+
if not allocate_public_ip and not azure_resources.is_eligible_private_subnet(
352+
network_client=network_client,
353+
resource_group=net_resource_group,
354+
network_name=network_name,
355+
subnet=subnet,
356+
):
357+
raise ComputeError(
358+
f"Subnet {subnet_name} in network {network_name} does not have outbound internet connectivity."
359+
" Ensure a NAT Gateway is attached or VNet peering is configured."
360+
)
361+
return net_resource_group, network_name, subnet_name
362+
332363
if vpc_ids is not None:
333364
vpc_id = vpc_ids.get(location)
334365
if vpc_id is None:
@@ -388,6 +419,11 @@ def _parse_config_vpc_id(vpc_id: str) -> Tuple[str, str]:
388419
return resource_group, network_name
389420

390421

422+
def _parse_config_subnet_id(subnet_id: str) -> Tuple[str, str, str]:
423+
resource_group, network_name, subnet_name = subnet_id.split("/")
424+
return resource_group, network_name, subnet_name
425+
426+
391427
class VMImageVariant(enum.Enum):
392428
GRID = enum.auto()
393429
CUDA = enum.auto()

‎src/dstack/_internal/core/backends/azure/configurator.py‎

Lines changed: 26 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -125,7 +125,7 @@ def create_backend(
125125
subscription_id=config.subscription_id,
126126
resource_group=config.resource_group,
127127
locations=config.regions,
128-
create_default_network=config.vpc_ids is None,
128+
create_default_network=config.vpc_ids is None and config.subnet_ids is None,
129129
)
130130
return BackendRecord(
131131
config=AzureStoredConfig(
@@ -226,23 +226,38 @@ def _check_config_vpc(
226226
if config.subscription_id is None:
227227
return None
228228
allocate_public_ip = config.public_ips if config.public_ips is not None else True
229-
if config.public_ips is False and config.vpc_ids is None:
230-
raise ServerClientError(msg="`vpc_ids` must be specified if `public_ips: false`.")
229+
if config.public_ips is False and config.vpc_ids is None and config.subnet_ids is None:
230+
raise ServerClientError(
231+
msg="`vpc_ids` or `subnet_ids` must be specified if `public_ips: false`."
232+
)
233+
if config.vpc_ids is not None and config.subnet_ids is not None:
234+
overlap = sorted(set(config.vpc_ids.keys()) & set(config.subnet_ids.keys()))
235+
if overlap:
236+
raise ServerClientError(
237+
f"Regions {overlap} are configured in both `vpc_ids` and `subnet_ids`."
238+
" Each region must be specified in only one of them."
239+
)
231240
locations = config.regions
232241
if locations is None:
233242
locations = DEFAULT_LOCATIONS
234-
if config.vpc_ids is not None:
235-
vpc_ids_locations = list(config.vpc_ids.keys())
236-
not_configured_locations = [loc for loc in locations if loc not in vpc_ids_locations]
243+
if config.vpc_ids is not None or config.subnet_ids is not None:
244+
configured_locations = set()
245+
if config.vpc_ids is not None:
246+
configured_locations |= set(config.vpc_ids.keys())
247+
if config.subnet_ids is not None:
248+
configured_locations |= set(config.subnet_ids.keys())
249+
not_configured_locations = [
250+
loc for loc in locations if loc not in configured_locations
251+
]
237252
if len(not_configured_locations) > 0:
238253
if config.regions is None:
239254
raise ServerClientError(
240-
f"`vpc_ids` not configured for regions {not_configured_locations}. "
241-
"Configure `vpc_ids` for all regions or specify `regions`."
255+
f"Networking not configured for regions {not_configured_locations}. "
256+
"Configure either `vpc_ids` or `subnet_ids` for all regions or specify `regions`."
242257
)
243258
raise ServerClientError(
244-
f"`vpc_ids` not configured for regions {not_configured_locations}. "
245-
"Configure `vpc_ids` for all regions specified in `regions`."
259+
f"Networking not configured for regions {not_configured_locations}. "
260+
"Configure either `vpc_ids` or `subnet_ids` for all regions specified in `regions`."
246261
)
247262
network_client = network_mgmt.NetworkManagementClient(
248263
credential=credential,
@@ -256,6 +271,7 @@ def _check_config_vpc(
256271
network_client=network_client,
257272
resource_group=None,
258273
vpc_ids=config.vpc_ids,
274+
subnet_ids=config.subnet_ids,
259275
location=location,
260276
allocate_public_ip=allocate_public_ip,
261277
)

‎src/dstack/_internal/core/backends/azure/models.py‎

Lines changed: 12 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -51,13 +51,23 @@ class AzureBackendConfig(CoreModel):
5151
)
5252
),
5353
] = None
54+
subnet_ids: Annotated[
55+
Optional[Dict[str, str]],
56+
Field(
57+
description=(
58+
"The mapping from configured Azure locations to subnet IDs."
59+
" A subnet ID must have a format `networkResourceGroup/networkName/subnetName`."
60+
" Cannot be configured for the same region as `vpc_ids`"
61+
)
62+
),
63+
] = None
5464
public_ips: Annotated[
5565
Optional[bool],
5666
Field(
5767
description=(
5868
"A flag to enable/disable public IP assigning on instances."
59-
" `public_ips: false` requires `vpc_ids` that specifies custom networks with outbound internet connectivity"
60-
" provided by NAT Gateway or other mechanism."
69+
" `public_ips: false` requires `vpc_ids` or `subnet_ids` that specifies custom networks"
70+
" with outbound internet connectivity provided by NAT Gateway or other mechanism."
6171
" Defaults to `true`"
6272
)
6373
),

‎src/dstack/_internal/core/backends/azure/resources.py‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -25,7 +25,7 @@ def get_network_subnets(
2525
)
2626
for subnet in subnets:
2727
if private:
28-
if _is_eligible_private_subnet(
28+
if is_eligible_private_subnet(
2929
network_client=network_client,
3030
resource_group=resource_group,
3131
network_name=network_name,
@@ -54,7 +54,7 @@ def _is_eligible_public_subnet(
5454
return True
5555

5656

57-
def _is_eligible_private_subnet(
57+
def is_eligible_private_subnet(
5858
network_client: network_mgmt.NetworkManagementClient,
5959
resource_group: str,
6060
network_name: str,

‎src/tests/_internal/core/backends/azure/test_configurator.py‎

Lines changed: 70 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,7 @@
1010
from dstack._internal.core.errors import (
1111
BackendAuthError,
1212
BackendInvalidCredentialsError,
13+
ServerClientError,
1314
)
1415

1516

@@ -59,3 +60,72 @@ def test_validate_config_invalid_creds(self):
5960
["creds", "client_id"],
6061
["creds", "client_secret"],
6162
]
63+
64+
65+
class TestCheckConfigVpc:
66+
def _make_config(self, **kwargs):
67+
return AzureBackendConfigWithCreds(
68+
creds=AzureClientCreds(tenant_id="t", client_id="c", client_secret="s"),
69+
tenant_id="ten1",
70+
subscription_id="sub1",
71+
**kwargs,
72+
)
73+
74+
def _check(self, config):
75+
with (
76+
patch("azure.mgmt.network.NetworkManagementClient"),
77+
patch(
78+
"dstack._internal.core.backends.azure.compute.get_resource_group_network_subnet_or_error"
79+
),
80+
):
81+
AzureConfigurator()._check_config_vpc(config, Mock())
82+
83+
def test_public_ips_false_requires_network_config(self):
84+
config = self._make_config(regions=["westeurope"], public_ips=False)
85+
with pytest.raises(ServerClientError, match="`vpc_ids` or `subnet_ids` must be specified"):
86+
AzureConfigurator()._check_config_vpc(config, Mock())
87+
88+
def test_public_ips_false_with_vpc_ids_ok(self):
89+
config = self._make_config(
90+
regions=["westeurope"], public_ips=False, vpc_ids={"westeurope": "rg/net"}
91+
)
92+
self._check(config)
93+
94+
def test_public_ips_false_with_subnet_ids_ok(self):
95+
config = self._make_config(
96+
regions=["westeurope"], public_ips=False, subnet_ids={"westeurope": "rg/net/subnet"}
97+
)
98+
self._check(config)
99+
100+
def test_overlap_raises(self):
101+
config = self._make_config(
102+
regions=["westeurope", "eastus"],
103+
vpc_ids={"westeurope": "rg/net", "eastus": "rg/net2"},
104+
subnet_ids={"westeurope": "rg/net/subnet"},
105+
)
106+
with pytest.raises(ServerClientError, match="westeurope"):
107+
AzureConfigurator()._check_config_vpc(config, Mock())
108+
109+
def test_uncovered_region_raises_with_vpc_ids(self):
110+
config = self._make_config(
111+
regions=["westeurope", "eastus"],
112+
vpc_ids={"westeurope": "rg/net"},
113+
)
114+
with pytest.raises(ServerClientError, match="eastus"):
115+
AzureConfigurator()._check_config_vpc(config, Mock())
116+
117+
def test_uncovered_region_raises_with_subnet_ids(self):
118+
config = self._make_config(
119+
regions=["westeurope", "eastus"],
120+
subnet_ids={"westeurope": "rg/net/subnet"},
121+
)
122+
with pytest.raises(ServerClientError, match="eastus"):
123+
AzureConfigurator()._check_config_vpc(config, Mock())
124+
125+
def test_mixed_vpc_and_subnet_ids_covers_all_regions(self):
126+
config = self._make_config(
127+
regions=["westeurope", "eastus"],
128+
vpc_ids={"westeurope": "rg/net"},
129+
subnet_ids={"eastus": "rg/net/subnet"},
130+
)
131+
self._check(config)

0 commit comments

Comments
 (0)