File size: 15,106 Bytes
2c76aad e77d2f1 2c76aad e77d2f1 2c76aad e77d2f1 2c76aad e77d2f1 2c76aad e77d2f1 2c76aad e77d2f1 2c76aad e77d2f1 2c76aad e77d2f1 2c76aad e77d2f1 2c76aad e77d2f1 2c76aad e77d2f1 2c76aad e77d2f1 2c76aad e77d2f1 2c76aad e77d2f1 2c76aad e77d2f1 2c76aad e77d2f1 2c76aad e77d2f1 2c76aad e77d2f1 2c76aad | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 | """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 collections.abc import Iterable
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",
}
OPTIONAL_FIELDS = {"tensor_manifest"}
KNOWN_FIELDS = REQUIRED_FIELDS | OPTIONAL_FIELDS
ALLOWED_DTYPES = {"f16", "bf16", "f32", "i8", "u8"}
DTYPE_BYTES = {"f16": 2, "bf16": 2, "f32": 4, "i8": 1, "u8": 1}
ALLOWED_LAYOUTS = {"layer-major-k-then-v"}
MAX_METADATA_BYTES = 1_048_576
MANIFEST_FIELDS = {"version", "byte_order", "batch_size", "cache_class", "attention_type", "tensors"}
TENSOR_FIELDS = {"name", "kind", "layer", "shape", "strides", "offset", "nbytes", "dtype"}
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 _exact_fields(value: Mapping[str, Any], expected: set[str], label: str) -> None:
missing = expected.difference(value)
extra = set(value).difference(expected)
if missing or extra:
raise CacheFormatError(f"{label} fields differ: missing={sorted(missing)}, extra={sorted(extra)}")
def _positive_integer(value: Any, label: str) -> int:
if not isinstance(value, int) or isinstance(value, bool) or value <= 0:
raise CacheFormatError(f"{label} must be a positive integer")
return value
def _contiguous_byte_strides(shape: list[int], item_bytes: int) -> list[int]:
strides = [item_bytes] * len(shape)
for index in range(len(shape) - 2, -1, -1):
strides[index] = strides[index + 1] * shape[index + 1]
return strides
def _validate_tensor_manifest(metadata: Mapping[str, Any], manifest: Any) -> int:
if not isinstance(manifest, Mapping):
raise CacheFormatError("tensor_manifest must be a mapping")
manifest = dict(manifest)
_exact_fields(manifest, MANIFEST_FIELDS, "tensor_manifest")
if manifest["version"] != 1:
raise CacheFormatError("tensor_manifest version must be 1")
if manifest["byte_order"] not in {"little", "big"}:
raise CacheFormatError("tensor_manifest byte_order must be little or big")
batch_size = _positive_integer(manifest["batch_size"], "tensor_manifest batch_size")
for field in ("cache_class", "attention_type"):
if not isinstance(manifest[field], str) or not manifest[field].strip():
raise CacheFormatError(f"tensor_manifest {field} must be a non-empty string")
tensors = manifest["tensors"]
if not isinstance(tensors, list):
raise CacheFormatError("tensor_manifest tensors must be a list")
expected_count = metadata["layers"] * 2
if len(tensors) != expected_count:
raise CacheFormatError(f"tensor_manifest must contain exactly {expected_count} key/value tensors")
expected_shape = [batch_size, metadata["kv_heads"], metadata["sequence_length"], metadata["head_dim"]]
expected_offset = 0
seen_names: set[str] = set()
seen_slots: set[tuple[int, str]] = set()
for index, raw_tensor in enumerate(tensors):
if not isinstance(raw_tensor, Mapping):
raise CacheFormatError(f"tensor_manifest tensor {index} must be a mapping")
tensor = dict(raw_tensor)
_exact_fields(tensor, TENSOR_FIELDS, f"tensor_manifest tensor {index}")
name = tensor["name"]
if not isinstance(name, str) or not name:
raise CacheFormatError(f"tensor_manifest tensor {index} name must be non-empty")
if name in seen_names:
raise CacheFormatError(f"tensor_manifest tensor name is duplicated: {name}")
seen_names.add(name)
kind = tensor["kind"]
if kind not in {"key", "value"}:
raise CacheFormatError(f"tensor_manifest tensor {index} kind must be key or value")
layer = tensor["layer"]
if not isinstance(layer, int) or isinstance(layer, bool) or not 0 <= layer < metadata["layers"]:
raise CacheFormatError(f"tensor_manifest tensor {index} layer is out of range")
slot = (layer, kind)
if slot in seen_slots:
raise CacheFormatError(f"tensor_manifest duplicates layer {layer} {kind}")
seen_slots.add(slot)
shape = tensor["shape"]
if shape != expected_shape:
raise CacheFormatError(f"tensor_manifest tensor {index} shape must equal {expected_shape}")
dtype = tensor["dtype"]
if dtype != metadata["dtype"]:
raise CacheFormatError(f"tensor_manifest tensor {index} dtype must match metadata dtype")
expected_nbytes = DTYPE_BYTES[dtype]
for dimension in shape:
expected_nbytes *= _positive_integer(dimension, f"tensor_manifest tensor {index} shape dimension")
if tensor["nbytes"] != expected_nbytes:
raise CacheFormatError(f"tensor_manifest tensor {index} nbytes must equal {expected_nbytes}")
expected_strides = _contiguous_byte_strides(shape, DTYPE_BYTES[dtype])
if tensor["strides"] != expected_strides:
raise CacheFormatError(
f"tensor_manifest tensor {index} must use C-contiguous byte strides {expected_strides}"
)
if tensor["offset"] != expected_offset:
raise CacheFormatError(
f"tensor_manifest tensor {index} offsets must be contiguous; expected {expected_offset}"
)
expected_offset += expected_nbytes
expected_slots = {(layer, kind) for layer in range(metadata["layers"]) for kind in ("key", "value")}
if seen_slots != expected_slots:
raise CacheFormatError("tensor_manifest must contain one key and one value tensor for every layer")
return expected_offset
def _validate_payload_length(metadata: Mapping[str, Any], payload_length: int) -> None:
manifest = metadata.get("tensor_manifest")
if manifest is None:
return
expected = _validate_tensor_manifest(metadata, manifest)
if payload_length != expected:
raise CacheFormatError(
f"tensor manifest declares {expected} payload bytes but container has {payload_length}"
)
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(KNOWN_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")
if "tensor_manifest" in value:
_validate_tensor_manifest(value, value["tensor_manifest"])
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")
return write_cache_stream(Path(path), metadata, (payload,))
def write_cache_stream(path: str | os.PathLike[str], metadata: Mapping[str, Any], chunks: Iterable[bytes | bytearray | memoryview]) -> None:
destination = Path(path)
metadata_dict = _plain_dict(metadata)
metadata_bytes = _metadata_bytes(metadata_dict)
destination.parent.mkdir(parents=True, exist_ok=True)
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
placeholder = PREFIX.pack(MAGIC, len(metadata_bytes), 0, b"\x00" * 32)
temporary.write(placeholder)
temporary.write(metadata_bytes)
sha = hashlib.sha256()
length = 0
for chunk in chunks:
if not isinstance(chunk, (bytes, bytearray, memoryview)):
raise CacheFormatError("chunk must be bytes, bytearray, or memoryview")
chunk_bytes = bytes(chunk)
temporary.write(chunk_bytes)
sha.update(chunk_bytes)
length += len(chunk_bytes)
_validate_payload_length(metadata_dict, length)
digest = sha.digest()
temporary.seek(0)
temporary.write(PREFIX.pack(MAGIC, len(metadata_bytes), length, digest))
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")
_validate_payload_length(metadata, len(payload))
return metadata, payload, actual_digest.hex()
def read_cache(path: str | os.PathLike[str], expected_identity: Mapping[str, Any] | None = None) -> tuple[dict[str, Any], bytes]:
metadata, payload, _ = _decode(path)
if expected_identity is not None:
if not isinstance(expected_identity, Mapping):
raise CacheCompatibilityError("expected_identity must be a mapping")
expected = dict(expected_identity)
unknown = set(expected).difference(KNOWN_FIELDS)
if unknown:
raise CacheCompatibilityError(f"expected_identity has unknown fields: {sorted(unknown)}")
try:
_plain_dict({**metadata, **expected})
except CacheFormatError as error:
raise CacheCompatibilityError(f"expected_identity is invalid: {error}") from error
differences = [field for field in sorted(expected) 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)
manifest = metadata.get("tensor_manifest")
receipt = {
"format": "KVC1",
"metadata": metadata,
"payload_bytes": len(payload),
"payload_sha256": digest,
}
if manifest is not None:
receipt["tensor_manifest_verified"] = True
receipt["tensor_count"] = len(manifest["tensors"])
return receipt
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())
|