from __future__ import annotations import json from typing import Any from jsonschema import Draft202012Validator, FormatChecker from hub_core.contracts import CONTRACT_VERSION, extension_contract_root from hub_core.runtime.models import RegistryRegistration class ContractValidator: """Validate runtime registration input against the packaged contract.""" def __init__(self, *, runtime_contract_version: str = CONTRACT_VERSION) -> None: self._runtime_contract_version = _parse_semver(runtime_contract_version) contract_root = extension_contract_root() schema_root = contract_root.joinpath("schemas") self._descriptor = _validator(schema_root.joinpath("hub-descriptor.schema.json")) self._manifest = _validator(schema_root.joinpath("hub-manifest.schema.json")) catalog = json.loads( contract_root.joinpath("catalogs", "event-types.json").read_text(encoding="utf-8") ) self._event_families = { entry["type"]: entry["family"] for entry in catalog["event_types"] } def validate_registration(self, registration: RegistryRegistration) -> None: self._descriptor.validate(registration.descriptor) self._manifest.validate(registration.manifest) descriptor_id = registration.descriptor.get("reuse_surface_id") manifest_id = registration.manifest.get("reuse_surface_id") if descriptor_id != manifest_id: raise ValueError("descriptor and manifest reuse_surface_id must match") self._negotiate_contract_version(registration.descriptor) def _negotiate_contract_version(self, descriptor: dict[str, Any]) -> None: version_min = _parse_semver(descriptor["contract_version_min"]) version_max = _parse_semver(descriptor["contract_version_max"]) if version_min > version_max: raise ValueError( "descriptor contract_version_min must not exceed contract_version_max" ) if not (version_min <= self._runtime_contract_version <= version_max): raise ValueError( "descriptor requires contract version range " f"{descriptor['contract_version_min']}-{descriptor['contract_version_max']}, " f"incompatible with runtime contract version {CONTRACT_VERSION}" ) def validate_event_family(self, event_type: str, expected_family: str) -> None: actual_family = self._event_families.get(event_type) if actual_family is None: raise ValueError(f"event type '{event_type}' is not cataloged") if actual_family != expected_family: raise ValueError( f"event type '{event_type}' belongs to '{actual_family}', not '{expected_family}'" ) def _validator(resource: Any) -> Draft202012Validator: schema = json.loads(resource.read_text(encoding="utf-8")) Draft202012Validator.check_schema(schema) return Draft202012Validator(schema, format_checker=FormatChecker()) def _parse_semver(value: str) -> tuple[int, int, int]: core = value.split("+", 1)[0].split("-", 1)[0] major, minor, patch = core.split(".") return (int(major), int(minor), int(patch))