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.
This commit is contained in:
Will Miao
2026-09-06 22:23:48 +08:00
parent a17399d667
commit 41302e75ba
6 changed files with 427 additions and 0 deletions
+41
View File
@@ -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)
@@ -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
@@ -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)]