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,