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
@@ -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)]