* Run the malware gate on the RAG embedding model before it loads Setting the RAG embedding model through PUT /api/settings/embedding-model persisted an arbitrary repo and later handed it straight to SentenceTransformer, which deserializes pickle weights. Unlike the normal model-load paths, this route never ran evaluate_file_security, and force skipped verification entirely, so a repo Hugging Face flags as unsafe (or any repo under force) could be downloaded and loaded in the backend process without a scan. Run the malware/pickle scan at both ends: the settings endpoint now scans before persisting and returns 409 on a flagged repo even under force (force still only skips the is-embedding-model type check for offline or local repos), and the embedder scans again at the load sink so a name that arrives via env or default is covered too. Local paths and unreachable scans fail open inside evaluate_file_security, and the sink never bricks the embedder on a gate error. * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Thread the load token into the embedding scan and hard-fail on a block The load-sink scan ran without a token, so evaluate_file_security (which passes token=False when none is given) could not reach a gated or private repo and failed open for exactly the model SentenceTransformer would still load. Resolve the loader's own token (HF_TOKEN env or the cached login) and pass it to the sink scan, and fall back to it in the settings endpoint when the request omits one. The sink previously raised a plain RuntimeError, which the llama-server fallback in encode() and _build_st_backend_or_fallback() swallowed as a routine ST failure, silently switching backends instead of blocking. Raise a distinct UnsafeEmbeddingModelError that both fallback paths re-raise, so a flagged model hard-fails. * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Scan sentence-transformers module dirs and scope the embedding pickle gate to the ST backend Extend the RAG embedding malware gate so a poisoned pickle under a SentenceTransformer module dir (for example 0_Transformer/pytorch_model.bin) blocks. Those dirs are read from the repo's modules.json and passed as load roots to evaluate_file_security at both the settings endpoint and the load sink, so such a pickle is treated as root-level there instead of an unreferenced nested shard that was previously allowed. Scope the ST pickle scan to the sentence-transformers backend. On the llama-server backend the embedder loads GGUF files (inert) from the -GGUF companion repo, never the ST repo's pickle, so a custom ST repo with a flagged pickle and a clean GGUF companion is no longer rejected. The existing GGUF availability checks already cover that path. Return 403 for the hard security block instead of 409. The settings UI routes every 409 into the forceable save-anyway flow, but this block cannot be bypassed by force, so it now uses a distinct status the client treats as non-forceable. * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Base the embedding pickle scan on the actual backend, not just the resolver _llama_backend_active only consulted the auto resolver, so on a GPU box where auto resolves to sentence-transformers but the process already fell back to the llama-server backend at runtime (a torch or CUDA load/encode failure), it returned False and the settings endpoint hard-blocked a save whose ST pickle is flagged even though the process loads only inert GGUF. Add active_backend_is_llama, which reflects the actual built backend (True when the cached backend is a LlamaServerBackend, including a runtime fallback) and otherwise defers to the resolver as a fresh process would, and delegate _llama_backend_active to it. * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Report the cached embedding backend verbatim, not the resolver active_backend_is_llama() fell through to the config resolver whenever a backend was already built but was not llama-server, so a live sentence-transformers backend could report llama=True once the resolver picked llama (GPU heuristic or a runtime config change) and wrongly skip its pickle scan. Once a backend exists, return isinstance(backend, LlamaServerBackend) directly; only defer to the resolver before any backend is built. --------- Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
114 lines
4.9 KiB
Python
114 lines
4.9 KiB
Python
# SPDX-License-Identifier: AGPL-3.0-only
|
|
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
|
|
|
|
"""Deterministic consistency guards for the model-load security gate.
|
|
|
|
The gate spans many parallel sites (validate/load/status, the inference/training/export
|
|
workers, the preflight route); past regressions were a fix at one site with a sibling
|
|
left behind. These guards enumerate the sites mechanically (AST + source) so a new site
|
|
that drops the token or mis-reports the requirement fails here, not in a later review.
|
|
"""
|
|
|
|
import ast
|
|
from pathlib import Path
|
|
|
|
_BACKEND = Path(__file__).resolve().parent.parent
|
|
|
|
# Probes read Hub config to classify a model; a token-less call 404s on a gated repo.
|
|
# Scan callers under routes/ and core/ (probe definitions live in utils/).
|
|
_PROBE_FUNCS = {"is_vision_model", "is_embedding_model", "detect_audio_type"}
|
|
_PROBE_CALLER_ROOTS = ("routes", "core")
|
|
|
|
|
|
def _iter_caller_files():
|
|
for root in _PROBE_CALLER_ROOTS:
|
|
yield from (_BACKEND / root).rglob("*.py")
|
|
|
|
|
|
def _passes_token(call: ast.Call) -> bool:
|
|
"""True if the call passes an hf_token (keyword, or the 2nd positional slot)."""
|
|
if any(kw.arg in ("hf_token", "token") for kw in call.keywords if kw.arg is not None):
|
|
return True
|
|
return len(call.args) >= 2
|
|
|
|
|
|
def _call_name(call: ast.Call):
|
|
fn = call.func
|
|
return fn.id if isinstance(fn, ast.Name) else getattr(fn, "attr", None)
|
|
|
|
|
|
def test_capability_probes_thread_the_hf_token():
|
|
"""Every capability-probe caller passes the token; a token-less probe misclassifies
|
|
a gated model (the /check-vision regression)."""
|
|
offenders = []
|
|
for path in _iter_caller_files():
|
|
try:
|
|
tree = ast.parse(path.read_text())
|
|
except SyntaxError:
|
|
continue
|
|
for node in ast.walk(tree):
|
|
if isinstance(node, ast.Call) and _call_name(node) in _PROBE_FUNCS:
|
|
if not _passes_token(node):
|
|
rel = path.relative_to(_BACKEND)
|
|
offenders.append(f"{rel}:{node.lineno} {_call_name(node)}() drops the hf_token")
|
|
assert not offenders, (
|
|
"A capability probe must pass the hf_token so gated/private models classify "
|
|
"correctly:\n " + "\n ".join(offenders)
|
|
)
|
|
|
|
|
|
def test_gguf_trust_remote_code_reported_inert_not_from_yaml():
|
|
"""GGUF never executes auto_map, so requires_trust_remote_code is reported via the
|
|
resolver or False, never the raw YAML bool() (the round-6 regression)."""
|
|
src = (_BACKEND / "routes" / "inference.py").read_text()
|
|
assert "requires_trust_remote_code = bool(" not in src, (
|
|
"Report requires_trust_remote_code via _resolve_loaded_trust_remote_code "
|
|
"(non-GGUF) or set it False (GGUF); never bool(inference_config.get(...))."
|
|
)
|
|
|
|
|
|
def test_capability_detection_caches_are_token_aware():
|
|
"""Every capability cache is keyed by (model, token_fingerprint) so an unauthenticated
|
|
miss cannot poison a later authenticated lookup (the audio-cache regression)."""
|
|
src = (_BACKEND / "utils" / "models" / "model_config.py").read_text()
|
|
offenders = []
|
|
for line in src.splitlines():
|
|
stripped = line.strip()
|
|
if "_detection_cache:" in stripped and stripped.endswith("= {}"):
|
|
if "Dict[Tuple" not in stripped and "Dict[tuple" not in stripped:
|
|
offenders.append(stripped)
|
|
assert not offenders, (
|
|
"A capability cache must be keyed by (model, token_fingerprint), not the bare "
|
|
"model name:\n " + "\n ".join(offenders)
|
|
)
|
|
|
|
|
|
def test_malware_and_consent_gates_cover_the_lora_base():
|
|
"""Every worker that runs a load gate also resolves the LoRA base, so a poisoned or
|
|
custom-code base is never skipped."""
|
|
gated_workers = [
|
|
"core/inference/worker.py",
|
|
"core/export/worker.py",
|
|
"core/training/worker.py",
|
|
]
|
|
offenders = []
|
|
for rel in gated_workers:
|
|
src = (_BACKEND / rel).read_text()
|
|
runs_gate = "evaluate_file_security(" in src or "evaluate_remote_code_consent" in src
|
|
resolves_base = "get_base_model_from_lora_identifier(" in src or "base_model" in src
|
|
if runs_gate and not resolves_base:
|
|
offenders.append(f"{rel} runs a load gate but never resolves the LoRA base")
|
|
assert not offenders, "\n".join(offenders)
|
|
|
|
|
|
def test_rag_embedding_path_runs_the_malware_gate():
|
|
"""The RAG embedding model is set through /settings and later loaded by
|
|
SentenceTransformer, which deserializes pickles; both sites must run the malware gate
|
|
or a flagged repo loads unscanned (bypassing the normal model-load protections)."""
|
|
offenders = []
|
|
for rel in ("routes/settings.py", "core/rag/embeddings.py"):
|
|
if "evaluate_file_security(" not in (_BACKEND / rel).read_text():
|
|
offenders.append(
|
|
f"{rel} loads/persists an embedding model without evaluate_file_security"
|
|
)
|
|
assert not offenders, "\n".join(offenders)
|