from __future__ import annotations import base64 import gzip import hashlib import os import pickle from pathlib import Path from typing import Any from .contracts import NodePayload ENCODING = "pickle+gzip+base64" FILE_ENCODING = "pickle+gzip+file" def _encode_bytes(raw: bytes) -> str: return base64.b64encode(gzip.compress(raw)).decode("ascii") def _decode_bytes(payload: str) -> bytes: return gzip.decompress(base64.b64decode(payload.encode("ascii"))) def _farm_root() -> Path: return Path(os.environ.get("KEY_FARM_ROOT", "/data/farm")) def _decode_file_payload(payload: str) -> bytes: path = Path(payload) if not path.is_absolute(): path = _farm_root() / path return gzip.decompress(path.read_bytes()) def pack_object(obj: Any) -> tuple[str, str]: raw = pickle.dumps(obj, protocol=pickle.HIGHEST_PROTOCOL) return _encode_bytes(raw), hashlib.sha256(raw).hexdigest() def unpack_object(payload: str, sha256: str | None = None) -> Any: raw = _decode_bytes(payload) digest = hashlib.sha256(raw).hexdigest() if sha256 and digest != sha256: raise ValueError(f"payload digest mismatch: {digest} != {sha256}") return pickle.loads(raw) def pack_node(node: Any) -> NodePayload: node_id = str(getattr(node, "id", "")) if not node_id: raise ValueError("node must expose a stable id before KEY generation dispatch") return pack_item(node, node_id) def pack_item(item: Any, item_id: str | None = None) -> NodePayload: payload, digest = pack_object(item) stable_id = item_id or str(getattr(item, "id", "")) or digest[:16] return NodePayload(node_id=stable_id, encoding=ENCODING, payload=payload, sha256=digest) def unpack_node(node_payload: NodePayload) -> Any: if node_payload.encoding == ENCODING: return unpack_object(node_payload.payload, node_payload.sha256) if node_payload.encoding == FILE_ENCODING: raw = _decode_file_payload(node_payload.payload) digest = hashlib.sha256(raw).hexdigest() if node_payload.sha256 and digest != node_payload.sha256: raise ValueError(f"payload digest mismatch: {digest} != {node_payload.sha256}") return pickle.loads(raw) else: raise ValueError(f"unsupported node encoding: {node_payload.encoding}")