310 lines
10 KiB
Python
310 lines
10 KiB
Python
"""Tests for scripts/vram_detect.py."""
|
|
|
|
from pathlib import Path
|
|
|
|
import pytest
|
|
|
|
import scripts.vram_detect as vram
|
|
|
|
|
|
def test_lookup_model_context() -> None:
|
|
assert vram._lookup_model_context("gpt-4o") == 128_000
|
|
assert vram._lookup_model_context("claude-3-5-sonnet") == 200_000
|
|
assert vram._lookup_model_context("unknown-model") == 0
|
|
|
|
|
|
def test_parse_token_value() -> None:
|
|
assert vram._parse_token_value("128k") == 128_000
|
|
assert vram._parse_token_value("128K") == 128_000
|
|
assert vram._parse_token_value("128000") == 128_000
|
|
assert vram._parse_token_value("nonsense") is None
|
|
|
|
|
|
def test_extract_value() -> None:
|
|
assert vram._extract_value("- **Model**: gpt-4o") == "gpt-4o"
|
|
assert vram._extract_value("model = gpt-4o # comment") == "gpt-4o"
|
|
assert vram._extract_value("- **Model**: auto") == "auto"
|
|
|
|
|
|
def test_parse_config_model(tmp_path: Path) -> None:
|
|
config = tmp_path / "config.md"
|
|
config.write_text(
|
|
"## VRAM Configuration\n"
|
|
"- **Auto-detect**: Yes\n"
|
|
"- **Target context**: 16k tokens\n"
|
|
"- **Headroom**: 25%\n"
|
|
"\n"
|
|
"## Model Configuration\n"
|
|
"- **Model**: gpt-4o\n"
|
|
"- **Override context window**: 128k\n"
|
|
)
|
|
model, override = vram._parse_config_model(config)
|
|
assert model == "gpt-4o"
|
|
assert override == 128_000
|
|
|
|
|
|
def test_parse_config_model_skips_code_blocks(tmp_path: Path) -> None:
|
|
"""Example code blocks should not be parsed as live config."""
|
|
config = tmp_path / "config.md"
|
|
config.write_text(
|
|
"## VRAM Configuration\n"
|
|
"- **Headroom**: 25%\n"
|
|
"```\n"
|
|
"- **Headroom**: 99%\n"
|
|
"```\n"
|
|
)
|
|
parsed = vram.parse_vram_config(config)
|
|
assert parsed["headroom_pct"] == 25
|
|
|
|
|
|
def test_parse_vram_config_manual(tmp_path: Path) -> None:
|
|
config = tmp_path / "config.md"
|
|
config.write_text(
|
|
"## VRAM Configuration\n"
|
|
"- **Auto-detect**: No\n"
|
|
"- **Target context**: 8k\n"
|
|
"- **Headroom**: 30%\n"
|
|
"- **Max peak context per sub-task**: 5.6k\n"
|
|
)
|
|
parsed = vram.parse_vram_config(config)
|
|
assert parsed["auto_detect"] is False
|
|
assert parsed["target_context_kb"] == 8000
|
|
assert parsed["headroom_pct"] == 30
|
|
assert parsed["max_peak_kb"] == 5600
|
|
|
|
|
|
def test_recommend_context_api_model() -> None:
|
|
config = {"auto_detect": True, "headroom_pct": 25, "target_context_kb": 0, "max_peak_kb": 0}
|
|
_, recommended_kb, max_peak_kb = vram.recommend_context(
|
|
gpu_vram_gb=0,
|
|
ram_gb=16,
|
|
model_context_kb=128_000,
|
|
overhead_tokens=4000,
|
|
config=config,
|
|
)
|
|
assert recommended_kb > 0
|
|
assert max_peak_kb > 0
|
|
# Context sizing fix (task fix-context-sizing): headroom is applied
|
|
# EXACTLY ONCE. recommended_kb is the raw budget net of overhead
|
|
# (no headroom); max_peak_kb is post-headroom.
|
|
assert recommended_kb == 128_000 - 4000
|
|
assert max_peak_kb == (128_000 - 4000) * 75 // 100
|
|
|
|
|
|
def test_recommend_context_manual_mode() -> None:
|
|
config = {"auto_detect": False, "headroom_pct": 30, "target_context_kb": 8000, "max_peak_kb": 0}
|
|
headroom, recommended_kb, max_peak_kb = vram.recommend_context(
|
|
gpu_vram_gb=0,
|
|
ram_gb=16,
|
|
model_context_kb=0,
|
|
overhead_tokens=0,
|
|
config=config,
|
|
)
|
|
assert headroom == 30
|
|
assert recommended_kb == 8000
|
|
assert max_peak_kb == 5600
|
|
|
|
|
|
def test_extract_model_from_file_respects_10kb_limit(tmp_path: Path) -> None:
|
|
"""Only the first 10KB of an API config file is scanned."""
|
|
env_file = tmp_path / ".env"
|
|
# Put the model name far beyond 10KB.
|
|
env_file.write_text("x" * 11_000 + "\nMODEL=far-away-model\n")
|
|
model = vram._extract_model_from_file(env_file)
|
|
assert model is None
|
|
|
|
|
|
# --- Cross-platform fixtures (module-level multiline string constants) ---
|
|
|
|
MEMINFO_LINUX_16GB = "MemTotal: 16384000 kB\nMemAvailable: 8192000 kB\n"
|
|
MEMINFO_LINUX_32GB = "MemTotal: 32768000 kB\nMemAvailable: 16384000 kB\n"
|
|
|
|
APPLE_M2_PROFILER = """\
|
|
Graphics/Displays:
|
|
Apple M2:
|
|
Chipset Model: Apple M2
|
|
Type: GPU
|
|
Bus: Built-In
|
|
Total Number of Cores: 10
|
|
Vendor: Apple (0x106b)
|
|
Metal: Supported, version 2
|
|
"""
|
|
|
|
WMIC_VIDEOCONTROLLER = """\
|
|
AdapterRAM=8589934592
|
|
Name=NVIDIA GeForce RTX 3060
|
|
"""
|
|
|
|
WMIC_COMPUTERSYSTEM = """\
|
|
TotalPhysicalMemory=34359738368
|
|
"""
|
|
|
|
OLLAMA_LIST = """\
|
|
NAME ID SIZE MODIFIED
|
|
llama-3.1-8b abc 4.7GB 2 days ago
|
|
"""
|
|
|
|
|
|
class _FakeResult:
|
|
"""Minimal stand-in for subprocess.CompletedProcess."""
|
|
|
|
def __init__(self, stdout: str = "", returncode: int = 0) -> None:
|
|
self.stdout = stdout
|
|
self.returncode = returncode
|
|
|
|
|
|
def _make_fake_run(responses: dict[str, str]):
|
|
"""Build a subprocess.run fake that maps cmd[0] (or joined PS string) to stdout."""
|
|
|
|
def fake_run(cmd, *args, **kwargs):
|
|
name = cmd[0]
|
|
if name == "powershell":
|
|
joined = cmd[-1] if len(cmd) > 4 else ""
|
|
for key, out in responses.items():
|
|
if key in joined:
|
|
return _FakeResult(out)
|
|
return _FakeResult("", 1)
|
|
return _FakeResult(responses.get(name, ""))
|
|
|
|
return fake_run
|
|
|
|
|
|
def _patch_meminfo(monkeypatch, text: str) -> None:
|
|
"""Patch Path.exists/read_text to serve `text` for /proc/meminfo only."""
|
|
orig_exists = Path.exists
|
|
orig_read_text = Path.read_text
|
|
|
|
def fake_exists(self):
|
|
if self.as_posix() == "/proc/meminfo":
|
|
return True
|
|
return orig_exists(self)
|
|
|
|
def fake_read_text(self, *args, **kwargs):
|
|
if self.as_posix() == "/proc/meminfo":
|
|
return text
|
|
return orig_read_text(self, *args, **kwargs)
|
|
|
|
monkeypatch.setattr(Path, "exists", fake_exists)
|
|
monkeypatch.setattr(Path, "read_text", fake_read_text)
|
|
|
|
|
|
def test_detect_ram_linux(monkeypatch) -> None:
|
|
monkeypatch.setattr(vram.platform, "system", lambda: "Linux")
|
|
_patch_meminfo(monkeypatch, MEMINFO_LINUX_16GB)
|
|
total, available = vram.detect_ram()
|
|
assert (total, available) == (16_384_000, 8_192_000)
|
|
|
|
|
|
def test_detect_ram_macos(monkeypatch) -> None:
|
|
monkeypatch.setattr(vram.platform, "system", lambda: "Darwin")
|
|
monkeypatch.setattr(vram.shutil, "which", lambda name: name)
|
|
monkeypatch.setattr(
|
|
vram.subprocess, "run", _make_fake_run({"sysctl": "34359738368"})
|
|
)
|
|
total, available = vram.detect_ram()
|
|
assert total == 33_554_432
|
|
assert available == total
|
|
|
|
|
|
def test_detect_ram_windows(monkeypatch) -> None:
|
|
monkeypatch.setattr(vram.platform, "system", lambda: "Windows")
|
|
monkeypatch.setattr(vram.shutil, "which", lambda name: name)
|
|
monkeypatch.setattr(
|
|
vram.subprocess, "run", _make_fake_run({"wmic": WMIC_COMPUTERSYSTEM})
|
|
)
|
|
total, available = vram.detect_ram()
|
|
assert total == 33_554_432
|
|
|
|
|
|
def test_detect_gpu_vram_nvidia_linux(monkeypatch) -> None:
|
|
monkeypatch.setattr(vram.platform, "system", lambda: "Linux")
|
|
monkeypatch.setattr(vram.shutil, "which", lambda name: name)
|
|
monkeypatch.setattr(
|
|
vram.subprocess, "run", _make_fake_run({"nvidia-smi": "24576\n"})
|
|
)
|
|
total, per_gpu, num = vram.detect_gpu_vram()
|
|
assert (total, per_gpu, num) == (25_165_824, 25_165_824, 1)
|
|
|
|
|
|
def test_detect_gpu_vram_apple_silicon(monkeypatch) -> None:
|
|
monkeypatch.setattr(vram.platform, "system", lambda: "Darwin")
|
|
monkeypatch.setattr(vram.shutil, "which", lambda name: name)
|
|
monkeypatch.setattr(
|
|
vram.subprocess,
|
|
"run",
|
|
_make_fake_run(
|
|
{"system_profiler": APPLE_M2_PROFILER, "sysctl": "17179869184"}
|
|
),
|
|
)
|
|
total, per_gpu, num = vram.detect_gpu_vram()
|
|
assert total > 0
|
|
assert num >= 1
|
|
|
|
|
|
def test_detect_gpu_vram_windows_wmic(monkeypatch) -> None:
|
|
monkeypatch.setattr(vram.platform, "system", lambda: "Windows")
|
|
monkeypatch.setattr(vram.shutil, "which", lambda name: name)
|
|
monkeypatch.setattr(
|
|
vram.subprocess, "run", _make_fake_run({"wmic": WMIC_VIDEOCONTROLLER})
|
|
)
|
|
total, per_gpu, num = vram.detect_gpu_vram()
|
|
assert total == 8_388_608
|
|
|
|
|
|
def test_lookup_model_context_unknown_returns_zero() -> None:
|
|
assert vram._lookup_model_context("completely-unknown-model") == 0
|
|
|
|
|
|
def test_lookup_model_context_prefix_match() -> None:
|
|
assert vram._lookup_model_context("deepseek-r1:7b") == 64_000
|
|
assert vram._lookup_model_context("llama-3.1-8b-instruct") == 128_000
|
|
|
|
|
|
def test_lookup_model_context_no_false_prefix_match() -> None:
|
|
"""phi-4-mini should NOT match phi-4 (different model, wrong context)."""
|
|
assert vram._lookup_model_context("phi-4-mini-instruct") == 0
|
|
assert vram._lookup_model_context("gpt-4o-foo-unknown") == 0
|
|
|
|
|
|
def test_detect_model_context_ollama_probe(monkeypatch, tmp_path: Path) -> None:
|
|
monkeypatch.setattr(vram.Path, "home", lambda: tmp_path)
|
|
monkeypatch.setattr(
|
|
vram.shutil, "which", lambda name: name if name == "ollama" else None
|
|
)
|
|
monkeypatch.setattr(
|
|
vram.subprocess, "run", _make_fake_run({"ollama": OLLAMA_LIST})
|
|
)
|
|
context = vram.detect_model_context(None, tmp_path)
|
|
assert context == 128_000
|
|
|
|
|
|
def test_run_command_windows_powershell_wrapper(monkeypatch) -> None:
|
|
monkeypatch.setattr(vram.platform, "system", lambda: "Windows")
|
|
monkeypatch.setattr(vram.shutil, "which", lambda name: name)
|
|
captured: dict = {}
|
|
|
|
def fake_run(cmd, *args, **kwargs):
|
|
captured["cmd"] = cmd
|
|
return _FakeResult("ok", 0)
|
|
|
|
monkeypatch.setattr(vram.subprocess, "run", fake_run)
|
|
result = vram.run_command(
|
|
["Get-CimInstance", "Win32_VideoController", "-Property", "AdapterRAM"]
|
|
)
|
|
assert result == "ok"
|
|
cmd = captured["cmd"]
|
|
assert cmd[0] == "powershell"
|
|
assert "-NoProfile" in cmd
|
|
assert "-NoLogo" in cmd
|
|
assert "-Command" in cmd
|
|
assert "Get-CimInstance" in cmd[-1]
|
|
assert "Win32_VideoController" in cmd[-1]
|
|
|
|
|
|
def test_detect_ram_linux_regression(monkeypatch) -> None:
|
|
"""LOCKED regression guard — exact match on (total_kb, available_kb)."""
|
|
monkeypatch.setattr(vram.platform, "system", lambda: "Linux")
|
|
_patch_meminfo(monkeypatch, MEMINFO_LINUX_32GB)
|
|
total, available = vram._detect_ram_linux()
|
|
assert (total, available) == (32_768_000, 16_384_000)
|