"""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)"), ), patch( "gpu_rent.capture.lookup_by_sha256", return_value=None, ): 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_404_hf_fallback(): from gpu_rent.huggingface import HfFileHit with patch( "gpu_rent.capture.fetch_model_version_by_hash", side_effect=CloudError("Civitai by-hash abc: HTTP 404 (хосты civitai.red)"), ), patch( "gpu_rent.capture.lookup_by_sha256", return_value=HfFileHit( repo_id="org/model", filename="m.safetensors", url="https://huggingface.co/org/model/resolve/main/m.safetensors", sha256="a" * 64, title="org/model", ), ): 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", hf_token="hf_x", ) assert out.status == "ok" assert out.item is not None assert out.item.source == "huggingface" assert "huggingface.co" in out.item.url 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( kind="lora", version_id=100, model_id=1, url="https://civitai.red/models/1?modelVersionId=100", title="old", ), ModelCaptureItem( kind="lora", version_id=200, model_id=2, url="https://civitai.red/models/2?modelVersionId=200", title="new", ), ModelCaptureItem( kind="checkpoint", version_id=300, model_id=3, url="https://civitai.red/models/3?modelVersionId=300", title="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( kind="lora", version_id=1, model_id=1, url="https://civitai.red/models/1?modelVersionId=1", title="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( kind="checkpoint", version_id=100, model_id=1, url="https://civitai.red/models/1?modelVersionId=100", title="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" )