|
| 1 | +""" |
| 2 | +Unit tests for ``patch_transformers_mistral_regex``. |
| 3 | +
|
| 4 | +Verifies that our wrapper around |
| 5 | +``transformers.PreTrainedTokenizerBase._patch_mistral_regex`` catches |
| 6 | +exceptions from the unconditional ``huggingface_hub.model_info()`` lookup |
| 7 | +and returns the tokenizer unchanged — matching the success-path behavior |
| 8 | +for non-Mistral repos (transformers 4.57.3, ``tokenization_utils_base.py:2503``). |
| 9 | +
|
| 10 | +NOTE: These tests mutate ``transformers.PreTrainedTokenizerBase`` globally; |
| 11 | +run serially, not under ``pytest-xdist`` with per-worker process isolation. |
| 12 | +""" |
| 13 | + |
| 14 | +import sys |
| 15 | +from pathlib import Path |
| 16 | + |
| 17 | +import pytest |
| 18 | + |
| 19 | +sys.path.insert(0, str(Path(__file__).parent.parent)) |
| 20 | + |
| 21 | +from huggingface_hub.errors import OfflineModeIsEnabled # noqa: E402 |
| 22 | +from transformers.tokenization_utils_base import PreTrainedTokenizerBase # noqa: E402 |
| 23 | + |
| 24 | +import utils.hf_offline_patch as hf_offline_patch # noqa: E402 |
| 25 | + |
| 26 | + |
| 27 | +@pytest.fixture(autouse=True) |
| 28 | +def restore_mistral_regex(): |
| 29 | + """Snapshot the current ``_patch_mistral_regex`` and restore after each test.""" |
| 30 | + saved = PreTrainedTokenizerBase.__dict__.get("_patch_mistral_regex") |
| 31 | + saved_flag = hf_offline_patch._mistral_regex_patched |
| 32 | + try: |
| 33 | + yield |
| 34 | + finally: |
| 35 | + if saved is not None: |
| 36 | + PreTrainedTokenizerBase._patch_mistral_regex = saved |
| 37 | + hf_offline_patch._mistral_regex_patched = saved_flag |
| 38 | + |
| 39 | + |
| 40 | +def _apply_patch(): |
| 41 | + hf_offline_patch._mistral_regex_patched = False |
| 42 | + hf_offline_patch.patch_transformers_mistral_regex() |
| 43 | + |
| 44 | + |
| 45 | +def test_suppresses_offline_mode_is_enabled(monkeypatch): |
| 46 | + _apply_patch() |
| 47 | + |
| 48 | + import huggingface_hub |
| 49 | + |
| 50 | + def raise_offline(*_args, **_kwargs): |
| 51 | + raise OfflineModeIsEnabled("offline") |
| 52 | + |
| 53 | + monkeypatch.setattr(huggingface_hub, "model_info", raise_offline) |
| 54 | + |
| 55 | + sentinel = object() |
| 56 | + result = PreTrainedTokenizerBase._patch_mistral_regex( |
| 57 | + sentinel, "Qwen/Qwen3-TTS-12Hz-1.7B-Base" |
| 58 | + ) |
| 59 | + assert result is sentinel |
| 60 | + |
| 61 | + |
| 62 | +def test_suppresses_connection_errors(monkeypatch): |
| 63 | + _apply_patch() |
| 64 | + |
| 65 | + import huggingface_hub |
| 66 | + |
| 67 | + def raise_connection(*_args, **_kwargs): |
| 68 | + raise ConnectionError("network unreachable") |
| 69 | + |
| 70 | + monkeypatch.setattr(huggingface_hub, "model_info", raise_connection) |
| 71 | + |
| 72 | + sentinel = object() |
| 73 | + result = PreTrainedTokenizerBase._patch_mistral_regex( |
| 74 | + sentinel, "Qwen/Qwen3-TTS-12Hz-1.7B-Base" |
| 75 | + ) |
| 76 | + assert result is sentinel |
| 77 | + |
| 78 | + |
| 79 | +def test_passthrough_on_success(monkeypatch): |
| 80 | + """When model_info returns non-Mistral tags the original falls through and returns the tokenizer unchanged.""" |
| 81 | + _apply_patch() |
| 82 | + |
| 83 | + import huggingface_hub |
| 84 | + |
| 85 | + class FakeInfo: |
| 86 | + tags = ["model-type:qwen", "language:en"] |
| 87 | + |
| 88 | + monkeypatch.setattr(huggingface_hub, "model_info", lambda *_a, **_kw: FakeInfo()) |
| 89 | + |
| 90 | + sentinel = object() |
| 91 | + result = PreTrainedTokenizerBase._patch_mistral_regex( |
| 92 | + sentinel, "Qwen/Qwen3-TTS-12Hz-1.7B-Base" |
| 93 | + ) |
| 94 | + assert result is sentinel |
| 95 | + |
| 96 | + |
| 97 | +def test_idempotent(): |
| 98 | + _apply_patch() |
| 99 | + first = PreTrainedTokenizerBase._patch_mistral_regex |
| 100 | + hf_offline_patch.patch_transformers_mistral_regex() |
| 101 | + second = PreTrainedTokenizerBase._patch_mistral_regex |
| 102 | + assert first.__func__ is second.__func__ |
| 103 | + |
| 104 | + |
| 105 | +def test_missing_method_is_noop(monkeypatch): |
| 106 | + monkeypatch.delattr(PreTrainedTokenizerBase, "_patch_mistral_regex", raising=False) |
| 107 | + hf_offline_patch._mistral_regex_patched = False |
| 108 | + hf_offline_patch.patch_transformers_mistral_regex() |
| 109 | + assert hf_offline_patch._mistral_regex_patched is False |
| 110 | + |
| 111 | + |
| 112 | +if __name__ == "__main__": |
| 113 | + pytest.main([__file__, "-v"]) |
0 commit comments