diff --git a/py/routes/handlers/model_handlers.py b/py/routes/handlers/model_handlers.py index 6b26c652..a09b2ffd 100644 --- a/py/routes/handlers/model_handlers.py +++ b/py/routes/handlers/model_handlers.py @@ -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")) diff --git a/py/services/model_update_service.py b/py/services/model_update_service.py index 3c3b12bc..00147828 100644 --- a/py/services/model_update_service.py +++ b/py/services/model_update_service.py @@ -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, diff --git a/tests/routes/test_model_update_handler.py b/tests/routes/test_model_update_handler.py index 3be9ff81..0837ae05 100644 --- a/tests/routes/test_model_update_handler.py +++ b/tests/routes/test_model_update_handler.py @@ -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={}) diff --git a/tests/services/test_model_update_service.py b/tests/services/test_model_update_service.py index 04e72ad7..61d48bb2 100644 --- a/tests/services/test_model_update_service.py +++ b/tests/services/test_model_update_service.py @@ -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"