mirror of
https://github.com/willmiao/ComfyUI-Lora-Manager.git
synced 2026-10-05 09:35:31 -03:00
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:
@@ -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"
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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>
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user