File size: 15,165 Bytes
2d5c26a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Verify and publish an existing training snapshot using implicit Hub login.

Use a new --output directory for each explicit retry. No model or GPU is loaded.
The original artifact staging is never modified; final releases are not replaced.
"""
import argparse
import hashlib
import logging
import os
from pathlib import Path, PurePosixPath
import re
import shutil
import signal
import subprocess
import sys
import tarfile
import tempfile

ROOT = Path(__file__).resolve().parents[1]
if str(ROOT) not in sys.path:
    sys.path.insert(0, str(ROOT))
from scripts.publish_hf_final import digest, read_json, write_json, utc, deadline_alarm, stop_requested

EXPORT_ROOT = Path.home() / 'ai/opensysone/exports'


def relative_name(value):
    path = PurePosixPath(value)
    if path.is_absolute() or '..' in path.parts or str(path) != value or not path.parts:
        raise ValueError('Unsafe artifact path')
    return value


def verify_artifacts(directory, repo_id):
    directory = Path(directory).resolve()
    manifest = read_json(directory / 'backup-manifest.json')
    if manifest.get('format') != 'opensysone-backup-v1' or manifest.get('target_repository') != repo_id:
        raise ValueError('Snapshot format or target repository mismatch')
    sums = {}
    for line in (directory / 'SHA256SUMS').read_text().splitlines():
        match = re.fullmatch(r'([0-9a-f]{64})  (.+)', line)
        if not match or match[2] in sums:
            raise ValueError('Invalid or duplicate checksum entry')
        sums[relative_name(match[2])] = match[1]
    expected = set(manifest['files']) | {'backup-manifest.json'}
    actual = {str(path.relative_to(directory)) for path in directory.rglob('*') if path.is_file()}
    if set(sums) != expected or actual != expected | {'SHA256SUMS'}:
        raise ValueError('Snapshot file inventory differs from manifest/checksum index')
    for name, expected_hash in sums.items():
        path = directory / relative_name(name)
        if path.is_symlink() or not path.resolve().is_relative_to(directory) or digest(path) != expected_hash:
            raise ValueError('Snapshot checksum or path mismatch')
        if name in manifest['files']:
            record = manifest['files'][name]
            if record['sha256'] != expected_hash or record['bytes'] != path.stat().st_size:
                raise ValueError('Snapshot manifest disagrees with checksum/size')
    return manifest


def git_bytes(source_root, *arguments):
    return subprocess.check_output(['git', *arguments], cwd=source_root, stderr=subprocess.DEVNULL, timeout=60)


def archive_source(source_root, revision, destination, browse=None):
    entries = git_bytes(source_root, 'ls-tree', '-r', revision).decode().splitlines()
    for entry in entries:
        metadata, name = entry.split('\t', 1)
        path = Path(relative_name(name))
        if (not metadata.startswith(('100644 blob ', '100755 blob ')) or
                path.suffix.lower() in ('.pt', '.pth', '.bin', '.safetensors', '.pem', '.key') or
                path.name in ('.env', 'token', 'id_rsa', 'id_ed25519', '.netrc') or path.name.startswith('.env.')):
            raise ValueError('Committed source includes a prohibited file or symlink')
    destination.mkdir(parents=True, exist_ok=True)
    archive = destination / 'source.tar.gz'
    subprocess.run(['git', 'archive', '--format=tar.gz', '--output=' + str(archive), revision],
                   cwd=source_root, check=True, capture_output=True, timeout=60)
    source_files = {}
    with tarfile.open(archive, 'r:gz') as bundle:
        for member in bundle:
            if member.isdir():
                continue
            name = relative_name(member.name)
            if not member.isfile():
                raise ValueError('Source archive has a nonregular file')
            content = bundle.extractfile(member).read()
            source_files[name] = {'size': len(content), 'sha256': hashlib.sha256(content).hexdigest()}
            if browse is not None:
                target = browse / name
                target.parent.mkdir(parents=True, exist_ok=True)
                target.write_bytes(content)
    write_json(destination / 'manifest.json', {'source_commit': revision, 'archive_sha256': digest(archive),
                                              'archive_size': archive.stat().st_size, 'files': source_files})
    return source_files


def prepare(artifacts, repo_id, source_root=ROOT, export_root=EXPORT_ROOT):
    artifacts = Path(artifacts).resolve()
    manifest = verify_artifacts(artifacts, repo_id)
    source_root, export_root = Path(source_root), Path(export_root)
    head = git_bytes(source_root, 'rev-parse', 'HEAD').decode().strip()
    revisions = {entry['source_commit'] for entry in manifest['artifacts']}
    if not revisions or any(not re.fullmatch('[0-9a-f]{40}', revision) for revision in revisions):
        raise ValueError('Snapshot lacks exact training source revisions')
    export_root.mkdir(parents=True, exist_ok=True)
    stage = Path(tempfile.mkdtemp(prefix=artifacts.name + '-publish-', dir=export_root))
    repository = stage / 'repository'
    repository.mkdir()
    snapshot_prefix = 'snapshots/' + artifacts.name
    shutil.copytree(artifacts, repository / snapshot_prefix)
    # Verify the independent copy, including every weight, before any upload.
    verify_artifacts(repository / snapshot_prefix, repo_id)
    sources = {}
    for revision in sorted(revisions | {head}):
        sources[revision] = archive_source(source_root, revision, repository / 'sources' / revision,
                                          repository / 'source' if revision == head else None)
    for artifact in manifest['artifacts']:
        for name, expected in artifact['source_file_sha256'].items():
            if sources[artifact['source_commit']].get(name, {}).get('sha256') != expected:
                raise ValueError('Archived training source differs from artifact provenance')
    card = repository / 'source/HF_MODEL_CARD.md'
    if not card.is_file():
        raise ValueError('Committed HEAD lacks HF_MODEL_CARD.md')
    files = {str(path.relative_to(repository)): {'size': path.stat().st_size, 'sha256': digest(path)}
             for path in sorted(repository.rglob('*')) if path.is_file()}
    proof_path = f'publications/{artifacts.name}/{head}/manifest.json'
    proof = {'format': 'opensysone-snapshot-publication-v1', 'created_utc': utc(), 'repo_id': repo_id,
             'snapshot_id': artifacts.name, 'snapshot_path': snapshot_prefix, 'source_commit': head,
             'training_source_commits': sorted(revisions), 'files': files}
    (repository / proof_path).parent.mkdir(parents=True)
    write_json(repository / proof_path, proof)
    write_json(stage / 'stage.json', {'repository': str(repository), 'manifest': proof_path})
    for path in repository.rglob('*'):
        if path.is_file():
            path.chmod(0o444)
    return stage


def verify_stage(stage):
    location = read_json(stage / 'stage.json')
    repository = stage / 'repository'
    proof_path = relative_name(location['manifest'])
    proof = read_json(repository / proof_path)
    actual = {str(path.relative_to(repository)) for path in repository.rglob('*') if path.is_file()}
    if actual != set(proof['files']) | {proof_path}:
        raise ValueError('Publication stage inventory changed')
    for name, expected in proof['files'].items():
        path = repository / relative_name(name)
        if (path.is_symlink() or not path.resolve().is_relative_to(repository.resolve()) or
                path.stat().st_size != expected['size'] or digest(path) != expected['sha256']):
            raise ValueError('Publication stage checksum mismatch')
    verify_artifacts(repository / proof['snapshot_path'], proof['repo_id'])
    return repository, proof_path, proof


def license_line(card):
    if not card.startswith('---\n') or '\n---' not in card[4:]:
        return None
    frontmatter = card[4:].split('\n---', 1)[0]
    values = re.findall(r'^license:\s*(.*?)\s*$', frontmatter, flags=re.MULTILINE)
    return values[0] if len(values) == 1 else None


def ensure_no_final(api, repo_id, revision):
    if api.get_paths_info(repo_id, ['FINAL_MODEL.json'], revision=revision, repo_type='model'):
        raise ValueError('A final release already exists; training snapshot pointer will not replace it')


def publish(api, repo_id, stage, progress, operation_factory=None):
    repository, proof_path, proof = verify_stage(Path(stage))
    if proof['repo_id'] != repo_id:
        raise ValueError('Publication stage repository mismatch')
    initial = api.repo_info(repo_id, repo_type='model', timeout=30)
    if initial.private is not True:
        raise ValueError('Expected the existing private repository')
    ensure_no_final(api, repo_id, initial.sha)
    progress('uploading_payload', stage_path=str(stage), source_commit=proof['source_commit'])
    uploaded = api.upload_folder(repo_id=repo_id, repo_type='model', folder_path=str(repository),
        commit_message='Back up verified OpenSysOne training snapshot and pinned source', parent_commit=initial.sha)
    progress('verifying_payload', payload_commit=uploaded.oid)
    files = {**proof['files'], proof_path: {'size': (repository / proof_path).stat().st_size,
                                          'sha256': digest(repository / proof_path)}}
    names = list(files)
    remote = []
    for start in range(0, len(names), 250):
        remote.extend(api.get_paths_info(repo_id, names[start:start + 250], revision=uploaded.oid, repo_type='model'))
    by_path = {item.path: item for item in remote}
    for name, expected in files.items():
        item = by_path.get(name)
        if item is None or item.size != expected['size']:
            raise ValueError('Remote payload size mismatch')
        lfs = getattr(item, 'lfs', None)
        if name.endswith('.pt') and lfs is None:
            raise ValueError('Remote weights lack LFS checksum metadata')
        actual = (lfs.get('sha256') if isinstance(lfs, dict) else lfs.sha256) if lfs else item.blob_id
        expected_hash = expected['sha256'] if lfs else digest(repository / name, 'sha1', git_blob=True)
        if actual != expected_hash:
            raise ValueError('Remote payload checksum mismatch')
    tiny = [proof_path, proof['snapshot_path'] + '/backup-manifest.json', proof['snapshot_path'] + '/SHA256SUMS']
    tiny.extend('sources/' + revision + '/manifest.json' for revision in set(proof['training_source_commits']) | {proof['source_commit']})
    for name in tiny:
        downloaded = api.hf_hub_download(repo_id, name, revision=uploaded.oid, repo_type='model', etag_timeout=30)
        if Path(downloaded).read_bytes() != (repository / name).read_bytes():
            raise ValueError('Remote manifest bytes mismatch')
    current = api.repo_info(repo_id, repo_type='model', timeout=30)
    if current.private != initial.private:
        raise ValueError('Repository visibility changed')
    ensure_no_final(api, repo_id, current.sha)
    card = (repository / 'source/HF_MODEL_CARD.md').read_text()
    if api.get_paths_info(repo_id, ['README.md'], revision=current.sha, repo_type='model'):
        previous = Path(api.hf_hub_download(repo_id, 'README.md', revision=current.sha,
                                          repo_type='model', etag_timeout=30)).read_text()
        if '<!-- opensysone-final:start -->' in previous:
            raise ValueError('Final model card is already published')
        if license_line(previous) != license_line(card):
            raise ValueError('Committed model card would change existing license metadata')
    pointer = {'repo_id': repo_id, 'snapshot_id': proof['snapshot_id'], 'path': proof['snapshot_path'],
               'payload_commit': uploaded.oid, 'source_commit': proof['source_commit'],
               'training_source_commits': proof['training_source_commits'], 'manifest_path': proof_path,
               'manifest_sha256': files[proof_path]['sha256'], 'published_utc': utc(), 'calibrated': False}
    if operation_factory is None:
        from huggingface_hub import CommitOperationAdd
        operation_factory = CommitOperationAdd
    import json
    pointer_bytes = (json.dumps(pointer, indent=2) + '\n').encode()
    progress('publishing_pointer')
    committed = api.create_commit(repo_id, repo_type='model', parent_commit=current.sha,
        operations=[operation_factory(path_in_repo='CURRENT_SNAPSHOT.json', path_or_fileobj=pointer_bytes),
                    operation_factory(path_in_repo='README.md', path_or_fileobj=card.encode())],
        commit_message='Publish verified OpenSysOne training snapshot pointer and model card')
    for name, content in (('CURRENT_SNAPSHOT.json', pointer_bytes), ('README.md', card.encode())):
        downloaded = api.hf_hub_download(repo_id, name, revision=committed.oid, repo_type='model', etag_timeout=30)
        if Path(downloaded).read_bytes() != content:
            raise ValueError('Published pointer/card bytes mismatch')
    return {**pointer, 'pointer_commit': committed.oid, 'repository_private': initial.private}


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument('--artifacts', required=True)
    parser.add_argument('--repo-id', required=True)
    parser.add_argument('--output', required=True, help='New directory for this publication attempt status')
    args = parser.parse_args()
    output = Path(args.output).resolve()
    output.mkdir(parents=True, exist_ok=False)
    state = {'status': 'running', 'pid': os.getpid(), 'command': [sys.executable, *sys.argv],
             'started_utc': utc(), 'repo_id': args.repo_id, 'artifacts': str(Path(args.artifacts).resolve())}
    def progress(stage, **values):
        state.update(stage=stage, heartbeat_utc=utc(), **values)
        write_json(output / 'state.json', state)
    code = 1
    try:
        Path('/proc/self/oom_score_adj').write_text('0')
        signal.signal(signal.SIGALRM, deadline_alarm)
        signal.signal(signal.SIGTERM, stop_requested)
        signal.alarm(1800)
        progress('verifying_and_staging')
        stage = prepare(args.artifacts, args.repo_id)
        os.environ.update(HF_HUB_DISABLE_PROGRESS_BARS='1', HF_HUB_DISABLE_XET='1',
                          HF_HUB_DOWNLOAD_TIMEOUT='60', HF_HUB_ETAG_TIMEOUT='30')
        from huggingface_hub import HfApi, set_client_factory
        import httpx
        for name in ('huggingface_hub', 'httpx', 'httpcore'):
            logging.getLogger(name).setLevel(logging.CRITICAL)
        set_client_factory(lambda: httpx.Client(timeout=httpx.Timeout(60, connect=15), follow_redirects=True))
        result = publish(HfApi(), args.repo_id, stage, progress)
        progress('complete', status='complete', finished_utc=utc(), **result)
        code = 0
    except BaseException as error:
        progress('failed', status='failed', error_type=type(error).__name__, finished_utc=utc())
    finally:
        signal.alarm(0)
        (output / 'exit_code').write_text(str(code) + '\n')
    return code


if __name__ == '__main__':
    raise SystemExit(main())