import copy import hashlib import inspect import io import json from pathlib import Path import tarfile import tempfile from types import SimpleNamespace import unittest from unittest.mock import patch from scripts import publish_publication as publication def sha(content): return hashlib.sha256(content).hexdigest() def git_blob(content): return hashlib.sha1(f'blob {len(content)}\0'.encode() + content).hexdigest() class FakeHub: def __init__(self, root): self.root = root self.sha = 'a' * 40 self.private = True self.files = {'README.md': b'---\nlicense: unknown\n---\nOld card\n', '.gitattributes': b'*.pt filter=lfs\n', 'source/obsolete.py': b'# old source\n'} for name in publication.PROTECTED_POINTERS: self.files[name] = b'{"pointer":"unchanged"}\n' for prefix in publication.PROTECTED_PREFIXES: self.files[prefix + 'existing.json'] = b'{"historical":true}\n' self.revisions = {self.sha: copy.deepcopy(self.files)} self.commits = [] self.downloads = 0 self.corrupt_payload = self.change_protected = self.concurrent = False def repo_info(self, *args, **kwargs): return SimpleNamespace(sha=self.sha, private=self.private) def list_repo_tree(self, repo_id, *, revision, **kwargs): from huggingface_hub.hf_api import RepoFile records = [] for name, content in self.revisions[revision].items(): oid = git_blob(content) if self.corrupt_payload and self.commits and name == 'docs/guide.md': oid = '0' * 40 records.append(RepoFile(path=name, size=len(content), oid=oid, lfs={'size': len(content), 'oid': sha(content), 'pointerSize': 128} if name.endswith('.pt') else None)) return records def create_commit(self, repo_id, operations, *, parent_commit, **kwargs): from huggingface_hub import CommitOperationDelete assert parent_commit == self.sha operations = list(operations) for operation in operations: if isinstance(operation, CommitOperationDelete): assert operation.is_folder is False del self.files[operation.path_in_repo] else: value = operation.path_or_fileobj self.files[operation.path_in_repo] = value if isinstance(value, bytes) else Path(value).read_bytes() self.commits.append([operation.path_in_repo for operation in operations]) if self.change_protected: self.files['final/existing.json'] = b'changed' self.sha = ('b' if len(self.commits) == 1 else 'c') * 40 self.revisions[self.sha] = copy.deepcopy(self.files) result = SimpleNamespace(oid=self.sha) if self.concurrent: self.sha = 'd' * 40 self.revisions[self.sha] = copy.deepcopy(self.files) return result def hf_hub_download(self, repo_id, filename, *, revision, **kwargs): assert not filename.endswith(('.pt', '.tar.gz')) self.downloads += 1 path = self.root / f'download-{self.downloads}' path.write_bytes(self.revisions[revision][filename]) return str(path) class PublicationTests(unittest.TestCase): def setUp(self): temp = tempfile.TemporaryDirectory() self.addCleanup(temp.cleanup) self.root = Path(temp.name) self.api = FakeHub(self.root) inventory = self.root / 'remote-inventory.json' self.original = copy.deepcopy(self.api.files) records = publication.remote_inventory(self.api, 'andyshu/opensysone', self.api.sha) inventory.write_text(json.dumps({'repo_id': 'andyshu/opensysone', 'revision': self.api.sha, 'private': True, 'files': [{'path': n, **r} for n, r in records.items()]})) self.plan = {'repo_id': 'andyshu/opensysone', 'expected_revision': self.api.sha, 'expected_private': True, 'source_commit': 'e' * 40, 'inventory_path': str(inventory), 'files': {}, 'delete_source_files': ['source/obsolete.py']} contents = {'README.md': b'---\nlicense: unknown\n---\nNew publication card\n', 'publication-manifest.json': b'{"publication":"fixture"}\n', 'model/model.pt': b'fixture adapter/head only', 'docs/guide.md': b'Guide\n', 'source/app.py': b'# source\n'} for name, content in contents.items(): path = self.root / 'stage' / name path.parent.mkdir(parents=True, exist_ok=True) path.write_bytes(content) self.plan['files'][name] = {'source': str(path), 'sha256': sha(content)} mocked_hash = patch.object(publication, 'MODEL_SHA256', sha(contents['model/model.pt'])) mocked_hash.start() self.addCleanup(mocked_hash.stop) def publish(self, progress=lambda stage, **values: None): return publication.publish(self.api, self.plan, progress, roots=(self.root,)) def test_two_commits_match_installed_hub_signatures_and_preserve_history(self): from huggingface_hub import HfApi for name in ('repo_info', 'list_repo_tree', 'create_commit', 'hf_hub_download'): original = getattr(self.api, name) signature = inspect.signature(getattr(HfApi, name)) def checked(*args, _original=original, _signature=signature, **kwargs): _signature.bind(self.api, *args, **kwargs) return _original(*args, **kwargs) setattr(self.api, name, checked) result = self.publish() self.assertEqual(result['payload_commit'], 'b' * 40) self.assertEqual(result['publication_commit'], 'c' * 40) self.assertNotIn('README.md', self.api.commits[0]) self.assertNotIn('PUBLICATION.json', self.api.commits[0]) self.assertEqual(self.api.commits[1], ['README.md', 'PUBLICATION.json']) self.assertNotIn('source/obsolete.py', self.api.files) for name, content in self.original.items(): if name not in {'README.md', 'source/obsolete.py'}: self.assertEqual(self.api.files[name], content) pointer = json.loads(self.api.files['PUBLICATION.json']) self.assertEqual(pointer['manifest_sha256'], sha(self.api.files['publication-manifest.json'])) self.assertTrue(self.api.private) def test_production_progress_signature_writes_phases_without_keyword_collision(self): state, phases = {}, [] def progress(stage, **values): state.update(stage=stage, heartbeat_utc=publication.utc(), **values) phases.append(stage) publication.write_json(self.root / 'state.json', state) self.publish(progress) self.assertEqual(phases, ['uploading_payload', 'verifying_payload', 'publishing_front']) self.assertEqual(publication.read_json(self.root / 'state.json')['stage'], 'publishing_front') def test_corrupt_payload_never_publishes_front(self): self.api.corrupt_payload = True with self.assertRaisesRegex(ValueError, 'checksum mismatch'): self.publish() self.assertEqual(len(self.api.commits), 1) self.assertEqual(self.api.files['README.md'], self.original['README.md']) self.assertNotIn('PUBLICATION.json', self.api.files) def test_changed_protected_file_never_publishes_front(self): self.api.change_protected = True with self.assertRaisesRegex(ValueError, 'Preserved remote file'): self.publish() self.assertEqual(len(self.api.commits), 1) self.assertNotIn('PUBLICATION.json', self.api.files) def test_protected_overwrites_and_deletes_block_before_upload(self): original = copy.deepcopy(self.plan) for name in ['final/existing.json', 'snapshots/new.json', 'publications/new.json', *publication.PROTECTED_POINTERS]: with self.subTest(overwrite=name): self.plan = copy.deepcopy(original) self.plan['files'][name] = self.plan['files']['publication-manifest.json'] with self.assertRaises(ValueError): self.publish() for name in ['final/existing.json', 'CURRENT_SNAPSHOT.json', 'source/', 'source/missing.py', 'source/app.py']: with self.subTest(delete=name): self.plan = copy.deepcopy(original) self.plan['delete_source_files'] = [name] with self.assertRaises(ValueError): self.publish() self.assertEqual(self.api.commits, []) def test_revision_private_and_license_guards(self): with self.subTest('initial revision'): self.api.sha = 'f' * 40 with self.assertRaisesRegex(ValueError, 'revision or visibility'): self.publish() self.api.sha = self.plan['expected_revision'] with self.subTest('privacy'): self.api.private = False with self.assertRaisesRegex(ValueError, 'revision or visibility'): self.publish() self.api.private = True with self.subTest('license'): card = Path(self.plan['files']['README.md']['source']) card.write_text('---\nlicense: apache-2.0\n---\nChanged license\n') self.plan['files']['README.md']['sha256'] = publication.digest(card) with self.assertRaisesRegex(ValueError, 'license metadata'): self.publish() self.assertEqual(self.api.commits, []) def test_intervening_commit_blocks_front(self): self.api.concurrent = True with self.assertRaisesRegex(ValueError, 'revision or visibility'): self.publish() self.assertEqual(len(self.api.commits), 1) def test_hash_and_exact_calibrated_model_guards(self): model = Path(self.plan['files']['model/model.pt']['source']) model.write_bytes(b'other weight') with self.assertRaisesRegex(ValueError, 'checksum mismatch'): self.publish() self.plan['files']['model/model.pt']['sha256'] = publication.digest(model) with self.assertRaisesRegex(ValueError, 'exact calibrated'): self.publish() self.assertEqual(self.api.commits, []) def test_unsafe_paths_disguised_sources_and_credential_content(self): original = copy.deepcopy(self.plan) for name in ['../README.md', '/README.md', 'docs//guide.md', 'docs/../guide.md', 'source/.env', 'source/model.safetensors']: with self.subTest(path=name): self.plan = copy.deepcopy(original) self.plan['files'][name] = self.plan['files']['docs/guide.md'] with self.assertRaises(ValueError): self.publish() for filename, content in [('credentials.json', b'{}'), ('note.txt', ('hf_' + 'x' * 40).encode())]: self.plan = copy.deepcopy(original) path = self.root / filename path.write_bytes(content) self.plan['files']['docs/safe' + path.suffix] = {'source': str(path), 'sha256': sha(content)} with self.assertRaises(ValueError): self.publish() self.assertEqual(self.api.commits, []) def test_archive_guard_rejects_base_weights(self): archive = self.root / 'source.tar.gz' with tarfile.open(archive, 'w:gz') as bundle: member = tarfile.TarInfo('model.safetensors') member.size = 4 bundle.addfile(member, io.BytesIO(b'base')) self.plan['files']['archive/source.tar.gz'] = {'source': str(archive), 'sha256': publication.digest(archive)} with self.assertRaisesRegex(ValueError, 'base-weight'): self.publish() self.assertEqual(self.api.commits, []) if __name__ == '__main__': unittest.main()