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
264
scripts/stream_model.py
Normal file
264
scripts/stream_model.py
Normal 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)
|
||||
Loading…
Add table
Add a link
Reference in a new issue