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:
tegwick 2026-09-14 00:17:15 +02:00
parent ad88d16a52
commit 8a65297d56
16 changed files with 805 additions and 145 deletions

View file

@ -1,27 +1,15 @@
#!/usr/bin/env python3
"""Collect one HF model into local staging, optionally upload to Scaleway.
"""Collect an HF model: diskless S3 streaming or explicitly opted-in local cache.
Usage:
python3 scripts/collect_model.py \\
--repo-id Qwen/Qwen3-8B --local-name Qwen__Qwen3-8B --revision main
python3 scripts/collect_model.py \\
--repo-id deepseek-ai/DeepSeek-V4-Flash-0731 \\
--local-name deepseek-ai__DeepSeek-V4-Flash-0731 \\
--tier s \\
Remote reserve (default when --s3-bucket is supplied; no local weight files):
python3 scripts/collect_model.py \
--repo-id deepseek-ai/DeepSeek-V4-Flash-0731 \
--local-name deepseek-ai__DeepSeek-V4-Flash-0731 --tier s \
--s3-bucket railiance-fi-open-weight-reserve
Layout (local staging, still required for huggingface_hub):
{base}/models/{local_name}/{revision}/ # R / W
{base}/strategic/{local_name}/{revision}/ # S
{base}/staging/...
S3 prefix (SoT once bucket is live):
s3://{bucket}/models|strategic/{local_name}/{revision}/
s3://{bucket}/manifests/{local_name}/{revision}/MANIFEST.json
Credentials: AWS_ACCESS_KEY_ID + AWS_SECRET_ACCESS_KEY, or SCW_ACCESS_KEY +
SCW_SECRET_KEY. Never commit keys.
Use --dry-run to inspect the immutable source revision and transfer size.
Use --local-download without --s3-bucket only for an intentional local cache.
See docs/streaming-reserve.md for the resource-limited remote worker.
"""
from __future__ import annotations
@ -36,8 +24,6 @@ from pathlib import Path
# Prefer sequential downloads on WSL+drvfs (xet parallel can OOM)
os.environ.setdefault("HF_HUB_DISABLE_XET", "1")
from huggingface_hub import snapshot_download # noqa: E402
DEFAULT_ALLOW = (
"*.safetensors",
"*.bin",
@ -112,8 +98,22 @@ def main() -> int:
default=os.environ.get("FI_S3_REGION", "nl-ams"),
)
ap.add_argument("--token", default=os.environ.get("HF_TOKEN") or None)
ap.add_argument("--dry-run", action="store_true", help="list pinned S3 transfer without uploading")
ap.add_argument("--part-mib", type=int, default=64, help="streaming S3 part buffer (5128 MiB)")
ap.add_argument("--local-download", action="store_true", help="explicitly permit a full local snapshot (no S3)")
args = ap.parse_args()
if args.s3_bucket:
if args.local_download:
ap.error("--local-download cannot be combined with --s3-bucket")
from stream_model import collect
return collect(args, DEFAULT_ALLOW, IGNORE)
if not args.local_download:
ap.error("choose --s3-bucket (diskless streaming) or explicitly --local-download")
if args.dry_run:
ap.error("--dry-run currently requires --s3-bucket")
from huggingface_hub import snapshot_download
base = Path(args.base)
dest_kind = "strategic" if args.tier == "s" else "models"
stage = base / "staging" / args.local_name / args.revision
@ -179,97 +179,9 @@ def main() -> int:
(final / "MANIFEST.json").write_text(json.dumps(manifest, indent=2) + "\n")
result = {"ok": True, "total_bytes": total, "path": str(final), "tier": args.tier}
if args.s3_bucket:
prefix = f"{dest_kind}/{args.local_name}/{args.revision}"
storage_class = "GLACIER" if args.tier == "s" else "ONEZONE_IA"
uploaded = upload_tree(
final,
bucket=args.s3_bucket,
prefix=prefix,
endpoint=args.s3_endpoint,
region=args.s3_region,
storage_class=storage_class,
)
man_key = f"manifests/{args.local_name}/{args.revision}/MANIFEST.json"
upload_file(
final / "MANIFEST.json",
bucket=args.s3_bucket,
key=man_key,
endpoint=args.s3_endpoint,
region=args.s3_region,
storage_class="STANDARD",
)
result["s3"] = {
"bucket": args.s3_bucket,
"prefix": f"s3://{args.s3_bucket}/{prefix}/",
"manifest": f"s3://{args.s3_bucket}/{man_key}",
"storage_class": storage_class,
"objects": uploaded,
}
print(json.dumps(result, indent=2))
return 0
def _s3_client(endpoint: str, region: str):
try:
import boto3
except ImportError as exc:
raise SystemExit("boto3 required for --s3-bucket (pip install boto3)") from exc
access = os.environ.get("AWS_ACCESS_KEY_ID") or os.environ.get("SCW_ACCESS_KEY")
secret = os.environ.get("AWS_SECRET_ACCESS_KEY") or os.environ.get("SCW_SECRET_KEY")
if not access or not secret:
raise SystemExit("set AWS_ACCESS_KEY_ID/AWS_SECRET_ACCESS_KEY or SCW_ACCESS_KEY/SCW_SECRET_KEY")
return boto3.client(
"s3",
region_name=region,
endpoint_url=endpoint,
aws_access_key_id=access,
aws_secret_access_key=secret,
)
def upload_file(
path: Path,
*,
bucket: str,
key: str,
endpoint: str,
region: str,
storage_class: str,
) -> None:
client = _s3_client(endpoint, region)
extra = {"StorageClass": storage_class} if storage_class else {}
client.upload_file(str(path), bucket, key, ExtraArgs=extra)
def upload_tree(
root: Path,
*,
bucket: str,
prefix: str,
endpoint: str,
region: str,
storage_class: str,
) -> int:
n = 0
for p in sorted(root.rglob("*")):
if not p.is_file() or p.name.startswith("."):
continue
rel = p.relative_to(root).as_posix()
key = f"{prefix.rstrip('/')}/{rel}"
upload_file(
p,
bucket=bucket,
key=key,
endpoint=endpoint,
region=region,
storage_class=storage_class,
)
n += 1
print(f"uploaded {key}", file=sys.stderr)
return n
if __name__ == "__main__":
sys.exit(main())

View file

@ -0,0 +1,25 @@
[Unit]
Description=Freedom Intelligence diskless model transfer to Scaleway
StartLimitIntervalSec=3600
StartLimitBurst=3
[Service]
Type=exec
WorkingDirectory=%h/.local/share/fi-reserve
Environment=PYTHONDONTWRITEBYTECODE=1
Environment=PYTHONUNBUFFERED=1
Environment=HF_HUB_DISABLE_TELEMETRY=1
Environment=FI_S3_CREDENTIAL_FILE=%t/fi-reserve/s3.json
ExecStart=/usr/bin/flock -n %t/fi-reserve/collector.lock %h/.local/share/fi-reserve/venv/bin/python %h/.local/share/fi-reserve/scripts/collect_model.py --repo-id deepseek-ai/DeepSeek-V4-Flash-0731 --local-name deepseek-ai__DeepSeek-V4-Flash-0731 --revision 7872f01b1d1fe23eabc4c98b48bffcef5a386062 --tier s --s3-bucket railiance-fi-open-weight-reserve
Restart=on-failure
RestartSec=60
TimeoutStopSec=210
MemoryHigh=384M
MemoryMax=512M
MemorySwapMax=0
CPUQuota=50%
TasksMax=32
UMask=0077
NoNewPrivileges=true
# Deliberately not enabled at boot: runtime credentials must be supplied first.

View file

@ -0,0 +1,66 @@
#!/usr/bin/env python3
"""Send only S3 keys from local OpenBao to remote runtime storage over SSH.
Never forwards the OpenBao token. Runtime credentials disappear at reboot.
The existing bootstrap fallback requires --allow-bootstrap explicitly.
"""
import argparse
import json
import os
from pathlib import Path
import shlex
import subprocess
import urllib.error
import urllib.request
def main():
ap = argparse.ArgumentParser(description=__doc__)
ap.add_argument('--host', default='railiance01')
ap.add_argument('--allow-bootstrap', action='store_true')
args = ap.parse_args()
token = os.environ.get('OPENBAO_TOKEN') or os.environ.get('VAULT_TOKEN')
if not token:
token = Path.home().joinpath('.vault-token').read_text().strip()
addr = os.environ.get('BAO_ADDR', 'https://bao.coulomb.social')
paths = ['workloads/railiance/freedom-intelligence/object-storage']
if args.allow_bootstrap:
paths.append('workloads/railiance/scaleway/bootstrap')
credentials = None
for path in paths:
request = urllib.request.Request(f'{addr}/v1/platform/data/{path}',
headers={'X-Vault-Token': token})
try:
with urllib.request.urlopen(request, timeout=20) as response:
data = json.load(response)['data']['data']
except urllib.error.HTTPError as exc:
if exc.code == 404:
continue
raise RuntimeError(f'OpenBao read failed: HTTP {exc.code}') from None
access = data.get('ACCESS_KEY') or data.get('access_key') or data.get('AWS_ACCESS_KEY_ID')
secret = data.get('SECRET_KEY') or data.get('secret_key') or data.get('AWS_SECRET_ACCESS_KEY')
if not access or not secret:
raise RuntimeError('OpenBao S3 secret has unsupported fields')
credentials = {'ACCESS_KEY': access, 'SECRET_KEY': secret}
print(f'credential source: platform/{path}')
break
if credentials is None:
raise RuntimeError('no S3 credential found')
program = '''import json, os, pathlib, sys
data = json.load(sys.stdin)
root = pathlib.Path('/run/user') / str(os.getuid()) / 'fi-reserve'
root.mkdir(mode=0o700, exist_ok=True)
root.chmod(0o700)
fd = os.open(root / 's3.json', os.O_WRONLY | os.O_CREAT | os.O_TRUNC | os.O_NOFOLLOW, 0o600)
os.fchmod(fd, 0o600)
with os.fdopen(fd, 'w') as f:
json.dump(data, f)
print('S3 credentials installed in remote runtime directory (0600)')
'''
subprocess.run(['ssh', '-o', 'BatchMode=yes', args.host,
'python3 -c ' + shlex.quote(program)],
input=json.dumps(credentials), text=True, check=True)
if __name__ == '__main__':
main()

View file

@ -0,0 +1,3 @@
boto3==1.42.49
huggingface-hub==0.34.4
requests==2.32.5

View file

@ -0,0 +1,24 @@
"""Read-only progress for the pinned V4-Flash collection; no credential output."""
import json
from stream_model import client, optional_json
s3 = client('https://s3.nl-ams.scw.cloud', 'nl-ams')
bucket = 'railiance-fi-open-weight-reserve'
identity = 'deepseek-ai__DeepSeek-V4-Flash-0731/7872f01b1d1fe23eabc4c98b48bffcef5a386062'
prefix = f'strategic/{identity}/'
objects = []
for page in s3.get_paginator('list_objects_v2').paginate(Bucket=bucket, Prefix=prefix):
objects.extend(page.get('Contents', []))
manifest = optional_json(s3, bucket, f'manifests/{identity}/MANIFEST.json')
active = []
for page in s3.get_paginator('list_multipart_uploads').paginate(Bucket=bucket, Prefix=prefix):
for upload in page.get('Uploads', []):
size = 0
for parts in s3.get_paginator('list_parts').paginate(
Bucket=bucket, Key=upload['Key'], UploadId=upload['UploadId']):
size += sum(p['Size'] for p in parts.get('Parts', []))
active.append({'key': upload['Key'], 'uploaded_part_bytes': size})
print(json.dumps({'completed_objects': len(objects), 'object_bytes': sum(o['Size'] for o in objects),
'complete_manifest': bool(manifest and manifest.get('complete')),
'active_uploads': active}, indent=2))

View file

@ -0,0 +1,50 @@
"""Small live S3 checks; deletes only the smoke object versions created here."""
import hashlib
import io
import os
import uuid
from botocore.exceptions import ClientError
from stream_model import client, stream_object
def main():
s3 = client('https://s3.nl-ams.scw.cloud', 'nl-ams')
bucket = os.environ.get('FI_S3_BUCKET', 'railiance-fi-open-weight-reserve')
prefix = f'staging/stream-smoke/{uuid.uuid4()}'
data = b'fi-stream-smoke\n' * 450000 # ~6.4 MiB, two parts
spec = dict(path='payload', bytes=len(data), source_algorithm='sha256',
source_digest=hashlib.sha256(data).hexdigest())
for storage_class in ('STANDARD', 'ONEZONE_IA', 'GLACIER'):
key = f'{prefix}/{storage_class}'
receipt = stream_object(s3, bucket, key, io.BytesIO(data), spec,
storage_class, 5 * 1024 * 1024, {'fi-smoke': 'true'})
try:
head = s3.head_object(Bucket=bucket, Key=key)
assert head.get('StorageClass', 'STANDARD') == storage_class
if storage_class != 'GLACIER':
response = s3.get_object(Bucket=bucket, Key=key)
with response['Body'] as body:
assert hashlib.sha256(body.read()).hexdigest() == spec['source_digest']
print(f'PASS {storage_class}: multipart + HEAD' +
(' + full readback SHA256' if storage_class != 'GLACIER' else ' (no restore)'))
finally:
s3.delete_object(Bucket=bucket, Key=key, VersionId=receipt['version_id'])
key = f'{prefix}/bad-md5'
upload = s3.create_multipart_upload(Bucket=bucket, Key=key)['UploadId']
try:
try:
result = s3.upload_part(Bucket=bucket, Key=key, UploadId=upload, PartNumber=1,
Body=b'payload', ContentMD5='AAAAAAAAAAAAAAAAAAAAAA==')
except ClientError as exc:
assert exc.response['Error']['Code'] == 'BadDigest', exc.response['Error']['Code']
print('PASS server rejects corrupted multipart Content-MD5')
else:
assert result['ETag'].strip('"') == hashlib.md5(b'payload').hexdigest()
print('NOTE server ignores Content-MD5; returned part ETag matches actual bytes. Collector checks ETags explicitly.')
finally:
s3.abort_multipart_upload(Bucket=bucket, Key=key, UploadId=upload)
if __name__ == '__main__':
main()

264
scripts/stream_model.py Normal file
View file

@ -0,0 +1,264 @@
"""Diskless HF → S3 transfer. One file/part at a time; durable per-file receipts.
Run through collect_model.py. A single worker must own each destination prefix.
No HF download/cache functions are used. Interrupted files restart from byte zero;
completed files resume from receipts after checking the remote object identity.
"""
from __future__ import annotations
import base64
from datetime import datetime, timezone
import fnmatch
import hashlib
import json
import math
import os
import re
import signal
import sys
import time
from pathlib import Path
def md5(data):
return base64.b64encode(hashlib.md5(data).digest()).decode()
def client(endpoint, region):
import boto3
from botocore.config import Config
credentials = {}
if os.environ.get("FI_S3_CREDENTIAL_FILE"):
credentials = json.loads(Path(os.environ["FI_S3_CREDENTIAL_FILE"]).read_text())
access = (credentials.get("ACCESS_KEY") or os.environ.get("AWS_ACCESS_KEY_ID")
or os.environ.get("SCW_ACCESS_KEY"))
secret = (credentials.get("SECRET_KEY") or os.environ.get("AWS_SECRET_ACCESS_KEY")
or os.environ.get("SCW_SECRET_KEY"))
if not access or not secret:
raise RuntimeError("S3 credentials missing")
return boto3.client(
"s3", endpoint_url=endpoint, region_name=region,
aws_access_key_id=access, aws_secret_access_key=secret,
aws_session_token=credentials.get("SESSION_TOKEN") or os.environ.get("AWS_SESSION_TOKEN"),
config=Config(connect_timeout=20, read_timeout=180,
retries={"max_attempts": 5, "mode": "standard"},
request_checksum_calculation="when_required",
response_checksum_validation="when_required"),
)
def put_json(s3, bucket, key, value):
data = (json.dumps(value, indent=2) + "\n").encode()
result = s3.put_object(Bucket=bucket, Key=key, Body=data, ContentMD5=md5(data),
ContentType="application/json", StorageClass="STANDARD")
if result["ETag"].strip('"') != hashlib.md5(data).hexdigest():
raise ValueError(f"metadata object checksum mismatch: {key}")
if optional_json(s3, bucket, key) != value:
raise ValueError(f"metadata object readback mismatch: {key}")
def optional_json(s3, bucket, key):
from botocore.exceptions import ClientError
try:
response = s3.get_object(Bucket=bucket, Key=key)
except ClientError as exc:
if exc.response["Error"]["Code"] in {"NoSuchKey", "404"}:
return None
raise
with response["Body"] as body:
return json.load(body)
def source_file(sibling):
lfs = sibling.lfs
if sibling.size is None:
raise ValueError(f"no source size for {sibling.rfilename}")
digest = lfs.sha256 if lfs else sibling.blob_id
algorithm = "sha256" if lfs else "git-sha1"
if not digest or not re.fullmatch(r"[0-9a-f]{64}" if lfs else r"[0-9a-f]{40}", digest):
raise ValueError(f"no usable upstream digest for {sibling.rfilename}")
return {"path": sibling.rfilename, "bytes": sibling.size,
"source_digest": digest, "source_algorithm": algorithm}
def stream_object(s3, bucket, key, body, spec, storage_class, part_size, metadata):
"""Verify source hash and returned part/composite MD5 ETags before receipt.
Scaleway's live API accepted incorrect Content-MD5 (2026-09-14), so sending
that header alone is insufficient. Fail closed if ETags are not MD5-shaped.
"""
if math.ceil(spec["bytes"] / part_size) > 10000:
raise ValueError("file exceeds 10,000 parts; increase --part-mib")
sha = hashlib.sha256()
git_sha = hashlib.sha1(f'blob {spec["bytes"]}\0'.encode())
upload_id = None
total = 0
parts = []
part_digests = []
try:
if spec["bytes"]:
upload_id = s3.create_multipart_upload(
Bucket=bucket, Key=key, StorageClass=storage_class, Metadata=metadata,
)["UploadId"]
while True:
# read(n) from urllib3's HTTPResponse fills n bytes except at EOF.
block = body.read(part_size)
if not block:
break
total += len(block)
if total > spec["bytes"]:
raise ValueError("source exceeded declared size")
sha.update(block)
git_sha.update(block)
result = s3.upload_part(
Bucket=bucket, Key=key, UploadId=upload_id,
PartNumber=len(parts) + 1, Body=block, ContentMD5=md5(block),
)
part_digest = hashlib.md5(block).digest()
if result['ETag'].strip('"') != part_digest.hex():
raise ValueError(f"destination part checksum mismatch: {key}")
part_digests.append(part_digest)
parts.append({"PartNumber": len(parts) + 1, "ETag": result["ETag"]})
print(f"part {key} {total}/{spec['bytes']}", flush=True)
del block
actual = sha.hexdigest() if spec["source_algorithm"] == "sha256" else git_sha.hexdigest()
if total != spec["bytes"] or actual != spec["source_digest"]:
raise ValueError(f"source size/checksum mismatch: {key}")
if upload_id:
expected_etag = hashlib.md5(b"".join(part_digests)).hexdigest() + f"-{len(parts)}"
result = s3.complete_multipart_upload(
Bucket=bucket, Key=key, UploadId=upload_id, MultipartUpload={"Parts": parts},
)
upload_id = None
else:
expected_etag = hashlib.md5(b"").hexdigest()
result = s3.put_object(Bucket=bucket, Key=key, Body=b"", ContentMD5=md5(b""),
StorageClass=storage_class, Metadata=metadata)
head = s3.head_object(Bucket=bucket, Key=key)
if (head["ContentLength"] != total or head["Metadata"] != metadata
or head["ETag"] != result["ETag"]
or head["ETag"].strip('"') != expected_etag):
raise ValueError(f"destination identity mismatch: {key}")
return dict(spec, sha256=sha.hexdigest(), key=key, etag=head["ETag"],
version_id=head.get("VersionId"), storage_class=storage_class,
verification="upstream-digest+part-md5-etags+composite-etag+head; no restore readback")
finally:
if upload_id:
s3.abort_multipart_upload(Bucket=bucket, Key=key, UploadId=upload_id)
def receipt_matches(s3, bucket, key, receipt, spec, metadata, storage_class):
from botocore.exceptions import ClientError
if not receipt or any(receipt.get(k) != v for k, v in spec.items()):
return False
if receipt.get("key") != key or receipt.get("storage_class") != storage_class:
return False
if not re.fullmatch(r"[0-9a-f]{64}", receipt.get("sha256", "")):
return False
try:
head = s3.head_object(Bucket=bucket, Key=key)
except ClientError as exc:
if exc.response["Error"]["Code"] in {"NoSuchKey", "404", "NotFound"}:
return False
raise
return (head["ContentLength"] == spec["bytes"]
and head["ETag"] == receipt.get("etag")
and head.get("VersionId") == receipt.get("version_id")
and head["Metadata"] == metadata
and head.get("StorageClass", "STANDARD") == storage_class)
def abort_orphans(s3, bucket, prefix):
# Only this collector's immutable model prefix; never bucket-wide cleanup.
paginator = s3.get_paginator("list_multipart_uploads")
for page in paginator.paginate(Bucket=bucket, Prefix=prefix + "/"):
for upload in page.get("Uploads", []):
s3.abort_multipart_upload(Bucket=bucket, Key=upload["Key"], UploadId=upload["UploadId"])
def collect(args, allow, ignore):
import requests
from huggingface_hub import HfApi, hf_hub_url
if not 5 <= args.part_mib <= 128:
raise ValueError("--part-mib must be between 5 and 128")
if not re.fullmatch(r"[A-Za-z0-9_.-]+", args.local_name) or args.local_name in {".", ".."}:
raise ValueError("--local-name must be one safe path component")
info = HfApi(token=args.token).model_info(args.repo_id, revision=args.revision, files_metadata=True)
if not re.fullmatch(r"[0-9a-f]{40}", info.sha):
raise ValueError("source did not resolve to an immutable commit")
files = [source_file(f) for f in info.siblings
if any(fnmatch.fnmatchcase(f.rfilename, p) for p in allow)
and not any(fnmatch.fnmatchcase(f.rfilename, p) for p in ignore)]
files.sort(key=lambda f: f["path"])
if not files:
raise ValueError("no matching source files")
part_size = args.part_mib * 1024 * 1024
if any(math.ceil(f["bytes"] / part_size) > 10000 for f in files):
raise ValueError("file exceeds 10,000 parts; increase --part-mib")
kind = "strategic" if args.tier == "s" else "models"
prefix = f"{kind}/{args.local_name}/{info.sha}"
manifest_prefix = f"manifests/{args.local_name}/{info.sha}"
plan = {"repo_id": args.repo_id, "revision": info.sha, "requested_revision": args.revision,
"storage_path": f"s3://{args.s3_bucket}/{prefix}/", "local_name": args.local_name,
"total_bytes": sum(f["bytes"] for f in files), "files": len(files),
"part_mib": args.part_mib, "local_weight_bytes": 0}
print(json.dumps(plan), flush=True)
if args.dry_run:
return 0
s3 = client(args.s3_endpoint, args.s3_region)
s3.head_bucket(Bucket=args.s3_bucket)
abort_orphans(s3, args.s3_bucket, prefix)
storage_class = "GLACIER" if args.tier == "s" else "ONEZONE_IA"
artifacts = []
with requests.Session() as session:
for spec in files:
key = f"{prefix}/{spec['path']}"
receipt_key = f"{manifest_prefix}/receipts/{hashlib.sha256(spec['path'].encode()).hexdigest()}.json"
metadata = {"fi-revision": info.sha, "fi-source-digest": spec["source_digest"],
"fi-source-algorithm": spec["source_algorithm"]}
receipt = optional_json(s3, args.s3_bucket, receipt_key)
if receipt_matches(s3, args.s3_bucket, key, receipt, spec, metadata, storage_class):
print(f"resume verified object {key}", flush=True)
artifacts.append(receipt)
continue
for attempt in range(3):
try:
headers = {"Accept-Encoding": "identity"}
if args.token:
headers["Authorization"] = f"Bearer {args.token}"
# requests strips Authorization on cross-host CDN redirects.
with session.get(hf_hub_url(args.repo_id, spec["path"], revision=info.sha),
headers=headers, stream=True, timeout=(20, 180)) as response:
response.raise_for_status()
if response.status_code != 200:
raise ValueError("expected a full source response")
if response.headers.get("Content-Encoding", "identity") != "identity":
raise ValueError("unexpected source content encoding")
receipt = stream_object(s3, args.s3_bucket, key, response.raw, spec,
storage_class, part_size, metadata)
put_json(s3, args.s3_bucket, receipt_key, receipt)
artifacts.append(receipt)
break
except Exception as exc:
# Never print signed source URLs or credentials in exception text.
print(f"retry {spec['path']} attempt={attempt + 1} error={type(exc).__name__}",
file=sys.stderr, flush=True)
if attempt == 2:
raise RuntimeError(f"transfer failed: {spec['path']} ({type(exc).__name__})") from None
time.sleep(5 * (attempt + 1))
manifest = dict(plan, artifacts=artifacts, completed_at=datetime.now(timezone.utc).isoformat(),
complete=True, transport="http-stream-to-s3-multipart")
manifest_key = f"{manifest_prefix}/MANIFEST.json"
put_json(s3, args.s3_bucket, manifest_key, manifest)
print(json.dumps({"ok": True, "manifest": f"s3://{args.s3_bucket}/{manifest_key}"}), flush=True)
return 0
def terminate(signum, frame):
raise SystemExit(128 + signum)
# SIGTERM unwinds the current multipart upload; SIGKILL is cleaned on next run.
signal.signal(signal.SIGTERM, terminate)

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