Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions changelog.d/fix-hf-gated-repo-token.fixed.md
Original file line number Diff line number Diff line change
@@ -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.
24 changes: 20 additions & 4 deletions policyengine_core/tools/hugging_face.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand All @@ -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
Expand All @@ -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(
Expand Down
147 changes: 147 additions & 0 deletions tests/core/tools/test_hugging_face.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down
Loading