Files
ComfyUI-Lora-Manager/tests/services/test_model_update_service.py
T
Will Miao 156d2e5eb9 fix(ui): explain a 0 Buzz threshold in the price alerts panel
The default threshold is 0 ("only tell me when a version becomes free") and the
settings copy says so, but the panel did not: a real instance with 52 priced paid
versions and an untouched threshold showed "Nothing is under your price threshold
right now" with only a small "Alert threshold: 0 Buzz" in the corner, which reads
as a broken feature.

- the payload now carries pricedCount, so the empty state can say how many paid
  versions already have a known price
- when the threshold is 0 and prices are known, the empty state says so and
  points at Settings - Library instead of implying there is nothing to show
- the read-time threshold comparison means setting one takes effect immediately;
  measured on a copy of that instance: 0 Buzz -> 0 alerts, 100 -> 44,
  500 -> 48, 5000 -> 51
2026-10-04 20:28:57 +08:00

2054 lines
68 KiB
Python

import logging
import sqlite3
from dataclasses import replace
from types import SimpleNamespace
import pytest
from py.services.errors import ResourceNotFoundError
from py.services.model_update_service import (
ModelUpdateRecord,
ModelUpdateService,
ModelVersionRecord,
)
class DummyScanner:
def __init__(self, raw_data):
self._cache = SimpleNamespace(raw_data=raw_data, version_index={})
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, *args, **kwargs):
return self._cache
class DummyProvider:
def __init__(self, response, *, support_bulk: bool = True):
self.response = response
self.calls: int = 0
self.bulk_calls: list[list[int]] = []
self.support_bulk = support_bulk
async def get_model_versions(self, model_id):
self.calls += 1
return self.response
async def get_model_versions_bulk(self, model_ids):
if not self.support_bulk:
raise NotImplementedError
self.bulk_calls.append(list(model_ids))
return {model_id: self.response for model_id in model_ids}
class NotFoundProvider:
def __init__(self):
self.calls = 0
self.bulk_calls: list[list[int]] = []
async def get_model_versions(self, model_id):
self.calls += 1
raise ResourceNotFoundError("Resource not found")
async def get_model_versions_bulk(self, model_ids):
self.bulk_calls.append(list(model_ids))
return {}
def make_version(
version_id,
*,
in_library,
base_model=None,
should_ignore=False,
early_access_ends_at=None,
is_early_access=False,
is_paid=False,
paid_access=None,
):
return ModelVersionRecord(
version_id=version_id,
name=None,
base_model=base_model,
released_at=None,
size_bytes=None,
preview_url=None,
is_in_library=in_library,
should_ignore=should_ignore,
early_access_ends_at=early_access_ends_at,
is_early_access=is_early_access,
is_paid=is_paid,
paid_access=paid_access,
)
def make_record(*versions, should_ignore_model=False):
return ModelUpdateRecord(
model_type="lora",
model_id=999,
versions=list(versions),
last_checked_at=None,
should_ignore_model=should_ignore_model,
)
def test_extract_size_bytes_prefers_primary_model_file(tmp_path):
db_path = tmp_path / "updates.sqlite"
service = ModelUpdateService(str(db_path))
response = {
"modelVersions": [
{
"id": 42,
"files": [
{"sizeKB": 2018.0400390625, "type": "Training Data", "primary": False},
{
"sizeKB": 1152322.3515625,
"type": "Model",
"primary": "True",
},
],
"images": [],
}
]
}
versions = service._extract_versions(response)
assert versions is not None
assert versions[0].size_bytes == int(1152322.3515625 * 1024)
def test_extract_size_bytes_falls_back_without_primary(tmp_path):
db_path = tmp_path / "updates.sqlite"
service = ModelUpdateService(str(db_path))
response = {
"modelVersions": [
{
"id": 43,
"files": [
{
"sizeKB": 2048,
"type": "Training Data",
"primary": True,
},
{"sizeKB": 1024, "type": "Archive", "primary": False},
],
"images": [],
}
]
}
versions = service._extract_versions(response)
assert versions is not None
assert versions[0].size_bytes == int(2048 * 1024)
def test_has_update_requires_newer_version_than_library():
record = make_record(
make_version(5, in_library=True),
make_version(4, in_library=False),
make_version(8, in_library=False, should_ignore=True),
)
assert record.has_update() is False
def test_has_update_detects_newer_remote_version():
record = make_record(
make_version(5, in_library=True),
make_version(7, in_library=False),
make_version(6, in_library=False, should_ignore=True),
)
assert record.has_update() is True
def test_has_update_for_base_matches_same_base_model():
record = make_record(
make_version(5, in_library=True, base_model="Pony"),
make_version(6, in_library=False, base_model="Pony"),
make_version(7, in_library=False, base_model="Flux.1"),
)
assert record.has_update_for_base(5, "Pony") is True
def test_has_update_for_base_rejects_other_base_models():
record = make_record(
make_version(10, in_library=True, base_model="Flux"),
make_version(20, in_library=False, base_model="SDXL"),
)
assert record.has_update_for_base(10, "Flux") is False
def test_has_update_for_local_bases_detects_same_base_newer_version():
record = make_record(
make_version(5, in_library=True, base_model="Pony"),
make_version(6, in_library=False, base_model="pony"),
)
assert record.has_update_for_local_bases() is True
def test_has_update_for_local_bases_rejects_cross_base_only_update():
"""Issue #1083: a newer remote version targeting another base model must
not count when the report is scoped like the Updates filter."""
record = make_record(
make_version(5, in_library=True, base_model="Pony"),
make_version(6, in_library=False, base_model="Flux.1"),
)
assert record.has_update_for_local_bases() is False
assert record.has_update() is True
def test_has_update_for_local_bases_hits_when_any_scope_qualifies():
record = make_record(
make_version(5, in_library=True, base_model="Pony"),
make_version(6, in_library=False, base_model="Flux.1"),
make_version(7, in_library=False, base_model="Pony"),
)
assert record.has_update_for_local_bases() is True
def test_has_update_for_local_bases_respects_ignore_and_hides():
ignored = make_record(
make_version(5, in_library=True, base_model="Pony"),
make_version(6, in_library=False, base_model="Pony", should_ignore=True),
)
assert ignored.has_update_for_local_bases() is False
paid = make_record(
make_version(5, in_library=True, base_model="Pony"),
make_version(
6,
in_library=False,
base_model="Pony",
is_paid=True,
paid_access='{"permanent": true, "endsAt": null}',
),
)
assert paid.has_update_for_local_bases() is True
assert paid.has_update_for_local_bases(hide_paid=True) is False
timed_early_access = make_record(
make_version(5, in_library=True, base_model="Pony"),
make_version(
6,
in_library=False,
base_model="Pony",
early_access_ends_at="2099-01-01T00:00:00Z",
is_early_access=True,
),
)
assert timed_early_access.has_update_for_local_bases() is True
assert timed_early_access.has_update_for_local_bases(hide_early_access=True) is False
def test_has_update_for_local_bases_falls_back_when_no_local_scopes_known():
"""Without any known in-library base (e.g. a local version delisted from
Civitai before its first refresh), the aggregate falls back to the
unscoped predicate so the summary cannot silently drop models the
item-level Updates filter may still flag from file metadata."""
delisted_local = make_record(
make_version(5, in_library=True, base_model=None),
make_version(6, in_library=False, base_model="Pony"),
)
assert delisted_local.has_update_for_local_bases() is True
remote_only = make_record(make_version(6, in_library=False, base_model="Pony"))
assert remote_only.has_update_for_local_bases() is True
def test_build_record_from_remote_synthesizes_base_from_local_map(tmp_path):
"""Versions missing from the remote listing are synthesized with the base
model collected from cache items, so same-base scoping survives a version
being delisted upstream."""
db_path = tmp_path / "updates.sqlite"
service = ModelUpdateService(str(db_path))
remote = [make_version(6, in_library=False, base_model="Pony")]
record = service._build_record_from_remote(
model_type="lora",
model_id=1,
local_versions=[5],
remote_versions=remote,
existing=None,
timestamp=1.0,
local_base_models={5: "Pony"},
)
v5 = next(v for v in record.versions if v.version_id == 5)
assert v5.base_model == "Pony"
record_without_map = service._build_record_from_remote(
model_type="lora",
model_id=1,
local_versions=[5],
remote_versions=remote,
existing=None,
timestamp=1.0,
)
v5_plain = next(v for v in record_without_map.versions if v.version_id == 5)
assert v5_plain.base_model is None
@pytest.mark.asyncio
async def test_refresh_synthesizes_base_model_from_cache_items(tmp_path):
"""End-to-end: a locally-held version absent from the remote listing keeps
its cache-item base model, keeping it countable under same-base scoping."""
db_path = tmp_path / "updates.sqlite"
service = ModelUpdateService(str(db_path), ttl_seconds=0)
raw_data = [{"civitai": {"modelId": 1, "id": 11}, "base_model": "Pony"}]
scanner = DummyScanner(raw_data)
provider = DummyProvider(
{
"modelVersions": [
{"id": 12, "baseModel": "Pony", "files": [], "images": []},
]
}
)
await service.refresh_for_model_type("lora", scanner, provider)
record = await service.get_record("lora", 1)
assert record is not None
v11 = next(v for v in record.versions if v.version_id == 11)
assert v11.base_model == "Pony"
assert record.has_update() is True
assert record.has_update_for_local_bases() is True
@pytest.mark.asyncio
async def test_refresh_persists_versions_and_uses_cache(tmp_path):
db_path = tmp_path / "updates.sqlite"
service = ModelUpdateService(str(db_path), ttl_seconds=3600)
raw_data = [
{"civitai": {"modelId": 1, "id": 11}},
{"civitai": {"modelId": 1, "id": 15}},
]
scanner = DummyScanner(raw_data)
provider = DummyProvider(
{
"modelVersions": [
{
"id": 11,
"name": "v1",
"baseModel": "SD15",
"publishedAt": "2024-01-01T00:00:00Z",
"files": [{"sizeKB": 1024}],
"images": [{"url": "https://example.com/1.png"}],
},
{
"id": 15,
"name": "v1.5",
"baseModel": "SD15",
"publishedAt": "2024-02-01T00:00:00Z",
"files": [{"sizeKB": 512}],
"images": [{"url": "https://example.com/2.png"}],
},
]
}
)
await service.refresh_for_model_type("lora", scanner, provider)
record = await service.get_record("lora", 1)
assert provider.calls == 0
assert provider.bulk_calls == [[1]]
assert record is not None
assert record.version_ids == [11, 15]
assert record.in_library_version_ids == [11, 15]
assert [version.name for version in record.versions] == ["v1", "v1.5"]
assert record.should_ignore_model is False
assert record.has_update() is False
await service.refresh_for_model_type("lora", scanner, provider)
assert provider.calls == 0, "provider should not be called again within TTL"
assert provider.bulk_calls == [[1]]
@pytest.mark.asyncio
async def test_refresh_filters_to_requested_models(tmp_path):
db_path = tmp_path / "updates.sqlite"
service = ModelUpdateService(str(db_path), ttl_seconds=3600)
raw_data = [
{"civitai": {"modelId": 1, "id": 11}},
{"civitai": {"modelId": 2, "id": 21}},
]
scanner = DummyScanner(raw_data)
provider = DummyProvider({"modelVersions": []})
result = await service.refresh_for_model_type(
"lora",
scanner,
provider,
target_model_ids=[2],
)
assert list(result.keys()) == [2]
assert provider.bulk_calls == [[2]]
@pytest.mark.asyncio
async def test_refresh_returns_empty_when_targets_missing(tmp_path):
db_path = tmp_path / "updates.sqlite"
service = ModelUpdateService(str(db_path), ttl_seconds=3600)
raw_data = [{"civitai": {"modelId": 1, "id": 11}}]
scanner = DummyScanner(raw_data)
provider = DummyProvider({"modelVersions": []})
result = await service.refresh_for_model_type(
"lora",
scanner,
provider,
target_model_ids=[5],
)
assert result == {}
assert provider.bulk_calls == []
@pytest.mark.asyncio
async def test_refresh_respects_ignore_flag(tmp_path):
db_path = tmp_path / "updates.sqlite"
service = ModelUpdateService(str(db_path), ttl_seconds=3600)
raw_data = [{"civitai": {"modelId": 2, "id": 21}}]
scanner = DummyScanner(raw_data)
provider = DummyProvider(
{
"modelVersions": [
{"id": 21, "files": [], "images": []},
{"id": 22, "files": [], "images": []},
]
}
)
await service.refresh_for_model_type("lora", scanner, provider)
await service.set_should_ignore("lora", 2, True)
provider.calls = 0
provider.bulk_calls = []
await service.refresh_for_model_type("lora", scanner, provider)
assert provider.calls == 0
assert provider.bulk_calls == []
record = await service.get_record("lora", 2)
assert record is not None
assert record.should_ignore_model is True
@pytest.mark.asyncio
async def test_refresh_marks_model_ignored_when_remote_missing(tmp_path):
db_path = tmp_path / "updates.sqlite"
service = ModelUpdateService(str(db_path), ttl_seconds=3600)
raw_data = [{"civitai": {"modelId": 5, "id": 51}}]
scanner = DummyScanner(raw_data)
provider = NotFoundProvider()
await service.refresh_for_model_type("lora", scanner, provider)
record = await service.get_record("lora", 5)
assert provider.bulk_calls == [[5]]
assert provider.calls == 1
assert record is not None
assert record.should_ignore_model is True
assert record.in_library_version_ids == [51]
assert record.last_checked_at is not None
@pytest.mark.asyncio
async def test_refresh_logs_info_for_missing_remote(tmp_path, caplog):
db_path = tmp_path / "updates.sqlite"
service = ModelUpdateService(str(db_path), ttl_seconds=3600)
raw_data = [{"civitai": {"modelId": 6, "id": 61}}]
scanner = DummyScanner(raw_data)
provider = NotFoundProvider()
with caplog.at_level(logging.INFO, logger="py.services.model_update_service"):
await service.refresh_for_model_type("lora", scanner, provider)
relevant = [
record for record in caplog.records if "Single lookup for model" in record.message
]
assert relevant, "expected single lookup log entry"
assert all(record.levelno == logging.INFO for record in relevant)
@pytest.mark.asyncio
async def test_refresh_falls_back_when_bulk_not_supported(tmp_path):
db_path = tmp_path / "updates.sqlite"
service = ModelUpdateService(str(db_path), ttl_seconds=3600)
raw_data = [{"civitai": {"modelId": 4, "id": 41}}]
scanner = DummyScanner(raw_data)
provider = DummyProvider(
{"modelVersions": [{"id": 41, "files": [], "images": []}]},
support_bulk=False,
)
await service.refresh_for_model_type("lora", scanner, provider)
record = await service.get_record("lora", 4)
assert record is not None
assert provider.calls == 1
assert provider.bulk_calls == []
@pytest.mark.asyncio
async def test_refresh_batches_large_collections(tmp_path):
db_path = tmp_path / "updates.sqlite"
service = ModelUpdateService(str(db_path), ttl_seconds=3600)
raw_data = [
{"civitai": {"modelId": idx, "id": idx * 10}}
for idx in range(1, 151)
]
scanner = DummyScanner(raw_data)
provider = DummyProvider({"modelVersions": []})
await service.refresh_for_model_type("lora", scanner, provider)
# Expect two batches: 100 ids and remaining 50 ids
assert len(provider.bulk_calls) == 2
assert len(provider.bulk_calls[0]) == 100
assert len(provider.bulk_calls[1]) == 50
@pytest.mark.asyncio
async def test_update_in_library_versions_changes_update_state(tmp_path):
db_path = tmp_path / "updates.sqlite"
service = ModelUpdateService(str(db_path), ttl_seconds=1)
raw_data = [{"civitai": {"modelId": 3, "id": 31}}]
scanner = DummyScanner(raw_data)
provider = DummyProvider(
{
"modelVersions": [
{"id": 31, "files": [], "images": []},
{"id": 35, "files": [], "images": []},
]
}
)
await service.refresh_for_model_type("lora", scanner, provider)
await service.update_in_library_versions("lora", 3, [31])
record = await service.get_record("lora", 3)
assert record is not None
assert record.has_update() is True
await service.update_in_library_versions("lora", 3, [31, 35])
record = await service.get_record("lora", 3)
assert record is not None
assert record.has_update() is False
@pytest.mark.asyncio
async def test_version_ignore_blocks_update_flag(tmp_path):
db_path = tmp_path / "updates.sqlite"
service = ModelUpdateService(str(db_path), ttl_seconds=1)
raw_data = [{"civitai": {"modelId": 5, "id": 51}}]
scanner = DummyScanner(raw_data)
provider = DummyProvider(
{
"modelVersions": [
{"id": 51, "files": [], "images": []},
{"id": 55, "files": [], "images": []},
]
}
)
await service.refresh_for_model_type("lora", scanner, provider)
record = await service.get_record("lora", 5)
assert record is not None
assert record.has_update() is True
await service.set_version_should_ignore("lora", 5, 55, True)
record = await service.get_record("lora", 5)
assert record is not None
assert record.has_update() is False
@pytest.mark.asyncio
async def test_has_updates_bulk_returns_mapping(tmp_path):
db_path = tmp_path / "updates.sqlite"
service = ModelUpdateService(str(db_path), ttl_seconds=3600)
raw_data = [{"civitai": {"modelId": 9, "id": 91}}]
scanner = DummyScanner(raw_data)
provider = DummyProvider(
{
"modelVersions": [
{"id": 91, "files": [], "images": []},
{"id": 92, "files": [], "images": []},
]
}
)
await service.refresh_for_model_type("lora", scanner, provider)
mapping = await service.has_updates_bulk("lora", [9, 9, 42])
assert mapping == {9: True, 42: False}
assert await service.has_update("lora", 9) is True
@pytest.mark.asyncio
async def test_has_updates_bulk_handles_more_than_sqlite_max_variables(tmp_path):
"""Bulk query with >999 model IDs must not raise 'too many SQL variables'."""
db_path = tmp_path / "updates.sqlite"
service = ModelUpdateService(str(db_path), ttl_seconds=3600)
model_ids = list(range(1, 1201))
with sqlite3.connect(str(db_path)) as conn:
conn.execute("INSERT INTO model_update_status (model_id, model_type) VALUES (?, ?)", (1, "lora"))
conn.execute("INSERT INTO model_update_versions (model_id, version_id, sort_index, name) VALUES (?, ?, ?, ?)", (1, 10, 0, "v1"))
mapping = await service.has_updates_bulk("lora", model_ids)
assert mapping[1] is True
assert len(mapping) == len(model_ids)
assert all(v is False for k, v in mapping.items() if k != 1)
@pytest.mark.asyncio
async def test_get_records_bulk_handles_more_than_sqlite_max_variables(tmp_path):
"""Bulk record fetch with >999 model IDs must not raise 'too many SQL variables'."""
db_path = tmp_path / "updates.sqlite"
service = ModelUpdateService(str(db_path), ttl_seconds=3600)
model_ids = list(range(1, 1201))
with sqlite3.connect(str(db_path)) as conn:
conn.execute("INSERT INTO model_update_status (model_id, model_type) VALUES (?, ?)", (1, "lora"))
conn.execute("INSERT INTO model_update_versions (model_id, version_id, sort_index, name) VALUES (?, ?, ?, ?)", (1, 10, 0, "v1"))
records = await service.get_records_bulk("lora", model_ids)
assert 1 in records
assert records[1].model_id == 1
assert len(records) == 1
@pytest.mark.asyncio
async def test_refresh_allows_duplicate_version_ids_across_models(tmp_path):
db_path = tmp_path / "updates.sqlite"
service = ModelUpdateService(str(db_path), ttl_seconds=0)
raw_data = [
{"civitai": {"modelId": 1, "id": 42}},
{"civitai": {"modelId": 2, "id": 42}},
]
scanner = DummyScanner(raw_data)
provider = DummyProvider(
{
"modelVersions": [
{
"id": 42,
"name": "shared",
"baseModel": "SD15",
"publishedAt": "2024-03-01T00:00:00Z",
"files": [{"sizeKB": 256}],
"images": [],
}
]
}
)
results = await service.refresh_for_model_type("lora", scanner, provider)
assert set(results.keys()) == {1, 2}
assert results[1].version_ids == [42]
assert results[2].version_ids == [42]
with sqlite3.connect(str(db_path)) as conn:
count = conn.execute(
"SELECT COUNT(*) FROM model_update_versions WHERE version_id = 42"
).fetchone()[0]
assert count == 2
@pytest.mark.asyncio
async def test_refresh_rewrites_remote_preview_urls(tmp_path):
db_path = tmp_path / "updates.sqlite"
service = ModelUpdateService(str(db_path), ttl_seconds=1)
raw_data = [{"civitai": {"modelId": 7, "id": 71}}]
scanner = DummyScanner(raw_data)
provider = DummyProvider(
{
"modelVersions": [
{
"id": 71,
"files": [],
"images": [
{
"url": "https://image.civitai.com/high/original=true/sample.png",
"nsfwLevel": 6,
"type": "image",
},
{
"url": "https://image.civitai.com/safe/original=true/preview.png",
"nsfwLevel": 1,
"type": "image",
},
],
}
]
}
)
await service.refresh_for_model_type("lora", scanner, provider)
record = await service.get_record("lora", 7)
assert record is not None
assert record.versions
preview_url = record.versions[0].preview_url
@pytest.mark.asyncio
async def test_update_in_library_versions_populates_metadata(tmp_path):
db_path = tmp_path / "updates.sqlite"
service = ModelUpdateService(str(db_path))
version_info = {
"id": 123,
"name": "v1.0",
"baseModel": "SD 1.5",
"publishedAt": "2024-03-01T00:00:00Z",
"files": [{"sizeKB": 1024, "type": "Model", "primary": True}],
"images": [{"url": "https://example.com/preview.png"}],
}
await service.update_in_library_versions("lora", 1, [123], version_info=version_info)
record = await service.get_record("lora", 1)
assert record is not None
assert len(record.versions) == 1
version = record.versions[0]
assert version.version_id == 123
assert version.name == "v1.0"
assert version.base_model == "SD 1.5"
assert version.size_bytes == 1024 * 1024
assert version.preview_url == "https://example.com/preview.png"
assert version.is_in_library is True
@pytest.mark.asyncio
async def test_refresh_folder_filter_considers_cross_folder_versions(tmp_path):
"""When refreshing by folder, versions in other folders must still be
considered in-library so they aren't reported as available updates."""
db_path = tmp_path / "updates.sqlite"
service = ModelUpdateService(str(db_path), ttl_seconds=0)
# Same model (modelId=1) in two folders with different versions
raw_data = [
{"civitai": {"modelId": 1, "id": 11}, "folder": "folder_a"},
{"civitai": {"modelId": 1, "id": 15}, "folder": "folder_b"},
]
scanner = DummyScanner(raw_data)
# Remote offers: 11 (in folder_a), 15 (in folder_b), 20 (truly new)
provider = DummyProvider(
{
"modelVersions": [
{"id": 11, "files": [], "images": []},
{"id": 15, "files": [], "images": []},
{"id": 20, "files": [], "images": []},
]
}
)
await service.refresh_for_model_type(
"lora", scanner, provider, folder_path="folder_a",
)
record = await service.get_record("lora", 1)
assert record is not None
# Version 15 is in folder_b — must be in_library even when filtering by folder_a
v15 = next(v for v in record.versions if v.version_id == 15)
assert v15.is_in_library is True
# Version 20 is truly new — should not be in_library
v20 = next(v for v in record.versions if v.version_id == 20)
assert v20.is_in_library is False
# has_update must be True (version 20 > max_in_library=15)
assert record.has_update() is True
def test_extract_single_version_paid_access_timed(tmp_path):
"""A timed paidAccess gate (permanent=False + future endsAt) is detected
as early access while availability stays 'Public'."""
db_path = tmp_path / "updates.sqlite"
service = ModelUpdateService(str(db_path))
entry = {
"id": 42,
"name": "v1 paid",
"availability": "Public",
"paidAccess": {
"permanent": False,
"endsAt": "2026-08-22T18:30:00.000Z",
},
"files": [],
"images": [],
}
version = service._extract_single_version(entry, index=0)
assert version is not None
assert version.is_early_access is True
assert version.early_access_ends_at == "2026-08-22T18:30:00.000Z"
assert version.is_paid is False
assert version.paid_access is not None
def test_extract_single_version_paid_access_permanent(tmp_path):
"""A permanent paidAccess gate (permanent=True, no endsAt) is detected and
flagged as paid but is NOT early access and carries no end date."""
db_path = tmp_path / "updates.sqlite"
service = ModelUpdateService(str(db_path))
entry = {
"id": 42,
"name": "v1 paid",
"availability": "Public",
"paidAccess": {"permanent": True, "endsAt": None},
"files": [],
"images": [],
}
version = service._extract_single_version(entry, index=0)
assert version is not None
assert version.is_early_access is False
assert version.is_paid is True
assert version.early_access_ends_at is None
assert version.paid_access is not None
def test_extract_single_version_paid_access_pending_end(tmp_path):
"""A timed gate whose window end is not recorded yet
({"permanent": false, "endsAt": null}) is still an active gate, so it is
early access with no known end date rather than a free version."""
db_path = tmp_path / "updates.sqlite"
service = ModelUpdateService(str(db_path))
entry = {
"id": 42,
"name": "v1 paid",
"availability": "Public",
"paidAccess": {"permanent": False, "endsAt": None},
"files": [],
"images": [],
}
version = service._extract_single_version(entry, index=0)
assert version is not None
assert version.is_early_access is True
assert version.is_paid is False
assert version.early_access_ends_at is None
assert version.paid_access == '{"permanent": false, "endsAt": null}'
assert ModelUpdateRecord._is_early_access_active(version) is True
def test_normalize_paid_access_accepts_json_string():
"""The by-hash enrichment path may hand paidAccess to _normalize_paid_access
as a JSON string; both the permanent and timed shapes must normalize."""
service = ModelUpdateService.__new__(ModelUpdateService)
permanent = ModelUpdateService._normalize_paid_access(
'{"permanent": true, "endsAt": null}'
)
assert permanent == {"permanent": True, "endsAt": None}
timed = ModelUpdateService._normalize_paid_access(
'{"permanent": false, "endsAt": "2026-08-22T18:30:00.000Z"}'
)
assert timed == {"permanent": False, "endsAt": "2026-08-22T18:30:00.000Z"}
# A timed gate whose window end is not recorded yet. CivitAI only returns a
# non-null DTO for an ACTIVE gate (tombstones come back as null) and enforces
# this shape too - the model page reports canDownload: false for it - so it
# must be kept. Dropping it was the #1060 class of bug.
pending_end = ModelUpdateService._normalize_paid_access(
'{"permanent": false, "endsAt": null}'
)
assert pending_end == {"permanent": False, "endsAt": None}
malformed = ModelUpdateService._normalize_paid_access("{not json")
assert malformed is None
# No gate keys at all is not a gate signal.
assert ModelUpdateService._normalize_paid_access("{}") is None
def test_has_update_for_base_hide_paid():
"""hide_paid also suppresses permanent paid versions in the same-base
update path (has_update_for_base)."""
record = make_record(
make_version(5, in_library=True, base_model="illustrious"),
make_version(
7,
in_library=False,
base_model="illustrious",
is_paid=True,
paid_access='{"permanent": true, "endsAt": null}',
),
)
assert record.has_update_for_base(5, "illustrious") is True
assert record.has_update_for_base(5, "illustrious", hide_paid=True) is False
def test_has_update_hide_paid():
"""hide_paid suppresses update flags raised by a permanent paid version."""
record = make_record(
make_version(5, in_library=True),
make_version(
7,
in_library=False,
is_paid=True,
paid_access='{"permanent": true, "endsAt": null}',
),
)
assert record.has_update() is True
assert record.has_update(hide_paid=True) is False
def test_has_update_hide_early_access_paid_timed():
"""hide_early_access suppresses a newer timed paidAccess version."""
record = make_record(
make_version(5, in_library=True),
make_version(
7,
in_library=False,
is_early_access=True,
early_access_ends_at="2099-01-01T00:00:00Z",
),
)
assert record.has_update() is True
assert record.has_update(hide_early_access=True) is False
def test_build_record_from_remote_preserves_paid_fields(tmp_path):
"""_build_record_from_remote must carry paid_access/is_paid from the
parsed remote versions into the rebuilt record, or the refresh path
silently drops paid data before persistence."""
db_path = tmp_path / "updates.sqlite"
service = ModelUpdateService(str(db_path))
remote_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='{"permanent": true, "endsAt": null}',
is_paid=True,
)
record = service._build_record_from_remote(
model_type="lora",
model_id=123,
local_versions=[],
remote_versions=[remote_version],
existing=None,
timestamp=1.0,
)
rebuilt = record.versions[0]
assert rebuilt.paid_access == '{"permanent": true, "endsAt": null}'
assert rebuilt.is_paid is True
def test_extract_file_count_counts_weight_files(tmp_path):
"""file_count counts only weight-type files; a missing files array stays
None (unknown) so the UI can distinguish it from "no weight files"."""
db_path = tmp_path / "updates.sqlite"
service = ModelUpdateService(str(db_path))
response = {
"modelVersions": [
{
"id": 42,
"files": [
{"sizeKB": 100, "type": "Model", "primary": True},
{"sizeKB": 10, "type": "Training Data"},
{"sizeKB": 50, "type": "Pruned Model"},
],
"images": [],
},
{"id": 43, "images": []},
{"id": 44, "files": [], "images": []},
]
}
versions = service._extract_versions(response)
assert versions is not None
assert versions[0].file_count == 2
assert versions[1].file_count is None
assert versions[2].file_count == 0
@pytest.mark.asyncio
async def test_refresh_persists_file_count(tmp_path):
db_path = tmp_path / "updates.sqlite"
service = ModelUpdateService(str(db_path), ttl_seconds=3600)
raw_data = [{"civitai": {"modelId": 1, "id": 11}}]
scanner = DummyScanner(raw_data)
provider = DummyProvider(
{
"modelVersions": [
{
"id": 11,
"name": "v1",
"baseModel": "SD15",
"files": [
{"sizeKB": 1024, "type": "Model", "primary": True},
{"sizeKB": 2048, "type": "Model"},
{"sizeKB": 128, "type": "Training Data"},
],
"images": [],
}
]
}
)
await service.refresh_for_model_type("lora", scanner, provider)
record = await service.get_record("lora", 1)
assert record is not None
assert record.versions[0].file_count == 2
def test_build_record_from_remote_preserves_file_count(tmp_path):
"""A remote payload without files data must not clobber the previously
persisted file_count; a populated payload wins."""
db_path = tmp_path / "updates.sqlite"
service = ModelUpdateService(str(db_path))
existing = make_record(
ModelVersionRecord(
version_id=7,
name="v7",
base_model=None,
released_at=None,
size_bytes=None,
preview_url=None,
is_in_library=True,
should_ignore=False,
file_count=3,
)
)
remote_without_count = 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,
file_count=None,
)
record = service._build_record_from_remote(
model_type="lora",
model_id=999,
local_versions=[7],
remote_versions=[remote_without_count],
existing=existing,
timestamp=1.0,
)
assert record.versions[0].file_count == 3
remote_with_count = 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,
file_count=1,
)
record = service._build_record_from_remote(
model_type="lora",
model_id=999,
local_versions=[7],
remote_versions=[remote_with_count],
existing=existing,
timestamp=2.0,
)
assert record.versions[0].file_count == 1
# --- Gate-state transitions and price persistence ---------------------------
def _remote_gated(version_id, *, price_checked_at=None, price_buzz=None):
return ModelVersionRecord(
version_id=version_id,
name=f"v{version_id}",
base_model=None,
released_at=None,
size_bytes=None,
preview_url=None,
is_in_library=False,
should_ignore=False,
paid_access='{"permanent": false, "endsAt": "2026-10-10T13:10:17.404Z"}',
is_early_access=True,
early_access_ends_at="2026-10-10T13:10:17.404Z",
price_checked_at=price_checked_at,
price_buzz=price_buzz,
)
def _remote_free(version_id):
return ModelVersionRecord(
version_id=version_id,
name=f"v{version_id}",
base_model=None,
released_at=None,
size_bytes=None,
preview_url=None,
is_in_library=False,
should_ignore=False,
)
def test_build_record_emits_became_free_event(tmp_path):
"""A version whose gate lapsed produces a became_free event, keeps a lapse
timestamp, and drops the now-meaningless stored price."""
service = ModelUpdateService(str(tmp_path / "updates.sqlite"))
existing = make_record(
replace(
_remote_gated(7, price_checked_at=100.0, price_buzz=500),
)
)
record = service._build_record_from_remote(
model_type="lora",
model_id=999,
local_versions=[],
remote_versions=[_remote_free(7)],
existing=existing,
timestamp=1000.0,
)
version = record.versions[0]
assert version.gate_lapsed_at is not None
assert version.price_buzz is None
assert version.price_checked_at is None
assert version.price_alert_state is False
assert [event["kind"] for event in record.events] == ["became_free"]
assert record.events[0]["versionId"] == 7
def test_build_record_emits_new_gate_event(tmp_path):
service = ModelUpdateService(str(tmp_path / "updates.sqlite"))
existing = make_record(_remote_free(7))
record = service._build_record_from_remote(
model_type="lora",
model_id=999,
local_versions=[],
remote_versions=[_remote_gated(7)],
existing=existing,
timestamp=1000.0,
)
assert [event["kind"] for event in record.events] == ["new_gate"]
assert record.versions[0].gate_lapsed_at is None
def test_build_record_no_events_without_previous_snapshot(tmp_path):
"""First sight of a model must not report every existing gate as a new one."""
service = ModelUpdateService(str(tmp_path / "updates.sqlite"))
record = service._build_record_from_remote(
model_type="lora",
model_id=999,
local_versions=[],
remote_versions=[_remote_gated(7), _remote_free(8)],
existing=None,
timestamp=1000.0,
)
assert record.events == []
def test_build_record_keeps_lapse_marker_and_skips_ignored(tmp_path):
"""An already-free version keeps its original lapse marker and an ignored
version reports nothing at all."""
service = ModelUpdateService(str(tmp_path / "updates.sqlite"))
lapsed = replace(_remote_free(7), gate_lapsed_at="2026-09-01T00:00:00.000Z")
existing = make_record(lapsed, replace(_remote_gated(8), should_ignore=True))
record = service._build_record_from_remote(
model_type="lora",
model_id=999,
local_versions=[],
remote_versions=[_remote_free(7), _remote_free(8)],
existing=existing,
timestamp=1000.0,
)
by_id = {version.version_id: version for version in record.versions}
assert by_id[7].gate_lapsed_at == "2026-09-01T00:00:00.000Z"
assert record.events == []
def test_build_record_preserves_price_when_fetch_skipped(tmp_path):
"""A refresh that did not run a price fetch (no price_checked_at) must keep the
stored price instead of wiping it."""
service = ModelUpdateService(str(tmp_path / "updates.sqlite"))
existing = make_record(
replace(_remote_gated(7, price_checked_at=100.0, price_buzz=500))
)
record = service._build_record_from_remote(
model_type="lora",
model_id=999,
local_versions=[],
remote_versions=[_remote_gated(7)],
existing=existing,
timestamp=1000.0,
)
assert record.versions[0].price_buzz == 500
assert record.versions[0].price_checked_at == 100.0
assert record.events == []
def test_build_record_applies_fresh_price(tmp_path):
service = ModelUpdateService(str(tmp_path / "updates.sqlite"))
existing = make_record(
replace(_remote_gated(7, price_checked_at=100.0, price_buzz=500))
)
record = service._build_record_from_remote(
model_type="lora",
model_id=999,
local_versions=[],
remote_versions=[_remote_gated(7, price_checked_at=200.0, price_buzz=250)],
existing=existing,
timestamp=1000.0,
)
assert record.versions[0].price_buzz == 250
assert record.versions[0].price_checked_at == 200.0
def _legacy_version_table_sql() -> str:
"""The model_update_versions schema before gate-lapse/price columns."""
return """
CREATE TABLE model_update_versions (
model_id INTEGER NOT NULL,
version_id INTEGER NOT NULL,
sort_index INTEGER NOT NULL DEFAULT 0,
name TEXT,
base_model TEXT,
released_at TEXT,
size_bytes INTEGER,
preview_url TEXT,
is_in_library INTEGER NOT NULL DEFAULT 0,
should_ignore INTEGER NOT NULL DEFAULT 0,
early_access_ends_at TEXT,
is_early_access INTEGER NOT NULL DEFAULT 0,
usage_control TEXT,
paid_access TEXT,
is_paid INTEGER NOT NULL DEFAULT 0,
file_count INTEGER,
PRIMARY KEY (model_id, version_id)
)
"""
def test_migration_adds_price_columns_to_legacy_db(tmp_path):
"""Opening a database written before this feature must add the columns and keep
the existing rows (the migration path real users hit)."""
db_path = tmp_path / "updates.sqlite"
conn = sqlite3.connect(db_path)
conn.execute(
"CREATE TABLE model_update_status ("
"model_id INTEGER PRIMARY KEY, model_type TEXT NOT NULL, "
"last_checked_at REAL, should_ignore_model INTEGER NOT NULL DEFAULT 0)"
)
conn.execute(_legacy_version_table_sql())
conn.execute(
"INSERT INTO model_update_status (model_id, model_type, last_checked_at) "
"VALUES (999, 'lora', 1.0)"
)
conn.execute(
"INSERT INTO model_update_versions ("
"model_id, version_id, is_in_library, is_early_access, paid_access, is_paid"
") VALUES (999, 7, 1, 1, '{\"permanent\": false, \"endsAt\": null}', 0)"
)
conn.commit()
conn.close()
service = ModelUpdateService(str(db_path))
columns = service._get_table_columns(
service._connect(), "model_update_versions"
)
for column in (
"gate_lapsed_at",
"price_buzz",
"list_price_buzz",
"generation_price_buzz",
"accepts_blue_buzz",
"price_sale_ends_at",
"price_checked_at",
"price_alert_state",
):
assert column in columns
record = service._get_record("lora", 999)
assert record is not None
assert [version.version_id for version in record.versions] == [7]
version = record.versions[0]
assert version.paid_access == '{"permanent": false, "endsAt": null}'
assert version.price_buzz is None
assert version.price_checked_at is None
assert version.gate_lapsed_at is None
assert version.accepts_blue_buzz is False
def test_price_fields_round_trip_through_sqlite(tmp_path):
"""Price and lapse columns survive an upsert/read cycle."""
service = ModelUpdateService(str(tmp_path / "updates.sqlite"))
version = replace(
_remote_gated(7, price_checked_at=123.5, price_buzz=500),
list_price_buzz=600,
generation_price_buzz=100,
accepts_blue_buzz=True,
price_sale_ends_at="2026-10-01T00:00:00.000Z",
price_alert_state=True,
is_in_library=True,
gate_lapsed_at=None,
)
service._upsert_record(make_record(version))
stored = service._get_record("lora", 999)
assert stored is not None
loaded = stored.versions[0]
assert loaded.price_buzz == 500
assert loaded.list_price_buzz == 600
assert loaded.generation_price_buzz == 100
assert loaded.accepts_blue_buzz is True
assert loaded.price_sale_ends_at == "2026-10-01T00:00:00.000Z"
assert loaded.price_checked_at == 123.5
assert loaded.price_alert_state is True
# --- Optional price capture --------------------------------------------------
class FakeSettings:
"""Minimal stand-in for SettingsManager (only `.get` is used by these paths)."""
def __init__(self, values=None):
self._values = dict(values or {})
def get(self, key, default=None):
return self._values.get(key, default)
def set(self, key, value):
self._values[key] = value
class PriceProvider(DummyProvider):
"""DummyProvider that can also serve prices."""
def __init__(self, response, *, prices=None, error=None):
super().__init__(response)
self.prices = prices
self.error = error
self.price_calls = 0
async def get_model_prices(self, model_id):
self.price_calls += 1
if self.error is not None:
raise self.error
return self.prices
GATED_RESPONSE = {
"modelVersions": [
{
"id": 12,
"baseModel": "Pony",
"availability": "Public",
"paidAccess": {"permanent": False, "endsAt": "2999-01-01T00:00:00.000Z"},
"files": [],
"images": [],
}
]
}
FREE_RESPONSE = {
"modelVersions": [
{"id": 12, "baseModel": "Pony", "availability": "Public", "files": [], "images": []}
]
}
LOCAL_RAW_DATA = [{"civitai": {"modelId": 1, "id": 11}, "base_model": "Pony"}]
PRICE_PAYLOAD = {
12: {
"price_buzz": 250,
"list_price_buzz": 500,
"generation_price_buzz": 100,
"accepts_blue_buzz": True,
"price_sale_ends_at": "2999-01-02T00:00:00.000Z",
}
}
def _price_service(tmp_path, **settings):
return ModelUpdateService(
str(tmp_path / "updates.sqlite"),
ttl_seconds=0,
settings_manager=FakeSettings(settings),
)
@pytest.mark.asyncio
async def test_price_capture_off_by_default(tmp_path):
service = _price_service(tmp_path)
scanner = DummyScanner(LOCAL_RAW_DATA)
provider = PriceProvider(GATED_RESPONSE, prices=PRICE_PAYLOAD)
await service.refresh_for_model_type("lora", scanner, provider)
record = await service.get_record("lora", 1)
assert provider.price_calls == 0
assert record.versions[0].price_checked_at is None
assert record.versions[0].price_buzz is None
@pytest.mark.asyncio
async def test_price_capture_stores_prices_for_gated_versions(tmp_path):
service = _price_service(tmp_path, price_tracking_enabled=True)
scanner = DummyScanner(LOCAL_RAW_DATA)
provider = PriceProvider(GATED_RESPONSE, prices=PRICE_PAYLOAD)
await service.refresh_for_model_type("lora", scanner, provider)
record = await service.get_record("lora", 1)
assert provider.price_calls == 1
version = next(v for v in record.versions if v.version_id == 12)
assert version.price_buzz == 250
assert version.list_price_buzz == 500
assert version.generation_price_buzz == 100
assert version.accepts_blue_buzz is True
assert version.price_sale_ends_at == "2999-01-02T00:00:00.000Z"
assert version.price_checked_at is not None
@pytest.mark.asyncio
async def test_price_capture_skips_ungated_models(tmp_path):
service = _price_service(tmp_path, price_tracking_enabled=True)
scanner = DummyScanner(LOCAL_RAW_DATA)
provider = PriceProvider(FREE_RESPONSE, prices=PRICE_PAYLOAD)
await service.refresh_for_model_type("lora", scanner, provider)
assert provider.price_calls == 0
@pytest.mark.asyncio
async def test_price_capture_failure_keeps_stored_price(tmp_path):
service = _price_service(tmp_path, price_tracking_enabled=True)
scanner = DummyScanner(LOCAL_RAW_DATA)
await service.refresh_for_model_type(
"lora", scanner, PriceProvider(GATED_RESPONSE, prices=PRICE_PAYLOAD)
)
first = await service.get_record("lora", 1)
stored_checked_at = next(
v for v in first.versions if v.version_id == 12
).price_checked_at
assert stored_checked_at is not None
# The page is unreadable this time (a changed gate still triggers a fetch):
# the price must survive untouched rather than being blanked.
failing = PriceProvider(
{
"modelVersions": [
{
"id": 12,
"baseModel": "Pony",
"availability": "Public",
"paidAccess": {
"permanent": False,
"endsAt": "2999-06-01T00:00:00.000Z",
},
"files": [],
"images": [],
}
]
},
prices=None,
)
await service.refresh_for_model_type("lora", scanner, failing)
record = await service.get_record("lora", 1)
version = next(v for v in record.versions if v.version_id == 12)
assert failing.price_calls == 1
assert version.price_buzz == 250
assert version.price_checked_at == stored_checked_at
@pytest.mark.asyncio
async def test_price_capture_respects_ttl(tmp_path):
"""A second refresh with an unchanged gate and a fresh price must not refetch."""
service = _price_service(
tmp_path, price_tracking_enabled=True, price_check_ttl_hours=24
)
scanner = DummyScanner(LOCAL_RAW_DATA)
provider = PriceProvider(GATED_RESPONSE, prices=PRICE_PAYLOAD)
await service.refresh_for_model_type("lora", scanner, provider)
await service.refresh_for_model_type("lora", scanner, provider)
assert provider.price_calls == 1
@pytest.mark.asyncio
async def test_price_capture_refetches_when_gate_changes(tmp_path):
service = _price_service(tmp_path, price_tracking_enabled=True)
scanner = DummyScanner(LOCAL_RAW_DATA)
provider = PriceProvider(GATED_RESPONSE, prices=PRICE_PAYLOAD)
await service.refresh_for_model_type("lora", scanner, provider)
changed_gate = {
"modelVersions": [
{
"id": 12,
"baseModel": "Pony",
"availability": "Public",
"paidAccess": {
"permanent": True,
"endsAt": None,
},
"files": [],
"images": [],
}
]
}
provider.response = changed_gate
await service.refresh_for_model_type("lora", scanner, provider)
assert provider.price_calls == 2
@pytest.mark.asyncio
async def test_price_capture_survives_provider_without_support(tmp_path):
"""A provider that cannot serve prices must not break the refresh."""
service = _price_service(tmp_path, price_tracking_enabled=True)
scanner = DummyScanner(LOCAL_RAW_DATA)
provider = DummyProvider(GATED_RESPONSE)
await service.refresh_for_model_type("lora", scanner, provider)
record = await service.get_record("lora", 1)
assert record is not None
version = next(v for v in record.versions if v.version_id == 12)
assert version.price_buzz is None
@pytest.mark.asyncio
async def test_price_capture_drops_price_when_version_becomes_free(tmp_path):
service = _price_service(tmp_path, price_tracking_enabled=True)
scanner = DummyScanner(LOCAL_RAW_DATA)
await service.refresh_for_model_type(
"lora", scanner, PriceProvider(GATED_RESPONSE, prices=PRICE_PAYLOAD)
)
await service.refresh_for_model_type("lora", scanner, PriceProvider(FREE_RESPONSE))
record = await service.get_record("lora", 1)
version = next(v for v in record.versions if v.version_id == 12)
assert version.price_buzz is None
assert version.price_checked_at is None
assert version.gate_lapsed_at is not None
# --- Price alerts ------------------------------------------------------------
def _gated_response(ends_at: str) -> dict:
"""Gated response with a specific EA end.
Varying the end date between refreshes is what forces a re-fetch inside a
single test (a price is otherwise considered fresh for the whole TTL).
"""
return {
"modelVersions": [
{
"id": 12,
"baseModel": "Pony",
"availability": "Public",
"paidAccess": {"permanent": False, "endsAt": ends_at},
"files": [],
"images": [],
}
]
}
def _prices(price_buzz: int) -> dict:
return {12: {"price_buzz": price_buzz, "list_price_buzz": price_buzz}}
@pytest.mark.asyncio
async def test_price_alert_state_persists_and_fires_on_crossing(tmp_path):
service = _price_service(
tmp_path, price_tracking_enabled=True, price_alert_threshold_buzz=300
)
scanner = DummyScanner(LOCAL_RAW_DATA)
# First sight: the state is recorded (so the alerts list shows it) but nothing
# is announced - otherwise enabling the feature would toast every cheap
# version in the library at once.
first_refresh = await service.refresh_for_model_type(
"lora",
scanner,
PriceProvider(_gated_response("2999-01-01T00:00:00.000Z"), prices=_prices(250)),
)
version = next(v for v in first_refresh[1].versions if v.version_id == 12)
assert version.price_alert_state is True
assert first_refresh[1].events == []
alerts = await service.get_price_alerts("lora")
assert [alert["versionId"] for alert in alerts] == [12]
assert alerts[0]["priceBuzz"] == 250
assert alerts[0]["modelId"] == 1
# Price rises above the threshold: the state resets silently.
raised = await service.refresh_for_model_type(
"lora",
scanner,
PriceProvider(_gated_response("2999-02-01T00:00:00.000Z"), prices=_prices(900)),
)
assert raised[1].versions[0].price_alert_state is False
assert raised[1].events == []
assert await service.get_price_alerts("lora") == []
# …and drops back under it: now it is news.
dropped = await service.refresh_for_model_type(
"lora",
scanner,
PriceProvider(_gated_response("2999-03-01T00:00:00.000Z"), prices=_prices(200)),
)
assert dropped[1].versions[0].price_alert_state is True
assert [event["kind"] for event in dropped[1].events] == ["price_drop"]
assert dropped[1].events[0]["priceBuzz"] == 200
# Events are derived, never persisted: a later read reports none.
stored = await service.get_record("lora", 1)
assert stored.events == []
assert stored.versions[0].price_alert_state is True
@pytest.mark.asyncio
async def test_price_alert_not_fired_above_threshold(tmp_path):
service = _price_service(
tmp_path, price_tracking_enabled=True, price_alert_threshold_buzz=100
)
scanner = DummyScanner(LOCAL_RAW_DATA)
provider = PriceProvider(GATED_RESPONSE, prices=PRICE_PAYLOAD)
await service.refresh_for_model_type("lora", scanner, provider)
record = await service.get_record("lora", 1)
assert record.versions[0].price_alert_state is False
assert record.events == []
assert await service.get_price_alerts("lora") == []
@pytest.mark.asyncio
async def test_price_alert_requires_price_tracking(tmp_path):
"""With tracking off there are no prices, so no price alert either."""
service = _price_service(tmp_path, price_alert_threshold_buzz=100000)
scanner = DummyScanner(LOCAL_RAW_DATA)
provider = PriceProvider(GATED_RESPONSE, prices=PRICE_PAYLOAD)
await service.refresh_for_model_type("lora", scanner, provider)
record = await service.get_record("lora", 1)
assert record.versions[0].price_alert_state is False
assert await service.get_price_alerts("lora") == []
@pytest.mark.asyncio
async def test_price_alerts_respect_ignores_and_model_type(tmp_path):
service = _price_service(
tmp_path, price_tracking_enabled=True, price_alert_threshold_buzz=1000
)
scanner = DummyScanner(LOCAL_RAW_DATA)
provider = PriceProvider(GATED_RESPONSE, prices=PRICE_PAYLOAD)
await service.refresh_for_model_type("lora", scanner, provider)
assert len(await service.get_price_alerts("lora")) == 1
await service.set_version_should_ignore("lora", 1, 12, True)
assert await service.get_price_alerts("lora") == []
assert await service.get_price_alerts("checkpoint") == []
@pytest.mark.asyncio
async def test_get_price_alerts_tolerates_bad_limit(tmp_path):
service = ModelUpdateService(str(tmp_path / "updates.sqlite"))
assert await service.get_price_alerts("lora", limit="not-a-number") == []
# --- Price alerts panel semantics (read-time threshold, across types) ---------
def _gated_response_for(version_id: int, ends_at: str) -> dict:
return {
"modelVersions": [
{
"id": version_id,
"baseModel": "Pony",
"availability": "Public",
"paidAccess": {"permanent": False, "endsAt": ends_at},
"files": [],
"images": [],
}
]
}
def _free_response_for(version_id: int) -> dict:
return {
"modelVersions": [
{
"id": version_id,
"baseModel": "Pony",
"availability": "Public",
"files": [],
"images": [],
}
]
}
@pytest.mark.asyncio
async def test_price_alerts_threshold_is_compared_at_read_time(tmp_path):
"""Changing the threshold must change panel membership without a refresh."""
service = _price_service(
tmp_path, price_tracking_enabled=True, price_alert_threshold_buzz=300
)
scanner = DummyScanner(LOCAL_RAW_DATA)
await service.refresh_for_model_type(
"lora", scanner, PriceProvider(GATED_RESPONSE, prices=_prices(250))
)
assert len(await service.get_price_alerts("lora")) == 1
# The stored alert state still says "hit"...
stored = await service.get_record("lora", 1)
assert stored.versions[0].price_alert_state is True
# ...but a tighter threshold excludes it immediately.
assert await service.get_price_alerts("lora", threshold_buzz=100) == []
assert len(await service.get_price_alerts("lora", threshold_buzz=300)) == 1
@pytest.mark.asyncio
async def test_price_alerts_span_all_model_types(tmp_path):
"""model_type=None is what the global panel uses: one list, all types."""
service = _price_service(
tmp_path, price_tracking_enabled=True, price_alert_threshold_buzz=1000
)
lora_scanner = DummyScanner([{"civitai": {"modelId": 1, "id": 11}}])
checkpoint_scanner = DummyScanner([{"civitai": {"modelId": 2, "id": 21}}])
await service.refresh_for_model_type(
"lora",
lora_scanner,
PriceProvider(_gated_response_for(12, "2999-01-01T00:00:00.000Z"), prices=_prices(250)),
)
await service.refresh_for_model_type(
"checkpoint",
checkpoint_scanner,
PriceProvider(
_gated_response_for(22, "2999-01-01T00:00:00.000Z"),
prices={22: {"price_buzz": 100}},
),
)
everything = await service.get_price_alerts()
assert {(a["modelType"], a["modelId"], a["versionId"]) for a in everything} == {
("lora", 1, 12),
("checkpoint", 2, 22),
}
# Cheapest first.
assert [a["priceBuzz"] for a in everything] == [100, 250]
assert [a["modelType"] for a in await service.get_price_alerts("checkpoint")] == [
"checkpoint"
]
@pytest.mark.asyncio
async def test_price_alert_since_tracks_the_crossing_and_clears(tmp_path):
service = _price_service(
tmp_path, price_tracking_enabled=True, price_alert_threshold_buzz=300
)
scanner = DummyScanner(LOCAL_RAW_DATA)
first = await service.refresh_for_model_type(
"lora",
scanner,
PriceProvider(_gated_response("2999-01-01T00:00:00.000Z"), prices=_prices(250)),
)
since = first[1].versions[0].price_alert_since
assert since is not None
# Still under the threshold: the original crossing time is preserved.
second = await service.refresh_for_model_type(
"lora",
scanner,
PriceProvider(_gated_response("2999-02-01T00:00:00.000Z"), prices=_prices(200)),
)
assert second[1].versions[0].price_alert_since == since
stored = await service.get_record("lora", 1)
assert stored.versions[0].price_alert_since == since
# Back above it: cleared, and the row leaves the panel.
third = await service.refresh_for_model_type(
"lora",
scanner,
PriceProvider(_gated_response("2999-03-01T00:00:00.000Z"), prices=_prices(900)),
)
assert third[1].versions[0].price_alert_since is None
assert await service.get_price_alerts("lora") == []
@pytest.mark.asyncio
async def test_price_alerts_report_became_free_without_price_tracking(tmp_path):
"""Gate transitions need no price data, so the "became free" half of the panel
works even while price tracking is switched off."""
service = _price_service(tmp_path) # tracking off
scanner = DummyScanner(LOCAL_RAW_DATA)
await service.refresh_for_model_type("lora", scanner, DummyProvider(GATED_RESPONSE))
assert await service.get_price_alerts("lora") == []
await service.refresh_for_model_type("lora", scanner, DummyProvider(FREE_RESPONSE))
alerts = await service.get_price_alerts("lora")
assert [alert["kind"] for alert in alerts] == ["became_free"]
assert alerts[0]["priceBuzz"] is None
assert alerts[0]["gateLapsedAt"] is not None
assert alerts[0]["versionId"] == 12
@pytest.mark.asyncio
async def test_newest_price_checked_at_reports_the_latest_fetch(tmp_path):
service = _price_service(tmp_path, price_tracking_enabled=True)
scanner = DummyScanner(LOCAL_RAW_DATA)
assert service.newest_price_checked_at() is None
await service.refresh_for_model_type(
"lora", scanner, PriceProvider(GATED_RESPONSE, prices=PRICE_PAYLOAD)
)
newest = service.newest_price_checked_at()
assert newest is not None
assert newest > 0
@pytest.mark.asyncio
async def test_failed_price_attempt_is_recorded_as_unavailable(tmp_path):
"""A gated version we tried to price and could not must be distinguishable
from one we never looked at — that is the honest "unavailable" state for
mature models, whose pages no host will serve anonymously."""
service = _price_service(tmp_path, price_tracking_enabled=True)
scanner = DummyScanner(LOCAL_RAW_DATA)
failing = PriceProvider(GATED_RESPONSE, prices=None)
await service.refresh_for_model_type("lora", scanner, failing)
record = await service.get_record("lora", 1)
version = next(v for v in record.versions if v.version_id == 12)
assert failing.price_calls == 1
assert version.price_buzz is None
assert version.price_checked_at is None
assert version.price_check_attempted_at is not None
assert service.count_unavailable_prices("lora") == 1
assert service.count_unavailable_prices("checkpoint") == 0
# Nothing to alert on, and no price alert state.
assert await service.get_price_alerts("lora") == []
assert version.price_alert_state is False
@pytest.mark.asyncio
async def test_successful_price_attempt_sets_both_markers(tmp_path):
service = _price_service(tmp_path, price_tracking_enabled=True)
scanner = DummyScanner(LOCAL_RAW_DATA)
await service.refresh_for_model_type(
"lora", scanner, PriceProvider(GATED_RESPONSE, prices=PRICE_PAYLOAD)
)
record = await service.get_record("lora", 1)
version = next(v for v in record.versions if v.version_id == 12)
assert version.price_buzz == 250
assert version.price_checked_at is not None
assert version.price_check_attempted_at is not None
assert service.count_unavailable_prices("lora") == 0
@pytest.mark.asyncio
async def test_no_price_attempt_is_recorded_while_tracking_is_off(tmp_path):
service = _price_service(tmp_path) # tracking off
scanner = DummyScanner(LOCAL_RAW_DATA)
await service.refresh_for_model_type("lora", scanner, DummyProvider(GATED_RESPONSE))
record = await service.get_record("lora", 1)
assert record.versions[0].price_check_attempted_at is None
assert service.count_unavailable_prices("lora") == 0
@pytest.mark.asyncio
async def test_unavailable_marker_clears_when_the_version_becomes_free(tmp_path):
service = _price_service(tmp_path, price_tracking_enabled=True)
scanner = DummyScanner(LOCAL_RAW_DATA)
await service.refresh_for_model_type(
"lora", scanner, PriceProvider(GATED_RESPONSE, prices=None)
)
assert service.count_unavailable_prices("lora") == 1
await service.refresh_for_model_type("lora", scanner, DummyProvider(FREE_RESPONSE))
record = await service.get_record("lora", 1)
assert record.versions[0].price_check_attempted_at is None
assert service.count_unavailable_prices("lora") == 0
def _long_ttl_service(tmp_path, **settings):
"""A service whose metadata TTL does not lapse during the test."""
return ModelUpdateService(
str(tmp_path / "updates.sqlite"),
ttl_seconds=86400,
settings_manager=FakeSettings(settings),
)
@pytest.mark.asyncio
async def test_prices_are_fetched_even_when_the_version_list_is_fresh(tmp_path):
"""Enabling price tracking must not wait for the metadata TTL.
The cached record already carries the gate, so a fresh version list is no
reason to skip the price: otherwise turning the feature on prices only the
handful of models that happened to need a metadata refresh that round.
"""
service = _long_ttl_service(tmp_path, price_tracking_enabled=False)
scanner = DummyScanner(LOCAL_RAW_DATA)
provider = PriceProvider(GATED_RESPONSE, prices=PRICE_PAYLOAD)
await service.refresh_for_model_type("lora", scanner, provider)
assert provider.price_calls == 0
metadata_calls = provider.calls
checked_at_before = (await service.get_record("lora", 1)).last_checked_at
service._settings.set("price_tracking_enabled", True)
await service.refresh_for_model_type("lora", scanner, provider)
record = await service.get_record("lora", 1)
version = next(v for v in record.versions if v.version_id == 12)
# The version list came from the cache this round ...
assert provider.calls == metadata_calls
# ... and the price was captured anyway.
assert provider.price_calls == 1
assert version.price_buzz == 250
assert version.price_checked_at is not None
assert version.price_check_attempted_at is not None
# A price-only pass must not extend the metadata TTL.
assert record.last_checked_at == checked_at_before
@pytest.mark.asyncio
async def test_failed_price_attempt_is_not_retried_within_the_ttl(tmp_path):
"""A mature model whose page no host will serve must not cost two requests
on every single update check."""
service = _long_ttl_service(tmp_path, price_tracking_enabled=True)
scanner = DummyScanner(LOCAL_RAW_DATA)
provider = PriceProvider(GATED_RESPONSE, prices=None)
await service.refresh_for_model_type("lora", scanner, provider)
assert provider.price_calls == 1
await service.refresh_for_model_type("lora", scanner, provider)
assert provider.price_calls == 1
assert service.count_unavailable_prices("lora") == 1
@pytest.mark.asyncio
async def test_priced_count_reports_known_prices(tmp_path):
service = _price_service(tmp_path, price_tracking_enabled=True)
scanner = DummyScanner(LOCAL_RAW_DATA)
provider = PriceProvider(GATED_RESPONSE, prices=PRICE_PAYLOAD)
assert service.count_priced_versions("lora") == 0
await service.refresh_for_model_type("lora", scanner, provider)
assert service.count_priced_versions("lora") == 1
assert service.count_priced_versions("checkpoint") == 0
@pytest.mark.asyncio
async def test_forced_refresh_reprices_within_the_ttl(tmp_path):
service = _long_ttl_service(tmp_path, price_tracking_enabled=True)
scanner = DummyScanner(LOCAL_RAW_DATA)
provider = PriceProvider(GATED_RESPONSE, prices=PRICE_PAYLOAD)
await service.refresh_for_model_type("lora", scanner, provider)
assert provider.price_calls == 1
await service.refresh_for_model_type(
"lora", scanner, provider, force_refresh=True
)
assert provider.price_calls == 2