freedom-intelligence/scripts/test_stream_model.py
tegwick 8a65297d56 Stream model reserve to Scaleway without local weight staging
Assistant: codex
Assistant-Model: gpt-6-astra
Assistant-Session: 01a09cbd-43c1-79f3-809e-1ee97b40b64d
2026-09-14 00:17:15 +02:00

190 lines
7.5 KiB
Python

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()