from __future__ import annotations import json import os import sys import tempfile import types import unittest from pathlib import Path from unittest.mock import 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]] = [] uploaded: list[tuple[str, str, str]] = [] 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 upload_file(self, *, repo_id: str, repo_type: str, path_or_fileobj: str, path_in_repo: str, commit_message: str): uploaded.append((repo_id, repo_type, path_in_repo)) return types.SimpleNamespace(commit_url="https://huggingface.co/repo/commit/abc", oid="abc") fake_module = types.SimpleNamespace(HfApi=FakeApi) 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(uploaded, [("owner/repo", "model", "runs/model/run-1/checkpoint_latest.pt")]) 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_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()