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())