| """Portable, fail-closed raw KV-cache container. |
| |
| KVC1 stores already-materialized key/value tensor bytes plus the exact model, |
| tokenizer, tensor geometry, RoPE, dtype, layout, and sequence-position identity |
| needed to decide whether a runtime may safely reuse them. Cross-architecture |
| translation is deliberately out of scope for this byte container. |
| """ |
|
|
| from __future__ import annotations |
|
|
| import argparse |
| import hashlib |
| import hmac |
| import json |
| import os |
| from pathlib import Path |
| import struct |
| import tempfile |
| from typing import Any, Mapping |
|
|
|
|
| MAGIC = b"KVC1" |
| PREFIX = struct.Struct(">4sIQ32s") |
| REQUIRED_FIELDS = { |
| "model_revision", |
| "tokenizer_sha256", |
| "rope_theta", |
| "layers", |
| "kv_heads", |
| "head_dim", |
| "dtype", |
| "layout", |
| "sequence_start", |
| "sequence_length", |
| } |
| ALLOWED_DTYPES = {"f16", "bf16", "f32", "i8", "u8"} |
| ALLOWED_LAYOUTS = {"layer-major-k-then-v"} |
| MAX_METADATA_BYTES = 1_048_576 |
|
|
|
|
| class CacheFormatError(ValueError): |
| """The container is malformed, incomplete, corrupt, or unsupported.""" |
|
|
|
|
| class CacheCompatibilityError(ValueError): |
| """The container is valid but does not match the requested runtime.""" |
|
|
|
|
| def _plain_dict(metadata: Mapping[str, Any]) -> dict[str, Any]: |
| if not isinstance(metadata, Mapping): |
| raise CacheFormatError("metadata must be a mapping") |
| value = dict(metadata) |
| missing = REQUIRED_FIELDS.difference(value) |
| extra = set(value).difference(REQUIRED_FIELDS) |
| if missing or extra: |
| raise CacheFormatError(f"metadata fields differ: missing={sorted(missing)}, extra={sorted(extra)}") |
| if not isinstance(value["model_revision"], str) or "@" not in value["model_revision"]: |
| raise CacheFormatError("model_revision must identify an immutable revision") |
| tokenizer_hash = value["tokenizer_sha256"] |
| if not isinstance(tokenizer_hash, str) or len(tokenizer_hash) != 64: |
| raise CacheFormatError("tokenizer_sha256 must contain 64 hexadecimal characters") |
| try: |
| int(tokenizer_hash, 16) |
| except ValueError as error: |
| raise CacheFormatError("tokenizer_sha256 is not hexadecimal") from error |
| if not isinstance(value["rope_theta"], (int, float)) or isinstance(value["rope_theta"], bool) or value["rope_theta"] <= 0: |
| raise CacheFormatError("rope_theta must be positive") |
| for field in ("layers", "kv_heads", "head_dim", "sequence_length"): |
| if not isinstance(value[field], int) or isinstance(value[field], bool) or value[field] <= 0: |
| raise CacheFormatError(f"{field} must be a positive integer") |
| if not isinstance(value["sequence_start"], int) or isinstance(value["sequence_start"], bool) or value["sequence_start"] < 0: |
| raise CacheFormatError("sequence_start must be a non-negative integer") |
| if value["dtype"] not in ALLOWED_DTYPES: |
| raise CacheFormatError("unsupported dtype") |
| if value["layout"] not in ALLOWED_LAYOUTS: |
| raise CacheFormatError("unsupported layout") |
| return value |
|
|
|
|
| def _metadata_bytes(metadata: Mapping[str, Any]) -> bytes: |
| try: |
| encoded = json.dumps(_plain_dict(metadata), sort_keys=True, separators=(",", ":"), allow_nan=False).encode("utf-8") |
| except (TypeError, ValueError) as error: |
| if isinstance(error, CacheFormatError): |
| raise |
| raise CacheFormatError("metadata is not canonical JSON") from error |
| if len(encoded) > MAX_METADATA_BYTES: |
| raise CacheFormatError("metadata is too large") |
| return encoded |
|
|
|
|
| def write_cache(path: str | os.PathLike[str], metadata: Mapping[str, Any], payload: bytes) -> None: |
| """Atomically publish one KVC1 generation.""" |
| if type(payload) is not bytes: |
| raise CacheFormatError("payload must be raw bytes") |
| destination = Path(path) |
| destination.parent.mkdir(parents=True, exist_ok=True) |
| metadata_bytes = _metadata_bytes(metadata) |
| digest = hashlib.sha256(payload).digest() |
| prefix = PREFIX.pack(MAGIC, len(metadata_bytes), len(payload), digest) |
| temporary_name: str | None = None |
| try: |
| with tempfile.NamedTemporaryFile( |
| mode="wb", |
| prefix=f".{destination.name}.", |
| suffix=".tmp", |
| dir=destination.parent, |
| delete=False, |
| ) as temporary: |
| temporary_name = temporary.name |
| temporary.write(prefix) |
| temporary.write(metadata_bytes) |
| temporary.write(payload) |
| temporary.flush() |
| os.fsync(temporary.fileno()) |
| os.replace(temporary_name, destination) |
| temporary_name = None |
| finally: |
| if temporary_name is not None: |
| try: |
| os.unlink(temporary_name) |
| except FileNotFoundError: |
| pass |
|
|
|
|
| def _decode(path: str | os.PathLike[str]) -> tuple[dict[str, Any], bytes, str]: |
| try: |
| raw = Path(path).read_bytes() |
| except OSError as error: |
| raise CacheFormatError(f"cache could not be read: {error}") from error |
| if len(raw) < PREFIX.size: |
| raise CacheFormatError("container is truncated") |
| try: |
| magic, metadata_length, payload_length, expected_digest = PREFIX.unpack_from(raw) |
| except struct.error as error: |
| raise CacheFormatError("container prefix is malformed") from error |
| if magic != MAGIC: |
| raise CacheFormatError("unsupported container magic or version") |
| if metadata_length == 0 or metadata_length > MAX_METADATA_BYTES: |
| raise CacheFormatError("metadata length is invalid") |
| expected_length = PREFIX.size + metadata_length + payload_length |
| if len(raw) != expected_length: |
| raise CacheFormatError("container length does not match its header") |
| metadata_raw = raw[PREFIX.size:PREFIX.size + metadata_length] |
| payload = raw[PREFIX.size + metadata_length:] |
| try: |
| decoded = json.loads(metadata_raw.decode("utf-8")) |
| except (UnicodeDecodeError, json.JSONDecodeError) as error: |
| raise CacheFormatError("metadata is not valid UTF-8 JSON") from error |
| metadata = _plain_dict(decoded) |
| if _metadata_bytes(metadata) != metadata_raw: |
| raise CacheFormatError("metadata is not in canonical form") |
| actual_digest = hashlib.sha256(payload).digest() |
| if not hmac.compare_digest(actual_digest, expected_digest): |
| raise CacheFormatError("payload checksum mismatch") |
| return metadata, payload, actual_digest.hex() |
|
|
|
|
| def read_cache(path: str | os.PathLike[str], expected_metadata: Mapping[str, Any] | None = None) -> tuple[dict[str, Any], bytes]: |
| metadata, payload, _ = _decode(path) |
| if expected_metadata is not None: |
| expected = _plain_dict(expected_metadata) |
| differences = [field for field in sorted(REQUIRED_FIELDS) if metadata[field] != expected[field]] |
| if differences: |
| raise CacheCompatibilityError(f"cache is incompatible: {', '.join(differences)}") |
| return metadata, payload |
|
|
|
|
| def inspect_cache(path: str | os.PathLike[str]) -> dict[str, Any]: |
| metadata, payload, digest = _decode(path) |
| return { |
| "format": "KVC1", |
| "metadata": metadata, |
| "payload_bytes": len(payload), |
| "payload_sha256": digest, |
| } |
|
|
|
|
| def main() -> int: |
| parser = argparse.ArgumentParser(description="Inspect a portable raw KV-cache container.") |
| subparsers = parser.add_subparsers(dest="command", required=True) |
| inspect_parser = subparsers.add_parser("inspect") |
| inspect_parser.add_argument("path") |
| args = parser.parse_args() |
| if args.command == "inspect": |
| try: |
| print(json.dumps(inspect_cache(args.path), sort_keys=True)) |
| except (CacheFormatError, CacheCompatibilityError) as error: |
| parser.exit(1, f"kvcache: {error}\n") |
| return 0 |
|
|
|
|
| if __name__ == "__main__": |
| raise SystemExit(main()) |
|
|