From 41302e75ba3671ca53adb82dd387bc80d138e9b7 Mon Sep 17 00:00:00 2001 From: Will Miao Date: Sun, 6 Sep 2026 22:23:48 +0800 Subject: [PATCH] fix(download): save multi-variant files under raw stored filenames (#1100) The public REST API rewrites files[].name to "{model}_{version}" for non-LoRA model types, so every precision variant of a multi-file version shared one name and landed on disk with a random short-hash suffix. Fetch the raw stored filename from the model-versions/mini endpoint (always pinned with modelFileId) and use it for the on-disk name and metadata when available; fall back silently to the REST name otherwise. CivArchive already serves raw names and is skipped. --- py/services/civitai_client.py | 44 +++++ py/services/download_manager.py | 55 ++++++ py/services/model_metadata_provider.py | 57 ++++++ tests/services/test_civitai_client.py | 41 +++++ tests/services/test_download_manager_basic.py | 171 ++++++++++++++++++ .../services/test_model_metadata_provider.py | 59 ++++++ 6 files changed, 427 insertions(+) diff --git a/py/services/civitai_client.py b/py/services/civitai_client.py index fbcdaa06..b668bfd9 100644 --- a/py/services/civitai_client.py +++ b/py/services/civitai_client.py @@ -505,6 +505,50 @@ class CivitaiClient: logger.warning(f"Failed to fetch version by id {version_id}") return None + async def get_version_file_mini( + self, version_id: int, file_id: int + ) -> Optional[Dict[str, Any]]: + """Fetch raw stored file info via the model-versions/mini endpoint. + + The public REST API rewrites ``files[].name`` to + ``"{model}_{version}"`` for non-LoRA model types, so every + precision variant of a multi-file version shares one name (#1100). + The mini endpoint returns the raw ``ModelFile.name`` in + ``fileName``. ``file_id`` is mandatory: without it mini picks a + file via its own primary-file logic, which can disagree with the + REST ``primary`` flag. + + Returns the mini payload dict on success, None on any failure. + """ + try: + success, data = await self._make_request( + "GET", + f"{self.base_url}/model-versions/mini/{version_id}", + params={"modelFileId": file_id}, + use_auth=True, + ) + if success and isinstance(data, dict): + return data + if is_expected_offline_error(data): + return None + logger.debug( + "Mini endpoint lookup failed for version %s file %s: %s", + version_id, + file_id, + data, + ) + return None + except RateLimitError: + raise + except Exception as exc: + logger.debug( + "Error fetching mini info for version %s file %s: %s", + version_id, + file_id, + exc, + ) + return None + async def _fetch_version_by_hash(self, model_hash: Optional[str]) -> Optional[Dict[str, Any]]: if not model_hash: return None diff --git a/py/services/download_manager.py b/py/services/download_manager.py index 0cdd760d..d28be353 100644 --- a/py/services/download_manager.py +++ b/py/services/download_manager.py @@ -35,6 +35,7 @@ from .service_registry import ServiceRegistry from .settings_manager import get_settings_manager from .metadata_service import get_default_metadata_provider, get_metadata_provider from .downloader import get_downloader, DownloadProgress, DownloadStreamControl +from .errors import RateLimitError from .aria2_downloader import Aria2Error, get_aria2_downloader from .aria2_transfer_state import Aria2TransferStateStore from .download_queue_service import DownloadQueueService @@ -929,6 +930,42 @@ class DownloadManager: return download_urls + async def _fetch_raw_file_name( + self, + metadata_provider, + version_id: Optional[int], + file_id: Any, + ) -> Optional[str]: + """Best-effort lookup of the raw stored filename via the CivitAI + model-versions/mini endpoint (#1100). Returns None on any failure so + the caller can fall back to the (possibly rewritten) REST name.""" + if version_id is None or file_id is None: + return None + fetch = getattr(metadata_provider, "get_version_file_mini", None) + if fetch is None: + return None + try: + mini_info = await fetch(int(version_id), int(file_id)) + except (TypeError, ValueError): + return None + except RateLimitError: + raise + except Exception as exc: + logger.debug( + "Mini endpoint lookup failed for version %s file %s: %s", + version_id, + file_id, + exc, + ) + return None + if not isinstance(mini_info, dict): + return None + raw_name = mini_info.get("fileName") + if not isinstance(raw_name, str) or not raw_name.strip(): + return None + # Defensive: never let a path component slip into the filename. + return os.path.basename(raw_name.strip()) or None + def _build_metadata_for_resume( self, *, @@ -1858,6 +1895,24 @@ class DownloadManager: if not download_urls: return {"success": False, "error": "No mirror URL found"} + # The public REST API rewrites files[].name to + # "{model}_{version}" for non-LoRA model types, so every + # precision variant of a multi-file version shares one name and + # lands on disk with a random short-hash suffix. The mini + # endpoint returns the raw stored filename (#1100). CivArchive + # already serves raw names. + if source != "civarchive": + raw_file_name = await self._fetch_raw_file_name( + metadata_provider, resolved_version_id, file_info.get("id") + ) + if raw_file_name and raw_file_name != file_info.get("name"): + logger.info( + "[download] Using raw stored filename '%s' instead of REST name '%s'", + raw_file_name, + file_info.get("name"), + ) + file_info = {**file_info, "name": raw_file_name} + # 3. Prepare download file_name = file_info.get("name", "") if not file_name: diff --git a/py/services/model_metadata_provider.py b/py/services/model_metadata_provider.py index a6b6b528..a45579bf 100644 --- a/py/services/model_metadata_provider.py +++ b/py/services/model_metadata_provider.py @@ -169,6 +169,17 @@ class ModelMetadataProvider(ABC): """Published model count for the user; None when unsupported.""" return None + async def get_version_file_mini( + self, version_id: int, file_id: int + ) -> Optional[Dict[str, Any]]: + """Fetch raw stored file info via CivitAI's model-versions/mini endpoint. + + Only the CivitAI provider implements this (#1100); other providers + already serve raw file names (CivArchive) or cannot resolve this + lookup (SQLite), so the default is None. + """ + return None + class CivitaiModelMetadataProvider(ModelMetadataProvider): """Provider that uses Civitai API for metadata""" @@ -203,6 +214,11 @@ class CivitaiModelMetadataProvider(ModelMetadataProvider): async def get_creator_model_count(self, username: str) -> Optional[int]: return await self.client.get_creator_model_count(username) + async def get_version_file_mini( + self, version_id: int, file_id: int + ) -> Optional[Dict[str, Any]]: + return await self.client.get_version_file_mini(version_id, file_id) + class CivArchiveModelMetadataProvider(ModelMetadataProvider): """Provider that uses CivArchive API for metadata""" @@ -700,6 +716,37 @@ class FallbackMetadataProvider(ModelMetadataProvider): continue return None + async def get_version_file_mini( + self, version_id: int, file_id: int + ) -> Optional[Dict[str, Any]]: + rate_limited = False + for provider, label in self._iter_providers(): + if rate_limited and label not in _LOCAL_PROVIDER_LABELS: + continue + try: + result = await self._call_with_rate_limit( + label, + provider.get_version_file_mini, + version_id, + file_id, + ) + if result: + return result + except RateLimitError as exc: + rate_limited = True + logger.warning( + "Provider %s is rate-limited (retry_after=%.0fs); not failing over to other network providers", + label, + exc.retry_after or 0, + ) + continue + except Exception as e: + logger.debug( + "Provider %s failed for get_version_file_mini: %s", label, e + ) + continue + return None + def _iter_providers(self): return zip(self.providers, self._provider_labels) @@ -791,6 +838,16 @@ class RateLimitRetryingProvider(ModelMetadataProvider): async def get_creator_model_count(self, username: str) -> Optional[int]: return await self._provider.get_creator_model_count(username) + async def get_version_file_mini( + self, version_id: int, file_id: int + ) -> Optional[Dict[str, Any]]: + return await self._rate_limit_helper.run( + self._label, + self._provider.get_version_file_mini, + version_id, + file_id, + ) + class ModelMetadataProviderManager: """Manager for selecting and using model metadata providers""" diff --git a/tests/services/test_civitai_client.py b/tests/services/test_civitai_client.py index b07cabdd..84abe602 100644 --- a/tests/services/test_civitai_client.py +++ b/tests/services/test_civitai_client.py @@ -818,3 +818,44 @@ async def test_get_model_by_hash_rejects_empty_placeholder_without_request(downl assert result is None assert error == "Model not found" assert requested == [] + + +async def test_get_version_file_mini_returns_payload(downloader): + """The mini endpoint returns the raw stored filename (#1100).""" + client = await CivitaiClient.get_instance() + + async def fake_make_request(method, url, use_auth=True, **kwargs): + assert method == "GET" + assert url.endswith("/model-versions/mini/3284136") + assert kwargs.get("params") == {"modelFileId": 3168412} + assert use_auth is True + return True, {"fileName": "CyberRealistic_zit_v8.0_bf16.safetensors"} + + downloader.make_request = fake_make_request + + result = await client.get_version_file_mini(3284136, 3168412) + + assert result == {"fileName": "CyberRealistic_zit_v8.0_bf16.safetensors"} + + +async def test_get_version_file_mini_returns_none_on_failure(downloader): + client = await CivitaiClient.get_instance() + + async def fake_make_request(method, url, use_auth=True, **kwargs): + return False, "Model file 2 not found in version 1" + + downloader.make_request = fake_make_request + + assert await client.get_version_file_mini(1, 2) is None + + +async def test_get_version_file_mini_propagates_rate_limit(downloader): + client = await CivitaiClient.get_instance() + + async def fake_make_request(method, url, use_auth=True, **kwargs): + return False, RateLimitError("limited", retry_after=1.0) + + downloader.make_request = fake_make_request + + with pytest.raises(RateLimitError): + await client.get_version_file_mini(1, 2) diff --git a/tests/services/test_download_manager_basic.py b/tests/services/test_download_manager_basic.py index c58503f4..62b1d78a 100644 --- a/tests/services/test_download_manager_basic.py +++ b/tests/services/test_download_manager_basic.py @@ -2098,3 +2098,174 @@ async def test_discard_cleared_downloads_stops_tracking_and_preserves_files( # Partial files are preserved for a future resume from disk. assert save_path.exists() assert control_path.exists() + + +@pytest.mark.asyncio +async def test_download_uses_raw_file_name_from_mini_endpoint( + monkeypatch, scanners, metadata_provider, tmp_path +): + """#1100: when the REST name is rewritten ("{model}_{version}"), the raw + stored filename from the mini endpoint wins for the on-disk name.""" + manager = DownloadManager() + get_settings_manager().settings["default_unet_root"] = str(tmp_path / "unet") + metadata_provider.payload = { + "id": 3284136, + "model": {"type": "Checkpoint", "tags": ["realistic"]}, + "baseModel": "ZImageTurbo", + "creator": {"username": "Author"}, + "files": [ + { + "id": 3168412, + "type": "Model", + "primary": True, + "name": "cyberrealisticZImage_v80.safetensors", + "downloadUrl": "https://civitai.com/api/download/models/3284136?fileId=3168412", + } + ], + } + metadata_provider.get_version_file_mini = AsyncMock( + return_value={"fileName": "CyberRealistic_zit_v8.0_bf16.safetensors"} + ) + + captured = {} + + async def fake_execute_download(self, **kwargs): + captured["download_urls"] = kwargs["download_urls"] + captured["file_path"] = kwargs["metadata"].file_path + return {"success": True} + + monkeypatch.setattr( + DownloadManager, "_execute_download", fake_execute_download, raising=False + ) + + result = await manager.download_from_civitai( + model_version_id=3284136, + save_dir=str(tmp_path), + use_default_paths=True, + progress_callback=None, + source=None, + ) + + assert result["success"] is True, result + metadata_provider.get_version_file_mini.assert_awaited_once_with(3284136, 3168412) + assert captured["file_path"].endswith("CyberRealistic_zit_v8.0_bf16.safetensors") + # The file's own pinned downloadUrl is untouched. + assert captured["download_urls"] == [ + "https://civitai.com/api/download/models/3284136?fileId=3168412" + ] + + +@pytest.mark.asyncio +async def test_download_falls_back_to_rest_name_when_mini_fails( + monkeypatch, scanners, metadata_provider, tmp_path +): + """A failed/absent mini lookup must keep the previous behavior.""" + manager = DownloadManager() + metadata_provider.payload = { + "id": 42, + "model": {"type": "Checkpoint", "tags": ["fantasy"]}, + "baseModel": "BaseModel", + "creator": {"username": "Author"}, + "files": [ + { + "id": 1001, + "type": "Model", + "primary": True, + "name": "rewritten_v10.safetensors", + "downloadUrl": "https://example.invalid/file.safetensors", + } + ], + } + metadata_provider.get_version_file_mini = AsyncMock(return_value=None) + + captured = {} + + async def fake_execute_download(self, **kwargs): + captured["file_path"] = kwargs["metadata"].file_path + return {"success": True} + + monkeypatch.setattr( + DownloadManager, "_execute_download", fake_execute_download, raising=False + ) + + result = await manager.download_from_civitai( + model_version_id=42, + save_dir=str(tmp_path), + use_default_paths=True, + progress_callback=None, + source=None, + ) + + assert result["success"] is True + assert captured["file_path"].endswith("rewritten_v10.safetensors") + + +@pytest.mark.asyncio +async def test_download_skips_mini_lookup_for_civarchive_source( + monkeypatch, scanners, metadata_provider, tmp_path +): + """CivArchive already serves raw stored names — no mini call.""" + manager = DownloadManager() + mini_mock = AsyncMock(return_value={"fileName": "should_not_be_used.safetensors"}) + metadata_provider.get_version_file_mini = mini_mock + + monkeypatch.setattr( + download_manager, + "get_metadata_provider", + AsyncMock(return_value=metadata_provider), + ) + + captured = {} + + async def fake_execute_download(self, **kwargs): + captured["file_path"] = kwargs["metadata"].file_path + return {"success": True} + + monkeypatch.setattr( + DownloadManager, "_execute_download", fake_execute_download, raising=False + ) + + result = await manager.download_from_civitai( + model_version_id=99, + save_dir=str(tmp_path), + use_default_paths=True, + progress_callback=None, + source="civarchive", + ) + + assert result["success"] is True + mini_mock.assert_not_called() + assert captured["file_path"].endswith("file.safetensors") + + +@pytest.mark.asyncio +async def test_fetch_raw_file_name_edge_cases(): + """_fetch_raw_file_name never raises and strips path components.""" + manager = DownloadManager() + + provider = SimpleNamespace() + + # Missing version id / file id short-circuit before any provider call. + provider.get_version_file_mini = AsyncMock() + assert await manager._fetch_raw_file_name(provider, None, 1) is None + assert await manager._fetch_raw_file_name(provider, 1, None) is None + provider.get_version_file_mini.assert_not_called() + + # Provider without the method (older mocks / non-CivitAI providers). + assert await manager._fetch_raw_file_name(object(), 1, 2) is None + + # Non-dict payload, empty fileName. + provider.get_version_file_mini = AsyncMock(return_value="oops") + assert await manager._fetch_raw_file_name(provider, 1, 2) is None + provider.get_version_file_mini = AsyncMock(return_value={"fileName": " "}) + assert await manager._fetch_raw_file_name(provider, 1, 2) is None + + # Path components are stripped defensively. + provider.get_version_file_mini = AsyncMock( + return_value={"fileName": "../evil/model.safetensors"} + ) + assert await manager._fetch_raw_file_name(provider, 1, 2) == "model.safetensors" + + # Provider exceptions degrade to None. + provider.get_version_file_mini = AsyncMock(side_effect=RuntimeError("boom")) + assert await manager._fetch_raw_file_name(provider, 1, 2) is None diff --git a/tests/services/test_model_metadata_provider.py b/tests/services/test_model_metadata_provider.py index bbde1f9e..d3f2edc6 100644 --- a/tests/services/test_model_metadata_provider.py +++ b/tests/services/test_model_metadata_provider.py @@ -207,3 +207,62 @@ async def test_retry_helper_retries_normally_for_small_retry_after(monkeypatch): result, _ = await helper.run("test", succeeding) assert result == {"ok": True} assert calls == 2 # Retried once (small retry_after) + + +class MiniCapableProvider(ModelMetadataProvider): + """Provider that serves raw file names via the mini endpoint (#1100).""" + + def __init__(self, payload=None) -> None: + self.payload = payload + self.calls = [] + + async def get_model_by_hash(self, model_hash: str): + return None, None + + async def get_model_versions(self, model_id: str): + return None + + async def get_model_version(self, model_id=None, version_id=None): + return None + + async def get_model_version_info(self, version_id: str): + return None, None + + async def get_user_models(self, username: str, cursor=None): + return None + + async def get_version_file_mini(self, version_id: int, file_id: int): + self.calls.append((version_id, file_id)) + return self.payload + + +@pytest.mark.asyncio +async def test_base_provider_get_version_file_mini_defaults_to_none(): + provider = TrackingProvider() + assert await provider.get_version_file_mini(1, 2) is None + + +@pytest.mark.asyncio +async def test_fallback_get_version_file_mini_returns_first_hit(): + primary = TrackingProvider() # base default: None + secondary = MiniCapableProvider({"fileName": "raw.safetensors"}) + + fallback = FallbackMetadataProvider( + [("primary", primary), ("secondary", secondary)], + ) + + result = await fallback.get_version_file_mini(10, 20) + + assert result == {"fileName": "raw.safetensors"} + assert secondary.calls == [(10, 20)] + + +@pytest.mark.asyncio +async def test_rate_limit_retrying_provider_delegates_get_version_file_mini(): + inner = MiniCapableProvider({"fileName": "raw.safetensors"}) + wrapper = RateLimitRetryingProvider(inner, label="inner") + + result = await wrapper.get_version_file_mini(10, 20) + + assert result == {"fileName": "raw.safetensors"} + assert inner.calls == [(10, 20)]