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
Original file line number Diff line number Diff line change
Expand Up @@ -11,14 +11,12 @@
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 (
Expand All @@ -28,9 +26,8 @@
parse_submission,
)
from data_designer.slurm.launcher.runner import CommandRunner, SubprocessRunner
from data_designer.slurm.state import SchedulerIdentity
from data_designer.slurm.state import SchedulerIdentity, SchedulerJobIdentity

_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

Expand Down Expand Up @@ -86,7 +83,7 @@ def submit_script(self, script: str) -> SlurmJobSubmissionReceipt:
)
return parse_submission(output)

def query_queue(self, selectors: Sequence[_JobSelector]) -> tuple[SlurmQueueEntry, ...]:
def query_queue(self, selectors: Sequence[SchedulerJobIdentity]) -> tuple[SlurmQueueEntry, ...]:
"""Return normalized active-queue rows for explicit managed jobs."""
requested = tuple(selectors)
jobs = _format_selectors(requested)
Expand All @@ -107,7 +104,7 @@ def query_queue(self, selectors: Sequence[_JobSelector]) -> tuple[SlurmQueueEntr
)
return tuple(entry for entry in entries if entry.job_identity not in ignored)

def query_accounting(self, selectors: Sequence[_JobSelector]) -> tuple[SlurmAccountingEntry, ...]:
def query_accounting(self, selectors: Sequence[SchedulerJobIdentity]) -> tuple[SlurmAccountingEntry, ...]:
"""Return normalized accounting rows for explicit managed jobs."""
requested = tuple(selectors)
jobs = _format_selectors(requested)
Expand All @@ -130,7 +127,7 @@ def query_accounting(self, selectors: Sequence[_JobSelector]) -> tuple[SlurmAcco
)
return tuple(entry for entry in entries if entry.job_identity not in ignored)

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

Expand Down Expand Up @@ -162,13 +159,13 @@ def _run(self, command: Sequence[str], *, input_text: str | None = None) -> str:
return stdout


def _format_selectors(selectors: Sequence[_JobSelector]) -> str:
def _format_selectors(selectors: Sequence[SchedulerJobIdentity]) -> 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:
def _format_selector(selector: SchedulerJobIdentity) -> str:
if isinstance(selector, SchedulerIdentity):
job_id = _format_job_id(selector.array_job_id)
if selector.array_task_id > _MAX_SLURM_INTEGER:
Expand All @@ -184,16 +181,16 @@ def _format_job_id(value: object) -> str:


def _validate_observed_job_identities(
job_identities: Sequence[SlurmObservedJobIdentity],
selectors: Sequence[_JobSelector],
job_identities: Sequence[SchedulerJobIdentity],
selectors: Sequence[SchedulerJobIdentity],
*,
command: str,
) -> frozenset[SlurmObservedJobIdentity]:
) -> frozenset[SchedulerJobIdentity]:
"""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()
ignored: set[SchedulerJobIdentity] = set()
for job_identity in job_identities:
if type(job_identity) is int and job_identity in selected_job_ids:
continue
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -6,11 +6,8 @@
from __future__ import annotations

from dataclasses import dataclass
from typing import TypeAlias

from data_designer.slurm.state import SchedulerIdentity, SchedulerState

SlurmObservedJobIdentity: TypeAlias = int | SchedulerIdentity
from data_designer.slurm.state import SchedulerJobIdentity, SchedulerState


@dataclass(frozen=True)
Expand All @@ -32,14 +29,14 @@ class SlurmProcessExitCode:
class SlurmQueueEntry:
"""One transient normalized active-queue entry."""

job_identity: SlurmObservedJobIdentity
job_identity: SchedulerJobIdentity
state: SchedulerState


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

job_identity: SlurmObservedJobIdentity
job_identity: SchedulerJobIdentity
state: SchedulerState
process_exit_code: SlurmProcessExitCode
Original file line number Diff line number Diff line change
Expand Up @@ -11,11 +11,10 @@
from data_designer.slurm.launcher.models import (
SlurmAccountingEntry,
SlurmJobSubmissionReceipt,
SlurmObservedJobIdentity,
SlurmProcessExitCode,
SlurmQueueEntry,
)
from data_designer.slurm.state import SchedulerIdentity, SchedulerState
from data_designer.slurm.state import SchedulerIdentity, SchedulerJobIdentity, SchedulerState

_ARRAY_ID_PATTERN = re.compile(r"^(?P<job>[1-9][0-9]*)_(?P<task>[0-9]+)$")
_JOB_ID_PATTERN = re.compile(r"^[1-9][0-9]*$")
Expand Down Expand Up @@ -71,7 +70,7 @@ def parse_submission(output: str) -> SlurmJobSubmissionReceipt:
def parse_queue(output: str) -> tuple[SlurmQueueEntry, ...]:
"""Parse ``squeue --format=%i|%T`` rows."""
entries: list[SlurmQueueEntry] = []
identities: set[SlurmObservedJobIdentity] = set()
identities: set[SchedulerJobIdentity] = set()
for line_number, line in _collect_nonempty_lines(output):
fields = line.split("|")
if len(fields) != 2:
Expand All @@ -85,7 +84,7 @@ def parse_queue(output: str) -> tuple[SlurmQueueEntry, ...]:
def parse_accounting(output: str) -> tuple[SlurmAccountingEntry, ...]:
"""Parse job and array-task rows from ``sacct --format=JobID,State,ExitCode``."""
entries: list[SlurmAccountingEntry] = []
identities: set[SlurmObservedJobIdentity] = set()
identities: set[SchedulerJobIdentity] = set()
for line_number, line in _collect_nonempty_lines(output):
fields = line.split("|")
if len(fields) != 3:
Expand Down Expand Up @@ -184,7 +183,7 @@ def _parse_array_identity(value: str, *, command: str, line_number: int) -> Sche
)


def _parse_job_identity(value: str, *, command: str, line_number: int) -> SlurmObservedJobIdentity:
def _parse_job_identity(value: str, *, command: str, line_number: int) -> SchedulerJobIdentity:
message = f"{command} line {line_number} contains an invalid job or array-task ID"
if _JOB_ID_PATTERN.fullmatch(value) is not None:
return _parse_decimal(value, message=message)
Expand Down Expand Up @@ -218,8 +217,8 @@ def _parse_decimal(value: str, *, message: str) -> int:


def _reject_duplicate(
job_identity: SlurmObservedJobIdentity,
identities: set[SlurmObservedJobIdentity],
job_identity: SchedulerJobIdentity,
identities: set[SchedulerJobIdentity],
*,
command: str,
line_number: int,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,7 @@
from data_designer.slurm.state.artifacts import compute_candidate_schema_digest
from data_designer.slurm.state.base import (
SchedulerIdentity,
SchedulerJobIdentity,
StateRecord,
StateValue,
)
Expand All @@ -38,6 +39,12 @@
RunManifest,
ShardManifest,
)
from data_designer.slurm.state.observation import (
SchedulerAccountingRecord,
SchedulerObservationClient,
SchedulerObservationCollector,
SchedulerQueueRecord,
)
from data_designer.slurm.state.outputs import (
CANDIDATE_OUTPUT_FORMAT,
MAXIMUM_CANDIDATE_OUTPUT_FILES,
Expand Down Expand Up @@ -66,6 +73,13 @@
SchedulerObservation,
SchedulerState,
)
from data_designer.slurm.state.status import (
AttemptStatus,
EffectiveRunState,
GenerationState,
RunStatus,
ShardStatus,
)
from data_designer.slurm.state.validation import (
StateContractError,
validate_attempt_manifest,
Expand All @@ -80,9 +94,11 @@
)

if TYPE_CHECKING:
from data_designer.slurm.state.observer import SlurmStateReconciler # noqa: F401
from data_designer.slurm.state.store import SlurmStateWriter # noqa: F401

_LAZY_IMPORTS: dict[str, tuple[str, str]] = {
"SlurmStateReconciler": ("data_designer.slurm.state.observer", "SlurmStateReconciler"),
"SlurmStateWriter": ("data_designer.slurm.state.store", "SlurmStateWriter"),
}

Expand All @@ -92,6 +108,7 @@
"AttemptManifest",
"AttemptId",
"AttemptReadiness",
"AttemptStatus",
"AttemptTerminalClassification",
"CandidateOutcome",
"CANDIDATE_OUTPUT_FORMAT",
Expand All @@ -105,23 +122,33 @@
"ContractValue",
"DeploymentReadiness",
"EffectiveAttemptState",
"EffectiveRunState",
"EndpointPublicationState",
"Identifier",
"GenerationState",
"ProbeEvidence",
"ProbeOutcome",
"ReadinessState",
"ReasonCode",
"RecordRange",
"RunManifest",
"RunStatus",
"ResumeWorkspace",
"SchedulerIdentity",
"SchedulerJobIdentity",
"SchedulerAccountingRecord",
"SchedulerObservationClient",
"SchedulerObservationCollector",
"SchedulerQueueRecord",
"SchedulerObservation",
"SchedulerState",
"Sha256Digest",
"ShardManifest",
"ShardStatus",
"ShardId",
"ShardWinner",
"SlurmStateError",
"SlurmStateReconciler",
"SlurmStateWriter",
"StateConflictError",
"StateContractError",
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@
from __future__ import annotations

from datetime import datetime, timedelta
from typing import TypeAlias

from pydantic import NonNegativeInt, PositiveInt

Expand Down Expand Up @@ -44,10 +45,14 @@ class SchedulerIdentity(StateValue):
array_task_id: NonNegativeInt


SchedulerJobIdentity: TypeAlias = PositiveInt | SchedulerIdentity


__all__ = [
"ArtifactReference",
"Identifier",
"SchedulerIdentity",
"SchedulerJobIdentity",
"Sha256Digest",
"StateRecord",
"StateValue",
Expand Down
Loading
Loading