Add workload SDK foundation
This commit is contained in:
@@ -958,6 +958,42 @@ workload logic into the service.
|
|||||||
- the existing single-operator demo remains usable through a documented local
|
- the existing single-operator demo remains usable through a documented local
|
||||||
development configuration.
|
development configuration.
|
||||||
|
|
||||||
|
### CTX-16 — Workload SDK foundation
|
||||||
|
|
||||||
|
**Goal:** Provide a strict Python authoring SDK for installed, allowlisted
|
||||||
|
scientific workloads while retaining the CTX-07 distributed protocol as a
|
||||||
|
compatible wire profile.
|
||||||
|
|
||||||
|
**Depends on:** CTX-07 and CTX-08. Coordinator-backed generalized scheduling
|
||||||
|
also depends on CTX-10 through CTX-14, but the Python contract and local
|
||||||
|
conformance runtime can land independently and must fail closed for unavailable
|
||||||
|
features.
|
||||||
|
|
||||||
|
**Acceptance criteria:**
|
||||||
|
|
||||||
|
- public manifest, workflow, task, artifact, resource, execution, provenance,
|
||||||
|
and verifier value objects are immutable, typed, JSON-safe, versioned, and
|
||||||
|
strict about unknown fields;
|
||||||
|
- installed workload discovery requires an administrator allowlist plus exact
|
||||||
|
workload version and package digest; job parameters cannot select code;
|
||||||
|
- compatibility negotiation covers SDK/protocol/profile/feature/environment
|
||||||
|
versions and occurs before planner invocation;
|
||||||
|
- plans/tasks pin package and manifest digests plus selected trust mode, and
|
||||||
|
quorum candidates carry coordinator-owned candidate/owner and scientific
|
||||||
|
binding identities;
|
||||||
|
- `core-batch-v1` has a trusted local conformance executor with atomic resource
|
||||||
|
reservation, sealed-output/provenance validation, declared verifier
|
||||||
|
invocation, and golden scientific parity;
|
||||||
|
- exact, canonical-record, and structured numeric-tolerance verifier
|
||||||
|
primitives return bounded sanitized decisions;
|
||||||
|
- the existing distributed `similarity-search` is available through an adapter
|
||||||
|
without changing its wire schema, worker alias boundary, or scientific
|
||||||
|
result, and parity is tested;
|
||||||
|
- advanced dynamic, stream, accelerator, gang, and side-effect profiles are
|
||||||
|
rejected unless an enforcing runtime advertises their required features;
|
||||||
|
- an author guide documents package entry points, security boundaries,
|
||||||
|
conformance tests, and current coordinator/Worker limitations.
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
## 10. Suggested assignment bundles
|
## 10. Suggested assignment bundles
|
||||||
@@ -973,6 +1009,7 @@ parallel unless one engineer owns integration.
|
|||||||
| Distributed computation | CTX-07, CTX-08, CTX-10 | Scientific Python engineer |
|
| Distributed computation | CTX-07, CTX-08, CTX-10 | Scientific Python engineer |
|
||||||
| Product surface | CTX-09, CTX-11 | Full-stack/backend engineer |
|
| Product surface | CTX-09, CTX-11 | Full-stack/backend engineer |
|
||||||
| Quality gate | CTX-12 | DevOps/QA engineer |
|
| Quality gate | CTX-12 | DevOps/QA engineer |
|
||||||
|
| Workload SDK | CTX-16 | Scientific Python/platform engineer |
|
||||||
|
|
||||||
Suggested order for a small team:
|
Suggested order for a small team:
|
||||||
|
|
||||||
|
|||||||
@@ -133,6 +133,27 @@ pytest
|
|||||||
|
|
||||||
The package separates common dataset parsing and fingerprints from independent workloads. Add future workloads through the workload registry without changing the main CLI.
|
The package separates common dataset parsing and fingerprints from independent workloads. Add future workloads through the workload registry without changing the main CLI.
|
||||||
|
|
||||||
|
## Workload SDK
|
||||||
|
|
||||||
|
`scimesh.sdk` implements the `core-batch-v1` authoring profile: strict and
|
||||||
|
immutable workload manifests, typed artifact ports, static map/reduce plans,
|
||||||
|
resource eligibility and local reservations, exact/canonical/numeric verifier
|
||||||
|
primitives, installed-package allowlisting, and a compatibility adapter for the
|
||||||
|
existing distributed `similarity-search`. See the
|
||||||
|
[SDK author guide](docs/workload-sdk.md), [contract](docs/scimesh-sdk-contract.md),
|
||||||
|
and [delivery roadmap](docs/scimesh-sdk-roadmap.md).
|
||||||
|
|
||||||
|
Dynamic workflows, real Worker concurrency, coordinator-backed GPU allocation,
|
||||||
|
streaming, and gang execution remain fail-closed until their versioned runtime
|
||||||
|
features are implemented; declaring those profiles does not silently enable
|
||||||
|
them.
|
||||||
|
|
||||||
|
The included `LocalCoreBatchExecutor` is a trusted, single-threaded in-process
|
||||||
|
conformance harness. It validates scientific parity, sealed outputs, provenance,
|
||||||
|
and limits, but intentionally refuses profiles that claim network/process
|
||||||
|
isolation, secrets, accelerators, gangs, checkpoints, or retries; those require
|
||||||
|
the future enforcing Agent runtime.
|
||||||
|
|
||||||
## Team
|
## Team
|
||||||
|
|
||||||
- [Emil](https://github.com/emil28092005) — Project Lead
|
- [Emil](https://github.com/emil28092005) — Project Lead
|
||||||
|
|||||||
@@ -1,7 +1,7 @@
|
|||||||
# SciMesh Status
|
# SciMesh Status
|
||||||
|
|
||||||
**Updated:** 2026-08-01
|
**Updated:** 2026-08-01
|
||||||
**Branch baseline:** `main` at `b9a975b` (self-service worker enrollment)
|
**Branch baseline:** `main`; this revision adds the Workload SDK foundation.
|
||||||
|
|
||||||
## Current state
|
## Current state
|
||||||
|
|
||||||
@@ -49,6 +49,7 @@ the complete result-artifact SHA-256 before a task is accepted.
|
|||||||
| CTX-11 Dashboard/operator view | Implemented | Protected live control room: recent-run/worker overview, real pipeline-stage visualization, shard attempts and safe failures, validated similarity-search upload, coordinator artifacts, final-result download, and bounded polling. |
|
| CTX-11 Dashboard/operator view | Implemented | Protected live control room: recent-run/worker overview, real pipeline-stage visualization, shard attempts and safe failures, validated similarity-search upload, coordinator artifacts, final-result download, and bounded polling. |
|
||||||
| CTX-12 Reliability, security, CI | In progress | Unit, race, PostgreSQL integration, and smoke checks exist; CI hardening remains. |
|
| CTX-12 Reliability, security, CI | In progress | Unit, race, PostgreSQL integration, and smoke checks exist; CI hardening remains. |
|
||||||
| CTX-15 User Service and access control | Implemented | User/owner scoping, verified contributors, worker keys, self-service enrollment, and quorum-backed untrusted workers are merged; local Go/Python and Docker/PostgreSQL checks passed. |
|
| CTX-15 User Service and access control | Implemented | User/owner scoping, verified contributors, worker keys, self-service enrollment, and quorum-backed untrusted workers are merged; local Go/Python and Docker/PostgreSQL checks passed. |
|
||||||
|
| CTX-16 Workload SDK foundation | Implemented | `scimesh.sdk` provides strict immutable manifests/plans/artifacts, digest/trust-pinned tasks, typed DAGs, compatibility negotiation, verifier primitives with owner/binding-safe quorum inputs, resource eligibility/local allocation, measured package discovery, a trusted local core-batch conformance harness, and a tested legacy similarity-search adapter. Enforcing coordinator/Worker profiles remain fail-closed. |
|
||||||
|
|
||||||
## Next recommended assignment
|
## Next recommended assignment
|
||||||
|
|
||||||
@@ -65,6 +66,14 @@ block-pair planning and reduction for `similarity-graph`.
|
|||||||
- The coordinator accepts uploaded distributed jobs only for
|
- The coordinator accepts uploaded distributed jobs only for
|
||||||
`similarity-search` with `query_smiles`. It rejects `similarity-graph` until
|
`similarity-search` with `query_smiles`. It rejects `similarity-graph` until
|
||||||
CTX-10 supplies cross-shard pair planning.
|
CTX-10 supplies cross-shard pair planning.
|
||||||
|
- The SDK can execute `core-batch-v1` locally, but the protocol-v1 coordinator
|
||||||
|
still has flat single-input/single-result tasks and no package/resource
|
||||||
|
leases. General DAG, concurrent-Agent, GPU, stream, and gang execution needs
|
||||||
|
a versioned coordinator/Worker rollout; unsupported features fail before
|
||||||
|
planner invocation.
|
||||||
|
- The local SDK executor is intentionally trusted and in-process. It does not
|
||||||
|
enforce process/network/timeout/credential isolation and rejects declarations
|
||||||
|
that would require those guarantees.
|
||||||
|
|
||||||
## Update rule
|
## Update rule
|
||||||
|
|
||||||
|
|||||||
@@ -1,13 +1,16 @@
|
|||||||
# SciMesh Workload SDK contract
|
# SciMesh Workload SDK contract
|
||||||
|
|
||||||
**Status:** design draft `0.1`; not implemented. Normative words **MUST**,
|
**Status:** contract `0.1`. The Python `core-batch-v1` foundation is implemented
|
||||||
**MUST NOT**, **SHOULD**, and **MAY** describe the intended future contract,
|
in `scimesh.sdk`; dynamic, streaming, accelerator, gang, side-effect, and
|
||||||
not capabilities of the current release.
|
coordinator protocol-v2 behavior remains a normative target. Normative words
|
||||||
|
**MUST**, **MUST NOT**, **SHOULD**, and **MAY** apply to an implementation only
|
||||||
|
when it advertises the affected profile or feature.
|
||||||
|
|
||||||
This document defines the compatibility boundary for approved SciMesh workload
|
This document defines the compatibility boundary for approved SciMesh workload
|
||||||
packages. The sequencing and unresolved product decisions remain in
|
packages. The sequencing and unresolved product decisions remain in
|
||||||
[`scimesh-sdk-roadmap.md`](scimesh-sdk-roadmap.md). The currently implemented
|
[`scimesh-sdk-roadmap.md`](scimesh-sdk-roadmap.md). The production coordinator
|
||||||
protocol is still [`ctx-07-distributed-workload-protocol.md`](ctx-07-distributed-workload-protocol.md).
|
wire compatibility profile remains
|
||||||
|
[`ctx-07-distributed-workload-protocol.md`](ctx-07-distributed-workload-protocol.md).
|
||||||
|
|
||||||
## 1. Scope and invariants
|
## 1. Scope and invariants
|
||||||
|
|
||||||
@@ -240,7 +243,11 @@ indefinitely while holding resources.
|
|||||||
`TaskSpec` is the concrete unit leased to a Worker Agent:
|
`TaskSpec` is the concrete unit leased to a Worker Agent:
|
||||||
|
|
||||||
```yaml
|
```yaml
|
||||||
task_schema_version: 1
|
schema_version: 1
|
||||||
|
workload: descriptor-batch@1.2.0
|
||||||
|
package_digest: sha256:...
|
||||||
|
manifest_digest: ...
|
||||||
|
trust_mode: trusted
|
||||||
task_key: calculate/shard-000042
|
task_key: calculate/shard-000042
|
||||||
stage_id: calculate
|
stage_id: calculate
|
||||||
parameters: {...validated JSON...}
|
parameters: {...validated JSON...}
|
||||||
@@ -279,9 +286,10 @@ represent a collection. A later protocol may persist collection edges directly.
|
|||||||
|
|
||||||
Before completion, a runner uploads an `OutputManifest` listing every declared
|
Before completion, a runner uploads an `OutputManifest` listing every declared
|
||||||
output artifact, checksum, schema, size, record/dimension summary, metrics, and
|
output artifact, checksum, schema, size, record/dimension summary, metrics, and
|
||||||
provenance. Provenance includes resolved versions, package/environment digest,
|
provenance. Provenance includes resolved versions, package/environment and
|
||||||
Worker runtime, allocated resource IDs, parameters digest, input collection
|
manifest digests, Worker runtime, allocated resource IDs, parameters digest,
|
||||||
digest, timestamps, random seed where applicable, and checkpoint lineage.
|
input collection digest, timestamps, random seed where applicable, and
|
||||||
|
checkpoint lineage.
|
||||||
|
|
||||||
Unexpected ports, missing required outputs, extra artifacts, schema failures,
|
Unexpected ports, missing required outputs, extra artifacts, schema failures,
|
||||||
or limit violations reject the Attempt. Logs and checkpoints are separate
|
or limit violations reject the Attempt. Logs and checkpoints are separate
|
||||||
@@ -419,6 +427,13 @@ declare an appropriate verifier/trust combination. Reducers consume only
|
|||||||
accepted partial outputs and MUST detect missing, duplicate, conflicting, or
|
accepted partial outputs and MUST detect missing, duplicate, conflicting, or
|
||||||
inconclusive inputs.
|
inconclusive inputs.
|
||||||
|
|
||||||
|
Quorum candidates MUST be coordinator-authenticated envelopes with unique
|
||||||
|
Attempt/candidate identity and an owner identity. A verifier counts at most one
|
||||||
|
vote per owner. It also receives a coordinator-owned binding for workload,
|
||||||
|
task, package/manifest/environment, parameters, and input-collection digests;
|
||||||
|
outputs from another job or code pin are invalid even when their result bytes
|
||||||
|
match.
|
||||||
|
|
||||||
## 8. Failure, retry, cancellation, and checkpoint semantics
|
## 8. Failure, retry, cancellation, and checkpoint semantics
|
||||||
|
|
||||||
Every failure has a stable sanitized code, category (`input`, `scientific`,
|
Every failure has a stable sanitized code, category (`input`, `scientific`,
|
||||||
@@ -472,26 +487,34 @@ isolation is ineligible rather than silently unsandboxed.
|
|||||||
|
|
||||||
## 10. SDK interfaces and conformance
|
## 10. SDK interfaces and conformance
|
||||||
|
|
||||||
The future Python API SHOULD expose protocols equivalent to:
|
The Python API exposes protocols equivalent to:
|
||||||
|
|
||||||
```python
|
```python
|
||||||
class Planner(Protocol):
|
class Planner(Protocol):
|
||||||
def validate(self, request: JobRequest) -> ValidatedJob: ...
|
def validate(self, request: JobRequest) -> ValidatedJob: ...
|
||||||
def plan(self, job: ValidatedJob, artifacts: ArtifactCatalog) -> WorkflowPlan: ...
|
def plan(self, job: ValidatedJob, context: PlanningContext) -> WorkflowPlan: ...
|
||||||
|
|
||||||
class Runner(Protocol):
|
class Runner(Protocol):
|
||||||
def run(self, context: TaskContext) -> OutputManifest: ...
|
def run(self, context: TaskContext) -> OutputManifest: ...
|
||||||
|
|
||||||
class Reducer(Protocol):
|
class Reducer(Protocol):
|
||||||
def reduce(self, context: ReduceContext, inputs: AcceptedOutputs) -> OutputManifest: ...
|
def reduce(self, context: ReduceContext) -> OutputManifest: ...
|
||||||
|
|
||||||
class Verifier(Protocol):
|
class Verifier(Protocol):
|
||||||
def verify(self, context: VerifyContext, candidates: CandidateOutputs) -> VerificationDecision: ...
|
def verify(self, context: VerifyContext, candidates: CandidateOutputs) -> VerificationDecision: ...
|
||||||
```
|
```
|
||||||
|
|
||||||
Concrete public value objects are immutable, typed, JSON-safe, schema-versioned,
|
Concrete public value objects are immutable, typed, JSON-safe, strict about
|
||||||
and reject unknown fields. Scientific cores SHOULD remain callable without a
|
unknown fields, and canonically serialized inside versioned wire contracts.
|
||||||
coordinator so the same implementation powers local and distributed adapters.
|
Scientific cores SHOULD remain callable without a coordinator so the same
|
||||||
|
implementation powers local and distributed adapters.
|
||||||
|
|
||||||
|
The shipped `LocalCoreBatchExecutor` is a trusted in-process conformance
|
||||||
|
harness, not the `core-batch-v1` production isolation boundary. It rejects
|
||||||
|
restricted-network, parallel-process/thread, accelerator, secret, checkpoint,
|
||||||
|
retry, gang, and advanced-stage declarations. Subprocess isolation, hard
|
||||||
|
timeouts, leases, and credential enforcement remain requirements for an Agent
|
||||||
|
runtime that advertises those guarantees.
|
||||||
|
|
||||||
An SDK conformance suite MUST test manifest/schema validation, deterministic
|
An SDK conformance suite MUST test manifest/schema validation, deterministic
|
||||||
planning, no local-path/URI leakage, output bounds, local/distributed parity,
|
planning, no local-path/URI leakage, output bounds, local/distributed parity,
|
||||||
|
|||||||
@@ -1,7 +1,11 @@
|
|||||||
# SciMesh Workload SDK roadmap
|
# SciMesh Workload SDK roadmap
|
||||||
|
|
||||||
**Status:** future design and sequencing document. No SDK package, commands, or
|
**Status:** active delivery roadmap. The Python `scimesh.sdk` package now
|
||||||
general verifier abstraction described here is implemented yet.
|
implements the `core-batch-v1` foundation, verifier primitives, resource
|
||||||
|
eligibility/local allocation, installed-package registry, local conformance
|
||||||
|
runtime, and legacy similarity-search adapter. Coordinator-backed generalized
|
||||||
|
DAG execution, Worker concurrency, accelerators, streaming, gang execution,
|
||||||
|
and authoring CLI commands remain future phases.
|
||||||
|
|
||||||
The normative future API, workflow, execution, resource, security, and failure
|
The normative future API, workflow, execution, resource, security, and failure
|
||||||
semantics are specified in the design-draft
|
semantics are specified in the design-draft
|
||||||
@@ -69,8 +73,10 @@ canonical, numeric, domain-specific, or trust-policy comparison. It processes
|
|||||||
structured manifests and bounded streams where practical, records sanitized
|
structured manifests and bounded streams where practical, records sanitized
|
||||||
evidence, and rejects inconsistent results.
|
evidence, and rejects inconsistent results.
|
||||||
|
|
||||||
Current untrusted quorum is only `ExactArtifactVerifier`: distinct owners must
|
Current untrusted quorum is only `ExactArtifactVerifier`: coordinator-created
|
||||||
produce whole files with identical SHA-256. Future modes are
|
candidate envelopes from distinct owners must share the exact workload, task,
|
||||||
|
package/manifest/environment, parameters, and input binding and produce whole
|
||||||
|
files with identical SHA-256. Future modes are
|
||||||
`CanonicalRecordVerifier`, `NumericToleranceVerifier`,
|
`CanonicalRecordVerifier`, `NumericToleranceVerifier`,
|
||||||
`DomainSpecificVerifier`, and `TrustedWorkerPolicy`. Canonical mode requires a
|
`DomainSpecificVerifier`, and `TrustedWorkerPolicy`. Canonical mode requires a
|
||||||
specified parser/schema/order/encoding/serialization; numeric mode compares
|
specified parser/schema/order/encoding/serialization; numeric mode compares
|
||||||
|
|||||||
@@ -0,0 +1,115 @@
|
|||||||
|
# Workload SDK handoff
|
||||||
|
|
||||||
|
**Audience:** the engineer/AI continuing SciMesh Workload SDK implementation.
|
||||||
|
**Date:** 2026-08-01. **Baseline:** uncommitted working tree on `main` (`11e9333`
|
||||||
|
plus the CTX-16 SDK changes); `python -m pytest -q` reports **225 passed**.
|
||||||
|
|
||||||
|
Read first, in this order: `AGENTS.md` (binding repo rules),
|
||||||
|
`docs/scimesh-sdk-roadmap.md` (delivery order — it governs, this file does not),
|
||||||
|
`docs/scimesh-sdk-contract.md` (normative target semantics),
|
||||||
|
`docs/workload-sdk.md` (author guide for what exists), and the CTX-16 entry in
|
||||||
|
`PLAN.md`.
|
||||||
|
|
||||||
|
## What is already done (do not redo)
|
||||||
|
|
||||||
|
CTX-16 "Workload SDK foundation" is complete and tested. `scimesh/sdk/`
|
||||||
|
implements the `core-batch-v1` profile:
|
||||||
|
|
||||||
|
- Immutable, JSON-strict value objects: `identity.py`, `artifacts.py`,
|
||||||
|
`workflow.py`, `manifest.py`, `plans.py`, `execution.py`, `resources.py`.
|
||||||
|
- Fail-closed compatibility negotiation: `runtime.py` (`negotiate_manifest`)
|
||||||
|
plus request-level checks in `registry.py`.
|
||||||
|
- Installed-package registry with administrator allowlist, exact version +
|
||||||
|
`sha256:` digest pinning, entry-point discovery with digest measured before
|
||||||
|
and after import: `registry.py`, `integrity.py`.
|
||||||
|
- Verifier primitives `ExactArtifactVerifier`, `CanonicalRecordVerifier`,
|
||||||
|
`NumericToleranceVerifier` with bounded sanitized evidence: `verification.py`.
|
||||||
|
- Local conformance harness: `LocalArtifactStore`, `LocalCoreBatchExecutor`,
|
||||||
|
`ResourcePool` (atomic all-or-nothing reservation): `conformance.py`.
|
||||||
|
- Legacy adapter exposing distributed `similarity-search` through the SDK
|
||||||
|
without changing its wire schema: `compat/distributed_v1.py`, `builtins.py`;
|
||||||
|
entry point `similarity-search@1.0.0` is declared in `pyproject.toml`.
|
||||||
|
- Tests: `tests/test_sdk_{models,resources,verification,compatibility,registry}.py`
|
||||||
|
including fail-closed rejection coverage for every advanced profile
|
||||||
|
declaration (gang, GPU modes, pools, checkpoints, retries, secrets, streams,
|
||||||
|
loops, side effects).
|
||||||
|
|
||||||
|
## What remains, in delivery order
|
||||||
|
|
||||||
|
1. **`descriptor-batch` reference workload** (roadmap step 3 — the recommended
|
||||||
|
next task; it is pure Python and needs no coordinator changes). Pinned RDKit
|
||||||
|
2D descriptors, canonical one-row-per-input CSV, shard-index concatenation
|
||||||
|
with one header, byte-identical local/distributed output, two-worker quorum.
|
||||||
|
Build it as an SDK-native package (manifest + planner/runner/reducer/
|
||||||
|
verifier handlers), not through the legacy adapter; reuse the
|
||||||
|
`similarity-search` adapter (`scimesh/sdk/compat/distributed_v1.py`) and
|
||||||
|
`builtins.py` as the structural template, and the
|
||||||
|
`tests/test_sdk_compatibility.py` fixtures as the test template. This is the
|
||||||
|
intended first `untrusted_quorum` candidate (byte_exact + exact-artifact@1).
|
||||||
|
2. **Distributed `similarity-graph`** (CTX-10, roadmap step 1). The coordinator
|
||||||
|
currently rejects `similarity-graph` uploads; it needs cross-shard block-pair
|
||||||
|
planning and duplicate-safe reduction. STATUS.md names this the next
|
||||||
|
recommended assignment overall.
|
||||||
|
3. **Coordinator/Worker protocol v2** (needs CTX-10, then CTX-13 in-worker CPU
|
||||||
|
parallelism and CTX-14 GPU execution; Go + Python). The protocol-v1
|
||||||
|
coordinator persists only flat one-input/one-result tasks: no resource
|
||||||
|
requirements, stage edges, package versions, device allocations, or gang
|
||||||
|
leases. Until a versioned rollout lands, SDK declarations for those features
|
||||||
|
must stay fail-closed — do not silently "enable" them.
|
||||||
|
4. **More chemistry workloads** (roadmap step 4): standardization, SMARTS
|
||||||
|
screening, fingerprint export, fixed-template reaction enumeration, then
|
||||||
|
reaction validation/descriptors.
|
||||||
|
5. **Composite artifacts and richer verifier policies** (roadmap step 5):
|
||||||
|
first-class ordered/keyed `ArtifactCollection` edges instead of composite
|
||||||
|
manifest artifacts; decide where verifiers execute (open decision in the
|
||||||
|
roadmap).
|
||||||
|
6. **Authoring CLI** (future tooling, does not exist today): `scimesh workload
|
||||||
|
init`, `validate`, `test-local`, `test-distributed`, `golden`, `package`.
|
||||||
|
Per AGENTS.md, keep CLI parsing in workload modules and register through
|
||||||
|
`scimesh/core/registry.py`; no workload-specific logic in the main CLI.
|
||||||
|
7. **Open decisions** (listed at the end of the roadmap): SDK distribution
|
||||||
|
split, Go↔Python planner bridge, verifier execution/attestation, trust-mode
|
||||||
|
governance, multi-user enablement. Do not pick one unilaterally — surface it.
|
||||||
|
|
||||||
|
## Known traps (cost the previous session real time)
|
||||||
|
|
||||||
|
- The legacy adapter pins its own manifest (`adapter.manifest`). If a test
|
||||||
|
changes limits/workflow on the manifest, the adapter's copy must be replaced
|
||||||
|
too, or `registry.plan` fails with "planner plan does not carry the selected
|
||||||
|
immutable workload pin".
|
||||||
|
- `WorkloadDefinition` validation: a PLAN stage's `entry_point` must equal
|
||||||
|
`planner.entry_point`; every non-REDUCE stage's `entry_point` must be a key
|
||||||
|
in `runners` (REDUCE → `reducers`); verifier handlers are keyed by
|
||||||
|
`ComponentRef.canonical` and must expose a matching `.identity`.
|
||||||
|
- Negotiation requires each triggering property's feature to be declared
|
||||||
|
separately: e.g. `PROCESS_POOL` needs `process-pools` **and** `multi-process`
|
||||||
|
for `max_processes > 1`. Runtime must also advertise every declared required
|
||||||
|
feature, or negotiation fails with `feature-unavailable`.
|
||||||
|
- `feature-fallback-disallowed` in `scimesh/sdk/registry.py` is currently
|
||||||
|
unreachable via `registry.plan` (the `feature-unavailable` check fires first
|
||||||
|
for any runtime that produced a fallback). Behavior is still fail-closed;
|
||||||
|
decide whether to reorder or delete the branch.
|
||||||
|
- The local executor is deliberately trusted/in-process: it rejects anything
|
||||||
|
but `TrustMode.TRUSTED`, `NetworkPolicy.TRUSTED`, single-threaded CPU
|
||||||
|
map/reduce without retries/gangs/accelerators/secrets/checkpoints. That is a
|
||||||
|
contract, not a bug — test rejections, don't "fix" them.
|
||||||
|
- `JobRequest` parameters and failure/evidence payloads reject local paths and
|
||||||
|
URIs by design; keep new payloads location-free.
|
||||||
|
- There is a stray nested clone `SciMesh/` in the repo root (same repo at an
|
||||||
|
older commit). Ignore it and never `git add` it; consider deleting it.
|
||||||
|
- The full ChEMBL extract `chembl_37_chemreps.txt` (~2.9M rows) makes the
|
||||||
|
single-threaded local executor run for many minutes; tests must use small
|
||||||
|
TSV fixtures (see `_write_tiny_dataset`).
|
||||||
|
|
||||||
|
## Working agreement
|
||||||
|
|
||||||
|
- Verify with `source .venv/bin/activate && python -m pytest -q`; the baseline
|
||||||
|
is 225 passing tests and it must stay green. Add a regression test for every
|
||||||
|
behavioral change; similarity code needs a brute-force/sorted reference and
|
||||||
|
determinism across block sizes.
|
||||||
|
- Legacy `similarity-search` wire schema, worker alias boundary, and scientific
|
||||||
|
output must not change. Worker code never talks to PostgreSQL directly;
|
||||||
|
results go through the coordinator; failures go to `/failure`, never as
|
||||||
|
`file://`/`worker://` result URIs.
|
||||||
|
- Do not commit datasets, generated CSV/PNG, tokens, or local worker artifacts.
|
||||||
|
- One CTX task per pull request; link the CTX item from `PLAN.md`.
|
||||||
@@ -0,0 +1,220 @@
|
|||||||
|
# SciMesh Workload SDK v1
|
||||||
|
|
||||||
|
SciMesh now ships a public Python SDK under `scimesh.sdk`. The implemented
|
||||||
|
authoring profile is **`core-batch-v1`**: installed and digest-pinned workload
|
||||||
|
definitions, strict JSON manifests, typed artifact ports and collections, a
|
||||||
|
static map/reduce workflow, CPU/memory/scratch eligibility, atomic local
|
||||||
|
resource reservation, exact/canonical/numeric verifier primitives, and a
|
||||||
|
compatibility adapter for the existing `DistributedWorkload` protocol. Its
|
||||||
|
local executor is deliberately a trusted, in-process conformance harness; the
|
||||||
|
production subprocess/lease sandbox remains a coordinator/Worker milestone.
|
||||||
|
|
||||||
|
The full target contract remains in
|
||||||
|
[`scimesh-sdk-contract.md`](scimesh-sdk-contract.md). Dynamic expansion,
|
||||||
|
streaming, accelerators, gang execution, and side effects have typed bounded
|
||||||
|
declarations, but the current coordinator/Worker runtime does not advertise
|
||||||
|
their features. Compatibility negotiation therefore rejects those workflows
|
||||||
|
before planner code runs.
|
||||||
|
|
||||||
|
## What authors import
|
||||||
|
|
||||||
|
The stable authoring surface is exported from `scimesh.sdk`:
|
||||||
|
|
||||||
|
- `WorkloadManifest`, `WorkloadId`, `VersionRange`, `PackageSpec`, and
|
||||||
|
`EnvironmentSpec` pin identity and compatibility;
|
||||||
|
- `ArtifactSchema`, `PortSpec`, `ArtifactRef`, and `ArtifactCollection` define
|
||||||
|
immutable data boundaries without transport URLs or local paths;
|
||||||
|
- `WorkflowSpec`, `StageSpec`, `ArtifactEdge`, `TaskSpec`, and `WorkflowPlan`
|
||||||
|
define a typed acyclic plan and pin package/manifest digests plus trust mode;
|
||||||
|
- `ResourceRequirements` and `ExecutionProfile` separate per-task resources
|
||||||
|
from Agent `max_concurrency`;
|
||||||
|
- `Planner`, `Runner`, `Reducer`, and `Verifier` are the package handler
|
||||||
|
protocols;
|
||||||
|
- `OutputManifest` and `Provenance` describe sealed durable results;
|
||||||
|
- `WorkloadRegistry` resolves an exact name, version, package digest, runtime,
|
||||||
|
environment, and feature set. It never selects an implicit latest version.
|
||||||
|
|
||||||
|
Persisted manifests, requests, plans, tasks, expansions, outputs, candidates,
|
||||||
|
decisions, and failures are frozen, recursively immutable, JSON-safe,
|
||||||
|
canonically serialized, and strict about unknown fields; their enclosing wire
|
||||||
|
contracts carry schema versions.
|
||||||
|
Artifact identities contain a coordinator-owned UUID, schema, checksum, media
|
||||||
|
type, and bounds; a scientific handler never persists a filesystem path.
|
||||||
|
|
||||||
|
## Try the built-in SDK workload
|
||||||
|
|
||||||
|
This example executes the current distributed `similarity-search` through the
|
||||||
|
SDK without starting PostgreSQL or the coordinator:
|
||||||
|
|
||||||
|
```python
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
from scimesh.sdk import (
|
||||||
|
ArtifactCollection,
|
||||||
|
JobRequest,
|
||||||
|
LocalArtifactStore,
|
||||||
|
LocalCoreBatchExecutor,
|
||||||
|
default_sdk_registry,
|
||||||
|
default_sdk_runtime,
|
||||||
|
similarity_search_sdk_adapter,
|
||||||
|
)
|
||||||
|
|
||||||
|
root = Path("sdk-run")
|
||||||
|
store = LocalArtifactStore(root / "artifacts")
|
||||||
|
adapter = similarity_search_sdk_adapter(shard_rows=1_000)
|
||||||
|
|
||||||
|
dataset = store.import_file(
|
||||||
|
Path("chembl_37_chemreps.txt"),
|
||||||
|
declaration=adapter.input_port.schema,
|
||||||
|
)
|
||||||
|
request = JobRequest(
|
||||||
|
workload=adapter.manifest.workload,
|
||||||
|
parameters={"query_smiles": "CCO", "top_k": 20},
|
||||||
|
inputs={"input": ArtifactCollection.single(dataset)},
|
||||||
|
)
|
||||||
|
|
||||||
|
result = LocalCoreBatchExecutor(
|
||||||
|
default_sdk_registry(shard_rows=1_000),
|
||||||
|
default_sdk_runtime(),
|
||||||
|
store,
|
||||||
|
root / "attempts",
|
||||||
|
).execute(request, adapter.manifest.package.digest)
|
||||||
|
|
||||||
|
result_ref = result.outputs["result"].items[0].artifact
|
||||||
|
print(store.materialize(result_ref))
|
||||||
|
```
|
||||||
|
|
||||||
|
`LocalCoreBatchExecutor` is a correctness/conformance runtime, not a substitute
|
||||||
|
for coordinator leases or multi-machine scheduling. It accepts only
|
||||||
|
`TrustMode.TRUSTED`, `NetworkPolicy.TRUSTED`, single-process/single-threaded CPU
|
||||||
|
map/reduce stages without secrets, checkpoints, retries, gangs, or
|
||||||
|
accelerators. It does not claim network, timeout, process, or credential
|
||||||
|
isolation. Unsupported declarations are rejected before a handler runs. The
|
||||||
|
harness uses the same legacy scientific planner, shard runner, and reducer as
|
||||||
|
the distributed `similarity-search`, and its parity is covered by automated
|
||||||
|
tests.
|
||||||
|
|
||||||
|
## Package shape and registration
|
||||||
|
|
||||||
|
An SDK distribution provides one explicit entry point per workload version:
|
||||||
|
|
||||||
|
```toml
|
||||||
|
[project.entry-points."scimesh.workloads"]
|
||||||
|
"descriptor-batch@1.0.0" = "scimesh_descriptors.sdk:workload_definition"
|
||||||
|
```
|
||||||
|
|
||||||
|
The factory returns a `WorkloadDefinition` containing its manifest and handler
|
||||||
|
objects. An administrator supplies an `AllowedPackage` with the same
|
||||||
|
distribution, exact `WorkloadId`, and `sha256:` package digest. Discovery
|
||||||
|
filters installed metadata before importing an entry point and fails
|
||||||
|
transactionally if an allowlisted definition is missing or mismatched. Job
|
||||||
|
parameters cannot name a module, entry point, package path, or executable.
|
||||||
|
The measured digest covers package payload files and installed entry-point
|
||||||
|
declarations and is checked before and after loading. It is a content pin, not
|
||||||
|
a signature or image attestation; production discovery should run in a fresh
|
||||||
|
trusted control-plane process so a pre-populated Python module cache is not an
|
||||||
|
integrity boundary.
|
||||||
|
|
||||||
|
Direct registration is useful for tests and embedded deployments:
|
||||||
|
|
||||||
|
```python
|
||||||
|
registry = WorkloadRegistry()
|
||||||
|
registry.register(definition, enabled=False)
|
||||||
|
registry.enable(
|
||||||
|
definition.manifest.workload.name,
|
||||||
|
definition.manifest.workload.version,
|
||||||
|
definition.manifest.package.digest,
|
||||||
|
)
|
||||||
|
```
|
||||||
|
|
||||||
|
Both version and digest are required when resolving or planning. Upgrading an
|
||||||
|
installed definition does not change the identity of an existing Job.
|
||||||
|
|
||||||
|
## Authoring rules
|
||||||
|
|
||||||
|
1. Keep the scientific core callable without a coordinator.
|
||||||
|
2. Inline a strict JSON parameter schema with `type: object` and
|
||||||
|
`additionalProperties: false`; the planner still performs domain validation.
|
||||||
|
3. Give every external and stage port an `ArtifactSchema` with a media type,
|
||||||
|
schema version, and byte/record/dimension bounds.
|
||||||
|
4. Connect stage ports with `ArtifactEdge` values. `WorkflowSpec` checks source
|
||||||
|
and target schemas, complete input bindings, declared dependencies, and
|
||||||
|
acyclicity.
|
||||||
|
5. Declare one `ResourceRequirements` and `ExecutionProfile` per stage. A task
|
||||||
|
cannot run until its entire request is eligible and atomically reserved.
|
||||||
|
6. Return only sink-sealed artifacts in `OutputManifest`; the local harness
|
||||||
|
binds task key/provenance itself and rejects fabricated references,
|
||||||
|
unexpected/missing ports, wrong schema/media type, and cumulative output or
|
||||||
|
artifact-limit violations.
|
||||||
|
7. Select a verifier compatible with determinism and trust. SDK v1 permits
|
||||||
|
`untrusted_quorum` only for `byte_exact` plus `exact-artifact@1`.
|
||||||
|
8. Add golden fixtures, local/distributed parity, retry/completion-order, and
|
||||||
|
verifier failure tests before enabling a package.
|
||||||
|
|
||||||
|
`ArtifactSink` and `ArtifactCatalog` are bridge-owned protocols. They let
|
||||||
|
scientific handlers materialize verified inputs and seal outputs without bearer
|
||||||
|
tokens, database credentials, upload URLs, or durable local paths.
|
||||||
|
|
||||||
|
## Verification
|
||||||
|
|
||||||
|
The SDK includes:
|
||||||
|
|
||||||
|
- `ExactArtifactVerifier`: compares logical port/collection/schema/content
|
||||||
|
digests while ignoring coordinator UUIDs, timestamps, metrics, and worker
|
||||||
|
identity. Quorum inputs use coordinator-created `CandidateOutput` envelopes,
|
||||||
|
count at most one vote per owner, and require a `VerificationBinding` for the
|
||||||
|
exact task, inputs, parameters, package, manifest, and environment;
|
||||||
|
- `CanonicalRecordVerifier`: applies a package-owned bounded canonicalizer and
|
||||||
|
compares length-framed canonical records;
|
||||||
|
- `NumericToleranceVerifier`: recursively checks structure plus explicit
|
||||||
|
absolute, relative, ULP, and NaN policy, returning bounded sanitized evidence.
|
||||||
|
|
||||||
|
Canonical and numeric objects expose direct bounded comparison methods. To use
|
||||||
|
them as manifest `Verifier` handlers, the package supplies an artifact-to-record
|
||||||
|
or artifact-to-structured-value loader; without one, verification returns
|
||||||
|
`inconclusive` rather than accepting bytes it did not parse.
|
||||||
|
|
||||||
|
A decision is `accepted`, `rejected`, or `inconclusive`; only `accepted`
|
||||||
|
satisfies a stage. Evidence is limited to 16 KiB and cannot contain local paths
|
||||||
|
or transport URLs.
|
||||||
|
|
||||||
|
## Resources and current runtime boundary
|
||||||
|
|
||||||
|
`ResourcePool` provides a lock-protected all-or-nothing local reservation for
|
||||||
|
CPU cores, memory, scratch, and accelerator device/partition IDs, including
|
||||||
|
whole-device versus partition conflict fencing. It enforces aggregate capacity
|
||||||
|
and execution-slot count. `ExecutionProfile` produces only
|
||||||
|
allocation-derived OpenMP/BLAS and device-visibility values; credentials never
|
||||||
|
belong to scientific parameters.
|
||||||
|
|
||||||
|
The current protocol-v1 coordinator stores one input/result per flat task and
|
||||||
|
does not persist resource requirements, device allocations, stage edges, or
|
||||||
|
package versions. The production Worker also remains serial. Consequently:
|
||||||
|
|
||||||
|
- SDK `core-batch-v1` can be authored, validated, tested, discovered, and run
|
||||||
|
through the trusted local conformance harness now;
|
||||||
|
- existing production `similarity-search` remains on its compatible v1 wire
|
||||||
|
path and is not renamed;
|
||||||
|
- real concurrent claims, GPU scheduling, multi-output DAG execution, dynamic
|
||||||
|
loops, streaming, and gang leases require the versioned coordinator/Worker
|
||||||
|
changes listed in [`scimesh-sdk-roadmap.md`](scimesh-sdk-roadmap.md);
|
||||||
|
- merely declaring a GPU or gang request never enables it. Missing runtime
|
||||||
|
features or inventory fail before the planner executes.
|
||||||
|
|
||||||
|
## Conformance commands
|
||||||
|
|
||||||
|
Install development tools and run the SDK suite:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
pip install -e '.[dev]'
|
||||||
|
pytest tests/test_sdk_models.py \
|
||||||
|
tests/test_sdk_resources.py \
|
||||||
|
tests/test_sdk_verification.py \
|
||||||
|
tests/test_sdk_compatibility.py \
|
||||||
|
tests/test_sdk_registry.py
|
||||||
|
```
|
||||||
|
|
||||||
|
Run `pytest` for the full legacy, Worker, local-science, and SDK regression
|
||||||
|
suite. Package authors can reuse `LocalArtifactStore`,
|
||||||
|
`LocalCoreBatchExecutor`, and `assert_manifest_round_trip` in their own golden
|
||||||
|
tests.
|
||||||
@@ -17,8 +17,12 @@ dev = ["pytest>=8"]
|
|||||||
scimesh = "scimesh.cli:main"
|
scimesh = "scimesh.cli:main"
|
||||||
scimesh-worker = "scimesh.worker.cli:main"
|
scimesh-worker = "scimesh.worker.cli:main"
|
||||||
|
|
||||||
|
[project.entry-points."scimesh.workloads"]
|
||||||
|
"similarity-search@1.0.0" = "scimesh.sdk.builtins:similarity_search_workload_definition"
|
||||||
|
|
||||||
[tool.setuptools.packages.find]
|
[tool.setuptools.packages.find]
|
||||||
include = ["scimesh*"]
|
include = ["scimesh*"]
|
||||||
|
namespaces = false
|
||||||
|
|
||||||
[tool.pytest.ini_options]
|
[tool.pytest.ini_options]
|
||||||
testpaths = ["tests"]
|
testpaths = ["tests"]
|
||||||
|
|||||||
@@ -0,0 +1,225 @@
|
|||||||
|
"""SciMesh Workload SDK v1.
|
||||||
|
|
||||||
|
The implemented profile is ``core-batch-v1``: strict manifests, typed artifact
|
||||||
|
ports, static map/reduce DAGs, resource eligibility/local reservation, exact
|
||||||
|
verification, and an adapter for existing distributed workloads. Advanced
|
||||||
|
dynamic, stream, accelerator, gang, and side-effect declarations are modeled
|
||||||
|
but fail compatibility negotiation unless an enforcing runtime advertises the
|
||||||
|
corresponding versioned features.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from .artifacts import (
|
||||||
|
ArtifactCollection,
|
||||||
|
ArtifactItem,
|
||||||
|
ArtifactRef,
|
||||||
|
ArtifactSchema,
|
||||||
|
Cardinality,
|
||||||
|
CollectionKind,
|
||||||
|
OutputManifest,
|
||||||
|
PortSpec,
|
||||||
|
Provenance,
|
||||||
|
)
|
||||||
|
from .builtins import (
|
||||||
|
current_environment_digest,
|
||||||
|
current_scimesh_package_digest,
|
||||||
|
default_sdk_registry,
|
||||||
|
default_sdk_runtime,
|
||||||
|
similarity_search_sdk_adapter,
|
||||||
|
)
|
||||||
|
from .conformance import (
|
||||||
|
CancellationFlag,
|
||||||
|
LocalArtifactStore,
|
||||||
|
LocalCoreBatchExecutor,
|
||||||
|
LocalPlanningContext,
|
||||||
|
LocalTaskContext,
|
||||||
|
assert_manifest_round_trip,
|
||||||
|
)
|
||||||
|
from .execution import (
|
||||||
|
CheckpointPolicy,
|
||||||
|
ExecutionProfile,
|
||||||
|
FailureCategory,
|
||||||
|
FailureReport,
|
||||||
|
NetworkPolicy,
|
||||||
|
ProcessModel,
|
||||||
|
RetryPolicy,
|
||||||
|
)
|
||||||
|
from .identity import (
|
||||||
|
MANIFEST_SCHEMA_VERSION,
|
||||||
|
OUTPUT_SCHEMA_VERSION,
|
||||||
|
SDK_API_VERSION,
|
||||||
|
TASK_SCHEMA_VERSION,
|
||||||
|
WORKFLOW_SCHEMA_VERSION,
|
||||||
|
ComponentRef,
|
||||||
|
FeatureRequirement,
|
||||||
|
SchemaRef,
|
||||||
|
VersionRange,
|
||||||
|
WorkloadId,
|
||||||
|
)
|
||||||
|
from .integrity import installed_distribution_digest
|
||||||
|
from .manifest import (
|
||||||
|
DeterminismProfile,
|
||||||
|
EnvironmentSpec,
|
||||||
|
PackageSpec,
|
||||||
|
TrustMode,
|
||||||
|
VerifierSpec,
|
||||||
|
WorkloadLimits,
|
||||||
|
WorkloadManifest,
|
||||||
|
)
|
||||||
|
from .plans import ExpansionManifest, JobRequest, TaskSpec, ValidatedJob, WorkflowPlan
|
||||||
|
from .protocols import (
|
||||||
|
ArtifactCatalog,
|
||||||
|
ArtifactSink,
|
||||||
|
CancellationToken,
|
||||||
|
Planner,
|
||||||
|
PlanningContext,
|
||||||
|
PlanningResources,
|
||||||
|
ReduceContext,
|
||||||
|
Reducer,
|
||||||
|
Runner,
|
||||||
|
TaskContext,
|
||||||
|
Verifier,
|
||||||
|
)
|
||||||
|
from .registry import (
|
||||||
|
AllowedPackage,
|
||||||
|
WorkloadDefinition,
|
||||||
|
WorkloadDescription,
|
||||||
|
WorkloadRegistry,
|
||||||
|
)
|
||||||
|
from .resources import (
|
||||||
|
AcceleratorDevice,
|
||||||
|
AcceleratorMode,
|
||||||
|
ResourceAllocation,
|
||||||
|
ResourceInventory,
|
||||||
|
ResourcePool,
|
||||||
|
ResourceRequirements,
|
||||||
|
ResourceUnavailableError,
|
||||||
|
)
|
||||||
|
from .runtime import (
|
||||||
|
CompatibilityError,
|
||||||
|
NegotiatedWorkload,
|
||||||
|
RuntimeCapabilities,
|
||||||
|
negotiate_manifest,
|
||||||
|
)
|
||||||
|
from .verification import (
|
||||||
|
CandidateOutput,
|
||||||
|
CandidateOutputs,
|
||||||
|
CanonicalRecordVerifier,
|
||||||
|
ExactArtifactVerifier,
|
||||||
|
NumericTolerance,
|
||||||
|
NumericToleranceVerifier,
|
||||||
|
VerificationDecision,
|
||||||
|
VerificationBinding,
|
||||||
|
VerificationStatus,
|
||||||
|
VerifyContext,
|
||||||
|
)
|
||||||
|
from .workflow import (
|
||||||
|
ArtifactEdge,
|
||||||
|
GangSpec,
|
||||||
|
LoopSpec,
|
||||||
|
PortRef,
|
||||||
|
SideEffectSpec,
|
||||||
|
StageKind,
|
||||||
|
StageSpec,
|
||||||
|
StreamSpec,
|
||||||
|
WorkflowFailurePolicy,
|
||||||
|
WorkflowSpec,
|
||||||
|
)
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"AcceleratorDevice",
|
||||||
|
"AcceleratorMode",
|
||||||
|
"AllowedPackage",
|
||||||
|
"ArtifactCatalog",
|
||||||
|
"ArtifactCollection",
|
||||||
|
"ArtifactEdge",
|
||||||
|
"ArtifactItem",
|
||||||
|
"ArtifactRef",
|
||||||
|
"ArtifactSchema",
|
||||||
|
"ArtifactSink",
|
||||||
|
"CancellationFlag",
|
||||||
|
"CancellationToken",
|
||||||
|
"CandidateOutput",
|
||||||
|
"CandidateOutputs",
|
||||||
|
"CanonicalRecordVerifier",
|
||||||
|
"Cardinality",
|
||||||
|
"CheckpointPolicy",
|
||||||
|
"CollectionKind",
|
||||||
|
"CompatibilityError",
|
||||||
|
"ComponentRef",
|
||||||
|
"DeterminismProfile",
|
||||||
|
"EnvironmentSpec",
|
||||||
|
"ExactArtifactVerifier",
|
||||||
|
"ExecutionProfile",
|
||||||
|
"ExpansionManifest",
|
||||||
|
"FailureCategory",
|
||||||
|
"FailureReport",
|
||||||
|
"FeatureRequirement",
|
||||||
|
"GangSpec",
|
||||||
|
"JobRequest",
|
||||||
|
"LocalArtifactStore",
|
||||||
|
"LocalCoreBatchExecutor",
|
||||||
|
"LocalPlanningContext",
|
||||||
|
"LocalTaskContext",
|
||||||
|
"LoopSpec",
|
||||||
|
"MANIFEST_SCHEMA_VERSION",
|
||||||
|
"NegotiatedWorkload",
|
||||||
|
"NetworkPolicy",
|
||||||
|
"NumericTolerance",
|
||||||
|
"NumericToleranceVerifier",
|
||||||
|
"OUTPUT_SCHEMA_VERSION",
|
||||||
|
"OutputManifest",
|
||||||
|
"PackageSpec",
|
||||||
|
"Planner",
|
||||||
|
"PlanningContext",
|
||||||
|
"PlanningResources",
|
||||||
|
"PortRef",
|
||||||
|
"PortSpec",
|
||||||
|
"ProcessModel",
|
||||||
|
"Provenance",
|
||||||
|
"ReduceContext",
|
||||||
|
"Reducer",
|
||||||
|
"ResourceAllocation",
|
||||||
|
"ResourceInventory",
|
||||||
|
"ResourcePool",
|
||||||
|
"ResourceRequirements",
|
||||||
|
"ResourceUnavailableError",
|
||||||
|
"RetryPolicy",
|
||||||
|
"Runner",
|
||||||
|
"RuntimeCapabilities",
|
||||||
|
"SDK_API_VERSION",
|
||||||
|
"SchemaRef",
|
||||||
|
"SideEffectSpec",
|
||||||
|
"StageKind",
|
||||||
|
"StageSpec",
|
||||||
|
"StreamSpec",
|
||||||
|
"TASK_SCHEMA_VERSION",
|
||||||
|
"TaskContext",
|
||||||
|
"TaskSpec",
|
||||||
|
"TrustMode",
|
||||||
|
"ValidatedJob",
|
||||||
|
"VerificationDecision",
|
||||||
|
"VerificationBinding",
|
||||||
|
"VerificationStatus",
|
||||||
|
"Verifier",
|
||||||
|
"VerifierSpec",
|
||||||
|
"VersionRange",
|
||||||
|
"VerifyContext",
|
||||||
|
"WORKFLOW_SCHEMA_VERSION",
|
||||||
|
"WorkflowFailurePolicy",
|
||||||
|
"WorkflowPlan",
|
||||||
|
"WorkflowSpec",
|
||||||
|
"WorkloadDefinition",
|
||||||
|
"WorkloadDescription",
|
||||||
|
"WorkloadId",
|
||||||
|
"WorkloadLimits",
|
||||||
|
"WorkloadManifest",
|
||||||
|
"WorkloadRegistry",
|
||||||
|
"assert_manifest_round_trip",
|
||||||
|
"current_environment_digest",
|
||||||
|
"current_scimesh_package_digest",
|
||||||
|
"default_sdk_registry",
|
||||||
|
"default_sdk_runtime",
|
||||||
|
"installed_distribution_digest",
|
||||||
|
"negotiate_manifest",
|
||||||
|
"similarity_search_sdk_adapter",
|
||||||
|
]
|
||||||
@@ -0,0 +1,380 @@
|
|||||||
|
"""Internal validation helpers for strict, JSON-safe SDK value objects."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
import math
|
||||||
|
import re
|
||||||
|
from types import MappingProxyType
|
||||||
|
from typing import Any, Mapping
|
||||||
|
from urllib.parse import unquote
|
||||||
|
from uuid import UUID
|
||||||
|
|
||||||
|
|
||||||
|
WORKLOAD_NAME_PATTERN = re.compile(r"^[a-z][a-z0-9]*(?:-[a-z0-9]+)*$")
|
||||||
|
IDENTIFIER_PATTERN = re.compile(r"^[a-z][a-z0-9]*(?:[-_.][a-z0-9]+)*$")
|
||||||
|
ENTRY_POINT_PATTERN = re.compile(
|
||||||
|
r"^[A-Za-z_][A-Za-z0-9_.]*:[A-Za-z_][A-Za-z0-9_.]*(?:@v[1-9][0-9]*)?$"
|
||||||
|
)
|
||||||
|
SEMVER_PATTERN = re.compile(
|
||||||
|
r"^(0|[1-9][0-9]*)\.(0|[1-9][0-9]*)\.(0|[1-9][0-9]*)"
|
||||||
|
r"(?:-([0-9A-Za-z-]+(?:\.[0-9A-Za-z-]+)*))?"
|
||||||
|
r"(?:\+([0-9A-Za-z-]+(?:\.[0-9A-Za-z-]+)*))?$"
|
||||||
|
)
|
||||||
|
_VERSION_PATTERN = re.compile(r"^(0|[1-9][0-9]*)(?:\.(0|[1-9][0-9]*))?(?:\.(0|[1-9][0-9]*))?$")
|
||||||
|
_VERSION_CLAUSE_PATTERN = re.compile(r"^(==|>=|<=|>|<)\s*(.+)$")
|
||||||
|
_FORBIDDEN_LOCATOR_PREFIXES = (
|
||||||
|
"file://",
|
||||||
|
"worker://",
|
||||||
|
"http://",
|
||||||
|
"https://",
|
||||||
|
"s3://",
|
||||||
|
"/",
|
||||||
|
)
|
||||||
|
_URI_SCHEME_PATTERN = re.compile(r"^[A-Za-z][A-Za-z0-9+.-]*:")
|
||||||
|
_WINDOWS_PATH_PATTERN = re.compile(r"^[A-Za-z]:(?:[\\/]|[^\s]*[\\/])")
|
||||||
|
_SECRET_ASSIGNMENT_PATTERN = re.compile(
|
||||||
|
r"(?i)(?:^|[^A-Za-z0-9_])"
|
||||||
|
r"(?:authorization|bearer|token|secret|password|api[-_]?key)\s*[:=]"
|
||||||
|
)
|
||||||
|
_PATH_ASSIGNMENT_PATTERN = re.compile(
|
||||||
|
r"(?i)(?:^|[^A-Za-z0-9_])"
|
||||||
|
r"(?:path|file|directory|dir|workspace|cwd|upload|download)\s*[:=]"
|
||||||
|
)
|
||||||
|
_PATH_SEGMENT_PATTERN = re.compile(r"^[A-Za-z0-9_.-]+$")
|
||||||
|
_FILE_NAME_PATTERN = re.compile(r"^[A-Za-z0-9_.-]+\.[A-Za-z0-9]{1,16}$")
|
||||||
|
_TASK_KEY_COMPONENT_PATTERN = re.compile(r"^[a-z0-9][a-z0-9_.-]*$")
|
||||||
|
|
||||||
|
|
||||||
|
def require_exact_keys(
|
||||||
|
value: Mapping[str, object],
|
||||||
|
expected: set[str],
|
||||||
|
label: str,
|
||||||
|
*,
|
||||||
|
optional: set[str] | None = None,
|
||||||
|
) -> None:
|
||||||
|
"""Reject unknown fields and report missing required fields."""
|
||||||
|
if any(not isinstance(key, str) for key in value):
|
||||||
|
raise ValueError(f"{label} must use string field names")
|
||||||
|
optional = optional or set()
|
||||||
|
actual = set(value)
|
||||||
|
missing = expected - actual
|
||||||
|
unknown = actual - expected - optional
|
||||||
|
if not missing and not unknown:
|
||||||
|
return
|
||||||
|
details: list[str] = []
|
||||||
|
if missing:
|
||||||
|
details.append("missing " + ", ".join(sorted(missing)))
|
||||||
|
if unknown:
|
||||||
|
details.append("unknown " + ", ".join(sorted(unknown)))
|
||||||
|
raise ValueError(f"{label} has invalid fields: {'; '.join(details)}")
|
||||||
|
|
||||||
|
|
||||||
|
def require_mapping(value: object, field: str) -> Mapping[str, object]:
|
||||||
|
if not isinstance(value, Mapping) or any(not isinstance(key, str) for key in value):
|
||||||
|
raise ValueError(f"{field} must be an object with string keys")
|
||||||
|
return value
|
||||||
|
|
||||||
|
|
||||||
|
def require_string(value: object, field: str, *, max_length: int = 256) -> str:
|
||||||
|
if not isinstance(value, str) or not value.strip() or len(value) > max_length:
|
||||||
|
raise ValueError(f"{field} must be a non-empty string of at most {max_length} characters")
|
||||||
|
if any(ord(character) < 32 for character in value):
|
||||||
|
raise ValueError(f"{field} must not contain control characters")
|
||||||
|
return value
|
||||||
|
|
||||||
|
|
||||||
|
def require_identifier(value: object, field: str) -> str:
|
||||||
|
text = require_string(value, field, max_length=128)
|
||||||
|
if not IDENTIFIER_PATTERN.fullmatch(text):
|
||||||
|
raise ValueError(f"{field} must be a canonical identifier")
|
||||||
|
return text
|
||||||
|
|
||||||
|
|
||||||
|
def require_workload_name(value: object, field: str = "workload.name") -> str:
|
||||||
|
text = require_string(value, field, max_length=128)
|
||||||
|
if not WORKLOAD_NAME_PATTERN.fullmatch(text):
|
||||||
|
raise ValueError(f"{field} must be a canonical hyphenated workload name")
|
||||||
|
return text
|
||||||
|
|
||||||
|
|
||||||
|
def require_entry_point(value: object, field: str) -> str:
|
||||||
|
text = require_string(value, field, max_length=256)
|
||||||
|
if not ENTRY_POINT_PATTERN.fullmatch(text):
|
||||||
|
raise ValueError(f"{field} must be a package-owned module:object entry point")
|
||||||
|
return text
|
||||||
|
|
||||||
|
|
||||||
|
def require_semver(value: object, field: str) -> str:
|
||||||
|
text = require_string(value, field, max_length=64)
|
||||||
|
match = SEMVER_PATTERN.fullmatch(text)
|
||||||
|
if match is None:
|
||||||
|
raise ValueError(f"{field} must be a semantic version such as 1.0.0")
|
||||||
|
prerelease = match.group(4)
|
||||||
|
if prerelease is not None and any(
|
||||||
|
identifier.isdigit() and len(identifier) > 1 and identifier.startswith("0")
|
||||||
|
for identifier in prerelease.split(".")
|
||||||
|
):
|
||||||
|
raise ValueError(f"{field} has a non-canonical numeric prerelease identifier")
|
||||||
|
return text
|
||||||
|
|
||||||
|
|
||||||
|
def require_uuid(value: object, field: str) -> str:
|
||||||
|
if not isinstance(value, str):
|
||||||
|
raise ValueError(f"{field} must be a UUID string")
|
||||||
|
try:
|
||||||
|
return str(UUID(value))
|
||||||
|
except ValueError as error:
|
||||||
|
raise ValueError(f"{field} must be a UUID string") from error
|
||||||
|
|
||||||
|
|
||||||
|
def require_sha256(value: object, field: str, *, prefixed: bool = False) -> str:
|
||||||
|
if not isinstance(value, str):
|
||||||
|
raise ValueError(f"{field} must be a SHA-256 digest")
|
||||||
|
digest = value[7:] if prefixed and value.startswith("sha256:") else value
|
||||||
|
if prefixed and not value.startswith("sha256:"):
|
||||||
|
raise ValueError(f"{field} must use the sha256:<hex> form")
|
||||||
|
if not re.fullmatch(r"[0-9a-f]{64}", digest):
|
||||||
|
raise ValueError(f"{field} must be a lowercase SHA-256 digest")
|
||||||
|
return f"sha256:{digest}" if prefixed else digest
|
||||||
|
|
||||||
|
|
||||||
|
def require_nonnegative_int(value: object, field: str) -> int:
|
||||||
|
if isinstance(value, bool) or not isinstance(value, int) or value < 0:
|
||||||
|
raise ValueError(f"{field} must be a non-negative integer")
|
||||||
|
return value
|
||||||
|
|
||||||
|
|
||||||
|
def require_positive_int(value: object, field: str) -> int:
|
||||||
|
result = require_nonnegative_int(value, field)
|
||||||
|
if result == 0:
|
||||||
|
raise ValueError(f"{field} must be a positive integer")
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
def require_schema_version(value: object, expected: int, field: str) -> int:
|
||||||
|
if isinstance(value, bool) or not isinstance(value, int) or value != expected:
|
||||||
|
raise ValueError(f"{field} must be the integer {expected}")
|
||||||
|
return value
|
||||||
|
|
||||||
|
|
||||||
|
def require_task_key(value: object, field: str = "task_key") -> str:
|
||||||
|
text = require_string(value, field, max_length=256)
|
||||||
|
parts = text.split("/")
|
||||||
|
if any(
|
||||||
|
not part or part in {".", ".."} or not _TASK_KEY_COMPONENT_PATTERN.fullmatch(part)
|
||||||
|
for part in parts
|
||||||
|
):
|
||||||
|
raise ValueError(f"{field} must be a canonical workflow-relative key")
|
||||||
|
return text
|
||||||
|
|
||||||
|
|
||||||
|
def contains_unsafe_location(value: str) -> bool:
|
||||||
|
stripped = value.strip()
|
||||||
|
variants = [stripped]
|
||||||
|
for _ in range(2):
|
||||||
|
decoded = unquote(variants[-1])
|
||||||
|
if decoded == variants[-1]:
|
||||||
|
break
|
||||||
|
variants.append(decoded)
|
||||||
|
for candidate in variants:
|
||||||
|
if _SECRET_ASSIGNMENT_PATTERN.search(candidate):
|
||||||
|
return True
|
||||||
|
fragments = (candidate,) + tuple(
|
||||||
|
fragment
|
||||||
|
for fragment in re.split(r"[\s=\"'()\[\]{}<>;,]+", candidate)
|
||||||
|
if fragment
|
||||||
|
)
|
||||||
|
for fragment in fragments:
|
||||||
|
lower = fragment.lower()
|
||||||
|
normalized = fragment.replace("\\", "/")
|
||||||
|
segments = normalized.split("/")
|
||||||
|
looks_relative = (
|
||||||
|
len(segments) >= 3
|
||||||
|
and all(_PATH_SEGMENT_PATTERN.fullmatch(segment) for segment in segments)
|
||||||
|
) or (
|
||||||
|
len(segments) >= 2
|
||||||
|
and all(_PATH_SEGMENT_PATTERN.fullmatch(segment) for segment in segments)
|
||||||
|
and bool(_FILE_NAME_PATTERN.fullmatch(segments[-1]))
|
||||||
|
)
|
||||||
|
if (
|
||||||
|
bool(_URI_SCHEME_PATTERN.match(fragment))
|
||||||
|
or lower.startswith(tuple(prefix.lower() for prefix in _FORBIDDEN_LOCATOR_PREFIXES))
|
||||||
|
or fragment.startswith(("./", "../", "~/", "\\\\"))
|
||||||
|
or bool(_WINDOWS_PATH_PATTERN.match(fragment))
|
||||||
|
or any(segment == ".." for segment in segments)
|
||||||
|
or looks_relative
|
||||||
|
or (
|
||||||
|
_PATH_ASSIGNMENT_PATTERN.search(candidate) is not None
|
||||||
|
and ("/" in fragment or "\\" in fragment)
|
||||||
|
)
|
||||||
|
):
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def require_safe_message(value: object, field: str, *, max_length: int = 512) -> str:
|
||||||
|
text = require_string(value, field, max_length=max_length)
|
||||||
|
tokens = (text,) + tuple(text.split())
|
||||||
|
if any(contains_unsafe_location(token.strip("'\"()[]{}<>,;")) for token in tokens):
|
||||||
|
raise ValueError(f"{field} must not contain a URI or local path")
|
||||||
|
return text
|
||||||
|
|
||||||
|
|
||||||
|
def require_opaque_resource_id(value: object, field: str) -> str:
|
||||||
|
"""Validate a non-secret resource handle without treating it as a locator."""
|
||||||
|
text = require_string(value, field, max_length=160)
|
||||||
|
if (
|
||||||
|
contains_unsafe_location(text)
|
||||||
|
or "/" in text
|
||||||
|
or "\\" in text
|
||||||
|
or "," in text
|
||||||
|
or any(character.isspace() for character in text)
|
||||||
|
):
|
||||||
|
raise ValueError(f"{field} must be an opaque single resource identifier")
|
||||||
|
return text
|
||||||
|
|
||||||
|
|
||||||
|
def require_finite_number(value: object, field: str) -> int | float:
|
||||||
|
if isinstance(value, bool) or not isinstance(value, (int, float)):
|
||||||
|
raise ValueError(f"{field} must be a finite number")
|
||||||
|
if isinstance(value, int):
|
||||||
|
if abs(value).bit_length() > 4096:
|
||||||
|
raise ValueError(f"{field} exceeds the 4096-bit integer bound")
|
||||||
|
return value
|
||||||
|
if not math.isfinite(value):
|
||||||
|
raise ValueError(f"{field} must be a finite number")
|
||||||
|
return value
|
||||||
|
|
||||||
|
|
||||||
|
def freeze_json(
|
||||||
|
value: object,
|
||||||
|
field: str,
|
||||||
|
*,
|
||||||
|
forbid_locations: bool = False,
|
||||||
|
_depth: int = 0,
|
||||||
|
) -> Any:
|
||||||
|
"""Return an immutable deep copy of a JSON value.
|
||||||
|
|
||||||
|
Scientific task parameters use ``forbid_locations`` so durable payloads
|
||||||
|
cannot smuggle worker-local paths or transport URLs. Manifests and verifier
|
||||||
|
evidence use ordinary JSON validation because JSON Schema keywords and
|
||||||
|
sanitized references may legitimately contain URI-shaped strings.
|
||||||
|
"""
|
||||||
|
if _depth > 64:
|
||||||
|
raise ValueError(f"{field} nesting exceeds 64 levels")
|
||||||
|
if value is None or isinstance(value, bool):
|
||||||
|
return value
|
||||||
|
if isinstance(value, int):
|
||||||
|
if abs(value).bit_length() > 4096:
|
||||||
|
raise ValueError(f"{field} contains an integer above the 4096-bit JSON bound")
|
||||||
|
return value
|
||||||
|
if isinstance(value, str):
|
||||||
|
if any(ord(character) < 32 for character in value):
|
||||||
|
raise ValueError(f"{field} must not contain control characters")
|
||||||
|
if forbid_locations and contains_unsafe_location(value):
|
||||||
|
raise ValueError(f"{field} must not contain a URI or local path")
|
||||||
|
return value
|
||||||
|
if isinstance(value, float):
|
||||||
|
if not math.isfinite(value):
|
||||||
|
raise ValueError(f"{field} must not contain NaN or infinity")
|
||||||
|
return value
|
||||||
|
if isinstance(value, Mapping):
|
||||||
|
frozen: dict[str, Any] = {}
|
||||||
|
for key, child in value.items():
|
||||||
|
if not isinstance(key, str):
|
||||||
|
raise ValueError(f"{field} must use string object keys")
|
||||||
|
frozen[key] = freeze_json(
|
||||||
|
child,
|
||||||
|
f"{field}.{key}",
|
||||||
|
forbid_locations=forbid_locations,
|
||||||
|
_depth=_depth + 1,
|
||||||
|
)
|
||||||
|
return MappingProxyType(frozen)
|
||||||
|
if isinstance(value, (list, tuple)):
|
||||||
|
return tuple(
|
||||||
|
freeze_json(
|
||||||
|
child,
|
||||||
|
f"{field}[]",
|
||||||
|
forbid_locations=forbid_locations,
|
||||||
|
_depth=_depth + 1,
|
||||||
|
)
|
||||||
|
for child in value
|
||||||
|
)
|
||||||
|
raise ValueError(f"{field} must contain only JSON-compatible values")
|
||||||
|
|
||||||
|
|
||||||
|
def freeze_json_mapping(
|
||||||
|
value: object,
|
||||||
|
field: str,
|
||||||
|
*,
|
||||||
|
forbid_locations: bool = False,
|
||||||
|
) -> Mapping[str, Any]:
|
||||||
|
mapping = require_mapping(value, field)
|
||||||
|
frozen = freeze_json(mapping, field, forbid_locations=forbid_locations)
|
||||||
|
assert isinstance(frozen, Mapping)
|
||||||
|
return frozen
|
||||||
|
|
||||||
|
|
||||||
|
def thaw_json(value: object) -> Any:
|
||||||
|
if isinstance(value, Mapping):
|
||||||
|
return {key: thaw_json(child) for key, child in value.items()}
|
||||||
|
if isinstance(value, tuple):
|
||||||
|
return [thaw_json(child) for child in value]
|
||||||
|
return value
|
||||||
|
|
||||||
|
|
||||||
|
def canonical_json(value: object) -> str:
|
||||||
|
return json.dumps(thaw_json(value), sort_keys=True, separators=(",", ":"), allow_nan=False)
|
||||||
|
|
||||||
|
|
||||||
|
def parse_release(value: object, field: str = "version") -> tuple[int, int, int]:
|
||||||
|
text = require_string(value, field, max_length=32)
|
||||||
|
match = _VERSION_PATTERN.fullmatch(text)
|
||||||
|
if match is None:
|
||||||
|
raise ValueError(f"{field} must contain one to three numeric release components")
|
||||||
|
return tuple(int(part or 0) for part in match.groups()) # type: ignore[return-value]
|
||||||
|
|
||||||
|
|
||||||
|
def validate_version_range(expression: object, field: str) -> str:
|
||||||
|
text = require_string(expression, field, max_length=128)
|
||||||
|
clauses = [clause.strip() for clause in text.split(",")]
|
||||||
|
if not clauses or any(not clause for clause in clauses):
|
||||||
|
raise ValueError(f"{field} must be an explicit version range")
|
||||||
|
canonical_clauses: list[str] = []
|
||||||
|
for clause in clauses:
|
||||||
|
match = _VERSION_CLAUSE_PATTERN.fullmatch(clause)
|
||||||
|
if match is None:
|
||||||
|
raise ValueError(f"{field} must use ==, >=, <=, >, or < clauses")
|
||||||
|
bound = match.group(2).strip()
|
||||||
|
parse_release(bound, field)
|
||||||
|
canonical_clauses.append(match.group(1) + bound)
|
||||||
|
return ",".join(canonical_clauses)
|
||||||
|
|
||||||
|
|
||||||
|
def version_in_range(version: object, expression: str) -> bool:
|
||||||
|
candidate = parse_release(version)
|
||||||
|
for clause in expression.split(","):
|
||||||
|
match = _VERSION_CLAUSE_PATTERN.fullmatch(clause)
|
||||||
|
assert match is not None
|
||||||
|
operator, raw_bound = match.groups()
|
||||||
|
bound = parse_release(raw_bound)
|
||||||
|
if operator == "==" and candidate != bound:
|
||||||
|
return False
|
||||||
|
if operator == ">=" and candidate < bound:
|
||||||
|
return False
|
||||||
|
if operator == "<=" and candidate > bound:
|
||||||
|
return False
|
||||||
|
if operator == ">" and candidate <= bound:
|
||||||
|
return False
|
||||||
|
if operator == "<" and candidate >= bound:
|
||||||
|
return False
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
|
def enum_value(enum_type: type[Any], value: object, field: str) -> Any:
|
||||||
|
try:
|
||||||
|
return enum_type(value)
|
||||||
|
except (TypeError, ValueError) as error:
|
||||||
|
allowed = ", ".join(member.value for member in enum_type)
|
||||||
|
raise ValueError(f"{field} must be one of: {allowed}") from error
|
||||||
@@ -0,0 +1,705 @@
|
|||||||
|
"""Typed artifact ports, immutable collections, and output provenance."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import hashlib
|
||||||
|
import json
|
||||||
|
from dataclasses import dataclass, field
|
||||||
|
from datetime import datetime, timezone
|
||||||
|
from enum import Enum
|
||||||
|
from types import MappingProxyType
|
||||||
|
from typing import Any, Mapping, Sequence
|
||||||
|
|
||||||
|
from ._validation import (
|
||||||
|
canonical_json,
|
||||||
|
enum_value,
|
||||||
|
freeze_json_mapping,
|
||||||
|
require_exact_keys,
|
||||||
|
require_finite_number,
|
||||||
|
require_identifier,
|
||||||
|
require_nonnegative_int,
|
||||||
|
require_opaque_resource_id,
|
||||||
|
require_positive_int,
|
||||||
|
require_sha256,
|
||||||
|
require_schema_version,
|
||||||
|
require_string,
|
||||||
|
require_task_key,
|
||||||
|
require_uuid,
|
||||||
|
thaw_json,
|
||||||
|
parse_release,
|
||||||
|
)
|
||||||
|
from .identity import ComponentRef, OUTPUT_SCHEMA_VERSION, SchemaRef, WorkloadId
|
||||||
|
|
||||||
|
|
||||||
|
class CollectionKind(str, Enum):
|
||||||
|
SINGLE = "single"
|
||||||
|
ORDERED = "ordered"
|
||||||
|
KEYED = "keyed"
|
||||||
|
SET = "set"
|
||||||
|
|
||||||
|
|
||||||
|
class Cardinality(str, Enum):
|
||||||
|
ONE = "one"
|
||||||
|
OPTIONAL = "optional"
|
||||||
|
MANY = "many"
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class ArtifactSchema:
|
||||||
|
"""Logical artifact shape and hard parsing bounds."""
|
||||||
|
|
||||||
|
ref: SchemaRef
|
||||||
|
media_type: str
|
||||||
|
encoding: str | None
|
||||||
|
max_bytes: int
|
||||||
|
validator: ComponentRef
|
||||||
|
validator_configuration: Mapping[str, Any] = field(default_factory=dict)
|
||||||
|
max_records: int | None = None
|
||||||
|
max_dimensions: tuple[int, ...] = ()
|
||||||
|
streaming: bool = False
|
||||||
|
canonicalizer: str | None = None
|
||||||
|
privacy_class: str = "project"
|
||||||
|
retention_class: str = "durable"
|
||||||
|
allow_nested_collections: bool = False
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
if not isinstance(self.ref, SchemaRef):
|
||||||
|
raise ValueError("artifact schema ref must be a SchemaRef")
|
||||||
|
object.__setattr__(self, "media_type", require_string(self.media_type, "media_type", max_length=128))
|
||||||
|
if "/" not in self.media_type or any(character.isspace() for character in self.media_type):
|
||||||
|
raise ValueError("media_type must be a valid type/subtype token")
|
||||||
|
if self.encoding is not None:
|
||||||
|
object.__setattr__(self, "encoding", require_identifier(self.encoding, "encoding"))
|
||||||
|
object.__setattr__(self, "max_bytes", require_positive_int(self.max_bytes, "max_bytes"))
|
||||||
|
if not isinstance(self.validator, ComponentRef):
|
||||||
|
raise ValueError("artifact schema validator must be a ComponentRef")
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"validator_configuration",
|
||||||
|
freeze_json_mapping(
|
||||||
|
self.validator_configuration,
|
||||||
|
"artifact validator_configuration",
|
||||||
|
forbid_locations=True,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
if self.max_records is not None:
|
||||||
|
object.__setattr__(self, "max_records", require_positive_int(self.max_records, "max_records"))
|
||||||
|
dimensions = tuple(self.max_dimensions)
|
||||||
|
if any(
|
||||||
|
isinstance(value, bool) or not isinstance(value, int) or value < 1
|
||||||
|
for value in dimensions
|
||||||
|
):
|
||||||
|
raise ValueError("max_dimensions must contain positive integers")
|
||||||
|
if len(dimensions) > 8:
|
||||||
|
raise ValueError("max_dimensions must contain at most 8 axes")
|
||||||
|
object.__setattr__(self, "max_dimensions", dimensions)
|
||||||
|
if self.canonicalizer is not None:
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"canonicalizer",
|
||||||
|
require_identifier(self.canonicalizer, "canonicalizer"),
|
||||||
|
)
|
||||||
|
object.__setattr__(self, "privacy_class", require_identifier(self.privacy_class, "privacy_class"))
|
||||||
|
object.__setattr__(self, "retention_class", require_identifier(self.retention_class, "retention_class"))
|
||||||
|
if not isinstance(self.streaming, bool) or not isinstance(self.allow_nested_collections, bool):
|
||||||
|
raise ValueError("streaming and allow_nested_collections must be booleans")
|
||||||
|
|
||||||
|
def to_dict(self) -> dict[str, object]:
|
||||||
|
return {
|
||||||
|
"ref": self.ref.canonical,
|
||||||
|
"media_type": self.media_type,
|
||||||
|
"encoding": self.encoding,
|
||||||
|
"max_bytes": self.max_bytes,
|
||||||
|
"validator": self.validator.canonical,
|
||||||
|
"validator_configuration": thaw_json(self.validator_configuration),
|
||||||
|
"max_records": self.max_records,
|
||||||
|
"max_dimensions": list(self.max_dimensions),
|
||||||
|
"streaming": self.streaming,
|
||||||
|
"canonicalizer": self.canonicalizer,
|
||||||
|
"privacy_class": self.privacy_class,
|
||||||
|
"retention_class": self.retention_class,
|
||||||
|
"allow_nested_collections": self.allow_nested_collections,
|
||||||
|
}
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_dict(cls, value: object) -> "ArtifactSchema":
|
||||||
|
if not isinstance(value, Mapping):
|
||||||
|
raise ValueError("artifact schema must be an object")
|
||||||
|
fields = {
|
||||||
|
"ref", "media_type", "encoding", "max_bytes", "validator",
|
||||||
|
"validator_configuration", "max_records",
|
||||||
|
"max_dimensions", "streaming", "canonicalizer", "privacy_class",
|
||||||
|
"retention_class", "allow_nested_collections",
|
||||||
|
}
|
||||||
|
require_exact_keys(value, fields, "artifact schema")
|
||||||
|
dimensions = value["max_dimensions"]
|
||||||
|
if not isinstance(dimensions, list):
|
||||||
|
raise ValueError("max_dimensions must be an array")
|
||||||
|
return cls(
|
||||||
|
ref=SchemaRef.from_dict(value["ref"]),
|
||||||
|
media_type=value["media_type"], # type: ignore[arg-type]
|
||||||
|
encoding=value["encoding"], # type: ignore[arg-type]
|
||||||
|
max_bytes=value["max_bytes"], # type: ignore[arg-type]
|
||||||
|
validator=ComponentRef.from_dict(value["validator"]),
|
||||||
|
validator_configuration=value["validator_configuration"], # type: ignore[arg-type]
|
||||||
|
max_records=value["max_records"], # type: ignore[arg-type]
|
||||||
|
max_dimensions=tuple(dimensions),
|
||||||
|
streaming=value["streaming"], # type: ignore[arg-type]
|
||||||
|
canonicalizer=value["canonicalizer"], # type: ignore[arg-type]
|
||||||
|
privacy_class=value["privacy_class"], # type: ignore[arg-type]
|
||||||
|
retention_class=value["retention_class"], # type: ignore[arg-type]
|
||||||
|
allow_nested_collections=value["allow_nested_collections"], # type: ignore[arg-type]
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class PortSpec:
|
||||||
|
schema: ArtifactSchema
|
||||||
|
cardinality: Cardinality = Cardinality.ONE
|
||||||
|
collection: CollectionKind = CollectionKind.SINGLE
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
if not isinstance(self.schema, ArtifactSchema):
|
||||||
|
raise ValueError("port schema must be an ArtifactSchema")
|
||||||
|
object.__setattr__(self, "cardinality", enum_value(Cardinality, self.cardinality, "cardinality"))
|
||||||
|
object.__setattr__(self, "collection", enum_value(CollectionKind, self.collection, "collection"))
|
||||||
|
if self.cardinality is Cardinality.MANY and self.collection is CollectionKind.SINGLE:
|
||||||
|
raise ValueError("many cardinality requires an ordered, keyed, or set collection")
|
||||||
|
if self.cardinality is not Cardinality.MANY and self.collection is not CollectionKind.SINGLE:
|
||||||
|
raise ValueError("one and optional cardinality require a single collection")
|
||||||
|
|
||||||
|
def validate_collection(self, value: "ArtifactCollection", field: str = "artifact collection") -> None:
|
||||||
|
if value.kind is not self.collection:
|
||||||
|
raise ValueError(f"{field} kind does not match its port declaration")
|
||||||
|
count = len(value.items)
|
||||||
|
if self.cardinality is Cardinality.ONE and count != 1:
|
||||||
|
raise ValueError(f"{field} must contain exactly one artifact")
|
||||||
|
if self.cardinality is Cardinality.OPTIONAL and count > 1:
|
||||||
|
raise ValueError(f"{field} must contain at most one artifact")
|
||||||
|
if self.cardinality is Cardinality.MANY and count < 1:
|
||||||
|
raise ValueError(f"{field} must contain at least one artifact")
|
||||||
|
for item in value.items:
|
||||||
|
artifact = item.artifact
|
||||||
|
if artifact.schema != self.schema.ref:
|
||||||
|
raise ValueError(f"{field} contains an artifact with the wrong schema")
|
||||||
|
if artifact.media_type != self.schema.media_type:
|
||||||
|
raise ValueError(f"{field} contains an artifact with the wrong media type")
|
||||||
|
if artifact.size_bytes > self.schema.max_bytes:
|
||||||
|
raise ValueError(f"{field} exceeds its per-artifact byte limit")
|
||||||
|
if self.schema.max_records is not None:
|
||||||
|
if artifact.records is None:
|
||||||
|
raise ValueError(f"{field} is missing its required record summary")
|
||||||
|
if artifact.records > self.schema.max_records:
|
||||||
|
raise ValueError(f"{field} exceeds its record limit")
|
||||||
|
if self.schema.max_dimensions:
|
||||||
|
if not artifact.dimensions:
|
||||||
|
raise ValueError(f"{field} is missing its required dimension summary")
|
||||||
|
if len(artifact.dimensions) != len(self.schema.max_dimensions) or any(
|
||||||
|
actual > maximum
|
||||||
|
for actual, maximum in zip(artifact.dimensions, self.schema.max_dimensions)
|
||||||
|
):
|
||||||
|
raise ValueError(f"{field} exceeds its dimension limits")
|
||||||
|
|
||||||
|
def to_dict(self) -> dict[str, object]:
|
||||||
|
return {
|
||||||
|
"schema": self.schema.to_dict(),
|
||||||
|
"cardinality": self.cardinality.value,
|
||||||
|
"collection": self.collection.value,
|
||||||
|
}
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_dict(cls, value: object) -> "PortSpec":
|
||||||
|
if not isinstance(value, Mapping):
|
||||||
|
raise ValueError("port specification must be an object")
|
||||||
|
require_exact_keys(value, {"schema", "cardinality", "collection"}, "port specification")
|
||||||
|
return cls(
|
||||||
|
schema=ArtifactSchema.from_dict(value["schema"]),
|
||||||
|
cardinality=value["cardinality"], # type: ignore[arg-type]
|
||||||
|
collection=value["collection"], # type: ignore[arg-type]
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class ArtifactRef:
|
||||||
|
"""Coordinator-owned artifact identity; transport URIs are intentionally absent."""
|
||||||
|
|
||||||
|
artifact_id: str
|
||||||
|
sha256: str
|
||||||
|
schema: SchemaRef
|
||||||
|
media_type: str
|
||||||
|
size_bytes: int
|
||||||
|
records: int | None = None
|
||||||
|
dimensions: tuple[int, ...] = ()
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
object.__setattr__(self, "artifact_id", require_uuid(self.artifact_id, "artifact_id"))
|
||||||
|
object.__setattr__(self, "sha256", require_sha256(self.sha256, "sha256"))
|
||||||
|
if not isinstance(self.schema, SchemaRef):
|
||||||
|
raise ValueError("artifact schema must be a SchemaRef")
|
||||||
|
object.__setattr__(self, "media_type", require_string(self.media_type, "media_type", max_length=128))
|
||||||
|
if "/" not in self.media_type or any(character.isspace() for character in self.media_type):
|
||||||
|
raise ValueError("media_type must be a valid type/subtype token")
|
||||||
|
object.__setattr__(self, "size_bytes", require_nonnegative_int(self.size_bytes, "size_bytes"))
|
||||||
|
if self.records is not None:
|
||||||
|
object.__setattr__(self, "records", require_nonnegative_int(self.records, "records"))
|
||||||
|
dimensions = tuple(self.dimensions)
|
||||||
|
if any(
|
||||||
|
isinstance(value, bool) or not isinstance(value, int) or value < 0
|
||||||
|
for value in dimensions
|
||||||
|
):
|
||||||
|
raise ValueError("dimensions must contain non-negative integers")
|
||||||
|
if len(dimensions) > 8:
|
||||||
|
raise ValueError("dimensions must contain at most 8 axes")
|
||||||
|
object.__setattr__(self, "dimensions", dimensions)
|
||||||
|
|
||||||
|
def to_dict(self) -> dict[str, object]:
|
||||||
|
return {
|
||||||
|
"artifact_id": self.artifact_id,
|
||||||
|
"sha256": self.sha256,
|
||||||
|
"schema": self.schema.canonical,
|
||||||
|
"media_type": self.media_type,
|
||||||
|
"size_bytes": self.size_bytes,
|
||||||
|
"records": self.records,
|
||||||
|
"dimensions": list(self.dimensions),
|
||||||
|
}
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_dict(cls, value: object) -> "ArtifactRef":
|
||||||
|
if not isinstance(value, Mapping):
|
||||||
|
raise ValueError("artifact reference must be an object")
|
||||||
|
require_exact_keys(
|
||||||
|
value,
|
||||||
|
{"artifact_id", "sha256", "schema", "media_type", "size_bytes", "records", "dimensions"},
|
||||||
|
"artifact reference",
|
||||||
|
)
|
||||||
|
dimensions = value["dimensions"]
|
||||||
|
if not isinstance(dimensions, list):
|
||||||
|
raise ValueError("artifact dimensions must be an array")
|
||||||
|
return cls(
|
||||||
|
artifact_id=value["artifact_id"], # type: ignore[arg-type]
|
||||||
|
sha256=value["sha256"], # type: ignore[arg-type]
|
||||||
|
schema=SchemaRef.from_dict(value["schema"]),
|
||||||
|
media_type=value["media_type"], # type: ignore[arg-type]
|
||||||
|
size_bytes=value["size_bytes"], # type: ignore[arg-type]
|
||||||
|
records=value["records"], # type: ignore[arg-type]
|
||||||
|
dimensions=tuple(dimensions),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class ArtifactItem:
|
||||||
|
artifact: ArtifactRef
|
||||||
|
key: str | None = None
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
if not isinstance(self.artifact, ArtifactRef):
|
||||||
|
raise ValueError("artifact item must contain an ArtifactRef")
|
||||||
|
if self.key is not None:
|
||||||
|
object.__setattr__(self, "key", require_identifier(self.key, "artifact key"))
|
||||||
|
|
||||||
|
def to_dict(self) -> dict[str, object]:
|
||||||
|
return {"key": self.key, "artifact": self.artifact.to_dict()}
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_dict(cls, value: object) -> "ArtifactItem":
|
||||||
|
if not isinstance(value, Mapping):
|
||||||
|
raise ValueError("artifact item must be an object")
|
||||||
|
require_exact_keys(value, {"key", "artifact"}, "artifact item")
|
||||||
|
return cls(artifact=ArtifactRef.from_dict(value["artifact"]), key=value["key"]) # type: ignore[arg-type]
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class ArtifactCollection:
|
||||||
|
kind: CollectionKind
|
||||||
|
items: tuple[ArtifactItem, ...]
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
object.__setattr__(self, "kind", enum_value(CollectionKind, self.kind, "collection.kind"))
|
||||||
|
items = tuple(self.items)
|
||||||
|
if any(not isinstance(item, ArtifactItem) for item in items):
|
||||||
|
raise ValueError("collection items must be ArtifactItem values")
|
||||||
|
if self.kind is CollectionKind.SINGLE:
|
||||||
|
if len(items) > 1 or any(item.key is not None for item in items):
|
||||||
|
raise ValueError("single collection contains at most one unkeyed artifact")
|
||||||
|
elif self.kind is CollectionKind.KEYED:
|
||||||
|
if any(item.key is None for item in items):
|
||||||
|
raise ValueError("keyed collection requires a key for every artifact")
|
||||||
|
keys = [item.key for item in items]
|
||||||
|
if len(keys) != len(set(keys)):
|
||||||
|
raise ValueError("keyed collection keys must be unique")
|
||||||
|
items = tuple(sorted(items, key=lambda item: item.key or ""))
|
||||||
|
else:
|
||||||
|
if any(item.key is not None for item in items):
|
||||||
|
raise ValueError("ordered and set collections must not use keys")
|
||||||
|
if self.kind is CollectionKind.SET:
|
||||||
|
identities = [
|
||||||
|
(item.artifact.schema, item.artifact.sha256, item.artifact.size_bytes)
|
||||||
|
for item in items
|
||||||
|
]
|
||||||
|
if len(identities) != len(set(identities)):
|
||||||
|
raise ValueError("set collection must not contain duplicate artifacts")
|
||||||
|
items = tuple(
|
||||||
|
sorted(
|
||||||
|
items,
|
||||||
|
key=lambda item: (
|
||||||
|
item.artifact.schema.canonical,
|
||||||
|
item.artifact.sha256,
|
||||||
|
item.artifact.size_bytes,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
object.__setattr__(self, "items", items)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def single(cls, artifact: ArtifactRef | None) -> "ArtifactCollection":
|
||||||
|
return cls(CollectionKind.SINGLE, () if artifact is None else (ArtifactItem(artifact),))
|
||||||
|
|
||||||
|
@property
|
||||||
|
def size_bytes(self) -> int:
|
||||||
|
return sum(item.artifact.size_bytes for item in self.items)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def digest(self) -> str:
|
||||||
|
payload = {
|
||||||
|
"kind": self.kind.value,
|
||||||
|
"items": [
|
||||||
|
{
|
||||||
|
"key": item.key,
|
||||||
|
"sha256": item.artifact.sha256,
|
||||||
|
"schema": item.artifact.schema.canonical,
|
||||||
|
"media_type": item.artifact.media_type,
|
||||||
|
"size_bytes": item.artifact.size_bytes,
|
||||||
|
}
|
||||||
|
for item in self.items
|
||||||
|
],
|
||||||
|
}
|
||||||
|
return hashlib.sha256(canonical_json(payload).encode("utf-8")).hexdigest()
|
||||||
|
|
||||||
|
def to_dict(self) -> dict[str, object]:
|
||||||
|
return {"kind": self.kind.value, "items": [item.to_dict() for item in self.items]}
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_dict(cls, value: object) -> "ArtifactCollection":
|
||||||
|
if not isinstance(value, Mapping):
|
||||||
|
raise ValueError("artifact collection must be an object")
|
||||||
|
require_exact_keys(value, {"kind", "items"}, "artifact collection")
|
||||||
|
items = value["items"]
|
||||||
|
if not isinstance(items, list):
|
||||||
|
raise ValueError("artifact collection items must be an array")
|
||||||
|
return cls(
|
||||||
|
kind=value["kind"], # type: ignore[arg-type]
|
||||||
|
items=tuple(ArtifactItem.from_dict(item) for item in items),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _timestamp(value: object, field: str) -> str:
|
||||||
|
text = require_string(value, field, max_length=64)
|
||||||
|
try:
|
||||||
|
parsed = datetime.fromisoformat(text.replace("Z", "+00:00"))
|
||||||
|
except ValueError as error:
|
||||||
|
raise ValueError(f"{field} must be an RFC 3339 timestamp") from error
|
||||||
|
if parsed.tzinfo is None:
|
||||||
|
raise ValueError(f"{field} must include a timezone")
|
||||||
|
return parsed.astimezone(timezone.utc).isoformat().replace("+00:00", "Z")
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class Provenance:
|
||||||
|
workload: WorkloadId
|
||||||
|
sdk_api_version: str
|
||||||
|
protocol_version: str
|
||||||
|
manifest_schema_version: int
|
||||||
|
workflow_schema_version: int
|
||||||
|
verifier: ComponentRef
|
||||||
|
artifact_schemas: tuple[SchemaRef, ...]
|
||||||
|
package_digest: str
|
||||||
|
manifest_digest: str
|
||||||
|
environment_digest: str
|
||||||
|
worker_runtime: Mapping[str, Any]
|
||||||
|
allocated_resource_ids: tuple[str, ...]
|
||||||
|
parameters_digest: str
|
||||||
|
input_collection_digest: str
|
||||||
|
execution_contract_digest: str
|
||||||
|
selected_features: Mapping[str, str]
|
||||||
|
optional_fallbacks: Mapping[str, str]
|
||||||
|
job_id: str
|
||||||
|
task_id: str
|
||||||
|
started_at: str
|
||||||
|
finished_at: str
|
||||||
|
trust_mode: str = "trusted"
|
||||||
|
random_seed: int | None = None
|
||||||
|
checkpoint_lineage: tuple[str, ...] = ()
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
if not isinstance(self.workload, WorkloadId):
|
||||||
|
raise ValueError("provenance workload must be a WorkloadId")
|
||||||
|
object.__setattr__(self, "sdk_api_version", require_string(self.sdk_api_version, "sdk_api_version"))
|
||||||
|
object.__setattr__(self, "protocol_version", require_string(self.protocol_version, "protocol_version"))
|
||||||
|
parse_release(self.sdk_api_version, "sdk_api_version")
|
||||||
|
parse_release(self.protocol_version, "protocol_version")
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"manifest_schema_version",
|
||||||
|
require_positive_int(self.manifest_schema_version, "manifest_schema_version"),
|
||||||
|
)
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"workflow_schema_version",
|
||||||
|
require_positive_int(self.workflow_schema_version, "workflow_schema_version"),
|
||||||
|
)
|
||||||
|
if not isinstance(self.verifier, ComponentRef):
|
||||||
|
raise ValueError("provenance verifier must be a ComponentRef")
|
||||||
|
schemas = tuple(self.artifact_schemas)
|
||||||
|
if not schemas or any(not isinstance(schema, SchemaRef) for schema in schemas):
|
||||||
|
raise ValueError("provenance artifact_schemas must contain SchemaRef values")
|
||||||
|
if len(schemas) != len(set(schemas)):
|
||||||
|
raise ValueError("provenance artifact_schemas must be unique")
|
||||||
|
if schemas != tuple(sorted(schemas, key=lambda schema: schema.canonical)):
|
||||||
|
raise ValueError("provenance artifact_schemas must be in canonical order")
|
||||||
|
object.__setattr__(self, "artifact_schemas", schemas)
|
||||||
|
object.__setattr__(self, "package_digest", require_sha256(self.package_digest, "package_digest", prefixed=True))
|
||||||
|
object.__setattr__(self, "manifest_digest", require_sha256(self.manifest_digest, "manifest_digest"))
|
||||||
|
object.__setattr__(self, "environment_digest", require_sha256(self.environment_digest, "environment_digest", prefixed=True))
|
||||||
|
runtime = freeze_json_mapping(self.worker_runtime, "worker_runtime", forbid_locations=True)
|
||||||
|
if len(canonical_json(runtime).encode("utf-8")) > 65_536:
|
||||||
|
raise ValueError("worker_runtime exceeds 64 KiB")
|
||||||
|
object.__setattr__(self, "worker_runtime", runtime)
|
||||||
|
resource_ids = tuple(
|
||||||
|
require_opaque_resource_id(value, "allocated_resource_id")
|
||||||
|
for value in self.allocated_resource_ids
|
||||||
|
)
|
||||||
|
if not resource_ids or len(resource_ids) != len(set(resource_ids)):
|
||||||
|
raise ValueError("allocated_resource_ids must be non-empty and unique")
|
||||||
|
object.__setattr__(self, "allocated_resource_ids", resource_ids)
|
||||||
|
object.__setattr__(self, "parameters_digest", require_sha256(self.parameters_digest, "parameters_digest"))
|
||||||
|
object.__setattr__(self, "input_collection_digest", require_sha256(self.input_collection_digest, "input_collection_digest"))
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"execution_contract_digest",
|
||||||
|
require_sha256(self.execution_contract_digest, "execution_contract_digest"),
|
||||||
|
)
|
||||||
|
selected_features = freeze_json_mapping(
|
||||||
|
self.selected_features,
|
||||||
|
"provenance.selected_features",
|
||||||
|
)
|
||||||
|
optional_fallbacks = freeze_json_mapping(
|
||||||
|
self.optional_fallbacks,
|
||||||
|
"provenance.optional_fallbacks",
|
||||||
|
)
|
||||||
|
for name, version in selected_features.items():
|
||||||
|
require_identifier(name, "provenance selected feature")
|
||||||
|
require_string(version, "provenance selected feature version", max_length=32)
|
||||||
|
parse_release(version, "provenance selected feature version")
|
||||||
|
for name, fallback in optional_fallbacks.items():
|
||||||
|
require_identifier(name, "provenance fallback feature")
|
||||||
|
require_identifier(fallback, "provenance fallback")
|
||||||
|
if set(selected_features).intersection(optional_fallbacks):
|
||||||
|
raise ValueError("provenance feature cannot be selected and fallbacked")
|
||||||
|
object.__setattr__(self, "selected_features", selected_features)
|
||||||
|
object.__setattr__(self, "optional_fallbacks", optional_fallbacks)
|
||||||
|
object.__setattr__(self, "job_id", require_uuid(self.job_id, "provenance.job_id"))
|
||||||
|
object.__setattr__(self, "task_id", require_uuid(self.task_id, "provenance.task_id"))
|
||||||
|
object.__setattr__(self, "started_at", _timestamp(self.started_at, "started_at"))
|
||||||
|
object.__setattr__(self, "finished_at", _timestamp(self.finished_at, "finished_at"))
|
||||||
|
if datetime.fromisoformat(self.finished_at.replace("Z", "+00:00")) < datetime.fromisoformat(
|
||||||
|
self.started_at.replace("Z", "+00:00")
|
||||||
|
):
|
||||||
|
raise ValueError("finished_at must not precede started_at")
|
||||||
|
trust_mode = require_identifier(self.trust_mode, "provenance.trust_mode")
|
||||||
|
if trust_mode not in {"trusted", "verified", "untrusted_quorum"}:
|
||||||
|
raise ValueError("provenance.trust_mode is unsupported")
|
||||||
|
object.__setattr__(self, "trust_mode", trust_mode)
|
||||||
|
if self.random_seed is not None and (isinstance(self.random_seed, bool) or not isinstance(self.random_seed, int)):
|
||||||
|
raise ValueError("random_seed must be an integer")
|
||||||
|
lineage = tuple(require_uuid(value, "checkpoint_lineage") for value in self.checkpoint_lineage)
|
||||||
|
if len(lineage) != len(set(lineage)):
|
||||||
|
raise ValueError("checkpoint_lineage must not contain duplicate artifacts")
|
||||||
|
object.__setattr__(self, "checkpoint_lineage", lineage)
|
||||||
|
|
||||||
|
def to_dict(self) -> dict[str, object]:
|
||||||
|
return {
|
||||||
|
"workload": self.workload.to_dict(),
|
||||||
|
"sdk_api_version": self.sdk_api_version,
|
||||||
|
"protocol_version": self.protocol_version,
|
||||||
|
"manifest_schema_version": self.manifest_schema_version,
|
||||||
|
"workflow_schema_version": self.workflow_schema_version,
|
||||||
|
"verifier": self.verifier.canonical,
|
||||||
|
"artifact_schemas": [schema.canonical for schema in self.artifact_schemas],
|
||||||
|
"package_digest": self.package_digest,
|
||||||
|
"manifest_digest": self.manifest_digest,
|
||||||
|
"environment_digest": self.environment_digest,
|
||||||
|
"worker_runtime": thaw_json(self.worker_runtime),
|
||||||
|
"allocated_resource_ids": list(self.allocated_resource_ids),
|
||||||
|
"parameters_digest": self.parameters_digest,
|
||||||
|
"input_collection_digest": self.input_collection_digest,
|
||||||
|
"execution_contract_digest": self.execution_contract_digest,
|
||||||
|
"selected_features": thaw_json(self.selected_features),
|
||||||
|
"optional_fallbacks": thaw_json(self.optional_fallbacks),
|
||||||
|
"job_id": self.job_id,
|
||||||
|
"task_id": self.task_id,
|
||||||
|
"started_at": self.started_at,
|
||||||
|
"finished_at": self.finished_at,
|
||||||
|
"trust_mode": self.trust_mode,
|
||||||
|
"random_seed": self.random_seed,
|
||||||
|
"checkpoint_lineage": list(self.checkpoint_lineage),
|
||||||
|
}
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_dict(cls, value: object) -> "Provenance":
|
||||||
|
if not isinstance(value, Mapping):
|
||||||
|
raise ValueError("provenance must be an object")
|
||||||
|
fields = {
|
||||||
|
"workload", "sdk_api_version", "protocol_version", "manifest_schema_version",
|
||||||
|
"workflow_schema_version", "verifier", "artifact_schemas", "package_digest",
|
||||||
|
"manifest_digest", "environment_digest", "worker_runtime", "allocated_resource_ids",
|
||||||
|
"parameters_digest", "input_collection_digest", "execution_contract_digest",
|
||||||
|
"selected_features", "optional_fallbacks",
|
||||||
|
"job_id", "task_id",
|
||||||
|
"started_at", "finished_at",
|
||||||
|
"trust_mode", "random_seed", "checkpoint_lineage",
|
||||||
|
}
|
||||||
|
require_exact_keys(value, fields, "provenance")
|
||||||
|
resource_ids = value["allocated_resource_ids"]
|
||||||
|
artifact_schemas = value["artifact_schemas"]
|
||||||
|
lineage = value["checkpoint_lineage"]
|
||||||
|
if not isinstance(resource_ids, list) or not isinstance(artifact_schemas, list) or not isinstance(lineage, list):
|
||||||
|
raise ValueError("provenance resource IDs and checkpoint lineage must be arrays")
|
||||||
|
return cls(
|
||||||
|
workload=WorkloadId.from_dict(value["workload"]),
|
||||||
|
sdk_api_version=value["sdk_api_version"], # type: ignore[arg-type]
|
||||||
|
protocol_version=value["protocol_version"], # type: ignore[arg-type]
|
||||||
|
manifest_schema_version=value["manifest_schema_version"], # type: ignore[arg-type]
|
||||||
|
workflow_schema_version=value["workflow_schema_version"], # type: ignore[arg-type]
|
||||||
|
verifier=ComponentRef.from_dict(value["verifier"]),
|
||||||
|
artifact_schemas=tuple(SchemaRef.from_dict(item) for item in artifact_schemas),
|
||||||
|
package_digest=value["package_digest"], # type: ignore[arg-type]
|
||||||
|
manifest_digest=value["manifest_digest"], # type: ignore[arg-type]
|
||||||
|
environment_digest=value["environment_digest"], # type: ignore[arg-type]
|
||||||
|
worker_runtime=value["worker_runtime"], # type: ignore[arg-type]
|
||||||
|
allocated_resource_ids=tuple(resource_ids),
|
||||||
|
parameters_digest=value["parameters_digest"], # type: ignore[arg-type]
|
||||||
|
input_collection_digest=value["input_collection_digest"], # type: ignore[arg-type]
|
||||||
|
execution_contract_digest=value["execution_contract_digest"], # type: ignore[arg-type]
|
||||||
|
selected_features=value["selected_features"], # type: ignore[arg-type]
|
||||||
|
optional_fallbacks=value["optional_fallbacks"], # type: ignore[arg-type]
|
||||||
|
job_id=value["job_id"], # type: ignore[arg-type]
|
||||||
|
task_id=value["task_id"], # type: ignore[arg-type]
|
||||||
|
started_at=value["started_at"], # type: ignore[arg-type]
|
||||||
|
finished_at=value["finished_at"], # type: ignore[arg-type]
|
||||||
|
trust_mode=value["trust_mode"], # type: ignore[arg-type]
|
||||||
|
random_seed=value["random_seed"], # type: ignore[arg-type]
|
||||||
|
checkpoint_lineage=tuple(lineage),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class OutputManifest:
|
||||||
|
task_key: str
|
||||||
|
outputs: Mapping[str, ArtifactCollection]
|
||||||
|
metrics: Mapping[str, int | float]
|
||||||
|
provenance: Provenance
|
||||||
|
schema_version: int = OUTPUT_SCHEMA_VERSION
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
require_schema_version(self.schema_version, OUTPUT_SCHEMA_VERSION, "output schema_version")
|
||||||
|
object.__setattr__(self, "task_key", require_task_key(self.task_key))
|
||||||
|
if not isinstance(self.outputs, Mapping) or not self.outputs:
|
||||||
|
raise ValueError("outputs must be a non-empty object")
|
||||||
|
outputs: dict[str, ArtifactCollection] = {}
|
||||||
|
for name, collection in self.outputs.items():
|
||||||
|
canonical = require_identifier(name, "output port")
|
||||||
|
if not isinstance(collection, ArtifactCollection):
|
||||||
|
raise ValueError("output values must be ArtifactCollection values")
|
||||||
|
outputs[canonical] = collection
|
||||||
|
object.__setattr__(self, "outputs", MappingProxyType(outputs))
|
||||||
|
if not isinstance(self.metrics, Mapping):
|
||||||
|
raise ValueError("metrics must be an object")
|
||||||
|
metrics: dict[str, int | float] = {}
|
||||||
|
for name, value in self.metrics.items():
|
||||||
|
canonical = require_identifier(name, "metric name")
|
||||||
|
metrics[canonical] = require_finite_number(value, "metric value")
|
||||||
|
if len(canonical_json(metrics).encode("utf-8")) > 16_384:
|
||||||
|
raise ValueError("output metrics exceed 16 KiB")
|
||||||
|
object.__setattr__(self, "metrics", MappingProxyType(metrics))
|
||||||
|
if not isinstance(self.provenance, Provenance):
|
||||||
|
raise ValueError("provenance must be a Provenance value")
|
||||||
|
|
||||||
|
def validate_against(
|
||||||
|
self,
|
||||||
|
expected: Mapping[str, PortSpec],
|
||||||
|
*,
|
||||||
|
max_output_bytes: int,
|
||||||
|
) -> "OutputManifest":
|
||||||
|
if set(self.outputs) != set(expected):
|
||||||
|
missing = sorted(set(expected) - set(self.outputs))
|
||||||
|
unexpected = sorted(set(self.outputs) - set(expected))
|
||||||
|
details = []
|
||||||
|
if missing:
|
||||||
|
details.append("missing " + ", ".join(missing))
|
||||||
|
if unexpected:
|
||||||
|
details.append("unexpected " + ", ".join(unexpected))
|
||||||
|
raise ValueError("output ports do not match the declaration: " + "; ".join(details))
|
||||||
|
total = 0
|
||||||
|
for name, port in expected.items():
|
||||||
|
if not isinstance(port, PortSpec):
|
||||||
|
raise ValueError("expected outputs must contain PortSpec values")
|
||||||
|
port.validate_collection(self.outputs[name], f"output {name}")
|
||||||
|
total += self.outputs[name].size_bytes
|
||||||
|
if total > require_positive_int(max_output_bytes, "max_output_bytes"):
|
||||||
|
raise ValueError("output manifest exceeds the total byte limit")
|
||||||
|
return self
|
||||||
|
|
||||||
|
@property
|
||||||
|
def digest(self) -> str:
|
||||||
|
payload = {
|
||||||
|
"outputs": {
|
||||||
|
name: {"kind": collection.kind.value, "digest": collection.digest}
|
||||||
|
for name, collection in sorted(self.outputs.items())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return hashlib.sha256(canonical_json(payload).encode("utf-8")).hexdigest()
|
||||||
|
|
||||||
|
@property
|
||||||
|
def manifest_digest(self) -> str:
|
||||||
|
"""Digest the complete audit manifest, including provenance and metrics."""
|
||||||
|
return hashlib.sha256(self.to_json().encode("utf-8")).hexdigest()
|
||||||
|
|
||||||
|
def to_dict(self) -> dict[str, object]:
|
||||||
|
return {
|
||||||
|
"schema_version": self.schema_version,
|
||||||
|
"task_key": self.task_key,
|
||||||
|
"outputs": {name: value.to_dict() for name, value in self.outputs.items()},
|
||||||
|
"metrics": dict(self.metrics),
|
||||||
|
"provenance": self.provenance.to_dict(),
|
||||||
|
}
|
||||||
|
|
||||||
|
def to_json(self) -> str:
|
||||||
|
return canonical_json(self.to_dict())
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_dict(cls, value: object) -> "OutputManifest":
|
||||||
|
if not isinstance(value, Mapping):
|
||||||
|
raise ValueError("output manifest must be an object")
|
||||||
|
require_exact_keys(
|
||||||
|
value,
|
||||||
|
{"schema_version", "task_key", "outputs", "metrics", "provenance"},
|
||||||
|
"output manifest",
|
||||||
|
)
|
||||||
|
outputs = value["outputs"]
|
||||||
|
if not isinstance(outputs, Mapping):
|
||||||
|
raise ValueError("outputs must be an object")
|
||||||
|
return cls(
|
||||||
|
schema_version=value["schema_version"], # type: ignore[arg-type]
|
||||||
|
task_key=value["task_key"], # type: ignore[arg-type]
|
||||||
|
outputs={name: ArtifactCollection.from_dict(item) for name, item in outputs.items()},
|
||||||
|
metrics=value["metrics"], # type: ignore[arg-type]
|
||||||
|
provenance=Provenance.from_dict(value["provenance"]),
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_json(cls, value: str) -> "OutputManifest":
|
||||||
|
try:
|
||||||
|
decoded = json.loads(value)
|
||||||
|
except (TypeError, json.JSONDecodeError, RecursionError) as error:
|
||||||
|
raise ValueError("output manifest must be valid JSON") from error
|
||||||
|
return cls.from_dict(decoded)
|
||||||
@@ -0,0 +1,146 @@
|
|||||||
|
"""SDK definitions for existing SciMesh workloads and local core runtime."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import hashlib
|
||||||
|
import os
|
||||||
|
import platform
|
||||||
|
import sys
|
||||||
|
|
||||||
|
from rdkit import rdBase
|
||||||
|
|
||||||
|
from scimesh.distributed.similarity_search import (
|
||||||
|
SimilaritySearchDistributedWorkload,
|
||||||
|
run_similarity_search_shard,
|
||||||
|
)
|
||||||
|
|
||||||
|
from .artifacts import ArtifactSchema, PortSpec
|
||||||
|
from .compat import LegacyDistributedWorkloadAdapter
|
||||||
|
from .identity import ComponentRef, SDK_API_VERSION, SchemaRef
|
||||||
|
from .integrity import installed_distribution_digest
|
||||||
|
from .registry import WorkloadRegistry
|
||||||
|
from .resources import ResourceInventory
|
||||||
|
from .runtime import RuntimeCapabilities
|
||||||
|
|
||||||
|
|
||||||
|
def current_scimesh_package_digest() -> str:
|
||||||
|
"""Hash installed SciMesh Python sources for the built-in trusted adapter.
|
||||||
|
|
||||||
|
This is a local immutable-code pin, not a package signature or container
|
||||||
|
attestation. Consequently the built-in compatibility manifest is trusted
|
||||||
|
only; an administrator must supply signed image metadata before enabling an
|
||||||
|
untrusted quorum policy.
|
||||||
|
"""
|
||||||
|
# Source/editable installs are allowed only for this explicit local
|
||||||
|
# development helper. Registry discovery keeps the secure default.
|
||||||
|
return installed_distribution_digest("scimesh", allow_editable=True)
|
||||||
|
|
||||||
|
|
||||||
|
def current_environment_digest() -> str:
|
||||||
|
payload = "\n".join(
|
||||||
|
(
|
||||||
|
current_scimesh_package_digest(),
|
||||||
|
f"python={sys.implementation.name}-{platform.python_version()}",
|
||||||
|
f"rdkit={rdBase.rdkitVersion}",
|
||||||
|
f"platform={sys.platform}-{platform.machine().lower()}",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
return "sha256:" + hashlib.sha256(payload.encode("utf-8")).hexdigest()
|
||||||
|
|
||||||
|
|
||||||
|
def similarity_search_sdk_adapter(*, shard_rows: int = 10_000) -> LegacyDistributedWorkloadAdapter:
|
||||||
|
dataset_schema = 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",
|
||||||
|
)
|
||||||
|
partial_schema = ArtifactSchema(
|
||||||
|
SchemaRef("similarity-search-partial", 1),
|
||||||
|
"text/csv",
|
||||||
|
"utf-8",
|
||||||
|
max_bytes=1024 * 1024 * 1024,
|
||||||
|
validator=ComponentRef("delimited-table", 1),
|
||||||
|
validator_configuration={
|
||||||
|
"columns": ["rank", "chembl_id", "canonical_smiles", "similarity"],
|
||||||
|
},
|
||||||
|
max_records=100_000,
|
||||||
|
canonicalizer="scimesh-search-partial-v1",
|
||||||
|
)
|
||||||
|
result_schema = ArtifactSchema(
|
||||||
|
SchemaRef("similarity-search-result", 1),
|
||||||
|
"text/csv",
|
||||||
|
"utf-8",
|
||||||
|
max_bytes=1024 * 1024 * 1024,
|
||||||
|
validator=ComponentRef("delimited-table", 1),
|
||||||
|
validator_configuration={
|
||||||
|
"columns": ["rank", "chembl_id", "canonical_smiles", "similarity"],
|
||||||
|
},
|
||||||
|
max_records=100_000,
|
||||||
|
canonicalizer="scimesh-search-result-v1",
|
||||||
|
)
|
||||||
|
parameters_schema = {
|
||||||
|
"type": "object",
|
||||||
|
"additionalProperties": False,
|
||||||
|
"properties": {
|
||||||
|
"query_id": {"type": "string", "minLength": 1, "maxLength": 200},
|
||||||
|
"query_smiles": {"type": "string", "minLength": 1, "maxLength": 200},
|
||||||
|
"top_k": {"type": "integer", "minimum": 1},
|
||||||
|
"threshold": {"type": "number", "minimum": 0, "maximum": 1},
|
||||||
|
"threshold_direction": {"enum": ["greater", "less"]},
|
||||||
|
"max_rows": {"type": "integer", "minimum": 1},
|
||||||
|
"progress_every": {"type": "integer", "minimum": 0},
|
||||||
|
},
|
||||||
|
"oneOf": [
|
||||||
|
{"required": ["query_id"], "not": {"required": ["query_smiles"]}},
|
||||||
|
{"required": ["query_smiles"], "not": {"required": ["query_id"]}},
|
||||||
|
],
|
||||||
|
}
|
||||||
|
return LegacyDistributedWorkloadAdapter(
|
||||||
|
SimilaritySearchDistributedWorkload(),
|
||||||
|
run_similarity_search_shard,
|
||||||
|
version="1.0.0",
|
||||||
|
package_digest=current_scimesh_package_digest(),
|
||||||
|
environment_digest=current_environment_digest(),
|
||||||
|
parameters_schema=parameters_schema,
|
||||||
|
input_port=PortSpec(dataset_schema),
|
||||||
|
partial_port=PortSpec(partial_schema),
|
||||||
|
output_port=PortSpec(result_schema),
|
||||||
|
resolved_parameter_names=("query_source", "fingerprint"),
|
||||||
|
shard_rows=shard_rows,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
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)
|
||||||
|
return registry
|
||||||
|
|
||||||
|
|
||||||
|
def similarity_search_workload_definition():
|
||||||
|
"""Installed entry-point factory for the default shard-size definition."""
|
||||||
|
return similarity_search_sdk_adapter().definition()
|
||||||
|
|
||||||
|
|
||||||
|
def default_sdk_runtime() -> RuntimeCapabilities:
|
||||||
|
architecture = platform.machine().lower() or "unknown"
|
||||||
|
return RuntimeCapabilities(
|
||||||
|
sdk_api_version=SDK_API_VERSION,
|
||||||
|
protocol_version="1.0.0",
|
||||||
|
profiles=("core-batch-v1",),
|
||||||
|
features={"artifact-collections": "1.0.0", "exact-verifier": "1.0.0"},
|
||||||
|
workload_capabilities=("similarity-search",),
|
||||||
|
inventory=ResourceInventory(
|
||||||
|
cpu_cores=max(os.cpu_count() or 1, 1),
|
||||||
|
memory_mb=4096,
|
||||||
|
scratch_mb=4096,
|
||||||
|
architecture=architecture,
|
||||||
|
environment_digests=(current_environment_digest(),),
|
||||||
|
),
|
||||||
|
)
|
||||||
@@ -0,0 +1,5 @@
|
|||||||
|
"""Adapters for versioned pre-SDK SciMesh workload contracts."""
|
||||||
|
|
||||||
|
from .distributed_v1 import LegacyDistributedWorkloadAdapter
|
||||||
|
|
||||||
|
__all__ = ["LegacyDistributedWorkloadAdapter"]
|
||||||
@@ -0,0 +1,369 @@
|
|||||||
|
"""Compatibility adapter for the CTX-07 ``DistributedWorkload`` protocol."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import hashlib
|
||||||
|
import shutil
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any, Callable, Mapping, Sequence
|
||||||
|
|
||||||
|
from scimesh.distributed.models import (
|
||||||
|
ArtifactReference as LegacyArtifactReference,
|
||||||
|
CompletedPartial,
|
||||||
|
FinalResult,
|
||||||
|
)
|
||||||
|
from scimesh.distributed.workload import DistributedWorkload
|
||||||
|
|
||||||
|
from ..artifacts import (
|
||||||
|
ArtifactCollection,
|
||||||
|
ArtifactItem,
|
||||||
|
ArtifactRef,
|
||||||
|
Cardinality,
|
||||||
|
CollectionKind,
|
||||||
|
OutputManifest,
|
||||||
|
PortSpec,
|
||||||
|
)
|
||||||
|
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
|
||||||
|
|
||||||
|
|
||||||
|
ShardRunner = Callable[[Path, Mapping[str, object], Path], Mapping[str, int | float]]
|
||||||
|
|
||||||
|
|
||||||
|
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()
|
||||||
|
|
||||||
|
|
||||||
|
class LegacyDistributedWorkloadAdapter:
|
||||||
|
"""Expose a legacy map/reduce workload through the SDK core-batch profile.
|
||||||
|
|
||||||
|
The adapter preserves the old wire schema. Local files are materialized and
|
||||||
|
sealed only through bridge-owned contexts, and no path is included in a
|
||||||
|
``TaskSpec`` or ``WorkflowPlan``.
|
||||||
|
"""
|
||||||
|
|
||||||
|
MAP_ENTRY_POINT = "scimesh.sdk.compat.distributed_v1:run_legacy@v1"
|
||||||
|
REDUCE_ENTRY_POINT = "scimesh.sdk.compat.distributed_v1:reduce_legacy@v1"
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
workload: DistributedWorkload,
|
||||||
|
shard_runner: ShardRunner,
|
||||||
|
*,
|
||||||
|
version: str,
|
||||||
|
package_digest: str,
|
||||||
|
environment_digest: str,
|
||||||
|
parameters_schema: Mapping[str, Any],
|
||||||
|
input_port: PortSpec,
|
||||||
|
partial_port: PortSpec,
|
||||||
|
output_port: PortSpec,
|
||||||
|
resolved_parameter_names: Sequence[str] = (),
|
||||||
|
shard_rows: int = 10_000,
|
||||||
|
resources: ResourceRequirements | None = None,
|
||||||
|
execution: ExecutionProfile | None = None,
|
||||||
|
limits: WorkloadLimits | None = None,
|
||||||
|
) -> None:
|
||||||
|
if not isinstance(workload.name, str) or not isinstance(workload.description, str):
|
||||||
|
raise ValueError("legacy workload must expose name and description")
|
||||||
|
if not callable(shard_runner):
|
||||||
|
raise ValueError("shard_runner must be callable")
|
||||||
|
if isinstance(shard_rows, bool) or not isinstance(shard_rows, int) or shard_rows < 1:
|
||||||
|
raise ValueError("shard_rows must be a positive integer")
|
||||||
|
self.workload = workload
|
||||||
|
self.shard_runner = shard_runner
|
||||||
|
self.shard_rows = shard_rows
|
||||||
|
self.input_port = input_port
|
||||||
|
self.partial_port = partial_port
|
||||||
|
self.output_port = output_port
|
||||||
|
resources = resources or ResourceRequirements(
|
||||||
|
profile="legacy-cpu-v1",
|
||||||
|
cpu_cores=1,
|
||||||
|
memory_mb=1024,
|
||||||
|
scratch_mb=1024,
|
||||||
|
max_duration_seconds=3600,
|
||||||
|
)
|
||||||
|
execution = execution or ExecutionProfile(
|
||||||
|
profile="legacy-python-process-v1",
|
||||||
|
network=NetworkPolicy.TRUSTED,
|
||||||
|
timeout_seconds=3600,
|
||||||
|
checkpoint=CheckpointPolicy(),
|
||||||
|
)
|
||||||
|
limits = limits or WorkloadLimits(
|
||||||
|
max_input_bytes=input_port.schema.max_bytes,
|
||||||
|
max_tasks=10_000,
|
||||||
|
max_output_bytes=output_port.schema.max_bytes,
|
||||||
|
)
|
||||||
|
parameter_names = tuple(sorted(parameters_schema.get("properties", {})))
|
||||||
|
reduce_parameter_names = tuple(sorted(set(parameter_names).union(resolved_parameter_names)))
|
||||||
|
map_stage = StageSpec(
|
||||||
|
stage_id="map",
|
||||||
|
kind=StageKind.MAP,
|
||||||
|
entry_point=self.MAP_ENTRY_POINT,
|
||||||
|
needs=(),
|
||||||
|
inputs={"input": input_port},
|
||||||
|
outputs={"partial": partial_port},
|
||||||
|
parameter_names=parameter_names,
|
||||||
|
resources=resources,
|
||||||
|
execution=execution,
|
||||||
|
retry=RetryPolicy(),
|
||||||
|
verifier=ComponentRef("exact-artifact", 1),
|
||||||
|
trust_modes=("trusted",),
|
||||||
|
max_fan_out=limits.max_tasks,
|
||||||
|
cacheable=True,
|
||||||
|
)
|
||||||
|
reduce_input = PortSpec(
|
||||||
|
schema=partial_port.schema,
|
||||||
|
cardinality=Cardinality.MANY,
|
||||||
|
collection=CollectionKind.KEYED,
|
||||||
|
)
|
||||||
|
reduce_stage = StageSpec(
|
||||||
|
stage_id="reduce",
|
||||||
|
kind=StageKind.REDUCE,
|
||||||
|
entry_point=self.REDUCE_ENTRY_POINT,
|
||||||
|
needs=("map",),
|
||||||
|
inputs={"partials": reduce_input},
|
||||||
|
outputs={"result": output_port},
|
||||||
|
parameter_names=reduce_parameter_names,
|
||||||
|
resources=resources,
|
||||||
|
execution=execution,
|
||||||
|
retry=RetryPolicy(),
|
||||||
|
verifier=ComponentRef("exact-artifact", 1),
|
||||||
|
trust_modes=("trusted",),
|
||||||
|
cacheable=True,
|
||||||
|
)
|
||||||
|
workflow = WorkflowSpec(
|
||||||
|
workflow_id="map-reduce-v1",
|
||||||
|
inputs={"input": 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(workload.name, version),
|
||||||
|
description=workload.description,
|
||||||
|
package=PackageSpec("scimesh", package_digest),
|
||||||
|
environment=EnvironmentSpec("python-process", environment_digest, {"adapter": "distributed-v1"}),
|
||||||
|
parameters_schema=parameters_schema,
|
||||||
|
workflow=workflow,
|
||||||
|
inputs={"input": input_port},
|
||||||
|
outputs={"result": output_port},
|
||||||
|
determinism=DeterminismProfile.BYTE_EXACT,
|
||||||
|
trust_modes=(TrustMode.TRUSTED,),
|
||||||
|
verifier=VerifierSpec(ComponentRef("exact-artifact", 1), {}),
|
||||||
|
limits=limits,
|
||||||
|
capabilities=(workload.name,),
|
||||||
|
conformance_profiles=("core-batch-v1",),
|
||||||
|
)
|
||||||
|
self._exact_verifier = ExactArtifactVerifier()
|
||||||
|
|
||||||
|
def definition(self) -> WorkloadDefinition:
|
||||||
|
return WorkloadDefinition(
|
||||||
|
manifest=self.manifest,
|
||||||
|
planner=self,
|
||||||
|
runners={self.MAP_ENTRY_POINT: self},
|
||||||
|
reducers={self.REDUCE_ENTRY_POINT: self},
|
||||||
|
verifiers={self._exact_verifier.identity.canonical: self._exact_verifier},
|
||||||
|
)
|
||||||
|
|
||||||
|
def validate(self, request: JobRequest) -> ValidatedJob:
|
||||||
|
if request.workload != self.manifest.workload:
|
||||||
|
raise ValueError("legacy adapter received a request for another workload")
|
||||||
|
self.workload.validate_job(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("legacy adapter 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)
|
||||||
|
legacy = self.workload.plan(
|
||||||
|
input_path,
|
||||||
|
input_artifact.artifact_id,
|
||||||
|
job.request.parameters,
|
||||||
|
self.shard_rows,
|
||||||
|
workspace,
|
||||||
|
)
|
||||||
|
if legacy.workload != self.workload.name:
|
||||||
|
raise ValueError("legacy planner returned a plan for another workload")
|
||||||
|
tasks: list[TaskSpec] = []
|
||||||
|
negotiated = context.negotiated
|
||||||
|
map_stage = self.manifest.workflow.stages[0]
|
||||||
|
used_paths: set[Path] = set()
|
||||||
|
for planned in legacy.tasks:
|
||||||
|
path = self._find_planned_file(workspace, planned.input_artifact.sha256, used_paths)
|
||||||
|
sealed = context.sink.seal(
|
||||||
|
path,
|
||||||
|
declaration=self.input_port.schema,
|
||||||
|
)
|
||||||
|
if sealed.sha256 != planned.input_artifact.sha256:
|
||||||
|
raise ValueError("artifact sink returned a checksum that differs from the legacy plan")
|
||||||
|
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/{planned.chunk_index:08d}",
|
||||||
|
stage_id="map",
|
||||||
|
parameters=planned.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=legacy.resolved_parameters,
|
||||||
|
tasks=tuple(tasks),
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _find_planned_file(workspace: Path, expected_sha256: str, used: set[Path]) -> Path:
|
||||||
|
for candidate in sorted(workspace.rglob("*")):
|
||||||
|
if candidate in used or not candidate.is_file() or candidate.is_symlink():
|
||||||
|
continue
|
||||||
|
if _sha256_file(candidate) == expected_sha256:
|
||||||
|
used.add(candidate)
|
||||||
|
return candidate
|
||||||
|
raise ValueError("legacy planner did not materialize its planned artifact")
|
||||||
|
|
||||||
|
def run(self, context: TaskContext) -> OutputManifest:
|
||||||
|
context.cancellation.raise_if_cancelled()
|
||||||
|
collection = context.task.inputs.get("input")
|
||||||
|
if collection is None:
|
||||||
|
raise ValueError("legacy map task requires one input collection")
|
||||||
|
self.input_port.validate_collection(collection, "legacy 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"
|
||||||
|
if source.resolve() != input_path.resolve():
|
||||||
|
shutil.copyfile(source, input_path)
|
||||||
|
metrics = self.shard_runner(input_path, context.task.parameters, output_path)
|
||||||
|
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("legacy reducer requires a non-empty keyed partial collection")
|
||||||
|
self.manifest.workflow.stages[1].inputs["partials"].validate_collection(
|
||||||
|
collection,
|
||||||
|
"legacy reducer partials",
|
||||||
|
)
|
||||||
|
workspace = context.workspace
|
||||||
|
workspace.mkdir(parents=True, exist_ok=True)
|
||||||
|
partials: list[CompletedPartial] = []
|
||||||
|
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("legacy partial key must use map.<eight-digit-index>")
|
||||||
|
indexed_items.append((int(raw_index), item))
|
||||||
|
indices = [index for index, _ in indexed_items]
|
||||||
|
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("legacy partial keys do not match the coordinator expected set")
|
||||||
|
if sorted(indices) != list(range(len(indexed_items))):
|
||||||
|
raise ValueError("legacy partial keys must be complete and contiguous")
|
||||||
|
for index, item in sorted(indexed_items):
|
||||||
|
artifact = 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")
|
||||||
|
partials.append(
|
||||||
|
CompletedPartial(
|
||||||
|
index,
|
||||||
|
LegacyArtifactReference(
|
||||||
|
artifact.artifact_id,
|
||||||
|
artifact.sha256,
|
||||||
|
artifact.media_type,
|
||||||
|
),
|
||||||
|
{},
|
||||||
|
)
|
||||||
|
)
|
||||||
|
result = self.workload.reduce(partials, context.task.parameters, workspace)
|
||||||
|
if not isinstance(result, FinalResult):
|
||||||
|
raise ValueError("legacy reducer must return a FinalResult")
|
||||||
|
path = self._find_planned_file(workspace, result.artifact.sha256, set())
|
||||||
|
sealed = context.sink.seal(
|
||||||
|
path,
|
||||||
|
declaration=self.output_port.schema,
|
||||||
|
)
|
||||||
|
if sealed.sha256 != result.artifact.sha256:
|
||||||
|
raise ValueError("artifact sink returned a checksum that differs from the legacy result")
|
||||||
|
return OutputManifest(
|
||||||
|
context.task.task_key,
|
||||||
|
{"result": ArtifactCollection.single(sealed)},
|
||||||
|
result.metrics,
|
||||||
|
context.provenance,
|
||||||
|
).validate_against(context.task.expected_outputs, max_output_bytes=self.manifest.limits.max_output_bytes)
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,345 @@
|
|||||||
|
"""Execution, retry, checkpoint, cancellation, and failure declarations."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from enum import Enum
|
||||||
|
from types import MappingProxyType
|
||||||
|
from typing import Any, Mapping
|
||||||
|
|
||||||
|
from ._validation import (
|
||||||
|
enum_value,
|
||||||
|
freeze_json_mapping,
|
||||||
|
require_exact_keys,
|
||||||
|
require_identifier,
|
||||||
|
require_nonnegative_int,
|
||||||
|
require_safe_message,
|
||||||
|
require_positive_int,
|
||||||
|
require_string,
|
||||||
|
thaw_json,
|
||||||
|
)
|
||||||
|
from .identity import SchemaRef
|
||||||
|
from .resources import ResourceAllocation, ResourceRequirements
|
||||||
|
|
||||||
|
|
||||||
|
class ProcessModel(str, Enum):
|
||||||
|
SINGLE = "single"
|
||||||
|
PROCESS_POOL = "process_pool"
|
||||||
|
THREAD_POOL = "thread_pool"
|
||||||
|
EXTERNAL_RUNTIME = "external_runtime"
|
||||||
|
|
||||||
|
|
||||||
|
class NetworkPolicy(str, Enum):
|
||||||
|
NONE = "none"
|
||||||
|
COORDINATOR_ARTIFACTS_ONLY = "coordinator_artifacts_only"
|
||||||
|
ALLOWLISTED_EGRESS = "allowlisted_egress"
|
||||||
|
TRUSTED = "trusted"
|
||||||
|
|
||||||
|
|
||||||
|
class FailureCategory(str, Enum):
|
||||||
|
INPUT = "input"
|
||||||
|
SCIENTIFIC = "scientific"
|
||||||
|
RESOURCE = "resource"
|
||||||
|
INFRASTRUCTURE = "infrastructure"
|
||||||
|
LEASE = "lease"
|
||||||
|
VERIFICATION = "verification"
|
||||||
|
POLICY = "policy"
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class RetryPolicy:
|
||||||
|
max_attempts: int = 1
|
||||||
|
retryable_categories: tuple[FailureCategory, ...] = ()
|
||||||
|
initial_backoff_seconds: int = 1
|
||||||
|
max_backoff_seconds: int = 60
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
object.__setattr__(self, "max_attempts", require_positive_int(self.max_attempts, "retry.max_attempts"))
|
||||||
|
categories = tuple(
|
||||||
|
enum_value(FailureCategory, value, "retryable_category")
|
||||||
|
for value in self.retryable_categories
|
||||||
|
)
|
||||||
|
if len(categories) != len(set(categories)):
|
||||||
|
raise ValueError("retryable_categories must be unique")
|
||||||
|
object.__setattr__(self, "retryable_categories", categories)
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"initial_backoff_seconds",
|
||||||
|
require_nonnegative_int(self.initial_backoff_seconds, "retry.initial_backoff_seconds"),
|
||||||
|
)
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"max_backoff_seconds",
|
||||||
|
require_nonnegative_int(self.max_backoff_seconds, "retry.max_backoff_seconds"),
|
||||||
|
)
|
||||||
|
if self.max_backoff_seconds < self.initial_backoff_seconds:
|
||||||
|
raise ValueError("retry max_backoff_seconds must not be less than initial_backoff_seconds")
|
||||||
|
if self.max_attempts == 1 and categories:
|
||||||
|
raise ValueError("a non-retrying policy must not list retryable categories")
|
||||||
|
|
||||||
|
def to_dict(self) -> dict[str, object]:
|
||||||
|
return {
|
||||||
|
"max_attempts": self.max_attempts,
|
||||||
|
"retryable_categories": [category.value for category in self.retryable_categories],
|
||||||
|
"initial_backoff_seconds": self.initial_backoff_seconds,
|
||||||
|
"max_backoff_seconds": self.max_backoff_seconds,
|
||||||
|
}
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_dict(cls, value: object) -> "RetryPolicy":
|
||||||
|
if not isinstance(value, Mapping):
|
||||||
|
raise ValueError("retry policy must be an object")
|
||||||
|
fields = {"max_attempts", "retryable_categories", "initial_backoff_seconds", "max_backoff_seconds"}
|
||||||
|
require_exact_keys(value, fields, "retry policy")
|
||||||
|
categories = value["retryable_categories"]
|
||||||
|
if not isinstance(categories, list):
|
||||||
|
raise ValueError("retryable_categories must be an array")
|
||||||
|
return cls(
|
||||||
|
max_attempts=value["max_attempts"], # type: ignore[arg-type]
|
||||||
|
retryable_categories=tuple(categories),
|
||||||
|
initial_backoff_seconds=value["initial_backoff_seconds"], # type: ignore[arg-type]
|
||||||
|
max_backoff_seconds=value["max_backoff_seconds"], # type: ignore[arg-type]
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class CheckpointPolicy:
|
||||||
|
enabled: bool = False
|
||||||
|
schema: SchemaRef | None = None
|
||||||
|
compatibility_version: int | None = None
|
||||||
|
interval_seconds: int | None = None
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
if not isinstance(self.enabled, bool):
|
||||||
|
raise ValueError("checkpoint.enabled must be a boolean")
|
||||||
|
if not self.enabled:
|
||||||
|
if any(value is not None for value in (self.schema, self.compatibility_version, self.interval_seconds)):
|
||||||
|
raise ValueError("disabled checkpoint policy must not declare checkpoint fields")
|
||||||
|
return
|
||||||
|
if not isinstance(self.schema, SchemaRef):
|
||||||
|
raise ValueError("enabled checkpoint policy requires a schema")
|
||||||
|
if self.compatibility_version is None:
|
||||||
|
raise ValueError("enabled checkpoint policy requires a compatibility_version")
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"compatibility_version",
|
||||||
|
require_positive_int(self.compatibility_version, "checkpoint.compatibility_version"),
|
||||||
|
)
|
||||||
|
if self.interval_seconds is not None:
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"interval_seconds",
|
||||||
|
require_positive_int(self.interval_seconds, "checkpoint.interval_seconds"),
|
||||||
|
)
|
||||||
|
|
||||||
|
def to_dict(self) -> dict[str, object]:
|
||||||
|
return {
|
||||||
|
"enabled": self.enabled,
|
||||||
|
"schema": self.schema.canonical if self.schema is not None else None,
|
||||||
|
"compatibility_version": self.compatibility_version,
|
||||||
|
"interval_seconds": self.interval_seconds,
|
||||||
|
}
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_dict(cls, value: object) -> "CheckpointPolicy":
|
||||||
|
if not isinstance(value, Mapping):
|
||||||
|
raise ValueError("checkpoint policy must be an object")
|
||||||
|
fields = {"enabled", "schema", "compatibility_version", "interval_seconds"}
|
||||||
|
require_exact_keys(value, fields, "checkpoint policy")
|
||||||
|
raw_schema = value["schema"]
|
||||||
|
return cls(
|
||||||
|
enabled=value["enabled"], # type: ignore[arg-type]
|
||||||
|
schema=None if raw_schema is None else SchemaRef.from_dict(raw_schema),
|
||||||
|
compatibility_version=value["compatibility_version"], # type: ignore[arg-type]
|
||||||
|
interval_seconds=value["interval_seconds"], # type: ignore[arg-type]
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class ExecutionProfile:
|
||||||
|
profile: str
|
||||||
|
process_model: ProcessModel = ProcessModel.SINGLE
|
||||||
|
max_processes: int = 1
|
||||||
|
threads_per_process: int = 1
|
||||||
|
native_threads: int = 1
|
||||||
|
nested_parallelism: bool = False
|
||||||
|
network: NetworkPolicy = NetworkPolicy.NONE
|
||||||
|
timeout_seconds: int = 3600
|
||||||
|
cancellation_grace_seconds: int = 10
|
||||||
|
checkpoint: CheckpointPolicy = CheckpointPolicy()
|
||||||
|
allowed_egress: tuple[str, ...] = ()
|
||||||
|
secret_handles: tuple[str, ...] = ()
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
object.__setattr__(self, "profile", require_identifier(self.profile, "execution.profile"))
|
||||||
|
object.__setattr__(self, "process_model", enum_value(ProcessModel, self.process_model, "process_model"))
|
||||||
|
object.__setattr__(self, "max_processes", require_positive_int(self.max_processes, "max_processes"))
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"threads_per_process",
|
||||||
|
require_positive_int(self.threads_per_process, "threads_per_process"),
|
||||||
|
)
|
||||||
|
object.__setattr__(self, "native_threads", require_positive_int(self.native_threads, "native_threads"))
|
||||||
|
if not isinstance(self.nested_parallelism, bool):
|
||||||
|
raise ValueError("nested_parallelism must be a boolean")
|
||||||
|
object.__setattr__(self, "network", enum_value(NetworkPolicy, self.network, "network"))
|
||||||
|
object.__setattr__(self, "timeout_seconds", require_positive_int(self.timeout_seconds, "timeout_seconds"))
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"cancellation_grace_seconds",
|
||||||
|
require_nonnegative_int(self.cancellation_grace_seconds, "cancellation_grace_seconds"),
|
||||||
|
)
|
||||||
|
if not isinstance(self.checkpoint, CheckpointPolicy):
|
||||||
|
raise ValueError("checkpoint must be a CheckpointPolicy")
|
||||||
|
egress = tuple(require_string(value, "allowed_egress", max_length=253) for value in self.allowed_egress)
|
||||||
|
if len(egress) != len(set(egress)):
|
||||||
|
raise ValueError("allowed_egress must be unique")
|
||||||
|
if self.network is NetworkPolicy.ALLOWLISTED_EGRESS and not egress:
|
||||||
|
raise ValueError("allowlisted egress policy requires at least one target")
|
||||||
|
if self.network is not NetworkPolicy.ALLOWLISTED_EGRESS and egress:
|
||||||
|
raise ValueError("allowed_egress is valid only for allowlisted egress")
|
||||||
|
object.__setattr__(self, "allowed_egress", egress)
|
||||||
|
handles = tuple(require_identifier(value, "secret_handle") for value in self.secret_handles)
|
||||||
|
if len(handles) != len(set(handles)):
|
||||||
|
raise ValueError("secret_handles must be unique")
|
||||||
|
if handles and self.network is NetworkPolicy.NONE:
|
||||||
|
raise ValueError("secret handles require an explicit network policy")
|
||||||
|
object.__setattr__(self, "secret_handles", handles)
|
||||||
|
if self.process_model is ProcessModel.SINGLE and (
|
||||||
|
self.max_processes != 1 or self.threads_per_process != 1
|
||||||
|
):
|
||||||
|
raise ValueError("single process model requires one process and one Python thread")
|
||||||
|
if not self.nested_parallelism and self.threads_per_process > 1 and self.native_threads > 1:
|
||||||
|
raise ValueError("nested thread pools require nested_parallelism=true")
|
||||||
|
|
||||||
|
@property
|
||||||
|
def maximum_cpu_threads(self) -> int:
|
||||||
|
return self.max_processes * self.threads_per_process * self.native_threads
|
||||||
|
|
||||||
|
def validate_resources(self, resources: ResourceRequirements) -> None:
|
||||||
|
if self.maximum_cpu_threads > resources.cpu_cores:
|
||||||
|
raise ValueError("execution profile can oversubscribe its CPU reservation")
|
||||||
|
if self.timeout_seconds > resources.max_duration_seconds:
|
||||||
|
raise ValueError("execution timeout exceeds the resource maximum duration")
|
||||||
|
|
||||||
|
def allocation_environment(self, allocation: ResourceAllocation) -> Mapping[str, str]:
|
||||||
|
"""Return only allocation-derived thread/device isolation variables."""
|
||||||
|
if not isinstance(allocation, ResourceAllocation):
|
||||||
|
raise ValueError("allocation must be a ResourceAllocation")
|
||||||
|
native = str(min(self.native_threads, allocation.cpu_cores))
|
||||||
|
values = {
|
||||||
|
"OMP_NUM_THREADS": native,
|
||||||
|
"OPENBLAS_NUM_THREADS": native,
|
||||||
|
"MKL_NUM_THREADS": native,
|
||||||
|
"NUMEXPR_NUM_THREADS": native,
|
||||||
|
"VECLIB_MAXIMUM_THREADS": native,
|
||||||
|
# Empty visibility explicitly prevents a CPU task from inheriting
|
||||||
|
# access to all host devices.
|
||||||
|
"CUDA_VISIBLE_DEVICES": ",".join(allocation.accelerator_ids),
|
||||||
|
"ROCR_VISIBLE_DEVICES": ",".join(allocation.accelerator_ids),
|
||||||
|
}
|
||||||
|
return MappingProxyType(values)
|
||||||
|
|
||||||
|
def to_dict(self) -> dict[str, object]:
|
||||||
|
return {
|
||||||
|
"profile": self.profile,
|
||||||
|
"process_model": self.process_model.value,
|
||||||
|
"max_processes": self.max_processes,
|
||||||
|
"threads_per_process": self.threads_per_process,
|
||||||
|
"native_threads": self.native_threads,
|
||||||
|
"nested_parallelism": self.nested_parallelism,
|
||||||
|
"network": self.network.value,
|
||||||
|
"timeout_seconds": self.timeout_seconds,
|
||||||
|
"cancellation_grace_seconds": self.cancellation_grace_seconds,
|
||||||
|
"checkpoint": self.checkpoint.to_dict(),
|
||||||
|
"allowed_egress": list(self.allowed_egress),
|
||||||
|
"secret_handles": list(self.secret_handles),
|
||||||
|
}
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_dict(cls, value: object) -> "ExecutionProfile":
|
||||||
|
if not isinstance(value, Mapping):
|
||||||
|
raise ValueError("execution profile must be an object")
|
||||||
|
fields = {
|
||||||
|
"profile", "process_model", "max_processes", "threads_per_process",
|
||||||
|
"native_threads", "nested_parallelism", "network", "timeout_seconds",
|
||||||
|
"cancellation_grace_seconds", "checkpoint", "allowed_egress", "secret_handles",
|
||||||
|
}
|
||||||
|
require_exact_keys(value, fields, "execution profile")
|
||||||
|
allowed_egress = value["allowed_egress"]
|
||||||
|
secret_handles = value["secret_handles"]
|
||||||
|
if not isinstance(allowed_egress, list) or not isinstance(secret_handles, list):
|
||||||
|
raise ValueError("execution allowed_egress and secret_handles must be arrays")
|
||||||
|
return cls(
|
||||||
|
profile=value["profile"], # type: ignore[arg-type]
|
||||||
|
process_model=value["process_model"], # type: ignore[arg-type]
|
||||||
|
max_processes=value["max_processes"], # type: ignore[arg-type]
|
||||||
|
threads_per_process=value["threads_per_process"], # type: ignore[arg-type]
|
||||||
|
native_threads=value["native_threads"], # type: ignore[arg-type]
|
||||||
|
nested_parallelism=value["nested_parallelism"], # type: ignore[arg-type]
|
||||||
|
network=value["network"], # type: ignore[arg-type]
|
||||||
|
timeout_seconds=value["timeout_seconds"], # type: ignore[arg-type]
|
||||||
|
cancellation_grace_seconds=value["cancellation_grace_seconds"], # type: ignore[arg-type]
|
||||||
|
checkpoint=CheckpointPolicy.from_dict(value["checkpoint"]),
|
||||||
|
allowed_egress=tuple(allowed_egress),
|
||||||
|
secret_handles=tuple(secret_handles),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class FailureReport:
|
||||||
|
code: str
|
||||||
|
category: FailureCategory
|
||||||
|
retryable: bool
|
||||||
|
message: str
|
||||||
|
evidence: Mapping[str, Any]
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
object.__setattr__(self, "code", require_identifier(self.code, "failure.code"))
|
||||||
|
object.__setattr__(self, "category", enum_value(FailureCategory, self.category, "failure.category"))
|
||||||
|
if not isinstance(self.retryable, bool):
|
||||||
|
raise ValueError("failure.retryable must be a boolean")
|
||||||
|
object.__setattr__(self, "message", require_safe_message(self.message, "failure.message", max_length=512))
|
||||||
|
evidence = freeze_json_mapping(self.evidence, "failure.evidence", forbid_locations=True)
|
||||||
|
import json
|
||||||
|
if len(json.dumps(thaw_json(evidence), allow_nan=False).encode("utf-8")) > 16_384:
|
||||||
|
raise ValueError("failure evidence exceeds 16 KiB")
|
||||||
|
object.__setattr__(self, "evidence", evidence)
|
||||||
|
|
||||||
|
def to_dict(self) -> dict[str, object]:
|
||||||
|
return {
|
||||||
|
"code": self.code,
|
||||||
|
"category": self.category.value,
|
||||||
|
"retryable": self.retryable,
|
||||||
|
"message": self.message,
|
||||||
|
"evidence": thaw_json(self.evidence),
|
||||||
|
}
|
||||||
|
|
||||||
|
def to_json(self) -> str:
|
||||||
|
from ._validation import canonical_json
|
||||||
|
|
||||||
|
return canonical_json(self.to_dict())
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_dict(cls, value: object) -> "FailureReport":
|
||||||
|
if not isinstance(value, Mapping):
|
||||||
|
raise ValueError("failure report must be an object")
|
||||||
|
fields = {"code", "category", "retryable", "message", "evidence"}
|
||||||
|
require_exact_keys(value, fields, "failure report")
|
||||||
|
return cls(
|
||||||
|
code=value["code"], # type: ignore[arg-type]
|
||||||
|
category=value["category"], # type: ignore[arg-type]
|
||||||
|
retryable=value["retryable"], # type: ignore[arg-type]
|
||||||
|
message=value["message"], # type: ignore[arg-type]
|
||||||
|
evidence=value["evidence"], # type: ignore[arg-type]
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_json(cls, value: str) -> "FailureReport":
|
||||||
|
import json
|
||||||
|
|
||||||
|
try:
|
||||||
|
decoded = json.loads(value)
|
||||||
|
except (TypeError, json.JSONDecodeError, RecursionError) as error:
|
||||||
|
raise ValueError("failure report must be valid JSON") from error
|
||||||
|
return cls.from_dict(decoded)
|
||||||
@@ -0,0 +1,169 @@
|
|||||||
|
"""Versioned identities used across the SciMesh workload SDK."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import Mapping
|
||||||
|
|
||||||
|
from ._validation import (
|
||||||
|
require_exact_keys,
|
||||||
|
require_identifier,
|
||||||
|
require_semver,
|
||||||
|
require_string,
|
||||||
|
require_workload_name,
|
||||||
|
validate_version_range,
|
||||||
|
version_in_range,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
SDK_API_VERSION = "1.0.0"
|
||||||
|
MANIFEST_SCHEMA_VERSION = 1
|
||||||
|
WORKFLOW_SCHEMA_VERSION = 1
|
||||||
|
TASK_SCHEMA_VERSION = 1
|
||||||
|
OUTPUT_SCHEMA_VERSION = 1
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class VersionRange:
|
||||||
|
"""A deliberately small, explicit compatibility range.
|
||||||
|
|
||||||
|
The v1 SDK accepts comma-separated comparisons such as ``>=1.0,<2.0``.
|
||||||
|
Wildcards and an omitted operator are rejected so a missing version can
|
||||||
|
never be interpreted as "latest".
|
||||||
|
"""
|
||||||
|
|
||||||
|
expression: str
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
object.__setattr__(self, "expression", validate_version_range(self.expression, "version range"))
|
||||||
|
|
||||||
|
def contains(self, version: str) -> bool:
|
||||||
|
return version_in_range(version, self.expression)
|
||||||
|
|
||||||
|
def to_dict(self) -> str:
|
||||||
|
return self.expression
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_dict(cls, value: object) -> "VersionRange":
|
||||||
|
return cls(value) # type: ignore[arg-type]
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class WorkloadId:
|
||||||
|
name: str
|
||||||
|
version: str
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
object.__setattr__(self, "name", require_workload_name(self.name))
|
||||||
|
object.__setattr__(self, "version", require_semver(self.version, "workload.version"))
|
||||||
|
|
||||||
|
def to_dict(self) -> dict[str, str]:
|
||||||
|
return {"name": self.name, "version": self.version}
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_dict(cls, value: object) -> "WorkloadId":
|
||||||
|
if not isinstance(value, Mapping):
|
||||||
|
raise ValueError("workload identity must be an object")
|
||||||
|
require_exact_keys(value, {"name", "version"}, "workload identity")
|
||||||
|
return cls(name=value["name"], version=value["version"]) # type: ignore[arg-type]
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class SchemaRef:
|
||||||
|
name: str
|
||||||
|
version: int
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
object.__setattr__(self, "name", require_identifier(self.name, "schema.name"))
|
||||||
|
if isinstance(self.version, bool) or not isinstance(self.version, int) or self.version < 1:
|
||||||
|
raise ValueError("schema.version must be a positive integer")
|
||||||
|
|
||||||
|
@property
|
||||||
|
def canonical(self) -> str:
|
||||||
|
return f"{self.name}@{self.version}"
|
||||||
|
|
||||||
|
def to_dict(self) -> dict[str, object]:
|
||||||
|
return {"name": self.name, "version": self.version}
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def parse(cls, value: object, field: str = "schema") -> "SchemaRef":
|
||||||
|
text = require_string(value, field, max_length=160)
|
||||||
|
name, separator, raw_version = text.rpartition("@")
|
||||||
|
if not separator or not raw_version.isdigit():
|
||||||
|
raise ValueError(f"{field} must use the name@version form")
|
||||||
|
return cls(name=name, version=int(raw_version))
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_dict(cls, value: object) -> "SchemaRef":
|
||||||
|
if isinstance(value, str):
|
||||||
|
return cls.parse(value)
|
||||||
|
if not isinstance(value, Mapping):
|
||||||
|
raise ValueError("schema reference must be a name@version string or object")
|
||||||
|
require_exact_keys(value, {"name", "version"}, "schema reference")
|
||||||
|
return cls(name=value["name"], version=value["version"]) # type: ignore[arg-type]
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class ComponentRef:
|
||||||
|
"""Versioned, package-owned planner/runner/reducer/verifier identity."""
|
||||||
|
|
||||||
|
name: str
|
||||||
|
version: int
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
object.__setattr__(self, "name", require_identifier(self.name, "component.name"))
|
||||||
|
if isinstance(self.version, bool) or not isinstance(self.version, int) or self.version < 1:
|
||||||
|
raise ValueError("component.version must be a positive integer")
|
||||||
|
|
||||||
|
@property
|
||||||
|
def canonical(self) -> str:
|
||||||
|
return f"{self.name}@{self.version}"
|
||||||
|
|
||||||
|
def to_dict(self) -> dict[str, object]:
|
||||||
|
return {"name": self.name, "version": self.version}
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_dict(cls, value: object) -> "ComponentRef":
|
||||||
|
if isinstance(value, str):
|
||||||
|
parsed = SchemaRef.parse(value, "component")
|
||||||
|
return cls(parsed.name, parsed.version)
|
||||||
|
if not isinstance(value, Mapping):
|
||||||
|
raise ValueError("component reference must be an object")
|
||||||
|
require_exact_keys(value, {"name", "version"}, "component reference")
|
||||||
|
return cls(name=value["name"], version=value["version"]) # type: ignore[arg-type]
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class FeatureRequirement:
|
||||||
|
name: str
|
||||||
|
versions: VersionRange
|
||||||
|
fallback: str | None = None
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
object.__setattr__(self, "name", require_identifier(self.name, "feature.name"))
|
||||||
|
if not isinstance(self.versions, VersionRange):
|
||||||
|
raise ValueError("feature.versions must be a VersionRange")
|
||||||
|
if self.fallback is not None:
|
||||||
|
object.__setattr__(self, "fallback", require_identifier(self.fallback, "feature.fallback"))
|
||||||
|
|
||||||
|
def to_dict(self) -> dict[str, object]:
|
||||||
|
result: dict[str, object] = {"name": self.name, "versions": self.versions.expression}
|
||||||
|
if self.fallback is not None:
|
||||||
|
result["fallback"] = self.fallback
|
||||||
|
return result
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_dict(cls, value: object) -> "FeatureRequirement":
|
||||||
|
if not isinstance(value, Mapping):
|
||||||
|
raise ValueError("feature requirement must be an object")
|
||||||
|
require_exact_keys(
|
||||||
|
value,
|
||||||
|
{"name", "versions"},
|
||||||
|
"feature requirement",
|
||||||
|
optional={"fallback"},
|
||||||
|
)
|
||||||
|
return cls(
|
||||||
|
name=value["name"], # type: ignore[arg-type]
|
||||||
|
versions=VersionRange.from_dict(value["versions"]),
|
||||||
|
fallback=value.get("fallback"), # type: ignore[arg-type]
|
||||||
|
)
|
||||||
@@ -0,0 +1,140 @@
|
|||||||
|
"""Independent installed-distribution content measurement for SDK allowlists."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import hashlib
|
||||||
|
import importlib.util
|
||||||
|
from importlib import metadata
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
|
||||||
|
def installed_distribution_digest(
|
||||||
|
distribution: metadata.Distribution | str,
|
||||||
|
*,
|
||||||
|
allow_editable: bool = False,
|
||||||
|
) -> str:
|
||||||
|
"""Hash installed package payload files using a stable path/length framing.
|
||||||
|
|
||||||
|
Distribution metadata is deliberately excluded: editable/non-editable
|
||||||
|
installers generate different RECORD and entry-point files for identical
|
||||||
|
package code. All source/native modules and package data below declared
|
||||||
|
top-level packages are included. Interpreter-generated ``__pycache__``
|
||||||
|
files are excluded because they are neither stable wheel payloads nor used
|
||||||
|
by the registry's cache-isolated discovery import.
|
||||||
|
"""
|
||||||
|
installed = metadata.distribution(distribution) if isinstance(distribution, str) else distribution
|
||||||
|
raw_top_level = installed.read_text("top_level.txt")
|
||||||
|
if raw_top_level is None:
|
||||||
|
raise ValueError("installed distribution does not declare top-level packages")
|
||||||
|
declared_top_levels = [line.strip() for line in raw_top_level.splitlines() if line.strip()]
|
||||||
|
if any(not value.isidentifier() for value in declared_top_levels):
|
||||||
|
raise ValueError("installed distribution declares an invalid top-level package")
|
||||||
|
top_levels = set(declared_top_levels)
|
||||||
|
if not top_levels:
|
||||||
|
raise ValueError("installed distribution has no measurable top-level package")
|
||||||
|
declared_files = tuple(installed.files or ())
|
||||||
|
editable_bootstrap = any(
|
||||||
|
Path(str(item)).name.startswith("__editable__") and Path(str(item)).suffix == ".pth"
|
||||||
|
for item in declared_files
|
||||||
|
)
|
||||||
|
if editable_bootstrap and not allow_editable:
|
||||||
|
raise ValueError("editable workload installations are not accepted for secure discovery")
|
||||||
|
for item in declared_files:
|
||||||
|
relative = Path(str(item))
|
||||||
|
suffix = relative.suffix.lower()
|
||||||
|
if suffix == ".pth" and not allow_editable:
|
||||||
|
raise ValueError("installed workload distribution declares a .pth bootstrap")
|
||||||
|
if suffix in {".pyc", ".pyo"} and "__pycache__" not in relative.parts:
|
||||||
|
raise ValueError("installed workload distribution declares sourceless bytecode")
|
||||||
|
|
||||||
|
selected: list[tuple[str, Path]] = []
|
||||||
|
for top_level in sorted(top_levels):
|
||||||
|
root = Path(installed.locate_file(top_level))
|
||||||
|
if not root.exists():
|
||||||
|
# PEP 660 editable distributions may expose source packages through
|
||||||
|
# a meta-path finder rather than a physical site-packages path.
|
||||||
|
spec = importlib.util.find_spec(top_level)
|
||||||
|
locations = tuple(spec.submodule_search_locations or ()) if spec is not None else ()
|
||||||
|
if len(locations) > 1:
|
||||||
|
raise ValueError("shared namespace packages are not supported for workload integrity")
|
||||||
|
if locations:
|
||||||
|
root = Path(locations[0])
|
||||||
|
if root.is_symlink():
|
||||||
|
raise ValueError("installed workload package root must not be a symbolic link")
|
||||||
|
if root.is_dir():
|
||||||
|
candidates = root.rglob("*")
|
||||||
|
for path in candidates:
|
||||||
|
if path.is_symlink():
|
||||||
|
raise ValueError("installed workload package contains a symbolic-link payload")
|
||||||
|
if not path.is_file():
|
||||||
|
continue
|
||||||
|
relative_parts = path.relative_to(root).parts
|
||||||
|
if "__pycache__" in relative_parts:
|
||||||
|
continue
|
||||||
|
if path.suffix.lower() in {".pyc", ".pyo"}:
|
||||||
|
raise ValueError("installed workload package contains sourceless bytecode")
|
||||||
|
relative = f"{top_level}/{path.relative_to(root).as_posix()}"
|
||||||
|
selected.append((relative, path))
|
||||||
|
continue
|
||||||
|
module = Path(installed.locate_file(top_level + ".py"))
|
||||||
|
if not module.exists():
|
||||||
|
spec = importlib.util.find_spec(top_level)
|
||||||
|
if spec is not None and spec.origin is not None:
|
||||||
|
module = Path(spec.origin)
|
||||||
|
if module.is_symlink() or not module.is_file():
|
||||||
|
raise ValueError("installed workload package contains a missing top-level payload")
|
||||||
|
selected.append((top_level + ".py", module))
|
||||||
|
# Include declared package data outside top-level import trees. Generated
|
||||||
|
# console wrappers and installer metadata are excluded; executable .pth and
|
||||||
|
# sourceless bytecode payloads were rejected above. Generated pycache
|
||||||
|
# entries are deliberately ignored and discovery imports from an empty
|
||||||
|
# cache prefix.
|
||||||
|
selected_names = {relative for relative, _ in selected}
|
||||||
|
metadata_root_names = {
|
||||||
|
Path(str(item)).parts[0]
|
||||||
|
for item in declared_files
|
||||||
|
if Path(str(item)).parts
|
||||||
|
and Path(str(item)).parts[0].endswith((".dist-info", ".egg-info"))
|
||||||
|
}
|
||||||
|
for item in declared_files:
|
||||||
|
relative = Path(str(item))
|
||||||
|
text = relative.as_posix()
|
||||||
|
if (
|
||||||
|
not relative.parts
|
||||||
|
or relative.parts[0] in metadata_root_names
|
||||||
|
or text.startswith("../../../bin/")
|
||||||
|
or "__pycache__" in relative.parts
|
||||||
|
):
|
||||||
|
continue
|
||||||
|
path = Path(installed.locate_file(item))
|
||||||
|
if path.is_symlink():
|
||||||
|
raise ValueError("installed workload distribution contains a symbolic-link payload")
|
||||||
|
if not path.is_file() or text in selected_names:
|
||||||
|
continue
|
||||||
|
selected.append((text, path))
|
||||||
|
selected_names.add(text)
|
||||||
|
|
||||||
|
entry_point_payloads = [
|
||||||
|
(
|
||||||
|
f".entry-points/{entry_point.group}/{entry_point.name}",
|
||||||
|
entry_point.value.encode("utf-8"),
|
||||||
|
)
|
||||||
|
for entry_point in installed.entry_points
|
||||||
|
]
|
||||||
|
if not selected:
|
||||||
|
raise ValueError("installed distribution has no measurable package payload")
|
||||||
|
digest = hashlib.sha256()
|
||||||
|
for relative, path in sorted(selected):
|
||||||
|
name = relative.encode("utf-8")
|
||||||
|
payload = path.read_bytes()
|
||||||
|
digest.update(len(name).to_bytes(4, "big"))
|
||||||
|
digest.update(name)
|
||||||
|
digest.update(len(payload).to_bytes(8, "big"))
|
||||||
|
digest.update(payload)
|
||||||
|
for relative, payload in sorted(entry_point_payloads):
|
||||||
|
name = relative.encode("utf-8")
|
||||||
|
digest.update(len(name).to_bytes(4, "big"))
|
||||||
|
digest.update(name)
|
||||||
|
digest.update(len(payload).to_bytes(8, "big"))
|
||||||
|
digest.update(payload)
|
||||||
|
return "sha256:" + digest.hexdigest()
|
||||||
@@ -0,0 +1,408 @@
|
|||||||
|
"""Installed-package manifest and cross-component compatibility contract."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
import re
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from enum import Enum
|
||||||
|
from types import MappingProxyType
|
||||||
|
from typing import Any, Mapping
|
||||||
|
|
||||||
|
from ._validation import (
|
||||||
|
canonical_json,
|
||||||
|
enum_value,
|
||||||
|
freeze_json_mapping,
|
||||||
|
require_exact_keys,
|
||||||
|
require_identifier,
|
||||||
|
require_positive_int,
|
||||||
|
require_sha256,
|
||||||
|
require_schema_version,
|
||||||
|
require_string,
|
||||||
|
thaw_json,
|
||||||
|
)
|
||||||
|
from .artifacts import PortSpec
|
||||||
|
from .identity import (
|
||||||
|
MANIFEST_SCHEMA_VERSION,
|
||||||
|
ComponentRef,
|
||||||
|
FeatureRequirement,
|
||||||
|
VersionRange,
|
||||||
|
WorkloadId,
|
||||||
|
)
|
||||||
|
from .workflow import StageKind, WorkflowSpec
|
||||||
|
from .schema import validate_schema_definition
|
||||||
|
|
||||||
|
|
||||||
|
class DeterminismProfile(str, Enum):
|
||||||
|
BYTE_EXACT = "byte_exact"
|
||||||
|
CANONICAL_EXACT = "canonical_exact"
|
||||||
|
NUMERIC_TOLERANCE = "numeric_tolerance"
|
||||||
|
SEEDED_STOCHASTIC = "seeded_stochastic"
|
||||||
|
SEARCH_OR_OPTIMIZATION = "search_or_optimization"
|
||||||
|
SIDE_EFFECTING = "side_effecting"
|
||||||
|
|
||||||
|
|
||||||
|
class TrustMode(str, Enum):
|
||||||
|
TRUSTED = "trusted"
|
||||||
|
VERIFIED = "verified"
|
||||||
|
UNTRUSTED_QUORUM = "untrusted_quorum"
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class PackageSpec:
|
||||||
|
distribution: str
|
||||||
|
digest: str
|
||||||
|
signature: str | None = None
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
distribution = require_string(self.distribution, "package.distribution", max_length=128).lower()
|
||||||
|
if not re.fullmatch(r"[a-z0-9]+(?:[-_.][a-z0-9]+)*", distribution):
|
||||||
|
raise ValueError("package.distribution must be a canonical Python distribution name")
|
||||||
|
object.__setattr__(self, "distribution", distribution.replace("_", "-"))
|
||||||
|
object.__setattr__(self, "digest", require_sha256(self.digest, "package.digest", prefixed=True))
|
||||||
|
if self.signature is not None:
|
||||||
|
object.__setattr__(self, "signature", require_string(self.signature, "package.signature", max_length=512))
|
||||||
|
|
||||||
|
def to_dict(self) -> dict[str, object]:
|
||||||
|
return {"distribution": self.distribution, "digest": self.digest, "signature": self.signature}
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_dict(cls, value: object) -> "PackageSpec":
|
||||||
|
if not isinstance(value, Mapping):
|
||||||
|
raise ValueError("package specification must be an object")
|
||||||
|
require_exact_keys(value, {"distribution", "digest", "signature"}, "package specification")
|
||||||
|
return cls(
|
||||||
|
distribution=value["distribution"], # type: ignore[arg-type]
|
||||||
|
digest=value["digest"], # type: ignore[arg-type]
|
||||||
|
signature=value["signature"], # type: ignore[arg-type]
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class EnvironmentSpec:
|
||||||
|
kind: str
|
||||||
|
digest: str
|
||||||
|
metadata: Mapping[str, Any]
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
object.__setattr__(self, "kind", require_identifier(self.kind, "environment.kind"))
|
||||||
|
object.__setattr__(self, "digest", require_sha256(self.digest, "environment.digest", prefixed=True))
|
||||||
|
object.__setattr__(self, "metadata", freeze_json_mapping(self.metadata, "environment.metadata"))
|
||||||
|
|
||||||
|
def to_dict(self) -> dict[str, object]:
|
||||||
|
return {"kind": self.kind, "digest": self.digest, "metadata": thaw_json(self.metadata)}
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_dict(cls, value: object) -> "EnvironmentSpec":
|
||||||
|
if not isinstance(value, Mapping):
|
||||||
|
raise ValueError("environment specification must be an object")
|
||||||
|
require_exact_keys(value, {"kind", "digest", "metadata"}, "environment specification")
|
||||||
|
return cls(
|
||||||
|
kind=value["kind"], # type: ignore[arg-type]
|
||||||
|
digest=value["digest"], # type: ignore[arg-type]
|
||||||
|
metadata=value["metadata"], # type: ignore[arg-type]
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class VerifierSpec:
|
||||||
|
verifier: ComponentRef
|
||||||
|
configuration: Mapping[str, Any]
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
if not isinstance(self.verifier, ComponentRef):
|
||||||
|
raise ValueError("verifier must be a ComponentRef")
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"configuration",
|
||||||
|
freeze_json_mapping(self.configuration, "verifier.configuration"),
|
||||||
|
)
|
||||||
|
|
||||||
|
def to_dict(self) -> dict[str, object]:
|
||||||
|
return {
|
||||||
|
"verifier": self.verifier.canonical,
|
||||||
|
"configuration": thaw_json(self.configuration),
|
||||||
|
}
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_dict(cls, value: object) -> "VerifierSpec":
|
||||||
|
if not isinstance(value, Mapping):
|
||||||
|
raise ValueError("verifier specification must be an object")
|
||||||
|
require_exact_keys(value, {"verifier", "configuration"}, "verifier specification")
|
||||||
|
return cls(
|
||||||
|
verifier=ComponentRef.from_dict(value["verifier"]),
|
||||||
|
configuration=value["configuration"], # type: ignore[arg-type]
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class WorkloadLimits:
|
||||||
|
max_input_bytes: int
|
||||||
|
max_tasks: int
|
||||||
|
max_output_bytes: int
|
||||||
|
max_parameter_bytes: int = 65_536
|
||||||
|
max_artifacts: int = 100_000
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
for field in (
|
||||||
|
"max_input_bytes", "max_tasks", "max_output_bytes", "max_parameter_bytes", "max_artifacts"
|
||||||
|
):
|
||||||
|
object.__setattr__(self, field, require_positive_int(getattr(self, field), f"limits.{field}"))
|
||||||
|
|
||||||
|
def to_dict(self) -> dict[str, int]:
|
||||||
|
return {
|
||||||
|
"max_input_bytes": self.max_input_bytes,
|
||||||
|
"max_tasks": self.max_tasks,
|
||||||
|
"max_output_bytes": self.max_output_bytes,
|
||||||
|
"max_parameter_bytes": self.max_parameter_bytes,
|
||||||
|
"max_artifacts": self.max_artifacts,
|
||||||
|
}
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_dict(cls, value: object) -> "WorkloadLimits":
|
||||||
|
if not isinstance(value, Mapping):
|
||||||
|
raise ValueError("workload limits must be an object")
|
||||||
|
fields = {
|
||||||
|
"max_input_bytes", "max_tasks", "max_output_bytes", "max_parameter_bytes", "max_artifacts",
|
||||||
|
}
|
||||||
|
require_exact_keys(value, fields, "workload limits")
|
||||||
|
return cls(**value) # type: ignore[arg-type]
|
||||||
|
|
||||||
|
|
||||||
|
def _ports(
|
||||||
|
value: Mapping[str, PortSpec], field: str, *, allow_empty: bool = False
|
||||||
|
) -> Mapping[str, PortSpec]:
|
||||||
|
if not isinstance(value, Mapping) or (not value and not allow_empty):
|
||||||
|
qualifier = "an object" if allow_empty else "a non-empty object"
|
||||||
|
raise ValueError(f"{field} must be {qualifier}")
|
||||||
|
result: dict[str, PortSpec] = {}
|
||||||
|
for name, port in value.items():
|
||||||
|
canonical = require_identifier(name, f"{field} port")
|
||||||
|
if not isinstance(port, PortSpec):
|
||||||
|
raise ValueError(f"{field} values must be PortSpec values")
|
||||||
|
result[canonical] = port
|
||||||
|
return MappingProxyType(result)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class WorkloadManifest:
|
||||||
|
sdk_api: VersionRange
|
||||||
|
protocol: VersionRange
|
||||||
|
workload: WorkloadId
|
||||||
|
description: str
|
||||||
|
package: PackageSpec
|
||||||
|
environment: EnvironmentSpec
|
||||||
|
parameters_schema: Mapping[str, Any]
|
||||||
|
workflow: WorkflowSpec
|
||||||
|
inputs: Mapping[str, PortSpec]
|
||||||
|
outputs: Mapping[str, PortSpec]
|
||||||
|
determinism: DeterminismProfile
|
||||||
|
trust_modes: tuple[TrustMode, ...]
|
||||||
|
verifier: VerifierSpec
|
||||||
|
limits: WorkloadLimits
|
||||||
|
capabilities: tuple[str, ...]
|
||||||
|
conformance_profiles: tuple[str, ...]
|
||||||
|
required_features: tuple[FeatureRequirement, ...] = ()
|
||||||
|
optional_features: tuple[FeatureRequirement, ...] = ()
|
||||||
|
manifest_schema_version: int = MANIFEST_SCHEMA_VERSION
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
require_schema_version(
|
||||||
|
self.manifest_schema_version,
|
||||||
|
MANIFEST_SCHEMA_VERSION,
|
||||||
|
"manifest_schema_version",
|
||||||
|
)
|
||||||
|
if not isinstance(self.sdk_api, VersionRange) or not isinstance(self.protocol, VersionRange):
|
||||||
|
raise ValueError("sdk_api and protocol must be explicit VersionRange values")
|
||||||
|
if not isinstance(self.workload, WorkloadId):
|
||||||
|
raise ValueError("workload must be a WorkloadId")
|
||||||
|
object.__setattr__(self, "description", require_string(self.description, "description", max_length=512))
|
||||||
|
if not isinstance(self.package, PackageSpec) or not isinstance(self.environment, EnvironmentSpec):
|
||||||
|
raise ValueError("manifest package and environment declarations are required")
|
||||||
|
schema = freeze_json_mapping(self.parameters_schema, "parameters_schema")
|
||||||
|
if schema.get("type") != "object" or schema.get("additionalProperties") is not False:
|
||||||
|
raise ValueError("parameters_schema must be an object schema with additionalProperties=false")
|
||||||
|
properties = schema.get("properties")
|
||||||
|
if not isinstance(properties, Mapping):
|
||||||
|
raise ValueError("parameters_schema.properties must be an object")
|
||||||
|
if len(canonical_json(schema).encode("utf-8")) > 1_048_576:
|
||||||
|
raise ValueError("parameters_schema exceeds 1 MiB")
|
||||||
|
validate_schema_definition(schema)
|
||||||
|
object.__setattr__(self, "parameters_schema", schema)
|
||||||
|
if not isinstance(self.workflow, WorkflowSpec):
|
||||||
|
raise ValueError("workflow must be a WorkflowSpec")
|
||||||
|
object.__setattr__(self, "inputs", _ports(self.inputs, "manifest.inputs", allow_empty=True))
|
||||||
|
object.__setattr__(self, "outputs", _ports(self.outputs, "manifest.outputs"))
|
||||||
|
if dict(self.inputs) != dict(self.workflow.inputs):
|
||||||
|
raise ValueError("manifest inputs must match workflow inputs")
|
||||||
|
if dict(self.outputs) != dict(self.workflow.output_ports()):
|
||||||
|
raise ValueError("manifest outputs must match workflow outputs")
|
||||||
|
object.__setattr__(self, "determinism", enum_value(DeterminismProfile, self.determinism, "determinism"))
|
||||||
|
modes = tuple(enum_value(TrustMode, mode, "trust_mode") for mode in self.trust_modes)
|
||||||
|
if not modes or len(modes) != len(set(modes)):
|
||||||
|
raise ValueError("trust_modes must be non-empty and unique")
|
||||||
|
object.__setattr__(self, "trust_modes", modes)
|
||||||
|
manifest_mode_values = {mode.value for mode in modes}
|
||||||
|
terminal_stage_ids = {
|
||||||
|
reference.stage_id
|
||||||
|
for reference in self.workflow.outputs.values()
|
||||||
|
if reference.stage_id is not None
|
||||||
|
}
|
||||||
|
for stage in self.workflow.stages:
|
||||||
|
if not set(stage.trust_modes).issubset(manifest_mode_values):
|
||||||
|
raise ValueError("stage trust modes must be a subset of manifest trust_modes")
|
||||||
|
if stage.verifier is None:
|
||||||
|
raise ValueError("every output-producing stage requires an acceptance verifier")
|
||||||
|
resource_sets = (stage.resources,) + (
|
||||||
|
(stage.gang.per_replica_resources,) if stage.gang is not None else ()
|
||||||
|
)
|
||||||
|
if any(
|
||||||
|
resources.environment_digest not in {None, self.environment.digest}
|
||||||
|
for resources in resource_sets
|
||||||
|
):
|
||||||
|
raise ValueError(
|
||||||
|
"stage resource environment must match the manifest environment pin"
|
||||||
|
)
|
||||||
|
if not isinstance(self.verifier, VerifierSpec):
|
||||||
|
raise ValueError("verifier must be a VerifierSpec")
|
||||||
|
for stage in self.workflow.stages:
|
||||||
|
if (
|
||||||
|
stage.stage_id in terminal_stage_ids
|
||||||
|
and stage.verifier != self.verifier.verifier
|
||||||
|
):
|
||||||
|
raise ValueError(
|
||||||
|
"terminal stage verifier must match the manifest acceptance verifier"
|
||||||
|
)
|
||||||
|
if not isinstance(self.limits, WorkloadLimits):
|
||||||
|
raise ValueError("limits must be WorkloadLimits")
|
||||||
|
if self.workflow.max_tasks > self.limits.max_tasks:
|
||||||
|
raise ValueError("workflow max_tasks exceeds the workload limit")
|
||||||
|
if self.workflow.max_output_bytes > self.limits.max_output_bytes:
|
||||||
|
raise ValueError("workflow max_output_bytes exceeds the workload limit")
|
||||||
|
capabilities = tuple(require_identifier(value, "capability") for value in self.capabilities)
|
||||||
|
if not capabilities or len(capabilities) != len(set(capabilities)):
|
||||||
|
raise ValueError("capabilities must be non-empty and unique")
|
||||||
|
if self.workload.name not in capabilities:
|
||||||
|
raise ValueError("capabilities must include the canonical workload name")
|
||||||
|
object.__setattr__(self, "capabilities", capabilities)
|
||||||
|
profiles = tuple(require_identifier(value, "conformance_profile") for value in self.conformance_profiles)
|
||||||
|
if "core-batch-v1" not in profiles or len(profiles) != len(set(profiles)):
|
||||||
|
raise ValueError("conformance_profiles must uniquely include core-batch-v1")
|
||||||
|
object.__setattr__(self, "conformance_profiles", profiles)
|
||||||
|
required = tuple(self.required_features)
|
||||||
|
optional = tuple(self.optional_features)
|
||||||
|
if any(not isinstance(item, FeatureRequirement) for item in required + optional):
|
||||||
|
raise ValueError("features must contain FeatureRequirement values")
|
||||||
|
names = [item.name for item in required + optional]
|
||||||
|
if len(names) != len(set(names)):
|
||||||
|
raise ValueError("required and optional feature names must be unique")
|
||||||
|
object.__setattr__(self, "required_features", required)
|
||||||
|
object.__setattr__(self, "optional_features", optional)
|
||||||
|
self._validate_acceptance_policy()
|
||||||
|
|
||||||
|
def _validate_acceptance_policy(self) -> None:
|
||||||
|
verifier = self.verifier.verifier
|
||||||
|
exact = verifier == ComponentRef("exact-artifact", 1)
|
||||||
|
canonical = verifier == ComponentRef("canonical-record", 1)
|
||||||
|
numeric = verifier == ComponentRef("numeric-tolerance", 1)
|
||||||
|
if self.determinism is DeterminismProfile.BYTE_EXACT and not exact:
|
||||||
|
raise ValueError("byte_exact workloads require exact-artifact verifier")
|
||||||
|
if self.determinism is DeterminismProfile.CANONICAL_EXACT and not canonical:
|
||||||
|
raise ValueError("canonical_exact workloads require canonical-record verifier")
|
||||||
|
if self.determinism is DeterminismProfile.NUMERIC_TOLERANCE and not numeric:
|
||||||
|
raise ValueError("numeric_tolerance workloads require numeric-tolerance verifier")
|
||||||
|
if TrustMode.UNTRUSTED_QUORUM in self.trust_modes:
|
||||||
|
if self.determinism is not DeterminismProfile.BYTE_EXACT or not exact:
|
||||||
|
raise ValueError("untrusted_quorum v1 requires byte_exact and exact-artifact")
|
||||||
|
if any(stage.kind is StageKind.SIDE_EFFECT for stage in self.workflow.stages):
|
||||||
|
raise ValueError("side-effect stages cannot use untrusted quorum")
|
||||||
|
if self.determinism is DeterminismProfile.SIDE_EFFECTING:
|
||||||
|
if self.trust_modes != (TrustMode.TRUSTED,):
|
||||||
|
raise ValueError("side_effecting workloads must be trusted-only")
|
||||||
|
if not any(stage.kind is StageKind.SIDE_EFFECT for stage in self.workflow.stages):
|
||||||
|
raise ValueError("side_effecting workload requires a side-effect stage")
|
||||||
|
|
||||||
|
@property
|
||||||
|
def digest(self) -> str:
|
||||||
|
import hashlib
|
||||||
|
return hashlib.sha256(self.to_json().encode("utf-8")).hexdigest()
|
||||||
|
|
||||||
|
def to_dict(self) -> dict[str, object]:
|
||||||
|
return {
|
||||||
|
"manifest_schema_version": self.manifest_schema_version,
|
||||||
|
"sdk_api": self.sdk_api.expression,
|
||||||
|
"protocol": self.protocol.expression,
|
||||||
|
"workload": self.workload.to_dict(),
|
||||||
|
"description": self.description,
|
||||||
|
"package": self.package.to_dict(),
|
||||||
|
"environment": self.environment.to_dict(),
|
||||||
|
"parameters_schema": thaw_json(self.parameters_schema),
|
||||||
|
"workflow": self.workflow.to_dict(),
|
||||||
|
"inputs": {name: port.to_dict() for name, port in self.inputs.items()},
|
||||||
|
"outputs": {name: port.to_dict() for name, port in self.outputs.items()},
|
||||||
|
"determinism": self.determinism.value,
|
||||||
|
"trust_modes": [mode.value for mode in self.trust_modes],
|
||||||
|
"verifier": self.verifier.to_dict(),
|
||||||
|
"limits": self.limits.to_dict(),
|
||||||
|
"capabilities": list(self.capabilities),
|
||||||
|
"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],
|
||||||
|
}
|
||||||
|
|
||||||
|
def to_json(self) -> str:
|
||||||
|
return canonical_json(self.to_dict())
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_dict(cls, value: object) -> "WorkloadManifest":
|
||||||
|
if not isinstance(value, Mapping):
|
||||||
|
raise ValueError("workload manifest must be an object")
|
||||||
|
fields = {
|
||||||
|
"manifest_schema_version", "sdk_api", "protocol", "workload", "description",
|
||||||
|
"package", "environment", "parameters_schema", "workflow", "inputs", "outputs",
|
||||||
|
"determinism", "trust_modes", "verifier", "limits", "capabilities",
|
||||||
|
"conformance_profiles", "required_features", "optional_features",
|
||||||
|
}
|
||||||
|
require_exact_keys(value, fields, "workload manifest")
|
||||||
|
inputs, outputs = value["inputs"], value["outputs"]
|
||||||
|
arrays = (
|
||||||
|
value["trust_modes"], value["capabilities"], value["conformance_profiles"],
|
||||||
|
value["required_features"], value["optional_features"],
|
||||||
|
)
|
||||||
|
if not isinstance(inputs, Mapping) or not isinstance(outputs, Mapping):
|
||||||
|
raise ValueError("manifest inputs and outputs must be objects")
|
||||||
|
if any(not isinstance(item, list) for item in arrays):
|
||||||
|
raise ValueError("manifest trust, capability, profile, and feature fields must be arrays")
|
||||||
|
return cls(
|
||||||
|
manifest_schema_version=value["manifest_schema_version"], # type: ignore[arg-type]
|
||||||
|
sdk_api=VersionRange.from_dict(value["sdk_api"]),
|
||||||
|
protocol=VersionRange.from_dict(value["protocol"]),
|
||||||
|
workload=WorkloadId.from_dict(value["workload"]),
|
||||||
|
description=value["description"], # type: ignore[arg-type]
|
||||||
|
package=PackageSpec.from_dict(value["package"]),
|
||||||
|
environment=EnvironmentSpec.from_dict(value["environment"]),
|
||||||
|
parameters_schema=value["parameters_schema"], # type: ignore[arg-type]
|
||||||
|
workflow=WorkflowSpec.from_dict(value["workflow"]),
|
||||||
|
inputs={name: PortSpec.from_dict(port) for name, port in inputs.items()},
|
||||||
|
outputs={name: PortSpec.from_dict(port) for name, port in outputs.items()},
|
||||||
|
determinism=value["determinism"], # type: ignore[arg-type]
|
||||||
|
trust_modes=tuple(value["trust_modes"]), # type: ignore[arg-type]
|
||||||
|
verifier=VerifierSpec.from_dict(value["verifier"]),
|
||||||
|
limits=WorkloadLimits.from_dict(value["limits"]),
|
||||||
|
capabilities=tuple(value["capabilities"]), # type: ignore[arg-type]
|
||||||
|
conformance_profiles=tuple(value["conformance_profiles"]), # type: ignore[arg-type]
|
||||||
|
required_features=tuple(
|
||||||
|
FeatureRequirement.from_dict(item) for item in value["required_features"] # type: ignore[union-attr]
|
||||||
|
),
|
||||||
|
optional_features=tuple(
|
||||||
|
FeatureRequirement.from_dict(item) for item in value["optional_features"] # type: ignore[union-attr]
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_json(cls, value: str) -> "WorkloadManifest":
|
||||||
|
try:
|
||||||
|
decoded = json.loads(value)
|
||||||
|
except (TypeError, json.JSONDecodeError, RecursionError) as error:
|
||||||
|
raise ValueError("workload manifest must be valid JSON") from error
|
||||||
|
return cls.from_dict(decoded)
|
||||||
@@ -0,0 +1,845 @@
|
|||||||
|
"""Strict job, task, workflow-plan, and expansion value objects."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import hashlib
|
||||||
|
import json
|
||||||
|
from dataclasses import dataclass, field
|
||||||
|
from types import MappingProxyType
|
||||||
|
from typing import Any, Mapping
|
||||||
|
|
||||||
|
from ._validation import (
|
||||||
|
canonical_json,
|
||||||
|
freeze_json_mapping,
|
||||||
|
require_exact_keys,
|
||||||
|
require_identifier,
|
||||||
|
require_nonnegative_int,
|
||||||
|
parse_release,
|
||||||
|
require_positive_int,
|
||||||
|
require_sha256,
|
||||||
|
require_schema_version,
|
||||||
|
require_string,
|
||||||
|
require_task_key,
|
||||||
|
require_uuid,
|
||||||
|
thaw_json,
|
||||||
|
)
|
||||||
|
from .artifacts import ArtifactCollection, Cardinality, CollectionKind, PortSpec
|
||||||
|
from .execution import ExecutionProfile
|
||||||
|
from .identity import ComponentRef, TASK_SCHEMA_VERSION, WorkloadId
|
||||||
|
from .manifest import TrustMode
|
||||||
|
from .resources import ResourceRequirements
|
||||||
|
from .workflow import StageKind, StageSpec, WorkflowSpec
|
||||||
|
|
||||||
|
|
||||||
|
def _collections(
|
||||||
|
value: Mapping[str, ArtifactCollection], field: str
|
||||||
|
) -> Mapping[str, ArtifactCollection]:
|
||||||
|
if not isinstance(value, Mapping):
|
||||||
|
raise ValueError(f"{field} must be an object")
|
||||||
|
result: dict[str, ArtifactCollection] = {}
|
||||||
|
for name, collection in value.items():
|
||||||
|
canonical = require_identifier(name, f"{field} port")
|
||||||
|
if not isinstance(collection, ArtifactCollection):
|
||||||
|
raise ValueError(f"{field} values must be ArtifactCollection values")
|
||||||
|
result[canonical] = collection
|
||||||
|
return MappingProxyType(result)
|
||||||
|
|
||||||
|
|
||||||
|
def _ports(value: Mapping[str, PortSpec], field: str) -> Mapping[str, PortSpec]:
|
||||||
|
if not isinstance(value, Mapping):
|
||||||
|
raise ValueError(f"{field} must be an object")
|
||||||
|
result: dict[str, PortSpec] = {}
|
||||||
|
for name, port in value.items():
|
||||||
|
canonical = require_identifier(name, f"{field} port")
|
||||||
|
if not isinstance(port, PortSpec):
|
||||||
|
raise ValueError(f"{field} values must be PortSpec values")
|
||||||
|
result[canonical] = port
|
||||||
|
return MappingProxyType(result)
|
||||||
|
|
||||||
|
|
||||||
|
def _feature_versions(value: object, field: str) -> Mapping[str, str]:
|
||||||
|
if not isinstance(value, Mapping):
|
||||||
|
raise ValueError(f"{field} must be an object")
|
||||||
|
result: dict[str, str] = {}
|
||||||
|
for name, version in value.items():
|
||||||
|
canonical = require_identifier(name, f"{field} feature")
|
||||||
|
text = require_string(version, f"{field} version", max_length=32)
|
||||||
|
parse_release(text, f"{field} version")
|
||||||
|
result[canonical] = text
|
||||||
|
return MappingProxyType(result)
|
||||||
|
|
||||||
|
|
||||||
|
def _fallbacks(value: object, field: str) -> Mapping[str, str]:
|
||||||
|
if not isinstance(value, Mapping):
|
||||||
|
raise ValueError(f"{field} must be an object")
|
||||||
|
return MappingProxyType(
|
||||||
|
{
|
||||||
|
require_identifier(name, f"{field} feature"): require_identifier(
|
||||||
|
fallback,
|
||||||
|
f"{field} fallback",
|
||||||
|
)
|
||||||
|
for name, fallback in value.items()
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class JobRequest:
|
||||||
|
workload: WorkloadId
|
||||||
|
parameters: Mapping[str, Any]
|
||||||
|
inputs: Mapping[str, ArtifactCollection]
|
||||||
|
required_features: tuple[str, ...] = ()
|
||||||
|
trust_mode: TrustMode = TrustMode.TRUSTED
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
if not isinstance(self.workload, WorkloadId):
|
||||||
|
raise ValueError("job workload must be a WorkloadId")
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"parameters",
|
||||||
|
freeze_json_mapping(self.parameters, "job.parameters", forbid_locations=True),
|
||||||
|
)
|
||||||
|
object.__setattr__(self, "inputs", _collections(self.inputs, "job.inputs"))
|
||||||
|
features = tuple(require_identifier(value, "required_feature") for value in self.required_features)
|
||||||
|
if len(features) != len(set(features)):
|
||||||
|
raise ValueError("required_features must be unique")
|
||||||
|
object.__setattr__(self, "required_features", features)
|
||||||
|
try:
|
||||||
|
trust_mode = TrustMode(self.trust_mode)
|
||||||
|
except (TypeError, ValueError) as error:
|
||||||
|
raise ValueError("job trust_mode is unsupported") from error
|
||||||
|
object.__setattr__(self, "trust_mode", trust_mode)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def parameters_digest(self) -> str:
|
||||||
|
return hashlib.sha256(canonical_json(self.parameters).encode("utf-8")).hexdigest()
|
||||||
|
|
||||||
|
def to_dict(self) -> dict[str, object]:
|
||||||
|
return {
|
||||||
|
"workload": self.workload.to_dict(),
|
||||||
|
"parameters": thaw_json(self.parameters),
|
||||||
|
"inputs": {name: value.to_dict() for name, value in self.inputs.items()},
|
||||||
|
"required_features": list(self.required_features),
|
||||||
|
"trust_mode": self.trust_mode.value,
|
||||||
|
}
|
||||||
|
|
||||||
|
def to_json(self) -> str:
|
||||||
|
return canonical_json(self.to_dict())
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_dict(cls, value: object) -> "JobRequest":
|
||||||
|
if not isinstance(value, Mapping):
|
||||||
|
raise ValueError("job request must be an object")
|
||||||
|
fields = {"workload", "parameters", "inputs", "required_features", "trust_mode"}
|
||||||
|
require_exact_keys(value, fields, "job request")
|
||||||
|
inputs = value["inputs"]
|
||||||
|
features = value["required_features"]
|
||||||
|
if not isinstance(inputs, Mapping) or not isinstance(features, list):
|
||||||
|
raise ValueError("job inputs must be an object and required_features an array")
|
||||||
|
return cls(
|
||||||
|
workload=WorkloadId.from_dict(value["workload"]),
|
||||||
|
parameters=value["parameters"], # type: ignore[arg-type]
|
||||||
|
inputs={name: ArtifactCollection.from_dict(item) for name, item in inputs.items()},
|
||||||
|
required_features=tuple(features),
|
||||||
|
trust_mode=value["trust_mode"], # type: ignore[arg-type]
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_json(cls, value: str) -> "JobRequest":
|
||||||
|
try:
|
||||||
|
decoded = json.loads(value)
|
||||||
|
except (TypeError, json.JSONDecodeError, RecursionError) as error:
|
||||||
|
raise ValueError("job request must be valid JSON") from error
|
||||||
|
return cls.from_dict(decoded)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class ValidatedJob:
|
||||||
|
request: JobRequest
|
||||||
|
resolved_parameters: Mapping[str, Any]
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
if not isinstance(self.request, JobRequest):
|
||||||
|
raise ValueError("validated job request must be a JobRequest")
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"resolved_parameters",
|
||||||
|
freeze_json_mapping(
|
||||||
|
self.resolved_parameters,
|
||||||
|
"resolved_parameters",
|
||||||
|
forbid_locations=True,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def parameters_digest(self) -> str:
|
||||||
|
return hashlib.sha256(canonical_json(self.resolved_parameters).encode("utf-8")).hexdigest()
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class TaskSpec:
|
||||||
|
workload: WorkloadId
|
||||||
|
package_digest: str
|
||||||
|
manifest_digest: str
|
||||||
|
trust_mode: TrustMode
|
||||||
|
sdk_api_version: str
|
||||||
|
protocol_version: str
|
||||||
|
manifest_schema_version: int
|
||||||
|
workflow_schema_version: int
|
||||||
|
environment_digest: str
|
||||||
|
verifier: ComponentRef
|
||||||
|
selected_features: Mapping[str, str]
|
||||||
|
optional_fallbacks: Mapping[str, str]
|
||||||
|
task_key: str
|
||||||
|
stage_id: str
|
||||||
|
parameters: Mapping[str, Any]
|
||||||
|
inputs: Mapping[str, ArtifactCollection]
|
||||||
|
expected_outputs: Mapping[str, PortSpec]
|
||||||
|
resources: ResourceRequirements
|
||||||
|
execution: ExecutionProfile
|
||||||
|
expected_input_keys: Mapping[str, tuple[str, ...]] = field(default_factory=dict)
|
||||||
|
schema_version: int = TASK_SCHEMA_VERSION
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
require_schema_version(self.schema_version, TASK_SCHEMA_VERSION, "task schema_version")
|
||||||
|
if not isinstance(self.workload, WorkloadId):
|
||||||
|
raise ValueError("task workload must be a WorkloadId")
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"package_digest",
|
||||||
|
require_sha256(self.package_digest, "task package_digest", prefixed=True),
|
||||||
|
)
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"manifest_digest",
|
||||||
|
require_sha256(self.manifest_digest, "task manifest_digest"),
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
trust_mode = TrustMode(self.trust_mode)
|
||||||
|
except (TypeError, ValueError) as error:
|
||||||
|
raise ValueError("task trust_mode is unsupported") from error
|
||||||
|
object.__setattr__(self, "trust_mode", trust_mode)
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"sdk_api_version",
|
||||||
|
require_string(self.sdk_api_version, "task sdk_api_version", max_length=32),
|
||||||
|
)
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"protocol_version",
|
||||||
|
require_string(self.protocol_version, "task protocol_version", max_length=32),
|
||||||
|
)
|
||||||
|
parse_release(self.sdk_api_version, "task sdk_api_version")
|
||||||
|
parse_release(self.protocol_version, "task protocol_version")
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"manifest_schema_version",
|
||||||
|
require_positive_int(self.manifest_schema_version, "task manifest_schema_version"),
|
||||||
|
)
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"workflow_schema_version",
|
||||||
|
require_positive_int(self.workflow_schema_version, "task workflow_schema_version"),
|
||||||
|
)
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"environment_digest",
|
||||||
|
require_sha256(self.environment_digest, "task environment_digest", prefixed=True),
|
||||||
|
)
|
||||||
|
if not isinstance(self.verifier, ComponentRef):
|
||||||
|
raise ValueError("task verifier must be a ComponentRef")
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"selected_features",
|
||||||
|
_feature_versions(self.selected_features, "task selected_features"),
|
||||||
|
)
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"optional_fallbacks",
|
||||||
|
_fallbacks(self.optional_fallbacks, "task optional_fallbacks"),
|
||||||
|
)
|
||||||
|
if set(self.selected_features).intersection(self.optional_fallbacks):
|
||||||
|
raise ValueError("one task feature cannot be selected and fallbacked")
|
||||||
|
object.__setattr__(self, "task_key", require_task_key(self.task_key))
|
||||||
|
object.__setattr__(self, "stage_id", require_identifier(self.stage_id, "stage_id"))
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"parameters",
|
||||||
|
freeze_json_mapping(self.parameters, "task.parameters", forbid_locations=True),
|
||||||
|
)
|
||||||
|
object.__setattr__(self, "inputs", _collections(self.inputs, "task.inputs"))
|
||||||
|
object.__setattr__(self, "expected_outputs", _ports(self.expected_outputs, "task.expected_outputs"))
|
||||||
|
if not self.expected_outputs:
|
||||||
|
raise ValueError("task expected_outputs must not be empty")
|
||||||
|
if not isinstance(self.resources, ResourceRequirements):
|
||||||
|
raise ValueError("task resources must be ResourceRequirements")
|
||||||
|
if not isinstance(self.execution, ExecutionProfile):
|
||||||
|
raise ValueError("task execution must be ExecutionProfile")
|
||||||
|
self.execution.validate_resources(self.resources)
|
||||||
|
if not isinstance(self.expected_input_keys, Mapping):
|
||||||
|
raise ValueError("expected_input_keys must be an object")
|
||||||
|
expected_keys: dict[str, tuple[str, ...]] = {}
|
||||||
|
for port_name, keys in self.expected_input_keys.items():
|
||||||
|
canonical_port = require_identifier(port_name, "expected input key port")
|
||||||
|
if not isinstance(keys, (list, tuple)):
|
||||||
|
raise ValueError("expected input keys must be arrays")
|
||||||
|
canonical_keys = tuple(sorted(
|
||||||
|
require_identifier(key, "expected input key") for key in keys
|
||||||
|
))
|
||||||
|
if not canonical_keys or len(canonical_keys) != len(set(canonical_keys)):
|
||||||
|
raise ValueError("expected input keys must be non-empty and unique")
|
||||||
|
expected_keys[canonical_port] = canonical_keys
|
||||||
|
object.__setattr__(self, "expected_input_keys", MappingProxyType(expected_keys))
|
||||||
|
|
||||||
|
def validate_stage(self, stage: StageSpec) -> "TaskSpec":
|
||||||
|
if not isinstance(stage, StageSpec) or stage.stage_id != self.stage_id:
|
||||||
|
raise ValueError("task stage does not match its StageSpec")
|
||||||
|
if set(self.inputs) != set(stage.inputs):
|
||||||
|
raise ValueError("task input ports do not match the stage")
|
||||||
|
for name, declaration in stage.inputs.items():
|
||||||
|
declaration.validate_collection(self.inputs[name], f"task input {name}")
|
||||||
|
for name, expected_keys in self.expected_input_keys.items():
|
||||||
|
declaration = stage.inputs.get(name)
|
||||||
|
if (
|
||||||
|
declaration is None
|
||||||
|
or declaration.cardinality is not Cardinality.MANY
|
||||||
|
or declaration.collection is not CollectionKind.KEYED
|
||||||
|
):
|
||||||
|
raise ValueError("expected input keys require a keyed-many stage input")
|
||||||
|
actual_keys = tuple(
|
||||||
|
item.key for item in self.inputs[name].items if item.key is not None
|
||||||
|
)
|
||||||
|
if set(actual_keys) != set(expected_keys):
|
||||||
|
raise ValueError("task keyed input does not match its coordinator expected keys")
|
||||||
|
keyed_many_ports = {
|
||||||
|
name
|
||||||
|
for name, declaration in stage.inputs.items()
|
||||||
|
if declaration.cardinality is Cardinality.MANY
|
||||||
|
and declaration.collection is CollectionKind.KEYED
|
||||||
|
}
|
||||||
|
if set(self.expected_input_keys) != keyed_many_ports:
|
||||||
|
raise ValueError("task must pin expected keys for every keyed-many input")
|
||||||
|
if dict(self.expected_outputs) != dict(stage.outputs):
|
||||||
|
raise ValueError("task expected outputs do not match the stage")
|
||||||
|
if not set(self.parameters).issubset(stage.parameter_names):
|
||||||
|
raise ValueError("task parameters are outside the stage projection")
|
||||||
|
if self.resources != stage.resources or self.execution != stage.execution:
|
||||||
|
raise ValueError("task execution requirements do not match the stage")
|
||||||
|
if self.verifier != stage.verifier:
|
||||||
|
raise ValueError("task verifier does not match the stage acceptance verifier")
|
||||||
|
if self.trust_mode.value not in stage.trust_modes:
|
||||||
|
raise ValueError("task trust mode is not allowed by the stage")
|
||||||
|
return self
|
||||||
|
|
||||||
|
def to_dict(self) -> dict[str, object]:
|
||||||
|
return {
|
||||||
|
"schema_version": self.schema_version,
|
||||||
|
"workload": self.workload.to_dict(),
|
||||||
|
"package_digest": self.package_digest,
|
||||||
|
"manifest_digest": self.manifest_digest,
|
||||||
|
"trust_mode": self.trust_mode.value,
|
||||||
|
"sdk_api_version": self.sdk_api_version,
|
||||||
|
"protocol_version": self.protocol_version,
|
||||||
|
"manifest_schema_version": self.manifest_schema_version,
|
||||||
|
"workflow_schema_version": self.workflow_schema_version,
|
||||||
|
"environment_digest": self.environment_digest,
|
||||||
|
"verifier": self.verifier.canonical,
|
||||||
|
"selected_features": dict(self.selected_features),
|
||||||
|
"optional_fallbacks": dict(self.optional_fallbacks),
|
||||||
|
"task_key": self.task_key,
|
||||||
|
"stage_id": self.stage_id,
|
||||||
|
"parameters": thaw_json(self.parameters),
|
||||||
|
"inputs": {name: value.to_dict() for name, value in self.inputs.items()},
|
||||||
|
"expected_outputs": {name: value.to_dict() for name, value in self.expected_outputs.items()},
|
||||||
|
"resources": self.resources.to_dict(),
|
||||||
|
"execution": self.execution.to_dict(),
|
||||||
|
"expected_input_keys": {
|
||||||
|
name: list(keys) for name, keys in self.expected_input_keys.items()
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
def to_json(self) -> str:
|
||||||
|
return canonical_json(self.to_dict())
|
||||||
|
|
||||||
|
@property
|
||||||
|
def digest(self) -> str:
|
||||||
|
"""Canonical digest used to pin a coordinator execution contract."""
|
||||||
|
return hashlib.sha256(self.to_json().encode("utf-8")).hexdigest()
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_dict(cls, value: object) -> "TaskSpec":
|
||||||
|
if not isinstance(value, Mapping):
|
||||||
|
raise ValueError("task specification must be an object")
|
||||||
|
fields = {
|
||||||
|
"schema_version", "workload", "package_digest", "manifest_digest", "trust_mode",
|
||||||
|
"sdk_api_version", "protocol_version", "manifest_schema_version",
|
||||||
|
"workflow_schema_version", "environment_digest", "verifier",
|
||||||
|
"selected_features", "optional_fallbacks",
|
||||||
|
"task_key", "stage_id", "parameters", "inputs",
|
||||||
|
"expected_outputs", "resources", "execution", "expected_input_keys",
|
||||||
|
}
|
||||||
|
require_exact_keys(value, fields, "task specification")
|
||||||
|
inputs, outputs = value["inputs"], value["expected_outputs"]
|
||||||
|
if not isinstance(inputs, Mapping) or not isinstance(outputs, Mapping):
|
||||||
|
raise ValueError("task inputs and expected_outputs must be objects")
|
||||||
|
return cls(
|
||||||
|
schema_version=value["schema_version"], # type: ignore[arg-type]
|
||||||
|
workload=WorkloadId.from_dict(value["workload"]),
|
||||||
|
package_digest=value["package_digest"], # type: ignore[arg-type]
|
||||||
|
manifest_digest=value["manifest_digest"], # type: ignore[arg-type]
|
||||||
|
trust_mode=value["trust_mode"], # type: ignore[arg-type]
|
||||||
|
sdk_api_version=value["sdk_api_version"], # type: ignore[arg-type]
|
||||||
|
protocol_version=value["protocol_version"], # type: ignore[arg-type]
|
||||||
|
manifest_schema_version=value["manifest_schema_version"], # type: ignore[arg-type]
|
||||||
|
workflow_schema_version=value["workflow_schema_version"], # type: ignore[arg-type]
|
||||||
|
environment_digest=value["environment_digest"], # type: ignore[arg-type]
|
||||||
|
verifier=ComponentRef.from_dict(value["verifier"]),
|
||||||
|
selected_features=value["selected_features"], # type: ignore[arg-type]
|
||||||
|
optional_fallbacks=value["optional_fallbacks"], # type: ignore[arg-type]
|
||||||
|
task_key=value["task_key"], # type: ignore[arg-type]
|
||||||
|
stage_id=value["stage_id"], # type: ignore[arg-type]
|
||||||
|
parameters=value["parameters"], # type: ignore[arg-type]
|
||||||
|
inputs={name: ArtifactCollection.from_dict(item) for name, item in inputs.items()},
|
||||||
|
expected_outputs={name: PortSpec.from_dict(item) for name, item in outputs.items()},
|
||||||
|
resources=ResourceRequirements.from_dict(value["resources"]),
|
||||||
|
execution=ExecutionProfile.from_dict(value["execution"]),
|
||||||
|
expected_input_keys=value["expected_input_keys"], # type: ignore[arg-type]
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_json(cls, value: str) -> "TaskSpec":
|
||||||
|
try:
|
||||||
|
decoded = json.loads(value)
|
||||||
|
except (TypeError, json.JSONDecodeError, RecursionError) as error:
|
||||||
|
raise ValueError("task specification must be valid JSON") from error
|
||||||
|
return cls.from_dict(decoded)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class WorkflowPlan:
|
||||||
|
workload: WorkloadId
|
||||||
|
package_digest: str
|
||||||
|
manifest_digest: str
|
||||||
|
trust_mode: TrustMode
|
||||||
|
sdk_api_version: str
|
||||||
|
protocol_version: str
|
||||||
|
manifest_schema_version: int
|
||||||
|
workflow_schema_version: int
|
||||||
|
environment_digest: str
|
||||||
|
verifier: ComponentRef
|
||||||
|
selected_features: Mapping[str, str]
|
||||||
|
optional_fallbacks: Mapping[str, str]
|
||||||
|
workflow_id: str
|
||||||
|
resolved_parameters: Mapping[str, Any]
|
||||||
|
tasks: tuple[TaskSpec, ...]
|
||||||
|
schema_version: int = 1
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
require_schema_version(self.schema_version, 1, "workflow plan schema_version")
|
||||||
|
if not isinstance(self.workload, WorkloadId):
|
||||||
|
raise ValueError("workflow plan workload must be a WorkloadId")
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"package_digest",
|
||||||
|
require_sha256(self.package_digest, "plan package_digest", prefixed=True),
|
||||||
|
)
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"manifest_digest",
|
||||||
|
require_sha256(self.manifest_digest, "plan manifest_digest"),
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
trust_mode = TrustMode(self.trust_mode)
|
||||||
|
except (TypeError, ValueError) as error:
|
||||||
|
raise ValueError("plan trust_mode is unsupported") from error
|
||||||
|
object.__setattr__(self, "trust_mode", trust_mode)
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"sdk_api_version",
|
||||||
|
require_string(self.sdk_api_version, "plan sdk_api_version", max_length=32),
|
||||||
|
)
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"protocol_version",
|
||||||
|
require_string(self.protocol_version, "plan protocol_version", max_length=32),
|
||||||
|
)
|
||||||
|
parse_release(self.sdk_api_version, "plan sdk_api_version")
|
||||||
|
parse_release(self.protocol_version, "plan protocol_version")
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"manifest_schema_version",
|
||||||
|
require_positive_int(self.manifest_schema_version, "plan manifest_schema_version"),
|
||||||
|
)
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"workflow_schema_version",
|
||||||
|
require_positive_int(self.workflow_schema_version, "plan workflow_schema_version"),
|
||||||
|
)
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"environment_digest",
|
||||||
|
require_sha256(self.environment_digest, "plan environment_digest", prefixed=True),
|
||||||
|
)
|
||||||
|
if not isinstance(self.verifier, ComponentRef):
|
||||||
|
raise ValueError("plan verifier must be a ComponentRef")
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"selected_features",
|
||||||
|
_feature_versions(self.selected_features, "plan selected_features"),
|
||||||
|
)
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"optional_fallbacks",
|
||||||
|
_fallbacks(self.optional_fallbacks, "plan optional_fallbacks"),
|
||||||
|
)
|
||||||
|
if set(self.selected_features).intersection(self.optional_fallbacks):
|
||||||
|
raise ValueError("one plan feature cannot be selected and fallbacked")
|
||||||
|
object.__setattr__(self, "workflow_id", require_identifier(self.workflow_id, "workflow_id"))
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"resolved_parameters",
|
||||||
|
freeze_json_mapping(
|
||||||
|
self.resolved_parameters,
|
||||||
|
"resolved_parameters",
|
||||||
|
forbid_locations=True,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
tasks = tuple(self.tasks)
|
||||||
|
if not tasks or any(not isinstance(task, TaskSpec) for task in tasks):
|
||||||
|
raise ValueError("workflow plan tasks must contain at least one TaskSpec")
|
||||||
|
keys = [task.task_key for task in tasks]
|
||||||
|
if keys != sorted(keys) or len(keys) != len(set(keys)):
|
||||||
|
raise ValueError("workflow plan task keys must be unique and ascending")
|
||||||
|
for task in tasks:
|
||||||
|
if (
|
||||||
|
task.workload != self.workload
|
||||||
|
or task.package_digest != self.package_digest
|
||||||
|
or task.manifest_digest != self.manifest_digest
|
||||||
|
or task.trust_mode is not self.trust_mode
|
||||||
|
or task.sdk_api_version != self.sdk_api_version
|
||||||
|
or task.protocol_version != self.protocol_version
|
||||||
|
or task.manifest_schema_version != self.manifest_schema_version
|
||||||
|
or task.workflow_schema_version != self.workflow_schema_version
|
||||||
|
or task.environment_digest != self.environment_digest
|
||||||
|
or task.selected_features != self.selected_features
|
||||||
|
or task.optional_fallbacks != self.optional_fallbacks
|
||||||
|
):
|
||||||
|
raise ValueError("workflow plan tasks must carry the plan's exact workload pin")
|
||||||
|
object.__setattr__(self, "tasks", tasks)
|
||||||
|
|
||||||
|
def validate_workflow(self, workflow: WorkflowSpec) -> "WorkflowPlan":
|
||||||
|
if workflow.workflow_id != self.workflow_id:
|
||||||
|
raise ValueError("workflow plan references another workflow")
|
||||||
|
if len(self.tasks) > workflow.max_tasks:
|
||||||
|
raise ValueError("workflow plan exceeds max_tasks")
|
||||||
|
stages = {stage.stage_id: stage for stage in workflow.stages}
|
||||||
|
task_counts: dict[str, int] = {}
|
||||||
|
for task in self.tasks:
|
||||||
|
try:
|
||||||
|
task.validate_stage(stages[task.stage_id])
|
||||||
|
except KeyError as error:
|
||||||
|
raise ValueError(f"workflow plan references unknown stage: {task.stage_id}") from error
|
||||||
|
task_counts[task.stage_id] = task_counts.get(task.stage_id, 0) + 1
|
||||||
|
if task_counts[task.stage_id] > stages[task.stage_id].max_fan_out:
|
||||||
|
raise ValueError(f"workflow plan exceeds max_fan_out for stage {task.stage_id}")
|
||||||
|
return self
|
||||||
|
|
||||||
|
@property
|
||||||
|
def digest(self) -> str:
|
||||||
|
return hashlib.sha256(self.to_json().encode("utf-8")).hexdigest()
|
||||||
|
|
||||||
|
def to_dict(self) -> dict[str, object]:
|
||||||
|
return {
|
||||||
|
"schema_version": self.schema_version,
|
||||||
|
"workload": self.workload.to_dict(),
|
||||||
|
"package_digest": self.package_digest,
|
||||||
|
"manifest_digest": self.manifest_digest,
|
||||||
|
"trust_mode": self.trust_mode.value,
|
||||||
|
"sdk_api_version": self.sdk_api_version,
|
||||||
|
"protocol_version": self.protocol_version,
|
||||||
|
"manifest_schema_version": self.manifest_schema_version,
|
||||||
|
"workflow_schema_version": self.workflow_schema_version,
|
||||||
|
"environment_digest": self.environment_digest,
|
||||||
|
"verifier": self.verifier.canonical,
|
||||||
|
"selected_features": dict(self.selected_features),
|
||||||
|
"optional_fallbacks": dict(self.optional_fallbacks),
|
||||||
|
"workflow_id": self.workflow_id,
|
||||||
|
"resolved_parameters": thaw_json(self.resolved_parameters),
|
||||||
|
"tasks": [task.to_dict() for task in self.tasks],
|
||||||
|
}
|
||||||
|
|
||||||
|
def to_json(self) -> str:
|
||||||
|
return canonical_json(self.to_dict())
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_dict(cls, value: object) -> "WorkflowPlan":
|
||||||
|
if not isinstance(value, Mapping):
|
||||||
|
raise ValueError("workflow plan must be an object")
|
||||||
|
fields = {
|
||||||
|
"schema_version", "workload", "package_digest", "manifest_digest", "trust_mode",
|
||||||
|
"sdk_api_version", "protocol_version", "manifest_schema_version",
|
||||||
|
"workflow_schema_version", "environment_digest", "verifier",
|
||||||
|
"selected_features", "optional_fallbacks",
|
||||||
|
"workflow_id", "resolved_parameters", "tasks",
|
||||||
|
}
|
||||||
|
require_exact_keys(value, fields, "workflow plan")
|
||||||
|
tasks = value["tasks"]
|
||||||
|
if not isinstance(tasks, list):
|
||||||
|
raise ValueError("workflow plan tasks must be an array")
|
||||||
|
return cls(
|
||||||
|
schema_version=value["schema_version"], # type: ignore[arg-type]
|
||||||
|
workload=WorkloadId.from_dict(value["workload"]),
|
||||||
|
package_digest=value["package_digest"], # type: ignore[arg-type]
|
||||||
|
manifest_digest=value["manifest_digest"], # type: ignore[arg-type]
|
||||||
|
trust_mode=value["trust_mode"], # type: ignore[arg-type]
|
||||||
|
sdk_api_version=value["sdk_api_version"], # type: ignore[arg-type]
|
||||||
|
protocol_version=value["protocol_version"], # type: ignore[arg-type]
|
||||||
|
manifest_schema_version=value["manifest_schema_version"], # type: ignore[arg-type]
|
||||||
|
workflow_schema_version=value["workflow_schema_version"], # type: ignore[arg-type]
|
||||||
|
environment_digest=value["environment_digest"], # type: ignore[arg-type]
|
||||||
|
verifier=ComponentRef.from_dict(value["verifier"]),
|
||||||
|
selected_features=value["selected_features"], # type: ignore[arg-type]
|
||||||
|
optional_fallbacks=value["optional_fallbacks"], # type: ignore[arg-type]
|
||||||
|
workflow_id=value["workflow_id"], # type: ignore[arg-type]
|
||||||
|
resolved_parameters=value["resolved_parameters"], # type: ignore[arg-type]
|
||||||
|
tasks=tuple(TaskSpec.from_dict(task) for task in tasks),
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_json(cls, value: str) -> "WorkflowPlan":
|
||||||
|
try:
|
||||||
|
decoded = json.loads(value)
|
||||||
|
except (TypeError, json.JSONDecodeError, RecursionError) as error:
|
||||||
|
raise ValueError("workflow plan must be valid JSON") from error
|
||||||
|
return cls.from_dict(decoded)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class ExpansionManifest:
|
||||||
|
job_id: str
|
||||||
|
parent_task_id: str
|
||||||
|
parent_task_key: str
|
||||||
|
parent_execution_contract_digest: str
|
||||||
|
tasks: tuple[TaskSpec, ...]
|
||||||
|
max_children: int
|
||||||
|
schema_version: int = 1
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
require_schema_version(self.schema_version, 1, "expansion manifest schema_version")
|
||||||
|
object.__setattr__(self, "job_id", require_uuid(self.job_id, "expansion job_id"))
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"parent_task_id",
|
||||||
|
require_uuid(self.parent_task_id, "expansion parent_task_id"),
|
||||||
|
)
|
||||||
|
object.__setattr__(self, "parent_task_key", require_task_key(self.parent_task_key))
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"parent_execution_contract_digest",
|
||||||
|
require_sha256(
|
||||||
|
self.parent_execution_contract_digest,
|
||||||
|
"expansion parent_execution_contract_digest",
|
||||||
|
),
|
||||||
|
)
|
||||||
|
object.__setattr__(self, "max_children", require_positive_int(self.max_children, "max_children"))
|
||||||
|
tasks = tuple(self.tasks)
|
||||||
|
if not tasks or len(tasks) > self.max_children:
|
||||||
|
raise ValueError("expansion tasks must be non-empty and within max_children")
|
||||||
|
keys = [task.task_key for task in tasks]
|
||||||
|
if keys != sorted(keys) or len(keys) != len(set(keys)):
|
||||||
|
raise ValueError("expansion child task keys must be unique and ascending")
|
||||||
|
if any(not key.startswith(self.parent_task_key + "/") for key in keys):
|
||||||
|
raise ValueError("expansion child task keys must be namespaced by the parent")
|
||||||
|
first = tasks[0]
|
||||||
|
if any(
|
||||||
|
task.workload != first.workload
|
||||||
|
or task.package_digest != first.package_digest
|
||||||
|
or task.manifest_digest != first.manifest_digest
|
||||||
|
or task.trust_mode is not first.trust_mode
|
||||||
|
or task.sdk_api_version != first.sdk_api_version
|
||||||
|
or task.protocol_version != first.protocol_version
|
||||||
|
or task.manifest_schema_version != first.manifest_schema_version
|
||||||
|
or task.workflow_schema_version != first.workflow_schema_version
|
||||||
|
or task.environment_digest != first.environment_digest
|
||||||
|
or task.selected_features != first.selected_features
|
||||||
|
or task.optional_fallbacks != first.optional_fallbacks
|
||||||
|
for task in tasks[1:]
|
||||||
|
):
|
||||||
|
raise ValueError("expansion child tasks must carry one exact workload pin")
|
||||||
|
object.__setattr__(self, "tasks", tasks)
|
||||||
|
|
||||||
|
def validate_against(
|
||||||
|
self,
|
||||||
|
parent: TaskSpec,
|
||||||
|
workflow: WorkflowSpec,
|
||||||
|
*,
|
||||||
|
job_id: str,
|
||||||
|
parent_task_id: str,
|
||||||
|
declared_max_children: int,
|
||||||
|
remaining_tasks: int,
|
||||||
|
authorized_inputs: Mapping[str, Mapping[str, ArtifactCollection]],
|
||||||
|
existing_stage_task_counts: Mapping[str, int],
|
||||||
|
) -> "ExpansionManifest":
|
||||||
|
"""Validate an expansion against coordinator-owned durable state.
|
||||||
|
|
||||||
|
The IDs and remaining budget are deliberately supplied by the
|
||||||
|
coordinator rather than trusted from the package-produced manifest.
|
||||||
|
"""
|
||||||
|
if not isinstance(parent, TaskSpec):
|
||||||
|
raise ValueError("expansion parent must be a TaskSpec")
|
||||||
|
if not isinstance(workflow, WorkflowSpec):
|
||||||
|
raise ValueError("expansion workflow must be a WorkflowSpec")
|
||||||
|
if self.job_id != require_uuid(job_id, "coordinator job_id"):
|
||||||
|
raise ValueError("expansion belongs to another job")
|
||||||
|
if self.parent_task_id != require_uuid(parent_task_id, "coordinator parent_task_id"):
|
||||||
|
raise ValueError("expansion belongs to another durable parent task")
|
||||||
|
if self.parent_task_key != parent.task_key:
|
||||||
|
raise ValueError("expansion parent task key does not match")
|
||||||
|
if self.parent_execution_contract_digest != parent.digest:
|
||||||
|
raise ValueError("expansion parent execution contract does not match")
|
||||||
|
|
||||||
|
remaining = require_nonnegative_int(remaining_tasks, "remaining_tasks")
|
||||||
|
stages = {stage.stage_id: stage for stage in workflow.stages}
|
||||||
|
try:
|
||||||
|
parent_stage = stages[parent.stage_id]
|
||||||
|
except KeyError as error:
|
||||||
|
raise ValueError("expansion parent references an unknown workflow stage") from error
|
||||||
|
parent.validate_stage(parent_stage)
|
||||||
|
if parent_stage.kind is not StageKind.PLAN:
|
||||||
|
raise ValueError("v1 expansion parent must be a plan stage")
|
||||||
|
declared_limit = require_positive_int(
|
||||||
|
declared_max_children,
|
||||||
|
"declared_max_children",
|
||||||
|
)
|
||||||
|
allowed_children = min(declared_limit, remaining)
|
||||||
|
if self.max_children > declared_limit or len(self.tasks) > allowed_children:
|
||||||
|
raise ValueError("expansion exceeds the coordinator child task budget")
|
||||||
|
|
||||||
|
if not isinstance(authorized_inputs, Mapping):
|
||||||
|
raise ValueError("authorized_inputs must be an object")
|
||||||
|
allowed_by_target: dict[str, dict[str, ArtifactCollection]] = {}
|
||||||
|
for stage_id, ports in authorized_inputs.items():
|
||||||
|
canonical_stage = require_identifier(stage_id, "authorized input stage")
|
||||||
|
if canonical_stage not in stages or not isinstance(ports, Mapping):
|
||||||
|
raise ValueError("authorized_inputs references an unknown stage")
|
||||||
|
allowed_ports: dict[str, ArtifactCollection] = {}
|
||||||
|
for port_name, collection in ports.items():
|
||||||
|
canonical_port = require_identifier(port_name, "authorized input port")
|
||||||
|
declaration = stages[canonical_stage].inputs.get(canonical_port)
|
||||||
|
if declaration is None or not isinstance(collection, ArtifactCollection):
|
||||||
|
raise ValueError("authorized_inputs references an unknown input port")
|
||||||
|
declaration.validate_collection(
|
||||||
|
collection,
|
||||||
|
f"authorized input {canonical_stage}.{canonical_port}",
|
||||||
|
)
|
||||||
|
allowed_ports[canonical_port] = collection
|
||||||
|
allowed_by_target[canonical_stage] = allowed_ports
|
||||||
|
|
||||||
|
raw_counts = existing_stage_task_counts
|
||||||
|
if not isinstance(raw_counts, Mapping):
|
||||||
|
raise ValueError("existing_stage_task_counts must be an object")
|
||||||
|
stage_counts: dict[str, int] = {}
|
||||||
|
for stage_id, count in raw_counts.items():
|
||||||
|
canonical = require_identifier(stage_id, "existing stage task count")
|
||||||
|
if canonical not in stages:
|
||||||
|
raise ValueError("existing task count references an unknown stage")
|
||||||
|
stage_counts[canonical] = require_nonnegative_int(
|
||||||
|
count,
|
||||||
|
"existing stage task count",
|
||||||
|
)
|
||||||
|
|
||||||
|
for task in self.tasks:
|
||||||
|
if (
|
||||||
|
task.workload != parent.workload
|
||||||
|
or task.package_digest != parent.package_digest
|
||||||
|
or task.manifest_digest != parent.manifest_digest
|
||||||
|
or task.trust_mode is not parent.trust_mode
|
||||||
|
or task.sdk_api_version != parent.sdk_api_version
|
||||||
|
or task.protocol_version != parent.protocol_version
|
||||||
|
or task.manifest_schema_version != parent.manifest_schema_version
|
||||||
|
or task.workflow_schema_version != parent.workflow_schema_version
|
||||||
|
or task.environment_digest != parent.environment_digest
|
||||||
|
or task.selected_features != parent.selected_features
|
||||||
|
or task.optional_fallbacks != parent.optional_fallbacks
|
||||||
|
):
|
||||||
|
raise ValueError("expansion child task does not share the parent workload pin")
|
||||||
|
try:
|
||||||
|
stage = stages[task.stage_id]
|
||||||
|
except KeyError as error:
|
||||||
|
raise ValueError("expansion child references an unknown workflow stage") from error
|
||||||
|
if parent.stage_id not in stage.needs:
|
||||||
|
raise ValueError("v1 expansion child must be a direct successor of its parent stage")
|
||||||
|
task.validate_stage(stage)
|
||||||
|
target_ports = allowed_by_target.get(task.stage_id, {})
|
||||||
|
for port_name, collection in task.inputs.items():
|
||||||
|
allowed = target_ports.get(port_name)
|
||||||
|
if allowed is None or collection.kind is not allowed.kind:
|
||||||
|
raise ValueError("expansion child input target is not coordinator-authorized")
|
||||||
|
if collection.kind is CollectionKind.ORDERED:
|
||||||
|
cursor = 0
|
||||||
|
for item in collection.items:
|
||||||
|
while cursor < len(allowed.items) and allowed.items[cursor] != item:
|
||||||
|
cursor += 1
|
||||||
|
if cursor == len(allowed.items):
|
||||||
|
raise ValueError(
|
||||||
|
"expansion child input is not an authorized ordered subsequence"
|
||||||
|
)
|
||||||
|
cursor += 1
|
||||||
|
elif any(item not in allowed.items for item in collection.items):
|
||||||
|
raise ValueError(
|
||||||
|
"expansion child input artifact is not coordinator-authorized"
|
||||||
|
)
|
||||||
|
stage_counts[task.stage_id] = stage_counts.get(task.stage_id, 0) + 1
|
||||||
|
if stage_counts[task.stage_id] > stage.max_fan_out:
|
||||||
|
raise ValueError(
|
||||||
|
f"expansion exceeds max_fan_out for stage {task.stage_id}"
|
||||||
|
)
|
||||||
|
return self
|
||||||
|
|
||||||
|
@property
|
||||||
|
def digest(self) -> str:
|
||||||
|
return hashlib.sha256(canonical_json(self.to_dict()).encode("utf-8")).hexdigest()
|
||||||
|
|
||||||
|
def to_dict(self) -> dict[str, object]:
|
||||||
|
return {
|
||||||
|
"schema_version": self.schema_version,
|
||||||
|
"job_id": self.job_id,
|
||||||
|
"parent_task_id": self.parent_task_id,
|
||||||
|
"parent_task_key": self.parent_task_key,
|
||||||
|
"parent_execution_contract_digest": self.parent_execution_contract_digest,
|
||||||
|
"max_children": self.max_children,
|
||||||
|
"tasks": [task.to_dict() for task in self.tasks],
|
||||||
|
}
|
||||||
|
|
||||||
|
def to_json(self) -> str:
|
||||||
|
return canonical_json(self.to_dict())
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_dict(cls, value: object) -> "ExpansionManifest":
|
||||||
|
if not isinstance(value, Mapping):
|
||||||
|
raise ValueError("expansion manifest must be an object")
|
||||||
|
fields = {
|
||||||
|
"schema_version", "job_id", "parent_task_id", "parent_task_key",
|
||||||
|
"parent_execution_contract_digest", "max_children", "tasks",
|
||||||
|
}
|
||||||
|
require_exact_keys(value, fields, "expansion manifest")
|
||||||
|
tasks = value["tasks"]
|
||||||
|
if not isinstance(tasks, list):
|
||||||
|
raise ValueError("expansion tasks must be an array")
|
||||||
|
return cls(
|
||||||
|
schema_version=value["schema_version"], # type: ignore[arg-type]
|
||||||
|
job_id=value["job_id"], # type: ignore[arg-type]
|
||||||
|
parent_task_id=value["parent_task_id"], # type: ignore[arg-type]
|
||||||
|
parent_task_key=value["parent_task_key"], # type: ignore[arg-type]
|
||||||
|
parent_execution_contract_digest=value["parent_execution_contract_digest"], # type: ignore[arg-type]
|
||||||
|
max_children=value["max_children"], # type: ignore[arg-type]
|
||||||
|
tasks=tuple(TaskSpec.from_dict(task) for task in tasks),
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_json(cls, value: str) -> "ExpansionManifest":
|
||||||
|
try:
|
||||||
|
decoded = json.loads(value)
|
||||||
|
except (TypeError, json.JSONDecodeError, RecursionError) as error:
|
||||||
|
raise ValueError("expansion manifest must be valid JSON") from error
|
||||||
|
return cls.from_dict(decoded)
|
||||||
@@ -0,0 +1,108 @@
|
|||||||
|
"""Author-facing planner, runner, reducer, and verifier protocols."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Mapping, Protocol, Sequence
|
||||||
|
|
||||||
|
from .artifacts import ArtifactCollection, ArtifactRef, ArtifactSchema, OutputManifest, Provenance
|
||||||
|
from .plans import JobRequest, TaskSpec, ValidatedJob, WorkflowPlan
|
||||||
|
from .runtime import NegotiatedWorkload
|
||||||
|
from .verification import CandidateOutputs, VerificationDecision, VerifyContext
|
||||||
|
|
||||||
|
|
||||||
|
class ArtifactCatalog(Protocol):
|
||||||
|
"""Bridge-owned, read-only access to durable input artifacts."""
|
||||||
|
|
||||||
|
def materialize(self, artifact: ArtifactRef) -> Path:
|
||||||
|
"""Return an attempt-scoped verified local copy without exposing credentials."""
|
||||||
|
|
||||||
|
|
||||||
|
class ArtifactSink(Protocol):
|
||||||
|
"""Agent/bridge-owned sealing boundary for scientific output files."""
|
||||||
|
|
||||||
|
def seal(
|
||||||
|
self,
|
||||||
|
path: Path,
|
||||||
|
*,
|
||||||
|
declaration: ArtifactSchema,
|
||||||
|
records: int | None = None,
|
||||||
|
dimensions: tuple[int, ...] = (),
|
||||||
|
) -> ArtifactRef:
|
||||||
|
"""Validate/upload bytes and return coordinator-owned immutable metadata."""
|
||||||
|
|
||||||
|
|
||||||
|
class CancellationToken(Protocol):
|
||||||
|
def cancelled(self) -> bool: ...
|
||||||
|
|
||||||
|
def raise_if_cancelled(self) -> None: ...
|
||||||
|
|
||||||
|
|
||||||
|
class PlanningResources(Protocol):
|
||||||
|
"""Caller-provided catalog, sink, and workspace for registry planning."""
|
||||||
|
|
||||||
|
@property
|
||||||
|
def catalog(self) -> ArtifactCatalog: ...
|
||||||
|
|
||||||
|
@property
|
||||||
|
def sink(self) -> ArtifactSink: ...
|
||||||
|
|
||||||
|
@property
|
||||||
|
def workspace(self) -> Path: ...
|
||||||
|
|
||||||
|
|
||||||
|
class PlanningContext(PlanningResources, Protocol):
|
||||||
|
"""Planner-facing resources augmented by completed negotiation."""
|
||||||
|
|
||||||
|
@property
|
||||||
|
def negotiated(self) -> NegotiatedWorkload:
|
||||||
|
"""Resolved optional fallbacks and the exact negotiated manifest."""
|
||||||
|
|
||||||
|
|
||||||
|
class TaskContext(Protocol):
|
||||||
|
@property
|
||||||
|
def task(self) -> TaskSpec: ...
|
||||||
|
|
||||||
|
@property
|
||||||
|
def catalog(self) -> ArtifactCatalog: ...
|
||||||
|
|
||||||
|
@property
|
||||||
|
def sink(self) -> ArtifactSink: ...
|
||||||
|
|
||||||
|
@property
|
||||||
|
def workspace(self) -> Path: ...
|
||||||
|
|
||||||
|
@property
|
||||||
|
def cancellation(self) -> CancellationToken: ...
|
||||||
|
|
||||||
|
@property
|
||||||
|
def provenance(self) -> Provenance: ...
|
||||||
|
|
||||||
|
|
||||||
|
class ReduceContext(TaskContext, Protocol):
|
||||||
|
@property
|
||||||
|
def accepted_inputs(self) -> Mapping[str, ArtifactCollection]: ...
|
||||||
|
|
||||||
|
|
||||||
|
class Planner(Protocol):
|
||||||
|
entry_point: str
|
||||||
|
|
||||||
|
def validate(self, request: JobRequest) -> ValidatedJob: ...
|
||||||
|
|
||||||
|
def plan(self, job: ValidatedJob, context: PlanningContext) -> WorkflowPlan: ...
|
||||||
|
|
||||||
|
|
||||||
|
class Runner(Protocol):
|
||||||
|
def run(self, context: TaskContext) -> OutputManifest: ...
|
||||||
|
|
||||||
|
|
||||||
|
class Reducer(Protocol):
|
||||||
|
def reduce(self, context: ReduceContext) -> OutputManifest: ...
|
||||||
|
|
||||||
|
|
||||||
|
class Verifier(Protocol):
|
||||||
|
def verify(
|
||||||
|
self,
|
||||||
|
context: VerifyContext,
|
||||||
|
candidates: CandidateOutputs,
|
||||||
|
) -> VerificationDecision: ...
|
||||||
@@ -0,0 +1,526 @@
|
|||||||
|
"""Explicit, digest-pinned workload package registry and safe discovery."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import re
|
||||||
|
import sys
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from importlib import machinery, util
|
||||||
|
from importlib import metadata
|
||||||
|
from pathlib import Path
|
||||||
|
from tempfile import TemporaryDirectory
|
||||||
|
from threading import RLock
|
||||||
|
from types import MappingProxyType
|
||||||
|
from typing import Any, Mapping
|
||||||
|
|
||||||
|
from ._validation import (
|
||||||
|
canonical_json,
|
||||||
|
require_semver,
|
||||||
|
require_sha256,
|
||||||
|
require_string,
|
||||||
|
require_workload_name,
|
||||||
|
)
|
||||||
|
from .identity import ComponentRef, WorkloadId
|
||||||
|
from .integrity import installed_distribution_digest
|
||||||
|
from .manifest import WorkloadManifest
|
||||||
|
from .plans import JobRequest, ValidatedJob, WorkflowPlan
|
||||||
|
from .protocols import Planner, PlanningContext, PlanningResources, Reducer, Runner, Verifier
|
||||||
|
from .runtime import CompatibilityError, NegotiatedWorkload, RuntimeCapabilities, negotiate_manifest
|
||||||
|
from .schema import validate_parameter_instance
|
||||||
|
from .workflow import StageKind
|
||||||
|
|
||||||
|
|
||||||
|
_DISCOVERY_IMPORT_LOCK = RLock()
|
||||||
|
|
||||||
|
|
||||||
|
def _normalized_distribution_name(value: str) -> str:
|
||||||
|
return re.sub(r"[-_.]+", "-", value).lower()
|
||||||
|
|
||||||
|
|
||||||
|
def _validate_entry_point_ownership(entry_point: metadata.EntryPoint) -> None:
|
||||||
|
"""Require the entry-point module to be payload of its own distribution."""
|
||||||
|
distribution = entry_point.dist
|
||||||
|
if distribution is None:
|
||||||
|
raise ValueError("workload entry point has no owning distribution")
|
||||||
|
module_name = getattr(entry_point, "module", None)
|
||||||
|
if not isinstance(module_name, str) or not module_name:
|
||||||
|
value = getattr(entry_point, "value", "")
|
||||||
|
module_name = value.partition(":")[0].strip() if isinstance(value, str) else ""
|
||||||
|
parts = module_name.split(".")
|
||||||
|
if not parts or any(not part.isidentifier() for part in parts):
|
||||||
|
raise ValueError("workload entry point has an invalid module path")
|
||||||
|
|
||||||
|
raw_top_level = distribution.read_text("top_level.txt")
|
||||||
|
declared = {
|
||||||
|
line.strip()
|
||||||
|
for line in raw_top_level.splitlines()
|
||||||
|
if line.strip()
|
||||||
|
} if raw_top_level is not None else set()
|
||||||
|
root_name = parts[0]
|
||||||
|
if root_name not in declared:
|
||||||
|
raise ValueError("workload entry point module is outside its distribution")
|
||||||
|
|
||||||
|
owners = metadata.packages_distributions().get(root_name, ())
|
||||||
|
normalized_owners = {_normalized_distribution_name(owner) for owner in owners}
|
||||||
|
expected_owner = _normalized_distribution_name(distribution.name)
|
||||||
|
if normalized_owners and normalized_owners != {expected_owner}:
|
||||||
|
raise ValueError("workload entry point top-level package is not uniquely owned")
|
||||||
|
|
||||||
|
package_root = Path(distribution.locate_file(root_name))
|
||||||
|
if not package_root.exists():
|
||||||
|
root_spec = util.find_spec(root_name)
|
||||||
|
locations = (
|
||||||
|
tuple(root_spec.submodule_search_locations or ())
|
||||||
|
if root_spec is not None
|
||||||
|
else ()
|
||||||
|
)
|
||||||
|
if len(locations) == 1:
|
||||||
|
package_root = Path(locations[0])
|
||||||
|
if package_root.is_dir():
|
||||||
|
module_base = package_root.joinpath(*parts[1:])
|
||||||
|
ownership_root = package_root.resolve()
|
||||||
|
candidates = [
|
||||||
|
*(Path(str(module_base) + suffix) for suffix in machinery.SOURCE_SUFFIXES),
|
||||||
|
*(Path(str(module_base) + suffix) for suffix in machinery.EXTENSION_SUFFIXES),
|
||||||
|
*(module_base / ("__init__" + suffix) for suffix in machinery.SOURCE_SUFFIXES),
|
||||||
|
*(module_base / ("__init__" + suffix) for suffix in machinery.EXTENSION_SUFFIXES),
|
||||||
|
]
|
||||||
|
else:
|
||||||
|
if len(parts) != 1:
|
||||||
|
raise ValueError("workload entry point module is outside its distribution")
|
||||||
|
ownership_root = Path(distribution.locate_file(".")).resolve()
|
||||||
|
module_base = Path(distribution.locate_file(root_name))
|
||||||
|
candidates = [
|
||||||
|
*(Path(str(module_base) + suffix) for suffix in machinery.SOURCE_SUFFIXES),
|
||||||
|
*(Path(str(module_base) + suffix) for suffix in machinery.EXTENSION_SUFFIXES),
|
||||||
|
]
|
||||||
|
existing = tuple(candidate for candidate in candidates if candidate.is_file())
|
||||||
|
if len(existing) != 1 or not existing[0].resolve().is_relative_to(ownership_root):
|
||||||
|
raise ValueError("workload entry point module is not an owned package payload")
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class WorkloadDefinition:
|
||||||
|
manifest: WorkloadManifest
|
||||||
|
planner: Planner
|
||||||
|
runners: Mapping[str, Runner]
|
||||||
|
reducers: Mapping[str, Reducer]
|
||||||
|
verifiers: Mapping[str, Verifier]
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
if not isinstance(self.manifest, WorkloadManifest):
|
||||||
|
raise ValueError("definition manifest must be a WorkloadManifest")
|
||||||
|
if not callable(getattr(self.planner, "validate", None)) or not callable(
|
||||||
|
getattr(self.planner, "plan", None)
|
||||||
|
):
|
||||||
|
raise ValueError("definition planner must implement validate and plan")
|
||||||
|
collections: list[tuple[str, Mapping[str, Any], str]] = [
|
||||||
|
("runners", self.runners, "run"),
|
||||||
|
("reducers", self.reducers, "reduce"),
|
||||||
|
("verifiers", self.verifiers, "verify"),
|
||||||
|
]
|
||||||
|
for field, values, method in collections:
|
||||||
|
if not isinstance(values, Mapping):
|
||||||
|
raise ValueError(f"definition {field} must be an object")
|
||||||
|
copied: dict[str, Any] = {}
|
||||||
|
for name, handler in values.items():
|
||||||
|
canonical = require_string(name, f"{field} entry point", max_length=256)
|
||||||
|
if not callable(getattr(handler, method, None)):
|
||||||
|
raise ValueError(f"definition {field} handler must implement {method}")
|
||||||
|
copied[canonical] = handler
|
||||||
|
object.__setattr__(self, field, MappingProxyType(copied))
|
||||||
|
for stage in self.manifest.workflow.stages:
|
||||||
|
if stage.kind is StageKind.PLAN:
|
||||||
|
if getattr(self.planner, "entry_point", None) != stage.entry_point:
|
||||||
|
raise ValueError(
|
||||||
|
"PLAN stage entry point must match planner.entry_point"
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
if stage.kind is StageKind.REDUCE:
|
||||||
|
handlers = self.reducers
|
||||||
|
else:
|
||||||
|
# A VERIFY node is still an executable DAG stage. Its
|
||||||
|
# ``entry_point`` is a Runner; ``stage.verifier`` selects the
|
||||||
|
# independent acceptance component applied to its output.
|
||||||
|
handlers = self.runners
|
||||||
|
if stage.entry_point not in handlers:
|
||||||
|
raise ValueError(
|
||||||
|
f"definition has no installed handler for stage entry point: {stage.entry_point}"
|
||||||
|
)
|
||||||
|
verifier_key = self.manifest.verifier.verifier.canonical
|
||||||
|
if verifier_key not in self.verifiers:
|
||||||
|
raise ValueError(f"definition has no installed manifest verifier: {verifier_key}")
|
||||||
|
for key, verifier in self.verifiers.items():
|
||||||
|
try:
|
||||||
|
declared_identity = ComponentRef.from_dict(key)
|
||||||
|
except ValueError as error:
|
||||||
|
raise ValueError("definition verifier keys must be component identities") from error
|
||||||
|
if declared_identity.canonical != key or getattr(verifier, "identity", None) != declared_identity:
|
||||||
|
raise ValueError("definition verifier handler identity does not match its key")
|
||||||
|
manifest_verifier = self.verifiers[verifier_key]
|
||||||
|
handler_configuration = getattr(manifest_verifier, "configuration", None)
|
||||||
|
if handler_configuration is None:
|
||||||
|
if self.manifest.verifier.configuration:
|
||||||
|
raise ValueError("manifest verifier configuration is not bound by its handler")
|
||||||
|
elif dict(handler_configuration) != dict(self.manifest.verifier.configuration):
|
||||||
|
raise ValueError("manifest verifier configuration does not match its handler")
|
||||||
|
for stage in self.manifest.workflow.stages:
|
||||||
|
if stage.verifier is not None and stage.verifier.canonical not in self.verifiers:
|
||||||
|
raise ValueError(
|
||||||
|
f"definition has no installed stage verifier: {stage.verifier.canonical}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class AllowedPackage:
|
||||||
|
distribution: str
|
||||||
|
workload: WorkloadId
|
||||||
|
digest: str
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
distribution = require_string(self.distribution, "distribution", max_length=128).lower()
|
||||||
|
if not re.fullmatch(r"[a-z0-9]+(?:[-_.][a-z0-9]+)*", distribution):
|
||||||
|
raise ValueError("distribution must be a canonical Python distribution name")
|
||||||
|
object.__setattr__(self, "distribution", distribution.replace("_", "-"))
|
||||||
|
if not isinstance(self.workload, WorkloadId):
|
||||||
|
raise ValueError("allowed workload must be a WorkloadId")
|
||||||
|
object.__setattr__(self, "digest", require_sha256(self.digest, "allowed digest", prefixed=True))
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class WorkloadDescription:
|
||||||
|
workload: WorkloadId
|
||||||
|
description: str
|
||||||
|
package_digest: str
|
||||||
|
enabled: bool
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class _NegotiatedPlanningContext:
|
||||||
|
base: PlanningResources
|
||||||
|
negotiated: NegotiatedWorkload
|
||||||
|
|
||||||
|
@property
|
||||||
|
def catalog(self):
|
||||||
|
return self.base.catalog
|
||||||
|
|
||||||
|
@property
|
||||||
|
def sink(self):
|
||||||
|
return self.base.sink
|
||||||
|
|
||||||
|
@property
|
||||||
|
def workspace(self) -> Path:
|
||||||
|
return self.base.workspace
|
||||||
|
|
||||||
|
|
||||||
|
class WorkloadRegistry:
|
||||||
|
"""Registry keyed by exact workload version and immutable package digest."""
|
||||||
|
|
||||||
|
ENTRY_POINT_GROUP = "scimesh.workloads"
|
||||||
|
|
||||||
|
def __init__(self) -> None:
|
||||||
|
self._definitions: dict[tuple[str, str], WorkloadDefinition] = {}
|
||||||
|
self._enabled: set[tuple[str, str, str]] = set()
|
||||||
|
self._lock = RLock()
|
||||||
|
|
||||||
|
def register(self, definition: WorkloadDefinition, *, enabled: bool = False) -> None:
|
||||||
|
if not isinstance(definition, WorkloadDefinition):
|
||||||
|
raise ValueError("definition must be a WorkloadDefinition")
|
||||||
|
workload = definition.manifest.workload
|
||||||
|
key = (workload.name, workload.version)
|
||||||
|
with self._lock:
|
||||||
|
if key in self._definitions:
|
||||||
|
raise ValueError(f"workload version already registered: {workload.name}@{workload.version}")
|
||||||
|
self._definitions[key] = definition
|
||||||
|
if enabled:
|
||||||
|
self._enabled.add((*key, definition.manifest.package.digest))
|
||||||
|
|
||||||
|
def enable(self, name: str, version: str, package_digest: str) -> None:
|
||||||
|
digest = require_sha256(package_digest, "package_digest", prefixed=True)
|
||||||
|
with self._lock:
|
||||||
|
definition = self._registered(name, version)
|
||||||
|
if digest != definition.manifest.package.digest:
|
||||||
|
raise ValueError("package digest does not match the registered manifest")
|
||||||
|
self._enabled.add((definition.manifest.workload.name, version, digest))
|
||||||
|
|
||||||
|
def disable(self, name: str, version: str, package_digest: str) -> None:
|
||||||
|
canonical = require_workload_name(name)
|
||||||
|
version = require_semver(version, "workload.version")
|
||||||
|
digest = require_sha256(package_digest, "package_digest", prefixed=True)
|
||||||
|
with self._lock:
|
||||||
|
self._enabled.discard((canonical, version, digest))
|
||||||
|
|
||||||
|
def _registered(self, name: str, version: str) -> WorkloadDefinition:
|
||||||
|
canonical = require_workload_name(name)
|
||||||
|
version = require_semver(version, "workload.version")
|
||||||
|
with self._lock:
|
||||||
|
try:
|
||||||
|
return self._definitions[(canonical, version)]
|
||||||
|
except KeyError as error:
|
||||||
|
raise ValueError(f"unknown workload version: {canonical}@{version}") from error
|
||||||
|
|
||||||
|
def require(
|
||||||
|
self,
|
||||||
|
name: str,
|
||||||
|
version: str,
|
||||||
|
package_digest: str,
|
||||||
|
*,
|
||||||
|
runtime: RuntimeCapabilities | None = None,
|
||||||
|
) -> tuple[WorkloadDefinition, NegotiatedWorkload | None]:
|
||||||
|
digest = require_sha256(package_digest, "package_digest", prefixed=True)
|
||||||
|
with self._lock:
|
||||||
|
definition = self._registered(name, version)
|
||||||
|
identity = (definition.manifest.workload.name, definition.manifest.workload.version, digest)
|
||||||
|
if digest != definition.manifest.package.digest or identity not in self._enabled:
|
||||||
|
raise ValueError("workload package digest is not enabled")
|
||||||
|
negotiated = negotiate_manifest(definition.manifest, runtime) if runtime is not None else None
|
||||||
|
return definition, negotiated
|
||||||
|
|
||||||
|
def plan(
|
||||||
|
self,
|
||||||
|
request: JobRequest,
|
||||||
|
package_digest: str,
|
||||||
|
runtime: RuntimeCapabilities,
|
||||||
|
context: PlanningResources,
|
||||||
|
) -> WorkflowPlan:
|
||||||
|
"""Negotiate first, then invoke only the pre-registered planner object."""
|
||||||
|
if not isinstance(request, JobRequest):
|
||||||
|
raise ValueError("request must be a JobRequest")
|
||||||
|
definition, negotiated = self.require(
|
||||||
|
request.workload.name,
|
||||||
|
request.workload.version,
|
||||||
|
package_digest,
|
||||||
|
runtime=runtime,
|
||||||
|
)
|
||||||
|
assert negotiated is not None
|
||||||
|
self._validate_request_compatibility(request, definition.manifest, runtime, negotiated)
|
||||||
|
self._validate_request_shape(request, definition.manifest)
|
||||||
|
validated = definition.planner.validate(request)
|
||||||
|
if not isinstance(validated, ValidatedJob) or validated.request != request:
|
||||||
|
raise ValueError("planner.validate must return a ValidatedJob for the same request")
|
||||||
|
plan = definition.planner.plan(
|
||||||
|
validated,
|
||||||
|
_NegotiatedPlanningContext(context, negotiated),
|
||||||
|
)
|
||||||
|
if not isinstance(plan, WorkflowPlan) or plan.workload != request.workload:
|
||||||
|
raise ValueError("planner.plan must return a WorkflowPlan for the requested workload")
|
||||||
|
if (
|
||||||
|
plan.package_digest != definition.manifest.package.digest
|
||||||
|
or plan.manifest_digest != definition.manifest.digest
|
||||||
|
or plan.trust_mode is not request.trust_mode
|
||||||
|
or plan.sdk_api_version != runtime.sdk_api_version
|
||||||
|
or plan.protocol_version != runtime.protocol_version
|
||||||
|
or plan.manifest_schema_version != definition.manifest.manifest_schema_version
|
||||||
|
or plan.workflow_schema_version != definition.manifest.workflow.schema_version
|
||||||
|
or plan.environment_digest != definition.manifest.environment.digest
|
||||||
|
or plan.verifier != definition.manifest.verifier.verifier
|
||||||
|
or plan.selected_features != negotiated.selected_features
|
||||||
|
or plan.optional_fallbacks != negotiated.optional_fallbacks
|
||||||
|
):
|
||||||
|
raise ValueError("planner plan does not carry the selected immutable workload pin")
|
||||||
|
plan.validate_workflow(definition.manifest.workflow)
|
||||||
|
self._validate_plan_limits(request, plan, definition.manifest)
|
||||||
|
return WorkflowPlan.from_json(plan.to_json())
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _validate_request_compatibility(
|
||||||
|
request: JobRequest,
|
||||||
|
manifest: WorkloadManifest,
|
||||||
|
runtime: RuntimeCapabilities,
|
||||||
|
negotiated: NegotiatedWorkload,
|
||||||
|
) -> None:
|
||||||
|
if request.trust_mode not in manifest.trust_modes:
|
||||||
|
raise CompatibilityError(
|
||||||
|
"trust-mode-undeclared",
|
||||||
|
"requested trust mode is not declared by the workload",
|
||||||
|
)
|
||||||
|
if request.trust_mode not in runtime.trust_modes:
|
||||||
|
raise CompatibilityError(
|
||||||
|
"trust-mode-unavailable",
|
||||||
|
"runtime cannot enforce the requested trust mode",
|
||||||
|
)
|
||||||
|
for stage in manifest.workflow.stages:
|
||||||
|
if request.trust_mode.value not in stage.trust_modes:
|
||||||
|
raise CompatibilityError(
|
||||||
|
"stage-trust-unavailable",
|
||||||
|
f"stage {stage.stage_id} does not support the requested trust mode",
|
||||||
|
)
|
||||||
|
declared = {
|
||||||
|
feature.name: feature
|
||||||
|
for feature in manifest.required_features + manifest.optional_features
|
||||||
|
}
|
||||||
|
for name in request.required_features:
|
||||||
|
requirement = declared.get(name)
|
||||||
|
if requirement is None:
|
||||||
|
raise CompatibilityError(
|
||||||
|
"feature-undeclared",
|
||||||
|
f"job requests a feature not declared by the workload: {name}",
|
||||||
|
)
|
||||||
|
version = runtime.features.get(name)
|
||||||
|
if version is None or not requirement.versions.contains(version):
|
||||||
|
raise CompatibilityError(
|
||||||
|
"feature-unavailable",
|
||||||
|
f"job-required feature is unavailable or incompatible: {name}",
|
||||||
|
)
|
||||||
|
if name in negotiated.optional_fallbacks:
|
||||||
|
raise CompatibilityError(
|
||||||
|
"feature-fallback-disallowed",
|
||||||
|
f"job-required feature cannot use its fallback: {name}",
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _validate_request_shape(request: JobRequest, manifest: WorkloadManifest) -> None:
|
||||||
|
if set(request.inputs) != set(manifest.inputs):
|
||||||
|
raise ValueError("job input ports do not match the manifest")
|
||||||
|
total_bytes = 0
|
||||||
|
artifact_references: dict[str, object] = {}
|
||||||
|
for name, port in manifest.inputs.items():
|
||||||
|
port.validate_collection(request.inputs[name], f"job input {name}")
|
||||||
|
total_bytes += request.inputs[name].size_bytes
|
||||||
|
for item in request.inputs[name].items:
|
||||||
|
existing = artifact_references.get(item.artifact.artifact_id)
|
||||||
|
if existing is not None and existing != item.artifact:
|
||||||
|
raise ValueError("job reuses an artifact ID with conflicting metadata")
|
||||||
|
artifact_references[item.artifact.artifact_id] = item.artifact
|
||||||
|
if total_bytes > manifest.limits.max_input_bytes:
|
||||||
|
raise ValueError("job inputs exceed the manifest byte limit")
|
||||||
|
if len(artifact_references) > manifest.limits.max_artifacts:
|
||||||
|
raise ValueError("job inputs exceed the manifest artifact limit")
|
||||||
|
import json
|
||||||
|
from ._validation import thaw_json
|
||||||
|
if len(json.dumps(thaw_json(request.parameters), allow_nan=False).encode("utf-8")) > manifest.limits.max_parameter_bytes:
|
||||||
|
raise ValueError("job parameters exceed the manifest byte limit")
|
||||||
|
validate_parameter_instance(request.parameters, manifest.parameters_schema)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _validate_plan_limits(
|
||||||
|
request: JobRequest,
|
||||||
|
plan: WorkflowPlan,
|
||||||
|
manifest: WorkloadManifest,
|
||||||
|
) -> None:
|
||||||
|
references = {
|
||||||
|
item.artifact.artifact_id: item.artifact
|
||||||
|
for collection in request.inputs.values()
|
||||||
|
for item in collection.items
|
||||||
|
}
|
||||||
|
for task in plan.tasks:
|
||||||
|
for collection in task.inputs.values():
|
||||||
|
for item in collection.items:
|
||||||
|
existing = references.get(item.artifact.artifact_id)
|
||||||
|
if existing is not None and existing != item.artifact:
|
||||||
|
raise ValueError("workflow plan reuses an artifact ID with conflicting metadata")
|
||||||
|
references[item.artifact.artifact_id] = item.artifact
|
||||||
|
if len(canonical_json(task.parameters).encode("utf-8")) > manifest.limits.max_parameter_bytes:
|
||||||
|
raise ValueError("planned task parameters exceed the manifest byte limit")
|
||||||
|
if len(references) > manifest.limits.max_artifacts:
|
||||||
|
raise ValueError("workflow plan exceeds the manifest artifact limit")
|
||||||
|
if len(canonical_json(plan.resolved_parameters).encode("utf-8")) > manifest.limits.max_parameter_bytes:
|
||||||
|
raise ValueError("resolved parameters exceed the manifest byte limit")
|
||||||
|
|
||||||
|
def descriptions(self) -> tuple[WorkloadDescription, ...]:
|
||||||
|
with self._lock:
|
||||||
|
result = []
|
||||||
|
for key, definition in sorted(self._definitions.items()):
|
||||||
|
digest = definition.manifest.package.digest
|
||||||
|
result.append(
|
||||||
|
WorkloadDescription(
|
||||||
|
definition.manifest.workload,
|
||||||
|
definition.manifest.description,
|
||||||
|
digest,
|
||||||
|
(*key, digest) in self._enabled,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
return tuple(result)
|
||||||
|
|
||||||
|
def discover_installed(self, allowlist: tuple[AllowedPackage, ...]) -> None:
|
||||||
|
"""Load only configured installed entry points; never accept job module paths."""
|
||||||
|
allowed: dict[tuple[str, str, str], AllowedPackage] = {}
|
||||||
|
for item in allowlist:
|
||||||
|
if not isinstance(item, AllowedPackage):
|
||||||
|
raise ValueError("allowlist must contain AllowedPackage values")
|
||||||
|
key = (item.distribution, item.workload.name, item.workload.version)
|
||||||
|
if key in allowed:
|
||||||
|
raise ValueError("allowlist identities must be unique")
|
||||||
|
allowed[key] = item
|
||||||
|
entry_points = metadata.entry_points()
|
||||||
|
selected = entry_points.select(group=self.ENTRY_POINT_GROUP)
|
||||||
|
discovered: set[tuple[str, str, str]] = set()
|
||||||
|
pending: list[WorkloadDefinition] = []
|
||||||
|
for entry_point in selected:
|
||||||
|
distribution = (
|
||||||
|
_normalized_distribution_name(entry_point.dist.name)
|
||||||
|
if entry_point.dist
|
||||||
|
else ""
|
||||||
|
)
|
||||||
|
for key, approval in allowed.items():
|
||||||
|
if _normalized_distribution_name(key[0]) != distribution:
|
||||||
|
continue
|
||||||
|
if entry_point.name != f"{approval.workload.name}@{approval.workload.version}":
|
||||||
|
continue
|
||||||
|
_validate_entry_point_ownership(entry_point)
|
||||||
|
# Import policy is process-global, so installed discovery is
|
||||||
|
# serialized and intended for application startup. An empty
|
||||||
|
# cache prefix prevents pre-existing package pyc files from
|
||||||
|
# being consumed, while dont_write_bytecode keeps the
|
||||||
|
# measured source tree unchanged during both load and factory.
|
||||||
|
with _DISCOVERY_IMPORT_LOCK, TemporaryDirectory(
|
||||||
|
prefix="scimesh-discovery-cache-"
|
||||||
|
) as cache_prefix:
|
||||||
|
measured_before = installed_distribution_digest(entry_point.dist)
|
||||||
|
if measured_before != approval.digest:
|
||||||
|
raise ValueError(
|
||||||
|
"installed package content does not match its allowlist digest"
|
||||||
|
)
|
||||||
|
previous_bytecode_policy = sys.dont_write_bytecode
|
||||||
|
previous_cache_prefix = sys.pycache_prefix
|
||||||
|
sys.dont_write_bytecode = True
|
||||||
|
sys.pycache_prefix = cache_prefix
|
||||||
|
try:
|
||||||
|
loaded = entry_point.load()
|
||||||
|
definition = (
|
||||||
|
loaded()
|
||||||
|
if callable(loaded) and not isinstance(loaded, WorkloadDefinition)
|
||||||
|
else loaded
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
sys.pycache_prefix = previous_cache_prefix
|
||||||
|
sys.dont_write_bytecode = previous_bytecode_policy
|
||||||
|
if installed_distribution_digest(entry_point.dist) != measured_before:
|
||||||
|
raise ValueError(
|
||||||
|
"installed package content changed while loading its entry point"
|
||||||
|
)
|
||||||
|
if not isinstance(definition, WorkloadDefinition):
|
||||||
|
raise ValueError("workload entry point must provide a WorkloadDefinition")
|
||||||
|
if definition.manifest.workload != approval.workload:
|
||||||
|
raise ValueError("discovered workload identity does not match its allowlist entry")
|
||||||
|
if definition.manifest.package.distribution != approval.distribution:
|
||||||
|
raise ValueError("discovered package identity does not match its allowlist entry")
|
||||||
|
if definition.manifest.package.digest != approval.digest:
|
||||||
|
raise ValueError("discovered package digest does not match its allowlist entry")
|
||||||
|
if key in discovered:
|
||||||
|
raise ValueError("multiple installed entry points match one allowlist entry")
|
||||||
|
pending.append(definition)
|
||||||
|
discovered.add(key)
|
||||||
|
break
|
||||||
|
missing = sorted(set(allowed) - discovered)
|
||||||
|
if missing:
|
||||||
|
identities = ", ".join(f"{name}@{version}" for _, name, version in missing)
|
||||||
|
raise ValueError("allowlisted workload entry points were not installed: " + identities)
|
||||||
|
pending_keys = [
|
||||||
|
(definition.manifest.workload.name, definition.manifest.workload.version)
|
||||||
|
for definition in pending
|
||||||
|
]
|
||||||
|
if len(pending_keys) != len(set(pending_keys)):
|
||||||
|
raise ValueError("multiple allowlisted distributions provide one workload version")
|
||||||
|
with self._lock:
|
||||||
|
conflicts = [key for key in pending_keys if key in self._definitions]
|
||||||
|
if conflicts:
|
||||||
|
name, version = conflicts[0]
|
||||||
|
raise ValueError(f"workload version already registered: {name}@{version}")
|
||||||
|
definitions = dict(self._definitions)
|
||||||
|
enabled = set(self._enabled)
|
||||||
|
for key, definition in zip(pending_keys, pending):
|
||||||
|
definitions[key] = definition
|
||||||
|
enabled.add((*key, definition.manifest.package.digest))
|
||||||
|
self._definitions = definitions
|
||||||
|
self._enabled = enabled
|
||||||
@@ -0,0 +1,463 @@
|
|||||||
|
"""Generic resource declarations, runtime inventory, and atomic local allocation."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from enum import Enum
|
||||||
|
from threading import Lock
|
||||||
|
from types import MappingProxyType
|
||||||
|
from typing import Mapping
|
||||||
|
from uuid import uuid4
|
||||||
|
|
||||||
|
from ._validation import (
|
||||||
|
enum_value,
|
||||||
|
freeze_json_mapping,
|
||||||
|
require_exact_keys,
|
||||||
|
require_identifier,
|
||||||
|
require_nonnegative_int,
|
||||||
|
require_opaque_resource_id,
|
||||||
|
require_positive_int,
|
||||||
|
require_sha256,
|
||||||
|
require_string,
|
||||||
|
thaw_json,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class AcceleratorMode(str, Enum):
|
||||||
|
NONE = "none"
|
||||||
|
EXCLUSIVE_DEVICE = "exclusive_device"
|
||||||
|
FRACTIONAL = "fractional"
|
||||||
|
PARTITION = "partition"
|
||||||
|
|
||||||
|
|
||||||
|
def _resource_id(value: object, field: str) -> str:
|
||||||
|
return require_opaque_resource_id(value, field)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class AcceleratorDevice:
|
||||||
|
kind: str
|
||||||
|
vendor: str
|
||||||
|
device_id: str
|
||||||
|
model: str
|
||||||
|
memory_mb: int
|
||||||
|
modes: tuple[AcceleratorMode, ...]
|
||||||
|
capabilities: Mapping[str, str]
|
||||||
|
topology_group: str | None = None
|
||||||
|
partition_id: str | None = None
|
||||||
|
healthy: bool = True
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
object.__setattr__(self, "kind", require_identifier(self.kind, "accelerator.kind"))
|
||||||
|
object.__setattr__(self, "vendor", require_identifier(self.vendor, "accelerator.vendor"))
|
||||||
|
object.__setattr__(self, "device_id", _resource_id(self.device_id, "accelerator.device_id"))
|
||||||
|
object.__setattr__(self, "model", require_string(self.model, "accelerator.model", max_length=160))
|
||||||
|
object.__setattr__(self, "memory_mb", require_positive_int(self.memory_mb, "accelerator.memory_mb"))
|
||||||
|
modes = tuple(enum_value(AcceleratorMode, mode, "accelerator.mode") for mode in self.modes)
|
||||||
|
if not modes or AcceleratorMode.NONE in modes or len(modes) != len(set(modes)):
|
||||||
|
raise ValueError("accelerator modes must contain unique allocation modes other than none")
|
||||||
|
object.__setattr__(self, "modes", modes)
|
||||||
|
capabilities = freeze_json_mapping(self.capabilities, "accelerator.capabilities")
|
||||||
|
if any(not isinstance(value, str) for value in capabilities.values()):
|
||||||
|
raise ValueError("accelerator capabilities must use string values")
|
||||||
|
object.__setattr__(self, "capabilities", capabilities)
|
||||||
|
if self.topology_group is not None:
|
||||||
|
object.__setattr__(self, "topology_group", _resource_id(self.topology_group, "topology_group"))
|
||||||
|
if self.partition_id is not None:
|
||||||
|
object.__setattr__(self, "partition_id", _resource_id(self.partition_id, "partition_id"))
|
||||||
|
if AcceleratorMode.PARTITION not in modes:
|
||||||
|
raise ValueError("a partition_id requires partition allocation support")
|
||||||
|
if AcceleratorMode.EXCLUSIVE_DEVICE in modes:
|
||||||
|
raise ValueError("an accelerator partition cannot be allocated as a whole device")
|
||||||
|
if not isinstance(self.healthy, bool):
|
||||||
|
raise ValueError("accelerator.healthy must be a boolean")
|
||||||
|
|
||||||
|
@property
|
||||||
|
def allocation_id(self) -> str:
|
||||||
|
return self.partition_id or self.device_id
|
||||||
|
|
||||||
|
def to_dict(self) -> dict[str, object]:
|
||||||
|
return {
|
||||||
|
"kind": self.kind,
|
||||||
|
"vendor": self.vendor,
|
||||||
|
"device_id": self.device_id,
|
||||||
|
"model": self.model,
|
||||||
|
"memory_mb": self.memory_mb,
|
||||||
|
"modes": [mode.value for mode in self.modes],
|
||||||
|
"capabilities": thaw_json(self.capabilities),
|
||||||
|
"topology_group": self.topology_group,
|
||||||
|
"partition_id": self.partition_id,
|
||||||
|
"healthy": self.healthy,
|
||||||
|
}
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_dict(cls, value: object) -> "AcceleratorDevice":
|
||||||
|
if not isinstance(value, Mapping):
|
||||||
|
raise ValueError("accelerator device must be an object")
|
||||||
|
fields = {
|
||||||
|
"kind", "vendor", "device_id", "model", "memory_mb", "modes",
|
||||||
|
"capabilities", "topology_group", "partition_id", "healthy",
|
||||||
|
}
|
||||||
|
require_exact_keys(value, fields, "accelerator device")
|
||||||
|
modes = value["modes"]
|
||||||
|
if not isinstance(modes, list):
|
||||||
|
raise ValueError("accelerator modes must be an array")
|
||||||
|
return cls(
|
||||||
|
kind=value["kind"], # type: ignore[arg-type]
|
||||||
|
vendor=value["vendor"], # type: ignore[arg-type]
|
||||||
|
device_id=value["device_id"], # type: ignore[arg-type]
|
||||||
|
model=value["model"], # type: ignore[arg-type]
|
||||||
|
memory_mb=value["memory_mb"], # type: ignore[arg-type]
|
||||||
|
modes=tuple(modes),
|
||||||
|
capabilities=value["capabilities"], # type: ignore[arg-type]
|
||||||
|
topology_group=value["topology_group"], # type: ignore[arg-type]
|
||||||
|
partition_id=value["partition_id"], # type: ignore[arg-type]
|
||||||
|
healthy=value["healthy"], # type: ignore[arg-type]
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class ResourceInventory:
|
||||||
|
cpu_cores: int
|
||||||
|
memory_mb: int
|
||||||
|
scratch_mb: int
|
||||||
|
architecture: str
|
||||||
|
accelerators: tuple[AcceleratorDevice, ...] = ()
|
||||||
|
environment_digests: tuple[str, ...] = ()
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
object.__setattr__(self, "cpu_cores", require_positive_int(self.cpu_cores, "inventory.cpu_cores"))
|
||||||
|
object.__setattr__(self, "memory_mb", require_positive_int(self.memory_mb, "inventory.memory_mb"))
|
||||||
|
object.__setattr__(self, "scratch_mb", require_nonnegative_int(self.scratch_mb, "inventory.scratch_mb"))
|
||||||
|
object.__setattr__(self, "architecture", require_identifier(self.architecture, "inventory.architecture"))
|
||||||
|
devices = tuple(self.accelerators)
|
||||||
|
if any(not isinstance(device, AcceleratorDevice) for device in devices):
|
||||||
|
raise ValueError("inventory accelerators must contain AcceleratorDevice values")
|
||||||
|
ids = [device.allocation_id for device in devices]
|
||||||
|
if len(ids) != len(set(ids)):
|
||||||
|
raise ValueError("inventory accelerator allocation IDs must be unique")
|
||||||
|
object.__setattr__(self, "accelerators", devices)
|
||||||
|
digests = tuple(
|
||||||
|
require_sha256(value, "environment_digest", prefixed=True)
|
||||||
|
for value in self.environment_digests
|
||||||
|
)
|
||||||
|
if len(digests) != len(set(digests)):
|
||||||
|
raise ValueError("environment_digests must be unique")
|
||||||
|
object.__setattr__(self, "environment_digests", digests)
|
||||||
|
|
||||||
|
def to_dict(self) -> dict[str, object]:
|
||||||
|
return {
|
||||||
|
"cpu_cores": self.cpu_cores,
|
||||||
|
"memory_mb": self.memory_mb,
|
||||||
|
"scratch_mb": self.scratch_mb,
|
||||||
|
"architecture": self.architecture,
|
||||||
|
"accelerators": [device.to_dict() for device in self.accelerators],
|
||||||
|
"environment_digests": list(self.environment_digests),
|
||||||
|
}
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_dict(cls, value: object) -> "ResourceInventory":
|
||||||
|
if not isinstance(value, Mapping):
|
||||||
|
raise ValueError("resource inventory must be an object")
|
||||||
|
fields = {
|
||||||
|
"cpu_cores", "memory_mb", "scratch_mb", "architecture",
|
||||||
|
"accelerators", "environment_digests",
|
||||||
|
}
|
||||||
|
require_exact_keys(value, fields, "resource inventory")
|
||||||
|
accelerators = value["accelerators"]
|
||||||
|
digests = value["environment_digests"]
|
||||||
|
if not isinstance(accelerators, list) or not isinstance(digests, list):
|
||||||
|
raise ValueError("inventory accelerators and environment_digests must be arrays")
|
||||||
|
return cls(
|
||||||
|
cpu_cores=value["cpu_cores"], # type: ignore[arg-type]
|
||||||
|
memory_mb=value["memory_mb"], # type: ignore[arg-type]
|
||||||
|
scratch_mb=value["scratch_mb"], # type: ignore[arg-type]
|
||||||
|
architecture=value["architecture"], # type: ignore[arg-type]
|
||||||
|
accelerators=tuple(AcceleratorDevice.from_dict(device) for device in accelerators),
|
||||||
|
environment_digests=tuple(digests),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class ResourceRequirements:
|
||||||
|
profile: str
|
||||||
|
cpu_cores: int
|
||||||
|
memory_mb: int
|
||||||
|
scratch_mb: int
|
||||||
|
accelerator_count: int = 0
|
||||||
|
accelerator_kind: str | None = None
|
||||||
|
accelerator_memory_mb: int = 0
|
||||||
|
accelerator_mode: AcceleratorMode = AcceleratorMode.NONE
|
||||||
|
architecture: str | None = None
|
||||||
|
topology_group: str | None = None
|
||||||
|
environment_digest: str | None = None
|
||||||
|
estimated_input_bytes: int = 0
|
||||||
|
estimated_output_bytes: int = 0
|
||||||
|
max_duration_seconds: int = 3600
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
object.__setattr__(self, "profile", require_identifier(self.profile, "resources.profile"))
|
||||||
|
object.__setattr__(self, "cpu_cores", require_positive_int(self.cpu_cores, "resources.cpu_cores"))
|
||||||
|
object.__setattr__(self, "memory_mb", require_positive_int(self.memory_mb, "resources.memory_mb"))
|
||||||
|
object.__setattr__(self, "scratch_mb", require_nonnegative_int(self.scratch_mb, "resources.scratch_mb"))
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"accelerator_count",
|
||||||
|
require_nonnegative_int(self.accelerator_count, "resources.accelerator_count"),
|
||||||
|
)
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"accelerator_memory_mb",
|
||||||
|
require_nonnegative_int(self.accelerator_memory_mb, "resources.accelerator_memory_mb"),
|
||||||
|
)
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"accelerator_mode",
|
||||||
|
enum_value(AcceleratorMode, self.accelerator_mode, "resources.accelerator_mode"),
|
||||||
|
)
|
||||||
|
if self.accelerator_count == 0:
|
||||||
|
if self.accelerator_kind is not None or self.accelerator_memory_mb or self.accelerator_mode is not AcceleratorMode.NONE:
|
||||||
|
raise ValueError("CPU-only resources must not declare accelerator constraints")
|
||||||
|
if self.topology_group is not None:
|
||||||
|
raise ValueError("CPU-only resources must not declare accelerator topology")
|
||||||
|
else:
|
||||||
|
if self.accelerator_kind is None:
|
||||||
|
raise ValueError("accelerator_kind is required when accelerator_count is non-zero")
|
||||||
|
object.__setattr__(self, "accelerator_kind", require_identifier(self.accelerator_kind, "accelerator_kind"))
|
||||||
|
if self.accelerator_mode is AcceleratorMode.NONE:
|
||||||
|
raise ValueError("accelerator_mode is required when accelerator_count is non-zero")
|
||||||
|
if self.architecture is not None:
|
||||||
|
object.__setattr__(self, "architecture", require_identifier(self.architecture, "resources.architecture"))
|
||||||
|
if self.topology_group is not None:
|
||||||
|
object.__setattr__(self, "topology_group", _resource_id(self.topology_group, "resources.topology_group"))
|
||||||
|
if self.environment_digest is not None:
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"environment_digest",
|
||||||
|
require_sha256(self.environment_digest, "resources.environment_digest", prefixed=True),
|
||||||
|
)
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"estimated_input_bytes",
|
||||||
|
require_nonnegative_int(self.estimated_input_bytes, "estimated_input_bytes"),
|
||||||
|
)
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"estimated_output_bytes",
|
||||||
|
require_nonnegative_int(self.estimated_output_bytes, "estimated_output_bytes"),
|
||||||
|
)
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"max_duration_seconds",
|
||||||
|
require_positive_int(self.max_duration_seconds, "max_duration_seconds"),
|
||||||
|
)
|
||||||
|
|
||||||
|
def eligibility_errors(self, inventory: ResourceInventory) -> tuple[str, ...]:
|
||||||
|
errors: list[str] = []
|
||||||
|
if self.cpu_cores > inventory.cpu_cores:
|
||||||
|
errors.append("insufficient-cpu")
|
||||||
|
if self.memory_mb > inventory.memory_mb:
|
||||||
|
errors.append("insufficient-memory")
|
||||||
|
if self.scratch_mb > inventory.scratch_mb:
|
||||||
|
errors.append("insufficient-scratch")
|
||||||
|
if self.architecture is not None and self.architecture != inventory.architecture:
|
||||||
|
errors.append("architecture-mismatch")
|
||||||
|
if self.environment_digest is not None and self.environment_digest not in inventory.environment_digests:
|
||||||
|
errors.append("environment-unavailable")
|
||||||
|
matches = self._matching_devices(inventory.accelerators)
|
||||||
|
if len(matches) < self.accelerator_count:
|
||||||
|
errors.append("accelerator-unavailable")
|
||||||
|
return tuple(errors)
|
||||||
|
|
||||||
|
def _matching_devices(
|
||||||
|
self,
|
||||||
|
devices: tuple[AcceleratorDevice, ...],
|
||||||
|
unavailable: set[str] | None = None,
|
||||||
|
) -> tuple[AcceleratorDevice, ...]:
|
||||||
|
unavailable = unavailable or set()
|
||||||
|
if self.accelerator_count == 0:
|
||||||
|
return ()
|
||||||
|
matches = [
|
||||||
|
device
|
||||||
|
for device in devices
|
||||||
|
if device.healthy
|
||||||
|
and device.allocation_id not in unavailable
|
||||||
|
and device.kind == self.accelerator_kind
|
||||||
|
and device.memory_mb >= self.accelerator_memory_mb
|
||||||
|
and self.accelerator_mode in device.modes
|
||||||
|
and (
|
||||||
|
(self.accelerator_mode is AcceleratorMode.PARTITION and device.partition_id is not None)
|
||||||
|
or (
|
||||||
|
self.accelerator_mode is AcceleratorMode.EXCLUSIVE_DEVICE
|
||||||
|
and device.partition_id is None
|
||||||
|
)
|
||||||
|
or self.accelerator_mode is AcceleratorMode.FRACTIONAL
|
||||||
|
)
|
||||||
|
and (self.topology_group is None or device.topology_group == self.topology_group)
|
||||||
|
]
|
||||||
|
if self.accelerator_count > 1 and self.topology_group is None:
|
||||||
|
groups: dict[str | None, list[AcceleratorDevice]] = {}
|
||||||
|
for device in matches:
|
||||||
|
groups.setdefault(device.topology_group, []).append(device)
|
||||||
|
sufficiently_large = [group for group in groups.values() if len(group) >= self.accelerator_count]
|
||||||
|
if sufficiently_large:
|
||||||
|
matches = min(sufficiently_large, key=lambda group: tuple(item.allocation_id for item in group))
|
||||||
|
return tuple(sorted(matches, key=lambda device: device.allocation_id))
|
||||||
|
|
||||||
|
def to_dict(self) -> dict[str, object]:
|
||||||
|
return {
|
||||||
|
"profile": self.profile,
|
||||||
|
"cpu_cores": self.cpu_cores,
|
||||||
|
"memory_mb": self.memory_mb,
|
||||||
|
"scratch_mb": self.scratch_mb,
|
||||||
|
"accelerator_count": self.accelerator_count,
|
||||||
|
"accelerator_kind": self.accelerator_kind,
|
||||||
|
"accelerator_memory_mb": self.accelerator_memory_mb,
|
||||||
|
"accelerator_mode": self.accelerator_mode.value,
|
||||||
|
"architecture": self.architecture,
|
||||||
|
"topology_group": self.topology_group,
|
||||||
|
"environment_digest": self.environment_digest,
|
||||||
|
"estimated_input_bytes": self.estimated_input_bytes,
|
||||||
|
"estimated_output_bytes": self.estimated_output_bytes,
|
||||||
|
"max_duration_seconds": self.max_duration_seconds,
|
||||||
|
}
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_dict(cls, value: object) -> "ResourceRequirements":
|
||||||
|
if not isinstance(value, Mapping):
|
||||||
|
raise ValueError("resource requirements must be an object")
|
||||||
|
fields = {
|
||||||
|
"profile", "cpu_cores", "memory_mb", "scratch_mb", "accelerator_count",
|
||||||
|
"accelerator_kind", "accelerator_memory_mb", "accelerator_mode", "architecture",
|
||||||
|
"topology_group", "environment_digest", "estimated_input_bytes",
|
||||||
|
"estimated_output_bytes", "max_duration_seconds",
|
||||||
|
}
|
||||||
|
require_exact_keys(value, fields, "resource requirements")
|
||||||
|
return cls(**value) # type: ignore[arg-type]
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class ResourceAllocation:
|
||||||
|
allocation_id: str
|
||||||
|
owner_id: str
|
||||||
|
cpu_cores: int
|
||||||
|
memory_mb: int
|
||||||
|
scratch_mb: int
|
||||||
|
accelerator_ids: tuple[str, ...]
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
object.__setattr__(self, "allocation_id", _resource_id(self.allocation_id, "allocation_id"))
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"owner_id",
|
||||||
|
require_string(self.owner_id, "reservation owner_id", max_length=256),
|
||||||
|
)
|
||||||
|
object.__setattr__(self, "cpu_cores", require_positive_int(self.cpu_cores, "allocation.cpu_cores"))
|
||||||
|
object.__setattr__(self, "memory_mb", require_positive_int(self.memory_mb, "allocation.memory_mb"))
|
||||||
|
object.__setattr__(self, "scratch_mb", require_nonnegative_int(self.scratch_mb, "allocation.scratch_mb"))
|
||||||
|
ids = tuple(_resource_id(value, "accelerator_id") for value in self.accelerator_ids)
|
||||||
|
if len(ids) != len(set(ids)):
|
||||||
|
raise ValueError("accelerator_ids must be unique")
|
||||||
|
object.__setattr__(self, "accelerator_ids", ids)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def task_key(self) -> str:
|
||||||
|
"""Compatibility alias; new callers must supply a globally unique attempt owner."""
|
||||||
|
return self.owner_id
|
||||||
|
|
||||||
|
|
||||||
|
class ResourceUnavailableError(RuntimeError):
|
||||||
|
"""Raised before execution when a complete atomic reservation is unavailable."""
|
||||||
|
|
||||||
|
|
||||||
|
class ResourcePool:
|
||||||
|
"""Lock-protected local allocator used by an Agent execution layer.
|
||||||
|
|
||||||
|
This object is intentionally coordinator-independent. A protocol-v2 Agent
|
||||||
|
will bind its returned allocation ID to a coordinator-owned reservation
|
||||||
|
token; the current protocol must not enable concurrent claims based only on
|
||||||
|
this local state.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, inventory: ResourceInventory, *, max_concurrency: int = 1) -> None:
|
||||||
|
if not isinstance(inventory, ResourceInventory):
|
||||||
|
raise ValueError("inventory must be a ResourceInventory")
|
||||||
|
self.inventory = inventory
|
||||||
|
self.max_concurrency = require_positive_int(max_concurrency, "max_concurrency")
|
||||||
|
self._lock = Lock()
|
||||||
|
self._allocations: dict[str, ResourceAllocation] = {}
|
||||||
|
self._allocated_devices: dict[str, tuple[AcceleratorDevice, ...]] = {}
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _devices_conflict(left: AcceleratorDevice, right: AcceleratorDevice) -> bool:
|
||||||
|
if left.device_id != right.device_id:
|
||||||
|
return False
|
||||||
|
if left.partition_id is None or right.partition_id is None:
|
||||||
|
return True
|
||||||
|
return left.partition_id == right.partition_id
|
||||||
|
|
||||||
|
def reserve(self, owner_id: str, requirements: ResourceRequirements) -> ResourceAllocation:
|
||||||
|
if not isinstance(requirements, ResourceRequirements):
|
||||||
|
raise ValueError("requirements must be ResourceRequirements")
|
||||||
|
owner_id = require_string(owner_id, "reservation owner_id", max_length=256)
|
||||||
|
if requirements.accelerator_mode is AcceleratorMode.FRACTIONAL:
|
||||||
|
raise ResourceUnavailableError("fractional-accelerator-unsupported")
|
||||||
|
with self._lock:
|
||||||
|
if any(allocation.owner_id == owner_id for allocation in self._allocations.values()):
|
||||||
|
raise ValueError("reservation owner already has an active resource allocation")
|
||||||
|
if len(self._allocations) >= self.max_concurrency:
|
||||||
|
raise ResourceUnavailableError("execution-slot-unavailable")
|
||||||
|
used_cpu = sum(allocation.cpu_cores for allocation in self._allocations.values())
|
||||||
|
used_memory = sum(allocation.memory_mb for allocation in self._allocations.values())
|
||||||
|
used_scratch = sum(allocation.scratch_mb for allocation in self._allocations.values())
|
||||||
|
if used_cpu + requirements.cpu_cores > self.inventory.cpu_cores:
|
||||||
|
raise ResourceUnavailableError("insufficient-cpu")
|
||||||
|
if used_memory + requirements.memory_mb > self.inventory.memory_mb:
|
||||||
|
raise ResourceUnavailableError("insufficient-memory")
|
||||||
|
if used_scratch + requirements.scratch_mb > self.inventory.scratch_mb:
|
||||||
|
raise ResourceUnavailableError("insufficient-scratch")
|
||||||
|
static_errors = tuple(
|
||||||
|
error
|
||||||
|
for error in requirements.eligibility_errors(self.inventory)
|
||||||
|
if error not in {"insufficient-cpu", "insufficient-memory", "insufficient-scratch", "accelerator-unavailable"}
|
||||||
|
)
|
||||||
|
if static_errors:
|
||||||
|
raise ResourceUnavailableError(static_errors[0])
|
||||||
|
reserved_devices = tuple(
|
||||||
|
device
|
||||||
|
for values in self._allocated_devices.values()
|
||||||
|
for device in values
|
||||||
|
)
|
||||||
|
available_devices = tuple(
|
||||||
|
device
|
||||||
|
for device in self.inventory.accelerators
|
||||||
|
if not any(self._devices_conflict(device, reserved) for reserved in reserved_devices)
|
||||||
|
)
|
||||||
|
devices = requirements._matching_devices(available_devices)
|
||||||
|
if len(devices) < requirements.accelerator_count:
|
||||||
|
raise ResourceUnavailableError("accelerator-unavailable")
|
||||||
|
selected = tuple(device.allocation_id for device in devices[: requirements.accelerator_count])
|
||||||
|
allocation = ResourceAllocation(
|
||||||
|
allocation_id=str(uuid4()),
|
||||||
|
owner_id=owner_id,
|
||||||
|
cpu_cores=requirements.cpu_cores,
|
||||||
|
memory_mb=requirements.memory_mb,
|
||||||
|
scratch_mb=requirements.scratch_mb,
|
||||||
|
accelerator_ids=selected,
|
||||||
|
)
|
||||||
|
self._allocations[allocation.allocation_id] = allocation
|
||||||
|
self._allocated_devices[allocation.allocation_id] = tuple(
|
||||||
|
devices[: requirements.accelerator_count]
|
||||||
|
)
|
||||||
|
return allocation
|
||||||
|
|
||||||
|
def release(self, allocation_id: str) -> bool:
|
||||||
|
allocation_id = _resource_id(allocation_id, "allocation_id")
|
||||||
|
with self._lock:
|
||||||
|
removed = self._allocations.pop(allocation_id, None)
|
||||||
|
self._allocated_devices.pop(allocation_id, None)
|
||||||
|
return removed is not None
|
||||||
|
|
||||||
|
def active_allocations(self) -> tuple[ResourceAllocation, ...]:
|
||||||
|
with self._lock:
|
||||||
|
return tuple(sorted(self._allocations.values(), key=lambda item: item.owner_id))
|
||||||
@@ -0,0 +1,269 @@
|
|||||||
|
"""Fail-closed SDK/profile/feature/resource compatibility negotiation."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from types import MappingProxyType
|
||||||
|
from typing import Mapping
|
||||||
|
|
||||||
|
from ._validation import require_identifier, require_string, validate_version_range, version_in_range
|
||||||
|
from .identity import SDK_API_VERSION
|
||||||
|
from .execution import NetworkPolicy, ProcessModel
|
||||||
|
from .manifest import TrustMode, WorkloadManifest
|
||||||
|
from .resources import AcceleratorMode, ResourceInventory
|
||||||
|
from .workflow import StageKind
|
||||||
|
|
||||||
|
|
||||||
|
class CompatibilityError(ValueError):
|
||||||
|
def __init__(self, code: str, message: str) -> None:
|
||||||
|
self.code = require_identifier(code, "compatibility error code")
|
||||||
|
super().__init__(message)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class RuntimeCapabilities:
|
||||||
|
sdk_api_version: str
|
||||||
|
protocol_version: str
|
||||||
|
profiles: tuple[str, ...]
|
||||||
|
features: Mapping[str, str]
|
||||||
|
workload_capabilities: tuple[str, ...]
|
||||||
|
inventory: ResourceInventory
|
||||||
|
trust_modes: tuple[TrustMode, ...] = (TrustMode.TRUSTED,)
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
object.__setattr__(self, "sdk_api_version", require_string(self.sdk_api_version, "sdk_api_version"))
|
||||||
|
object.__setattr__(self, "protocol_version", require_string(self.protocol_version, "protocol_version"))
|
||||||
|
# Parsing as an equality range provides the same numeric release rules
|
||||||
|
# used by manifest ranges without accepting an implicit/latest value.
|
||||||
|
validate_version_range(f"=={self.sdk_api_version}", "sdk_api_version")
|
||||||
|
validate_version_range(f"=={self.protocol_version}", "protocol_version")
|
||||||
|
profiles = tuple(require_identifier(value, "runtime profile") for value in self.profiles)
|
||||||
|
if len(profiles) != len(set(profiles)):
|
||||||
|
raise ValueError("runtime profiles must be unique")
|
||||||
|
object.__setattr__(self, "profiles", profiles)
|
||||||
|
if not isinstance(self.features, Mapping):
|
||||||
|
raise ValueError("runtime features must be an object")
|
||||||
|
features: dict[str, str] = {}
|
||||||
|
for name, version in self.features.items():
|
||||||
|
canonical = require_identifier(name, "runtime feature")
|
||||||
|
text = require_string(version, "runtime feature version", max_length=32)
|
||||||
|
validate_version_range(f"=={text}", "runtime feature version")
|
||||||
|
features[canonical] = text
|
||||||
|
object.__setattr__(self, "features", MappingProxyType(features))
|
||||||
|
capabilities = tuple(require_identifier(value, "workload capability") for value in self.workload_capabilities)
|
||||||
|
if len(capabilities) != len(set(capabilities)):
|
||||||
|
raise ValueError("workload_capabilities must be unique")
|
||||||
|
object.__setattr__(self, "workload_capabilities", capabilities)
|
||||||
|
if not isinstance(self.inventory, ResourceInventory):
|
||||||
|
raise ValueError("runtime inventory must be a ResourceInventory")
|
||||||
|
try:
|
||||||
|
trust_modes = tuple(TrustMode(value) for value in self.trust_modes)
|
||||||
|
except (TypeError, ValueError) as error:
|
||||||
|
raise ValueError("runtime trust_modes contain an unsupported value") from error
|
||||||
|
if not trust_modes or len(trust_modes) != len(set(trust_modes)):
|
||||||
|
raise ValueError("runtime trust_modes must be non-empty and unique")
|
||||||
|
object.__setattr__(self, "trust_modes", trust_modes)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class NegotiatedWorkload:
|
||||||
|
manifest: WorkloadManifest
|
||||||
|
optional_fallbacks: Mapping[str, str]
|
||||||
|
sdk_api_version: str
|
||||||
|
protocol_version: str
|
||||||
|
selected_features: Mapping[str, str]
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
if not isinstance(self.manifest, WorkloadManifest):
|
||||||
|
raise ValueError("negotiated manifest must be a WorkloadManifest")
|
||||||
|
object.__setattr__(self, "optional_fallbacks", MappingProxyType(dict(self.optional_fallbacks)))
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"sdk_api_version",
|
||||||
|
require_string(self.sdk_api_version, "negotiated sdk_api_version", max_length=32),
|
||||||
|
)
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"protocol_version",
|
||||||
|
require_string(self.protocol_version, "negotiated protocol_version", max_length=32),
|
||||||
|
)
|
||||||
|
validate_version_range(f"=={self.sdk_api_version}", "negotiated sdk_api_version")
|
||||||
|
validate_version_range(f"=={self.protocol_version}", "negotiated protocol_version")
|
||||||
|
selected: dict[str, str] = {}
|
||||||
|
for name, version in self.selected_features.items():
|
||||||
|
selected[require_identifier(name, "negotiated feature")] = require_string(
|
||||||
|
version,
|
||||||
|
"negotiated feature version",
|
||||||
|
max_length=32,
|
||||||
|
)
|
||||||
|
validate_version_range(
|
||||||
|
f"=={selected[name]}",
|
||||||
|
"negotiated feature version",
|
||||||
|
)
|
||||||
|
object.__setattr__(self, "selected_features", MappingProxyType(selected))
|
||||||
|
|
||||||
|
|
||||||
|
def negotiate_manifest(
|
||||||
|
manifest: WorkloadManifest,
|
||||||
|
runtime: RuntimeCapabilities,
|
||||||
|
) -> NegotiatedWorkload:
|
||||||
|
"""Resolve compatibility before any package handler or planner is invoked."""
|
||||||
|
if not isinstance(manifest, WorkloadManifest) or not isinstance(runtime, RuntimeCapabilities):
|
||||||
|
raise ValueError("negotiation requires WorkloadManifest and RuntimeCapabilities")
|
||||||
|
if runtime.sdk_api_version != SDK_API_VERSION:
|
||||||
|
raise CompatibilityError(
|
||||||
|
"runtime-sdk-mismatch",
|
||||||
|
"runtime SDK declaration does not match this SDK implementation",
|
||||||
|
)
|
||||||
|
if not manifest.sdk_api.contains(runtime.sdk_api_version):
|
||||||
|
raise CompatibilityError("sdk-api-mismatch", "runtime SDK API is outside the manifest range")
|
||||||
|
if not manifest.protocol.contains(runtime.protocol_version):
|
||||||
|
raise CompatibilityError("protocol-mismatch", "runtime protocol is outside the manifest range")
|
||||||
|
missing_profiles = sorted(set(manifest.conformance_profiles) - set(runtime.profiles))
|
||||||
|
if missing_profiles:
|
||||||
|
raise CompatibilityError(
|
||||||
|
"profile-unavailable",
|
||||||
|
"runtime does not support required profiles: " + ", ".join(missing_profiles),
|
||||||
|
)
|
||||||
|
if manifest.workload.name not in runtime.workload_capabilities:
|
||||||
|
raise CompatibilityError(
|
||||||
|
"workload-unavailable",
|
||||||
|
"runtime does not advertise the canonical workload capability",
|
||||||
|
)
|
||||||
|
if manifest.environment.digest not in runtime.inventory.environment_digests:
|
||||||
|
raise CompatibilityError("environment-unavailable", "pinned workload environment is unavailable")
|
||||||
|
for feature in manifest.required_features:
|
||||||
|
version = runtime.features.get(feature.name)
|
||||||
|
if version is None or not feature.versions.contains(version):
|
||||||
|
raise CompatibilityError(
|
||||||
|
"feature-unavailable",
|
||||||
|
f"required feature is unavailable or incompatible: {feature.name}",
|
||||||
|
)
|
||||||
|
fallbacks: dict[str, str] = {}
|
||||||
|
selected_features: dict[str, str] = {}
|
||||||
|
for feature in manifest.required_features:
|
||||||
|
version = runtime.features.get(feature.name)
|
||||||
|
if version is not None and feature.versions.contains(version):
|
||||||
|
selected_features[feature.name] = version
|
||||||
|
for feature in manifest.optional_features:
|
||||||
|
version = runtime.features.get(feature.name)
|
||||||
|
if version is None or not feature.versions.contains(version):
|
||||||
|
if feature.fallback is None:
|
||||||
|
raise CompatibilityError(
|
||||||
|
"optional-feature-unavailable",
|
||||||
|
f"optional feature has no declared fallback: {feature.name}",
|
||||||
|
)
|
||||||
|
fallbacks[feature.name] = feature.fallback
|
||||||
|
else:
|
||||||
|
selected_features[feature.name] = version
|
||||||
|
required_by_shape: dict[StageKind, str] = {
|
||||||
|
StageKind.PLAN: "dynamic-expansion",
|
||||||
|
StageKind.LOOP_CONTROLLER: "bounded-loops",
|
||||||
|
StageKind.STREAM: "stream-checkpoints",
|
||||||
|
StageKind.SERVICE: "services",
|
||||||
|
StageKind.SIDE_EFFECT: "side-effect",
|
||||||
|
}
|
||||||
|
declared_required = {feature.name for feature in manifest.required_features}
|
||||||
|
|
||||||
|
def require_declared(condition: bool, feature: str, message: str) -> None:
|
||||||
|
if condition and feature not in declared_required:
|
||||||
|
raise CompatibilityError("feature-undeclared", message + f" requires {feature}")
|
||||||
|
|
||||||
|
for stage in manifest.workflow.stages:
|
||||||
|
shape_feature = required_by_shape.get(stage.kind)
|
||||||
|
if shape_feature is not None and shape_feature not in declared_required:
|
||||||
|
raise CompatibilityError(
|
||||||
|
"feature-undeclared",
|
||||||
|
f"stage {stage.stage_id} requires declared feature {shape_feature}",
|
||||||
|
)
|
||||||
|
if stage.gang is not None and "gang-leases" not in declared_required:
|
||||||
|
raise CompatibilityError("feature-undeclared", "gang execution requires gang-leases")
|
||||||
|
execution = stage.execution
|
||||||
|
require_declared(
|
||||||
|
execution.process_model is ProcessModel.PROCESS_POOL,
|
||||||
|
"process-pools",
|
||||||
|
f"stage {stage.stage_id} process pool",
|
||||||
|
)
|
||||||
|
require_declared(
|
||||||
|
execution.process_model is ProcessModel.THREAD_POOL,
|
||||||
|
"thread-pools",
|
||||||
|
f"stage {stage.stage_id} thread pool",
|
||||||
|
)
|
||||||
|
require_declared(
|
||||||
|
execution.process_model is ProcessModel.EXTERNAL_RUNTIME,
|
||||||
|
"external-runtimes",
|
||||||
|
f"stage {stage.stage_id} external runtime",
|
||||||
|
)
|
||||||
|
require_declared(
|
||||||
|
execution.max_processes > 1,
|
||||||
|
"multi-process",
|
||||||
|
f"stage {stage.stage_id} multi-process execution",
|
||||||
|
)
|
||||||
|
require_declared(
|
||||||
|
execution.threads_per_process > 1,
|
||||||
|
"python-threads",
|
||||||
|
f"stage {stage.stage_id} Python threading",
|
||||||
|
)
|
||||||
|
require_declared(
|
||||||
|
execution.native_threads > 1,
|
||||||
|
"native-threads",
|
||||||
|
f"stage {stage.stage_id} native threading",
|
||||||
|
)
|
||||||
|
require_declared(
|
||||||
|
execution.nested_parallelism,
|
||||||
|
"nested-parallelism",
|
||||||
|
f"stage {stage.stage_id} nested parallelism",
|
||||||
|
)
|
||||||
|
network_features = {
|
||||||
|
NetworkPolicy.NONE: "network-isolation",
|
||||||
|
NetworkPolicy.COORDINATOR_ARTIFACTS_ONLY: "artifact-network-policy",
|
||||||
|
NetworkPolicy.ALLOWLISTED_EGRESS: "egress-allowlist",
|
||||||
|
}
|
||||||
|
network_feature = network_features.get(execution.network)
|
||||||
|
if network_feature is not None:
|
||||||
|
require_declared(
|
||||||
|
True,
|
||||||
|
network_feature,
|
||||||
|
f"stage {stage.stage_id} network policy",
|
||||||
|
)
|
||||||
|
require_declared(
|
||||||
|
execution.checkpoint.enabled,
|
||||||
|
"checkpoints",
|
||||||
|
f"stage {stage.stage_id} checkpoint policy",
|
||||||
|
)
|
||||||
|
require_declared(
|
||||||
|
stage.retry.max_attempts > 1,
|
||||||
|
"retries",
|
||||||
|
f"stage {stage.stage_id} retry policy",
|
||||||
|
)
|
||||||
|
require_declared(
|
||||||
|
bool(execution.secret_handles),
|
||||||
|
"secret-injection",
|
||||||
|
f"stage {stage.stage_id} secret handles",
|
||||||
|
)
|
||||||
|
resource_sets = (stage.resources,) + (
|
||||||
|
(stage.gang.per_replica_resources,) if stage.gang is not None else ()
|
||||||
|
)
|
||||||
|
for resources in resource_sets:
|
||||||
|
if resources.accelerator_count:
|
||||||
|
if resources.accelerator_mode is AcceleratorMode.EXCLUSIVE_DEVICE:
|
||||||
|
feature = "gpu-exclusive"
|
||||||
|
elif resources.accelerator_mode is AcceleratorMode.PARTITION:
|
||||||
|
feature = "gpu-mig"
|
||||||
|
else:
|
||||||
|
feature = "accelerator-fractional"
|
||||||
|
if feature not in declared_required:
|
||||||
|
raise CompatibilityError(
|
||||||
|
"feature-undeclared",
|
||||||
|
f"accelerator stage requires declared feature {feature}",
|
||||||
|
)
|
||||||
|
errors = resources.eligibility_errors(runtime.inventory)
|
||||||
|
if errors:
|
||||||
|
raise CompatibilityError("resource-ineligible", errors[0])
|
||||||
|
return NegotiatedWorkload(
|
||||||
|
manifest,
|
||||||
|
fallbacks,
|
||||||
|
runtime.sdk_api_version,
|
||||||
|
runtime.protocol_version,
|
||||||
|
selected_features,
|
||||||
|
)
|
||||||
@@ -0,0 +1,380 @@
|
|||||||
|
"""Bounded JSON Schema subset used for SDK v1 public parameters."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import math
|
||||||
|
import re
|
||||||
|
from fractions import Fraction
|
||||||
|
from typing import Mapping, Sequence
|
||||||
|
|
||||||
|
|
||||||
|
_ANNOTATIONS = {
|
||||||
|
"$schema",
|
||||||
|
"title",
|
||||||
|
"description",
|
||||||
|
"default",
|
||||||
|
"examples",
|
||||||
|
"deprecated",
|
||||||
|
"readOnly",
|
||||||
|
"writeOnly",
|
||||||
|
}
|
||||||
|
_KEYWORDS = _ANNOTATIONS | {
|
||||||
|
"type",
|
||||||
|
"enum",
|
||||||
|
"const",
|
||||||
|
"properties",
|
||||||
|
"additionalProperties",
|
||||||
|
"required",
|
||||||
|
"minProperties",
|
||||||
|
"maxProperties",
|
||||||
|
"items",
|
||||||
|
"minItems",
|
||||||
|
"maxItems",
|
||||||
|
"uniqueItems",
|
||||||
|
"minLength",
|
||||||
|
"maxLength",
|
||||||
|
"pattern",
|
||||||
|
"minimum",
|
||||||
|
"maximum",
|
||||||
|
"exclusiveMinimum",
|
||||||
|
"exclusiveMaximum",
|
||||||
|
"multipleOf",
|
||||||
|
"allOf",
|
||||||
|
"anyOf",
|
||||||
|
"oneOf",
|
||||||
|
"not",
|
||||||
|
}
|
||||||
|
_TYPES = {"null", "boolean", "object", "array", "number", "integer", "string"}
|
||||||
|
|
||||||
|
|
||||||
|
class ParameterValidationError(ValueError):
|
||||||
|
"""Sanitized public-parameter schema failure."""
|
||||||
|
|
||||||
|
|
||||||
|
def _schema_error(message: str) -> ValueError:
|
||||||
|
return ValueError("unsupported or invalid parameters_schema: " + message)
|
||||||
|
|
||||||
|
|
||||||
|
def _json_equal(left: object, right: object) -> bool:
|
||||||
|
"""Compare values using the JSON data model rather than Python coercion.
|
||||||
|
|
||||||
|
Python considers ``True == 1`` while JSON has distinct boolean and number
|
||||||
|
types. JSON Schema does, however, treat integral and non-integral syntax for
|
||||||
|
the same mathematical number (for example ``1`` and ``1.0``) as equal.
|
||||||
|
"""
|
||||||
|
if isinstance(left, bool) or isinstance(right, bool):
|
||||||
|
return isinstance(left, bool) and isinstance(right, bool) and left is right
|
||||||
|
if isinstance(left, (int, float)) and isinstance(right, (int, float)):
|
||||||
|
return left == right
|
||||||
|
if left is None or right is None:
|
||||||
|
return left is None and right is None
|
||||||
|
if isinstance(left, str) or isinstance(right, str):
|
||||||
|
return isinstance(left, str) and isinstance(right, str) and left == right
|
||||||
|
if isinstance(left, Mapping) and isinstance(right, Mapping):
|
||||||
|
return set(left) == set(right) and all(
|
||||||
|
_json_equal(left[key], right[key]) for key in left
|
||||||
|
)
|
||||||
|
if isinstance(left, (list, tuple)) and isinstance(right, (list, tuple)):
|
||||||
|
return len(left) == len(right) and all(
|
||||||
|
_json_equal(left_item, right_item)
|
||||||
|
for left_item, right_item in zip(left, right)
|
||||||
|
)
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def _json_key(value: object, depth: int = 0) -> object:
|
||||||
|
"""Build a hashable JSON-type-aware key in linear time."""
|
||||||
|
if depth > 64:
|
||||||
|
raise ValueError("JSON value nesting exceeds 64 levels")
|
||||||
|
if value is None:
|
||||||
|
return ("null",)
|
||||||
|
if isinstance(value, bool):
|
||||||
|
return ("boolean", value)
|
||||||
|
if isinstance(value, (int, float)):
|
||||||
|
return ("number", Fraction(value) if isinstance(value, int) else Fraction.from_float(value))
|
||||||
|
if isinstance(value, str):
|
||||||
|
return ("string", value)
|
||||||
|
if isinstance(value, Mapping):
|
||||||
|
return (
|
||||||
|
"object",
|
||||||
|
tuple(
|
||||||
|
(key, _json_key(child, depth + 1))
|
||||||
|
for key, child in sorted(value.items())
|
||||||
|
),
|
||||||
|
)
|
||||||
|
if isinstance(value, (list, tuple)):
|
||||||
|
return ("array", tuple(_json_key(child, depth + 1) for child in value))
|
||||||
|
raise ValueError("value is not JSON-compatible")
|
||||||
|
|
||||||
|
|
||||||
|
def _validate_safe_pattern(pattern: str) -> None:
|
||||||
|
"""Accept only the v1 linear-time regex subset.
|
||||||
|
|
||||||
|
Groups, alternation, backreferences, and repetition operators are excluded;
|
||||||
|
literals, anchors, character classes, escapes, and ``.`` remain available.
|
||||||
|
"""
|
||||||
|
escaped = False
|
||||||
|
in_class = False
|
||||||
|
for character in pattern:
|
||||||
|
if escaped:
|
||||||
|
if character.isdigit():
|
||||||
|
raise _schema_error("pattern backreferences are not supported")
|
||||||
|
escaped = False
|
||||||
|
continue
|
||||||
|
if character == "\\":
|
||||||
|
escaped = True
|
||||||
|
continue
|
||||||
|
if character == "[" and not in_class:
|
||||||
|
in_class = True
|
||||||
|
continue
|
||||||
|
if character == "]" and in_class:
|
||||||
|
in_class = False
|
||||||
|
continue
|
||||||
|
if not in_class and character in "()|*+?{}":
|
||||||
|
raise _schema_error("pattern uses an unbounded regex operator")
|
||||||
|
if escaped or in_class:
|
||||||
|
# ``re.compile`` will provide the canonical invalid-regex error below.
|
||||||
|
return
|
||||||
|
|
||||||
|
|
||||||
|
def _is_json_multiple(value: int | float, divisor: int | float) -> bool:
|
||||||
|
"""Evaluate ``multipleOf`` without converting arbitrary integers to float."""
|
||||||
|
if isinstance(value, int) and isinstance(divisor, int):
|
||||||
|
return value % divisor == 0
|
||||||
|
value_fraction = Fraction(value) if isinstance(value, int) else Fraction(str(value))
|
||||||
|
divisor_fraction = (
|
||||||
|
Fraction(divisor) if isinstance(divisor, int) else Fraction(str(divisor))
|
||||||
|
)
|
||||||
|
return (value_fraction / divisor_fraction).denominator == 1
|
||||||
|
|
||||||
|
|
||||||
|
def validate_schema_definition(schema: Mapping[str, object], *, _depth: int = 0) -> None:
|
||||||
|
if _depth > 64:
|
||||||
|
raise _schema_error("nesting exceeds 64 levels")
|
||||||
|
if not isinstance(schema, Mapping):
|
||||||
|
raise _schema_error("each schema node must be an object")
|
||||||
|
unknown = set(schema) - _KEYWORDS
|
||||||
|
if unknown:
|
||||||
|
raise _schema_error("unknown keyword " + sorted(unknown)[0])
|
||||||
|
raw_type = schema.get("type")
|
||||||
|
if raw_type is not None:
|
||||||
|
declared = (raw_type,) if isinstance(raw_type, str) else raw_type
|
||||||
|
if not isinstance(declared, (list, tuple)) or not declared:
|
||||||
|
raise _schema_error("type must be a string or non-empty array")
|
||||||
|
if any(not isinstance(value, str) or value not in _TYPES for value in declared):
|
||||||
|
raise _schema_error("type contains an unsupported JSON type")
|
||||||
|
if len(declared) != len(set(declared)):
|
||||||
|
raise _schema_error("type alternatives must be unique")
|
||||||
|
properties = schema.get("properties")
|
||||||
|
if properties is not None:
|
||||||
|
if not isinstance(properties, Mapping) or any(not isinstance(name, str) for name in properties):
|
||||||
|
raise _schema_error("properties must be an object")
|
||||||
|
for child in properties.values():
|
||||||
|
validate_schema_definition(child, _depth=_depth + 1) # type: ignore[arg-type]
|
||||||
|
additional = schema.get("additionalProperties")
|
||||||
|
if additional is not None and not isinstance(additional, (bool, Mapping)):
|
||||||
|
raise _schema_error("additionalProperties must be a boolean or schema")
|
||||||
|
if isinstance(additional, Mapping):
|
||||||
|
validate_schema_definition(additional, _depth=_depth + 1)
|
||||||
|
required = schema.get("required")
|
||||||
|
if required is not None:
|
||||||
|
if not isinstance(required, (list, tuple)) or any(not isinstance(name, str) for name in required):
|
||||||
|
raise _schema_error("required must be an array of strings")
|
||||||
|
if len(required) != len(set(required)):
|
||||||
|
raise _schema_error("required names must be unique")
|
||||||
|
for keyword in ("items", "not"):
|
||||||
|
child = schema.get(keyword)
|
||||||
|
if child is not None:
|
||||||
|
validate_schema_definition(child, _depth=_depth + 1) # type: ignore[arg-type]
|
||||||
|
for keyword in ("allOf", "anyOf", "oneOf"):
|
||||||
|
children = schema.get(keyword)
|
||||||
|
if children is None:
|
||||||
|
continue
|
||||||
|
if not isinstance(children, (list, tuple)) or not children:
|
||||||
|
raise _schema_error(f"{keyword} must be a non-empty array")
|
||||||
|
for child in children:
|
||||||
|
validate_schema_definition(child, _depth=_depth + 1) # type: ignore[arg-type]
|
||||||
|
enum = schema.get("enum")
|
||||||
|
if enum is not None and (not isinstance(enum, (list, tuple)) or not enum):
|
||||||
|
raise _schema_error("enum must be a non-empty array")
|
||||||
|
if isinstance(enum, (list, tuple)):
|
||||||
|
seen_enum: set[object] = set()
|
||||||
|
for item in enum:
|
||||||
|
key = _json_key(item)
|
||||||
|
if key in seen_enum:
|
||||||
|
raise _schema_error("enum values must be unique")
|
||||||
|
seen_enum.add(key)
|
||||||
|
for keyword in (
|
||||||
|
"minProperties", "maxProperties", "minItems", "maxItems", "minLength", "maxLength"
|
||||||
|
):
|
||||||
|
value = schema.get(keyword)
|
||||||
|
if value is not None and (isinstance(value, bool) or not isinstance(value, int) or value < 0):
|
||||||
|
raise _schema_error(f"{keyword} must be a non-negative integer")
|
||||||
|
for minimum, maximum in (
|
||||||
|
("minProperties", "maxProperties"),
|
||||||
|
("minItems", "maxItems"),
|
||||||
|
("minLength", "maxLength"),
|
||||||
|
):
|
||||||
|
if minimum in schema and maximum in schema and schema[minimum] > schema[maximum]: # type: ignore[operator]
|
||||||
|
raise _schema_error(f"{minimum} must not exceed {maximum}")
|
||||||
|
for keyword in (
|
||||||
|
"minimum", "maximum", "exclusiveMinimum", "exclusiveMaximum", "multipleOf"
|
||||||
|
):
|
||||||
|
value = schema.get(keyword)
|
||||||
|
if value is not None and (
|
||||||
|
isinstance(value, bool)
|
||||||
|
or not isinstance(value, (int, float))
|
||||||
|
or (isinstance(value, float) and not math.isfinite(value))
|
||||||
|
):
|
||||||
|
raise _schema_error(f"{keyword} must be a finite number")
|
||||||
|
if "multipleOf" in schema and schema["multipleOf"] <= 0: # type: ignore[operator]
|
||||||
|
raise _schema_error("multipleOf must be positive")
|
||||||
|
pattern = schema.get("pattern")
|
||||||
|
if pattern is not None:
|
||||||
|
if not isinstance(pattern, str) or len(pattern) > 1024:
|
||||||
|
raise _schema_error("pattern must be a string of at most 1024 characters")
|
||||||
|
_validate_safe_pattern(pattern)
|
||||||
|
try:
|
||||||
|
re.compile(pattern)
|
||||||
|
except re.error as error:
|
||||||
|
raise _schema_error("pattern is not a valid regular expression") from error
|
||||||
|
for keyword in ("uniqueItems", "deprecated", "readOnly", "writeOnly"):
|
||||||
|
if keyword in schema and not isinstance(schema[keyword], bool):
|
||||||
|
raise _schema_error(f"{keyword} must be a boolean")
|
||||||
|
|
||||||
|
|
||||||
|
def _type_matches(value: object, expected: str) -> bool:
|
||||||
|
if expected == "null":
|
||||||
|
return value is None
|
||||||
|
if expected == "boolean":
|
||||||
|
return isinstance(value, bool)
|
||||||
|
if expected == "object":
|
||||||
|
return isinstance(value, Mapping)
|
||||||
|
if expected == "array":
|
||||||
|
return isinstance(value, (list, tuple))
|
||||||
|
if expected == "integer":
|
||||||
|
return isinstance(value, int) and not isinstance(value, bool)
|
||||||
|
if expected == "number":
|
||||||
|
return isinstance(value, (int, float)) and not isinstance(value, bool)
|
||||||
|
if expected == "string":
|
||||||
|
return isinstance(value, str)
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def _failure(path: str, reason: str) -> ParameterValidationError:
|
||||||
|
return ParameterValidationError(f"job parameters violate their schema at {path}: {reason}")
|
||||||
|
|
||||||
|
|
||||||
|
def validate_parameter_instance(
|
||||||
|
value: object,
|
||||||
|
schema: Mapping[str, object],
|
||||||
|
*,
|
||||||
|
path: str = "$",
|
||||||
|
_depth: int = 0,
|
||||||
|
) -> None:
|
||||||
|
if _depth > 64:
|
||||||
|
raise _failure(path, "nesting exceeds 64 levels")
|
||||||
|
raw_type = schema.get("type")
|
||||||
|
if raw_type is not None:
|
||||||
|
expected = (raw_type,) if isinstance(raw_type, str) else tuple(raw_type) # type: ignore[arg-type]
|
||||||
|
if not any(_type_matches(value, item) for item in expected):
|
||||||
|
raise _failure(path, "type mismatch")
|
||||||
|
if "enum" in schema and not any(
|
||||||
|
_json_equal(value, candidate) for candidate in schema["enum"] # type: ignore[union-attr]
|
||||||
|
):
|
||||||
|
raise _failure(path, "value is outside enum")
|
||||||
|
if "const" in schema and not _json_equal(value, schema["const"]):
|
||||||
|
raise _failure(path, "value does not match const")
|
||||||
|
for keyword in ("allOf", "anyOf", "oneOf"):
|
||||||
|
children = schema.get(keyword)
|
||||||
|
if children is None:
|
||||||
|
continue
|
||||||
|
matches = 0
|
||||||
|
for child in children: # type: ignore[union-attr]
|
||||||
|
try:
|
||||||
|
validate_parameter_instance(value, child, path=path, _depth=_depth + 1)
|
||||||
|
except ParameterValidationError:
|
||||||
|
continue
|
||||||
|
matches += 1
|
||||||
|
if keyword == "allOf" and matches != len(children): # type: ignore[arg-type]
|
||||||
|
raise _failure(path, "allOf did not match")
|
||||||
|
if keyword == "anyOf" and matches == 0:
|
||||||
|
raise _failure(path, "anyOf did not match")
|
||||||
|
if keyword == "oneOf" and matches != 1:
|
||||||
|
raise _failure(path, "oneOf did not match exactly once")
|
||||||
|
excluded = schema.get("not")
|
||||||
|
if excluded is not None:
|
||||||
|
try:
|
||||||
|
validate_parameter_instance(value, excluded, path=path, _depth=_depth + 1) # type: ignore[arg-type]
|
||||||
|
except ParameterValidationError:
|
||||||
|
pass
|
||||||
|
else:
|
||||||
|
raise _failure(path, "value matches a forbidden schema")
|
||||||
|
if isinstance(value, Mapping):
|
||||||
|
required = schema.get("required", ())
|
||||||
|
missing = set(required) - set(value) # type: ignore[arg-type]
|
||||||
|
if missing:
|
||||||
|
raise _failure(path, "missing required field " + sorted(missing)[0])
|
||||||
|
minimum = schema.get("minProperties")
|
||||||
|
maximum = schema.get("maxProperties")
|
||||||
|
if minimum is not None and len(value) < minimum: # type: ignore[operator]
|
||||||
|
raise _failure(path, "too few properties")
|
||||||
|
if maximum is not None and len(value) > maximum: # type: ignore[operator]
|
||||||
|
raise _failure(path, "too many properties")
|
||||||
|
properties = schema.get("properties", {})
|
||||||
|
additional = schema.get("additionalProperties", True)
|
||||||
|
for name, child in value.items():
|
||||||
|
if name in properties: # type: ignore[operator]
|
||||||
|
validate_parameter_instance(
|
||||||
|
child,
|
||||||
|
properties[name], # type: ignore[index]
|
||||||
|
path=f"{path}.{name}",
|
||||||
|
_depth=_depth + 1,
|
||||||
|
)
|
||||||
|
elif additional is False:
|
||||||
|
raise _failure(path, f"unknown field {name}")
|
||||||
|
elif isinstance(additional, Mapping):
|
||||||
|
validate_parameter_instance(child, additional, path=f"{path}.{name}", _depth=_depth + 1)
|
||||||
|
if isinstance(value, (list, tuple)):
|
||||||
|
minimum = schema.get("minItems")
|
||||||
|
maximum = schema.get("maxItems")
|
||||||
|
if minimum is not None and len(value) < minimum: # type: ignore[operator]
|
||||||
|
raise _failure(path, "too few items")
|
||||||
|
if maximum is not None and len(value) > maximum: # type: ignore[operator]
|
||||||
|
raise _failure(path, "too many items")
|
||||||
|
if schema.get("uniqueItems"):
|
||||||
|
seen_items: set[object] = set()
|
||||||
|
for item in value:
|
||||||
|
key = _json_key(item)
|
||||||
|
if key in seen_items:
|
||||||
|
raise _failure(path, "items must be unique")
|
||||||
|
seen_items.add(key)
|
||||||
|
child_schema = schema.get("items")
|
||||||
|
if child_schema is not None:
|
||||||
|
for index, item in enumerate(value):
|
||||||
|
validate_parameter_instance(
|
||||||
|
item,
|
||||||
|
child_schema, # type: ignore[arg-type]
|
||||||
|
path=f"{path}[{index}]",
|
||||||
|
_depth=_depth + 1,
|
||||||
|
)
|
||||||
|
if isinstance(value, str):
|
||||||
|
if "minLength" in schema and len(value) < schema["minLength"]: # type: ignore[operator]
|
||||||
|
raise _failure(path, "string is too short")
|
||||||
|
if "maxLength" in schema and len(value) > schema["maxLength"]: # type: ignore[operator]
|
||||||
|
raise _failure(path, "string is too long")
|
||||||
|
if "pattern" in schema and re.search(schema["pattern"], value) is None: # type: ignore[arg-type]
|
||||||
|
raise _failure(path, "string does not match pattern")
|
||||||
|
if isinstance(value, (int, float)) and not isinstance(value, bool):
|
||||||
|
checks = (
|
||||||
|
("minimum", lambda actual, bound: actual >= bound),
|
||||||
|
("maximum", lambda actual, bound: actual <= bound),
|
||||||
|
("exclusiveMinimum", lambda actual, bound: actual > bound),
|
||||||
|
("exclusiveMaximum", lambda actual, bound: actual < bound),
|
||||||
|
)
|
||||||
|
for keyword, predicate in checks:
|
||||||
|
if keyword in schema and not predicate(value, schema[keyword]):
|
||||||
|
raise _failure(path, f"number violates {keyword}")
|
||||||
|
if "multipleOf" in schema:
|
||||||
|
if not _is_json_multiple(value, schema["multipleOf"]): # type: ignore[arg-type]
|
||||||
|
raise _failure(path, "number violates multipleOf")
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,614 @@
|
|||||||
|
"""Versioned workflow DAG and bounded advanced-stage declarations."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from enum import Enum
|
||||||
|
from types import MappingProxyType
|
||||||
|
from typing import Mapping
|
||||||
|
|
||||||
|
from ._validation import (
|
||||||
|
enum_value,
|
||||||
|
require_entry_point,
|
||||||
|
require_exact_keys,
|
||||||
|
require_identifier,
|
||||||
|
require_nonnegative_int,
|
||||||
|
require_positive_int,
|
||||||
|
require_schema_version,
|
||||||
|
require_string,
|
||||||
|
)
|
||||||
|
from .artifacts import PortSpec
|
||||||
|
from .execution import ExecutionProfile, NetworkPolicy, RetryPolicy
|
||||||
|
from .identity import ComponentRef, SchemaRef, WORKFLOW_SCHEMA_VERSION
|
||||||
|
from .resources import ResourceRequirements
|
||||||
|
|
||||||
|
|
||||||
|
class StageKind(str, Enum):
|
||||||
|
PLAN = "plan"
|
||||||
|
MAP = "map"
|
||||||
|
REDUCE = "reduce"
|
||||||
|
VERIFY = "verify"
|
||||||
|
LOOP_CONTROLLER = "loop-controller"
|
||||||
|
STREAM = "stream"
|
||||||
|
SERVICE = "service"
|
||||||
|
SIDE_EFFECT = "side-effect"
|
||||||
|
|
||||||
|
|
||||||
|
class WorkflowFailurePolicy(str, Enum):
|
||||||
|
FAIL_FAST = "fail_fast"
|
||||||
|
CONTINUE_INDEPENDENT = "continue_independent"
|
||||||
|
ALLOW_PARTIAL = "allow_partial"
|
||||||
|
COMPENSATE = "compensate"
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class LoopSpec:
|
||||||
|
state_schema: SchemaRef
|
||||||
|
max_iterations: int
|
||||||
|
max_wall_seconds: int
|
||||||
|
body_workflow: str
|
||||||
|
continue_when: ComponentRef
|
||||||
|
checkpoint_every: int
|
||||||
|
on_limit: str = "fail"
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
if not isinstance(self.state_schema, SchemaRef):
|
||||||
|
raise ValueError("loop state_schema must be a SchemaRef")
|
||||||
|
object.__setattr__(self, "max_iterations", require_positive_int(self.max_iterations, "loop.max_iterations"))
|
||||||
|
object.__setattr__(self, "max_wall_seconds", require_positive_int(self.max_wall_seconds, "loop.max_wall_seconds"))
|
||||||
|
object.__setattr__(self, "body_workflow", require_identifier(self.body_workflow, "loop.body_workflow"))
|
||||||
|
if not isinstance(self.continue_when, ComponentRef):
|
||||||
|
raise ValueError("loop continue_when must be a ComponentRef")
|
||||||
|
object.__setattr__(self, "checkpoint_every", require_positive_int(self.checkpoint_every, "loop.checkpoint_every"))
|
||||||
|
if self.checkpoint_every > self.max_iterations:
|
||||||
|
raise ValueError("loop checkpoint_every must not exceed max_iterations")
|
||||||
|
if self.on_limit not in {"fail", "accept-best", "return-inconclusive"}:
|
||||||
|
raise ValueError("loop on_limit must be fail, accept-best, or return-inconclusive")
|
||||||
|
|
||||||
|
def to_dict(self) -> dict[str, object]:
|
||||||
|
return {
|
||||||
|
"state_schema": self.state_schema.canonical,
|
||||||
|
"max_iterations": self.max_iterations,
|
||||||
|
"max_wall_seconds": self.max_wall_seconds,
|
||||||
|
"body_workflow": self.body_workflow,
|
||||||
|
"continue_when": self.continue_when.canonical,
|
||||||
|
"checkpoint_every": self.checkpoint_every,
|
||||||
|
"on_limit": self.on_limit,
|
||||||
|
}
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_dict(cls, value: object) -> "LoopSpec":
|
||||||
|
if not isinstance(value, Mapping):
|
||||||
|
raise ValueError("loop specification must be an object")
|
||||||
|
fields = {
|
||||||
|
"state_schema", "max_iterations", "max_wall_seconds", "body_workflow",
|
||||||
|
"continue_when", "checkpoint_every", "on_limit",
|
||||||
|
}
|
||||||
|
require_exact_keys(value, fields, "loop specification")
|
||||||
|
return cls(
|
||||||
|
state_schema=SchemaRef.from_dict(value["state_schema"]),
|
||||||
|
max_iterations=value["max_iterations"], # type: ignore[arg-type]
|
||||||
|
max_wall_seconds=value["max_wall_seconds"], # type: ignore[arg-type]
|
||||||
|
body_workflow=value["body_workflow"], # type: ignore[arg-type]
|
||||||
|
continue_when=ComponentRef.from_dict(value["continue_when"]),
|
||||||
|
checkpoint_every=value["checkpoint_every"], # type: ignore[arg-type]
|
||||||
|
on_limit=value["on_limit"], # type: ignore[arg-type]
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class StreamSpec:
|
||||||
|
source: str
|
||||||
|
partitioning: str
|
||||||
|
checkpoint_schema: SchemaRef
|
||||||
|
window_seconds: int
|
||||||
|
watermark_seconds: int
|
||||||
|
backpressure_limit: int
|
||||||
|
delivery_guarantee: str
|
||||||
|
max_windows: int
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
object.__setattr__(self, "source", require_identifier(self.source, "stream.source"))
|
||||||
|
object.__setattr__(self, "partitioning", require_identifier(self.partitioning, "stream.partitioning"))
|
||||||
|
if not isinstance(self.checkpoint_schema, SchemaRef):
|
||||||
|
raise ValueError("stream checkpoint_schema must be a SchemaRef")
|
||||||
|
object.__setattr__(self, "window_seconds", require_positive_int(self.window_seconds, "stream.window_seconds"))
|
||||||
|
object.__setattr__(self, "watermark_seconds", require_nonnegative_int(self.watermark_seconds, "stream.watermark_seconds"))
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"backpressure_limit",
|
||||||
|
require_positive_int(self.backpressure_limit, "stream.backpressure_limit"),
|
||||||
|
)
|
||||||
|
if self.delivery_guarantee not in {"at_least_once", "exactly_once"}:
|
||||||
|
raise ValueError("stream delivery_guarantee must be at_least_once or exactly_once")
|
||||||
|
object.__setattr__(self, "max_windows", require_positive_int(self.max_windows, "stream.max_windows"))
|
||||||
|
|
||||||
|
def to_dict(self) -> dict[str, object]:
|
||||||
|
return {
|
||||||
|
"source": self.source,
|
||||||
|
"partitioning": self.partitioning,
|
||||||
|
"checkpoint_schema": self.checkpoint_schema.canonical,
|
||||||
|
"window_seconds": self.window_seconds,
|
||||||
|
"watermark_seconds": self.watermark_seconds,
|
||||||
|
"backpressure_limit": self.backpressure_limit,
|
||||||
|
"delivery_guarantee": self.delivery_guarantee,
|
||||||
|
"max_windows": self.max_windows,
|
||||||
|
}
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_dict(cls, value: object) -> "StreamSpec":
|
||||||
|
if not isinstance(value, Mapping):
|
||||||
|
raise ValueError("stream specification must be an object")
|
||||||
|
fields = {
|
||||||
|
"source", "partitioning", "checkpoint_schema", "window_seconds",
|
||||||
|
"watermark_seconds", "backpressure_limit", "delivery_guarantee", "max_windows",
|
||||||
|
}
|
||||||
|
require_exact_keys(value, fields, "stream specification")
|
||||||
|
return cls(
|
||||||
|
source=value["source"], # type: ignore[arg-type]
|
||||||
|
partitioning=value["partitioning"], # type: ignore[arg-type]
|
||||||
|
checkpoint_schema=SchemaRef.from_dict(value["checkpoint_schema"]),
|
||||||
|
window_seconds=value["window_seconds"], # type: ignore[arg-type]
|
||||||
|
watermark_seconds=value["watermark_seconds"], # type: ignore[arg-type]
|
||||||
|
backpressure_limit=value["backpressure_limit"], # type: ignore[arg-type]
|
||||||
|
delivery_guarantee=value["delivery_guarantee"], # type: ignore[arg-type]
|
||||||
|
max_windows=value["max_windows"], # type: ignore[arg-type]
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class GangSpec:
|
||||||
|
replicas: int
|
||||||
|
per_replica_resources: ResourceRequirements
|
||||||
|
same_topology_group: bool = False
|
||||||
|
bandwidth_class: str | None = None
|
||||||
|
failure_mode: str = "fail_all"
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
object.__setattr__(self, "replicas", require_positive_int(self.replicas, "gang.replicas"))
|
||||||
|
if self.replicas < 2:
|
||||||
|
raise ValueError("gang execution requires at least two replicas")
|
||||||
|
if not isinstance(self.per_replica_resources, ResourceRequirements):
|
||||||
|
raise ValueError("gang per_replica_resources must be ResourceRequirements")
|
||||||
|
if not isinstance(self.same_topology_group, bool):
|
||||||
|
raise ValueError("gang same_topology_group must be a boolean")
|
||||||
|
if self.bandwidth_class is not None:
|
||||||
|
object.__setattr__(self, "bandwidth_class", require_identifier(self.bandwidth_class, "gang.bandwidth_class"))
|
||||||
|
if self.failure_mode != "fail_all":
|
||||||
|
raise ValueError("SDK v1 gang failure_mode must be fail_all")
|
||||||
|
|
||||||
|
def to_dict(self) -> dict[str, object]:
|
||||||
|
return {
|
||||||
|
"replicas": self.replicas,
|
||||||
|
"per_replica_resources": self.per_replica_resources.to_dict(),
|
||||||
|
"same_topology_group": self.same_topology_group,
|
||||||
|
"bandwidth_class": self.bandwidth_class,
|
||||||
|
"failure_mode": self.failure_mode,
|
||||||
|
}
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_dict(cls, value: object) -> "GangSpec":
|
||||||
|
if not isinstance(value, Mapping):
|
||||||
|
raise ValueError("gang specification must be an object")
|
||||||
|
fields = {
|
||||||
|
"replicas", "per_replica_resources", "same_topology_group",
|
||||||
|
"bandwidth_class", "failure_mode",
|
||||||
|
}
|
||||||
|
require_exact_keys(value, fields, "gang specification")
|
||||||
|
return cls(
|
||||||
|
replicas=value["replicas"], # type: ignore[arg-type]
|
||||||
|
per_replica_resources=ResourceRequirements.from_dict(value["per_replica_resources"]),
|
||||||
|
same_topology_group=value["same_topology_group"], # type: ignore[arg-type]
|
||||||
|
bandwidth_class=value["bandwidth_class"], # type: ignore[arg-type]
|
||||||
|
failure_mode=value["failure_mode"], # type: ignore[arg-type]
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class SideEffectSpec:
|
||||||
|
target: str
|
||||||
|
idempotency_key_parameter: str
|
||||||
|
credential_scope: str
|
||||||
|
compensation: str
|
||||||
|
manual_approval: bool = True
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
object.__setattr__(self, "target", require_identifier(self.target, "side_effect.target"))
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"idempotency_key_parameter",
|
||||||
|
require_identifier(self.idempotency_key_parameter, "side_effect.idempotency_key_parameter"),
|
||||||
|
)
|
||||||
|
object.__setattr__(self, "credential_scope", require_identifier(self.credential_scope, "side_effect.credential_scope"))
|
||||||
|
object.__setattr__(self, "compensation", require_identifier(self.compensation, "side_effect.compensation"))
|
||||||
|
if not isinstance(self.manual_approval, bool):
|
||||||
|
raise ValueError("side_effect.manual_approval must be a boolean")
|
||||||
|
|
||||||
|
def to_dict(self) -> dict[str, object]:
|
||||||
|
return {
|
||||||
|
"target": self.target,
|
||||||
|
"idempotency_key_parameter": self.idempotency_key_parameter,
|
||||||
|
"credential_scope": self.credential_scope,
|
||||||
|
"compensation": self.compensation,
|
||||||
|
"manual_approval": self.manual_approval,
|
||||||
|
}
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_dict(cls, value: object) -> "SideEffectSpec":
|
||||||
|
if not isinstance(value, Mapping):
|
||||||
|
raise ValueError("side-effect specification must be an object")
|
||||||
|
fields = {
|
||||||
|
"target", "idempotency_key_parameter", "credential_scope", "compensation", "manual_approval",
|
||||||
|
}
|
||||||
|
require_exact_keys(value, fields, "side-effect specification")
|
||||||
|
return cls(**value) # type: ignore[arg-type]
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class PortRef:
|
||||||
|
"""A stage port, or an external workflow input when ``stage_id`` is None."""
|
||||||
|
|
||||||
|
port: str
|
||||||
|
stage_id: str | None = None
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
object.__setattr__(self, "port", require_identifier(self.port, "port reference"))
|
||||||
|
if self.stage_id is not None:
|
||||||
|
object.__setattr__(self, "stage_id", require_identifier(self.stage_id, "stage reference"))
|
||||||
|
|
||||||
|
def to_dict(self) -> dict[str, object]:
|
||||||
|
return {"stage_id": self.stage_id, "port": self.port}
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_dict(cls, value: object) -> "PortRef":
|
||||||
|
if not isinstance(value, Mapping):
|
||||||
|
raise ValueError("port reference must be an object")
|
||||||
|
require_exact_keys(value, {"stage_id", "port"}, "port reference")
|
||||||
|
return cls(stage_id=value["stage_id"], port=value["port"]) # type: ignore[arg-type]
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class ArtifactEdge:
|
||||||
|
source: PortRef
|
||||||
|
target: PortRef
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
if not isinstance(self.source, PortRef) or not isinstance(self.target, PortRef):
|
||||||
|
raise ValueError("artifact edge endpoints must be PortRef values")
|
||||||
|
if self.target.stage_id is None:
|
||||||
|
raise ValueError("artifact edge target must be a stage input")
|
||||||
|
|
||||||
|
def to_dict(self) -> dict[str, object]:
|
||||||
|
return {"source": self.source.to_dict(), "target": self.target.to_dict()}
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_dict(cls, value: object) -> "ArtifactEdge":
|
||||||
|
if not isinstance(value, Mapping):
|
||||||
|
raise ValueError("artifact edge must be an object")
|
||||||
|
require_exact_keys(value, {"source", "target"}, "artifact edge")
|
||||||
|
return cls(source=PortRef.from_dict(value["source"]), target=PortRef.from_dict(value["target"]))
|
||||||
|
|
||||||
|
|
||||||
|
def _port_mapping(value: Mapping[str, PortSpec], field: str) -> Mapping[str, PortSpec]:
|
||||||
|
if not isinstance(value, Mapping):
|
||||||
|
raise ValueError(f"{field} must be an object")
|
||||||
|
ports: dict[str, PortSpec] = {}
|
||||||
|
for name, port in value.items():
|
||||||
|
canonical = require_identifier(name, f"{field} port")
|
||||||
|
if not isinstance(port, PortSpec):
|
||||||
|
raise ValueError(f"{field} values must be PortSpec values")
|
||||||
|
ports[canonical] = port
|
||||||
|
return MappingProxyType(ports)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class StageSpec:
|
||||||
|
stage_id: str
|
||||||
|
kind: StageKind
|
||||||
|
entry_point: str
|
||||||
|
needs: tuple[str, ...]
|
||||||
|
inputs: Mapping[str, PortSpec]
|
||||||
|
outputs: Mapping[str, PortSpec]
|
||||||
|
parameter_names: tuple[str, ...]
|
||||||
|
resources: ResourceRequirements
|
||||||
|
execution: ExecutionProfile
|
||||||
|
retry: RetryPolicy
|
||||||
|
verifier: ComponentRef | None = None
|
||||||
|
trust_modes: tuple[str, ...] = ("trusted",)
|
||||||
|
max_fan_out: int = 1
|
||||||
|
cacheable: bool = False
|
||||||
|
loop: LoopSpec | None = None
|
||||||
|
stream: StreamSpec | None = None
|
||||||
|
gang: GangSpec | None = None
|
||||||
|
side_effect: SideEffectSpec | None = None
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
object.__setattr__(self, "stage_id", require_identifier(self.stage_id, "stage_id"))
|
||||||
|
object.__setattr__(self, "kind", enum_value(StageKind, self.kind, "stage.kind"))
|
||||||
|
object.__setattr__(self, "entry_point", require_entry_point(self.entry_point, "stage.entry_point"))
|
||||||
|
needs = tuple(require_identifier(value, "stage.needs") for value in self.needs)
|
||||||
|
if self.stage_id in needs or len(needs) != len(set(needs)):
|
||||||
|
raise ValueError("stage.needs must contain unique other stage IDs")
|
||||||
|
object.__setattr__(self, "needs", needs)
|
||||||
|
object.__setattr__(self, "inputs", _port_mapping(self.inputs, "stage.inputs"))
|
||||||
|
object.__setattr__(self, "outputs", _port_mapping(self.outputs, "stage.outputs"))
|
||||||
|
if not self.outputs:
|
||||||
|
raise ValueError("a stage must declare at least one output port")
|
||||||
|
names = tuple(require_identifier(value, "parameter_name") for value in self.parameter_names)
|
||||||
|
if len(names) != len(set(names)):
|
||||||
|
raise ValueError("parameter_names must be unique")
|
||||||
|
object.__setattr__(self, "parameter_names", names)
|
||||||
|
if not isinstance(self.resources, ResourceRequirements):
|
||||||
|
raise ValueError("stage.resources must be ResourceRequirements")
|
||||||
|
if not isinstance(self.execution, ExecutionProfile):
|
||||||
|
raise ValueError("stage.execution must be ExecutionProfile")
|
||||||
|
self.execution.validate_resources(self.resources)
|
||||||
|
if not isinstance(self.retry, RetryPolicy):
|
||||||
|
raise ValueError("stage.retry must be RetryPolicy")
|
||||||
|
if self.verifier is not None and not isinstance(self.verifier, ComponentRef):
|
||||||
|
raise ValueError("stage.verifier must be a ComponentRef")
|
||||||
|
modes = tuple(require_identifier(value, "trust_mode") for value in self.trust_modes)
|
||||||
|
if not modes or len(modes) != len(set(modes)):
|
||||||
|
raise ValueError("stage.trust_modes must be non-empty and unique")
|
||||||
|
if not set(modes).issubset({"trusted", "verified", "untrusted_quorum"}):
|
||||||
|
raise ValueError("stage.trust_modes contains an unsupported trust mode")
|
||||||
|
object.__setattr__(self, "trust_modes", modes)
|
||||||
|
object.__setattr__(self, "max_fan_out", require_positive_int(self.max_fan_out, "stage.max_fan_out"))
|
||||||
|
if not isinstance(self.cacheable, bool):
|
||||||
|
raise ValueError("stage.cacheable must be a boolean")
|
||||||
|
advanced = {
|
||||||
|
StageKind.LOOP_CONTROLLER: self.loop,
|
||||||
|
StageKind.STREAM: self.stream,
|
||||||
|
StageKind.SIDE_EFFECT: self.side_effect,
|
||||||
|
}
|
||||||
|
expected_types = {
|
||||||
|
StageKind.LOOP_CONTROLLER: LoopSpec,
|
||||||
|
StageKind.STREAM: StreamSpec,
|
||||||
|
StageKind.SIDE_EFFECT: SideEffectSpec,
|
||||||
|
}
|
||||||
|
for kind, declaration in advanced.items():
|
||||||
|
if self.kind is kind and declaration is None:
|
||||||
|
raise ValueError(f"{kind.value} stage requires its bounded declaration")
|
||||||
|
if self.kind is not kind and declaration is not None:
|
||||||
|
raise ValueError(f"{kind.value} declaration is valid only for a {kind.value} stage")
|
||||||
|
if declaration is not None and not isinstance(declaration, expected_types[kind]):
|
||||||
|
raise ValueError(f"{kind.value} declaration has the wrong type")
|
||||||
|
if self.gang is not None and not isinstance(self.gang, GangSpec):
|
||||||
|
raise ValueError("stage.gang must be a GangSpec")
|
||||||
|
if self.gang is not None:
|
||||||
|
self.execution.validate_resources(self.gang.per_replica_resources)
|
||||||
|
if self.kind is StageKind.SIDE_EFFECT:
|
||||||
|
raise ValueError("side-effect stages cannot use gang execution")
|
||||||
|
if self.kind is StageKind.SIDE_EFFECT:
|
||||||
|
if self.cacheable:
|
||||||
|
raise ValueError("side-effect stages cannot be cached")
|
||||||
|
if self.execution.network not in {NetworkPolicy.ALLOWLISTED_EGRESS, NetworkPolicy.TRUSTED}:
|
||||||
|
raise ValueError("side-effect stages require explicit egress")
|
||||||
|
assert self.side_effect is not None
|
||||||
|
if self.side_effect.idempotency_key_parameter not in self.parameter_names:
|
||||||
|
raise ValueError("side-effect idempotency key must be projected into the stage")
|
||||||
|
|
||||||
|
def to_dict(self) -> dict[str, object]:
|
||||||
|
return {
|
||||||
|
"stage_id": self.stage_id,
|
||||||
|
"kind": self.kind.value,
|
||||||
|
"entry_point": self.entry_point,
|
||||||
|
"needs": list(self.needs),
|
||||||
|
"inputs": {name: port.to_dict() for name, port in self.inputs.items()},
|
||||||
|
"outputs": {name: port.to_dict() for name, port in self.outputs.items()},
|
||||||
|
"parameter_names": list(self.parameter_names),
|
||||||
|
"resources": self.resources.to_dict(),
|
||||||
|
"execution": self.execution.to_dict(),
|
||||||
|
"retry": self.retry.to_dict(),
|
||||||
|
"verifier": self.verifier.canonical if self.verifier is not None else None,
|
||||||
|
"trust_modes": list(self.trust_modes),
|
||||||
|
"max_fan_out": self.max_fan_out,
|
||||||
|
"cacheable": self.cacheable,
|
||||||
|
"loop": self.loop.to_dict() if self.loop is not None else None,
|
||||||
|
"stream": self.stream.to_dict() if self.stream is not None else None,
|
||||||
|
"gang": self.gang.to_dict() if self.gang is not None else None,
|
||||||
|
"side_effect": self.side_effect.to_dict() if self.side_effect is not None else None,
|
||||||
|
}
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_dict(cls, value: object) -> "StageSpec":
|
||||||
|
if not isinstance(value, Mapping):
|
||||||
|
raise ValueError("stage specification must be an object")
|
||||||
|
fields = {
|
||||||
|
"stage_id", "kind", "entry_point", "needs", "inputs", "outputs",
|
||||||
|
"parameter_names", "resources", "execution", "retry", "verifier",
|
||||||
|
"trust_modes", "max_fan_out", "cacheable", "loop", "stream", "gang", "side_effect",
|
||||||
|
}
|
||||||
|
require_exact_keys(value, fields, "stage specification")
|
||||||
|
arrays = (value["needs"], value["parameter_names"], value["trust_modes"])
|
||||||
|
if any(not isinstance(item, list) for item in arrays):
|
||||||
|
raise ValueError("stage needs, parameter_names, and trust_modes must be arrays")
|
||||||
|
inputs, outputs = value["inputs"], value["outputs"]
|
||||||
|
if not isinstance(inputs, Mapping) or not isinstance(outputs, Mapping):
|
||||||
|
raise ValueError("stage inputs and outputs must be objects")
|
||||||
|
return cls(
|
||||||
|
stage_id=value["stage_id"], # type: ignore[arg-type]
|
||||||
|
kind=value["kind"], # type: ignore[arg-type]
|
||||||
|
entry_point=value["entry_point"], # type: ignore[arg-type]
|
||||||
|
needs=tuple(value["needs"]), # type: ignore[arg-type]
|
||||||
|
inputs={name: PortSpec.from_dict(port) for name, port in inputs.items()},
|
||||||
|
outputs={name: PortSpec.from_dict(port) for name, port in outputs.items()},
|
||||||
|
parameter_names=tuple(value["parameter_names"]), # type: ignore[arg-type]
|
||||||
|
resources=ResourceRequirements.from_dict(value["resources"]),
|
||||||
|
execution=ExecutionProfile.from_dict(value["execution"]),
|
||||||
|
retry=RetryPolicy.from_dict(value["retry"]),
|
||||||
|
verifier=None if value["verifier"] is None else ComponentRef.from_dict(value["verifier"]),
|
||||||
|
trust_modes=tuple(value["trust_modes"]), # type: ignore[arg-type]
|
||||||
|
max_fan_out=value["max_fan_out"], # type: ignore[arg-type]
|
||||||
|
cacheable=value["cacheable"], # type: ignore[arg-type]
|
||||||
|
loop=None if value["loop"] is None else LoopSpec.from_dict(value["loop"]),
|
||||||
|
stream=None if value["stream"] is None else StreamSpec.from_dict(value["stream"]),
|
||||||
|
gang=None if value["gang"] is None else GangSpec.from_dict(value["gang"]),
|
||||||
|
side_effect=None if value["side_effect"] is None else SideEffectSpec.from_dict(value["side_effect"]),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class WorkflowSpec:
|
||||||
|
workflow_id: str
|
||||||
|
inputs: Mapping[str, PortSpec]
|
||||||
|
stages: tuple[StageSpec, ...]
|
||||||
|
edges: tuple[ArtifactEdge, ...]
|
||||||
|
outputs: Mapping[str, PortRef]
|
||||||
|
failure_policy: WorkflowFailurePolicy = WorkflowFailurePolicy.FAIL_FAST
|
||||||
|
max_tasks: int = 10_000
|
||||||
|
max_output_bytes: int = 10 * 1024 * 1024 * 1024
|
||||||
|
schema_version: int = WORKFLOW_SCHEMA_VERSION
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
require_schema_version(self.schema_version, WORKFLOW_SCHEMA_VERSION, "workflow schema_version")
|
||||||
|
object.__setattr__(self, "workflow_id", require_identifier(self.workflow_id, "workflow_id"))
|
||||||
|
object.__setattr__(self, "inputs", _port_mapping(self.inputs, "workflow.inputs"))
|
||||||
|
stages = tuple(self.stages)
|
||||||
|
if not stages or any(not isinstance(stage, StageSpec) for stage in stages):
|
||||||
|
raise ValueError("workflow stages must contain at least one StageSpec")
|
||||||
|
stage_by_id = {stage.stage_id: stage for stage in stages}
|
||||||
|
if len(stage_by_id) != len(stages):
|
||||||
|
raise ValueError("workflow stage IDs must be unique")
|
||||||
|
object.__setattr__(self, "stages", stages)
|
||||||
|
edges = tuple(self.edges)
|
||||||
|
if any(not isinstance(edge, ArtifactEdge) for edge in edges):
|
||||||
|
raise ValueError("workflow edges must contain ArtifactEdge values")
|
||||||
|
if len({(edge.source, edge.target) for edge in edges}) != len(edges):
|
||||||
|
raise ValueError("workflow edges must be unique")
|
||||||
|
object.__setattr__(self, "edges", edges)
|
||||||
|
if not isinstance(self.outputs, Mapping) or not self.outputs:
|
||||||
|
raise ValueError("workflow outputs must be a non-empty object")
|
||||||
|
outputs: dict[str, PortRef] = {}
|
||||||
|
for name, reference in self.outputs.items():
|
||||||
|
canonical = require_identifier(name, "workflow output")
|
||||||
|
if not isinstance(reference, PortRef) or reference.stage_id is None:
|
||||||
|
raise ValueError("workflow outputs must reference stage output ports")
|
||||||
|
outputs[canonical] = reference
|
||||||
|
object.__setattr__(self, "outputs", MappingProxyType(outputs))
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"failure_policy",
|
||||||
|
enum_value(WorkflowFailurePolicy, self.failure_policy, "workflow.failure_policy"),
|
||||||
|
)
|
||||||
|
object.__setattr__(self, "max_tasks", require_positive_int(self.max_tasks, "workflow.max_tasks"))
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"max_output_bytes",
|
||||||
|
require_positive_int(self.max_output_bytes, "workflow.max_output_bytes"),
|
||||||
|
)
|
||||||
|
self._validate_graph(stage_by_id)
|
||||||
|
|
||||||
|
def _source_port(self, reference: PortRef, stages: Mapping[str, StageSpec]) -> PortSpec:
|
||||||
|
if reference.stage_id is None:
|
||||||
|
try:
|
||||||
|
return self.inputs[reference.port]
|
||||||
|
except KeyError as error:
|
||||||
|
raise ValueError(f"unknown workflow input port: {reference.port}") from error
|
||||||
|
try:
|
||||||
|
stage = stages[reference.stage_id]
|
||||||
|
return stage.outputs[reference.port]
|
||||||
|
except KeyError as error:
|
||||||
|
raise ValueError(
|
||||||
|
f"unknown source stage output: {reference.stage_id}.{reference.port}"
|
||||||
|
) from error
|
||||||
|
|
||||||
|
def _validate_graph(self, stages: Mapping[str, StageSpec]) -> None:
|
||||||
|
incoming: dict[tuple[str, str], ArtifactEdge] = {}
|
||||||
|
dependencies: dict[str, set[str]] = {stage_id: set() for stage_id in stages}
|
||||||
|
for edge in self.edges:
|
||||||
|
source_port = self._source_port(edge.source, stages)
|
||||||
|
assert edge.target.stage_id is not None
|
||||||
|
try:
|
||||||
|
target_stage = stages[edge.target.stage_id]
|
||||||
|
target_port = target_stage.inputs[edge.target.port]
|
||||||
|
except KeyError as error:
|
||||||
|
raise ValueError(
|
||||||
|
f"unknown target stage input: {edge.target.stage_id}.{edge.target.port}"
|
||||||
|
) from error
|
||||||
|
target_key = (edge.target.stage_id, edge.target.port)
|
||||||
|
if target_key in incoming:
|
||||||
|
raise ValueError("each stage input must have exactly one artifact edge")
|
||||||
|
incoming[target_key] = edge
|
||||||
|
same_schema = source_port.schema == target_port.schema
|
||||||
|
direct_match = source_port == target_port
|
||||||
|
map_fan_in = (
|
||||||
|
source_port.cardinality.value == "one"
|
||||||
|
and target_port.cardinality.value == "many"
|
||||||
|
and target_port.collection.value in {"ordered", "keyed", "set"}
|
||||||
|
)
|
||||||
|
if not same_schema or not (direct_match or map_fan_in):
|
||||||
|
raise ValueError("artifact edge source and target port declarations are incompatible")
|
||||||
|
if edge.source.stage_id is not None:
|
||||||
|
dependencies[edge.target.stage_id].add(edge.source.stage_id)
|
||||||
|
for stage in stages.values():
|
||||||
|
missing = [name for name in stage.inputs if (stage.stage_id, name) not in incoming]
|
||||||
|
if missing:
|
||||||
|
raise ValueError(
|
||||||
|
f"stage {stage.stage_id} has unbound inputs: {', '.join(sorted(missing))}"
|
||||||
|
)
|
||||||
|
if dependencies[stage.stage_id] != set(stage.needs):
|
||||||
|
raise ValueError(f"stage {stage.stage_id} needs do not match its artifact edges")
|
||||||
|
remaining = {name: set(values) for name, values in dependencies.items()}
|
||||||
|
ready = sorted(name for name, values in remaining.items() if not values)
|
||||||
|
visited: list[str] = []
|
||||||
|
while ready:
|
||||||
|
current = ready.pop(0)
|
||||||
|
visited.append(current)
|
||||||
|
for name, values in remaining.items():
|
||||||
|
if current in values:
|
||||||
|
values.remove(current)
|
||||||
|
if not values and name not in visited and name not in ready:
|
||||||
|
ready.append(name)
|
||||||
|
ready.sort()
|
||||||
|
if len(visited) != len(stages):
|
||||||
|
raise ValueError("workflow graph must be acyclic")
|
||||||
|
for reference in self.outputs.values():
|
||||||
|
self._source_port(reference, stages)
|
||||||
|
|
||||||
|
def output_ports(self) -> Mapping[str, PortSpec]:
|
||||||
|
stages = {stage.stage_id: stage for stage in self.stages}
|
||||||
|
return MappingProxyType({
|
||||||
|
name: self._source_port(reference, stages)
|
||||||
|
for name, reference in self.outputs.items()
|
||||||
|
})
|
||||||
|
|
||||||
|
def to_dict(self) -> dict[str, object]:
|
||||||
|
return {
|
||||||
|
"schema_version": self.schema_version,
|
||||||
|
"workflow_id": self.workflow_id,
|
||||||
|
"inputs": {name: port.to_dict() for name, port in self.inputs.items()},
|
||||||
|
"stages": [stage.to_dict() for stage in self.stages],
|
||||||
|
"edges": [edge.to_dict() for edge in self.edges],
|
||||||
|
"outputs": {name: reference.to_dict() for name, reference in self.outputs.items()},
|
||||||
|
"failure_policy": self.failure_policy.value,
|
||||||
|
"max_tasks": self.max_tasks,
|
||||||
|
"max_output_bytes": self.max_output_bytes,
|
||||||
|
}
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_dict(cls, value: object) -> "WorkflowSpec":
|
||||||
|
if not isinstance(value, Mapping):
|
||||||
|
raise ValueError("workflow specification must be an object")
|
||||||
|
fields = {
|
||||||
|
"schema_version", "workflow_id", "inputs", "stages", "edges",
|
||||||
|
"outputs", "failure_policy", "max_tasks", "max_output_bytes",
|
||||||
|
}
|
||||||
|
require_exact_keys(value, fields, "workflow specification")
|
||||||
|
inputs, outputs = value["inputs"], value["outputs"]
|
||||||
|
stages, edges = value["stages"], value["edges"]
|
||||||
|
if not isinstance(inputs, Mapping) or not isinstance(outputs, Mapping):
|
||||||
|
raise ValueError("workflow inputs and outputs must be objects")
|
||||||
|
if not isinstance(stages, list) or not isinstance(edges, list):
|
||||||
|
raise ValueError("workflow stages and edges must be arrays")
|
||||||
|
return cls(
|
||||||
|
schema_version=value["schema_version"], # type: ignore[arg-type]
|
||||||
|
workflow_id=value["workflow_id"], # type: ignore[arg-type]
|
||||||
|
inputs={name: PortSpec.from_dict(port) for name, port in inputs.items()},
|
||||||
|
stages=tuple(StageSpec.from_dict(stage) for stage in stages),
|
||||||
|
edges=tuple(ArtifactEdge.from_dict(edge) for edge in edges),
|
||||||
|
outputs={name: PortRef.from_dict(reference) for name, reference in outputs.items()},
|
||||||
|
failure_policy=value["failure_policy"], # type: ignore[arg-type]
|
||||||
|
max_tasks=value["max_tasks"], # type: ignore[arg-type]
|
||||||
|
max_output_bytes=value["max_output_bytes"], # type: ignore[arg-type]
|
||||||
|
)
|
||||||
@@ -0,0 +1,664 @@
|
|||||||
|
"""Compatibility tests for the built-in SDK bridge and scientific reference."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import csv
|
||||||
|
from dataclasses import replace
|
||||||
|
from pathlib import Path
|
||||||
|
from uuid import uuid4
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from scimesh.chemistry.dataset import find_molecule_by_id
|
||||||
|
from scimesh.sdk import (
|
||||||
|
ArtifactCollection,
|
||||||
|
ArtifactSchema,
|
||||||
|
CheckpointPolicy,
|
||||||
|
CompatibilityError,
|
||||||
|
ComponentRef,
|
||||||
|
DeterminismProfile,
|
||||||
|
FeatureRequirement,
|
||||||
|
GangSpec,
|
||||||
|
JobRequest,
|
||||||
|
LocalArtifactStore,
|
||||||
|
LocalCoreBatchExecutor,
|
||||||
|
LocalPlanningContext,
|
||||||
|
NetworkPolicy,
|
||||||
|
PortRef,
|
||||||
|
ProcessModel,
|
||||||
|
RetryPolicy,
|
||||||
|
SchemaRef,
|
||||||
|
StageKind,
|
||||||
|
TrustMode,
|
||||||
|
VerificationDecision,
|
||||||
|
VerificationStatus,
|
||||||
|
VersionRange,
|
||||||
|
WorkloadDefinition,
|
||||||
|
WorkloadRegistry,
|
||||||
|
assert_manifest_round_trip,
|
||||||
|
default_sdk_registry,
|
||||||
|
default_sdk_runtime,
|
||||||
|
similarity_search_sdk_adapter,
|
||||||
|
)
|
||||||
|
from scimesh.workloads.similarity_search import search_similar, write_search_results
|
||||||
|
|
||||||
|
|
||||||
|
def _write_tiny_dataset(path: Path) -> None:
|
||||||
|
path.write_text(
|
||||||
|
"chembl_id\tcanonical_smiles\textra\n"
|
||||||
|
"QUERY\tCCO\tquery\n"
|
||||||
|
"ALCOHOL\tCCCO\talcohol\n"
|
||||||
|
"ALKANE\tCCCC\talkane\n"
|
||||||
|
"BROKEN\tnot-a-smiles\tinvalid\n"
|
||||||
|
"DUPLICATE\tCCO\tduplicate\n"
|
||||||
|
"AMINE\tCCN\tamine\n",
|
||||||
|
encoding="utf-8",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _registered_similarity_search(shard_rows: int = 2):
|
||||||
|
registry = default_sdk_registry(shard_rows=shard_rows)
|
||||||
|
runtime = default_sdk_runtime()
|
||||||
|
descriptions = registry.descriptions()
|
||||||
|
assert len(descriptions) == 1
|
||||||
|
description = descriptions[0]
|
||||||
|
definition, negotiated = registry.require(
|
||||||
|
description.workload.name,
|
||||||
|
description.workload.version,
|
||||||
|
description.package_digest,
|
||||||
|
runtime=runtime,
|
||||||
|
)
|
||||||
|
return registry, runtime, description, definition, negotiated
|
||||||
|
|
||||||
|
|
||||||
|
def _request_for(
|
||||||
|
dataset: Path,
|
||||||
|
artifact_store: LocalArtifactStore,
|
||||||
|
definition: WorkloadDefinition,
|
||||||
|
) -> JobRequest:
|
||||||
|
input_port = definition.manifest.inputs["input"]
|
||||||
|
dataset_artifact = artifact_store.import_file(
|
||||||
|
dataset,
|
||||||
|
declaration=input_port.schema,
|
||||||
|
)
|
||||||
|
return JobRequest(
|
||||||
|
workload=definition.manifest.workload,
|
||||||
|
parameters={"query_id": "QUERY", "top_k": 3, "progress_every": 0},
|
||||||
|
inputs={"input": ArtifactCollection.single(dataset_artifact)},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_builtin_similarity_search_manifest_is_registered_and_negotiable() -> None:
|
||||||
|
_, _, description, definition, negotiated = _registered_similarity_search()
|
||||||
|
manifest = definition.manifest
|
||||||
|
|
||||||
|
assert description.enabled is True
|
||||||
|
assert manifest.workload.name == "similarity-search"
|
||||||
|
assert manifest.workload.version == "1.0.0"
|
||||||
|
assert manifest.determinism is DeterminismProfile.BYTE_EXACT
|
||||||
|
assert manifest.conformance_profiles == ("core-batch-v1",)
|
||||||
|
assert [stage.kind for stage in manifest.workflow.stages] == [
|
||||||
|
StageKind.MAP,
|
||||||
|
StageKind.REDUCE,
|
||||||
|
]
|
||||||
|
assert set(definition.runners) == {manifest.workflow.stages[0].entry_point}
|
||||||
|
assert set(definition.reducers) == {manifest.workflow.stages[1].entry_point}
|
||||||
|
assert negotiated is not None
|
||||||
|
assert negotiated.manifest == manifest
|
||||||
|
assert_manifest_round_trip(manifest)
|
||||||
|
|
||||||
|
|
||||||
|
def test_local_sdk_executor_matches_similarity_search_reference(tmp_path: Path) -> None:
|
||||||
|
dataset = tmp_path / "molecules.tsv"
|
||||||
|
_write_tiny_dataset(dataset)
|
||||||
|
registry, runtime, description, definition, _ = _registered_similarity_search()
|
||||||
|
artifact_store = LocalArtifactStore(tmp_path / "artifacts")
|
||||||
|
request = _request_for(dataset, artifact_store, definition)
|
||||||
|
|
||||||
|
result = LocalCoreBatchExecutor(
|
||||||
|
registry,
|
||||||
|
runtime,
|
||||||
|
artifact_store,
|
||||||
|
tmp_path / "sdk-work",
|
||||||
|
).execute(request, description.package_digest)
|
||||||
|
result_artifact = result.outputs["result"].items[0].artifact
|
||||||
|
|
||||||
|
reference_path = tmp_path / "reference.csv"
|
||||||
|
query = find_molecule_by_id(dataset, "QUERY")
|
||||||
|
reference = search_similar(dataset, query, top_k=3, progress_every=0)
|
||||||
|
write_search_results(reference_path, reference.matches)
|
||||||
|
|
||||||
|
assert artifact_store.materialize(result_artifact).read_bytes() == reference_path.read_bytes()
|
||||||
|
assert result.task_key == "reduce/final"
|
||||||
|
assert result.metrics == {"matches_emitted": 3, "partial_count": 3}
|
||||||
|
|
||||||
|
|
||||||
|
def test_legacy_adapter_planning_is_deterministic_ordered_and_path_free(
|
||||||
|
tmp_path: Path,
|
||||||
|
) -> None:
|
||||||
|
dataset = tmp_path / "molecules.tsv"
|
||||||
|
_write_tiny_dataset(dataset)
|
||||||
|
registry, runtime, description, definition, _ = _registered_similarity_search()
|
||||||
|
artifact_store = LocalArtifactStore(tmp_path / "artifacts")
|
||||||
|
request = _request_for(dataset, artifact_store, definition)
|
||||||
|
input_artifact = request.inputs["input"].items[0].artifact
|
||||||
|
|
||||||
|
first = registry.plan(
|
||||||
|
request,
|
||||||
|
description.package_digest,
|
||||||
|
runtime,
|
||||||
|
LocalPlanningContext(
|
||||||
|
artifact_store,
|
||||||
|
artifact_store,
|
||||||
|
tmp_path / "first-plan",
|
||||||
|
allowed_artifacts=(input_artifact,),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
second = registry.plan(
|
||||||
|
request,
|
||||||
|
description.package_digest,
|
||||||
|
runtime,
|
||||||
|
LocalPlanningContext(
|
||||||
|
artifact_store,
|
||||||
|
artifact_store,
|
||||||
|
tmp_path / "second-plan",
|
||||||
|
allowed_artifacts=(input_artifact,),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert first.to_json() == second.to_json()
|
||||||
|
assert first.digest == second.digest
|
||||||
|
assert first.package_digest == definition.manifest.package.digest
|
||||||
|
assert first.manifest_digest == definition.manifest.digest
|
||||||
|
assert first.trust_mode is request.trust_mode
|
||||||
|
assert JobRequest.from_json(request.to_json()) == request
|
||||||
|
assert [task.task_key for task in first.tasks] == [
|
||||||
|
"map/00000000",
|
||||||
|
"map/00000001",
|
||||||
|
"map/00000002",
|
||||||
|
]
|
||||||
|
assert all(task.stage_id == "map" for task in first.tasks)
|
||||||
|
assert all(task.package_digest == first.package_digest for task in first.tasks)
|
||||||
|
assert all(task.manifest_digest == first.manifest_digest for task in first.tasks)
|
||||||
|
assert all(task.trust_mode is first.trust_mode for task in first.tasks)
|
||||||
|
assert all("query_id" not in task.parameters for task in first.tasks)
|
||||||
|
assert all(task.parameters["query_smiles"] == "CCO" for task in first.tasks)
|
||||||
|
|
||||||
|
shard_ids: list[list[str]] = []
|
||||||
|
for task in first.tasks:
|
||||||
|
artifact = task.inputs["input"].items[0].artifact
|
||||||
|
with artifact_store.materialize(artifact).open(encoding="utf-8", newline="") as source:
|
||||||
|
shard_ids.append(
|
||||||
|
[row["chembl_id"] for row in csv.DictReader(source, delimiter="\t")]
|
||||||
|
)
|
||||||
|
assert set(artifact.to_dict()) == {
|
||||||
|
"artifact_id",
|
||||||
|
"sha256",
|
||||||
|
"schema",
|
||||||
|
"media_type",
|
||||||
|
"size_bytes",
|
||||||
|
"records",
|
||||||
|
"dimensions",
|
||||||
|
}
|
||||||
|
assert shard_ids == [
|
||||||
|
["QUERY", "ALCOHOL"],
|
||||||
|
["ALKANE", "BROKEN"],
|
||||||
|
["DUPLICATE", "AMINE"],
|
||||||
|
]
|
||||||
|
|
||||||
|
wire_payload = first.to_json()
|
||||||
|
assert str(tmp_path) not in wire_payload
|
||||||
|
assert "file://" not in wire_payload
|
||||||
|
assert "worker://" not in wire_payload
|
||||||
|
assert "workspace" not in wire_payload
|
||||||
|
|
||||||
|
|
||||||
|
def test_local_context_sink_cannot_seal_files_outside_the_attempt(
|
||||||
|
tmp_path: Path,
|
||||||
|
) -> None:
|
||||||
|
_, _, _, definition, _ = _registered_similarity_search()
|
||||||
|
artifact_store = LocalArtifactStore(tmp_path / "artifacts")
|
||||||
|
workspace = tmp_path / "attempt"
|
||||||
|
context = LocalPlanningContext(artifact_store, artifact_store, workspace)
|
||||||
|
outside = tmp_path / "private.txt"
|
||||||
|
outside.write_text("private", encoding="utf-8")
|
||||||
|
schema = definition.manifest.inputs["input"].schema
|
||||||
|
|
||||||
|
with pytest.raises(ValueError, match="inside its workspace"):
|
||||||
|
context.sink.seal(outside, declaration=schema)
|
||||||
|
|
||||||
|
workspace.mkdir(parents=True, exist_ok=True)
|
||||||
|
link = workspace / "result"
|
||||||
|
link.symlink_to(outside)
|
||||||
|
with pytest.raises(ValueError, match="real workspace directories"):
|
||||||
|
context.sink.seal(link, declaration=schema)
|
||||||
|
|
||||||
|
|
||||||
|
def test_local_store_rejects_malformed_content_before_publishing(
|
||||||
|
tmp_path: Path,
|
||||||
|
) -> None:
|
||||||
|
malformed = tmp_path / "malformed.json"
|
||||||
|
malformed.write_text('{"unfinished":', encoding="utf-8")
|
||||||
|
declaration = ArtifactSchema(
|
||||||
|
SchemaRef("json-result", 1),
|
||||||
|
"application/json",
|
||||||
|
"utf-8",
|
||||||
|
max_bytes=1_024,
|
||||||
|
validator=ComponentRef("json-document", 1),
|
||||||
|
)
|
||||||
|
store = LocalArtifactStore(tmp_path / "artifacts")
|
||||||
|
|
||||||
|
with pytest.raises(ValueError, match="not a valid bounded document"):
|
||||||
|
store.import_file(malformed, declaration=declaration)
|
||||||
|
assert tuple(path for path in store.root.iterdir() if not path.name.startswith(".seal-")) == ()
|
||||||
|
|
||||||
|
|
||||||
|
def test_delimited_validator_rejects_headerless_data_and_enforces_record_limit(
|
||||||
|
tmp_path: Path,
|
||||||
|
) -> None:
|
||||||
|
declaration = ArtifactSchema(
|
||||||
|
SchemaRef("bounded-table", 1),
|
||||||
|
"text/csv",
|
||||||
|
"utf-8",
|
||||||
|
max_bytes=1_024,
|
||||||
|
validator=ComponentRef("delimited-table", 1),
|
||||||
|
validator_configuration={"columns": ["value"]},
|
||||||
|
max_records=1,
|
||||||
|
)
|
||||||
|
store = LocalArtifactStore(tmp_path / "artifacts")
|
||||||
|
headerless = tmp_path / "headerless.csv"
|
||||||
|
headerless.write_text("1\n2\n", encoding="utf-8")
|
||||||
|
oversized = tmp_path / "oversized.csv"
|
||||||
|
oversized.write_text("value\n1\n2\n", encoding="utf-8")
|
||||||
|
|
||||||
|
with pytest.raises(ValueError, match="header does not match"):
|
||||||
|
store.import_file(headerless, declaration=declaration)
|
||||||
|
with pytest.raises(ValueError, match="record limit"):
|
||||||
|
store.import_file(oversized, declaration=declaration)
|
||||||
|
|
||||||
|
|
||||||
|
def test_custom_artifact_inspector_is_bound_to_schema_and_validator_identity(
|
||||||
|
tmp_path: Path,
|
||||||
|
) -> None:
|
||||||
|
schema_ref = SchemaRef("matrix-result", 1)
|
||||||
|
validator = ComponentRef("matrix-inspector", 1)
|
||||||
|
declaration = ArtifactSchema(
|
||||||
|
schema_ref,
|
||||||
|
"application/x-matrix",
|
||||||
|
None,
|
||||||
|
max_bytes=1_024,
|
||||||
|
validator=validator,
|
||||||
|
validator_configuration={"layout": "row-major"},
|
||||||
|
max_records=1,
|
||||||
|
max_dimensions=(2, 2),
|
||||||
|
)
|
||||||
|
source = tmp_path / "matrix.bin"
|
||||||
|
source.write_bytes(b"matrix")
|
||||||
|
wrong = LocalArtifactStore(
|
||||||
|
tmp_path / "wrong-store",
|
||||||
|
inspectors={
|
||||||
|
schema_ref.canonical: (
|
||||||
|
ComponentRef("other-inspector", 1),
|
||||||
|
lambda _path, _configuration: (1, (2, 2)),
|
||||||
|
)
|
||||||
|
},
|
||||||
|
)
|
||||||
|
with pytest.raises(ValueError, match="no matching registered validator"):
|
||||||
|
wrong.import_file(source, declaration=declaration)
|
||||||
|
|
||||||
|
def inspect(_path: Path, configuration):
|
||||||
|
assert dict(configuration) == {"layout": "row-major"}
|
||||||
|
return 1, (2, 2)
|
||||||
|
|
||||||
|
store = LocalArtifactStore(
|
||||||
|
tmp_path / "store",
|
||||||
|
inspectors={schema_ref.canonical: (validator, inspect)},
|
||||||
|
)
|
||||||
|
artifact = store.import_file(source, declaration=declaration)
|
||||||
|
assert artifact.records == 1
|
||||||
|
assert artifact.dimensions == (2, 2)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
("forgery", "message"),
|
||||||
|
(
|
||||||
|
("artifact", "artifacts sealed by its attempt"),
|
||||||
|
("provenance", "provenance does not match"),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
def test_local_executor_rejects_handler_forged_outputs(
|
||||||
|
tmp_path: Path,
|
||||||
|
forgery: str,
|
||||||
|
message: str,
|
||||||
|
) -> None:
|
||||||
|
dataset = tmp_path / "molecules.tsv"
|
||||||
|
_write_tiny_dataset(dataset)
|
||||||
|
_, runtime, description, original, _ = _registered_similarity_search()
|
||||||
|
map_stage = next(stage for stage in original.manifest.workflow.stages if stage.kind is StageKind.MAP)
|
||||||
|
inner = original.runners[map_stage.entry_point]
|
||||||
|
|
||||||
|
class ForgingRunner:
|
||||||
|
def run(self, context):
|
||||||
|
result = inner.run(context)
|
||||||
|
if forgery == "provenance":
|
||||||
|
forged = replace(
|
||||||
|
result.provenance,
|
||||||
|
worker_runtime={"kind": "forged-runtime"},
|
||||||
|
)
|
||||||
|
return replace(result, provenance=forged)
|
||||||
|
original_ref = result.outputs["partial"].items[0].artifact
|
||||||
|
forged_ref = replace(original_ref, artifact_id=str(uuid4()))
|
||||||
|
return replace(
|
||||||
|
result,
|
||||||
|
outputs={"partial": ArtifactCollection.single(forged_ref)},
|
||||||
|
)
|
||||||
|
|
||||||
|
definition = WorkloadDefinition(
|
||||||
|
original.manifest,
|
||||||
|
original.planner,
|
||||||
|
{map_stage.entry_point: ForgingRunner()},
|
||||||
|
original.reducers,
|
||||||
|
original.verifiers,
|
||||||
|
)
|
||||||
|
registry = WorkloadRegistry()
|
||||||
|
registry.register(definition, enabled=True)
|
||||||
|
store = LocalArtifactStore(tmp_path / "artifacts")
|
||||||
|
request = _request_for(dataset, store, definition)
|
||||||
|
|
||||||
|
with pytest.raises(ValueError, match=message):
|
||||||
|
LocalCoreBatchExecutor(registry, runtime, store, tmp_path / "work").execute(
|
||||||
|
request,
|
||||||
|
description.package_digest,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_local_executor_rejects_profiles_that_claim_network_isolation(tmp_path: Path) -> None:
|
||||||
|
dataset = tmp_path / "molecules.tsv"
|
||||||
|
_write_tiny_dataset(dataset)
|
||||||
|
_, runtime, description, original, _ = _registered_similarity_search()
|
||||||
|
stages = tuple(
|
||||||
|
replace(stage, execution=replace(stage.execution, network=NetworkPolicy.NONE))
|
||||||
|
for stage in original.manifest.workflow.stages
|
||||||
|
)
|
||||||
|
workflow = replace(original.manifest.workflow, stages=stages)
|
||||||
|
manifest = replace(original.manifest, workflow=workflow)
|
||||||
|
definition = WorkloadDefinition(
|
||||||
|
manifest,
|
||||||
|
original.planner,
|
||||||
|
original.runners,
|
||||||
|
original.reducers,
|
||||||
|
original.verifiers,
|
||||||
|
)
|
||||||
|
registry = WorkloadRegistry()
|
||||||
|
registry.register(definition, enabled=True)
|
||||||
|
store = LocalArtifactStore(tmp_path / "artifacts")
|
||||||
|
request = _request_for(dataset, store, definition)
|
||||||
|
|
||||||
|
with pytest.raises(CompatibilityError) as raised:
|
||||||
|
LocalCoreBatchExecutor(registry, runtime, store, tmp_path / "work").execute(
|
||||||
|
request,
|
||||||
|
description.package_digest,
|
||||||
|
)
|
||||||
|
assert raised.value.code == "feature-undeclared"
|
||||||
|
|
||||||
|
|
||||||
|
def test_local_executor_rejects_aliased_terminal_outputs_before_planning(
|
||||||
|
tmp_path: Path,
|
||||||
|
) -> None:
|
||||||
|
dataset = tmp_path / "molecules.tsv"
|
||||||
|
_write_tiny_dataset(dataset)
|
||||||
|
_, runtime, description, original, _ = _registered_similarity_search()
|
||||||
|
reducer = next(
|
||||||
|
stage for stage in original.manifest.workflow.stages if stage.kind is StageKind.REDUCE
|
||||||
|
)
|
||||||
|
internal_name = next(iter(reducer.outputs))
|
||||||
|
workflow = replace(
|
||||||
|
original.manifest.workflow,
|
||||||
|
outputs={"aliased": PortRef(internal_name, reducer.stage_id)},
|
||||||
|
)
|
||||||
|
manifest = replace(
|
||||||
|
original.manifest,
|
||||||
|
workflow=workflow,
|
||||||
|
outputs={"aliased": reducer.outputs[internal_name]},
|
||||||
|
)
|
||||||
|
definition = WorkloadDefinition(
|
||||||
|
manifest,
|
||||||
|
original.planner,
|
||||||
|
original.runners,
|
||||||
|
original.reducers,
|
||||||
|
original.verifiers,
|
||||||
|
)
|
||||||
|
registry = WorkloadRegistry()
|
||||||
|
registry.register(definition, enabled=True)
|
||||||
|
store = LocalArtifactStore(tmp_path / "artifacts")
|
||||||
|
request = _request_for(dataset, store, definition)
|
||||||
|
|
||||||
|
with pytest.raises(ValueError, match="identity-mapped reducer outputs"):
|
||||||
|
LocalCoreBatchExecutor(registry, runtime, store, tmp_path / "work").execute(
|
||||||
|
request,
|
||||||
|
description.package_digest,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_local_executor_rejects_non_trusted_trust_modes(tmp_path: Path) -> None:
|
||||||
|
dataset = tmp_path / "molecules.tsv"
|
||||||
|
_write_tiny_dataset(dataset)
|
||||||
|
_, runtime, _, original, _ = _registered_similarity_search()
|
||||||
|
stages = tuple(
|
||||||
|
replace(stage, trust_modes=("trusted", "verified"))
|
||||||
|
for stage in original.manifest.workflow.stages
|
||||||
|
)
|
||||||
|
manifest = replace(
|
||||||
|
original.manifest,
|
||||||
|
workflow=replace(original.manifest.workflow, stages=stages),
|
||||||
|
trust_modes=(TrustMode.TRUSTED, TrustMode.VERIFIED),
|
||||||
|
)
|
||||||
|
definition = WorkloadDefinition(
|
||||||
|
manifest,
|
||||||
|
original.planner,
|
||||||
|
original.runners,
|
||||||
|
original.reducers,
|
||||||
|
original.verifiers,
|
||||||
|
)
|
||||||
|
registry = WorkloadRegistry()
|
||||||
|
registry.register(definition, enabled=True)
|
||||||
|
store = LocalArtifactStore(tmp_path / "artifacts")
|
||||||
|
request = replace(
|
||||||
|
_request_for(dataset, store, definition),
|
||||||
|
trust_mode=TrustMode.VERIFIED,
|
||||||
|
)
|
||||||
|
runtime = replace(runtime, trust_modes=(TrustMode.TRUSTED, TrustMode.VERIFIED))
|
||||||
|
|
||||||
|
with pytest.raises(ValueError, match="supports only trusted workloads"):
|
||||||
|
LocalCoreBatchExecutor(registry, runtime, store, tmp_path / "work").execute(
|
||||||
|
request,
|
||||||
|
manifest.package.digest,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _advanced_execution_manifest(
|
||||||
|
original: WorkloadDefinition,
|
||||||
|
case: str,
|
||||||
|
) -> tuple[tuple[str, ...], object]:
|
||||||
|
"""Declare one negotiable advanced profile the local executor cannot enforce."""
|
||||||
|
stages = original.manifest.workflow.stages
|
||||||
|
if case == "process-pool":
|
||||||
|
features = ("process-pools", "multi-process")
|
||||||
|
changed = tuple(
|
||||||
|
replace(
|
||||||
|
stage,
|
||||||
|
resources=replace(stage.resources, cpu_cores=2),
|
||||||
|
execution=replace(
|
||||||
|
stage.execution,
|
||||||
|
process_model=ProcessModel.PROCESS_POOL,
|
||||||
|
max_processes=2,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
for stage in stages
|
||||||
|
)
|
||||||
|
elif case == "checkpoints":
|
||||||
|
features = ("checkpoints",)
|
||||||
|
changed = tuple(
|
||||||
|
replace(
|
||||||
|
stage,
|
||||||
|
execution=replace(
|
||||||
|
stage.execution,
|
||||||
|
checkpoint=CheckpointPolicy(
|
||||||
|
enabled=True,
|
||||||
|
schema=SchemaRef("task-state", 1),
|
||||||
|
compatibility_version=1,
|
||||||
|
),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
for stage in stages
|
||||||
|
)
|
||||||
|
elif case == "retries":
|
||||||
|
features = ("retries",)
|
||||||
|
changed = tuple(
|
||||||
|
replace(stage, retry=RetryPolicy(max_attempts=2))
|
||||||
|
for stage in stages
|
||||||
|
)
|
||||||
|
elif case == "secrets":
|
||||||
|
features = ("secret-injection",)
|
||||||
|
changed = tuple(
|
||||||
|
replace(
|
||||||
|
stage,
|
||||||
|
execution=replace(stage.execution, secret_handles=("db-credential",)),
|
||||||
|
)
|
||||||
|
for stage in stages
|
||||||
|
)
|
||||||
|
elif case == "gang":
|
||||||
|
features = ("gang-leases",)
|
||||||
|
changed = tuple(
|
||||||
|
replace(
|
||||||
|
stage,
|
||||||
|
gang=GangSpec(replicas=2, per_replica_resources=stage.resources),
|
||||||
|
)
|
||||||
|
for stage in stages
|
||||||
|
)
|
||||||
|
elif case == "network-isolation":
|
||||||
|
features = ("network-isolation",)
|
||||||
|
changed = tuple(
|
||||||
|
replace(
|
||||||
|
stage,
|
||||||
|
execution=replace(stage.execution, network=NetworkPolicy.NONE),
|
||||||
|
)
|
||||||
|
for stage in stages
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
assert case == "service-stage"
|
||||||
|
features = ("services",)
|
||||||
|
changed = (replace(stages[0], kind=StageKind.SERVICE),) + stages[1:]
|
||||||
|
manifest = replace(
|
||||||
|
original.manifest,
|
||||||
|
workflow=replace(original.manifest.workflow, stages=changed),
|
||||||
|
required_features=original.manifest.required_features
|
||||||
|
+ tuple(
|
||||||
|
FeatureRequirement(feature, VersionRange(">=1,<2")) for feature in features
|
||||||
|
),
|
||||||
|
)
|
||||||
|
return features, manifest
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
("case", "message"),
|
||||||
|
(
|
||||||
|
("process-pool", "one non-nested host thread"),
|
||||||
|
("checkpoints", "cannot enforce this stage profile"),
|
||||||
|
("retries", "does not implement retries"),
|
||||||
|
("secrets", "cannot enforce this stage profile"),
|
||||||
|
("gang", "cannot enforce this stage profile"),
|
||||||
|
("network-isolation", "cannot enforce a restricted network policy"),
|
||||||
|
("service-stage", "does not implement advanced stages"),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
def test_local_executor_rejects_profiles_it_cannot_enforce(
|
||||||
|
tmp_path: Path,
|
||||||
|
case: str,
|
||||||
|
message: str,
|
||||||
|
) -> None:
|
||||||
|
dataset = tmp_path / "molecules.tsv"
|
||||||
|
_write_tiny_dataset(dataset)
|
||||||
|
_, runtime, _, original, _ = _registered_similarity_search()
|
||||||
|
features, manifest = _advanced_execution_manifest(original, case)
|
||||||
|
definition = WorkloadDefinition(
|
||||||
|
manifest,
|
||||||
|
original.planner,
|
||||||
|
original.runners,
|
||||||
|
original.reducers,
|
||||||
|
original.verifiers,
|
||||||
|
)
|
||||||
|
registry = WorkloadRegistry()
|
||||||
|
registry.register(definition, enabled=True)
|
||||||
|
store = LocalArtifactStore(tmp_path / "artifacts")
|
||||||
|
request = _request_for(dataset, store, definition)
|
||||||
|
runtime = replace(
|
||||||
|
runtime,
|
||||||
|
features={**runtime.features, **{feature: "1.0.0" for feature in features}},
|
||||||
|
)
|
||||||
|
|
||||||
|
with pytest.raises(ValueError, match=message):
|
||||||
|
LocalCoreBatchExecutor(registry, runtime, store, tmp_path / "work").execute(
|
||||||
|
request,
|
||||||
|
manifest.package.digest,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_local_executor_fails_when_the_declared_verifier_rejects(tmp_path: Path) -> None:
|
||||||
|
dataset = tmp_path / "molecules.tsv"
|
||||||
|
_write_tiny_dataset(dataset)
|
||||||
|
_, runtime, _, original, _ = _registered_similarity_search()
|
||||||
|
|
||||||
|
class RejectingVerifier:
|
||||||
|
identity = ComponentRef("exact-artifact", 1)
|
||||||
|
|
||||||
|
def verify(self, context, candidates):
|
||||||
|
return VerificationDecision(
|
||||||
|
VerificationStatus.REJECTED,
|
||||||
|
self.identity,
|
||||||
|
"forced-rejection",
|
||||||
|
{},
|
||||||
|
)
|
||||||
|
|
||||||
|
definition = WorkloadDefinition(
|
||||||
|
original.manifest,
|
||||||
|
original.planner,
|
||||||
|
original.runners,
|
||||||
|
original.reducers,
|
||||||
|
{ComponentRef("exact-artifact", 1).canonical: RejectingVerifier()},
|
||||||
|
)
|
||||||
|
registry = WorkloadRegistry()
|
||||||
|
registry.register(definition, enabled=True)
|
||||||
|
store = LocalArtifactStore(tmp_path / "artifacts")
|
||||||
|
request = _request_for(dataset, store, definition)
|
||||||
|
|
||||||
|
with pytest.raises(ValueError, match="did not pass its declared verifier"):
|
||||||
|
LocalCoreBatchExecutor(registry, runtime, store, tmp_path / "work").execute(
|
||||||
|
request,
|
||||||
|
original.manifest.package.digest,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_local_executor_enforces_the_declared_output_byte_budget(tmp_path: Path) -> None:
|
||||||
|
dataset = tmp_path / "molecules.tsv"
|
||||||
|
_write_tiny_dataset(dataset)
|
||||||
|
runtime = default_sdk_runtime()
|
||||||
|
adapter = similarity_search_sdk_adapter(shard_rows=2)
|
||||||
|
# The planner pins its own manifest into every task, so the budget cut must
|
||||||
|
# be applied to the adapter's manifest for plan and definition to agree.
|
||||||
|
adapter.manifest = replace(
|
||||||
|
adapter.manifest,
|
||||||
|
workflow=replace(adapter.manifest.workflow, max_output_bytes=256),
|
||||||
|
limits=replace(adapter.manifest.limits, max_output_bytes=256),
|
||||||
|
)
|
||||||
|
definition = adapter.definition()
|
||||||
|
registry = WorkloadRegistry()
|
||||||
|
registry.register(definition, enabled=True)
|
||||||
|
store = LocalArtifactStore(tmp_path / "artifacts")
|
||||||
|
request = _request_for(dataset, store, definition)
|
||||||
|
|
||||||
|
with pytest.raises(ValueError, match="bytes exceed their sink limit"):
|
||||||
|
LocalCoreBatchExecutor(registry, runtime, store, tmp_path / "work").execute(
|
||||||
|
request,
|
||||||
|
definition.manifest.package.digest,
|
||||||
|
)
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,609 @@
|
|||||||
|
"""Security and version-pinning tests for the installed SDK registry."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from dataclasses import replace
|
||||||
|
from importlib import metadata
|
||||||
|
from pathlib import Path
|
||||||
|
import py_compile
|
||||||
|
from types import SimpleNamespace
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from scimesh.sdk import (
|
||||||
|
AllowedPackage,
|
||||||
|
ArtifactCollection,
|
||||||
|
ArtifactRef,
|
||||||
|
CompatibilityError,
|
||||||
|
FeatureRequirement,
|
||||||
|
JobRequest,
|
||||||
|
LocalArtifactStore,
|
||||||
|
LocalPlanningContext,
|
||||||
|
PackageSpec,
|
||||||
|
TrustMode,
|
||||||
|
VersionRange,
|
||||||
|
WorkloadDefinition,
|
||||||
|
WorkloadId,
|
||||||
|
WorkloadRegistry,
|
||||||
|
current_scimesh_package_digest,
|
||||||
|
default_sdk_runtime,
|
||||||
|
installed_distribution_digest,
|
||||||
|
similarity_search_sdk_adapter,
|
||||||
|
)
|
||||||
|
from scimesh.sdk.schema import (
|
||||||
|
ParameterValidationError,
|
||||||
|
validate_parameter_instance,
|
||||||
|
validate_schema_definition,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _definition(*, version: str = "1.0.0", digest_character: str = "a") -> WorkloadDefinition:
|
||||||
|
original = similarity_search_sdk_adapter(shard_rows=2).definition()
|
||||||
|
manifest = replace(
|
||||||
|
original.manifest,
|
||||||
|
workload=WorkloadId("similarity-search", version),
|
||||||
|
package=PackageSpec("scimesh", "sha256:" + digest_character * 64),
|
||||||
|
)
|
||||||
|
return WorkloadDefinition(
|
||||||
|
manifest,
|
||||||
|
original.planner,
|
||||||
|
original.runners,
|
||||||
|
original.reducers,
|
||||||
|
original.verifiers,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_registry_requires_an_explicit_enabled_version_and_digest() -> None:
|
||||||
|
first = _definition(version="1.0.0", digest_character="a")
|
||||||
|
second = _definition(version="2.0.0", digest_character="b")
|
||||||
|
registry = WorkloadRegistry()
|
||||||
|
registry.register(first, enabled=True)
|
||||||
|
registry.register(second)
|
||||||
|
|
||||||
|
resolved, _ = registry.require("similarity-search", "1.0.0", "sha256:" + "a" * 64)
|
||||||
|
assert resolved is first
|
||||||
|
with pytest.raises(ValueError, match="unknown workload version"):
|
||||||
|
registry.require("similarity-search", "3.0.0", "sha256:" + "a" * 64)
|
||||||
|
with pytest.raises(ValueError, match="not enabled"):
|
||||||
|
registry.require("similarity-search", "2.0.0", "sha256:" + "b" * 64)
|
||||||
|
with pytest.raises(ValueError, match="not enabled"):
|
||||||
|
registry.require("similarity-search", "1.0.0", "sha256:" + "c" * 64)
|
||||||
|
with pytest.raises(ValueError, match="already registered"):
|
||||||
|
registry.register(first)
|
||||||
|
|
||||||
|
registry.enable("similarity-search", "2.0.0", "sha256:" + "b" * 64)
|
||||||
|
assert [item.workload.version for item in registry.descriptions()] == ["1.0.0", "2.0.0"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_compatibility_failure_occurs_before_planner_invocation(
|
||||||
|
tmp_path: Path,
|
||||||
|
) -> None:
|
||||||
|
original = similarity_search_sdk_adapter(shard_rows=2).definition()
|
||||||
|
|
||||||
|
class CountingPlanner:
|
||||||
|
calls = 0
|
||||||
|
|
||||||
|
def validate(self, request):
|
||||||
|
self.calls += 1
|
||||||
|
return original.planner.validate(request)
|
||||||
|
|
||||||
|
def plan(self, job, context):
|
||||||
|
self.calls += 1
|
||||||
|
return original.planner.plan(job, context)
|
||||||
|
|
||||||
|
planner = CountingPlanner()
|
||||||
|
definition = WorkloadDefinition(
|
||||||
|
original.manifest,
|
||||||
|
planner,
|
||||||
|
original.runners,
|
||||||
|
original.reducers,
|
||||||
|
original.verifiers,
|
||||||
|
)
|
||||||
|
registry = WorkloadRegistry()
|
||||||
|
registry.register(definition, enabled=True)
|
||||||
|
input_port = definition.manifest.inputs["input"]
|
||||||
|
artifact = ArtifactRef(
|
||||||
|
"11111111-1111-4111-8111-111111111111",
|
||||||
|
"a" * 64,
|
||||||
|
input_port.schema.ref,
|
||||||
|
input_port.schema.media_type,
|
||||||
|
1,
|
||||||
|
)
|
||||||
|
request = JobRequest(
|
||||||
|
definition.manifest.workload,
|
||||||
|
{"query_smiles": "CCO"},
|
||||||
|
{"input": ArtifactCollection.single(artifact)},
|
||||||
|
)
|
||||||
|
incompatible = replace(default_sdk_runtime(), protocol_version="2.0.0")
|
||||||
|
store = LocalArtifactStore(tmp_path / "artifacts")
|
||||||
|
|
||||||
|
with pytest.raises(CompatibilityError) as raised:
|
||||||
|
registry.plan(
|
||||||
|
request,
|
||||||
|
definition.manifest.package.digest,
|
||||||
|
incompatible,
|
||||||
|
LocalPlanningContext(store, store, tmp_path / "plan"),
|
||||||
|
)
|
||||||
|
assert raised.value.code == "protocol-mismatch"
|
||||||
|
assert planner.calls == 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
("request_changes", "error_code"),
|
||||||
|
(
|
||||||
|
({"required_features": ("undeclared-feature",)}, "feature-undeclared"),
|
||||||
|
({"trust_mode": TrustMode.VERIFIED}, "trust-mode-undeclared"),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
def test_job_selected_features_and_trust_mode_fail_closed_before_planning(
|
||||||
|
tmp_path: Path,
|
||||||
|
request_changes: dict[str, object],
|
||||||
|
error_code: str,
|
||||||
|
) -> None:
|
||||||
|
definition = similarity_search_sdk_adapter(shard_rows=2).definition()
|
||||||
|
registry = WorkloadRegistry()
|
||||||
|
registry.register(definition, enabled=True)
|
||||||
|
input_port = definition.manifest.inputs["input"]
|
||||||
|
artifact = ArtifactRef(
|
||||||
|
"11111111-1111-4111-8111-111111111111",
|
||||||
|
"a" * 64,
|
||||||
|
input_port.schema.ref,
|
||||||
|
input_port.schema.media_type,
|
||||||
|
1,
|
||||||
|
records=1,
|
||||||
|
)
|
||||||
|
values: dict[str, object] = {
|
||||||
|
"workload": definition.manifest.workload,
|
||||||
|
"parameters": {"query_smiles": "CCO"},
|
||||||
|
"inputs": {"input": ArtifactCollection.single(artifact)},
|
||||||
|
}
|
||||||
|
values.update(request_changes)
|
||||||
|
request = JobRequest(**values) # type: ignore[arg-type]
|
||||||
|
store = LocalArtifactStore(tmp_path / "artifacts")
|
||||||
|
|
||||||
|
with pytest.raises(CompatibilityError) as raised:
|
||||||
|
registry.plan(
|
||||||
|
request,
|
||||||
|
definition.manifest.package.digest,
|
||||||
|
default_sdk_runtime(),
|
||||||
|
LocalPlanningContext(store, store, tmp_path / "plan"),
|
||||||
|
)
|
||||||
|
assert raised.value.code == error_code
|
||||||
|
|
||||||
|
|
||||||
|
class _EntryPoints(tuple):
|
||||||
|
def select(self, *, group: str):
|
||||||
|
assert group == WorkloadRegistry.ENTRY_POINT_GROUP
|
||||||
|
return self
|
||||||
|
|
||||||
|
|
||||||
|
def test_discovery_imports_only_an_exact_allowlisted_installed_entry_point(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
definition = similarity_search_sdk_adapter().definition()
|
||||||
|
loaded: list[str] = []
|
||||||
|
|
||||||
|
class EntryPoint:
|
||||||
|
def __init__(self, name: str, distribution: str) -> None:
|
||||||
|
self.name = name
|
||||||
|
self.dist = (
|
||||||
|
metadata.distribution("scimesh")
|
||||||
|
if distribution == "scimesh"
|
||||||
|
else SimpleNamespace(name=distribution)
|
||||||
|
)
|
||||||
|
self.value = "scimesh.sdk.builtins:similarity_search_sdk_adapter"
|
||||||
|
|
||||||
|
@property
|
||||||
|
def module(self) -> str:
|
||||||
|
return self.value.partition(":")[0]
|
||||||
|
|
||||||
|
def load(self):
|
||||||
|
loaded.append(self.name)
|
||||||
|
return lambda: definition
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"scimesh.sdk.registry.metadata.entry_points",
|
||||||
|
lambda: _EntryPoints(
|
||||||
|
(
|
||||||
|
EntryPoint("evil-workload@1.0.0", "unapproved"),
|
||||||
|
EntryPoint("similarity-search@1.0.0", "scimesh"),
|
||||||
|
)
|
||||||
|
),
|
||||||
|
)
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"scimesh.sdk.registry.installed_distribution_digest",
|
||||||
|
lambda _distribution: definition.manifest.package.digest,
|
||||||
|
)
|
||||||
|
registry = WorkloadRegistry()
|
||||||
|
registry.discover_installed(
|
||||||
|
(
|
||||||
|
AllowedPackage(
|
||||||
|
"scimesh",
|
||||||
|
definition.manifest.workload,
|
||||||
|
definition.manifest.package.digest,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
assert loaded == ["similarity-search@1.0.0"]
|
||||||
|
assert registry.descriptions()[0].enabled
|
||||||
|
|
||||||
|
|
||||||
|
def test_discovery_measures_package_before_importing_entry_point(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
definition = similarity_search_sdk_adapter().definition()
|
||||||
|
loaded = False
|
||||||
|
|
||||||
|
class EntryPoint:
|
||||||
|
name = "similarity-search@1.0.0"
|
||||||
|
dist = metadata.distribution("scimesh")
|
||||||
|
value = "scimesh.sdk.builtins:similarity_search_sdk_adapter"
|
||||||
|
module = "scimesh.sdk.builtins"
|
||||||
|
|
||||||
|
def load(self):
|
||||||
|
nonlocal loaded
|
||||||
|
loaded = True
|
||||||
|
return lambda: definition
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"scimesh.sdk.registry.metadata.entry_points",
|
||||||
|
lambda: _EntryPoints((EntryPoint(),)),
|
||||||
|
)
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"scimesh.sdk.registry.installed_distribution_digest",
|
||||||
|
lambda _distribution: "sha256:" + "f" * 64,
|
||||||
|
)
|
||||||
|
|
||||||
|
with pytest.raises(ValueError, match="content does not match"):
|
||||||
|
WorkloadRegistry().discover_installed(
|
||||||
|
(
|
||||||
|
AllowedPackage(
|
||||||
|
"scimesh",
|
||||||
|
definition.manifest.workload,
|
||||||
|
definition.manifest.package.digest,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
assert loaded is False
|
||||||
|
|
||||||
|
|
||||||
|
def test_installed_digest_is_stable_when_python_generates_a_pycache(
|
||||||
|
tmp_path: Path,
|
||||||
|
) -> None:
|
||||||
|
package = tmp_path / "fixture_pkg"
|
||||||
|
package.mkdir()
|
||||||
|
source = package / "__init__.py"
|
||||||
|
source.write_text("VALUE = 1\n", encoding="utf-8")
|
||||||
|
|
||||||
|
class FixtureDistribution:
|
||||||
|
name = "fixture-dist"
|
||||||
|
files = (Path("fixture_pkg/__init__.py"),)
|
||||||
|
entry_points = ()
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def read_text(name: str) -> str | None:
|
||||||
|
return "fixture_pkg\n" if name == "top_level.txt" else None
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def locate_file(value: object) -> Path:
|
||||||
|
return tmp_path / str(value)
|
||||||
|
|
||||||
|
distribution = FixtureDistribution()
|
||||||
|
before = installed_distribution_digest(distribution) # type: ignore[arg-type]
|
||||||
|
py_compile.compile(str(source), doraise=True)
|
||||||
|
|
||||||
|
assert installed_distribution_digest(distribution) == before # type: ignore[arg-type]
|
||||||
|
|
||||||
|
|
||||||
|
def test_discovery_rejects_entry_point_module_owned_by_another_distribution(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
tmp_path: Path,
|
||||||
|
) -> None:
|
||||||
|
owned = tmp_path / "owned_pkg"
|
||||||
|
owned.mkdir()
|
||||||
|
(owned / "__init__.py").write_text("", encoding="utf-8")
|
||||||
|
loaded = False
|
||||||
|
|
||||||
|
class Distribution:
|
||||||
|
name = "allowed-dist"
|
||||||
|
files = (Path("owned_pkg/__init__.py"),)
|
||||||
|
entry_points = ()
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def read_text(name: str) -> str | None:
|
||||||
|
return "owned_pkg\n" if name == "top_level.txt" else None
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def locate_file(value: object) -> Path:
|
||||||
|
return tmp_path / str(value)
|
||||||
|
|
||||||
|
class EntryPoint:
|
||||||
|
name = "similarity-search@1.0.0"
|
||||||
|
dist = Distribution()
|
||||||
|
value = "foreign_pkg.workload:factory"
|
||||||
|
module = "foreign_pkg.workload"
|
||||||
|
|
||||||
|
def load(self):
|
||||||
|
nonlocal loaded
|
||||||
|
loaded = True
|
||||||
|
raise AssertionError("foreign entry point must not load")
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"scimesh.sdk.registry.metadata.entry_points",
|
||||||
|
lambda: _EntryPoints((EntryPoint(),)),
|
||||||
|
)
|
||||||
|
definition = similarity_search_sdk_adapter().definition()
|
||||||
|
with pytest.raises(ValueError, match="outside its distribution"):
|
||||||
|
WorkloadRegistry().discover_installed(
|
||||||
|
(
|
||||||
|
AllowedPackage(
|
||||||
|
"allowed-dist",
|
||||||
|
definition.manifest.workload,
|
||||||
|
"sha256:" + "a" * 64,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
assert loaded is False
|
||||||
|
|
||||||
|
|
||||||
|
def test_missing_allowlisted_entry_point_fails_without_loading_or_registering(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
loaded: list[str] = []
|
||||||
|
|
||||||
|
class EntryPoint:
|
||||||
|
name = "job-selected-module@1.0.0"
|
||||||
|
dist = SimpleNamespace(name="unapproved")
|
||||||
|
|
||||||
|
def load(self):
|
||||||
|
loaded.append(self.name)
|
||||||
|
raise AssertionError("unapproved entry point must not load")
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"scimesh.sdk.registry.metadata.entry_points",
|
||||||
|
lambda: _EntryPoints((EntryPoint(),)),
|
||||||
|
)
|
||||||
|
registry = WorkloadRegistry()
|
||||||
|
with pytest.raises(ValueError, match="were not installed"):
|
||||||
|
registry.discover_installed(
|
||||||
|
(
|
||||||
|
AllowedPackage(
|
||||||
|
"scimesh",
|
||||||
|
WorkloadId("similarity-search", "1.0.0"),
|
||||||
|
current_scimesh_package_digest(),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
assert loaded == []
|
||||||
|
assert registry.descriptions() == ()
|
||||||
|
|
||||||
|
|
||||||
|
def test_parameter_schema_accepts_finite_big_integer_bounds() -> None:
|
||||||
|
bound = 10**400
|
||||||
|
schema = {"type": "integer", "minimum": -bound, "maximum": bound}
|
||||||
|
|
||||||
|
validate_schema_definition(schema)
|
||||||
|
validate_parameter_instance(bound, schema)
|
||||||
|
|
||||||
|
with pytest.raises(ParameterValidationError, match="violates maximum"):
|
||||||
|
validate_parameter_instance(bound + 1, schema)
|
||||||
|
|
||||||
|
|
||||||
|
def test_job_parameters_reject_unbounded_json_integers_early() -> None:
|
||||||
|
with pytest.raises(ValueError, match="4096-bit JSON bound"):
|
||||||
|
JobRequest(
|
||||||
|
WorkloadId("similarity-search", "1.0.0"),
|
||||||
|
{"value": 10**2_000},
|
||||||
|
{},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
("value", "multiple", "accepted"),
|
||||||
|
[
|
||||||
|
(3 * 10**400, 3, True),
|
||||||
|
(10**400, 3, False),
|
||||||
|
(10**400, 0.1, True),
|
||||||
|
(0.3, 0.1, True),
|
||||||
|
(0.31, 0.1, False),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_parameter_schema_multiple_of_is_exact_without_float_overflow(
|
||||||
|
value: int | float,
|
||||||
|
multiple: int | float,
|
||||||
|
accepted: bool,
|
||||||
|
) -> None:
|
||||||
|
schema = {"type": "number", "multipleOf": multiple}
|
||||||
|
validate_schema_definition(schema)
|
||||||
|
|
||||||
|
if accepted:
|
||||||
|
validate_parameter_instance(value, schema)
|
||||||
|
else:
|
||||||
|
with pytest.raises(ParameterValidationError, match="violates multipleOf"):
|
||||||
|
validate_parameter_instance(value, schema)
|
||||||
|
|
||||||
|
|
||||||
|
def test_parameter_schema_equality_uses_json_types() -> None:
|
||||||
|
validate_schema_definition({"enum": [True, 1]})
|
||||||
|
with pytest.raises(ValueError, match="enum values must be unique"):
|
||||||
|
validate_schema_definition({"enum": [1, 1.0]})
|
||||||
|
|
||||||
|
validate_parameter_instance(True, {"enum": [True]})
|
||||||
|
with pytest.raises(ParameterValidationError, match="outside enum"):
|
||||||
|
validate_parameter_instance(1, {"enum": [True]})
|
||||||
|
validate_parameter_instance(1.0, {"enum": [1]})
|
||||||
|
|
||||||
|
validate_parameter_instance({"enabled": True}, {"const": {"enabled": True}})
|
||||||
|
with pytest.raises(ParameterValidationError, match="does not match const"):
|
||||||
|
validate_parameter_instance({"enabled": 1}, {"const": {"enabled": True}})
|
||||||
|
|
||||||
|
unique = {"type": "array", "uniqueItems": True}
|
||||||
|
validate_parameter_instance([True, 1, {"enabled": True}, {"enabled": 1}], unique)
|
||||||
|
with pytest.raises(ParameterValidationError, match="items must be unique"):
|
||||||
|
validate_parameter_instance([1, 1.0], unique)
|
||||||
|
|
||||||
|
|
||||||
|
def test_disabled_workload_is_not_resolvable_until_re_enabled() -> None:
|
||||||
|
definition = _definition()
|
||||||
|
registry = WorkloadRegistry()
|
||||||
|
registry.register(definition, enabled=True)
|
||||||
|
resolved, _ = registry.require("similarity-search", "1.0.0", "sha256:" + "a" * 64)
|
||||||
|
assert resolved is definition
|
||||||
|
|
||||||
|
registry.disable("similarity-search", "1.0.0", "sha256:" + "a" * 64)
|
||||||
|
with pytest.raises(ValueError, match="not enabled"):
|
||||||
|
registry.require("similarity-search", "1.0.0", "sha256:" + "a" * 64)
|
||||||
|
|
||||||
|
registry.enable("similarity-search", "1.0.0", "sha256:" + "a" * 64)
|
||||||
|
resolved, _ = registry.require("similarity-search", "1.0.0", "sha256:" + "a" * 64)
|
||||||
|
assert resolved is definition
|
||||||
|
|
||||||
|
|
||||||
|
def test_discovery_rechecks_the_package_digest_after_loading(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
definition = similarity_search_sdk_adapter().definition()
|
||||||
|
digests = iter((definition.manifest.package.digest, "sha256:" + "e" * 64))
|
||||||
|
|
||||||
|
class EntryPoint:
|
||||||
|
name = "similarity-search@1.0.0"
|
||||||
|
dist = metadata.distribution("scimesh")
|
||||||
|
value = "scimesh.sdk.builtins:similarity_search_sdk_adapter"
|
||||||
|
module = "scimesh.sdk.builtins"
|
||||||
|
|
||||||
|
def load(self):
|
||||||
|
return lambda: definition
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"scimesh.sdk.registry.metadata.entry_points",
|
||||||
|
lambda: _EntryPoints((EntryPoint(),)),
|
||||||
|
)
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"scimesh.sdk.registry.installed_distribution_digest",
|
||||||
|
lambda _distribution: next(digests),
|
||||||
|
)
|
||||||
|
|
||||||
|
registry = WorkloadRegistry()
|
||||||
|
with pytest.raises(ValueError, match="changed while loading"):
|
||||||
|
registry.discover_installed(
|
||||||
|
(
|
||||||
|
AllowedPackage(
|
||||||
|
"scimesh",
|
||||||
|
definition.manifest.workload,
|
||||||
|
definition.manifest.package.digest,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
assert registry.descriptions() == ()
|
||||||
|
|
||||||
|
|
||||||
|
def test_request_trust_mode_must_be_enforceable_by_runtime_and_stages(tmp_path: Path) -> None:
|
||||||
|
original = similarity_search_sdk_adapter(shard_rows=2).definition()
|
||||||
|
manifest = replace(
|
||||||
|
original.manifest,
|
||||||
|
trust_modes=(TrustMode.TRUSTED, TrustMode.VERIFIED),
|
||||||
|
)
|
||||||
|
definition = WorkloadDefinition(
|
||||||
|
manifest,
|
||||||
|
original.planner,
|
||||||
|
original.runners,
|
||||||
|
original.reducers,
|
||||||
|
original.verifiers,
|
||||||
|
)
|
||||||
|
registry = WorkloadRegistry()
|
||||||
|
registry.register(definition, enabled=True)
|
||||||
|
input_port = definition.manifest.inputs["input"]
|
||||||
|
artifact = ArtifactRef(
|
||||||
|
"11111111-1111-4111-8111-111111111111",
|
||||||
|
"a" * 64,
|
||||||
|
input_port.schema.ref,
|
||||||
|
input_port.schema.media_type,
|
||||||
|
1,
|
||||||
|
records=1,
|
||||||
|
)
|
||||||
|
request = JobRequest(
|
||||||
|
definition.manifest.workload,
|
||||||
|
{"query_smiles": "CCO"},
|
||||||
|
{"input": ArtifactCollection.single(artifact)},
|
||||||
|
trust_mode=TrustMode.VERIFIED,
|
||||||
|
)
|
||||||
|
store = LocalArtifactStore(tmp_path / "artifacts")
|
||||||
|
|
||||||
|
with pytest.raises(CompatibilityError) as raised:
|
||||||
|
registry.plan(
|
||||||
|
request,
|
||||||
|
manifest.package.digest,
|
||||||
|
default_sdk_runtime(),
|
||||||
|
LocalPlanningContext(store, store, tmp_path / "runtime-plan"),
|
||||||
|
)
|
||||||
|
assert raised.value.code == "trust-mode-unavailable"
|
||||||
|
|
||||||
|
runtime = replace(
|
||||||
|
default_sdk_runtime(),
|
||||||
|
trust_modes=(TrustMode.TRUSTED, TrustMode.VERIFIED),
|
||||||
|
)
|
||||||
|
with pytest.raises(CompatibilityError) as raised:
|
||||||
|
registry.plan(
|
||||||
|
request,
|
||||||
|
manifest.package.digest,
|
||||||
|
runtime,
|
||||||
|
LocalPlanningContext(store, store, tmp_path / "stage-plan"),
|
||||||
|
)
|
||||||
|
assert raised.value.code == "stage-trust-unavailable"
|
||||||
|
|
||||||
|
|
||||||
|
def test_job_cannot_require_a_feature_outside_the_runtime(tmp_path: Path) -> None:
|
||||||
|
original = similarity_search_sdk_adapter(shard_rows=2).definition()
|
||||||
|
manifest = replace(
|
||||||
|
original.manifest,
|
||||||
|
optional_features=(
|
||||||
|
FeatureRequirement("gpu-fastpath", VersionRange(">=1,<2"), "cpu-fallback"),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
definition = WorkloadDefinition(
|
||||||
|
manifest,
|
||||||
|
original.planner,
|
||||||
|
original.runners,
|
||||||
|
original.reducers,
|
||||||
|
original.verifiers,
|
||||||
|
)
|
||||||
|
registry = WorkloadRegistry()
|
||||||
|
registry.register(definition, enabled=True)
|
||||||
|
input_port = definition.manifest.inputs["input"]
|
||||||
|
artifact = ArtifactRef(
|
||||||
|
"11111111-1111-4111-8111-111111111111",
|
||||||
|
"a" * 64,
|
||||||
|
input_port.schema.ref,
|
||||||
|
input_port.schema.media_type,
|
||||||
|
1,
|
||||||
|
records=1,
|
||||||
|
)
|
||||||
|
request = JobRequest(
|
||||||
|
definition.manifest.workload,
|
||||||
|
{"query_smiles": "CCO"},
|
||||||
|
{"input": ArtifactCollection.single(artifact)},
|
||||||
|
required_features=("gpu-fastpath",),
|
||||||
|
)
|
||||||
|
|
||||||
|
negotiated = registry.require(
|
||||||
|
"similarity-search",
|
||||||
|
"1.0.0",
|
||||||
|
manifest.package.digest,
|
||||||
|
runtime=default_sdk_runtime(),
|
||||||
|
)[1]
|
||||||
|
assert negotiated is not None
|
||||||
|
assert negotiated.optional_fallbacks == {"gpu-fastpath": "cpu-fallback"}
|
||||||
|
|
||||||
|
with pytest.raises(CompatibilityError) as raised:
|
||||||
|
registry.plan(
|
||||||
|
request,
|
||||||
|
manifest.package.digest,
|
||||||
|
default_sdk_runtime(),
|
||||||
|
LocalPlanningContext(
|
||||||
|
LocalArtifactStore(tmp_path / "artifacts"),
|
||||||
|
LocalArtifactStore(tmp_path / "artifacts"),
|
||||||
|
tmp_path / "plan",
|
||||||
|
),
|
||||||
|
)
|
||||||
|
assert raised.value.code == "feature-unavailable"
|
||||||
@@ -0,0 +1,227 @@
|
|||||||
|
"""Resource inventory and atomic local allocation tests for the SDK Agent layer."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from concurrent.futures import ThreadPoolExecutor
|
||||||
|
from threading import Barrier
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from scimesh.sdk import (
|
||||||
|
AcceleratorDevice,
|
||||||
|
AcceleratorMode,
|
||||||
|
ResourceInventory,
|
||||||
|
ResourcePool,
|
||||||
|
ResourceRequirements,
|
||||||
|
ResourceUnavailableError,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
ENVIRONMENT_DIGEST = "sha256:" + "d" * 64
|
||||||
|
|
||||||
|
|
||||||
|
def gpu(device_id: str, *, topology_group: str = "socket-0") -> AcceleratorDevice:
|
||||||
|
return AcceleratorDevice(
|
||||||
|
kind="gpu",
|
||||||
|
vendor="nvidia",
|
||||||
|
device_id=device_id,
|
||||||
|
model="Test GPU",
|
||||||
|
memory_mb=16_384,
|
||||||
|
modes=(AcceleratorMode.EXCLUSIVE_DEVICE,),
|
||||||
|
capabilities={"compute": "9.0", "driver": "test"},
|
||||||
|
topology_group=topology_group,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def cpu_requirements(*, cpu_cores: int = 1, memory_mb: int = 256) -> ResourceRequirements:
|
||||||
|
return ResourceRequirements(
|
||||||
|
profile="cpu-v1",
|
||||||
|
cpu_cores=cpu_cores,
|
||||||
|
memory_mb=memory_mb,
|
||||||
|
scratch_mb=128,
|
||||||
|
architecture="x86-64",
|
||||||
|
environment_digest=ENVIRONMENT_DIGEST,
|
||||||
|
max_duration_seconds=120,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def gpu_requirements(*, accelerator_count: int) -> ResourceRequirements:
|
||||||
|
return ResourceRequirements(
|
||||||
|
profile="gpu-v1",
|
||||||
|
cpu_cores=1,
|
||||||
|
memory_mb=512,
|
||||||
|
scratch_mb=128,
|
||||||
|
accelerator_count=accelerator_count,
|
||||||
|
accelerator_kind="gpu",
|
||||||
|
accelerator_memory_mb=8_192,
|
||||||
|
accelerator_mode=AcceleratorMode.EXCLUSIVE_DEVICE,
|
||||||
|
architecture="x86-64",
|
||||||
|
environment_digest=ENVIRONMENT_DIGEST,
|
||||||
|
max_duration_seconds=120,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("device_id", ("GPU-0,GPU-1", "file:/dev/gpu0", "/dev/gpu0"))
|
||||||
|
def test_accelerator_ids_are_opaque_visibility_tokens(device_id: str) -> None:
|
||||||
|
with pytest.raises(ValueError, match="opaque"):
|
||||||
|
gpu(device_id)
|
||||||
|
|
||||||
|
|
||||||
|
def test_resource_inventory_and_requirements_round_trip_without_mutable_aliases() -> None:
|
||||||
|
capabilities = {"compute": "9.0"}
|
||||||
|
device = AcceleratorDevice(
|
||||||
|
kind="gpu",
|
||||||
|
vendor="nvidia",
|
||||||
|
device_id="gpu-0",
|
||||||
|
model="Test GPU",
|
||||||
|
memory_mb=16_384,
|
||||||
|
modes=(AcceleratorMode.EXCLUSIVE_DEVICE,),
|
||||||
|
capabilities=capabilities,
|
||||||
|
topology_group="socket-0",
|
||||||
|
)
|
||||||
|
inventory = ResourceInventory(
|
||||||
|
cpu_cores=8,
|
||||||
|
memory_mb=32_768,
|
||||||
|
scratch_mb=8_192,
|
||||||
|
architecture="x86-64",
|
||||||
|
accelerators=(device,),
|
||||||
|
environment_digests=(ENVIRONMENT_DIGEST,),
|
||||||
|
)
|
||||||
|
requirements = gpu_requirements(accelerator_count=1)
|
||||||
|
|
||||||
|
capabilities["compute"] = "mutated"
|
||||||
|
assert device.capabilities["compute"] == "9.0"
|
||||||
|
with pytest.raises(TypeError):
|
||||||
|
device.capabilities["compute"] = "mutated"
|
||||||
|
assert ResourceInventory.from_dict(inventory.to_dict()) == inventory
|
||||||
|
assert ResourceRequirements.from_dict(requirements.to_dict()) == requirements
|
||||||
|
assert requirements.eligibility_errors(inventory) == ()
|
||||||
|
|
||||||
|
incompatible = ResourceRequirements.from_dict(
|
||||||
|
{**requirements.to_dict(), "architecture": "arm64"}
|
||||||
|
)
|
||||||
|
assert incompatible.eligibility_errors(inventory) == ("architecture-mismatch",)
|
||||||
|
|
||||||
|
|
||||||
|
def test_failed_multi_accelerator_reservation_is_atomic_and_releases_nothing_partial() -> None:
|
||||||
|
inventory = ResourceInventory(
|
||||||
|
cpu_cores=4,
|
||||||
|
memory_mb=4_096,
|
||||||
|
scratch_mb=2_048,
|
||||||
|
architecture="x86-64",
|
||||||
|
accelerators=(gpu("gpu-0"), gpu("gpu-1")),
|
||||||
|
environment_digests=(ENVIRONMENT_DIGEST,),
|
||||||
|
)
|
||||||
|
pool = ResourcePool(inventory, max_concurrency=3)
|
||||||
|
first = pool.reserve("task/first", gpu_requirements(accelerator_count=1))
|
||||||
|
|
||||||
|
with pytest.raises(ResourceUnavailableError, match="accelerator-unavailable"):
|
||||||
|
pool.reserve("task/gang", gpu_requirements(accelerator_count=2))
|
||||||
|
|
||||||
|
assert pool.active_allocations() == (first,)
|
||||||
|
assert pool.release(first.allocation_id)
|
||||||
|
gang = pool.reserve("task/gang", gpu_requirements(accelerator_count=2))
|
||||||
|
assert gang.accelerator_ids == ("gpu-0", "gpu-1")
|
||||||
|
assert pool.active_allocations() == (gang,)
|
||||||
|
|
||||||
|
|
||||||
|
def test_resource_pool_enforces_aggregate_limits_under_concurrent_reservations() -> None:
|
||||||
|
inventory = ResourceInventory(
|
||||||
|
cpu_cores=4,
|
||||||
|
memory_mb=1_024,
|
||||||
|
scratch_mb=512,
|
||||||
|
architecture="x86-64",
|
||||||
|
environment_digests=(ENVIRONMENT_DIGEST,),
|
||||||
|
)
|
||||||
|
pool = ResourcePool(inventory, max_concurrency=8)
|
||||||
|
barrier = Barrier(8)
|
||||||
|
|
||||||
|
def attempt(index: int):
|
||||||
|
barrier.wait()
|
||||||
|
try:
|
||||||
|
return pool.reserve(f"task/{index}", cpu_requirements())
|
||||||
|
except ResourceUnavailableError:
|
||||||
|
return None
|
||||||
|
|
||||||
|
with ThreadPoolExecutor(max_workers=8) as executor:
|
||||||
|
results = tuple(executor.map(attempt, range(8)))
|
||||||
|
|
||||||
|
successful = tuple(result for result in results if result is not None)
|
||||||
|
assert len(successful) == 4
|
||||||
|
assert sum(item.cpu_cores for item in successful) == inventory.cpu_cores
|
||||||
|
assert sum(item.memory_mb for item in successful) <= inventory.memory_mb
|
||||||
|
assert sum(item.scratch_mb for item in successful) <= inventory.scratch_mb
|
||||||
|
assert pool.active_allocations() == tuple(sorted(successful, key=lambda item: item.task_key))
|
||||||
|
|
||||||
|
|
||||||
|
def test_resource_pool_slot_and_task_identity_limits_do_not_leak_capacity() -> None:
|
||||||
|
inventory = ResourceInventory(
|
||||||
|
cpu_cores=4,
|
||||||
|
memory_mb=2_048,
|
||||||
|
scratch_mb=1_024,
|
||||||
|
architecture="x86-64",
|
||||||
|
environment_digests=(ENVIRONMENT_DIGEST,),
|
||||||
|
)
|
||||||
|
pool = ResourcePool(inventory, max_concurrency=1)
|
||||||
|
allocation = pool.reserve("task/one", cpu_requirements())
|
||||||
|
|
||||||
|
with pytest.raises(ValueError, match="already has"):
|
||||||
|
pool.reserve("task/one", cpu_requirements())
|
||||||
|
with pytest.raises(ResourceUnavailableError, match="execution-slot-unavailable"):
|
||||||
|
pool.reserve("task/two", cpu_requirements())
|
||||||
|
assert pool.active_allocations() == (allocation,)
|
||||||
|
|
||||||
|
assert pool.release(allocation.allocation_id)
|
||||||
|
assert not pool.release(allocation.allocation_id)
|
||||||
|
replacement = pool.reserve("task/two", cpu_requirements())
|
||||||
|
assert replacement.task_key == "task/two"
|
||||||
|
|
||||||
|
|
||||||
|
def test_exclusive_gpu_and_its_partitions_share_one_conflict_domain() -> None:
|
||||||
|
full = AcceleratorDevice(
|
||||||
|
kind="gpu",
|
||||||
|
vendor="nvidia",
|
||||||
|
device_id="gpu-0",
|
||||||
|
model="Test GPU",
|
||||||
|
memory_mb=16_384,
|
||||||
|
modes=(AcceleratorMode.EXCLUSIVE_DEVICE, AcceleratorMode.PARTITION),
|
||||||
|
capabilities={},
|
||||||
|
)
|
||||||
|
partitions = tuple(
|
||||||
|
AcceleratorDevice(
|
||||||
|
kind="gpu",
|
||||||
|
vendor="nvidia",
|
||||||
|
device_id="gpu-0",
|
||||||
|
partition_id=f"mig-{index}",
|
||||||
|
model="Test MIG",
|
||||||
|
memory_mb=8_192,
|
||||||
|
modes=(AcceleratorMode.PARTITION,),
|
||||||
|
capabilities={},
|
||||||
|
)
|
||||||
|
for index in range(2)
|
||||||
|
)
|
||||||
|
inventory = ResourceInventory(
|
||||||
|
cpu_cores=4,
|
||||||
|
memory_mb=4_096,
|
||||||
|
scratch_mb=2_048,
|
||||||
|
architecture="x86-64",
|
||||||
|
accelerators=(full, *partitions),
|
||||||
|
environment_digests=(ENVIRONMENT_DIGEST,),
|
||||||
|
)
|
||||||
|
pool = ResourcePool(inventory, max_concurrency=3)
|
||||||
|
exclusive = pool.reserve("task/exclusive", gpu_requirements(accelerator_count=1))
|
||||||
|
partition_request = ResourceRequirements(
|
||||||
|
**{
|
||||||
|
**gpu_requirements(accelerator_count=1).to_dict(),
|
||||||
|
"accelerator_mode": AcceleratorMode.PARTITION,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
with pytest.raises(ResourceUnavailableError, match="accelerator-unavailable"):
|
||||||
|
pool.reserve("task/partition", partition_request)
|
||||||
|
pool.release(exclusive.allocation_id)
|
||||||
|
|
||||||
|
first = pool.reserve("task/partition-0", partition_request)
|
||||||
|
second = pool.reserve("task/partition-1", partition_request)
|
||||||
|
assert set(first.accelerator_ids + second.accelerator_ids) == {"mig-0", "mig-1"}
|
||||||
|
with pytest.raises(ResourceUnavailableError, match="accelerator-unavailable"):
|
||||||
|
pool.reserve("task/full", gpu_requirements(accelerator_count=1))
|
||||||
@@ -0,0 +1,744 @@
|
|||||||
|
"""Contract tests for SDK verifier decisions and built-in verifiers."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import hashlib
|
||||||
|
import json
|
||||||
|
import math
|
||||||
|
from collections.abc import Callable
|
||||||
|
from dataclasses import replace
|
||||||
|
from uuid import NAMESPACE_URL, uuid5
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from scimesh.sdk import (
|
||||||
|
ArtifactCollection,
|
||||||
|
ArtifactRef,
|
||||||
|
ArtifactSchema,
|
||||||
|
CandidateOutput,
|
||||||
|
CandidateOutputs,
|
||||||
|
CanonicalRecordVerifier,
|
||||||
|
ComponentRef,
|
||||||
|
ExactArtifactVerifier,
|
||||||
|
NumericTolerance,
|
||||||
|
NumericToleranceVerifier,
|
||||||
|
OutputManifest,
|
||||||
|
PortSpec,
|
||||||
|
Provenance,
|
||||||
|
SchemaRef,
|
||||||
|
TrustMode,
|
||||||
|
VerificationDecision,
|
||||||
|
VerificationBinding,
|
||||||
|
VerificationStatus,
|
||||||
|
VerifyContext,
|
||||||
|
WorkloadId,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _sha256(seed: str) -> str:
|
||||||
|
return hashlib.sha256(seed.encode("utf-8")).hexdigest()
|
||||||
|
|
||||||
|
|
||||||
|
OUTPUT_SCHEMA = ArtifactSchema(
|
||||||
|
ref=SchemaRef("verification-result", 1),
|
||||||
|
media_type="application/json",
|
||||||
|
encoding="utf-8",
|
||||||
|
max_bytes=1_024,
|
||||||
|
validator=ComponentRef("json-document", 1),
|
||||||
|
)
|
||||||
|
OUTPUT_PORT = PortSpec(OUTPUT_SCHEMA)
|
||||||
|
AUTHENTICATION_KEY = b"sdk-verification-test-key-000001"
|
||||||
|
JOB_ID = str(uuid5(NAMESPACE_URL, "verification-job"))
|
||||||
|
TASK_ID = str(uuid5(NAMESPACE_URL, "verification-task"))
|
||||||
|
EXECUTION_CONTRACT_DIGEST = _sha256("execution-contract")
|
||||||
|
|
||||||
|
|
||||||
|
def _artifact(seed: str, *, size_bytes: int = 16) -> ArtifactRef:
|
||||||
|
return ArtifactRef(
|
||||||
|
artifact_id=str(uuid5(NAMESPACE_URL, f"artifact:{seed}")),
|
||||||
|
sha256=_sha256(seed),
|
||||||
|
schema=OUTPUT_SCHEMA.ref,
|
||||||
|
media_type=OUTPUT_SCHEMA.media_type,
|
||||||
|
size_bytes=size_bytes,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _provenance(attempt: str) -> Provenance:
|
||||||
|
return Provenance(
|
||||||
|
workload=WorkloadId("verification-fixture", "1.0.0"),
|
||||||
|
sdk_api_version="1.0.0",
|
||||||
|
protocol_version="1.0",
|
||||||
|
manifest_schema_version=1,
|
||||||
|
workflow_schema_version=1,
|
||||||
|
verifier=ComponentRef("exact-artifact", 1),
|
||||||
|
artifact_schemas=(OUTPUT_SCHEMA.ref,),
|
||||||
|
package_digest=f"sha256:{_sha256('package')}",
|
||||||
|
manifest_digest=_sha256("manifest"),
|
||||||
|
environment_digest=f"sha256:{_sha256('environment')}",
|
||||||
|
worker_runtime={"attempt": attempt},
|
||||||
|
allocated_resource_ids=(f"cpu-{attempt}",),
|
||||||
|
parameters_digest=_sha256("parameters"),
|
||||||
|
input_collection_digest=_sha256("inputs"),
|
||||||
|
execution_contract_digest=EXECUTION_CONTRACT_DIGEST,
|
||||||
|
selected_features={"exact-verifier": "1.0.0"},
|
||||||
|
optional_fallbacks={},
|
||||||
|
job_id=JOB_ID,
|
||||||
|
task_id=TASK_ID,
|
||||||
|
started_at="2026-08-01T10:00:00Z",
|
||||||
|
finished_at="2026-08-01T10:00:01Z",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _manifest(
|
||||||
|
output_seed: str,
|
||||||
|
attempt: str,
|
||||||
|
*,
|
||||||
|
port_name: str = "result",
|
||||||
|
size_bytes: int = 16,
|
||||||
|
) -> OutputManifest:
|
||||||
|
return OutputManifest(
|
||||||
|
task_key="verify/0",
|
||||||
|
outputs={port_name: ArtifactCollection.single(_artifact(output_seed, size_bytes=size_bytes))},
|
||||||
|
metrics={"elapsed_seconds": 1.0},
|
||||||
|
provenance=_provenance(attempt),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _candidate(
|
||||||
|
output_seed: str,
|
||||||
|
attempt: str,
|
||||||
|
*,
|
||||||
|
owner: str | None,
|
||||||
|
candidate_id: str | None = None,
|
||||||
|
port_name: str = "result",
|
||||||
|
size_bytes: int = 16,
|
||||||
|
authenticated: bool = True,
|
||||||
|
) -> CandidateOutput:
|
||||||
|
resolved_candidate_id = candidate_id or str(
|
||||||
|
uuid5(NAMESPACE_URL, f"candidate:{attempt}")
|
||||||
|
)
|
||||||
|
resolved_owner_id = None if owner is None else str(uuid5(NAMESPACE_URL, f"owner:{owner}"))
|
||||||
|
manifest = _manifest(
|
||||||
|
output_seed,
|
||||||
|
attempt,
|
||||||
|
port_name=port_name,
|
||||||
|
size_bytes=size_bytes,
|
||||||
|
)
|
||||||
|
if resolved_owner_id is not None and authenticated:
|
||||||
|
return CandidateOutput.from_coordinator_record(
|
||||||
|
resolved_candidate_id,
|
||||||
|
resolved_owner_id,
|
||||||
|
manifest,
|
||||||
|
AUTHENTICATION_KEY,
|
||||||
|
)
|
||||||
|
return CandidateOutput(resolved_candidate_id, resolved_owner_id, manifest)
|
||||||
|
|
||||||
|
|
||||||
|
def _context(
|
||||||
|
*,
|
||||||
|
minimum_matches: int = 1,
|
||||||
|
reference: OutputManifest | None = None,
|
||||||
|
require_distinct_owners: bool = False,
|
||||||
|
trust_mode: TrustMode = TrustMode.TRUSTED,
|
||||||
|
) -> VerifyContext:
|
||||||
|
provenance = _provenance("binding")
|
||||||
|
return VerifyContext(
|
||||||
|
expected_outputs={"result": OUTPUT_PORT},
|
||||||
|
max_output_bytes=1_024,
|
||||||
|
minimum_matches=minimum_matches,
|
||||||
|
reference=reference,
|
||||||
|
require_distinct_owners=require_distinct_owners,
|
||||||
|
binding=VerificationBinding(
|
||||||
|
workload=provenance.workload,
|
||||||
|
task_key="verify/0",
|
||||||
|
package_digest=provenance.package_digest,
|
||||||
|
manifest_digest=provenance.manifest_digest,
|
||||||
|
environment_digest=provenance.environment_digest,
|
||||||
|
parameters_digest=provenance.parameters_digest,
|
||||||
|
input_collection_digest=provenance.input_collection_digest,
|
||||||
|
execution_contract_digest=provenance.execution_contract_digest,
|
||||||
|
selected_features=provenance.selected_features,
|
||||||
|
optional_fallbacks=provenance.optional_fallbacks,
|
||||||
|
job_id=provenance.job_id,
|
||||||
|
task_id=provenance.task_id,
|
||||||
|
verifier=provenance.verifier,
|
||||||
|
sdk_api_version=provenance.sdk_api_version,
|
||||||
|
protocol_version=provenance.protocol_version,
|
||||||
|
manifest_schema_version=provenance.manifest_schema_version,
|
||||||
|
workflow_schema_version=provenance.workflow_schema_version,
|
||||||
|
artifact_schemas=provenance.artifact_schemas,
|
||||||
|
trust_mode=trust_mode,
|
||||||
|
),
|
||||||
|
trust_mode=trust_mode,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_verification_decision_is_strict_immutable_and_round_trips() -> None:
|
||||||
|
source = {"summary": {"counts": [1, 2]}}
|
||||||
|
decision = VerificationDecision(
|
||||||
|
VerificationStatus.REJECTED,
|
||||||
|
ComponentRef("test-verifier", 1),
|
||||||
|
"comparison-failed",
|
||||||
|
source,
|
||||||
|
)
|
||||||
|
|
||||||
|
source["summary"]["counts"].append(3) # type: ignore[index, union-attr]
|
||||||
|
|
||||||
|
assert decision.evidence["summary"]["counts"] == (1, 2)
|
||||||
|
assert VerificationDecision.from_dict(decision.to_dict()) == decision
|
||||||
|
with pytest.raises(TypeError):
|
||||||
|
decision.evidence["new"] = True # type: ignore[index]
|
||||||
|
with pytest.raises(TypeError):
|
||||||
|
decision.evidence["summary"]["new"] = True # type: ignore[index]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
("status", "accepted_digest"),
|
||||||
|
[
|
||||||
|
(VerificationStatus.ACCEPTED, None),
|
||||||
|
(VerificationStatus.REJECTED, "a" * 64),
|
||||||
|
(VerificationStatus.INCONCLUSIVE, "a" * 64),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_verification_decision_enforces_digest_status_invariant(
|
||||||
|
status: VerificationStatus,
|
||||||
|
accepted_digest: str | None,
|
||||||
|
) -> None:
|
||||||
|
with pytest.raises(ValueError, match="accepted_digest|accepted verification"):
|
||||||
|
VerificationDecision(
|
||||||
|
status,
|
||||||
|
ComponentRef("test-verifier", 1),
|
||||||
|
"test-result",
|
||||||
|
{},
|
||||||
|
accepted_digest,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"unsafe",
|
||||||
|
[
|
||||||
|
"/home/worker/private.log",
|
||||||
|
"https://worker.invalid/evidence",
|
||||||
|
"run-123/tasks/map/result.csv",
|
||||||
|
"path=attempts/job-123/private.txt",
|
||||||
|
"https%253A%252F%252Fworker.invalid%252Fevidence%253Ftoken%253Dsecret",
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_verification_decision_rejects_private_locations(unsafe: str) -> None:
|
||||||
|
with pytest.raises(ValueError, match="URI or local path"):
|
||||||
|
VerificationDecision(
|
||||||
|
VerificationStatus.REJECTED,
|
||||||
|
ComponentRef("test-verifier", 1),
|
||||||
|
"unsafe-evidence",
|
||||||
|
{"detail": unsafe},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_verification_decision_bounds_evidence_and_rejects_unknown_fields() -> None:
|
||||||
|
with pytest.raises(ValueError, match="exceeds 16 KiB"):
|
||||||
|
VerificationDecision(
|
||||||
|
VerificationStatus.REJECTED,
|
||||||
|
ComponentRef("test-verifier", 1),
|
||||||
|
"oversized-evidence",
|
||||||
|
{"detail": "x" * 17_000},
|
||||||
|
)
|
||||||
|
|
||||||
|
payload = VerificationDecision(
|
||||||
|
VerificationStatus.REJECTED,
|
||||||
|
ComponentRef("test-verifier", 1),
|
||||||
|
"test-result",
|
||||||
|
{},
|
||||||
|
).to_dict()
|
||||||
|
payload["unexpected"] = True
|
||||||
|
with pytest.raises(ValueError, match="unknown unexpected"):
|
||||||
|
VerificationDecision.from_dict(payload)
|
||||||
|
|
||||||
|
|
||||||
|
def test_candidate_output_envelope_is_strict_and_round_trips() -> None:
|
||||||
|
candidate = _candidate("output", "attempt-one", owner="owner-one")
|
||||||
|
|
||||||
|
decoded_candidate = CandidateOutput.from_dict(candidate.to_dict())
|
||||||
|
assert decoded_candidate.to_dict() == candidate.to_dict()
|
||||||
|
assert not decoded_candidate.coordinator_authenticated
|
||||||
|
candidates = CandidateOutputs(candidates=(candidate,))
|
||||||
|
decoded = CandidateOutputs.from_dict(candidates.to_dict())
|
||||||
|
assert not decoded.candidates[0].coordinator_authenticated
|
||||||
|
assert CandidateOutputs.from_authenticated_dict(
|
||||||
|
candidates.to_dict(), AUTHENTICATION_KEY
|
||||||
|
) == candidates
|
||||||
|
with pytest.raises(ValueError, match="authentication failed"):
|
||||||
|
CandidateOutputs.from_authenticated_dict(candidates.to_dict(), b"x" * 32)
|
||||||
|
assert candidates.manifests == (candidate.manifest,)
|
||||||
|
|
||||||
|
with pytest.raises(ValueError, match="opaque coordinator identity"):
|
||||||
|
CandidateOutput("../worker-path", candidate.owner_id, candidate.manifest)
|
||||||
|
|
||||||
|
|
||||||
|
def test_trusted_single_manifest_compatibility_uses_an_anonymous_envelope() -> None:
|
||||||
|
manifest = _manifest("output", "trusted")
|
||||||
|
candidates = CandidateOutputs((manifest,))
|
||||||
|
|
||||||
|
assert candidates.manifests == (manifest,)
|
||||||
|
assert candidates.candidates[0].owner_id is None
|
||||||
|
decision = ExactArtifactVerifier().verify(_context(), candidates)
|
||||||
|
assert decision.status is VerificationStatus.ACCEPTED
|
||||||
|
|
||||||
|
|
||||||
|
def test_raw_manifests_cannot_form_a_quorum() -> None:
|
||||||
|
with pytest.raises(ValueError, match="one trusted candidate"):
|
||||||
|
CandidateOutputs(
|
||||||
|
(
|
||||||
|
_manifest("output", "trusted-one"),
|
||||||
|
_manifest("output", "trusted-two"),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_multi_vote_context_automatically_requires_distinct_owners() -> None:
|
||||||
|
assert _context(minimum_matches=2).require_distinct_owners
|
||||||
|
assert _context(require_distinct_owners=True).require_distinct_owners
|
||||||
|
|
||||||
|
with pytest.raises(ValueError, match="boolean"):
|
||||||
|
_context(require_distinct_owners=1) # type: ignore[arg-type]
|
||||||
|
|
||||||
|
with pytest.raises(ValueError, match="coordinator binding"):
|
||||||
|
VerifyContext(
|
||||||
|
expected_outputs={"result": OUTPUT_PORT},
|
||||||
|
max_output_bytes=1_024,
|
||||||
|
minimum_matches=2,
|
||||||
|
)
|
||||||
|
|
||||||
|
with pytest.raises(ValueError, match="at least two"):
|
||||||
|
_context(trust_mode=TrustMode.UNTRUSTED_QUORUM)
|
||||||
|
|
||||||
|
|
||||||
|
def test_exact_verifier_reports_no_candidates_as_inconclusive() -> None:
|
||||||
|
decision = ExactArtifactVerifier().verify(_context(), CandidateOutputs(()))
|
||||||
|
|
||||||
|
assert decision.status is VerificationStatus.INCONCLUSIVE
|
||||||
|
assert decision.reason_code == "no-candidates"
|
||||||
|
assert decision.evidence == {"candidate_count": 0, "invalid_count": 0}
|
||||||
|
assert decision.accepted_digest is None
|
||||||
|
|
||||||
|
|
||||||
|
def test_exact_verifier_rejects_candidates_that_violate_output_contract() -> None:
|
||||||
|
invalid = _manifest("same-output", "invalid", port_name="undeclared")
|
||||||
|
|
||||||
|
decision = ExactArtifactVerifier().verify(_context(), CandidateOutputs((invalid,)))
|
||||||
|
|
||||||
|
assert decision.status is VerificationStatus.REJECTED
|
||||||
|
assert decision.reason_code == "no-valid-candidates"
|
||||||
|
assert decision.evidence == {"candidate_count": 1, "invalid_count": 1}
|
||||||
|
|
||||||
|
|
||||||
|
def test_exact_verifier_rejects_candidate_from_another_scientific_binding() -> None:
|
||||||
|
candidate = _candidate("same-output", "other-job", owner="owner-one")
|
||||||
|
forged_provenance = replace(
|
||||||
|
candidate.manifest.provenance,
|
||||||
|
parameters_digest=_sha256("different-parameters"),
|
||||||
|
)
|
||||||
|
forged = replace(
|
||||||
|
candidate,
|
||||||
|
manifest=replace(candidate.manifest, provenance=forged_provenance),
|
||||||
|
)
|
||||||
|
|
||||||
|
decision = ExactArtifactVerifier().verify(_context(), CandidateOutputs((forged,)))
|
||||||
|
|
||||||
|
assert decision.status is VerificationStatus.REJECTED
|
||||||
|
assert decision.reason_code == "no-valid-candidates"
|
||||||
|
|
||||||
|
|
||||||
|
def test_exact_verifier_accepts_unique_quorum_and_ignores_invalid_candidates() -> None:
|
||||||
|
first = _candidate("same-output", "one", owner="owner-one")
|
||||||
|
second = _candidate("same-output", "two", owner="owner-two")
|
||||||
|
minority = _candidate("different-output", "three", owner="owner-three")
|
||||||
|
invalid = _candidate(
|
||||||
|
"same-output",
|
||||||
|
"four",
|
||||||
|
owner="owner-four",
|
||||||
|
port_name="undeclared",
|
||||||
|
)
|
||||||
|
|
||||||
|
decision = ExactArtifactVerifier().verify(
|
||||||
|
_context(minimum_matches=2),
|
||||||
|
CandidateOutputs((first, second, minority, invalid)),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert first.manifest.digest == second.manifest.digest
|
||||||
|
assert first.manifest.manifest_digest != second.manifest.manifest_digest
|
||||||
|
assert decision.status is VerificationStatus.ACCEPTED
|
||||||
|
assert decision.reason_code == "quorum-match"
|
||||||
|
assert decision.accepted_digest == first.manifest.digest
|
||||||
|
assert decision.evidence == {
|
||||||
|
"matched": 2,
|
||||||
|
"required": 2,
|
||||||
|
"distinct_digests": 2,
|
||||||
|
"invalid_count": 1,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def test_exact_verifier_rejects_conflicting_quorums() -> None:
|
||||||
|
candidates = CandidateOutputs(
|
||||||
|
(
|
||||||
|
_candidate("group-a", "a-one", owner="a-one"),
|
||||||
|
_candidate("group-a", "a-two", owner="a-two"),
|
||||||
|
_candidate("group-b", "b-one", owner="b-one"),
|
||||||
|
_candidate("group-b", "b-two", owner="b-two"),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
decision = ExactArtifactVerifier().verify(_context(minimum_matches=2), candidates)
|
||||||
|
|
||||||
|
assert decision.status is VerificationStatus.REJECTED
|
||||||
|
assert decision.reason_code == "conflicting-quorums"
|
||||||
|
assert decision.accepted_digest is None
|
||||||
|
assert decision.evidence["largest_group"] == 2
|
||||||
|
assert decision.evidence["distinct_digests"] == 2
|
||||||
|
|
||||||
|
|
||||||
|
def test_exact_verifier_distinguishes_insufficient_evidence_from_reference_mismatch() -> None:
|
||||||
|
reference = _manifest("reference", "reference")
|
||||||
|
one_mismatch = CandidateOutputs((_candidate("other", "one", owner="owner-one"),))
|
||||||
|
two_mismatches = CandidateOutputs(
|
||||||
|
(
|
||||||
|
_candidate("other-a", "two", owner="owner-two"),
|
||||||
|
_candidate("other-b", "three", owner="owner-three"),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
insufficient = ExactArtifactVerifier().verify(
|
||||||
|
_context(minimum_matches=2, reference=reference),
|
||||||
|
one_mismatch,
|
||||||
|
)
|
||||||
|
rejected = ExactArtifactVerifier().verify(
|
||||||
|
_context(minimum_matches=2, reference=reference),
|
||||||
|
two_mismatches,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert (insufficient.status, insufficient.reason_code) == (
|
||||||
|
VerificationStatus.INCONCLUSIVE,
|
||||||
|
"insufficient-evidence",
|
||||||
|
)
|
||||||
|
assert (rejected.status, rejected.reason_code) == (
|
||||||
|
VerificationStatus.REJECTED,
|
||||||
|
"reference-mismatch",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_exact_verifier_accepts_declared_reference_quorum() -> None:
|
||||||
|
reference = _manifest("reference", "reference")
|
||||||
|
candidates = CandidateOutputs(
|
||||||
|
(
|
||||||
|
_candidate("reference", "worker-one", owner="owner-one"),
|
||||||
|
_candidate("reference", "worker-two", owner="owner-two"),
|
||||||
|
_candidate("other", "worker-three", owner="owner-three"),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
decision = ExactArtifactVerifier().verify(
|
||||||
|
_context(minimum_matches=2, reference=reference),
|
||||||
|
candidates,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert decision.status is VerificationStatus.ACCEPTED
|
||||||
|
assert decision.reason_code == "reference-match"
|
||||||
|
assert decision.accepted_digest == reference.digest
|
||||||
|
assert decision.evidence == {"matched": 2, "required": 2, "invalid_count": 0}
|
||||||
|
|
||||||
|
|
||||||
|
def test_candidate_outputs_rejects_duplicate_candidate_ids() -> None:
|
||||||
|
candidate = _candidate("output", "one", owner="owner-one")
|
||||||
|
replay = _candidate(
|
||||||
|
"different-output",
|
||||||
|
"two",
|
||||||
|
owner="owner-two",
|
||||||
|
candidate_id=candidate.candidate_id,
|
||||||
|
)
|
||||||
|
|
||||||
|
with pytest.raises(ValueError, match="candidate_id values must be unique"):
|
||||||
|
CandidateOutputs((candidate, replay))
|
||||||
|
|
||||||
|
|
||||||
|
def test_exact_quorum_rejects_candidates_without_authenticated_owners() -> None:
|
||||||
|
candidates = CandidateOutputs(
|
||||||
|
(
|
||||||
|
_candidate("output", "one", owner=None),
|
||||||
|
_candidate("output", "two", owner="owner-two"),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
decision = ExactArtifactVerifier().verify(
|
||||||
|
_context(minimum_matches=2),
|
||||||
|
candidates,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert decision.status is VerificationStatus.REJECTED
|
||||||
|
assert decision.reason_code == "coordinator-authentication-required"
|
||||||
|
assert decision.evidence == {"candidate_count": 2, "unauthenticated_count": 1}
|
||||||
|
|
||||||
|
|
||||||
|
def test_exact_verifier_counts_at_most_one_vote_per_owner() -> None:
|
||||||
|
candidates = CandidateOutputs(
|
||||||
|
(
|
||||||
|
_candidate("output", "one", owner="same-owner"),
|
||||||
|
_candidate("output", "two", owner="same-owner"),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
decision = ExactArtifactVerifier().verify(
|
||||||
|
_context(minimum_matches=2),
|
||||||
|
candidates,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert decision.status is VerificationStatus.INCONCLUSIVE
|
||||||
|
assert decision.reason_code == "insufficient-evidence"
|
||||||
|
assert decision.evidence == {
|
||||||
|
"largest_group": 1,
|
||||||
|
"required": 2,
|
||||||
|
"distinct_digests": 1,
|
||||||
|
"invalid_count": 0,
|
||||||
|
"duplicate_owner_candidates": 1,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def test_exact_verifier_rejects_owner_equivocation_without_leaking_identity() -> None:
|
||||||
|
owner = "equivocating-owner"
|
||||||
|
candidates = CandidateOutputs(
|
||||||
|
(
|
||||||
|
_candidate("output-a", "one", owner=owner),
|
||||||
|
_candidate("output-b", "two", owner=owner),
|
||||||
|
_candidate("output-a", "three", owner="honest-owner"),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
decision = ExactArtifactVerifier().verify(
|
||||||
|
_context(minimum_matches=2),
|
||||||
|
candidates,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert decision.status is VerificationStatus.REJECTED
|
||||||
|
assert decision.reason_code == "owner-equivocation"
|
||||||
|
assert decision.accepted_digest is None
|
||||||
|
assert decision.evidence == {
|
||||||
|
"candidate_count": 3,
|
||||||
|
"equivocating_owner_count": 1,
|
||||||
|
}
|
||||||
|
assert candidates.candidates[0].owner_id not in json.dumps(decision.to_dict())
|
||||||
|
|
||||||
|
|
||||||
|
def test_numeric_verifier_accepts_nested_values_with_absolute_and_relative_tolerance() -> None:
|
||||||
|
verifier = NumericToleranceVerifier(NumericTolerance(absolute=0.001, relative=0.01))
|
||||||
|
|
||||||
|
decision = verifier.verify_values(
|
||||||
|
{"energies": [10.0, 0.05], "converged": True},
|
||||||
|
{"energies": [10.05, 0.0505], "converged": True},
|
||||||
|
)
|
||||||
|
|
||||||
|
assert decision.status is VerificationStatus.ACCEPTED
|
||||||
|
assert decision.reason_code == "within-tolerance"
|
||||||
|
assert decision.accepted_digest is not None
|
||||||
|
assert decision.evidence == {
|
||||||
|
"absolute": 0.001,
|
||||||
|
"relative": 0.01,
|
||||||
|
"max_ulps": 0,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def test_numeric_verifier_supports_ulp_tolerance() -> None:
|
||||||
|
adjacent = math.nextafter(1.0, 2.0)
|
||||||
|
verifier = NumericToleranceVerifier(NumericTolerance(max_ulps=1))
|
||||||
|
|
||||||
|
decision = verifier.verify_values(1.0, adjacent)
|
||||||
|
|
||||||
|
assert decision.status is VerificationStatus.ACCEPTED
|
||||||
|
|
||||||
|
|
||||||
|
def test_numeric_verifier_reports_bounded_location_and_error_evidence() -> None:
|
||||||
|
verifier = NumericToleranceVerifier(NumericTolerance(absolute=0.1))
|
||||||
|
|
||||||
|
decision = verifier.verify_values(
|
||||||
|
{"matrix": [[1.0, 2.0]]},
|
||||||
|
{"matrix": [[1.0, 2.5]]},
|
||||||
|
)
|
||||||
|
|
||||||
|
assert decision.status is VerificationStatus.REJECTED
|
||||||
|
assert decision.reason_code == "numeric-mismatch"
|
||||||
|
assert decision.evidence["location"] == "$.matrix[0][1]"
|
||||||
|
assert decision.evidence["absolute_error"] == pytest.approx(0.5)
|
||||||
|
assert decision.evidence["allowed_error"] == pytest.approx(0.1)
|
||||||
|
assert isinstance(decision.evidence["ulp_distance"], int)
|
||||||
|
|
||||||
|
|
||||||
|
def test_numeric_verifier_rejects_shape_changes_before_value_comparison() -> None:
|
||||||
|
verifier = NumericToleranceVerifier(NumericTolerance())
|
||||||
|
|
||||||
|
decision = verifier.verify_values(
|
||||||
|
{"energy": 1.0, "iterations": 4},
|
||||||
|
{"energy": 1.0, "converged": True},
|
||||||
|
)
|
||||||
|
|
||||||
|
assert decision.status is VerificationStatus.REJECTED
|
||||||
|
assert decision.reason_code == "shape-mismatch"
|
||||||
|
assert decision.evidence == {
|
||||||
|
"location": "$",
|
||||||
|
"missing_keys": ("iterations",),
|
||||||
|
"extra_keys": ("converged",),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def test_numeric_verifier_applies_declared_nan_policy() -> None:
|
||||||
|
reject = NumericToleranceVerifier(NumericTolerance(nan_policy="reject"))
|
||||||
|
equal = NumericToleranceVerifier(NumericTolerance(nan_policy="equal"))
|
||||||
|
|
||||||
|
rejected = reject.verify_values(float("nan"), float("nan"))
|
||||||
|
accepted = equal.verify_values(float("nan"), float("nan"))
|
||||||
|
|
||||||
|
assert (rejected.status, rejected.reason_code) == (
|
||||||
|
VerificationStatus.REJECTED,
|
||||||
|
"nan-policy",
|
||||||
|
)
|
||||||
|
assert accepted.status is VerificationStatus.ACCEPTED
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"factory",
|
||||||
|
[
|
||||||
|
lambda: NumericTolerance(absolute=-0.1),
|
||||||
|
lambda: NumericTolerance(relative=float("inf")),
|
||||||
|
lambda: NumericTolerance(max_ulps=True),
|
||||||
|
lambda: NumericTolerance(nan_policy="propagate"),
|
||||||
|
lambda: NumericTolerance(max_elements=0),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_numeric_tolerance_rejects_ambiguous_or_non_finite_policy(
|
||||||
|
factory: Callable[[], NumericTolerance],
|
||||||
|
) -> None:
|
||||||
|
with pytest.raises(ValueError, match="numeric tolerance"):
|
||||||
|
factory()
|
||||||
|
|
||||||
|
|
||||||
|
def test_numeric_verifier_rejects_unrepresentable_integer_without_raising() -> None:
|
||||||
|
verifier = NumericToleranceVerifier(NumericTolerance())
|
||||||
|
|
||||||
|
decision = verifier.verify_values(10**10_000, 10**10_000)
|
||||||
|
|
||||||
|
assert decision.status is VerificationStatus.REJECTED
|
||||||
|
assert decision.reason_code == "numeric-range"
|
||||||
|
assert decision.evidence["location"] == "$"
|
||||||
|
|
||||||
|
|
||||||
|
def test_numeric_verifier_rejects_shapes_above_its_manifest_bound() -> None:
|
||||||
|
verifier = NumericToleranceVerifier(NumericTolerance(max_elements=3))
|
||||||
|
|
||||||
|
decision = verifier.verify_values([1, 2, 3], [1, 2, 3])
|
||||||
|
|
||||||
|
assert decision.status is VerificationStatus.REJECTED
|
||||||
|
assert decision.reason_code == "element-limit"
|
||||||
|
assert decision.evidence == {"max_elements": 3}
|
||||||
|
assert verifier.configuration["max_elements"] == 3
|
||||||
|
|
||||||
|
|
||||||
|
def test_numeric_verifier_does_not_round_mixed_integer_and_float_values() -> None:
|
||||||
|
verifier = NumericToleranceVerifier(NumericTolerance())
|
||||||
|
|
||||||
|
decision = verifier.verify_values(2**53 + 1, float(2**53))
|
||||||
|
|
||||||
|
assert decision.status is VerificationStatus.REJECTED
|
||||||
|
assert decision.reason_code == "numeric-mismatch"
|
||||||
|
|
||||||
|
|
||||||
|
def test_numeric_verifier_rejects_non_json_mapping_keys() -> None:
|
||||||
|
verifier = NumericToleranceVerifier(NumericTolerance())
|
||||||
|
|
||||||
|
decision = verifier.verify_values({1: 2.0}, {1: 2.0})
|
||||||
|
|
||||||
|
assert decision.status is VerificationStatus.REJECTED
|
||||||
|
assert decision.reason_code == "type-mismatch"
|
||||||
|
|
||||||
|
|
||||||
|
def test_numeric_verifier_implements_manifest_verifier_protocol_with_loader() -> None:
|
||||||
|
reference = _manifest("reference", "reference")
|
||||||
|
candidate = _manifest("candidate", "candidate")
|
||||||
|
values = {
|
||||||
|
reference.outputs["result"].items[0].artifact.sha256: {"energy": 1.0},
|
||||||
|
candidate.outputs["result"].items[0].artifact.sha256: {"energy": 1.0001},
|
||||||
|
}
|
||||||
|
|
||||||
|
def load(manifest: OutputManifest) -> object:
|
||||||
|
return values[manifest.outputs["result"].items[0].artifact.sha256]
|
||||||
|
|
||||||
|
verifier = NumericToleranceVerifier(NumericTolerance(absolute=0.001), load)
|
||||||
|
decision = verifier.verify(
|
||||||
|
_context(reference=reference),
|
||||||
|
CandidateOutputs((candidate,)),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert decision.status is VerificationStatus.ACCEPTED
|
||||||
|
assert decision.accepted_digest == reference.digest
|
||||||
|
|
||||||
|
|
||||||
|
def test_manifest_verifiers_fail_closed_without_artifact_loaders() -> None:
|
||||||
|
candidate = CandidateOutputs((_manifest("candidate", "candidate"),))
|
||||||
|
|
||||||
|
numeric = NumericToleranceVerifier(NumericTolerance()).verify(_context(), candidate)
|
||||||
|
canonical = CanonicalRecordVerifier(_canonical_json_record).verify(_context(), candidate)
|
||||||
|
|
||||||
|
assert numeric.reason_code == "loader-unavailable"
|
||||||
|
assert canonical.reason_code == "loader-unavailable"
|
||||||
|
|
||||||
|
|
||||||
|
def _canonical_json_record(record: object) -> bytes:
|
||||||
|
return json.dumps(record, sort_keys=True, separators=(",", ":")).encode("utf-8")
|
||||||
|
|
||||||
|
|
||||||
|
def test_canonical_record_verifier_accepts_normalized_records() -> None:
|
||||||
|
verifier = CanonicalRecordVerifier(_canonical_json_record)
|
||||||
|
expected = [{"name": "molecule", "score": 0.75}, {"id": 2}]
|
||||||
|
actual = [{"score": 0.75, "name": "molecule"}, {"id": 2}]
|
||||||
|
|
||||||
|
decision = verifier.verify_records(expected, actual)
|
||||||
|
|
||||||
|
assert decision.status is VerificationStatus.ACCEPTED
|
||||||
|
assert decision.reason_code == "canonical-match"
|
||||||
|
assert decision.evidence == {"records": 2}
|
||||||
|
assert decision.accepted_digest is not None
|
||||||
|
|
||||||
|
|
||||||
|
def test_canonical_record_verifier_rejects_content_or_count_mismatch() -> None:
|
||||||
|
verifier = CanonicalRecordVerifier(_canonical_json_record)
|
||||||
|
|
||||||
|
decision = verifier.verify_records([{"id": 1}, {"id": 2}], [{"id": 1}])
|
||||||
|
|
||||||
|
assert decision.status is VerificationStatus.REJECTED
|
||||||
|
assert decision.reason_code == "canonical-mismatch"
|
||||||
|
assert decision.evidence == {"expected_records": 2, "actual_records": 1}
|
||||||
|
|
||||||
|
|
||||||
|
def test_canonical_record_verifier_enforces_record_limit() -> None:
|
||||||
|
verifier = CanonicalRecordVerifier(_canonical_json_record, max_records=2)
|
||||||
|
|
||||||
|
decision = verifier.verify_records([1, 2, 3], [1, 2, 3])
|
||||||
|
|
||||||
|
assert decision.status is VerificationStatus.REJECTED
|
||||||
|
assert decision.reason_code == "record-limit"
|
||||||
|
assert decision.evidence == {"max_records": 2}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"canonicalizer",
|
||||||
|
[
|
||||||
|
lambda _record: "not-bytes",
|
||||||
|
lambda _record: (_ for _ in ()).throw(ValueError("/private/worker/path")),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_canonical_record_verifier_sanitizes_canonicalization_failures(
|
||||||
|
canonicalizer: Callable[[object], object],
|
||||||
|
) -> None:
|
||||||
|
verifier = CanonicalRecordVerifier(canonicalizer) # type: ignore[arg-type]
|
||||||
|
|
||||||
|
decision = verifier.verify_records([{"id": 1}], [{"id": 1}])
|
||||||
|
|
||||||
|
assert decision.status is VerificationStatus.REJECTED
|
||||||
|
assert decision.reason_code == "canonicalization-failed"
|
||||||
|
assert decision.evidence == {}
|
||||||
|
assert decision.accepted_digest is None
|
||||||
Reference in New Issue
Block a user