"""Copy mounted inference files sequentially and verify Hub content hashes.""" import hashlib import json from pathlib import Path PROTOCOL = 'sequential-local-copy-hub-hash-v1' def stage_checkpoint(mount, destination, manifest_path): mount, destination, manifest_path = map(Path, (mount, destination, manifest_path)) manifest_bytes = manifest_path.read_bytes() manifest = json.loads(manifest_bytes) destination.mkdir(parents=True, exist_ok=True) total = 0 for item in manifest['files']: name = item['name'] assert not Path(name).is_absolute() and '..' not in Path(name).parts, 'Unsafe checkpoint relative path' (destination/name).parent.mkdir(parents=True,exist_ok=True) source = mount/name assert source.stat().st_size == item['size'], 'Mounted file size differs: '+name print('SERVER_MODEL_STAGE_BEGIN', name, item['size'], flush=True) if item['hash_algorithm'] == 'sha256': digest = hashlib.sha256() else: assert item['hash_algorithm'] == 'git-sha1' digest = hashlib.sha1() digest.update(('blob '+str(item['size'])+'\0').encode()) temporary = destination/(name+'.partial') count = 0; next_progress = 512*1024*1024 try: with source.open('rb') as reader, temporary.open('wb') as writer: while block := reader.read(8*1024*1024): writer.write(block); digest.update(block); count += len(block) if count >= next_progress: print('SERVER_MODEL_STAGE_PROGRESS', name, count, item['size'], flush=True) next_progress += 512*1024*1024 assert count == item['size'], 'Incomplete file copy: '+name assert digest.hexdigest() == item['hash'], 'Content hash differs: '+name temporary.replace(destination/name) except Exception: temporary.unlink(missing_ok=True) raise total += count print('SERVER_MODEL_FILE_VERIFIED', name, count, flush=True) print('SERVER_MODEL_STAGE_COMPLETE', len(manifest['files']), total, flush=True) return dict(protocol=PROTOCOL, manifest_sha256=hashlib.sha256(manifest_bytes).hexdigest(), source_files_revision=manifest['source_files_revision'], verified_files=len(manifest['files']), verified_bytes=total)