Add descriptor-batch SDK reference workload

This commit is contained in:
Emil
2026-08-01 23:27:53 +03:00
parent c43af32495
commit 96169086f0
9 changed files with 1363 additions and 14 deletions
+7 -3
View File
@@ -48,7 +48,9 @@ def current_environment_digest() -> str:
return "sha256:" + hashlib.sha256(payload.encode("utf-8")).hexdigest()
def similarity_search_sdk_adapter(*, shard_rows: int = 10_000) -> LegacyDistributedWorkloadAdapter:
def similarity_search_sdk_adapter(
*, shard_rows: int = 10_000
) -> LegacyDistributedWorkloadAdapter:
dataset_schema = ArtifactSchema(
SchemaRef("molecule-table", 1),
"text/tab-separated-values",
@@ -119,7 +121,9 @@ def similarity_search_sdk_adapter(*, shard_rows: int = 10_000) -> LegacyDistribu
def default_sdk_registry(*, shard_rows: int = 10_000) -> WorkloadRegistry:
registry = WorkloadRegistry()
registry.register(similarity_search_sdk_adapter(shard_rows=shard_rows).definition(), enabled=True)
registry.register(
similarity_search_sdk_adapter(shard_rows=shard_rows).definition(), enabled=True
)
return registry
@@ -135,7 +139,7 @@ def default_sdk_runtime() -> RuntimeCapabilities:
protocol_version="1.0.0",
profiles=("core-batch-v1",),
features={"artifact-collections": "1.0.0", "exact-verifier": "1.0.0"},
workload_capabilities=("similarity-search",),
workload_capabilities=("similarity-search", "descriptor-batch"),
inventory=ResourceInventory(
cpu_cores=max(os.cpu_count() or 1, 1),
memory_mb=4096,
+41
View File
@@ -0,0 +1,41 @@
"""SDK-native ``descriptor-batch`` reference workload.
See ``core.py`` for the pinned scientific contract and ``definition.py`` for
the manifest-backed planner/runner/reducer handlers.
"""
from .core import (
DESCRIPTOR_COLUMNS,
DESCRIPTOR_NAMES,
DescriptorRow,
compute_descriptor_batch,
concatenate_descriptor_shards,
descriptor_calculator,
validate_descriptor_names,
write_descriptor_rows,
write_descriptor_shards,
)
from .definition import (
MAP_ENTRY_POINT,
REDUCE_ENTRY_POINT,
DescriptorBatchWorkload,
descriptor_batch_sdk_definition,
workload_definition,
)
__all__ = [
"DESCRIPTOR_COLUMNS",
"DESCRIPTOR_NAMES",
"MAP_ENTRY_POINT",
"REDUCE_ENTRY_POINT",
"DescriptorBatchWorkload",
"DescriptorRow",
"compute_descriptor_batch",
"concatenate_descriptor_shards",
"descriptor_batch_sdk_definition",
"descriptor_calculator",
"validate_descriptor_names",
"workload_definition",
"write_descriptor_rows",
"write_descriptor_shards",
]
+314
View File
@@ -0,0 +1,314 @@
"""Pinned RDKit 2D descriptor computation for the descriptor-batch workload.
The scientific contract of ``descriptor-batch`` is deliberately small and
fully pinned:
- exactly one output row per valid input molecule, in input order;
- RDKit canonical SMILES recomputed with ``MolToSmiles(..., canonical=True)``;
- the descriptor set is an explicit, versioned tuple of RDKit
``Descriptors.descList`` names (2D only), not a scan of installed names;
- float values are serialized with fixed ``%.6f`` formatting so that the
output is byte-identical for identical inputs and a pinned environment;
- invalid SMILES rows are either skipped (counted) or fail the run, selected
by the explicit ``skip_invalid`` parameter.
"""
from __future__ import annotations
import csv
from dataclasses import dataclass
from functools import lru_cache
from pathlib import Path
from typing import Any, Iterator, Mapping, Sequence
from rdkit import Chem
from rdkit.ML.Descriptors.MoleculeDescriptors import MolecularDescriptorCalculator
from scimesh.chemistry.dataset import iter_rows
# Explicit pinned list. Names must exist in the installed RDKit ``descList``;
# the list itself is the reproducibility contract and must change version
# together with the workload (descriptor-batch@1.0.0).
DESCRIPTOR_NAMES: tuple[str, ...] = (
"ExactMolWt",
"MolWt",
"HeavyAtomMolWt",
"HeavyAtomCount",
"NumHDonors",
"NumHAcceptors",
"NumRotatableBonds",
"NumHeteroatoms",
"NumRadicalElectrons",
"NumValenceElectrons",
"FractionCSP3",
"RingCount",
"NumAromaticRings",
"NumSaturatedRings",
"NumAliphaticRings",
"NumAromaticHeterocycles",
"NumSaturatedHeterocycles",
"NumAliphaticHeterocycles",
"NumAromaticCarbocycles",
"NumSaturatedCarbocycles",
"NumAliphaticCarbocycles",
"TPSA",
"LabuteASA",
"MolLogP",
"MolMR",
"BalabanJ",
"BertzCT",
"HallKierAlpha",
"Kappa1",
"Kappa2",
"Kappa3",
"Chi0",
"Chi1",
"Chi0n",
"Chi1n",
"Chi2n",
"Chi3n",
"Chi4n",
"Chi0v",
"Chi1v",
"Chi2v",
"Chi3v",
"Chi4v",
"PEOE_VSA1",
"PEOE_VSA2",
"PEOE_VSA3",
"PEOE_VSA4",
"PEOE_VSA5",
"PEOE_VSA6",
"PEOE_VSA7",
"PEOE_VSA8",
"PEOE_VSA9",
"PEOE_VSA10",
"PEOE_VSA11",
"PEOE_VSA12",
"PEOE_VSA13",
"PEOE_VSA14",
"SMR_VSA1",
"SMR_VSA2",
"SMR_VSA3",
"SMR_VSA4",
"SMR_VSA5",
"SMR_VSA6",
"SMR_VSA7",
"SMR_VSA8",
"SMR_VSA9",
"SMR_VSA10",
"SlogP_VSA1",
"SlogP_VSA2",
"SlogP_VSA3",
"SlogP_VSA4",
"SlogP_VSA5",
"SlogP_VSA6",
"SlogP_VSA7",
"SlogP_VSA8",
"SlogP_VSA9",
"SlogP_VSA10",
"SlogP_VSA11",
"SlogP_VSA12",
"NHOHCount",
"NOCount",
)
DESCRIPTOR_COLUMNS: tuple[str, ...] = (
"chembl_id",
"canonical_smiles",
) + DESCRIPTOR_NAMES
@lru_cache(maxsize=1)
def descriptor_calculator() -> MolecularDescriptorCalculator:
"""Build the pinned calculator once per process."""
return MolecularDescriptorCalculator(DESCRIPTOR_NAMES)
def validate_descriptor_names() -> None:
"""Fail fast when the pinned list is unavailable in the installed RDKit."""
from rdkit.Chem import Descriptors
available = {name for name, _ in Descriptors.descList}
missing = [name for name in DESCRIPTOR_NAMES if name not in available]
if missing:
raise ValueError(
"pinned descriptor-batch descriptors are missing from RDKit: "
+ ", ".join(missing)
)
@dataclass(frozen=True)
class DescriptorRow:
"""One canonical descriptor row for a valid input molecule."""
molecule_id: str
canonical_smiles: str
values: tuple[float, ...]
class DescriptorStats:
"""Row counters collected while computing a descriptor batch."""
def __init__(self) -> None:
self.scanned = 0
self.invalid = 0
self.emitted = 0
def as_metrics(self) -> dict[str, int]:
return {
"rows_scanned": self.scanned,
"invalid_rows": self.invalid,
"rows_emitted": self.emitted,
}
def iter_descriptor_rows(
input_path: Path,
*,
skip_invalid: bool = True,
) -> tuple[Iterator[DescriptorRow], DescriptorStats]:
"""Yield canonical descriptor rows in input order with streaming stats."""
calculator = descriptor_calculator()
stats = DescriptorStats()
def generate() -> Iterator[DescriptorRow]:
for row in iter_rows(input_path):
stats.scanned += 1
smiles = row.get("canonical_smiles", "")
molecule = Chem.MolFromSmiles(smiles)
if molecule is None:
stats.invalid += 1
if not skip_invalid:
raise ValueError(
f"row {stats.scanned} has an invalid canonical_smiles"
)
continue
canonical = Chem.MolToSmiles(molecule, canonical=True)
values = tuple(
float(value) for value in calculator.CalcDescriptors(molecule)
)
stats.emitted += 1
yield DescriptorRow(row.get("chembl_id", ""), canonical, values)
return generate(), stats
def write_descriptor_rows(
output_path: Path,
rows: Sequence[DescriptorRow],
) -> None:
"""Write a canonical one-row-per-input descriptor CSV with one header."""
output_path.parent.mkdir(parents=True, exist_ok=True)
with output_path.open("w", encoding="utf-8", newline="") as destination:
writer = csv.DictWriter(
destination, fieldnames=list(DESCRIPTOR_COLUMNS), lineterminator="\n"
)
writer.writeheader()
for row in rows:
writer.writerow(
{
"chembl_id": row.molecule_id,
"canonical_smiles": row.canonical_smiles,
**{
name: f"{value:.6f}"
for name, value in zip(DESCRIPTOR_NAMES, row.values)
},
}
)
def compute_descriptor_batch(
input_path: Path,
output_path: Path,
*,
skip_invalid: bool = True,
) -> dict[str, int]:
"""Single-process reference: read the whole input and write the CSV."""
rows, stats = iter_descriptor_rows(input_path, skip_invalid=skip_invalid)
materialized = list(rows)
write_descriptor_rows(output_path, materialized)
return stats.as_metrics()
def write_descriptor_shards(
input_path: Path,
workspace: Path,
shard_rows: int,
) -> list[Path]:
"""Split the input TSV into deterministic row-bounded shards with headers."""
if (
isinstance(shard_rows, bool)
or not isinstance(shard_rows, int)
or shard_rows < 1
):
raise ValueError("shard_rows must be a positive integer")
paths: list[Path] = []
current: Path | None = None
destination = None
writer = None
rows_in_shard = 0
try:
with input_path.open("r", encoding="utf-8", newline="") as source:
reader = csv.DictReader(source, delimiter="\t")
fieldnames = tuple(reader.fieldnames or ())
if not {"chembl_id", "canonical_smiles"}.issubset(set(fieldnames)):
raise ValueError(
"dataset is missing required columns: chembl_id, canonical_smiles"
)
for row in reader:
if destination is None or rows_in_shard == shard_rows:
if destination is not None:
destination.close()
current = workspace / f"shard-{len(paths)}.tsv"
destination = current.open("w", encoding="utf-8", newline="")
writer = csv.DictWriter(
destination,
fieldnames=list(fieldnames),
delimiter="\t",
lineterminator="\n",
)
writer.writeheader()
paths.append(current)
rows_in_shard = 0
assert writer is not None
writer.writerow(row)
rows_in_shard += 1
finally:
if destination is not None:
destination.close()
if not paths:
raise ValueError("dataset has no data rows")
return paths
def concatenate_descriptor_shards(
partial_paths: Sequence[Path],
output_path: Path,
) -> dict[str, int]:
"""Merge shard partial CSVs by shard index with exactly one header.
Every partial is a full CSV with the same header. The first partial is
copied verbatim; each later partial contributes only its data rows, so the
merged file is byte-identical to the single-process reference for the same
input rows.
"""
if not partial_paths:
raise ValueError("descriptor reducer requires at least one partial")
output_path.parent.mkdir(parents=True, exist_ok=True)
rows_emitted = 0
with output_path.open("w", encoding="utf-8", newline="") as destination:
for index, partial in enumerate(partial_paths):
with partial.open("r", encoding="utf-8", newline="") as source:
for line_index, line in enumerate(source):
if line_index == 0:
if index > 0:
continue
if line.rstrip("\r\n") != ",".join(DESCRIPTOR_COLUMNS):
raise ValueError(
"partial descriptor CSV has an invalid header"
)
destination.write(line)
if line_index > 0:
rows_emitted += 1
return {"partial_count": len(partial_paths), "rows_emitted": rows_emitted}
+443
View File
@@ -0,0 +1,443 @@
"""SDK-native ``descriptor-batch`` workload definition and handlers.
This module is the first non-adapter reference workload built directly on the
``core-batch-v1`` profile: an explicit immutable manifest, a static map/reduce
workflow, a row-bounded planner, pinned descriptor computation, deterministic
shard concatenation, and the exact-artifact verifier. It is the intended
first ``untrusted_quorum`` candidate: ``byte_exact`` determinism with whole
file SHA-256 agreement from distinct owners.
"""
from __future__ import annotations
import hashlib
import shutil
from pathlib import Path
from typing import Any, Mapping, Sequence
from ..artifacts import (
ArtifactCollection,
ArtifactItem,
ArtifactRef,
ArtifactSchema,
Cardinality,
CollectionKind,
OutputManifest,
PortSpec,
)
from ..builtins import current_environment_digest, current_scimesh_package_digest
from ..execution import (
CheckpointPolicy,
ExecutionProfile,
NetworkPolicy,
RetryPolicy,
)
from ..identity import ComponentRef, SchemaRef, VersionRange, WorkloadId
from ..manifest import (
DeterminismProfile,
EnvironmentSpec,
PackageSpec,
TrustMode,
VerifierSpec,
WorkloadLimits,
WorkloadManifest,
)
from ..plans import JobRequest, TaskSpec, ValidatedJob, WorkflowPlan
from ..protocols import PlanningContext, ReduceContext, TaskContext
from ..registry import WorkloadDefinition
from ..resources import ResourceRequirements
from ..verification import ExactArtifactVerifier
from ..workflow import ArtifactEdge, PortRef, StageKind, StageSpec, WorkflowSpec
from .core import (
DESCRIPTOR_COLUMNS,
compute_descriptor_batch,
concatenate_descriptor_shards,
validate_descriptor_names,
write_descriptor_shards,
)
MAP_ENTRY_POINT = "scimesh.sdk.descriptors.definition:map_descriptors@v1"
REDUCE_ENTRY_POINT = "scimesh.sdk.descriptors.definition:reduce_descriptors@v1"
_DESCRIPTOR_PARAMETERS = ("skip_invalid",)
def _sha256_file(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as source:
for block in iter(lambda: source.read(1024 * 1024), b""):
digest.update(block)
return digest.hexdigest()
def _parameters_schema() -> dict[str, Any]:
return {
"type": "object",
"additionalProperties": False,
"properties": {
"skip_invalid": {
"type": "boolean",
"default": True,
"description": "Skip rows with invalid SMILES instead of failing",
},
},
}
def _input_schema() -> ArtifactSchema:
return ArtifactSchema(
SchemaRef("molecule-table", 1),
"text/tab-separated-values",
"utf-8",
max_bytes=10 * 1024 * 1024 * 1024,
validator=ComponentRef("delimited-table", 1),
validator_configuration={
"required_columns": ["canonical_smiles", "chembl_id"],
},
max_records=100_000_000,
canonicalizer="scimesh-tsv-v1",
)
def _descriptor_schema() -> ArtifactSchema:
return ArtifactSchema(
SchemaRef("descriptor-table", 1),
"text/csv",
"utf-8",
max_bytes=100 * 1024 * 1024 * 1024,
validator=ComponentRef("delimited-table", 1),
validator_configuration={
"columns": list(DESCRIPTOR_COLUMNS),
},
max_records=100_000_000,
canonicalizer="descriptor-table-v1",
)
class DescriptorBatchWorkload:
"""Manifest-backed planner, runner, and reducer for descriptor-batch.
The class follows the legacy adapter's structural pattern (one object
registered under each stage entry point) while remaining fully SDK-native:
sharding is explicit and deterministic, every artifact is sealed through
the bridge-owned sink, and no filesystem path ever enters a plan or task.
"""
def __init__(
self,
*,
shard_rows: int,
package_digest: str,
environment_digest: str,
) -> None:
if (
isinstance(shard_rows, bool)
or not isinstance(shard_rows, int)
or shard_rows < 1
):
raise ValueError("shard_rows must be a positive integer")
validate_descriptor_names()
self.entry_point = MAP_ENTRY_POINT
self.shard_rows = shard_rows
self.input_port = PortSpec(_input_schema())
self.partial_port = PortSpec(_descriptor_schema())
self.output_port = PortSpec(_descriptor_schema())
resources = ResourceRequirements(
profile="descriptor-cpu-v1",
cpu_cores=1,
memory_mb=1024,
scratch_mb=1024,
max_duration_seconds=3600,
)
execution = ExecutionProfile(
profile="descriptor-python-process-v1",
network=NetworkPolicy.TRUSTED,
timeout_seconds=3600,
checkpoint=CheckpointPolicy(),
)
limits = WorkloadLimits(
max_input_bytes=self.input_port.schema.max_bytes,
max_tasks=10_000,
max_output_bytes=self.output_port.schema.max_bytes,
)
trust_modes = ("trusted", "untrusted_quorum")
map_stage = StageSpec(
stage_id="map",
kind=StageKind.MAP,
entry_point=MAP_ENTRY_POINT,
needs=(),
inputs={"input": self.input_port},
outputs={"partial": self.partial_port},
parameter_names=_DESCRIPTOR_PARAMETERS,
resources=resources,
execution=execution,
retry=RetryPolicy(),
verifier=ComponentRef("exact-artifact", 1),
trust_modes=trust_modes,
max_fan_out=limits.max_tasks,
cacheable=True,
)
reduce_input = PortSpec(
schema=self.partial_port.schema,
cardinality=Cardinality.MANY,
collection=CollectionKind.KEYED,
)
reduce_stage = StageSpec(
stage_id="reduce",
kind=StageKind.REDUCE,
entry_point=REDUCE_ENTRY_POINT,
needs=("map",),
inputs={"partials": reduce_input},
outputs={"result": self.output_port},
parameter_names=_DESCRIPTOR_PARAMETERS,
resources=resources,
execution=execution,
retry=RetryPolicy(),
verifier=ComponentRef("exact-artifact", 1),
trust_modes=trust_modes,
max_fan_out=1,
cacheable=True,
)
workflow = WorkflowSpec(
workflow_id="descriptor-map-reduce-v1",
inputs={"input": self.input_port},
stages=(map_stage, reduce_stage),
edges=(
ArtifactEdge(PortRef("input"), PortRef("input", "map")),
ArtifactEdge(PortRef("partial", "map"), PortRef("partials", "reduce")),
),
outputs={"result": PortRef("result", "reduce")},
max_tasks=limits.max_tasks,
max_output_bytes=limits.max_output_bytes,
)
self.manifest = WorkloadManifest(
sdk_api=VersionRange(">=1.0,<2.0"),
protocol=VersionRange(">=1,<2"),
workload=WorkloadId("descriptor-batch", "1.0.0"),
description=(
"Compute a pinned set of RDKit 2D descriptors, one canonical "
"CSV row per input molecule, in deterministic input order."
),
package=PackageSpec("scimesh", package_digest),
environment=EnvironmentSpec(
"python-process",
environment_digest,
{"adapter": "sdk-native"},
),
parameters_schema=_parameters_schema(),
workflow=workflow,
inputs={"input": self.input_port},
outputs={"result": self.output_port},
determinism=DeterminismProfile.BYTE_EXACT,
trust_modes=(TrustMode.TRUSTED, TrustMode.UNTRUSTED_QUORUM),
verifier=VerifierSpec(ComponentRef("exact-artifact", 1), {}),
limits=limits,
capabilities=("descriptor-batch",),
conformance_profiles=("core-batch-v1",),
)
self._exact_verifier = ExactArtifactVerifier()
def definition(self) -> WorkloadDefinition:
return WorkloadDefinition(
manifest=self.manifest,
planner=self,
runners={MAP_ENTRY_POINT: self},
reducers={REDUCE_ENTRY_POINT: self},
verifiers={self._exact_verifier.identity.canonical: self._exact_verifier},
)
@staticmethod
def _skip_invalid(parameters: Mapping[str, Any]) -> bool:
value = parameters.get("skip_invalid", True)
if not isinstance(value, bool):
raise ValueError("skip_invalid must be a boolean")
return value
def validate(self, request: JobRequest) -> ValidatedJob:
if request.workload != self.manifest.workload:
raise ValueError("descriptor-batch received a request for another workload")
self._skip_invalid(request.parameters)
return ValidatedJob(request, request.parameters)
def plan(self, job: ValidatedJob, context: PlanningContext) -> WorkflowPlan:
if not isinstance(job, ValidatedJob):
raise ValueError("job must be a ValidatedJob")
collection = job.request.inputs.get("input")
if collection is None:
raise ValueError("descriptor-batch requires the input port")
self.input_port.validate_collection(collection, "job input")
input_artifact = collection.items[0].artifact
input_path = context.catalog.materialize(input_artifact)
workspace = context.workspace
workspace.mkdir(parents=True, exist_ok=True)
shard_paths = write_descriptor_shards(
input_path,
workspace,
self.shard_rows,
)
negotiated = context.negotiated
map_stage = self.manifest.workflow.stages[0]
assert map_stage.verifier is not None
tasks: list[TaskSpec] = []
for index, path in enumerate(shard_paths):
sealed = context.sink.seal(
path,
declaration=self.input_port.schema,
)
tasks.append(
TaskSpec(
workload=self.manifest.workload,
package_digest=self.manifest.package.digest,
manifest_digest=self.manifest.digest,
trust_mode=job.request.trust_mode,
sdk_api_version=negotiated.sdk_api_version,
protocol_version=negotiated.protocol_version,
manifest_schema_version=self.manifest.manifest_schema_version,
workflow_schema_version=self.manifest.workflow.schema_version,
environment_digest=self.manifest.environment.digest,
verifier=map_stage.verifier,
selected_features=negotiated.selected_features,
optional_fallbacks=negotiated.optional_fallbacks,
task_key=f"map/{index:08d}",
stage_id="map",
parameters=job.resolved_parameters,
inputs={"input": ArtifactCollection.single(sealed)},
expected_outputs={"partial": self.partial_port},
resources=map_stage.resources,
execution=map_stage.execution,
)
)
return WorkflowPlan(
workload=self.manifest.workload,
package_digest=self.manifest.package.digest,
manifest_digest=self.manifest.digest,
trust_mode=job.request.trust_mode,
sdk_api_version=negotiated.sdk_api_version,
protocol_version=negotiated.protocol_version,
manifest_schema_version=self.manifest.manifest_schema_version,
workflow_schema_version=self.manifest.workflow.schema_version,
environment_digest=self.manifest.environment.digest,
verifier=self.manifest.verifier.verifier,
selected_features=negotiated.selected_features,
optional_fallbacks=negotiated.optional_fallbacks,
workflow_id=self.manifest.workflow.workflow_id,
resolved_parameters=job.resolved_parameters,
tasks=tuple(tasks),
)
def run(self, context: TaskContext) -> OutputManifest:
context.cancellation.raise_if_cancelled()
collection = context.task.inputs.get("input")
if collection is None:
raise ValueError("descriptor map task requires one input collection")
self.input_port.validate_collection(collection, "descriptor map input")
source = context.catalog.materialize(collection.items[0].artifact)
workspace = context.workspace
workspace.mkdir(parents=True, exist_ok=True)
input_path = workspace / "input"
output_path = workspace / "result.csv"
if source.resolve() != input_path.resolve():
shutil.copyfile(source, input_path)
metrics = compute_descriptor_batch(
input_path,
output_path,
skip_invalid=self._skip_invalid(context.task.parameters),
)
context.cancellation.raise_if_cancelled()
sealed = context.sink.seal(
output_path,
declaration=self.partial_port.schema,
)
return OutputManifest(
context.task.task_key,
{"partial": ArtifactCollection.single(sealed)},
metrics,
context.provenance,
).validate_against(
context.task.expected_outputs,
max_output_bytes=self.manifest.limits.max_output_bytes,
)
def reduce(self, context: ReduceContext) -> OutputManifest:
context.cancellation.raise_if_cancelled()
collection = context.accepted_inputs.get("partials")
if (
collection is None
or collection.kind is not CollectionKind.KEYED
or not collection.items
):
raise ValueError(
"descriptor reducer requires a non-empty keyed partial collection"
)
self.manifest.workflow.stages[1].inputs["partials"].validate_collection(
collection,
"descriptor reducer partials",
)
workspace = context.workspace
workspace.mkdir(parents=True, exist_ok=True)
indexed_items: list[tuple[int, ArtifactItem]] = []
for item in collection.items:
key = item.key or ""
prefix = "map."
raw_index = key[len(prefix) :] if key.startswith(prefix) else ""
if len(raw_index) != 8 or not raw_index.isdigit():
raise ValueError(
"descriptor partial key must use map.<eight-digit-index>"
)
indexed_items.append((int(raw_index), item))
expected_keys = context.task.expected_input_keys.get("partials")
if expected_keys is None or {item.key for item in collection.items} != set(
expected_keys
):
raise ValueError(
"descriptor partial keys do not match the coordinator expected set"
)
if sorted(index for index, _ in indexed_items) != list(
range(len(indexed_items))
):
raise ValueError("descriptor partial keys must be complete and contiguous")
partial_paths: list[Path] = []
for index, item in sorted(indexed_items):
artifact: ArtifactRef = item.artifact
source = context.catalog.materialize(artifact)
target = workspace / artifact.artifact_id
if source.resolve() != target.resolve():
shutil.copyfile(source, target)
if _sha256_file(target) != artifact.sha256:
raise ValueError("materialized partial checksum does not match")
partial_paths.append(target)
result_path = workspace / "result.csv"
metrics = concatenate_descriptor_shards(partial_paths, result_path)
context.cancellation.raise_if_cancelled()
sealed = context.sink.seal(
result_path,
declaration=self.output_port.schema,
)
return OutputManifest(
context.task.task_key,
{"result": ArtifactCollection.single(sealed)},
metrics,
context.provenance,
).validate_against(
context.task.expected_outputs,
max_output_bytes=self.manifest.limits.max_output_bytes,
)
def descriptor_batch_sdk_definition(
*,
shard_rows: int = 10_000,
package_digest: str | None = None,
environment_digest: str | None = None,
) -> DescriptorBatchWorkload:
"""Build the default local descriptor-batch definition for tests."""
return DescriptorBatchWorkload(
shard_rows=shard_rows,
package_digest=package_digest or current_scimesh_package_digest(),
environment_digest=environment_digest or current_environment_digest(),
)
def workload_definition() -> WorkloadDefinition:
"""Installed entry-point factory for the default descriptor-batch definition."""
return descriptor_batch_sdk_definition().definition()