opensysone / source /tests /test_publish_publication.py
andyshu's picture
Organize verified OpenSysOne publication payload
294f8ea verified
Raw History Blame
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()