# SPDX-License-Identifier: AGPL-3.0-only # Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 """Tests for the MTP auto-detection path (llama.cpp #22673). Pins three contracts: name-based detector, user-override detector, and the _already_in_target_state mirror that prevents needless reloads. """ from __future__ import annotations import struct import sys import types as _types from pathlib import Path _BACKEND_DIR = str(Path(__file__).resolve().parent.parent) if _BACKEND_DIR not in sys.path: sys.path.insert(0, _BACKEND_DIR) _loggers_stub = _types.ModuleType("loggers") _loggers_stub.get_logger = lambda name: __import__("logging").getLogger(name) sys.modules.setdefault("loggers", _loggers_stub) _structlog_stub = _types.ModuleType("structlog") _structlog_stub.get_logger = lambda *a, **k: __import__("logging").getLogger("stub") sys.modules.setdefault("structlog", _structlog_stub) _httpx_stub = _types.ModuleType("httpx") for _exc in ( "ConnectError", "TimeoutException", "ReadTimeout", "ReadError", "RemoteProtocolError", "CloseError", ): setattr(_httpx_stub, _exc, type(_exc, (Exception,), {})) _httpx_stub.Timeout = type("T", (), {"__init__": lambda s, *a, **k: None}) _httpx_stub.Client = type( "C", (), { "__init__": lambda s, **kw: None, "__enter__": lambda s: s, "__exit__": lambda s, *a: None, }, ) sys.modules.setdefault("httpx", _httpx_stub) import pytest from core.inference.llama_cpp import ( LlamaCppBackend, _extra_args_set_spec_type, _is_mtp_model_name, ) # Synthetic GGUF helper (mirrors test_gguf_metadata.py). _GGUF_MAGIC = 0x46554747 _VTYPE_STRING = 8 _VTYPE_UINT32 = 4 def _enc_string(s: str) -> bytes: b = s.encode("utf-8") return struct.pack(" bytes: return _enc_string(key) + struct.pack(" bytes: return ( _enc_string(key) + struct.pack(" Path: """Header-only GGUF with arch + optional nextn_predict_layers.""" extra_uint32 = dict(extra_uint32 or {}) body = _enc_kv_string("general.architecture", arch) kv_count = 1 if nextn is not None: body += _enc_kv_uint32(f"{arch}.nextn_predict_layers", nextn) kv_count += 1 for k, v in extra_uint32.items(): body += _enc_kv_uint32(k, v) kv_count += 1 header = struct.pack("0 should match. ("qwen3moe", 2), ("hypothetical_future_arch", 4), ], ) def test_read_gguf_metadata_captures_nextn_predict_layers(tmp_path, arch, nextn): gguf = _write_minimal_gguf( tmp_path / "model.gguf", arch = arch, nextn = nextn, extra_uint32 = {f"{arch}.block_count": 4}, ) backend = LlamaCppBackend() backend._read_gguf_metadata(str(gguf)) assert backend._nextn_predict_layers == nextn def test_read_gguf_metadata_leaves_nextn_unset_for_non_mtp_arch(tmp_path): gguf = _write_minimal_gguf( tmp_path / "model.gguf", arch = "qwen3", nextn = None, extra_uint32 = {"qwen3.block_count": 4}, ) backend = LlamaCppBackend() backend._read_gguf_metadata(str(gguf)) assert backend._nextn_predict_layers is None def test_read_gguf_metadata_zero_nextn_is_falsy(tmp_path): # bool(0) is False, so the spec block short-circuits. gguf = _write_minimal_gguf( tmp_path / "model.gguf", arch = "qwen35", nextn = 0, extra_uint32 = {"qwen35.block_count": 4}, ) backend = LlamaCppBackend() backend._read_gguf_metadata(str(gguf)) assert backend._nextn_predict_layers == 0 assert bool(backend._nextn_predict_layers) is False def test_unload_resets_nextn_predict_layers(): # MTP state from a previous load must not bleed into the next load. backend = LlamaCppBackend() backend._nextn_predict_layers = 1 backend.unload_model() assert backend._nextn_predict_layers is None # llama-server capability probe. def _make_fake_llama_server(path: Path, help_text: str) -> Path: """Bash stub that prints `help_text` on --help.""" path.write_text("#!/usr/bin/env bash\n" f"cat <<'EOF'\n{help_text}\nEOF\n") path.chmod(0o755) return path def _clear_caps_cache(): LlamaCppBackend._capability_cache.clear() def test_probe_server_capabilities_detects_draft_mtp(tmp_path): # Original naming from llama.cpp #22673. fake = _make_fake_llama_server( tmp_path / "llama-server", "--spec-type none,draft-simple,draft-eagle3,draft-mtp," "ngram-simple,ngram-map-k,ngram-map-k4v,ngram-mod,ngram-cache", ) _clear_caps_cache() caps = LlamaCppBackend.probe_server_capabilities(str(fake)) assert caps["found"] is True assert caps["mtp_token"] == "draft-mtp" assert caps["supports_mtp"] is True def test_probe_server_capabilities_detects_renamed_mtp(tmp_path): # Renamed upstream: draft-mtp -> mtp. fake = _make_fake_llama_server( tmp_path / "llama-server", "--spec-type [none|mtp|ngram-cache|ngram-simple|ngram-map-k|" "ngram-map-k4v|ngram-mod]", ) _clear_caps_cache() caps = LlamaCppBackend.probe_server_capabilities(str(fake)) assert caps["mtp_token"] == "mtp" assert caps["supports_mtp"] is True def test_probe_server_capabilities_reports_outdated_binary(tmp_path): # Pre-MTP llama.cpp: only ngram variants. fake = _make_fake_llama_server( tmp_path / "llama-server", "--spec-type none,ngram-simple,ngram-mod", ) _clear_caps_cache() caps = LlamaCppBackend.probe_server_capabilities(str(fake)) assert caps["found"] is True assert caps["mtp_token"] is None assert caps["supports_mtp"] is False def test_probe_server_capabilities_handles_missing_binary(): _clear_caps_cache() caps = LlamaCppBackend.probe_server_capabilities("/no/such/llama-server") assert caps["found"] is False assert caps["supports_mtp"] is False def test_probe_server_capabilities_caches_by_mtime(tmp_path): # Same (path, mtime) -> cache hit. Bumped mtime -> re-probe. fake = _make_fake_llama_server( tmp_path / "llama-server", "--spec-type none,ngram-mod", ) _clear_caps_cache() caps1 = LlamaCppBackend.probe_server_capabilities(str(fake)) assert caps1["supports_mtp"] is False import os import time _make_fake_llama_server( fake, "--spec-type none,draft-mtp,ngram-mod", ) new_mtime = int(time.time()) + 2 os.utime(fake, (new_mtime, new_mtime)) caps2 = LlamaCppBackend.probe_server_capabilities(str(fake)) assert caps2["mtp_token"] == "draft-mtp" assert caps2["supports_mtp"] is True