fix(update): align update-check summary count with Updates filter scope (#1083)

This commit is contained in:
Will Miao
2026-08-25 20:36:46 +08:00
parent 74f889f160
commit 08895f77ff
4 changed files with 558 additions and 24 deletions
+31 -1
View File
@@ -2654,10 +2654,20 @@ class ModelUpdateHandler:
except Exception:
pass
same_base_scope = self._uses_same_base_update_scope()
serialized_records = []
for record in records.values():
has_update_fn = getattr(record, "has_update", None)
if callable(has_update_fn) and has_update_fn(
if not callable(has_update_fn):
continue
scoped_fn = (
getattr(record, "has_update_for_local_bases", None)
if same_base_scope
else None
)
qualifies_fn = scoped_fn if callable(scoped_fn) else has_update_fn
if qualifies_fn(
hide_early_access=hide_early_access,
hide_paid=hide_paid,
):
@@ -2670,6 +2680,26 @@ class ModelUpdateHandler:
}
)
def _uses_same_base_update_scope(self) -> bool:
"""Return True when update reporting must honor same-base scoping.
Mirrors ``BaseModelService._annotate_update_flags``: the Updates filter
evaluates updates per local base model when ``version_grouping`` is
``same_base`` (its default). The refresh summary counts with the same
scope so the "Found N update(s)" toast matches what the filter
displays. See issue #1083.
"""
if self._settings is None:
return True
try:
strategy_value = self._settings.get("version_grouping")
except Exception:
return True
if isinstance(strategy_value, str) and strategy_value.strip():
return strategy_value.strip().lower() == "same_base"
return True
async def set_model_update_ignore(self, request: web.Request) -> web.Response:
payload = await self._read_json(request)
model_id = self._normalize_model_id(payload.get("modelId"))
+133 -21
View File
@@ -13,7 +13,7 @@ import sqlite3
import time
from dataclasses import dataclass, replace
from datetime import datetime, timezone
from typing import Any, Dict, Iterable, List, Mapping, Optional, Sequence
from typing import Any, Dict, Iterable, Iterator, List, Mapping, Optional, Sequence
from .errors import RateLimitError, ResourceNotFoundError
from .settings_manager import get_settings_manager
@@ -250,6 +250,51 @@ class ModelUpdateRecord:
return False
def has_update_for_local_bases(
self,
hide_early_access: bool = False,
hide_non_downloadable: bool = True,
hide_paid: bool = False,
) -> bool:
"""Return True when any locally-held base model scope has an update.
Aggregates :meth:`has_update_for_base` across every distinct base model
present among in-library versions. This mirrors the per-item evaluation
performed by ``BaseModelService._annotate_update_flags`` when the
``version_grouping`` setting is ``same_base``, so callers reporting
"how many models have updates" stay aligned with what the Updates
filter displays. Use this instead of :meth:`has_update` for such
summaries; see issue #1083.
When no local base model is known (nothing held locally, or versions
never seen in any remote listing), falls back to :meth:`has_update` so
a model the item-level filter may still flag is not silently dropped
from summaries.
"""
bases = {
_normalize_base_model(version.base_model)
for version in self.versions
if version.is_in_library
}
bases.discard(None)
if not bases:
return self.has_update(
hide_early_access=hide_early_access,
hide_non_downloadable=hide_non_downloadable,
hide_paid=hide_paid,
)
return any(
self.has_update_for_base(
None,
base,
hide_early_access=hide_early_access,
hide_non_downloadable=hide_non_downloadable,
hide_paid=hide_paid,
)
for base in bases
)
class ModelUpdateService:
"""Persist and query remote model version metadata."""
@@ -786,6 +831,11 @@ class ModelUpdateService:
target_model_ids=target_filter,
)
local_base_models = await self._collect_local_version_bases(
scanner,
target_model_ids=target_filter,
)
results: Dict[int, ModelUpdateRecord] = {}
prefetched: Dict[int, Mapping[Any, Any]] = {}
@@ -838,6 +888,7 @@ class ModelUpdateService:
force_refresh=force_refresh,
prefetched_response=prefetched.get(model_id),
all_local_version_ids=all_vids,
local_base_models=local_base_models,
)
if scanner.is_cancelled():
logger.info(f"{model_type.capitalize()} Update Service: Refresh cancelled by user")
@@ -872,12 +923,14 @@ class ModelUpdateService:
local_versions = await self._collect_local_versions(scanner)
version_ids = local_versions.get(model_id, [])
local_base_models = await self._collect_local_version_bases(scanner)
return await self._refresh_single_model(
model_type,
model_id,
version_ids,
metadata_provider,
force_refresh=force_refresh,
local_base_models=local_base_models,
)
async def update_in_library_versions(
@@ -1053,6 +1106,7 @@ class ModelUpdateService:
force_refresh: bool = False,
prefetched_response: Optional[Mapping[str, Any]] = None,
all_local_version_ids: Optional[Sequence[int]] = None,
local_base_models: Optional[Mapping[int, str]] = None,
) -> Optional[ModelUpdateRecord]:
normalized_local = self._normalize_sequence(local_versions)
# When folder-filtering, this carries the cross-folder version set
@@ -1177,6 +1231,7 @@ class ModelUpdateService:
existing,
now,
all_local_version_ids=normalized_all,
local_base_models=local_base_models,
)
else:
record = self._merge_with_local_versions(
@@ -1383,27 +1438,17 @@ class ModelUpdateService:
await self._enrich_version_entries(metadata_provider, aggregated)
return aggregated
async def _collect_local_versions(
self,
scanner,
@staticmethod
def _iter_local_civitai_items(
cache,
*,
target_model_ids: Optional[Sequence[int]] = None,
folder_path: Optional[str] = None,
) -> Dict[int, List[int]]:
cache = await scanner.get_cached_data()
mapping: Dict[int, set[int]] = {}
target_set: Optional[set[int]] = None,
normalized_folder: Optional[str] = None,
) -> Iterator[tuple[int, int, Any]]:
"""Yield ``(modelId, versionId, base_model)`` for each scannable item."""
if not cache or not getattr(cache, "raw_data", None):
return {}
target_set = None
if target_model_ids:
target_set = set(target_model_ids)
if not target_set:
return {}
normalized_folder = None
if folder_path is not None:
normalized_folder = folder_path.replace("\\", "/").strip("/")
return
for item in cache.raw_data:
# Apply folder filter first (cheapest check)
@@ -1423,10 +1468,75 @@ class ModelUpdateService:
continue
if target_set is not None and model_id not in target_set:
continue
yield model_id, version_id, item.get("base_model")
def _prepare_collection_filters(
self,
target_model_ids: Optional[Sequence[int]],
folder_path: Optional[str],
) -> tuple[Optional[set[int]], Optional[str]]:
target_set: Optional[set[int]] = None
if target_model_ids:
target_set = set(target_model_ids)
normalized_folder = None
if folder_path is not None:
normalized_folder = folder_path.replace("\\", "/").strip("/")
return target_set, normalized_folder
async def _collect_local_versions(
self,
scanner,
*,
target_model_ids: Optional[Sequence[int]] = None,
folder_path: Optional[str] = None,
) -> Dict[int, List[int]]:
cache = await scanner.get_cached_data()
mapping: Dict[int, set[int]] = {}
target_set, normalized_folder = self._prepare_collection_filters(
target_model_ids, folder_path
)
if target_model_ids and not target_set:
return {}
for model_id, version_id, _base_model in self._iter_local_civitai_items(
cache, target_set=target_set, normalized_folder=normalized_folder
):
mapping.setdefault(model_id, set()).add(version_id)
return {model_id: sorted(ids) for model_id, ids in mapping.items()}
async def _collect_local_version_bases(
self,
scanner,
*,
target_model_ids: Optional[Sequence[int]] = None,
) -> Dict[int, str]:
"""Map version id -> base model from cache items.
Deliberately unfiltered by folder: synthesized in-library entries must
carry a base regardless of which folder triggered the refresh.
"""
cache = await scanner.get_cached_data()
bases: Dict[int, str] = {}
target_set, _normalized_folder = self._prepare_collection_filters(
target_model_ids, None
)
if target_model_ids and not target_set:
return {}
for _model_id, version_id, base_model in self._iter_local_civitai_items(
cache, target_set=target_set
):
normalized_base = _normalize_string(base_model)
if normalized_base:
bases[version_id] = normalized_base
return bases
def _merge_with_local_versions(
self,
existing: Optional[ModelUpdateRecord],
@@ -1506,6 +1616,7 @@ class ModelUpdateService:
timestamp: float,
*,
all_local_version_ids: Optional[Sequence[int]] = None,
local_base_models: Optional[Mapping[int, str]] = None,
) -> ModelUpdateRecord:
local_set = set(local_versions)
# When folder-filtering, also consider versions in other folders
@@ -1552,6 +1663,7 @@ class ModelUpdateService:
missing_local = local_set - seen_ids
if missing_local:
item_base_models = local_base_models or {}
for version_id in sorted(missing_local):
existing_version = existing_map.get(version_id)
if existing_version:
@@ -1566,7 +1678,7 @@ class ModelUpdateService:
ModelVersionRecord(
version_id=version_id,
name=None,
base_model=None,
base_model=item_base_models.get(version_id),
released_at=None,
size_bytes=None,
preview_url=None,
+256 -2
View File
@@ -233,16 +233,26 @@ async def test_refresh_model_updates_filters_records_without_updates():
model_type="lora",
model_id=1,
versions=[
ModelVersionRecord(
version_id=8,
name="v0",
base_model="Pony",
released_at=None,
size_bytes=None,
preview_url=None,
is_in_library=True,
should_ignore=False,
),
ModelVersionRecord(
version_id=10,
name="v1",
base_model=None,
base_model="Pony",
released_at=None,
size_bytes=None,
preview_url=None,
is_in_library=False,
should_ignore=False,
)
),
],
last_checked_at=None,
should_ignore_model=False,
@@ -309,6 +319,250 @@ async def test_refresh_model_updates_filters_records_without_updates():
assert call["target_model_ids"] is None
@pytest.mark.asyncio
async def test_refresh_model_updates_same_base_scope_excludes_cross_base_updates():
"""Issue #1083: with version_grouping=same_base (the default), a newer
remote version targeting another base model must not be counted, matching
what the Updates filter displays."""
cache = SimpleNamespace(version_index={})
service = DummyService(cache)
cross_base_only = ModelUpdateRecord(
model_type="lora",
model_id=1,
versions=[
ModelVersionRecord(
version_id=5,
name="v0",
base_model="Pony",
released_at=None,
size_bytes=None,
preview_url=None,
is_in_library=True,
should_ignore=False,
),
ModelVersionRecord(
version_id=20,
name="v2",
base_model="Flux.1",
released_at=None,
size_bytes=None,
preview_url=None,
is_in_library=False,
should_ignore=False,
),
],
last_checked_at=None,
should_ignore_model=False,
)
update_service = DummyUpdateService({1: cross_base_only})
async def metadata_selector(name):
assert name == "civitai_api"
return object()
handler = ModelUpdateHandler(
service=service,
update_service=update_service,
metadata_provider_selector=metadata_selector,
settings_service=SimpleNamespace(get=lambda *_: False),
logger=logging.getLogger(__name__),
)
class DummyRequest:
can_read_body = True
query = {}
async def json(self):
return {}
response = await handler.refresh_model_updates(
DummyRequest() # pyright: ignore[reportArgumentType]
)
assert response.status == 200
text = response.text
assert text is not None
payload = json.loads(text)
assert payload["success"] is True
assert payload["records"] == []
@pytest.mark.asyncio
async def test_refresh_model_updates_any_grouping_counts_cross_base_updates():
"""With version_grouping=any the unscoped predicate applies, so a newer
remote version on any base model is counted."""
cache = SimpleNamespace(version_index={})
service = DummyService(cache)
cross_base_only = ModelUpdateRecord(
model_type="lora",
model_id=1,
versions=[
ModelVersionRecord(
version_id=5,
name="v0",
base_model="Pony",
released_at=None,
size_bytes=None,
preview_url=None,
is_in_library=True,
should_ignore=False,
),
ModelVersionRecord(
version_id=20,
name="v2",
base_model="Flux.1",
released_at=None,
size_bytes=None,
preview_url=None,
is_in_library=False,
should_ignore=False,
),
],
last_checked_at=None,
should_ignore_model=False,
)
update_service = DummyUpdateService({1: cross_base_only})
async def metadata_selector(name):
assert name == "civitai_api"
return object()
settings = SimpleNamespace(
get=lambda key, default=None: (
"any" if key == "version_grouping" else default
)
)
handler = ModelUpdateHandler(
service=service,
update_service=update_service,
metadata_provider_selector=metadata_selector,
settings_service=settings,
logger=logging.getLogger(__name__),
)
class DummyRequest:
can_read_body = True
query = {}
async def json(self):
return {}
response = await handler.refresh_model_updates(
DummyRequest() # pyright: ignore[reportArgumentType]
)
assert response.status == 200
text = response.text
assert text is not None
payload = json.loads(text)
assert payload["success"] is True
assert [record["modelId"] for record in payload["records"]] == [1]
def _make_cross_base_only_record() -> ModelUpdateRecord:
return ModelUpdateRecord(
model_type="lora",
model_id=1,
versions=[
ModelVersionRecord(
version_id=5,
name="v0",
base_model="Pony",
released_at=None,
size_bytes=None,
preview_url=None,
is_in_library=True,
should_ignore=False,
),
ModelVersionRecord(
version_id=20,
name="v2",
base_model="Flux.1",
released_at=None,
size_bytes=None,
preview_url=None,
is_in_library=False,
should_ignore=False,
),
],
last_checked_at=None,
should_ignore_model=False,
)
def _base_handler(update_service, settings_service):
async def metadata_selector(name):
assert name == "civitai_api"
return object()
return ModelUpdateHandler(
service=DummyService(SimpleNamespace(version_index={})),
update_service=update_service,
metadata_provider_selector=metadata_selector,
settings_service=settings_service,
logger=logging.getLogger(__name__),
)
@pytest.mark.asyncio
async def test_refresh_model_updates_explicit_same_base_setting_excludes_cross_base():
"""The literal "same_base" string (any casing/whitespace) selects the
scoped predicate, mirroring BaseModelService's strategy parsing."""
update_service = DummyUpdateService({1: _make_cross_base_only_record()})
settings = SimpleNamespace(
get=lambda key, default=None: (
" Same_Base " if key == "version_grouping" else default
)
)
handler = _base_handler(update_service, settings)
class DummyRequest:
can_read_body = True
query = {}
async def json(self):
return {}
response = await handler.refresh_model_updates(
DummyRequest() # pyright: ignore[reportArgumentType]
)
text = response.text
assert text is not None
payload = json.loads(text)
assert payload["records"] == []
@pytest.mark.asyncio
async def test_refresh_model_updates_falls_back_without_scoped_predicate():
"""A record type without has_update_for_local_bases (pre-change callers /
fakes) still counts via the unscoped predicate under same_base scope."""
legacy_record = _make_cross_base_only_record()
legacy_record.__dict__["has_update_for_local_bases"] = None
update_service = DummyUpdateService({1: legacy_record})
handler = _base_handler(update_service, SimpleNamespace(get=lambda *_: False))
class DummyRequest:
can_read_body = True
query = {}
async def json(self):
return {}
response = await handler.refresh_model_updates(
DummyRequest() # pyright: ignore[reportArgumentType]
)
text = response.text
assert text is not None
payload = json.loads(text)
assert payload["success"] is True
assert [record["modelId"] for record in payload["records"]] == [1]
@pytest.mark.asyncio
async def test_refresh_model_updates_with_target_ids():
cache = SimpleNamespace(version_index={})
+138
View File
@@ -187,6 +187,144 @@ def test_has_update_for_base_rejects_other_base_models():
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"