Stream model reserve to Scaleway without local weight staging
Assistant: codex Assistant-Model: gpt-6-astra Assistant-Session: 01a09cbd-43c1-79f3-809e-1ee97b40b64d
This commit is contained in:
parent
ad88d16a52
commit
8a65297d56
16 changed files with 805 additions and 145 deletions
190
scripts/test_stream_model.py
Normal file
190
scripts/test_stream_model.py
Normal file
|
|
@ -0,0 +1,190 @@
|
|||
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()
|
||||
Loading…
Add table
Add a link
Reference in a new issue