File size: 5,410 Bytes
490f3fe
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
from __future__ import annotations

import ast
import importlib.util
import json
from pathlib import Path

import pytest


ROOT = Path(__file__).resolve().parents[1]
SOURCE_ROOT = ROOT / "src" / "tmcra_benchmark"
REPOSITORY_ROOT = ROOT.parents[1]
SECURITY_SOURCE_ROOTS = (SOURCE_ROOT, REPOSITORY_ROOT / "code")


def _dotted_name(node: ast.AST) -> str:
    if isinstance(node, ast.Name):
        return node.id
    if isinstance(node, ast.Attribute):
        prefix = _dotted_name(node.value)
        return f"{prefix}.{node.attr}" if prefix else node.attr
    return ""


def _keyword(call: ast.Call, name: str) -> ast.AST | None:
    return next((item.value for item in call.keywords if item.arg == name), None)


def test_every_torch_load_is_weights_only() -> None:
    violations: list[str] = []
    for source_root in SECURITY_SOURCE_ROOTS:
        for path in source_root.rglob("*.py"):
            tree = ast.parse(path.read_text(encoding="utf-8"), filename=str(path))
            torch_aliases = {
                alias.asname or alias.name
                for node in ast.walk(tree)
                if isinstance(node, ast.Import)
                for alias in node.names
                if alias.name == "torch"
            }
            torch_aliases.update({"torch", "torch_module"})
            for node in ast.walk(tree):
                if not isinstance(node, ast.Call) or not isinstance(node.func, ast.Attribute):
                    continue
                if node.func.attr != "load":
                    continue
                receiver = _dotted_name(node.func.value)
                if receiver not in torch_aliases and not receiver.endswith(".torch"):
                    continue
                value = _keyword(node, "weights_only")
                if not isinstance(value, ast.Constant) or value.value is not True:
                    violations.append(f"{path.relative_to(REPOSITORY_ROOT)}:{node.lineno}")
    assert not violations, "torch.load must use weights_only=True: " + ", ".join(violations)


def test_local_transformer_models_disable_remote_code() -> None:
    violations: list[str] = []
    for path in SOURCE_ROOT.rglob("*.py"):
        tree = ast.parse(path.read_text(encoding="utf-8"), filename=str(path))
        for node in ast.walk(tree):
            if not isinstance(node, ast.Call) or not isinstance(node.func, ast.Attribute):
                continue
            if node.func.attr != "from_pretrained":
                continue
            owner = _dotted_name(node.func.value)
            if not (owner.startswith("AutoTokenizer") or owner.startswith("AutoModel")):
                continue
            local_only = _keyword(node, "local_files_only")
            remote_code = _keyword(node, "trust_remote_code")
            if not isinstance(local_only, ast.Constant) or local_only.value is not True:
                violations.append(f"{path.relative_to(ROOT)}:{node.lineno}:local_files_only")
            if not isinstance(remote_code, ast.Constant) or remote_code.value is not False:
                violations.append(f"{path.relative_to(ROOT)}:{node.lineno}:trust_remote_code")
    assert not violations, "local model loading is not hardened: " + ", ".join(violations)


def test_public_code_never_enables_remote_code() -> None:
    violations: list[str] = []
    for source_root in SECURITY_SOURCE_ROOTS:
        for path in source_root.rglob("*.py"):
            tree = ast.parse(path.read_text(encoding="utf-8"), filename=str(path))
            for node in ast.walk(tree):
                if not isinstance(node, ast.Call) or node.func is None:
                    continue
                remote_code = _keyword(node, "trust_remote_code")
                if isinstance(remote_code, ast.Constant) and remote_code.value is True:
                    violations.append(f"{path.relative_to(REPOSITORY_ROOT)}:{node.lineno}")
    assert not violations, "public code enables trust_remote_code: " + ", ".join(violations)


def _load_model_assets_module():
    path = SOURCE_ROOT / "legacy" / "tmcra_model_assets.py"
    spec = importlib.util.spec_from_file_location(path.stem, path)
    assert spec and spec.loader
    module = importlib.util.module_from_spec(spec)
    spec.loader.exec_module(module)
    return module


def test_official_hf_model_does_not_require_tmcra_manifest(tmp_path: Path) -> None:
    module = _load_model_assets_module()
    (tmp_path / "config.json").write_text("{}\n", encoding="utf-8")

    assert module.load_optional_hf_model_manifest(tmp_path) == {}

    for relative in ("tmcra_v3_reranker.py", "tmcra_v3_online_runtime.py"):
        source = (SOURCE_ROOT / "legacy" / relative).read_text(encoding="utf-8")
        assert "load_optional_hf_model_manifest" in source
        assert "pinned cross model manifest is required" not in source
        assert "pinned model manifest is required" not in source


def test_optional_tmcra_manifest_is_validated(tmp_path: Path) -> None:
    module = _load_model_assets_module()
    manifest = tmp_path / "TMCRA_MODEL_MANIFEST.json"
    manifest.write_text(json.dumps({"repo_id": "BAAI/bge-reranker-v2-m3", "revision": "abc"}), encoding="utf-8")
    assert module.load_optional_hf_model_manifest(tmp_path)["revision"] == "abc"

    manifest.write_text("[]", encoding="utf-8")
    with pytest.raises(ValueError, match="JSON object"):
        module.load_optional_hf_model_manifest(tmp_path)