Skip to content

Commit cfceac9

Browse files
committed
feat(gpu): add GPU count resource requests
Signed-off-by: Evan Lezar <elezar@nvidia.com>
1 parent f514fec commit cfceac9

17 files changed

Lines changed: 503 additions & 64 deletions

File tree

crates/openshell-cli/src/main.rs

Lines changed: 55 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1029,6 +1029,7 @@ enum DoctorCommands {
10291029
Check,
10301030
}
10311031

1032+
#[allow(clippy::large_enum_variant)]
10321033
#[derive(Subcommand, Debug)]
10331034
enum SandboxCommands {
10341035
/// Create a sandbox.
@@ -1086,10 +1087,14 @@ enum SandboxCommands {
10861087

10871088
/// Target a driver-specific GPU device. Docker and Podman use CDI device IDs
10881089
/// (for example "nvidia.com/gpu=0"); VM uses a PCI BDF or index.
1089-
/// Implies a GPU request. When omitted with --gpu, the driver uses its default GPU selection.
1090-
#[arg(long)]
1090+
/// Implies a GPU request. Mutually exclusive with --gpu-count.
1091+
#[arg(long, conflicts_with = "gpu_count")]
10911092
gpu_device: Option<String>,
10921093

1094+
/// Request a specific number of GPUs. Mutually exclusive with --gpu-device.
1095+
#[arg(long, value_parser = clap::value_parser!(u32).range(1..), conflicts_with = "gpu_device")]
1096+
gpu_count: Option<u32>,
1097+
10931098
/// CPU limit for the sandbox (for example: 500m, 1, 2.5).
10941099
#[arg(long)]
10951100
cpu: Option<String>,
@@ -2373,6 +2378,7 @@ async fn main() -> Result<()> {
23732378
editor,
23742379
gpu,
23752380
gpu_device,
2381+
gpu_count,
23762382
cpu,
23772383
memory,
23782384
providers,
@@ -2441,6 +2447,7 @@ async fn main() -> Result<()> {
24412447
keep,
24422448
gpu,
24432449
gpu_device.as_deref(),
2450+
gpu_count,
24442451
cpu.as_deref(),
24452452
memory.as_deref(),
24462453
editor,
@@ -3713,6 +3720,52 @@ mod tests {
37133720
}
37143721
}
37153722

3723+
#[test]
3724+
fn sandbox_create_gpu_count_parses_without_gpu_flag() {
3725+
let cli = Cli::try_parse_from(["openshell", "sandbox", "create", "--gpu-count", "2"])
3726+
.expect("sandbox create --gpu-count should parse");
3727+
3728+
if let Some(Commands::Sandbox {
3729+
command: Some(SandboxCommands::Create { gpu, gpu_count, .. }),
3730+
..
3731+
}) = cli.command
3732+
{
3733+
assert!(!gpu);
3734+
assert_eq!(gpu_count, Some(2));
3735+
} else {
3736+
panic!("expected SandboxCommands::Create");
3737+
}
3738+
}
3739+
3740+
#[test]
3741+
fn sandbox_create_gpu_count_rejects_zero() {
3742+
let result = Cli::try_parse_from(["openshell", "sandbox", "create", "--gpu-count", "0"]);
3743+
3744+
assert!(
3745+
result.is_err(),
3746+
"sandbox create --gpu-count 0 should be rejected"
3747+
);
3748+
}
3749+
3750+
#[test]
3751+
fn sandbox_create_gpu_count_conflicts_with_gpu_device() {
3752+
let result = Cli::try_parse_from([
3753+
"openshell",
3754+
"sandbox",
3755+
"create",
3756+
"--gpu",
3757+
"--gpu-device",
3758+
"0",
3759+
"--gpu-count",
3760+
"2",
3761+
]);
3762+
3763+
assert!(
3764+
result.is_err(),
3765+
"sandbox create should reject --gpu-count with --gpu-device"
3766+
);
3767+
}
3768+
37163769
#[test]
37173770
fn service_expose_accepts_positional_target_port_and_service() {
37183771
let cli = Cli::try_parse_from([

crates/openshell-cli/src/run.rs

Lines changed: 45 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -1613,6 +1613,7 @@ pub async fn sandbox_create(
16131613
keep: bool,
16141614
gpu: bool,
16151615
gpu_device: Option<&str>,
1616+
gpu_count: Option<u32>,
16161617
cpu: Option<&str>,
16171618
memory: Option<&str>,
16181619
editor: Option<Editor>,
@@ -1667,6 +1668,7 @@ pub async fn sandbox_create(
16671668
};
16681669
let requested_gpu = gpu
16691670
|| gpu_device.is_some_and(|device_id| !device_id.is_empty())
1671+
|| gpu_count.is_some()
16701672
|| image.as_deref().is_some_and(image_requests_gpu);
16711673

16721674
let inferred_types: Vec<String> = inferred_provider_type(command).into_iter().collect();
@@ -1693,7 +1695,7 @@ pub async fn sandbox_create(
16931695

16941696
let request = CreateSandboxRequest {
16951697
spec: Some(SandboxSpec {
1696-
placement: resource_requirements_from_cli(requested_gpu, gpu_device),
1698+
placement: resource_requirements_from_cli(requested_gpu, gpu_device, gpu_count),
16971699
policy,
16981700
providers: configured_providers,
16991701
template,
@@ -2086,13 +2088,22 @@ pub async fn sandbox_create(
20862088
fn resource_requirements_from_cli(
20872089
requested_gpu: bool,
20882090
gpu_device: Option<&str>,
2091+
gpu_count: Option<u32>,
20892092
) -> Option<ResourceRequirements> {
2093+
let requested_gpu = requested_gpu
2094+
|| gpu_device.is_some_and(|device_id| !device_id.is_empty())
2095+
|| gpu_count.is_some();
20902096
requested_gpu.then(|| ResourceRequirements {
20912097
gpu: Some(GpuSpec {
2092-
device_ids: gpu_device
2093-
.filter(|device_id| !device_id.is_empty())
2094-
.map(|device_id| vec![device_id.to_string()])
2095-
.unwrap_or_default(),
2098+
device_ids: if gpu_count.is_none() {
2099+
gpu_device
2100+
.filter(|device_id| !device_id.is_empty())
2101+
.map(|device_id| vec![device_id.to_string()])
2102+
.unwrap_or_default()
2103+
} else {
2104+
Vec::new()
2105+
},
2106+
count: gpu_count,
20962107
}),
20972108
})
20982109
}
@@ -6296,7 +6307,7 @@ mod tests {
62966307
parse_credential_pairs, plaintext_gateway_is_remote, progress_step_from_metadata,
62976308
provisioning_timeout_message, ready_false_condition_message, resolve_from,
62986309
resource_requirements_from_cli, sandbox_should_persist, service_expose_status_error,
6299-
service_url_for_gateway, source_requests_gpu,
6310+
service_url_for_gateway,
63006311
};
63016312
use crate::TEST_ENV_LOCK;
63026313
use hyper::StatusCode;
@@ -6644,32 +6655,54 @@ mod tests {
66446655
}
66456656

66466657
#[test]
6647-
fn source_requests_gpu_detects_known_community_gpu_name() {
6648-
assert!(source_requests_gpu("nvidia-gpu"));
6649-
assert!(!source_requests_gpu("base"));
6658+
fn image_requests_gpu_detects_known_community_gpu_name() {
6659+
assert!(image_requests_gpu("nvidia-gpu"));
6660+
assert!(!image_requests_gpu("base"));
66506661
}
66516662

66526663
#[test]
66536664
fn resource_requirements_from_cli_uses_presence_with_empty_device_ids_for_default_gpu() {
6654-
let request = resource_requirements_from_cli(true, None)
6665+
let request = resource_requirements_from_cli(true, None, None)
66556666
.expect("resource requirements should be present");
66566667
let gpu = request.gpu.expect("gpu request should be present");
66576668

66586669
assert!(gpu.device_ids.is_empty());
6670+
assert_eq!(gpu.count, None);
66596671
}
66606672

66616673
#[test]
66626674
fn resource_requirements_from_cli_maps_gpu_device_to_one_device_id() {
6663-
let request = resource_requirements_from_cli(true, Some("0000:2d:00.0"))
6675+
let request = resource_requirements_from_cli(true, Some("0000:2d:00.0"), None)
66646676
.expect("resource requirements should be present");
66656677
let gpu = request.gpu.expect("gpu request should be present");
66666678

66676679
assert_eq!(gpu.device_ids, vec!["0000:2d:00.0"]);
6680+
assert_eq!(gpu.count, None);
6681+
}
6682+
6683+
#[test]
6684+
fn resource_requirements_from_cli_maps_gpu_device_without_explicit_gpu_flag() {
6685+
let request = resource_requirements_from_cli(false, Some("0000:2d:00.0"), None)
6686+
.expect("resource requirements should be present");
6687+
let gpu = request.gpu.expect("gpu request should be present");
6688+
6689+
assert_eq!(gpu.device_ids, vec!["0000:2d:00.0"]);
6690+
assert_eq!(gpu.count, None);
6691+
}
6692+
6693+
#[test]
6694+
fn resource_requirements_from_cli_maps_gpu_count() {
6695+
let request = resource_requirements_from_cli(false, None, Some(2))
6696+
.expect("resource requirements should be present");
6697+
let gpu = request.gpu.expect("gpu request should be present");
6698+
6699+
assert!(gpu.device_ids.is_empty());
6700+
assert_eq!(gpu.count, Some(2));
66686701
}
66696702

66706703
#[test]
66716704
fn resource_requirements_from_cli_omits_placement_when_not_requested() {
6672-
assert!(resource_requirements_from_cli(false, Some("0")).is_none());
6705+
assert!(resource_requirements_from_cli(false, None, None).is_none());
66736706
}
66746707

66756708
#[test]

crates/openshell-cli/tests/sandbox_create_lifecycle_integration.rs

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -732,6 +732,7 @@ async fn sandbox_create_keeps_command_sessions_by_default() {
732732
None,
733733
None,
734734
None,
735+
None,
735736
&[],
736737
None,
737738
None,
@@ -770,6 +771,7 @@ async fn sandbox_create_sends_cpu_and_memory_limits_only() {
770771
true,
771772
false,
772773
None,
774+
None,
773775
Some("500m"),
774776
Some("2Gi"),
775777
None,
@@ -858,6 +860,7 @@ async fn sandbox_create_returns_vm_error_without_waiting_for_timeout() {
858860
None,
859861
None,
860862
None,
863+
None,
861864
&[],
862865
None,
863866
None,
@@ -910,6 +913,7 @@ async fn sandbox_create_keeps_waiting_while_vm_progress_arrives() {
910913
None,
911914
None,
912915
None,
916+
None,
913917
&[],
914918
None,
915919
None,
@@ -954,6 +958,7 @@ async fn sandbox_create_times_out_when_only_logs_arrive() {
954958
None,
955959
None,
956960
None,
961+
None,
957962
&[],
958963
None,
959964
None,
@@ -994,6 +999,7 @@ async fn sandbox_create_deletes_command_sessions_with_no_keep() {
994999
None,
9951000
None,
9961001
None,
1002+
None,
9971003
&[],
9981004
None,
9991005
None,
@@ -1038,6 +1044,7 @@ async fn sandbox_create_deletes_shell_sessions_with_no_keep() {
10381044
None,
10391045
None,
10401046
None,
1047+
None,
10411048
&[],
10421049
None,
10431050
None,
@@ -1082,6 +1089,7 @@ async fn sandbox_create_keeps_sandbox_with_hidden_keep_flag() {
10821089
None,
10831090
None,
10841091
None,
1092+
None,
10851093
&[],
10861094
None,
10871095
None,
@@ -1126,6 +1134,7 @@ async fn sandbox_create_keeps_sandbox_with_forwarding() {
11261134
None,
11271135
None,
11281136
None,
1137+
None,
11291138
&[],
11301139
None,
11311140
Some(openshell_core::forward::ForwardSpec::new(forward_port)),

crates/openshell-core/src/gpu.rs

Lines changed: 6 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -30,7 +30,10 @@ mod tests {
3030

3131
#[test]
3232
fn cdi_gpu_device_ids_defaults_empty_request_to_all_gpus() {
33-
let request = GpuSpec { device_ids: vec![] };
33+
let request = GpuSpec {
34+
device_ids: vec![],
35+
count: None,
36+
};
3437

3538
assert_eq!(
3639
cdi_gpu_device_ids(Some(&request)),
@@ -42,6 +45,7 @@ mod tests {
4245
fn cdi_gpu_device_ids_passes_single_device_id_through() {
4346
let request = GpuSpec {
4447
device_ids: vec!["nvidia.com/gpu=0".to_string()],
48+
count: None,
4549
};
4650

4751
assert_eq!(
@@ -57,6 +61,7 @@ mod tests {
5761
"nvidia.com/gpu=0".to_string(),
5862
"nvidia.com/gpu=1".to_string(),
5963
],
64+
count: None,
6065
};
6166

6267
assert_eq!(

crates/openshell-driver-docker/README.md

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -30,7 +30,7 @@ contract:
3030
| `cap_add` | Grants supervisor-only capabilities required for namespace setup and process inspection. |
3131
| `apparmor=unconfined` | Avoids Docker's default profile blocking required mount operations. |
3232
| `restart_policy = unless-stopped` | Keeps managed sandboxes resumable across daemon or gateway restarts. |
33-
| CDI GPU request | Uses explicit placement GPU device IDs when set; otherwise requests all NVIDIA GPUs for GPU placement requests when daemon CDI support is detected. |
33+
| CDI GPU request | Uses explicit placement GPU device IDs when set; otherwise requests all NVIDIA GPUs for GPU placement requests when daemon CDI support is detected. GPU count requests are rejected by this driver. |
3434

3535
The agent child process does not retain these supervisor privileges.
3636

crates/openshell-driver-docker/src/lib.rs

Lines changed: 16 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -381,6 +381,21 @@ impl DockerComputeDriver {
381381
}
382382

383383
fn validate_gpu_request(gpu: Option<&GpuSpec>, supports_gpu: bool) -> Result<(), Status> {
384+
if let Some(gpu) = gpu {
385+
if gpu.count == Some(0) {
386+
return Err(Status::invalid_argument("gpu.count must be greater than 0"));
387+
}
388+
if gpu.count.is_some() && !gpu.device_ids.is_empty() {
389+
return Err(Status::invalid_argument(
390+
"gpu.count is mutually exclusive with gpu.device_ids",
391+
));
392+
}
393+
if gpu.count.is_some() {
394+
return Err(Status::invalid_argument(
395+
"docker compute driver does not support GPU count requests",
396+
));
397+
}
398+
}
384399
if gpu.is_some() && !supports_gpu {
385400
return Err(Status::failed_precondition(
386401
"docker GPU sandboxes require Docker CDI support. Enable CDI on the Docker daemon, then restart the OpenShell gateway/server so GPU capability is detected.",
@@ -1057,7 +1072,7 @@ fn build_container_create_body(
10571072
.as_ref()
10581073
.and_then(|placement| placement.gpu.as_ref()),
10591074
),
1060-
mounts: Some(build_mounts(config)),
1075+
binds: Some(build_binds(config)),
10611076
restart_policy: Some(RestartPolicy {
10621077
name: Some(RestartPolicyNameEnum::UNLESS_STOPPED),
10631078
maximum_retry_count: None,

0 commit comments

Comments
 (0)