diff --git a/changelog.d/warn-hf-no-token.changed.md b/changelog.d/warn-hf-no-token.changed.md new file mode 100644 index 00000000..bc20e829 --- /dev/null +++ b/changelog.d/warn-hf-no-token.changed.md @@ -0,0 +1 @@ +Warn when a Hugging Face repo that requires authentication (private or gated) is downloaded with no HUGGING_FACE_TOKEN available, so that a 401 from a missing token, or a 403 from a cached token that is not approved for a gated repo, is easy to trace. diff --git a/policyengine_core/tools/hugging_face.py b/policyengine_core/tools/hugging_face.py index 0f8a114f..a47615fc 100644 --- a/policyengine_core/tools/hugging_face.py +++ b/policyengine_core/tools/hugging_face.py @@ -73,6 +73,16 @@ def download_huggingface_dataset( Returns: Path to the downloaded local file. + + Warns: + UserWarning: If the repo requires authentication but no + HUGGING_FACE_TOKEN was available. The download still runs with + token=None, so huggingface_hub applies its own cached token + (HF_TOKEN, or the file written by `hf auth login`, which is + `huggingface-cli login` before huggingface_hub 0.34) if it has + one, unless HF_HUB_DISABLE_IMPLICIT_TOKEN is set; the warning + explains a 401 that follows when there is no such token, or a + 403 when a gated repo has not approved it. """ # 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) @@ -104,6 +114,25 @@ def download_huggingface_dataset( authentication_token: str | None = None if requires_authentication: authentication_token = get_or_prompt_hf_token() + if authentication_token is None: + # Deliberately not an error: huggingface_hub resolves its own + # cached token when token=None, so `hf auth login` users still + # work. Warn so that the bare 401 huggingface_hub raises when + # that fallback is empty too can be traced back here (#529). + warnings.warn( + f"Hugging Face repo '{repo}' requires authentication, but no " + "HUGGING_FACE_TOKEN was available (the environment variable " + "is unset or empty, and no token was entered at a prompt). " + "huggingface_hub normally falls back to its own cached token " + "if it has one (the HF_TOKEN environment variable, or the file " + "written by `hf auth login`, which is `huggingface-cli login` " + "before huggingface_hub 0.34). If the download that follows " + "fails with RepositoryNotFoundError or GatedRepoError (a 401 " + "when no token was sent, or a 403 when a gated repo has not " + "approved the token), set HUGGING_FACE_TOKEN to a token whose " + "account has access.", + stacklevel=2, + ) return hf_hub_download( repo_id=repo, diff --git a/tests/core/tools/test_hugging_face.py b/tests/core/tools/test_hugging_face.py index d9905adf..9c127d3d 100644 --- a/tests/core/tools/test_hugging_face.py +++ b/tests/core/tools/test_hugging_face.py @@ -1,4 +1,6 @@ +import itertools import os +import warnings import pytest from unittest.mock import patch, MagicMock from huggingface_hub import ModelInfo @@ -126,9 +128,12 @@ def test_download_private_repo_no_token(self, environ): "Test error", response=mock_response ) - result = download_huggingface_dataset( - test_repo, test_filename, test_version, test_dir - ) + with pytest.warns( + UserWarning, match="no HUGGING_FACE_TOKEN" + ): + result = download_huggingface_dataset( + test_repo, test_filename, test_version, test_dir + ) assert result is mock_download.return_value mock_getpass.assert_not_called() @@ -233,9 +238,12 @@ def test_download_gated_repo_non_interactive_without_token(self): id=test_repo, private=False, gated="manual" ) - download_huggingface_dataset( - test_repo, test_filename, test_version, test_dir - ) + with pytest.warns( + UserWarning, match="no HUGGING_FACE_TOKEN" + ): + download_huggingface_dataset( + test_repo, test_filename, test_version, test_dir + ) mock_getpass.assert_not_called() mock_download.assert_called_once_with( @@ -264,6 +272,10 @@ def test_download_public_ungated_repo_never_prompts(self, gated): with patch( "policyengine_core.tools.hugging_face.getpass" ) as mock_getpass: + # A string, so that a widened predicate fails on + # assert_not_called() below rather than on storing a + # MagicMock in os.environ. + mock_getpass.return_value = "prompted_token" with patch( "policyengine_core.tools.hugging_face.hf_hub_download" ) as mock_download: @@ -406,3 +418,324 @@ def test_deep_subdirectory(self): def test_invalid_url_too_short(self): with pytest.raises(ValueError, match="Invalid hf:// URL format"): parse_hf_url("hf://owner/repo") + + +class TestNoTokenWarning: + """download_huggingface_dataset warns when it passes token=None for a + repo that needs authentication. + + Core deliberately does not raise or prompt in that case (#422): with + token=None, huggingface_hub falls back to its own cached token (for + example HF_TOKEN or the `hf auth login` file) and raises its own 401 if + that is missing too. The warning is what makes that 401 traceable to a + missing HUGGING_FACE_TOKEN (#529). + """ + + repo = "test_owner/test_repo" + filename = "test_filename" + version = "test_version" + local_dir = "test_dir" + + def _download(self): + return download_huggingface_dataset( + self.repo, self.filename, self.version, self.local_dir + ) + + def _assert_downloaded_with(self, mock_download, token): + mock_download.assert_called_once_with( + repo_id=self.repo, + repo_type="model", + filename=self.filename, + revision=self.version, + token=token, + local_dir=self.local_dir, + ) + + @staticmethod + def _lookup_response(lookup): + """Configure model_info for the given repo visibility. + + "public": the repo is public, so no token is ever needed. + "private-flag": model_info answers with private=True. + "gated": model_info answers with private=False, gated="manual", as + policyengine/policyengine-uk-data-private does. + "not-found": model_info raises RepositoryNotFoundError, which core + treats as "probably private". + """ + if lookup == "public": + return {"return_value": ModelInfo(id="test_repo", private=False)} + if lookup == "private-flag": + return {"return_value": ModelInfo(id="test_repo", private=True)} + if lookup == "gated": + return { + "return_value": ModelInfo(id="test_repo", private=False, gated="manual") + } + assert lookup == "not-found" + mock_response = MagicMock() + mock_response.status_code = 404 + mock_response.headers = {} + return { + "side_effect": RepositoryNotFoundError("Test error", response=mock_response) + } + + @pytest.mark.parametrize("lookup", ["private-flag", "gated", "not-found"]) + @pytest.mark.parametrize( + "environ", + [{}, {"HUGGING_FACE_TOKEN": ""}, {"HF_TOKEN": "hf_cached_token"}], + ids=["token-unset", "token-empty", "hf-token-only"], + ) + def test_warns_when_no_token_resolved_non_interactively(self, lookup, environ): + """No HUGGING_FACE_TOKEN, no TTY: warn, then pass token=None through. + + The hf-token-only case pins that the warning still fires when only + huggingface_hub's own HF_TOKEN is set: core resolved nothing, and + the warning itself says the fallback will be used if present. + """ + model_info_config = self._lookup_response(lookup) + + with patch.dict(os.environ, environ, clear=True): + with patch("os.isatty", return_value=False): + with patch( + "policyengine_core.tools.hugging_face.getpass" + ) as mock_getpass: + mock_getpass.return_value = "prompted_token" + with patch( + "policyengine_core.tools.hugging_face.hf_hub_download" + ) as mock_download: + with patch( + "policyengine_core.tools.hugging_face.model_info", + **model_info_config, + ): + with pytest.warns( + UserWarning, match="no HUGGING_FACE_TOKEN" + ) as record: + result = self._download() + + # Behaviour is unchanged: no prompt, no raise, token=None passed on. + assert result is mock_download.return_value + mock_getpass.assert_not_called() + self._assert_downloaded_with(mock_download, token=None) + + # Exactly one warning, naming the repo, the fallback (including the + # pre-0.34 login command), and both errors that can follow with + # their 401 and 403 causes. + assert len(record) == 1 + message = str(record[0].message) + assert self.repo in message + assert "HF_TOKEN" in message + assert "hf auth login" in message + assert "huggingface-cli login" in message + assert "RepositoryNotFoundError" in message + assert "GatedRepoError" in message + assert "401" in message + assert "403" in message + # stacklevel=2: the warning points at the caller, not at core. + assert record[0].filename == __file__ + + def test_warns_when_interactive_prompt_left_empty(self): + """TTY present but the user enters nothing: same warning, token=None.""" + with patch.dict(os.environ, {}, clear=True): + with patch("os.isatty", return_value=True): + with patch( + "policyengine_core.tools.hugging_face.getpass", + return_value="", + ) 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", + **self._lookup_response("private-flag"), + ): + with pytest.warns( + UserWarning, match="no HUGGING_FACE_TOKEN" + ): + self._download() + + mock_getpass.assert_called_once() + self._assert_downloaded_with(mock_download, token=None) + + @pytest.mark.parametrize( + ("lookup", "environ", "isatty", "prompted", "expected_token"), + [ + pytest.param("public", {}, False, None, None, id="public-repo"), + pytest.param( + "public", + {"HUGGING_FACE_TOKEN": "env_token"}, + False, + None, + None, + id="public-repo-ignores-env-token", + ), + pytest.param( + "private-flag", + {"HUGGING_FACE_TOKEN": "env_token"}, + False, + None, + "env_token", + id="private-flag-env-token", + ), + pytest.param( + "gated", + {"HUGGING_FACE_TOKEN": "env_token"}, + False, + None, + "env_token", + id="gated-env-token", + ), + pytest.param( + "not-found", + {"HUGGING_FACE_TOKEN": "env_token"}, + False, + None, + "env_token", + id="not-found-env-token", + ), + pytest.param( + "private-flag", + {}, + True, + "prompted_token", + "prompted_token", + id="private-flag-prompted-token", + ), + ], + ) + def test_no_warning_when_a_token_is_passed_or_not_needed( + self, lookup, environ, isatty, prompted, expected_token + ): + """Public repos pass token=None without warning; a resolved token + never warns. Guards against warning on every public download.""" + with patch.dict(os.environ, environ, clear=True): + with patch("os.isatty", return_value=isatty): + with patch( + "policyengine_core.tools.hugging_face.getpass", + return_value=prompted, + ): + with patch( + "policyengine_core.tools.hugging_face.hf_hub_download" + ) as mock_download: + with patch( + "policyengine_core.tools.hugging_face.model_info", + **self._lookup_response(lookup), + ): + with warnings.catch_warnings(): + warnings.simplefilter("error", UserWarning) + self._download() + + self._assert_downloaded_with(mock_download, token=expected_token) + + +class TestTokenRoutingInvariants: + """Exhaustive check of download_huggingface_dataset's token contract. + + Every combination of repo state, environment, TTY and prompt entry is run + through the real function (with model_info, hf_hub_download and getpass + mocked) and compared with the spec written out in _expected(): + + - Only private, gated or not-found repos require authentication. + - Such a repo gets HUGGING_FACE_TOKEN if it is non-empty; otherwise a + prompt on a TTY; otherwise None. A public, ungated repo always gets + None and never prompts. + - The token passed on is None or a non-empty string, never "". + - Exactly one no-token warning fires when a repo requiring + authentication ends up with None, and none fires otherwise. + """ + + repo = "test_owner/test_repo" + filename = "test_filename" + + REPO_STATES = { + "ungated": dict(private=False, gated=False), + "gated-none": dict(private=False, gated=None), + "fields-missing": dict(), + "gated-auto": dict(private=False, gated="auto"), + "gated-manual": dict(private=False, gated="manual"), + "private": dict(private=True, gated=False), + "private-gated": dict(private=True, gated="manual"), + "not-found": None, + } + ENVIRONS = { + "token-unset": {}, + "token-empty": {"HUGGING_FACE_TOKEN": ""}, + "hf-token-only": {"HF_TOKEN": "hf_cached_token"}, + "token-set": {"HUGGING_FACE_TOKEN": "env_token"}, + } + # itertools.product, not nested fors: only a comprehension's first + # iterable can see class attributes. + CASES = [ + pytest.param(state, environ, isatty, entry, id=f"{state}-{environ}-{tty}-{e}") + for state, environ, (isatty, tty), (entry, e) in itertools.product( + REPO_STATES, + ENVIRONS, + [(False, "no-tty"), (True, "tty")], + [("", "empty-entry"), ("prompted_token", "entry")], + ) + ] + + @staticmethod + def _expected(state, environ, isatty, entry): + requires_authentication = state in ( + "gated-auto", + "gated-manual", + "private", + "private-gated", + "not-found", + ) + if not requires_authentication: + return None, False, 0 + env_token = environ.get("HUGGING_FACE_TOKEN") or None + if env_token is not None: + return env_token, False, 0 + if isatty: + token = entry or None + return token, True, int(token is None) + return None, False, 1 + + @pytest.mark.parametrize(("state", "environ", "isatty", "entry"), CASES) + def test_token_routing_matches_spec(self, state, environ, isatty, entry): + fields = self.REPO_STATES[state] + if fields is None: + mock_response = MagicMock() + mock_response.status_code = 404 + mock_response.headers = {} + model_info_config = { + "side_effect": RepositoryNotFoundError( + "Test error", response=mock_response + ) + } + else: + model_info_config = {"return_value": ModelInfo(id="test_repo", **fields)} + env = self.ENVIRONS[environ] + + with patch.dict(os.environ, env, clear=True): + with patch("os.isatty", return_value=isatty): + with patch( + "policyengine_core.tools.hugging_face.getpass", + return_value=entry, + ) 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", + **model_info_config, + ): + with warnings.catch_warnings(record=True) as caught: + warnings.simplefilter("always") + result = download_huggingface_dataset( + self.repo, self.filename + ) + + expected_token, expected_prompt, expected_warnings = self._expected( + state, env, isatty, entry + ) + assert result is mock_download.return_value + mock_download.assert_called_once() + token = mock_download.call_args.kwargs["token"] + assert token == expected_token + assert token is None or (isinstance(token, str) and token != "") + assert mock_getpass.call_count == int(expected_prompt) + user_warnings = [w for w in caught if issubclass(w.category, UserWarning)] + assert len(user_warnings) == expected_warnings + assert all("no HUGGING_FACE_TOKEN" in str(w.message) for w in user_warnings)