Download source/tests/test_publish_publication.py from andyshu/opensysone: direct link, hf CLI and curl.
- Browser
- Download file 12 kB
-
https://huggingface.co/andyshu/opensysone/resolve/294f8ea1b877ac86188aa88eade4b81f5f190293/source/tests/test_publish_publication.py
- Command line
-
hf download hf://andyshu/opensysone@294f8ea1b877ac86188aa88eade4b81f5f190293/source/tests/test_publish_publication.py
-
curl -L -o test_publish_publication.py https://huggingface.co/andyshu/opensysone/resolve/294f8ea1b877ac86188aa88eade4b81f5f190293/source/tests/test_publish_publication.py
12 kB
| 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() | |