Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
23 commits
Select commit Hold shift + click to select a range
d7a0542
feat: add Slurm command client and renderer
nabinchha Aug 25, 2026
3a75c8c
fix: tighten Slurm launcher boundaries
nabinchha Aug 25, 2026
7c9ab7c
fix: use native sacct output fields
nabinchha Aug 25, 2026
996ac7f
test: require expanded Slurm array queries
nabinchha Aug 25, 2026
0bf2142
fix: sanitize Slurm command diagnostics
nabinchha Aug 25, 2026
65804af
fix: reject invalid visible GPU memory
nabinchha Aug 25, 2026
926a221
fix: hash rendered inputs without manifests
nabinchha Aug 25, 2026
c5df7d6
fix: make Slurm output decoding stable
nabinchha Aug 25, 2026
6d20a8c
fix: validate finite command timeouts
nabinchha Aug 25, 2026
e2297b1
fix: harden Slurm boundary inputs
nabinchha Aug 25, 2026
1ea40f0
fix: reject malformed Slurm GRES lists
nabinchha Aug 25, 2026
1f0c70b
fix: isolate submitted Slurm environments
nabinchha Aug 25, 2026
edb5140
test: verify launcher wheel import
nabinchha Aug 25, 2026
15b9dc3
fix: pin batch host tool path
nabinchha Aug 25, 2026
3474843
test: cover Slurm launcher boundaries
nabinchha Aug 25, 2026
3f57641
fix: correlate Slurm query responses
nabinchha Aug 25, 2026
ab73562
fix: harden Slurm command boundaries
nabinchha Aug 25, 2026
57a9377
fix: fall back from empty command path
nabinchha Aug 25, 2026
ddad8f0
fix: align Slurm launcher semantics
nabinchha Aug 26, 2026
709963c
fix Slurm job observation contracts
nabinchha Aug 26, 2026
87a520f
refine Slurm launcher boundaries
nabinchha Aug 27, 2026
ef82bfe
fix Slurm selector filtering ownership
nabinchha Aug 27, 2026
d5db22e
simplify Slurm launcher structure
nabinchha Aug 27, 2026
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
Original file line number Diff line number Diff line change
@@ -0,0 +1,6 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0

"""Internal Slurm submission, observation, and batch-rendering helpers."""

from __future__ import annotations
Original file line number Diff line number Diff line change
@@ -0,0 +1,215 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0

"""Internal typed argument-vector client for Slurm command-line tools."""

from __future__ import annotations

import re
import subprocess
import unicodedata
from collections.abc import Sequence
from dataclasses import dataclass
from pathlib import Path
from typing import TypeAlias

from data_designer.slurm.contracts import Identifier
from data_designer.slurm.launcher.errors import SlurmCommandError, SlurmCommandOutputError
from data_designer.slurm.launcher.models import (
SlurmAccountingEntry,
SlurmJobSubmissionReceipt,
SlurmObservedJobIdentity,
SlurmQueueEntry,
)
from data_designer.slurm.launcher.parsing import (
parse_accounting,
parse_gpu_counts,
parse_queue,
parse_submission,
)
from data_designer.slurm.launcher.runner import CommandRunner, SubprocessRunner
from data_designer.slurm.state import SchedulerIdentity

_JobSelector: TypeAlias = int | SchedulerIdentity
_IDENTIFIER_PATTERN = re.compile(r"^[A-Za-z0-9][A-Za-z0-9._-]{0,127}$")
_MAX_SLURM_INTEGER = (1 << 32) - 1


@dataclass(frozen=True)
class SlurmExecutables:
"""Executable paths used for bounded Slurm operations."""

sbatch: str = "sbatch"
squeue: str = "squeue"
sacct: str = "sacct"
scancel: str = "scancel"
sinfo: str = "sinfo"

def __post_init__(self) -> None:
for executable in (self.sbatch, self.squeue, self.sacct, self.scancel, self.sinfo):
_validate_argument(executable, field_name="Slurm executable")
if any(character.isspace() for character in executable):
raise ValueError("Slurm executable must be one argument-vector token")


class SlurmCommandClient:
"""Submit, observe, and cancel Slurm jobs through structured commands."""

_executables: SlurmExecutables
_runner: CommandRunner

def __init__(
self,
runner: CommandRunner | None = None,
*,
executables: SlurmExecutables | None = None,
) -> None:
self._runner = runner if runner is not None else SubprocessRunner()
self._executables = executables if executables is not None else SlurmExecutables()

def submit(self, script_path: str | Path) -> SlurmJobSubmissionReceipt:
"""Submit one rendered batch script and return its assigned job ID."""
path = str(script_path)
_validate_argument(path, field_name="batch script path")
if path.startswith("-"):
raise ValueError("batch script path must not begin with '-'; prefix relative paths with './'")
output = self._run((self._executables.sbatch, "--parsable", "--export=NIL", path))
return parse_submission(output)

def query_queue(self, selectors: Sequence[_JobSelector]) -> tuple[SlurmQueueEntry, ...]:
"""Return normalized active-queue rows for explicit managed jobs."""
requested = tuple(selectors)
jobs = _format_selectors(requested)
output = self._run(
(
self._executables.squeue,
"--noheader",
"--array",
"--format=%i|%T",
f"--jobs={jobs}",
)
)
entries = parse_queue(output)
ignored = _validate_observed_job_identities(
tuple(entry.job_identity for entry in entries),
requested,
command="squeue",
)
return tuple(entry for entry in entries if entry.job_identity not in ignored)

def query_accounting(self, selectors: Sequence[_JobSelector]) -> tuple[SlurmAccountingEntry, ...]:
"""Return normalized accounting rows for explicit managed jobs."""
requested = tuple(selectors)
jobs = _format_selectors(requested)
output = self._run(
(
self._executables.sacct,
"--noheader",
"--array",
"--allocations",
"--parsable2",
"--format=JobID,State,ExitCode",
f"--jobs={jobs}",
)
)
entries = parse_accounting(output)
ignored = _validate_observed_job_identities(
tuple(entry.job_identity for entry in entries),
requested,
command="sacct",
)
return tuple(entry for entry in entries if entry.job_identity not in ignored)

def cancel(self, selector: _JobSelector) -> None:
"""Cancel one managed Slurm job, array, or array task."""
self._run((self._executables.scancel, _format_selector(selector)))

def query_gpu_counts(self, *, partition: Identifier | None = None) -> tuple[int, ...]:
"""Return configured GPU counts reported for eligible node groups."""
command = [self._executables.sinfo, "--noheader", "--format=%G"]
if partition is not None:
if type(partition) is not str or _IDENTIFIER_PATTERN.fullmatch(partition) is None:
raise ValueError("Slurm partition must be a valid identifier")
command.append(f"--partition={partition}")
return parse_gpu_counts(self._run(command))

def _run(self, command: Sequence[str]) -> str:
command_name = Path(command[0]).name
try:
completed = self._runner.run(command)
except (OSError, subprocess.SubprocessError) as error:
raise SlurmCommandError(f"{command_name} could not be executed: {_format_error_detail(error)}") from error
returncode = getattr(completed, "returncode", None)
stdout = getattr(completed, "stdout", None)
stderr = getattr(completed, "stderr", None)
if type(returncode) is not int or not isinstance(stdout, str) or not isinstance(stderr, str):
raise SlurmCommandError(f"{command_name} returned a malformed process result")
if returncode:
detail = _normalize_bounded_text(stderr) or "no diagnostic output"
raise SlurmCommandError(f"{command_name} failed with exit code {returncode}: {detail}")
return stdout


def _format_selectors(selectors: Sequence[_JobSelector]) -> str:
if not selectors:
raise ValueError("at least one managed Slurm job selector is required")
return ",".join(dict.fromkeys(_format_selector(selector) for selector in selectors))


def _format_selector(selector: _JobSelector) -> str:
if isinstance(selector, SchedulerIdentity):
job_id = _format_job_id(selector.array_job_id)
if selector.array_task_id > _MAX_SLURM_INTEGER:
raise ValueError("Slurm array-task IDs must be non-negative 32-bit integers")
return f"{job_id}_{selector.array_task_id}"
return _format_job_id(selector)


def _format_job_id(value: object) -> str:
if type(value) is not int or not 0 < value <= _MAX_SLURM_INTEGER:
raise ValueError("Slurm job IDs must be positive 32-bit integers")
return str(value)


def _validate_observed_job_identities(
job_identities: Sequence[SlurmObservedJobIdentity],
selectors: Sequence[_JobSelector],
*,
command: str,
) -> frozenset[SlurmObservedJobIdentity]:
"""Validate result correlation and return unselected aggregate rows."""
selected_job_ids = {selector for selector in selectors if type(selector) is int}
selected_array_tasks = {selector for selector in selectors if isinstance(selector, SchedulerIdentity)}
selected_array_job_ids = {selector.array_job_id for selector in selected_array_tasks}
ignored: set[SlurmObservedJobIdentity] = set()
for job_identity in job_identities:
if type(job_identity) is int and job_identity in selected_job_ids:
continue
if type(job_identity) is int and job_identity in selected_array_job_ids:
ignored.add(job_identity)
continue
if isinstance(job_identity, SchedulerIdentity) and (
job_identity in selected_array_tasks or job_identity.array_job_id in selected_job_ids
):
continue
raise SlurmCommandOutputError(f"{command} returned an unrequested job or array-task ID")
return frozenset(ignored)


def _validate_argument(value: str, *, field_name: str) -> None:
if type(value) is not str or not value:
raise ValueError(f"{field_name} must not be empty")
if any(ord(character) < 32 or ord(character) == 127 for character in value):
raise ValueError(f"{field_name} must not contain control characters")


def _normalize_bounded_text(value: str, *, limit: int = 512) -> str:
sanitized = "".join(" " if unicodedata.category(character).startswith("C") else character for character in value)
normalized = " ".join(sanitized.split())
return normalized if len(normalized) <= limit else f"{normalized[: limit - 3]}..."


def _format_error_detail(error: BaseException) -> str:
if isinstance(error, subprocess.TimeoutExpired):
return "command timed out"
return _normalize_bounded_text(str(error)) or error.__class__.__name__
Original file line number Diff line number Diff line change
@@ -0,0 +1,22 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0

"""Internal normalized errors for the Slurm launcher boundary."""

from __future__ import annotations


class SlurmLauncherError(RuntimeError):
"""Base error for structured Slurm launcher operations."""


class SlurmCommandError(SlurmLauncherError):
"""A Slurm command could not be executed successfully."""


class SlurmCommandOutputError(SlurmLauncherError, ValueError):
"""A Slurm command returned output that violates its requested format."""


class SlurmBatchRenderError(SlurmLauncherError, ValueError):
"""A resolved plan cannot be rendered as a safe batch script."""
Original file line number Diff line number Diff line change
@@ -0,0 +1,45 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0

"""Transient typed values returned by Slurm commands."""

from __future__ import annotations

from dataclasses import dataclass
from typing import TypeAlias

from data_designer.slurm.state import SchedulerIdentity, SchedulerState

SlurmObservedJobIdentity: TypeAlias = int | SchedulerIdentity


@dataclass(frozen=True)
class SlurmJobSubmissionReceipt:
"""Job identity returned for one accepted non-federated submission."""

job_id: int


@dataclass(frozen=True)
class SlurmProcessExitCode:
"""Slurm's process status and terminating signal pair."""

exit_status: int
termination_signal: int


@dataclass(frozen=True)
class SlurmQueueEntry:
"""One transient normalized active-queue entry."""

job_identity: SlurmObservedJobIdentity
state: SchedulerState


@dataclass(frozen=True)
class SlurmAccountingEntry:
"""One transient normalized accounting entry."""

job_identity: SlurmObservedJobIdentity
state: SchedulerState
process_exit_code: SlurmProcessExitCode
Loading
Loading