Files
automaton/tests/test_vram_detect.py
T

300 lines
9.6 KiB
Python
Raw Normal View History

"""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
assert recommended_kb <= 128_000 * 0.75 # after headroom
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_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)