from __future__ import annotations import json import os import sys import tempfile import types import unittest from pathlib import Path from unittest.mock import Mock, patch from airfrans_frontier.training.hf_upload import HfArtifactUploader, resolve_resume_checkpoint class HuggingFaceUploadTests(unittest.TestCase): def test_uploader_creates_repo_uploads_file_and_writes_manifest(self) -> None: created: list[tuple[str, str, bool]] = [] committed: list[tuple[str, str, tuple[str, ...], str]] = [] class FakeCommitOperationAdd: def __init__(self, *, path_in_repo: str, path_or_fileobj: str) -> None: self.path_in_repo = path_in_repo self.path_or_fileobj = path_or_fileobj class FakeApi: def __init__(self, token: str) -> None: self.token = token def create_repo(self, *, repo_id: str, repo_type: str, private: bool, exist_ok: bool) -> None: created.append((repo_id, repo_type, private)) def create_commit(self, *, repo_id: str, repo_type: str, operations: list[FakeCommitOperationAdd], commit_message: str): committed.append((repo_id, repo_type, tuple(operation.path_in_repo for operation in operations), commit_message)) return types.SimpleNamespace(commit_url="https://huggingface.co/repo/commit/abc", oid="abc") fake_module = types.SimpleNamespace(HfApi=FakeApi, CommitOperationAdd=FakeCommitOperationAdd) with tempfile.TemporaryDirectory() as tmp, patch.dict(sys.modules, {"huggingface_hub": fake_module}), patch.dict(os.environ, {"HF_TOKEN": "token"}): root = Path(tmp) (root / "checkpoint_latest.pt").write_bytes(b"checkpoint") uploader = HfArtifactUploader( enabled=True, run_dir=root, repo_id="owner/repo", repo_type="model", path_in_repo="runs/model/run-1", private=False, ) result = uploader.upload_files(("checkpoint_latest.pt",), commit_message="upload checkpoint") self.assertEqual(result["uploaded"], ["runs/model/run-1/checkpoint_latest.pt"]) self.assertEqual(created, [("owner/repo", "model", False)]) self.assertEqual(committed, [("owner/repo", "model", ("runs/model/run-1/checkpoint_latest.pt",), "upload checkpoint")]) manifest = json.loads((root / "hf_upload_manifest.json").read_text()) self.assertTrue(manifest["enabled"]) self.assertEqual(manifest["repo_id"], "owner/repo") self.assertIn("runs/model/run-1/checkpoint_latest.pt", manifest["uploaded_paths"]) def test_uploader_suppresses_uploads_after_hf_retry_after_limit(self) -> None: class FakeRateLimitError(RuntimeError): def __init__(self) -> None: super().__init__("429 Too Many Requests: Retry after 600 seconds") self.response = types.SimpleNamespace(headers={"Retry-After": "600"}) class FakeCommitOperationAdd: def __init__(self, *, path_in_repo: str, path_or_fileobj: str) -> None: self.path_in_repo = path_in_repo self.path_or_fileobj = path_or_fileobj fake_module = types.SimpleNamespace(CommitOperationAdd=FakeCommitOperationAdd) with tempfile.TemporaryDirectory() as tmp, patch.dict(sys.modules, {"huggingface_hub": fake_module}): root = Path(tmp) (root / "metrics.jsonl").write_text("{}\n") uploader = HfArtifactUploader( enabled=True, run_dir=root, repo_id="owner/repo", repo_type="model", path_in_repo="runs/model/run-1", private=False, max_rate_limit_sleep_seconds=0, ) fake_api = types.SimpleNamespace(create_commit=Mock(side_effect=FakeRateLimitError())) uploader._api = fake_api with self.assertRaises(FakeRateLimitError): uploader.upload_files(("metrics.jsonl",), commit_message="first") suppressed = uploader.upload_files(("metrics.jsonl",), commit_message="second") self.assertTrue(suppressed["rate_limited"]) self.assertEqual(fake_api.create_commit.call_count, 1) manifest = json.loads((root / "hf_upload_manifest.json").read_text()) self.assertGreater(manifest["rate_limit_until"], 0) self.assertEqual(manifest["rate_limit_retry_after_seconds"], 600.0) self.assertEqual(len(manifest["suppressed_uploads"]), 2) def test_resolve_resume_checkpoint_downloads_hf_uri(self) -> None: calls: list[tuple[str, str]] = [] def fake_download(*, repo_id: str, repo_type: str, filename: str, token: str, local_dir: str) -> str: calls.append((repo_id, filename)) path = Path(local_dir) / filename path.parent.mkdir(parents=True, exist_ok=True) path.write_bytes(b"checkpoint") return str(path) fake_module = types.SimpleNamespace(hf_hub_download=fake_download) with patch.dict(sys.modules, {"huggingface_hub": fake_module}), patch.dict(os.environ, {"HF_TOKEN": "token"}): path, info = resolve_resume_checkpoint("hf://owner/repo/runs/model/checkpoint_latest.pt") self.assertIsNotNone(path) assert path is not None self.assertTrue(path.is_file()) self.assertEqual(calls, [("owner/repo", "runs/model/checkpoint_latest.pt")]) self.assertTrue(info["resume_downloaded"]) self.assertEqual(info["resume_source"], "hf://owner/repo/runs/model/checkpoint_latest.pt") if __name__ == "__main__": unittest.main()