File size: 6,662 Bytes
e476358
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Offline, bounded-memory file/index integrity check. Does not execute model code."""
import argparse
import hashlib
import json
import math
from pathlib import Path, PurePosixPath
import re
import struct


def unique(pairs):
    result = {}
    for key, value in pairs:
        if key in result:
            raise ValueError('duplicate JSON key')
        result[key] = value
    return result


def read_json(path):
    return json.loads(path.read_text(encoding='utf-8'), object_pairs_hook=unique)


def safe_path(root, name):
    if not isinstance(name, str) or not name or '\\' in name:
        raise ValueError('unsafe file path')
    p = PurePosixPath(name)
    if p.is_absolute() or '..' in p.parts or str(p) != name:
        raise ValueError('unsafe file path')
    target = root.joinpath(*p.parts)
    if any(parent.is_symlink() for parent in (target, *target.parents) if parent != root.parent):
        raise ValueError('symlink path not allowed')
    target.resolve().relative_to(root.resolve())
    return target


def file_hashes(path):
    size = path.stat().st_size
    sha = hashlib.sha256()
    blob = hashlib.sha1(f'blob {size}\0'.encode())
    with path.open('rb') as stream:
        for chunk in iter(lambda: stream.read(8 * 1024**2), b''):
            sha.update(chunk)
            blob.update(chunk)
    return size, sha.hexdigest(), blob.hexdigest()


def tensor_header(path):
    size = path.stat().st_size
    with path.open('rb') as stream:
        prefix = stream.read(8)
        if len(prefix) != 8:
            raise ValueError('short safetensors prefix')
        n = struct.unpack('<Q', prefix)[0]
        if n > 16 * 1024**2 or 8+n > size:
            raise ValueError('invalid header boundary')
        header = json.loads(stream.read(n), object_pairs_hook=unique)
    widths = {'BF16': 2, 'F16': 2, 'F32': 4, 'F64': 8, 'U32': 4, 'I32': 4,
              'U8': 1, 'I8': 1, 'U16': 2, 'I16': 2, 'U64': 8, 'I64': 8, 'BOOL': 1}
    spans, names = [], set()
    for name, value in header.items():
        if name == '__metadata__':
            continue
        shape = value['shape']
        if not isinstance(shape, list) or any(type(x) is not int or x < 0 for x in shape):
            raise ValueError('invalid tensor shape')
        start, end = value['data_offsets']
        if type(start) is not int or type(end) is not int or not 0 <= start <= end <= size-8-n:
            raise ValueError('invalid tensor payload boundary')
        if end-start != math.prod(shape)*widths[value['dtype']]:
            raise ValueError('dtype/shape byte count mismatch')
        spans.append((start, end))
        names.add(name)
    cursor = 0
    for start, end in sorted(spans):
        if start != cursor:
            raise ValueError('payload gap or overlap')
        cursor = end
    if cursor != size-8-n:
        raise ValueError('unreferenced or truncated payload')
    return names


def verify(root, expected_manifest_sha256=None):
    root = Path(root).absolute()
    errors, checked = [], []
    try:
        manifest_path = safe_path(root, 'UPLOAD_MANIFEST.json')
        _, manifest_sha, _ = file_hashes(manifest_path)
        if expected_manifest_sha256 and expected_manifest_sha256 != manifest_sha:
            raise ValueError('trusted manifest hash mismatch')
        manifest = read_json(manifest_path)
        entries = manifest['files']
        names = [e['path'] for e in entries]
        if len(names) != len(set(names)):
            raise ValueError('duplicate manifest file')
        if any(type(e['bytes']) is not int or e['bytes'] < 0 for e in entries):
            raise ValueError('invalid manifest size')
        if len(entries) != manifest['file_count'] or sum(e['bytes'] for e in entries) != manifest['total_bytes']:
            raise ValueError('manifest aggregates mismatch')
        for e in entries:
            if e['path'] == 'UPLOAD_MANIFEST.json' or not re.fullmatch('[a-f0-9]{64}', e['sha256']):
                raise ValueError('invalid manifest entry')
            p = safe_path(root, e['path'])
            if not p.is_file():
                errors.append({'file': e['path'], 'reason': 'missing file'})
                continue
            size, sha, blob = file_hashes(p)
            if size != e['bytes'] or sha != e['sha256'] or (e.get('git_blob_sha1') and blob != e['git_blob_sha1']):
                errors.append({'file': e['path'], 'reason': 'size or full-file hash mismatch'})
                continue
            if p.suffix == '.json':
                read_json(p)
            checked.append(e['path'])
        index_name = 'model.safetensors.index.json'
        if index_name not in names:
            errors.append({'file': index_name, 'reason': 'required index not in manifest'})
        elif index_name in checked:
            index = read_json(root/index_name)['weight_map']
            if not isinstance(index, dict) or not index:
                raise ValueError('invalid weight map')
            shards = set(index.values())
            expected_shards = {n for n in names if n.endswith('.safetensors')}
            if shards != expected_shards:
                errors.append({'reason': 'index/manifest shard set mismatch'})
            actual = {}
            for shard in sorted(shards):
                safe_path(root, shard)
                if shard not in checked:
                    errors.append({'file': shard, 'reason': 'referenced shard missing or not hash-verified'})
                    continue
                for tensor in tensor_header(root/shard):
                    if tensor in actual:
                        raise ValueError('duplicate tensor across shards')
                    actual[tensor] = shard
            if actual != index:
                errors.append({'reason': 'actual tensor inventory differs from weight map'})
    except (OSError, ValueError, TypeError, KeyError, AttributeError) as exc:
        errors.append({'reason': 'invalid or incomplete package', 'exception_type': type(exc).__name__})
    return {'status': 'FAIL' if errors else 'PASS', 'errors': errors,
            'fully_hashed_files': len(checked), 'manifest_pinned': bool(expected_manifest_sha256),
            'scope': 'File integrity and tensor index only; not generation, finite-value or quality validation'}


if __name__ == '__main__':
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument('snapshot', type=Path)
    parser.add_argument('--manifest-sha256')
    args = parser.parse_args()
    result = verify(args.snapshot, args.manifest_sha256)
    print(json.dumps(result, indent=2))
    raise SystemExit(0 if result['status'] == 'PASS' else 1)