Add SDK-declared workload UI elements and reduction metadata

This commit is contained in:
Emil
2026-08-02 17:47:43 +03:00
parent a18b8b8ae4
commit 5d738e0a14
10 changed files with 397 additions and 18 deletions
+1
View File
@@ -59,6 +59,7 @@ from .manifest import (
WorkloadLimits,
WorkloadManifest,
)
from .ui import UIElement
from .plans import ExpansionManifest, JobRequest, TaskSpec, ValidatedJob, WorkflowPlan
from .protocols import (
ArtifactCatalog,
+12 -2
View File
@@ -19,8 +19,9 @@ from __future__ import annotations
import hashlib
import shutil
from collections.abc import Mapping, Sequence
from pathlib import Path
from typing import Any, Mapping, Sequence
from typing import Any
from .artifacts import (
ArtifactCollection,
@@ -50,6 +51,7 @@ from .plans import JobRequest, TaskSpec, ValidatedJob, WorkflowPlan
from .protocols import PlanningContext, ReduceContext, TaskContext
from .registry import WorkloadDefinition
from .resources import ResourceRequirements
from .ui import UIElement
from .verification import ExactArtifactVerifier
from .workflow import ArtifactEdge, PortRef, StageKind, StageSpec, WorkflowSpec
@@ -133,7 +135,9 @@ class MapReduceWorkload:
port), ``map_parameter_names``, ``reduce_parameter_names``, ``capabilities``,
``trust_modes``, ``workflow_id``, ``limits``, ``resources``, ``execution``,
``map_entry_point``, ``reduce_entry_point``, ``shard_rows`` (default 1000,
used only by the default ``partition_input``).
used only by the default ``partition_input``), ``ui_elements`` (tuple of
``UIElement`` declarations that shape the operator "new job" form in the
coordinator UI; each ``field`` must name a ``parameters_schema`` property).
"""
workload_id: WorkloadId
@@ -155,6 +159,9 @@ class MapReduceWorkload:
map_entry_point: str | None = None
reduce_entry_point: str | None = None
shard_rows: int = 1_000
ui_elements: tuple[UIElement, ...] = ()
reduction: str = "ordered-concat"
upload_ready: bool = True
def __init__(
self,
@@ -291,6 +298,9 @@ class MapReduceWorkload:
limits=limits,
capabilities=self.capabilities,
conformance_profiles=("core-batch-v1",),
ui_elements=tuple(self.ui_elements),
reduction=self.reduction,
upload_ready=self.upload_ready,
)
self._exact_verifier = _EXACT_VERIFIER
self._limits = limits
+47 -9
View File
@@ -4,10 +4,11 @@ from __future__ import annotations
import json
import re
from collections.abc import Mapping
from dataclasses import dataclass
from enum import Enum
from types import MappingProxyType
from typing import Any, Mapping
from typing import Any
from ._validation import (
canonical_json,
@@ -16,8 +17,8 @@ from ._validation import (
require_exact_keys,
require_identifier,
require_positive_int,
require_sha256,
require_schema_version,
require_sha256,
require_string,
thaw_json,
)
@@ -29,8 +30,9 @@ from .identity import (
VersionRange,
WorkloadId,
)
from .workflow import StageKind, WorkflowSpec
from .schema import validate_schema_definition
from .ui import UIElement, ui_elements_from_list
from .workflow import StageKind, WorkflowSpec
class DeterminismProfile(str, Enum):
@@ -100,7 +102,7 @@ class PackageSpec:
}
@classmethod
def from_dict(cls, value: object) -> "PackageSpec":
def from_dict(cls, value: object) -> PackageSpec:
if not isinstance(value, Mapping):
raise ValueError("package specification must be an object")
require_exact_keys(
@@ -146,7 +148,7 @@ class EnvironmentSpec:
}
@classmethod
def from_dict(cls, value: object) -> "EnvironmentSpec":
def from_dict(cls, value: object) -> EnvironmentSpec:
if not isinstance(value, Mapping):
raise ValueError("environment specification must be an object")
require_exact_keys(
@@ -186,7 +188,7 @@ class VerifierSpec:
}
@classmethod
def from_dict(cls, value: object) -> "VerifierSpec":
def from_dict(cls, value: object) -> VerifierSpec:
if not isinstance(value, Mapping):
raise ValueError("verifier specification must be an object")
require_exact_keys(
@@ -236,7 +238,7 @@ class WorkloadLimits:
}
@classmethod
def from_dict(cls, value: object) -> "WorkloadLimits":
def from_dict(cls, value: object) -> WorkloadLimits:
if not isinstance(value, Mapping):
raise ValueError("workload limits must be an object")
fields = {
@@ -265,6 +267,15 @@ def _ports(
return MappingProxyType(result)
_REDUCTION_MODES = ("top-k", "ordered-concat")
def _reduction(value: object) -> str:
if value not in _REDUCTION_MODES:
raise ValueError(f"reduction must be one of: {', '.join(_REDUCTION_MODES)}")
return value # type: ignore[return-value]
@dataclass(frozen=True, slots=True)
class WorkloadManifest:
"""The installed workload's complete, immutable declaration.
@@ -294,6 +305,9 @@ class WorkloadManifest:
conformance_profiles: tuple[str, ...]
required_features: tuple[FeatureRequirement, ...] = ()
optional_features: tuple[FeatureRequirement, ...] = ()
ui_elements: tuple[UIElement, ...] = ()
reduction: str = "ordered-concat"
upload_ready: bool = True
manifest_schema_version: int = MANIFEST_SCHEMA_VERSION
def __post_init__(self) -> None:
@@ -336,6 +350,17 @@ class WorkloadManifest:
raise ValueError("parameters_schema exceeds 1 MiB")
validate_schema_definition(schema)
object.__setattr__(self, "parameters_schema", schema)
properties = dict(schema["properties"])
ui_elements = ui_elements_from_list(self.ui_elements, "ui_elements")
for element in ui_elements:
if element.field not in properties:
raise ValueError(
f"ui element {element.field!r} is not a declared parameter"
)
object.__setattr__(self, "ui_elements", ui_elements)
object.__setattr__(self, "reduction", _reduction(self.reduction))
if not isinstance(self.upload_ready, bool):
raise ValueError("upload_ready must be a boolean")
if not isinstance(self.workflow, WorkflowSpec):
raise ValueError("workflow must be a WorkflowSpec")
object.__setattr__(
@@ -485,13 +510,16 @@ class WorkloadManifest:
"conformance_profiles": list(self.conformance_profiles),
"required_features": [item.to_dict() for item in self.required_features],
"optional_features": [item.to_dict() for item in self.optional_features],
"ui_elements": [element.to_dict() for element in self.ui_elements],
"reduction": self.reduction,
"upload_ready": self.upload_ready,
}
def to_json(self) -> str:
return canonical_json(self.to_dict())
@classmethod
def from_dict(cls, value: object) -> "WorkloadManifest":
def from_dict(cls, value: object) -> WorkloadManifest:
if not isinstance(value, Mapping):
raise ValueError("workload manifest must be an object")
fields = {
@@ -514,6 +542,9 @@ class WorkloadManifest:
"conformance_profiles",
"required_features",
"optional_features",
"ui_elements",
"reduction",
"upload_ready",
}
require_exact_keys(value, fields, "workload manifest")
inputs, outputs = value["inputs"], value["outputs"]
@@ -523,6 +554,7 @@ class WorkloadManifest:
value["conformance_profiles"],
value["required_features"],
value["optional_features"],
value["ui_elements"],
)
if not isinstance(inputs, Mapping) or not isinstance(outputs, Mapping):
raise ValueError("manifest inputs and outputs must be objects")
@@ -556,10 +588,16 @@ class WorkloadManifest:
FeatureRequirement.from_dict(item)
for item in value["optional_features"] # type: ignore[union-attr]
),
ui_elements=ui_elements_from_list(
tuple(value["ui_elements"]),
"ui_elements", # type: ignore[arg-type]
),
reduction=value["reduction"], # type: ignore[arg-type]
upload_ready=value["upload_ready"], # type: ignore[arg-type]
)
@classmethod
def from_json(cls, value: str) -> "WorkloadManifest":
def from_json(cls, value: str) -> WorkloadManifest:
try:
decoded = json.loads(value)
except (TypeError, json.JSONDecodeError, RecursionError) as error:
+137
View File
@@ -0,0 +1,137 @@
"""Workload-declared UI elements for the operator "new job" form.
Workloads may declare how their parameters should be rendered in the
coordinator UI. These declarations are presentation metadata only: the
strict parameter schema remains the authoritative validation contract, and
the UI falls back to schema-derived controls when a workload declares no
elements.
"""
from __future__ import annotations
from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from typing import Any
from ._validation import (
freeze_json,
require_exact_keys,
require_identifier,
require_string,
)
_WIDGETS = {"text", "textarea", "number", "select", "checkbox"}
def _optional_string(value: object, field: str) -> str:
if value == "":
return ""
return require_string(value, field)
@dataclass(frozen=True, slots=True)
class UIElement:
"""One form control bound to a workload parameter.
``field`` must name a property of the workload's ``parameters_schema``;
``widget`` is one of ``text``, ``textarea``, ``number``, ``select``, or
``checkbox``; ``options`` are required for ``select``. ``default`` must be
JSON-safe (``None``, boolean, number, or string).
"""
field: str
widget: str
label: str
help: str = ""
placeholder: str = ""
options: tuple[str, ...] = ()
default: Any = None
order: int = 0
group: str = ""
def __post_init__(self) -> None:
object.__setattr__(self, "field", require_identifier(self.field, "ui.field"))
object.__setattr__(self, "widget", require_identifier(self.widget, "ui.widget"))
if self.widget not in _WIDGETS:
raise ValueError(f"ui.widget must be one of: {', '.join(sorted(_WIDGETS))}")
object.__setattr__(self, "label", require_string(self.label, "ui.label"))
object.__setattr__(self, "help", _optional_string(self.help, "ui.help"))
object.__setattr__(
self,
"placeholder",
_optional_string(self.placeholder, "ui.placeholder"),
)
options = tuple(require_string(value, "ui.option") for value in self.options)
if len(options) != len(set(options)):
raise ValueError("ui.options must be unique")
if self.widget == "select" and not options:
raise ValueError("select ui elements require options")
object.__setattr__(self, "options", options)
freeze_json(self.default, "ui.default")
object.__setattr__(self, "order", self.order)
if isinstance(self.order, bool) or not isinstance(self.order, int):
raise ValueError("ui.order must be an integer")
object.__setattr__(self, "group", _optional_string(self.group, "ui.group"))
def to_dict(self) -> dict[str, Any]:
return {
"field": self.field,
"widget": self.widget,
"label": self.label,
"help": self.help,
"placeholder": self.placeholder,
"options": list(self.options),
"default": self.default,
"order": self.order,
"group": self.group,
}
@classmethod
def from_dict(cls, value: object) -> UIElement:
if not isinstance(value, Mapping):
raise ValueError("ui element must be an object")
require_exact_keys(
value,
{
"field",
"widget",
"label",
"help",
"placeholder",
"options",
"default",
"order",
"group",
},
"ui element",
)
options = value["options"]
if not isinstance(options, list):
raise ValueError("ui.options must be an array")
return cls(
field=value["field"], # type: ignore[arg-type]
widget=value["widget"], # type: ignore[arg-type]
label=value["label"], # type: ignore[arg-type]
help=value["help"], # type: ignore[arg-type]
placeholder=value["placeholder"], # type: ignore[arg-type]
options=tuple(options),
default=value["default"],
order=value["order"], # type: ignore[arg-type]
group=value["group"], # type: ignore[arg-type]
)
def ui_elements_from_list(value: Sequence[object], field: str) -> tuple[UIElement, ...]:
"""Validate and freeze a manifest ``ui_elements`` declaration."""
elements: list[UIElement] = []
for item in value:
if isinstance(item, UIElement):
elements.append(item)
elif isinstance(item, Mapping):
elements.append(UIElement.from_dict(item))
else:
raise ValueError(f"{field} must contain UIElement values")
names = [element.field for element in elements]
if len(names) != len(set(names)):
raise ValueError(f"{field} fields must be unique")
return tuple(elements)
+13 -1
View File
@@ -9,13 +9,15 @@ partition, compute, and merge.
from __future__ import annotations
from collections.abc import Mapping
from pathlib import Path
from typing import Any, Mapping
from typing import Any
from scimesh.sdk.artifacts import ArtifactSchema, ComponentRef, PortSpec
from scimesh.sdk.batch import MapReduceWorkload
from scimesh.sdk.identity import SchemaRef, WorkloadId
from scimesh.sdk.registry import WorkloadDefinition
from scimesh.sdk.ui import UIElement
from ..environment import current_environment_digest, current_scimesh_package_digest
from .core import (
@@ -88,6 +90,16 @@ class DescriptorBatchWorkload(MapReduceWorkload):
reduce_parameter_names = ("skip_invalid",)
map_entry_point = MAP_ENTRY_POINT
reduce_entry_point = REDUCE_ENTRY_POINT
ui_elements = (
UIElement(
"skip_invalid",
"checkbox",
"Skip invalid molecules",
help="Skip rows with invalid SMILES instead of failing the shard.",
default=True,
order=1,
),
)
def __init__(
self,
+30 -1
View File
@@ -9,8 +9,9 @@ byte-identical to the local brute-force reference.
from __future__ import annotations
from collections.abc import Mapping, Sequence
from pathlib import Path
from typing import Any, Mapping, Sequence
from typing import Any
from scimesh.sdk.artifacts import (
ArtifactCollection,
@@ -23,6 +24,7 @@ from scimesh.sdk.identity import SchemaRef, WorkloadId
from scimesh.sdk.plans import TaskSpec, ValidatedJob
from scimesh.sdk.protocols import PlanningContext
from scimesh.sdk.registry import WorkloadDefinition
from scimesh.sdk.ui import UIElement
from scimesh.sdk.workflow import StageSpec
from ..environment import current_environment_digest, current_scimesh_package_digest
@@ -109,8 +111,35 @@ class SimilarityGraphSDKWorkload(MapReduceWorkload):
"max_rows",
)
workflow_id = "graph-block-pairs-v1"
upload_ready = False
map_entry_point = MAP_ENTRY_POINT
reduce_entry_point = REDUCE_ENTRY_POINT
ui_elements = (
UIElement(
"threshold",
"number",
"Similarity threshold",
help="Minimum (greater) or maximum (less) edge similarity. Required.",
order=1,
),
UIElement(
"threshold_direction",
"select",
"Direction",
help="Whether to keep edges above (greater) or below (less) the threshold.",
options=("greater", "less"),
default="greater",
order=2,
),
UIElement(
"block_size",
"number",
"Block size",
help="Deterministic block size for pair sharding.",
default=100,
order=3,
),
)
def domain_validate(self, parameters: Mapping[str, Any]) -> None:
unknown = set(parameters) - {
+29 -1
View File
@@ -8,13 +8,15 @@ header-preserving concatenation, so nothing else is needed.
from __future__ import annotations
from collections.abc import Mapping
from pathlib import Path
from typing import Any, Mapping
from typing import Any
from scimesh.sdk.artifacts import ArtifactSchema, ComponentRef, PortSpec
from scimesh.sdk.batch import MapReduceWorkload
from scimesh.sdk.identity import SchemaRef, WorkloadId
from scimesh.sdk.registry import WorkloadDefinition
from scimesh.sdk.ui import UIElement
from ..environment import current_environment_digest, current_scimesh_package_digest
from .core import MOLWT_COLUMNS, filter_molecules_by_molwt
@@ -93,6 +95,32 @@ class MolwtFilterWorkload(MapReduceWorkload):
reduce_parameter_names = ("min_molwt", "max_molwt", "skip_invalid")
map_entry_point = MAP_ENTRY_POINT
reduce_entry_point = REDUCE_ENTRY_POINT
ui_elements = (
UIElement(
"min_molwt",
"number",
"Minimum molecular weight",
help="Keep molecules with MolWt at least this value. Optional.",
placeholder="e.g. 100",
order=1,
),
UIElement(
"max_molwt",
"number",
"Maximum molecular weight",
help="Keep molecules with MolWt at most this value. Optional.",
placeholder="e.g. 600",
order=2,
),
UIElement(
"skip_invalid",
"checkbox",
"Skip invalid molecules",
help="Skip rows with invalid SMILES instead of failing the shard.",
default=True,
order=3,
),
)
def __init__(
self,
+45 -1
View File
@@ -9,8 +9,9 @@ and merge.
from __future__ import annotations
from collections.abc import Mapping, Sequence
from pathlib import Path
from typing import Any, Mapping, Sequence
from typing import Any
from rdkit import Chem
@@ -21,6 +22,7 @@ from scimesh.sdk.batch import MapReduceWorkload
from scimesh.sdk.identity import SchemaRef, WorkloadId
from scimesh.sdk.plans import JobRequest, ValidatedJob
from scimesh.sdk.registry import WorkloadDefinition
from scimesh.sdk.ui import UIElement
from ..environment import current_environment_digest, current_scimesh_package_digest
from .core import merge_search_partials, run_search_shard, write_search_shards
@@ -112,6 +114,48 @@ class SimilaritySearchSDKWorkload(MapReduceWorkload):
reduce_parameter_names = _MAP_PARAMETERS + ("query_source", "fingerprint")
map_entry_point = MAP_ENTRY_POINT
reduce_entry_point = REDUCE_ENTRY_POINT
reduction = "top-k"
ui_elements = (
UIElement(
"query_id",
"text",
"Query molecule id",
help="ChEMBL id of the query molecule. Provide exactly one of id or SMILES.",
order=1,
),
UIElement(
"query_smiles",
"text",
"Query molecule SMILES",
help="SMILES of the query molecule. Provide exactly one of id or SMILES.",
order=2,
),
UIElement(
"top_k",
"number",
"Top k",
help="Number of most similar molecules to keep per shard (global merge keeps the best of these).",
default=20,
order=3,
),
UIElement(
"threshold_direction",
"select",
"Direction",
help="Keep molecules with similarity greater or less than the threshold.",
options=("greater", "less"),
default="greater",
order=4,
),
UIElement(
"threshold",
"number",
"Similarity threshold",
help="Optional similarity bound: results are filtered to this direction.",
placeholder="e.g. 0.8",
order=5,
),
)
def __init__(
self,
+44 -1
View File
@@ -50,6 +50,12 @@ class WorkloadCLI:
)
export_parser.set_defaults(workload_handler=self.export_workloads)
allowlist_parser = subparsers.add_parser(
"allowlist",
help="Print the installed-package workload allowlist for worker configuration.",
)
allowlist_parser.set_defaults(workload_handler=self.export_allowlist)
run_parser = subparsers.add_parser(
"run", help="Run one SDK workload locally against an input file."
)
@@ -147,6 +153,11 @@ class WorkloadCLI:
"verifier": manifest.verifier.verifier.canonical,
"enabled": item.enabled,
"parameters_schema": thaw_json(manifest.parameters_schema),
"ui_elements": [
element.to_dict() for element in manifest.ui_elements
],
"reduction": manifest.reduction,
"upload_ready": manifest.upload_ready,
"inputs": {
name: port.schema.to_dict()
for name, port in manifest.inputs.items()
@@ -158,7 +169,7 @@ class WorkloadCLI:
}
)
payload: dict[str, object] = {
"schema_version": 1,
"schema_version": 2,
"generated_by": "scimesh workload export",
"workloads": workloads,
}
@@ -169,6 +180,38 @@ class WorkloadCLI:
print(f"Exported {len(workloads)} workloads to {args.output}")
return 0
def export_allowlist(self, args: argparse.Namespace) -> int:
"""Print the allowlist JSON that worker environments consume.
The printed array feeds ``SCIMESH_WORKLOAD_ALLOWLIST`` on workers and
mirrors the digest pins of the installed distribution.
"""
import json
registry = self._registry(args)
payload = []
for item in sorted(
registry.descriptions(), key=lambda value: value.workload.name
):
if not item.enabled:
continue
definition, _ = registry.require(
item.workload.name,
item.workload.version,
item.package_digest,
)
manifest = definition.manifest
payload.append(
{
"distribution": manifest.package.distribution,
"name": manifest.workload.name,
"version": manifest.workload.version,
"digest": manifest.package.digest,
}
)
print(json.dumps(payload, indent=2, sort_keys=True))
return 0
def run_workload(self, args: argparse.Namespace) -> int:
registry = self._registry(args)
descriptions = registry.descriptions()
+39 -2
View File
@@ -241,7 +241,7 @@ def test_workload_cli_exports_the_library_as_json(tmp_path: Path) -> None:
output = tmp_path / "workloads.json"
assert main(["workload", "export", "-o", str(output)]) == 0
payload = json.loads(output.read_text(encoding="utf-8"))
assert payload["schema_version"] == 1
assert payload["schema_version"] == 2
names = [item["name"] for item in payload["workloads"]]
assert names == sorted(
["descriptor-batch", "molwt-filter", "similarity-graph", "similarity-search"]
@@ -253,7 +253,9 @@ def test_workload_cli_exports_the_library_as_json(tmp_path: Path) -> None:
assert item["verifier"] == "exact-artifact@1"
assert "input" in item["inputs"]
assert "result" in item["outputs"]
molwt = next(item for item in payload["workloads"] if item["name"] == "molwt-filter")
molwt = next(
item for item in payload["workloads"] if item["name"] == "molwt-filter"
)
assert molwt["parameters_schema"]["properties"] == {
"min_molwt": {
"type": "number",
@@ -271,3 +273,38 @@ def test_workload_cli_exports_the_library_as_json(tmp_path: Path) -> None:
"description": "Skip rows with invalid SMILES instead of failing",
},
}
assert molwt["ui_elements"] == [
{
"field": "min_molwt",
"widget": "number",
"label": "Minimum molecular weight",
"help": "Keep molecules with MolWt at least this value. Optional.",
"placeholder": "e.g. 100",
"options": [],
"default": None,
"order": 1,
"group": "",
},
{
"field": "max_molwt",
"widget": "number",
"label": "Maximum molecular weight",
"help": "Keep molecules with MolWt at most this value. Optional.",
"placeholder": "e.g. 600",
"options": [],
"default": None,
"order": 2,
"group": "",
},
{
"field": "skip_invalid",
"widget": "checkbox",
"label": "Skip invalid molecules",
"help": "Skip rows with invalid SMILES instead of failing the shard.",
"placeholder": "",
"options": [],
"default": True,
"order": 3,
"group": "",
},
]