diff --git a/changelog.d/fix-hf-gated-repo-token.fixed.md b/changelog.d/fix-hf-gated-repo-token.fixed.md new file mode 100644 index 00000000..488fbd72 --- /dev/null +++ b/changelog.d/fix-hf-gated-repo-token.fixed.md @@ -0,0 +1 @@ +Send the HUGGING_FACE_TOKEN when downloading from public but gated Hugging Face repos, not only private ones, so gated dataset downloads no longer fail with a 401 gated-repo error. diff --git a/policyengine_core/tools/hugging_face.py b/policyengine_core/tools/hugging_face.py index b48dd647..0f8a114f 100644 --- a/policyengine_core/tools/hugging_face.py +++ b/policyengine_core/tools/hugging_face.py @@ -53,6 +53,14 @@ def download_huggingface_dataset( """ Download a PolicyEngine dataset file from the Hugging Face Hub. + Private repos and public-but-gated repos both require authentication, + so for either the token is resolved by get_or_prompt_hf_token() + (HUGGING_FACE_TOKEN, else an interactive prompt when stdin is a TTY, + else None) and passed explicitly to hf_hub_download. Public, ungated + repos are requested with token=None and never prompt; huggingface_hub + may still apply its own implicitly configured token (HF_TOKEN or the + login file) in that case. + Args: repo: Hugging Face model repository identifier in ``owner/name`` format. @@ -68,15 +76,23 @@ def download_huggingface_dataset( """ # Attempt connection to Hugging Face model_info endpoint # (https://huggingface.co/docs/huggingface_hub/v0.26.5/en/package_reference/hf_api#huggingface_hub.HfApi.model_info) - # Attempt to fetch model info to determine if repo is private + # Attempt to fetch model info to determine if the repo requires + # authentication. Testing `private` alone is not enough: a public but + # gated repo still answers model_info() without a token, but the file + # download returns 401 unless a gate-approved token is sent. + # ModelInfo.gated is False for ungated repos, "auto" or "manual" for + # gated repos, and None if the field is absent, so a truthiness test + # covers every state (both gate modes need an authenticated request). # A RepositoryNotFoundError & 401 likely means the repo is private, # but this error will also surface for public repos with malformed URL, etc. try: fetched_model_info: ModelInfo = model_info(repo) - is_repo_private = bool(fetched_model_info.private) + requires_authentication: bool = bool(fetched_model_info.private) or bool( + fetched_model_info.gated + ) except RepositoryNotFoundError as e: # If this error type arises, it's likely the repo is private; see docs above - is_repo_private = True + requires_authentication = True pass except Exception as e: # Otherwise, there probably is just a download error @@ -86,7 +102,7 @@ def download_huggingface_dataset( ) authentication_token: str | None = None - if is_repo_private: + if requires_authentication: authentication_token = get_or_prompt_hf_token() return hf_hub_download( diff --git a/tests/core/tools/test_hugging_face.py b/tests/core/tools/test_hugging_face.py index 8740eb40..d9905adf 100644 --- a/tests/core/tools/test_hugging_face.py +++ b/tests/core/tools/test_hugging_face.py @@ -141,6 +141,153 @@ def test_download_private_repo_no_token(self, environ): local_dir=test_dir, ) + @pytest.mark.parametrize("gated", ["manual", "auto"]) + def test_download_gated_public_repo_passes_env_token(self, gated): + """A public but gated repo must receive HUGGING_FACE_TOKEN. + + Regression test for PolicyEngine/policyengine-core#529: `private` + is False for a gated repo, but the file download still needs a + gate-approved token or the Hub answers 401. + """ + test_repo = "test_repo" + test_filename = "test_filename" + test_version = "test_version" + test_dir = "test_dir" + test_token = "gated_repo_test_token" + + with patch.dict(os.environ, {"HUGGING_FACE_TOKEN": test_token}, clear=True): + with patch( + "policyengine_core.tools.hugging_face.hf_hub_download" + ) as mock_download: + with patch( + "policyengine_core.tools.hugging_face.model_info" + ) as mock_model_info: + mock_model_info.return_value = ModelInfo( + id=test_repo, private=False, gated=gated + ) + + download_huggingface_dataset( + test_repo, test_filename, test_version, test_dir + ) + + mock_download.assert_called_once_with( + repo_id=test_repo, + repo_type="model", + filename=test_filename, + revision=test_version, + token=test_token, + local_dir=test_dir, + ) + + def test_download_private_flag_repo_passes_env_token(self): + """A repo reported as private=True by model_info still gets the token.""" + test_repo = "test_repo" + test_filename = "test_filename" + test_version = "test_version" + test_dir = "test_dir" + test_token = "private_repo_test_token" + + with patch.dict(os.environ, {"HUGGING_FACE_TOKEN": test_token}, clear=True): + with patch( + "policyengine_core.tools.hugging_face.hf_hub_download" + ) as mock_download: + with patch( + "policyengine_core.tools.hugging_face.model_info" + ) as mock_model_info: + mock_model_info.return_value = ModelInfo( + id=test_repo, private=True, gated=False + ) + + download_huggingface_dataset( + test_repo, test_filename, test_version, test_dir + ) + + mock_download.assert_called_once_with( + repo_id=test_repo, + repo_type="model", + filename=test_filename, + revision=test_version, + token=test_token, + local_dir=test_dir, + ) + + def test_download_gated_repo_non_interactive_without_token(self): + """Gated repo in CI without secrets: no prompt, token=None is passed.""" + test_repo = "test_repo" + test_filename = "test_filename" + test_version = "test_version" + test_dir = "test_dir" + + with patch.dict(os.environ, {}, clear=True): + with patch("os.isatty", return_value=False): + with patch( + "policyengine_core.tools.hugging_face.getpass" + ) as mock_getpass: + with patch( + "policyengine_core.tools.hugging_face.hf_hub_download" + ) as mock_download: + with patch( + "policyengine_core.tools.hugging_face.model_info" + ) as mock_model_info: + mock_model_info.return_value = ModelInfo( + id=test_repo, private=False, gated="manual" + ) + + download_huggingface_dataset( + test_repo, test_filename, test_version, test_dir + ) + + mock_getpass.assert_not_called() + mock_download.assert_called_once_with( + repo_id=test_repo, + repo_type="model", + filename=test_filename, + revision=test_version, + token=None, + local_dir=test_dir, + ) + + @pytest.mark.parametrize("gated", [False, None]) + def test_download_public_ungated_repo_never_prompts(self, gated): + """Public, ungated repo: no token lookup and no interactive prompt. + + Guards against widening the predicate so far that every public + download on a developer machine asks for a token. + """ + test_repo = "test_repo" + test_filename = "test_filename" + test_version = "test_version" + test_dir = "test_dir" + + with patch.dict(os.environ, {}, clear=True): + with patch("os.isatty", return_value=True): + with patch( + "policyengine_core.tools.hugging_face.getpass" + ) as mock_getpass: + with patch( + "policyengine_core.tools.hugging_face.hf_hub_download" + ) as mock_download: + with patch( + "policyengine_core.tools.hugging_face.model_info" + ) as mock_model_info: + mock_model_info.return_value = ModelInfo( + id=test_repo, private=False, gated=gated + ) + + download_huggingface_dataset( + test_repo, test_filename, test_version, test_dir + ) + + mock_getpass.assert_not_called() + mock_download.assert_called_once_with( + repo_id=test_repo, + repo_type="model", + filename=test_filename, + revision=test_version, + token=None, + local_dir=test_dir, + ) + class TestGetOrPromptHfToken: def test_get_token_from_environment(self):