- Introduced GPU probing functionality to gather and store GPU specifications in `/mnt/swarm_data/.gpu-rent-gpu.json`, aiding in performance tuning. - Updated `install_ollama.sh` and `install_llamacpp.sh` to utilize GPU information for configuring optimal runtime parameters. - Enhanced `provision.py` to include GPU probing and performance tuning logic, ensuring better resource allocation for LLM operations. - Improved documentation in `decisions.md`, `llm.md`, and `swarmui.md` to reflect changes in GPU handling and performance tuning processes. - Added new tests to validate the GPU probing and model resolution logic, ensuring robustness in handling various GPU configurations.
237 lines
7.4 KiB
Python
237 lines
7.4 KiB
Python
"""Unit tests for capture merge / URL builders (no SSH)."""
|
|
|
|
from pathlib import Path
|
|
from unittest.mock import patch
|
|
|
|
from gpu_rent.capture import (
|
|
ExtCaptureItem,
|
|
ModelCaptureItem,
|
|
merge_extensions_yaml,
|
|
merge_models_yaml,
|
|
resolve_model_item,
|
|
strip_git_auth,
|
|
)
|
|
from gpu_rent.civitai import civitai_model_url, version_ids_from_payload
|
|
from gpu_rent.errors import CloudError
|
|
from gpu_rent.manifests import parse_extensions, parse_models
|
|
|
|
|
|
def test_civitai_model_url():
|
|
assert (
|
|
civitai_model_url(2786499, 3231611, "civitai.red")
|
|
== "https://civitai.red/models/2786499?modelVersionId=3231611"
|
|
)
|
|
|
|
|
|
def test_version_ids_from_payload():
|
|
vid, mid = version_ids_from_payload({"id": 9, "modelId": 1})
|
|
assert vid == 9 and mid == 1
|
|
vid2, mid2 = version_ids_from_payload({"id": 3, "model": {"id": 7}})
|
|
assert vid2 == 3 and mid2 == 7
|
|
|
|
|
|
def test_resolve_from_sidecar():
|
|
out = resolve_model_item(
|
|
{
|
|
"kind": "lora",
|
|
"rel": "Lora/foo.safetensors",
|
|
"name": "foo.safetensors",
|
|
"version_id": 3107521,
|
|
"model_id": 2187487,
|
|
"sha256": None,
|
|
},
|
|
token="",
|
|
api_host="civitai.red",
|
|
link_host="civitai.red",
|
|
)
|
|
assert out.status == "ok"
|
|
assert out.item is not None
|
|
assert out.item.version_id == 3107521
|
|
assert "modelVersionId=3107521" in out.item.url
|
|
|
|
|
|
def test_resolve_vid_only_fetches_model_id():
|
|
with patch(
|
|
"gpu_rent.capture.fetch_model_version",
|
|
return_value=("civitai.red", {"id": 10, "modelId": 20, "name": "X"}),
|
|
) as mocked:
|
|
out = resolve_model_item(
|
|
{
|
|
"kind": "lora",
|
|
"rel": "Lora/x.safetensors",
|
|
"name": "x.safetensors",
|
|
"version_id": 10,
|
|
"model_id": None,
|
|
"sha256": None,
|
|
},
|
|
token="tok",
|
|
api_host="civitai.red",
|
|
link_host="civitai.red",
|
|
)
|
|
mocked.assert_called_once()
|
|
assert out.status == "ok"
|
|
assert out.item is not None
|
|
assert out.item.model_id == 20
|
|
assert "models/20?modelVersionId=10" in out.item.url
|
|
|
|
|
|
def test_resolve_by_hash_404_is_unknown():
|
|
with patch(
|
|
"gpu_rent.capture.fetch_model_version_by_hash",
|
|
side_effect=CloudError("Civitai by-hash abc: HTTP 404 (хосты civitai.red)"),
|
|
):
|
|
out = resolve_model_item(
|
|
{
|
|
"kind": "lora",
|
|
"rel": "Lora/m.safetensors",
|
|
"name": "m.safetensors",
|
|
"sha256": "a" * 64,
|
|
},
|
|
token="",
|
|
api_host="civitai.red",
|
|
link_host="civitai.red",
|
|
)
|
|
assert out.status == "unknown"
|
|
assert out.item is None
|
|
|
|
|
|
def test_resolve_by_hash_network_is_api_error():
|
|
with patch(
|
|
"gpu_rent.capture.fetch_model_version_by_hash",
|
|
side_effect=CloudError(
|
|
"Civitai by-hash abc: Connection timeout (хосты civitai.red)"
|
|
),
|
|
):
|
|
out = resolve_model_item(
|
|
{
|
|
"kind": "lora",
|
|
"rel": "Lora/m.safetensors",
|
|
"name": "m.safetensors",
|
|
"sha256": "b" * 64,
|
|
},
|
|
token="",
|
|
api_host="civitai.red",
|
|
link_host="civitai.red",
|
|
)
|
|
assert out.status == "api_error"
|
|
assert "timeout" in out.detail.lower() or "Connection" in out.detail
|
|
|
|
|
|
def test_merge_models_dedupe(tmp_path: Path):
|
|
path = tmp_path / "models.yaml"
|
|
path.write_text(
|
|
"lora:\n - url: https://civitai.red/models/1?modelVersionId=100\n",
|
|
encoding="utf-8",
|
|
)
|
|
items = [
|
|
ModelCaptureItem(
|
|
"lora", 100, 1, "https://civitai.red/models/1?modelVersionId=100", "old"
|
|
),
|
|
ModelCaptureItem(
|
|
"lora", 200, 2, "https://civitai.red/models/2?modelVersionId=200", "new"
|
|
),
|
|
ModelCaptureItem(
|
|
"checkpoint",
|
|
300,
|
|
3,
|
|
"https://civitai.red/models/3?modelVersionId=300",
|
|
"ckpt",
|
|
),
|
|
]
|
|
added, skipped = merge_models_yaml(path, items, dry_run=False)
|
|
assert len(added) == 2
|
|
assert any("100" in s for s in skipped)
|
|
entries = parse_models(path)
|
|
vids = {e.version_id for e in entries}
|
|
assert vids == {100, 200, 300}
|
|
assert (tmp_path / "models.yaml.bak").is_file()
|
|
|
|
|
|
def test_merge_models_dry_run(tmp_path: Path):
|
|
path = tmp_path / "models.yaml"
|
|
path.write_text("lora: []\n", encoding="utf-8")
|
|
before = path.read_text(encoding="utf-8")
|
|
items = [
|
|
ModelCaptureItem(
|
|
"lora", 1, 1, "https://civitai.red/models/1?modelVersionId=1", "x"
|
|
),
|
|
]
|
|
added, _ = merge_models_yaml(path, items, dry_run=True)
|
|
assert len(added) == 1
|
|
assert path.read_text(encoding="utf-8") == before
|
|
|
|
|
|
def test_merge_models_kind_scoped_dedupe(tmp_path: Path):
|
|
"""Same version_id under different kinds can both be kept."""
|
|
path = tmp_path / "models.yaml"
|
|
path.write_text(
|
|
"lora:\n - url: https://civitai.red/models/1?modelVersionId=100\n",
|
|
encoding="utf-8",
|
|
)
|
|
items = [
|
|
ModelCaptureItem(
|
|
"checkpoint",
|
|
100,
|
|
1,
|
|
"https://civitai.red/models/1?modelVersionId=100",
|
|
"as-ckpt",
|
|
),
|
|
]
|
|
added, skipped = merge_models_yaml(path, items, dry_run=False)
|
|
assert len(added) == 1
|
|
assert skipped == []
|
|
entries = parse_models(path)
|
|
assert {(e.kind, e.version_id) for e in entries} == {("lora", 100), ("checkpoint", 100)}
|
|
|
|
|
|
def test_merge_extensions_dedupe(tmp_path: Path):
|
|
path = tmp_path / "extensions.yaml"
|
|
path.write_text(
|
|
"swarmui:\n - url: https://github.com/org/A.git\n ref: main\n dir: A\n",
|
|
encoding="utf-8",
|
|
)
|
|
items = [
|
|
ExtCaptureItem("swarmui", "https://github.com/org/A.git", "main", "A"),
|
|
ExtCaptureItem("comfy", "https://github.com/org/B.git", "v1", "B"),
|
|
ExtCaptureItem(
|
|
"swarmui",
|
|
"https://x-access-token:SECRET@github.com/org/C.git",
|
|
"main",
|
|
"C",
|
|
),
|
|
]
|
|
added, updated, skipped = merge_extensions_yaml(path, items, dry_run=False)
|
|
assert len(added) == 2
|
|
assert updated == []
|
|
assert len(skipped) == 1
|
|
repos = parse_extensions(path)
|
|
urls = [r.url for r in repos]
|
|
assert "https://github.com/org/B.git" in urls
|
|
assert all("SECRET" not in u for u in urls)
|
|
|
|
|
|
def test_merge_extensions_updates_url_same_dir(tmp_path: Path):
|
|
path = tmp_path / "extensions.yaml"
|
|
path.write_text(
|
|
"swarmui:\n - url: https://github.com/old/A.git\n ref: main\n dir: A\n",
|
|
encoding="utf-8",
|
|
)
|
|
items = [
|
|
ExtCaptureItem("swarmui", "https://github.com/new/A.git", "develop", "A"),
|
|
]
|
|
added, updated, skipped = merge_extensions_yaml(path, items, dry_run=False)
|
|
assert added == []
|
|
assert len(updated) == 1
|
|
assert skipped == []
|
|
repos = parse_extensions(path)
|
|
assert len(repos) == 1
|
|
assert repos[0].url == "https://github.com/new/A.git"
|
|
assert repos[0].ref == "develop"
|
|
|
|
|
|
def test_strip_git_auth():
|
|
assert (
|
|
strip_git_auth("https://x-access-token:tok@github.com/org/r.git")
|
|
== "https://github.com/org/r.git"
|
|
)
|