"""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: 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