import hashlib import io import unittest from unittest.mock import Mock, patch from types import SimpleNamespace from contextlib import ExitStack from stream_model import collect, md5, receipt_matches, stream_object class MemoryS3: def __init__(self): self.parts = [] self.aborted = False self.completed = False self.metadata = {} def create_multipart_upload(self, **kw): self.metadata = kw['Metadata'] return {'UploadId': 'upload'} def upload_part(self, **kw): assert kw['ContentMD5'] == md5(kw['Body']) self.parts.append(kw['Body']) return {'ETag': hashlib.md5(kw['Body']).hexdigest()} def complete_multipart_upload(self, **kw): self.completed = True return {'ETag': self.etag()} def abort_multipart_upload(self, **kw): self.aborted = True def put_object(self, **kw): self.metadata = kw['Metadata'] return {'ETag': self.etag()} def etag(self): if not self.parts: return hashlib.md5(b'').hexdigest() return hashlib.md5(b''.join(hashlib.md5(p).digest() for p in self.parts)).hexdigest() + f'-{len(self.parts)}' def head_object(self, **kw): return {'ContentLength': sum(map(len, self.parts)), 'ETag': self.etag(), 'Metadata': self.metadata, 'VersionId': 'v1', 'StorageClass': 'GLACIER'} class BoundedReader(io.BytesIO): def read(self, n=-1): assert 0 < n <= 5, 'unbounded read' return super().read(n) class TransferTests(unittest.TestCase): def spec(self, data, algorithm='sha256'): digest = (hashlib.sha256(data).hexdigest() if algorithm == 'sha256' else hashlib.sha1(f'blob {len(data)}\0'.encode() + data).hexdigest()) return dict(path='weights.bin', bytes=len(data), source_algorithm=algorithm, source_digest=digest) def run_transfer(self, data, spec=None, body=None, s3=None): s3 = s3 or MemoryS3() with patch('builtins.print'): receipt = stream_object(s3, 'bucket', 'key', body or BoundedReader(data), spec or self.spec(data), 'GLACIER', 5, {'commit': 'pinned'}) return s3, receipt def test_bounded_multipart_and_digest(self): data = b'abcdefghijkl' s3, receipt = self.run_transfer(data) self.assertEqual(s3.parts, [b'abcde', b'fghij', b'kl']) self.assertEqual(receipt['sha256'], hashlib.sha256(data).hexdigest()) self.assertTrue(s3.completed) self.assertFalse(s3.aborted) def test_git_blob_digest(self): self.run_transfer(b'config', self.spec(b'config', 'git-sha1')) def test_empty_file(self): s3, receipt = self.run_transfer(b'') self.assertEqual(receipt['bytes'], 0) self.assertEqual(s3.parts, []) def test_wrong_digest_never_completes(self): s3 = MemoryS3() with self.assertRaises(ValueError): self.run_transfer(b'wrong', self.spec(b'right'), s3=s3) self.assertTrue(s3.aborted) self.assertFalse(s3.completed) def test_short_source_aborts(self): s3 = MemoryS3() with self.assertRaises(ValueError): self.run_transfer(b'short', self.spec(b'longer'), s3=s3) self.assertTrue(s3.aborted) def test_long_source_aborts(self): s3 = MemoryS3() with self.assertRaises(ValueError): self.run_transfer(b'longer', self.spec(b'short'), s3=s3) self.assertTrue(s3.aborted) def test_interrupt_aborts(self): s3 = MemoryS3() body = Mock() body.read.side_effect = [b'abcde', SystemExit(143)] with self.assertRaises(SystemExit): self.run_transfer(b'abcdefgh', body=body, s3=s3) self.assertTrue(s3.aborted) self.assertFalse(s3.completed) def test_failed_part_aborts(self): s3 = MemoryS3() s3.upload_part = Mock(side_effect=OSError('network')) with self.assertRaises(OSError): self.run_transfer(b'payload', s3=s3) self.assertTrue(s3.aborted) def test_corrupted_destination_part_aborts(self): s3 = MemoryS3() s3.upload_part = Mock(return_value={'ETag': '0' * 32}) with self.assertRaises(ValueError): self.run_transfer(b'payload', s3=s3) self.assertTrue(s3.aborted) self.assertFalse(s3.completed) def test_resume_requires_same_remote_version(self): s3, receipt = self.run_transfer(b'payload') spec = self.spec(b'payload') self.assertTrue(receipt_matches(s3, 'bucket', 'key', receipt, spec, {'commit': 'pinned'}, 'GLACIER')) receipt['version_id'] = 'replaced' self.assertFalse(receipt_matches(s3, 'bucket', 'key', receipt, spec, {'commit': 'pinned'}, 'GLACIER')) def test_part_limit_checked_before_start(self): s3 = MemoryS3() spec = self.spec(b'') spec['bytes'] = 50001 with self.assertRaises(ValueError): self.run_transfer(b'', spec, s3=s3) self.assertEqual(s3.metadata, {}) def test_failed_transfer_cannot_publish_manifest(self): with self.collect_mocks() as mocks: mocks['stream_object'].side_effect = OSError('broken source') with self.assertRaises(RuntimeError): collect(self.args(), ('*',), ()) mocks['put_json'].assert_not_called() def test_success_publishes_complete_manifest_last(self): with self.collect_mocks() as mocks: mocks['stream_object'].return_value = {'sha256': 'a' * 64} self.assertEqual(collect(self.args(), ('*',), ()), 0) calls = mocks['put_json'].call_args_list self.assertEqual(len(calls), 2) self.assertTrue(calls[-1].args[2].endswith('/MANIFEST.json')) self.assertTrue(calls[-1].args[3]['complete']) self.assertEqual(calls[-1].args[3]['revision'], 'a' * 40) def args(self): return SimpleNamespace(part_mib=64, local_name='model', token=None, repo_id='org/model', revision='main', tier='s', s3_bucket='bucket', s3_endpoint='endpoint', s3_region='region', dry_run=False) def collect_mocks(self): from contextlib import contextmanager @contextmanager def setup(): with ExitStack() as stack: api = stack.enter_context(patch('huggingface_hub.HfApi')) api.return_value.model_info.return_value = SimpleNamespace( sha='a' * 40, siblings=[SimpleNamespace(rfilename='README.md', size=0, blob_id=hashlib.sha1(b'blob 0\0').hexdigest(), lfs=None)]) session = stack.enter_context(patch('requests.Session')) response = session.return_value.__enter__.return_value.get.return_value.__enter__.return_value response.status_code = 200 response.headers = {} mocks = {name: stack.enter_context(patch('stream_model.' + name)) for name in ('client', 'abort_orphans', 'optional_json', 'receipt_matches', 'stream_object', 'put_json')} mocks['optional_json'].return_value = None mocks['receipt_matches'].return_value = False stack.enter_context(patch('stream_model.time.sleep')) stack.enter_context(patch('builtins.print')) yield mocks return setup() if __name__ == '__main__': unittest.main()