Files
ComfyUI-Lora-Manager/tests/routes/test_model_update_handler.py
T
Will Miao e42d649df0 refactor(update)!: make obtainability a property of updates, not a surface
The owner could not tell from the UI what "Buzz Price Tracking" enabled, what the
"Price alert threshold" number meant, or what "Price Alerts" was alerting about.
That was not a copy problem: the implementation exposed our mechanism (a page
scrape) and our SQL predicates as the user's concepts. Two concrete defects came
from the same root:

- the alert population included versions the user already owns (neither the event
  generator nor the query filtered on is_in_library; in the owner's library 28 of
  52 gated versions were already downloaded, so most "alerts" were about files
  already on disk, which cannot become cheaper *for them*);
- a threshold-filtered state list lived in a notification surface, so an empty
  panel had three indistinguishable causes and read as a broken feature.

The information model is now the version plus ownership: cost is shown only where
a decision exists. Owned -> nothing. Not owned and free -> nothing. Not owned and
gated -> the price when it is known, `Paid` without a number when it is not, and
early access keeps its countdown because "free on <date>" decides between waiting
and paying. The numeric threshold has no place in that model: every decision is
categorical (wait / pay / skip), so the setting, the comparison and the whole
alert-state machine are gone.

- both alert-state columns are removed from the schema rather than left dead; a
  database created by an unreleased build has them dropped on open (native
  ALTER TABLE ... DROP COLUMN, guarded), which is a no-op for everyone else
- gate events are emitted only for versions the user does not have, and the
  price-drop event goes with the threshold it belonged to
- both alert endpoints, PriceAlertsHandler and the service-registry adapter field
  it needed are removed: events already reach the UI through the refresh response
- the bell tab, panel, CSS, both entry points, the unread watermark and their
  locales are removed; the setting keeps only the enable flag and the refresh
  interval and is framed as plumbing
- "Price unavailable" is replaced by `Paid`: the gate is certain from the public
  API, only the number is best-effort, and that is our plumbing, not the user's
  problem

Verified against a copy of the owner's real database: 52 gated versions ->
28 owned (now silent) + 24 the feature is actually about; the drop migration ran
and both removed endpoints 404.
2026-10-05 08:54:08 +08:00

1195 lines
37 KiB
Python

import copy
import json
import logging
from types import SimpleNamespace
from typing import Any
from unittest.mock import AsyncMock
import pytest
from py.config import config
from py.routes.handlers.model_handlers import (
ModelCivitaiHandler,
ModelHandlerSet,
ModelManagementHandler,
ModelUpdateHandler,
)
from py.routes.model_route_registrar import COMMON_ROUTE_DEFINITIONS
from py.services.service_registry import ServiceRegistry
from py.utils.metadata_manager import MetadataManager
from py.services.model_update_service import ModelUpdateRecord, ModelVersionRecord
class DummyScanner:
def __init__(self, cache):
self._cache = cache
self._cancelled = False
def is_cancelled(self) -> bool:
return self._cancelled
def reset_cancellation(self) -> None:
self._cancelled = False
async def get_cached_data(self):
return self._cache
class DummyService:
def __init__(self, cache):
self.model_type = "lora"
self.scanner = DummyScanner(cache)
class DummyUpdateService:
def __init__(self, records):
self.records = records
self.calls = []
async def refresh_for_model_type(
self,
model_type,
scanner,
provider,
*,
force_refresh=False,
target_model_ids=None,
folder_path=None,
):
self.calls.append(
{
"model_type": model_type,
"scanner": scanner,
"provider": provider,
"force_refresh": force_refresh,
"target_model_ids": target_model_ids,
"folder_path": folder_path,
}
)
return self.records
@pytest.mark.asyncio
async def test_build_version_context_includes_static_urls():
cache = SimpleNamespace(version_index={123: {"preview_url": "/tmp/previews/example.png"}})
service = DummyService(cache)
handler = ModelUpdateHandler(
service=service,
update_service=SimpleNamespace(),
metadata_provider_selector=lambda *_: None,
settings_service=SimpleNamespace(get=lambda *_: False),
logger=logging.getLogger(__name__),
)
record = ModelUpdateRecord(
model_type="lora",
model_id=42,
versions=[
ModelVersionRecord(
version_id=123,
name=None,
base_model=None,
released_at=None,
size_bytes=None,
preview_url=None,
is_in_library=True,
should_ignore=False,
)
],
last_checked_at=None,
should_ignore_model=False,
)
overrides = await handler._build_version_context(record)
expected = config.get_preview_static_url("/tmp/previews/example.png")
assert overrides == {
123: {
"file_path": None,
"file_name": None,
"preview_override": expected,
"has_been_downloaded": False,
}
}
@pytest.mark.asyncio
async def test_build_version_context_includes_download_history(monkeypatch):
cache = SimpleNamespace(version_index={})
service = DummyService(cache)
handler = ModelUpdateHandler(
service=service,
update_service=SimpleNamespace(),
metadata_provider_selector=lambda *_: None,
settings_service=SimpleNamespace(get=lambda *_: False),
logger=logging.getLogger(__name__),
)
class DummyHistoryService:
async def get_downloaded_version_ids(self, model_type, model_id):
assert model_type == "lora"
assert model_id == 42
return [123]
async def fake_history_service_factory():
return DummyHistoryService()
monkeypatch.setattr(
ServiceRegistry,
"get_downloaded_version_history_service",
staticmethod(fake_history_service_factory),
)
record = ModelUpdateRecord(
model_type="lora",
model_id=42,
versions=[
ModelVersionRecord(
version_id=123,
name="Downloaded",
base_model=None,
released_at=None,
size_bytes=None,
preview_url=None,
is_in_library=False,
should_ignore=False,
),
ModelVersionRecord(
version_id=124,
name="Fresh",
base_model=None,
released_at=None,
size_bytes=None,
preview_url=None,
is_in_library=False,
should_ignore=False,
),
],
last_checked_at=None,
should_ignore_model=False,
)
overrides = await handler._build_version_context(record)
assert overrides[123]["has_been_downloaded"] is True
assert overrides[124]["has_been_downloaded"] is False
@pytest.mark.asyncio
async def test_get_civitai_versions_degrades_when_download_history_unavailable(monkeypatch):
cache = SimpleNamespace(version_index={})
service = DummyService(cache)
class DummyProvider:
async def get_model_versions(self, model_id):
assert model_id == "42"
return {
"type": "lora",
"modelVersions": [
{
"id": 7,
"name": "Version 7",
"files": [],
}
],
}
async def fake_history_service_factory():
raise RuntimeError("download history unavailable")
monkeypatch.setattr(
ServiceRegistry,
"get_downloaded_version_history_service",
staticmethod(fake_history_service_factory),
)
async def metadata_provider_factory():
return DummyProvider()
handler = ModelCivitaiHandler(
service=service,
settings_service=SimpleNamespace(get=lambda *_: False), # pyright: ignore[reportArgumentType]
ws_manager=SimpleNamespace(), # pyright: ignore[reportArgumentType]
logger=logging.getLogger(__name__),
metadata_provider_factory=metadata_provider_factory,
validate_model_type=lambda *_: True,
expected_model_types=lambda: "LoRA",
find_model_file=lambda *_: None,
metadata_sync=SimpleNamespace(), # pyright: ignore[reportArgumentType]
metadata_refresh_use_case=SimpleNamespace(), # pyright: ignore[reportArgumentType]
metadata_progress_callback=lambda *_args, **_kwargs: None, # pyright: ignore[reportArgumentType]
)
response = await handler.get_civitai_versions(
SimpleNamespace(match_info={"model_id": "42"}) # pyright: ignore[reportArgumentType]
)
text = response.text
assert text is not None
payload = json.loads(text)
assert response.status == 200
assert payload[0]["id"] == 7
assert payload[0]["existsLocally"] is False
assert payload[0]["hasBeenDownloaded"] is False
@pytest.mark.asyncio
async def test_refresh_model_updates_filters_records_without_updates():
cache = SimpleNamespace(version_index={})
service = DummyService(cache)
record_with_update = ModelUpdateRecord(
model_type="lora",
model_id=1,
versions=[
ModelVersionRecord(
version_id=8,
name="v0",
base_model="Pony",
released_at=None,
size_bytes=None,
preview_url=None,
is_in_library=True,
should_ignore=False,
),
ModelVersionRecord(
version_id=10,
name="v1",
base_model="Pony",
released_at=None,
size_bytes=None,
preview_url=None,
is_in_library=False,
should_ignore=False,
),
],
last_checked_at=None,
should_ignore_model=False,
)
record_without_update = ModelUpdateRecord(
model_type="lora",
model_id=2,
versions=[
ModelVersionRecord(
version_id=20,
name="v2",
base_model=None,
released_at=None,
size_bytes=None,
preview_url=None,
is_in_library=True,
should_ignore=False,
)
],
last_checked_at=None,
should_ignore_model=False,
)
update_service = DummyUpdateService({1: record_with_update, 2: record_without_update})
async def metadata_selector(name):
assert name == "civitai_api"
return object()
handler = ModelUpdateHandler(
service=service,
update_service=update_service,
metadata_provider_selector=metadata_selector,
settings_service=SimpleNamespace(get=lambda *_: False),
logger=logging.getLogger(__name__),
)
class DummyRequest:
can_read_body = True
query = {}
async def json(self):
return {}
response = await handler.refresh_model_updates(
DummyRequest() # pyright: ignore[reportArgumentType]
)
assert response.status == 200
text = response.text
assert text is not None
payload = json.loads(text)
assert payload["success"] is True
assert len(payload["records"]) == 1
assert payload["records"][0]["modelId"] == 1
assert payload["records"][0]["hasUpdate"] is True
assert len(update_service.calls) == 1
call = update_service.calls[0]
assert call["model_type"] == "lora"
assert call["scanner"] is service.scanner
assert call["force_refresh"] is False
assert call["provider"] is not None
assert call["target_model_ids"] is None
@pytest.mark.asyncio
async def test_refresh_model_updates_reports_gate_events_for_all_records():
"""Gate transitions are reported even for records that do not qualify as
updates (a version already in the library that became free)."""
cache = SimpleNamespace(version_index={})
service = DummyService(cache)
record = ModelUpdateRecord(
model_type="lora",
model_id=1,
versions=[
ModelVersionRecord(
version_id=11,
name="v11",
base_model=None,
released_at=None,
size_bytes=None,
preview_url=None,
is_in_library=True,
should_ignore=False,
)
],
last_checked_at=None,
should_ignore_model=False,
events=[{"versionId": 11, "kind": "became_free", "versionName": "v11", "isInLibrary": True}],
)
update_service = DummyUpdateService({1: record})
metadata_selector = AsyncMock(return_value=SimpleNamespace())
handler = ModelUpdateHandler(
service=service,
update_service=update_service,
metadata_provider_selector=metadata_selector,
settings_service=SimpleNamespace(get=lambda *_: False),
logger=logging.getLogger(__name__),
)
class DummyRequest:
can_read_body = True
query = {}
async def json(self):
return {}
response = await handler.refresh_model_updates(
DummyRequest() # pyright: ignore[reportArgumentType]
)
payload = json.loads(response.text)
# The record itself does not qualify as an update...
assert payload["records"] == []
# ...but its transition is still surfaced.
assert payload["events"] == [
{
"modelId": 1,
"modelType": "lora",
"versionId": 11,
"kind": "became_free",
"versionName": "v11",
"isInLibrary": True,
}
]
@pytest.mark.asyncio
async def test_refresh_model_updates_same_base_scope_excludes_cross_base_updates():
"""Issue #1083: with version_grouping=same_base (the default), a newer
remote version targeting another base model must not be counted, matching
what the Updates filter displays."""
cache = SimpleNamespace(version_index={})
service = DummyService(cache)
cross_base_only = ModelUpdateRecord(
model_type="lora",
model_id=1,
versions=[
ModelVersionRecord(
version_id=5,
name="v0",
base_model="Pony",
released_at=None,
size_bytes=None,
preview_url=None,
is_in_library=True,
should_ignore=False,
),
ModelVersionRecord(
version_id=20,
name="v2",
base_model="Flux.1",
released_at=None,
size_bytes=None,
preview_url=None,
is_in_library=False,
should_ignore=False,
),
],
last_checked_at=None,
should_ignore_model=False,
)
update_service = DummyUpdateService({1: cross_base_only})
async def metadata_selector(name):
assert name == "civitai_api"
return object()
handler = ModelUpdateHandler(
service=service,
update_service=update_service,
metadata_provider_selector=metadata_selector,
settings_service=SimpleNamespace(get=lambda *_: False),
logger=logging.getLogger(__name__),
)
class DummyRequest:
can_read_body = True
query = {}
async def json(self):
return {}
response = await handler.refresh_model_updates(
DummyRequest() # pyright: ignore[reportArgumentType]
)
assert response.status == 200
text = response.text
assert text is not None
payload = json.loads(text)
assert payload["success"] is True
assert payload["records"] == []
@pytest.mark.asyncio
async def test_refresh_model_updates_any_grouping_counts_cross_base_updates():
"""With version_grouping=any the unscoped predicate applies, so a newer
remote version on any base model is counted."""
cache = SimpleNamespace(version_index={})
service = DummyService(cache)
cross_base_only = ModelUpdateRecord(
model_type="lora",
model_id=1,
versions=[
ModelVersionRecord(
version_id=5,
name="v0",
base_model="Pony",
released_at=None,
size_bytes=None,
preview_url=None,
is_in_library=True,
should_ignore=False,
),
ModelVersionRecord(
version_id=20,
name="v2",
base_model="Flux.1",
released_at=None,
size_bytes=None,
preview_url=None,
is_in_library=False,
should_ignore=False,
),
],
last_checked_at=None,
should_ignore_model=False,
)
update_service = DummyUpdateService({1: cross_base_only})
async def metadata_selector(name):
assert name == "civitai_api"
return object()
settings = SimpleNamespace(
get=lambda key, default=None: (
"any" if key == "version_grouping" else default
)
)
handler = ModelUpdateHandler(
service=service,
update_service=update_service,
metadata_provider_selector=metadata_selector,
settings_service=settings,
logger=logging.getLogger(__name__),
)
class DummyRequest:
can_read_body = True
query = {}
async def json(self):
return {}
response = await handler.refresh_model_updates(
DummyRequest() # pyright: ignore[reportArgumentType]
)
assert response.status == 200
text = response.text
assert text is not None
payload = json.loads(text)
assert payload["success"] is True
assert [record["modelId"] for record in payload["records"]] == [1]
def _make_cross_base_only_record() -> ModelUpdateRecord:
return ModelUpdateRecord(
model_type="lora",
model_id=1,
versions=[
ModelVersionRecord(
version_id=5,
name="v0",
base_model="Pony",
released_at=None,
size_bytes=None,
preview_url=None,
is_in_library=True,
should_ignore=False,
),
ModelVersionRecord(
version_id=20,
name="v2",
base_model="Flux.1",
released_at=None,
size_bytes=None,
preview_url=None,
is_in_library=False,
should_ignore=False,
),
],
last_checked_at=None,
should_ignore_model=False,
)
def _base_handler(update_service, settings_service):
async def metadata_selector(name):
assert name == "civitai_api"
return object()
return ModelUpdateHandler(
service=DummyService(SimpleNamespace(version_index={})),
update_service=update_service,
metadata_provider_selector=metadata_selector,
settings_service=settings_service,
logger=logging.getLogger(__name__),
)
@pytest.mark.asyncio
async def test_refresh_model_updates_explicit_same_base_setting_excludes_cross_base():
"""The literal "same_base" string (any casing/whitespace) selects the
scoped predicate, mirroring BaseModelService's strategy parsing."""
update_service = DummyUpdateService({1: _make_cross_base_only_record()})
settings = SimpleNamespace(
get=lambda key, default=None: (
" Same_Base " if key == "version_grouping" else default
)
)
handler = _base_handler(update_service, settings)
class DummyRequest:
can_read_body = True
query = {}
async def json(self):
return {}
response = await handler.refresh_model_updates(
DummyRequest() # pyright: ignore[reportArgumentType]
)
text = response.text
assert text is not None
payload = json.loads(text)
assert payload["records"] == []
@pytest.mark.asyncio
async def test_refresh_model_updates_falls_back_without_scoped_predicate():
"""A record type without has_update_for_local_bases (pre-change callers /
fakes) still counts via the unscoped predicate under same_base scope."""
legacy_record = _make_cross_base_only_record()
legacy_record.__dict__["has_update_for_local_bases"] = None
update_service = DummyUpdateService({1: legacy_record})
handler = _base_handler(update_service, SimpleNamespace(get=lambda *_: False))
class DummyRequest:
can_read_body = True
query = {}
async def json(self):
return {}
response = await handler.refresh_model_updates(
DummyRequest() # pyright: ignore[reportArgumentType]
)
text = response.text
assert text is not None
payload = json.loads(text)
assert payload["success"] is True
assert [record["modelId"] for record in payload["records"]] == [1]
@pytest.mark.asyncio
async def test_refresh_model_updates_with_target_ids():
cache = SimpleNamespace(version_index={})
service = DummyService(cache)
record_with_update = ModelUpdateRecord(
model_type="lora",
model_id=1,
versions=[
ModelVersionRecord(
version_id=10,
name="v1",
base_model=None,
released_at=None,
size_bytes=None,
preview_url=None,
is_in_library=False,
should_ignore=False,
)
],
last_checked_at=None,
should_ignore_model=False,
)
update_service = DummyUpdateService({1: record_with_update})
async def metadata_selector(name):
assert name == "civitai_api"
return object()
handler = ModelUpdateHandler(
service=service,
update_service=update_service,
metadata_provider_selector=metadata_selector,
settings_service=SimpleNamespace(get=lambda *_: False),
logger=logging.getLogger(__name__),
)
class DummyRequest:
can_read_body = True
query = {}
async def json(self):
return {"modelIds": [1, "2", None]}
response = await handler.refresh_model_updates(
DummyRequest() # pyright: ignore[reportArgumentType]
)
assert response.status == 200
call = update_service.calls[0]
assert call["target_model_ids"] == [1, 2]
@pytest.mark.asyncio
async def test_refresh_model_updates_accepts_snake_case_ids():
cache = SimpleNamespace(version_index={})
service = DummyService(cache)
record_with_update = ModelUpdateRecord(
model_type="lora",
model_id=3,
versions=[
ModelVersionRecord(
version_id=30,
name="v3",
base_model=None,
released_at=None,
size_bytes=None,
preview_url=None,
is_in_library=False,
should_ignore=False,
)
],
last_checked_at=None,
should_ignore_model=False,
)
update_service = DummyUpdateService({3: record_with_update})
async def metadata_selector(name):
assert name == "civitai_api"
return object()
handler = ModelUpdateHandler(
service=service,
update_service=update_service,
metadata_provider_selector=metadata_selector,
settings_service=SimpleNamespace(get=lambda *_: False),
logger=logging.getLogger(__name__),
)
class DummyRequest:
can_read_body = True
query = {}
async def json(self):
return {"model_ids": [3, "4", "abc", None]}
response = await handler.refresh_model_updates(
DummyRequest() # pyright: ignore[reportArgumentType]
)
assert response.status == 200
call = update_service.calls[0]
assert call["target_model_ids"] == [3, 4]
@pytest.mark.asyncio
async def test_fetch_missing_license_data_updates_metadata(monkeypatch):
cache = SimpleNamespace(
raw_data=[
{"file_path": "/tmp/model1.safetensors", "civitai": {"modelId": 10}},
{"file_path": "/tmp/model2.safetensors", "civitai": {"modelId": 10}},
{"file_path": "/tmp/model3.safetensors", "civitai": {"modelId": 20}},
],
version_index={},
)
metadata_store = {
"/tmp/model1.safetensors": {"civitai": {"model": {}}},
"/tmp/model2.safetensors": {"civitai": {"model": {}}},
"/tmp/model3.safetensors": {"civitai": {"model": {}}},
}
async def fake_load(path: str):
data = metadata_store.get(path)
if data is None:
return None, False
return SimpleNamespace(to_dict=lambda: copy.deepcopy(data)), False
saved: list[tuple[str, dict[str, Any]]] = []
async def fake_save(path: str, metadata: dict[str, Any]):
saved.append((path, copy.deepcopy(metadata)))
return True
monkeypatch.setattr(MetadataManager, "load_metadata", staticmethod(fake_load))
monkeypatch.setattr(MetadataManager, "save_metadata", staticmethod(fake_save))
provider_calls: list[list[int]] = []
async def fake_bulk(model_ids):
provider_calls.append(list(model_ids))
return {
10: {
"allowNoCredit": True,
"allowCommercialUse": ["Sell"],
"allowDerivatives": True,
"allowDifferentLicense": True,
},
20: {
"allowNoCredit": False,
"allowCommercialUse": ["Image"],
"allowDerivatives": False,
"allowDifferentLicense": False,
},
}
provider = SimpleNamespace()
provider.get_model_versions_bulk = fake_bulk
async def metadata_selector(name):
assert name == "civitai_api"
return provider
handler = ModelUpdateHandler(
service=DummyService(cache),
update_service=SimpleNamespace(),
metadata_provider_selector=metadata_selector,
settings_service=SimpleNamespace(get=lambda *_: False),
logger=logging.getLogger(__name__),
)
class DummyRequest:
can_read_body = True
query = {}
async def json(self):
return {}
response = await handler.fetch_missing_civitai_license_data(
DummyRequest() # pyright: ignore[reportArgumentType]
)
assert response.status == 200
text = response.text
assert text is not None
payload = json.loads(text)
assert payload["success"] is True
assert len(payload["updated"]) == 3
assert provider_calls == [[10, 20]]
assert len(saved) == 3
first_metadata = saved[0][1]
assert first_metadata["civitai"]["model"]["allowNoCredit"] is True
assert first_metadata["civitai"]["model"]["allowCommercialUse"] == ["Sell"]
assert "missingModelIds" not in payload
assert "errors" not in payload
@pytest.mark.asyncio
async def test_fetch_missing_license_data_filters_model_ids(monkeypatch):
cache = SimpleNamespace(
raw_data=[
{"file_path": "/tmp/model1.safetensors", "civitai": {"modelId": 10}},
{"file_path": "/tmp/model2.safetensors", "civitai": {"modelId": 20}},
],
version_index={},
)
metadata_store = {
"/tmp/model1.safetensors": {"civitai": {"model": {}}},
"/tmp/model2.safetensors": {"civitai": {"model": {}}},
}
async def fake_load(path: str):
data = metadata_store.get(path)
if data is None:
return None, False
return SimpleNamespace(to_dict=lambda: copy.deepcopy(data)), False
saved: list[tuple[str, dict[str, Any]]] = []
async def fake_save(path: str, metadata: dict[str, Any]):
saved.append((path, copy.deepcopy(metadata)))
return True
monkeypatch.setattr(MetadataManager, "load_metadata", staticmethod(fake_load))
monkeypatch.setattr(MetadataManager, "save_metadata", staticmethod(fake_save))
provider_calls: list[list[int]] = []
async def fake_bulk(model_ids):
provider_calls.append(list(model_ids))
return {
10: {
"allowNoCredit": True,
"allowCommercialUse": ["Sell"],
"allowDerivatives": True,
"allowDifferentLicense": True,
},
20: {
"allowNoCredit": False,
"allowCommercialUse": ["Image"],
"allowDerivatives": False,
"allowDifferentLicense": False,
},
}
provider = SimpleNamespace()
provider.get_model_versions_bulk = fake_bulk
async def metadata_selector(name):
assert name == "civitai_api"
return provider
handler = ModelUpdateHandler(
service=DummyService(cache),
update_service=SimpleNamespace(),
metadata_provider_selector=metadata_selector,
settings_service=SimpleNamespace(get=lambda *_: False),
logger=logging.getLogger(__name__),
)
class DummyRequest:
can_read_body = True
query = {}
async def json(self):
return {"modelIds": [20]}
response = await handler.fetch_missing_civitai_license_data(
DummyRequest() # pyright: ignore[reportArgumentType]
)
assert response.status == 200
text = response.text
assert text is not None
payload = json.loads(text)
assert payload["success"] is True
assert len(payload["updated"]) == 1
assert provider_calls == [[20]]
assert len(saved) == 1
def test_serialize_version_permanent_paid_is_not_early_access():
"""Permanent paid versions (is_paid, no end date) must not be flagged as
early access, mirroring _is_early_access_active in the update service."""
version = ModelVersionRecord(
version_id=7, name="v7", base_model=None, released_at=None, size_bytes=None,
preview_url=None, is_in_library=False, should_ignore=False,
early_access_ends_at=None, is_early_access=True, usage_control="Download",
paid_access=json.dumps({"permanent": True, "endsAt": None}), is_paid=True,
)
serialized = ModelUpdateHandler._serialize_version(version, None)
assert serialized["isEarlyAccess"] is False
assert serialized["isPaid"] is True
assert serialized["paidAccess"] == {"permanent": True, "endsAt": None}
def test_serialize_version_timed_paid_is_early_access():
"""Timed paid gates (endsAt in the future) stay flagged as early access."""
version = ModelVersionRecord(
version_id=8, name="v8", base_model=None, released_at=None, size_bytes=None,
preview_url=None, is_in_library=False, should_ignore=False,
early_access_ends_at="2099-01-01T00:00:00.000Z", is_early_access=True,
usage_control="Download",
paid_access=json.dumps({"permanent": False, "endsAt": "2099-01-01T00:00:00.000Z"}),
is_paid=False,
)
serialized = ModelUpdateHandler._serialize_version(version, None)
assert serialized["isEarlyAccess"] is True
assert serialized["isPaid"] is False
def test_serialize_version_malformed_paid_access_does_not_crash():
"""A malformed paid_access row must degrade to None instead of failing
the whole versions-list response."""
version = ModelVersionRecord(
version_id=10, name="v10", base_model=None, released_at=None, size_bytes=None,
preview_url=None, is_in_library=False, should_ignore=False,
early_access_ends_at=None, is_early_access=True, usage_control=None,
paid_access="{not json", is_paid=False,
)
serialized = ModelUpdateHandler._serialize_version(version, None)
assert serialized["paidAccess"] is None
assert serialized["isEarlyAccess"] is True
async def test_enrich_early_access_details_skips_permanent_paid(monkeypatch):
"""Permanent paid versions must not trigger per-version CivitAI fetches in
_enrich_early_access_details: they are not early access and can never get
an end time, so enriching them is wasted API traffic."""
record = ModelUpdateRecord(
model_type="lora",
model_id=1,
versions=[
ModelVersionRecord(
version_id=100, name="paid", base_model=None, released_at=None,
size_bytes=None, preview_url=None, is_in_library=False,
should_ignore=False, early_access_ends_at=None,
is_early_access=True, usage_control="Download",
paid_access='{"permanent": true, "endsAt": null}', is_paid=True,
),
ModelVersionRecord(
version_id=200, name="ea", base_model=None, released_at=None,
size_bytes=None, preview_url=None, is_in_library=False,
should_ignore=False, early_access_ends_at=None,
is_early_access=True, usage_control="Download",
paid_access=None, is_paid=False,
),
],
last_checked_at=1.0,
should_ignore_model=False,
)
fetched: list[int] = []
async def fake_version_info(version_id: str):
fetched.append(int(version_id))
return {"earlyAccessEndsAt": "2099-01-01T00:00:00.000Z"}, None
provider = SimpleNamespace(get_model_version_info=fake_version_info)
async def metadata_selector(name):
assert name == "civitai_api"
return provider
handler = ModelUpdateHandler(
service=DummyService(SimpleNamespace(raw_data=[], version_index={})),
update_service=SimpleNamespace(),
metadata_provider_selector=metadata_selector,
settings_service=SimpleNamespace(get=lambda *_: False),
logger=logging.getLogger(__name__),
)
enriched = await handler._enrich_early_access_details(record)
# Only the timed EA version (200) is fetched; the permanent paid one (100) is skipped.
assert fetched == [200]
enriched_map = {v.version_id: v for v in enriched.versions}
assert enriched_map[200].early_access_ends_at == "2099-01-01T00:00:00.000Z"
assert enriched_map[100].early_access_ends_at is None
def test_serialize_version_includes_file_count():
version = ModelVersionRecord(
version_id=11, name="v11", base_model=None, released_at=None, size_bytes=None,
preview_url=None, is_in_library=True, should_ignore=False, file_count=2,
)
serialized = ModelUpdateHandler._serialize_version(version, None)
assert serialized["fileCount"] == 2
def test_serialize_version_file_count_defaults_to_none():
version = ModelVersionRecord(
version_id=12, name="v12", base_model=None, released_at=None, size_bytes=None,
preview_url=None, is_in_library=False, should_ignore=False,
)
serialized = ModelUpdateHandler._serialize_version(version, None)
assert serialized["fileCount"] is None
def _build_relink_handler(metadata_sync):
service = SimpleNamespace(
scanner=SimpleNamespace(update_single_model_cache=AsyncMock())
)
return ModelManagementHandler(
service=service,
logger=logging.getLogger(__name__),
metadata_sync=metadata_sync,
preview_service=SimpleNamespace(),
tag_update_service=SimpleNamespace(),
lifecycle_service=SimpleNamespace(),
)
@pytest.mark.asyncio
async def test_relink_civitai_rejects_unsupported_source():
metadata_sync = SimpleNamespace(
load_local_metadata=AsyncMock(return_value={}),
relink_metadata=AsyncMock(),
)
handler = _build_relink_handler(metadata_sync)
request = SimpleNamespace(
json=AsyncMock(
return_value={
"file_path": "/tmp/model.safetensors",
"model_id": "123",
"model_version_id": "456",
"source": "huggingface",
}
)
)
response = await handler.relink_civitai(request)
assert response.status == 400
payload = json.loads(response.text)
assert payload["success"] is False
assert "Unsupported relink source" in payload["error"]
metadata_sync.relink_metadata.assert_not_awaited()
@pytest.mark.asyncio
async def test_relink_civitai_passes_provider_name_for_civarchive_source():
metadata_sync = SimpleNamespace(
load_local_metadata=AsyncMock(return_value={"model_name": "Local"}),
relink_metadata=AsyncMock(
return_value={"model_name": "Archived", "sha256": "abc"}
),
)
handler = _build_relink_handler(metadata_sync)
request = SimpleNamespace(
json=AsyncMock(
return_value={
"file_path": "/tmp/model.safetensors",
"model_id": "123",
"model_version_id": "456",
"source": "civarchive",
}
)
)
response = await handler.relink_civitai(request)
assert response.status == 200
payload = json.loads(response.text)
assert payload["success"] is True
assert "CivArchive" in payload["message"]
metadata_sync.relink_metadata.assert_awaited_once_with(
file_path="/tmp/model.safetensors",
metadata={"model_name": "Local"},
model_id=123,
model_version_id=456,
provider_name="civarchive_api",
)
@pytest.mark.asyncio
async def test_relink_civitai_surfaces_provider_unavailable_without_500():
metadata_sync = SimpleNamespace(
load_local_metadata=AsyncMock(return_value={}),
relink_metadata=AsyncMock(
side_effect=ValueError(
"CivitArchive is not available or not enabled. "
"Enable the CivitArchive API in settings to relink via CivArchive."
)
),
)
handler = _build_relink_handler(metadata_sync)
request = SimpleNamespace(
json=AsyncMock(
return_value={
"file_path": "/tmp/model.safetensors",
"model_id": "123",
"model_version_id": None,
"source": "civarchive",
}
)
)
response = await handler.relink_civitai(request)
assert response.status == 400
payload = json.loads(response.text)
assert payload["success"] is False
assert "CivitArchive" in payload["error"]
def test_every_common_route_definition_resolves_to_a_handler():
"""Guard the declarative route table against a missing mapping entry.
Adding a RouteDefinition without registering it in
``ModelHandlerSet.to_route_mapping`` fails at request time with a bare
KeyError from the handler lookup (and only on a live server), so assert the
whole table resolves here instead.
"""
class AnyHandler:
def __getattr__(self, _name):
return lambda request: None
handler_set = ModelHandlerSet(
page_view=AnyHandler(),
listing=AnyHandler(),
management=AnyHandler(),
query=AnyHandler(),
download=AnyHandler(),
civitai=AnyHandler(),
move=AnyHandler(),
auto_organize=AnyHandler(),
filename_template=AnyHandler(),
updates=AnyHandler(),
)
mapping = handler_set.to_route_mapping()
missing = [
definition.handler_name
for definition in COMMON_ROUTE_DEFINITIONS
if definition.handler_name not in mapping
]
assert missing == []
def test_price_alert_endpoints_are_not_registered():
"""The redesign dropped the standalone alerts surface: obtainability rides on
the update surfaces, so no per-type alert route may remain."""
assert [
definition.path_template
for definition in COMMON_ROUTE_DEFINITIONS
if "price-alert" in definition.path_template
] == []