Skip to content

Commit c19ebb7

Browse files
committed
Merge remote-tracking branch 'origin/master' into pr_kubernetes_offers
2 parents c055576 + 22ac385 commit c19ebb7

8 files changed

Lines changed: 25 additions & 4 deletions

File tree

src/dstack/_internal/core/backends/base/compute.py

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -128,6 +128,7 @@ def run_job(
128128
project_ssh_private_key: str,
129129
volumes: List[Volume],
130130
placement_group: Optional[PlacementGroup],
131+
requirements: Requirements,
131132
) -> JobProvisioningData:
132133
"""
133134
Launches a new instance for the job. It should return `JobProvisioningData` ASAP.
@@ -307,6 +308,7 @@ def run_job(
307308
project_ssh_private_key: str,
308309
volumes: List[Volume],
309310
placement_group: Optional[PlacementGroup],
311+
requirements: Requirements,
310312
) -> JobProvisioningData:
311313
"""
312314
The default `run_job()` implementation for all backends that support `create_instance()`.
@@ -358,6 +360,7 @@ def run_jobs(
358360
project_ssh_public_key: str,
359361
project_ssh_private_key: str,
360362
placement_group: Optional[PlacementGroup],
363+
requirements: Requirements,
361364
) -> ComputeGroupProvisioningData:
362365
pass
363366

src/dstack/_internal/core/backends/kubernetes/compute.py

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -182,6 +182,7 @@ def run_job(
182182
project_ssh_private_key: str,
183183
volumes: list[Volume],
184184
placement_group: Optional[PlacementGroup],
185+
requirements: Requirements,
185186
) -> JobProvisioningData:
186187
cluster = self.region_cluster_map.get(instance_offer.region)
187188
if cluster is None:
@@ -256,6 +257,7 @@ def run_job(
256257
run_spec=run.run_spec,
257258
job_spec=job.job_spec,
258259
volumes=volumes,
260+
requirements=requirements,
259261
authorized_keys=authorized_keys,
260262
)
261263
exit_stack.callback(
@@ -1112,6 +1114,7 @@ def _create_job_pod(
11121114
run_spec: RunSpec,
11131115
job_spec: JobSpec,
11141116
volumes: list[Volume],
1117+
requirements: Requirements,
11151118
authorized_keys: list[str],
11161119
) -> None:
11171120
node_affinity: Optional[client.V1NodeAffinity] = None
@@ -1120,7 +1123,7 @@ def _create_job_pod(
11201123
volume_mounts: list[client.V1VolumeMount] = []
11211124
env_vars: list[client.V1EnvVar] = []
11221125

1123-
resources_spec = job_spec.requirements.resources
1126+
resources_spec = requirements.resources
11241127
resource_requests = ResourceRequests.from_resources_spec(resources_spec)
11251128
resource_limits = ResourceLimits.from_resources_spec(resources_spec)
11261129
gpu_resource: Optional[AnyKubernetesGPUResource] = None

src/dstack/_internal/core/backends/runpod/compute.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -130,6 +130,7 @@ def run_job(
130130
project_ssh_private_key: str,
131131
volumes: List[Volume],
132132
placement_group: Optional[PlacementGroup],
133+
requirements: Requirements,
133134
) -> JobProvisioningData:
134135
assert run.run_spec.ssh_key_pub is not None
135136
instance_config = InstanceConfiguration(
@@ -245,6 +246,7 @@ def run_jobs(
245246
project_ssh_public_key: str,
246247
project_ssh_private_key: str,
247248
placement_group: Optional[PlacementGroup],
249+
requirements: Requirements,
248250
) -> ComputeGroupProvisioningData:
249251
master_job_configuration = job_configurations[0]
250252
master_job = master_job_configuration.job

src/dstack/_internal/core/backends/slurm/compute.py

Lines changed: 6 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -127,12 +127,14 @@ def run_job(
127127
project_ssh_private_key: str,
128128
volumes: list[Volume],
129129
placement_group: Optional[PlacementGroup],
130+
requirements: Requirements,
130131
) -> JobProvisioningData:
131132
compute_provisioning_data = self._run_slurm_job(
132133
run=run,
133134
job=job,
134135
instance_offer=instance_offer,
135136
project_ssh_public_key=project_ssh_public_key,
137+
requirements=requirements,
136138
)
137139
return compute_provisioning_data.job_provisioning_datas[0]
138140

@@ -144,13 +146,15 @@ def run_jobs(
144146
project_ssh_public_key: str,
145147
project_ssh_private_key: str,
146148
placement_group: Optional[PlacementGroup],
149+
requirements: Requirements,
147150
) -> ComputeGroupProvisioningData:
148151
master_job = job_configurations[0].job
149152
return self._run_slurm_job(
150153
run=run,
151154
job=master_job,
152155
instance_offer=instance_offer,
153156
project_ssh_public_key=project_ssh_public_key,
157+
requirements=requirements,
154158
)
155159

156160
def terminate_instance(
@@ -172,6 +176,7 @@ def _run_slurm_job(
172176
job: Job,
173177
instance_offer: InstanceOfferWithAvailability,
174178
project_ssh_public_key: str,
179+
requirements: Requirements,
175180
) -> ComputeGroupProvisioningData:
176181
if job.job_spec.registry_auth is not None:
177182
self._skip_offer_cache.add(run, job, instance_offer)
@@ -196,7 +201,7 @@ def _run_slurm_job(
196201
authorized_keys = [project_ssh_public_key.strip(), run.run_spec.ssh_key_pub.strip()]
197202

198203
node_count = job.job_spec.jobs_per_replica
199-
resources_spec = job.job_spec.requirements.resources
204+
resources_spec = requirements.resources
200205
requested_resources = get_requested_resources_from_resources_spec(resources_spec)
201206

202207
partitions = _get_cluster_partitions(cluster, requested_resources)

src/dstack/_internal/core/backends/template/compute.py.jinja

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -84,6 +84,7 @@ class {{ backend_name }}Compute(
8484
project_ssh_private_key: str,
8585
volumes: List[Volume],
8686
placement_group: Optional[PlacementGroup],
87+
requirements: Requirements,
8788
) -> JobProvisioningData:
8889
# TODO: Implement if create_instance() is not implemented. Delete otherwise.
8990
raise NotImplementedError()

src/dstack/_internal/core/backends/vastai/compute.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -118,6 +118,7 @@ def run_job(
118118
project_ssh_private_key: str,
119119
volumes: List[Volume],
120120
placement_group: Optional[PlacementGroup],
121+
requirements: Requirements,
121122
) -> JobProvisioningData:
122123
instance_name = generate_unique_instance_name_for_job(
123124
run, job, max_length=MAX_INSTANCE_NAME_LEN

src/dstack/_internal/core/services/repos.py

Lines changed: 6 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -2,7 +2,7 @@
22
from contextlib import suppress
33
from pathlib import Path
44
from tempfile import NamedTemporaryFile
5-
from typing import Optional
5+
from typing import Optional, cast
66

77
import git
88
import git.cmd
@@ -205,8 +205,12 @@ def _get_repo_default_branch(url: str, env: dict[str, str]) -> Optional[str]:
205205
# See: https://github.com/git/git/commit/3d4355712b9fe77a96ad4ad877d92dc7ff6e0874
206206
# See: https://gist.github.com/ChrisTollefson/ab9c0a5d1dd4dd615217345c6936a307
207207
_git = git.cmd.Git()(c="credential.helper=")
208+
# Type cast is required since GitPython 3.1.51 where Git.ls_remote() was implemented as
209+
# an actual method wrapping Git.execute() but no proper @overload signatures were added.
210+
# Our call is translated to:
211+
# Git.execute(..., with_extended_output=False, as_process=False, stdout_as_string=True) -> str
212+
output = cast(str, _git.ls_remote("--symref", url, "HEAD", env=env))
208213
# output example: "ref: refs/heads/dev\tHEAD\n545344f77c0df78367085952a97fc3a058eb4c65\tHEAD"
209-
output: str = _git.ls_remote("--symref", url, "HEAD", env=env)
210214
for line in output.splitlines():
211215
# line format: `<oid> TAB <ref> LF`
212216
oid, _, ref = line.partition("\t")

src/dstack/_internal/server/background/pipeline_tasks/jobs_submitted.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2292,6 +2292,7 @@ async def _provision_new_capacity(
22922292
project_ssh_public_key,
22932293
project_ssh_private_key,
22942294
placement_group_model_to_placement_group_optional(placement_group_model),
2295+
requirements,
22952296
)
22962297
return _ProvisionNewCapacityResult(
22972298
provisioning_data=compute_group_provisioning_data,
@@ -2315,6 +2316,7 @@ async def _provision_new_capacity(
23152316
project_ssh_private_key,
23162317
offer_volumes,
23172318
placement_group_model_to_placement_group_optional(placement_group_model),
2319+
requirements,
23182320
)
23192321
return _ProvisionNewCapacityResult(
23202322
provisioning_data=job_provisioning_data,

0 commit comments

Comments
 (0)