mirror of
https://github.com/willmiao/ComfyUI-Lora-Manager.git
synced 2026-09-21 03:01:27 -03:00
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:
@@ -505,6 +505,50 @@ class CivitaiClient:
|
|||||||
logger.warning(f"Failed to fetch version by id {version_id}")
|
logger.warning(f"Failed to fetch version by id {version_id}")
|
||||||
return None
|
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]]:
|
async def _fetch_version_by_hash(self, model_hash: Optional[str]) -> Optional[Dict[str, Any]]:
|
||||||
if not model_hash:
|
if not model_hash:
|
||||||
return None
|
return None
|
||||||
|
|||||||
@@ -35,6 +35,7 @@ from .service_registry import ServiceRegistry
|
|||||||
from .settings_manager import get_settings_manager
|
from .settings_manager import get_settings_manager
|
||||||
from .metadata_service import get_default_metadata_provider, get_metadata_provider
|
from .metadata_service import get_default_metadata_provider, get_metadata_provider
|
||||||
from .downloader import get_downloader, DownloadProgress, DownloadStreamControl
|
from .downloader import get_downloader, DownloadProgress, DownloadStreamControl
|
||||||
|
from .errors import RateLimitError
|
||||||
from .aria2_downloader import Aria2Error, get_aria2_downloader
|
from .aria2_downloader import Aria2Error, get_aria2_downloader
|
||||||
from .aria2_transfer_state import Aria2TransferStateStore
|
from .aria2_transfer_state import Aria2TransferStateStore
|
||||||
from .download_queue_service import DownloadQueueService
|
from .download_queue_service import DownloadQueueService
|
||||||
@@ -929,6 +930,42 @@ class DownloadManager:
|
|||||||
|
|
||||||
return download_urls
|
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(
|
def _build_metadata_for_resume(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
@@ -1858,6 +1895,24 @@ class DownloadManager:
|
|||||||
if not download_urls:
|
if not download_urls:
|
||||||
return {"success": False, "error": "No mirror URL found"}
|
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
|
# 3. Prepare download
|
||||||
file_name = file_info.get("name", "")
|
file_name = file_info.get("name", "")
|
||||||
if not file_name:
|
if not file_name:
|
||||||
|
|||||||
@@ -169,6 +169,17 @@ class ModelMetadataProvider(ABC):
|
|||||||
"""Published model count for the user; None when unsupported."""
|
"""Published model count for the user; None when unsupported."""
|
||||||
return None
|
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):
|
class CivitaiModelMetadataProvider(ModelMetadataProvider):
|
||||||
"""Provider that uses Civitai API for metadata"""
|
"""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]:
|
async def get_creator_model_count(self, username: str) -> Optional[int]:
|
||||||
return await self.client.get_creator_model_count(username)
|
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):
|
class CivArchiveModelMetadataProvider(ModelMetadataProvider):
|
||||||
"""Provider that uses CivArchive API for metadata"""
|
"""Provider that uses CivArchive API for metadata"""
|
||||||
|
|
||||||
@@ -700,6 +716,37 @@ class FallbackMetadataProvider(ModelMetadataProvider):
|
|||||||
continue
|
continue
|
||||||
return None
|
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):
|
def _iter_providers(self):
|
||||||
return zip(self.providers, self._provider_labels)
|
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]:
|
async def get_creator_model_count(self, username: str) -> Optional[int]:
|
||||||
return await self._provider.get_creator_model_count(username)
|
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:
|
class ModelMetadataProviderManager:
|
||||||
"""Manager for selecting and using model metadata providers"""
|
"""Manager for selecting and using model metadata providers"""
|
||||||
|
|
||||||
|
|||||||
@@ -818,3 +818,44 @@ async def test_get_model_by_hash_rejects_empty_placeholder_without_request(downl
|
|||||||
assert result is None
|
assert result is None
|
||||||
assert error == "Model not found"
|
assert error == "Model not found"
|
||||||
assert requested == []
|
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.
|
# Partial files are preserved for a future resume from disk.
|
||||||
assert save_path.exists()
|
assert save_path.exists()
|
||||||
assert control_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)
|
result, _ = await helper.run("test", succeeding)
|
||||||
assert result == {"ok": True}
|
assert result == {"ok": True}
|
||||||
assert calls == 2 # Retried once (small retry_after)
|
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)]
|
||||||
|
|||||||
Reference in New Issue
Block a user