feat(update): track buzz prices and alert below a threshold

CivitAI's public API deliberately omits prices — paidAccess is trimmed to
{permanent, endsAt} because "pricing belongs to the purchase flow" — but the
public model page embeds the site's own model.getById result, including
paidAccess.terms, in its server-rendered payload. That is read anonymously
(no API key, no internal endpoint, no forged Origin), one request per gated
model, so only the ~2% of models that actually carry a gate pay for it.

- optional capture, off by default: price_tracking_enabled,
  price_alert_threshold_buzz (0 = alert on "became free" only) and
  price_check_ttl_hours; prices refresh on their own TTL and immediately when a
  gate changes, and a failed fetch keeps the stored price instead of blanking it
- versions that stop carrying a gate are marked free (persisted gate_lapsed_at)
  and gate transitions are reported as events on the refresh response, so a
  version already in the library can announce that it became free
- price_alert_state plus a price_drop edge event; new
  GET /api/lm/{type}/updates/price-alerts lists what is under the threshold
- versions tab shows the price (effective, with the list price struck through
  and a Blue Buzz note) and a Free Now badge; an update check toasts the
  transitions in one message
- the parser and the alerts query are unit-tested against a trimmed page
  fixture, and every route definition is now asserted to resolve to a handler

Plan, verification notes and the deviations from it are in
docs/plans/paid-model-price-tracking.md.
This commit is contained in:
Will Miao
2026-10-04 08:53:24 +08:00
parent 5f4054265d
commit ad2402724b
31 changed files with 2829 additions and 33 deletions
+149
View File
@@ -10,9 +10,11 @@ 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
@@ -324,6 +326,105 @@ async def test_refresh_model_updates_filters_records_without_updates():
assert call["target_model_ids"] is None
@pytest.mark.asyncio
async def test_get_price_alerts_returns_flagged_versions():
"""The price alert list is a GET (companion-extension friendly) and reports
the threshold it was computed with."""
class AlertsService:
def __init__(self):
self.calls = []
async def get_price_alerts(self, model_type, limit=200):
self.calls.append((model_type, limit))
return [{"modelId": 1, "versionId": 12, "priceBuzz": 250}]
update_service = AlertsService()
settings = {"price_tracking_enabled": True, "price_alert_threshold_buzz": 300}
handler = ModelUpdateHandler(
service=DummyService(SimpleNamespace(version_index={})),
update_service=update_service,
metadata_provider_selector=lambda *_: None,
settings_service=SimpleNamespace(get=lambda key, default=None: settings.get(key, default)),
logger=logging.getLogger(__name__),
)
request = SimpleNamespace(query={"limit": "50"})
response = await handler.get_price_alerts(request) # pyright: ignore[reportArgumentType]
assert response.status == 200
payload = json.loads(response.text)
assert payload["success"] is True
assert payload["enabled"] is True
assert payload["thresholdBuzz"] == 300
assert payload["alerts"] == [{"modelId": 1, "versionId": 12, "priceBuzz": 250}]
assert update_service.calls == [("lora", 50)]
@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
@@ -1077,3 +1178,51 @@ async def test_relink_civitai_surfaces_provider_unavailable_without_500():
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_get_price_alerts_is_registered_as_a_route():
"""The alerts endpoint must be declared and reachable via GET."""
definition = next(
d
for d in COMMON_ROUTE_DEFINITIONS
if d.handler_name == "get_price_alerts"
)
assert definition.method == "GET"
assert definition.build_path("loras") == "/api/lm/loras/updates/price-alerts"
+62
View File
@@ -859,3 +859,65 @@ async def test_get_version_file_mini_propagates_rate_limit(downloader):
with pytest.raises(RateLimitError):
await client.get_version_file_mini(1, 2)
async def test_get_model_prices_parses_public_page(downloader):
"""Prices come from the page payload, read anonymously (no API key)."""
client = await CivitaiClient.get_instance()
page_html = (
'<script id="__NEXT_DATA__" type="application/json">'
'{"props":{"pageProps":{"trpcState":{"json":{"queries":['
'{"queryKey":[["model","getById"],{"input":{"id":7}}],"state":{"data":'
'{"modelVersions":[{"id":42,"paidAccess":{"endsAt":null,'
'"timeframeDays":null,"terms":{"download":{"price":5000}},'
'"sale":null}}]}}}]}}}}}</script>'
)
async def fake_make_request(method, url, use_auth=True, **kwargs):
assert method == "GET"
assert "/models/7" in url
assert use_auth is False
assert kwargs.get("custom_headers", {}).get("Accept") == "text/html"
return True, page_html
downloader.make_request = fake_make_request
result = await client.get_model_prices(7)
assert result is not None
assert result[42]["price_buzz"] == 5000
async def test_get_model_prices_returns_none_on_unusable_page(downloader):
client = await CivitaiClient.get_instance()
async def fake_make_request(method, url, use_auth=True, **kwargs):
return True, "<html><body>challenge</body></html>"
downloader.make_request = fake_make_request
assert await client.get_model_prices(7) is None
async def test_get_model_prices_rejects_json_body(downloader):
"""A JSON response is not the page; it must not be parsed as one."""
client = await CivitaiClient.get_instance()
async def fake_make_request(method, url, use_auth=True, **kwargs):
return True, {"error": "nope"}
downloader.make_request = fake_make_request
assert await client.get_model_prices(7) is None
async def test_get_model_prices_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_model_prices(7)
+631
View File
@@ -1,5 +1,6 @@
import logging
import sqlite3
from dataclasses import replace
from types import SimpleNamespace
import pytest
@@ -1092,3 +1093,633 @@ def test_build_record_from_remote_preserves_file_count(tmp_path):
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)
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") == []
@@ -0,0 +1,108 @@
<!DOCTYPE html>
<!--
Trimmed stand-in for a CivitAI model page, mirroring the shape verified against
civitai.com / civitai.red. Values are synthetic; the structure (unrelated query
first, then the model.getById query, tombstones included) is what matters.
-->
<html>
<head>
<title>Fixture model</title>
</head>
<body>
<script id="__NEXT_DATA__" type="application/json">
{
"props": {
"pageProps": {
"trpcState": {
"json": {
"queries": [
{
"queryKey": [["common", "getEntityAccess"], { "entityId": 1001 }],
"state": { "data": [] }
},
{
"queryKey": [
["model", "getById"],
{ "input": { "id": 4242, "browsingLevel": 1 }, "type": "query" }
],
"state": {
"data": {
"id": 4242,
"name": "Fixture model",
"modelVersions": [
{
"id": 1001,
"name": "permanent paid",
"canDownload": false,
"paidAccess": {
"endsAt": null,
"timeframeDays": null,
"terms": {
"download": { "price": 5000 },
"generation": { "price": 100, "trialLimit": 5 }
},
"sale": null
}
},
{
"id": 1002,
"name": "timed gate on sale",
"canDownload": false,
"paidAccess": {
"endsAt": "2999-10-10T13:10:17.404Z",
"timeframeDays": 7,
"terms": {
"download": { "price": 125 },
"generation": { "trialLimit": 10 },
"acceptsBlueBuzz": true
},
"sale": {
"listTerms": { "download": { "price": 125 } },
"buyerTerms": { "download": { "price": 100 } },
"endsAt": "2999-09-01T00:00:00.000Z",
"discountType": "percent",
"discountAmount": 20
}
}
},
{
"id": 1003,
"name": "timed gate with no recorded end",
"canDownload": false,
"paidAccess": {
"endsAt": null,
"timeframeDays": 7,
"terms": { "download": { "price": 300 } },
"sale": null
}
},
{
"id": 1004,
"name": "lapsed gate (tombstone)",
"canDownload": true,
"paidAccess": {
"endsAt": "2000-08-27T16:56:57.438Z",
"timeframeDays": 3,
"terms": { "download": { "price": 200 } },
"sale": null
}
},
{
"id": 1005,
"name": "never gated",
"canDownload": true,
"paidAccess": null
}
]
}
}
}
]
}
}
}
}
}
</script>
</body>
</html>
+91
View File
@@ -0,0 +1,91 @@
"""Tests for the CivitAI model page price parser.
The fixture mirrors a real page payload; the cases that matter are the ones that
decide whether a user gets a price at all, and whether a lapsed gate is mistaken
for a live one.
"""
from pathlib import Path
import pytest
from py.utils.civitai_page_prices import MAX_PAGE_BYTES, parse_model_page_prices
FIXTURE = Path(__file__).parent / "fixtures" / "civitai_model_page_paid.html"
@pytest.fixture(scope="module")
def fixture_html() -> str:
return FIXTURE.read_text(encoding="utf-8")
def test_parses_prices_for_gated_versions(fixture_html):
prices = parse_model_page_prices(fixture_html)
assert prices is not None
# Lapsed gates and never-gated versions must not appear.
assert set(prices) == {1001, 1002, 1003}
def test_permanent_gate_price(fixture_html):
prices = parse_model_page_prices(fixture_html)
permanent = prices[1001]
assert permanent["price_buzz"] == 5000
assert permanent["list_price_buzz"] == 5000
assert permanent["generation_price_buzz"] == 100
assert permanent["accepts_blue_buzz"] is False
assert permanent["price_sale_ends_at"] is None
def test_timed_gate_uses_sale_price_as_effective(fixture_html):
prices = parse_model_page_prices(fixture_html)
timed = prices[1002]
# The stored/list price stays visible for a strikethrough, the effective price
# is what a buyer pays now.
assert timed["price_buzz"] == 100
assert timed["list_price_buzz"] == 125
assert timed["accepts_blue_buzz"] is True
assert timed["price_sale_ends_at"] == "2999-09-01T00:00:00.000Z"
# Generation is bundled with the download tier here (`trialLimit` only).
assert timed["generation_price_buzz"] is None
def test_timed_gate_without_recorded_end_is_still_priced(fixture_html):
prices = parse_model_page_prices(fixture_html)
assert prices[1003]["price_buzz"] == 300
assert prices[1003]["generation_price_buzz"] is None
@pytest.mark.parametrize(
"html",
[
None,
"",
"<html><body>no payload here</body></html>",
'<script id="__NEXT_DATA__" type="application/json">{not json</script>',
# Valid JSON, but not the shape we need.
'<script id="__NEXT_DATA__" type="application/json">{"props":{}}</script>',
# Model query present, but no versions.
'<script id="__NEXT_DATA__" type="application/json">'
'{"props":{"pageProps":{"trpcState":{"json":{"queries":['
'{"queryKey":[["model","getById"]],"state":{"data":{}}}]}}}}}</script>',
],
)
def test_unusable_payloads_return_none(html):
assert parse_model_page_prices(html) is None
def test_oversized_page_is_skipped():
oversized = "x" * (MAX_PAGE_BYTES + 1)
assert parse_model_page_prices(oversized) is None
def test_unrelated_first_query_does_not_win(fixture_html):
"""Query order varies per page; selection is by procedure name."""
prices = parse_model_page_prices(fixture_html)
assert prices is not None
assert 1001 in prices