mirror of
https://github.com/willmiao/ComfyUI-Lora-Manager.git
synced 2026-08-25 23:11:26 -03:00
fix(update): align update-check summary count with Updates filter scope (#1083)
This commit is contained in:
@@ -2654,10 +2654,20 @@ class ModelUpdateHandler:
|
|||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
same_base_scope = self._uses_same_base_update_scope()
|
||||||
|
|
||||||
serialized_records = []
|
serialized_records = []
|
||||||
for record in records.values():
|
for record in records.values():
|
||||||
has_update_fn = getattr(record, "has_update", None)
|
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_early_access=hide_early_access,
|
||||||
hide_paid=hide_paid,
|
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:
|
async def set_model_update_ignore(self, request: web.Request) -> web.Response:
|
||||||
payload = await self._read_json(request)
|
payload = await self._read_json(request)
|
||||||
model_id = self._normalize_model_id(payload.get("modelId"))
|
model_id = self._normalize_model_id(payload.get("modelId"))
|
||||||
|
|||||||
@@ -13,7 +13,7 @@ import sqlite3
|
|||||||
import time
|
import time
|
||||||
from dataclasses import dataclass, replace
|
from dataclasses import dataclass, replace
|
||||||
from datetime import datetime, timezone
|
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 .errors import RateLimitError, ResourceNotFoundError
|
||||||
from .settings_manager import get_settings_manager
|
from .settings_manager import get_settings_manager
|
||||||
@@ -250,6 +250,51 @@ class ModelUpdateRecord:
|
|||||||
|
|
||||||
return False
|
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:
|
class ModelUpdateService:
|
||||||
"""Persist and query remote model version metadata."""
|
"""Persist and query remote model version metadata."""
|
||||||
@@ -786,6 +831,11 @@ class ModelUpdateService:
|
|||||||
target_model_ids=target_filter,
|
target_model_ids=target_filter,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
local_base_models = await self._collect_local_version_bases(
|
||||||
|
scanner,
|
||||||
|
target_model_ids=target_filter,
|
||||||
|
)
|
||||||
|
|
||||||
results: Dict[int, ModelUpdateRecord] = {}
|
results: Dict[int, ModelUpdateRecord] = {}
|
||||||
prefetched: Dict[int, Mapping[Any, Any]] = {}
|
prefetched: Dict[int, Mapping[Any, Any]] = {}
|
||||||
|
|
||||||
@@ -838,6 +888,7 @@ class ModelUpdateService:
|
|||||||
force_refresh=force_refresh,
|
force_refresh=force_refresh,
|
||||||
prefetched_response=prefetched.get(model_id),
|
prefetched_response=prefetched.get(model_id),
|
||||||
all_local_version_ids=all_vids,
|
all_local_version_ids=all_vids,
|
||||||
|
local_base_models=local_base_models,
|
||||||
)
|
)
|
||||||
if scanner.is_cancelled():
|
if scanner.is_cancelled():
|
||||||
logger.info(f"{model_type.capitalize()} Update Service: Refresh cancelled by user")
|
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)
|
local_versions = await self._collect_local_versions(scanner)
|
||||||
version_ids = local_versions.get(model_id, [])
|
version_ids = local_versions.get(model_id, [])
|
||||||
|
local_base_models = await self._collect_local_version_bases(scanner)
|
||||||
return await self._refresh_single_model(
|
return await self._refresh_single_model(
|
||||||
model_type,
|
model_type,
|
||||||
model_id,
|
model_id,
|
||||||
version_ids,
|
version_ids,
|
||||||
metadata_provider,
|
metadata_provider,
|
||||||
force_refresh=force_refresh,
|
force_refresh=force_refresh,
|
||||||
|
local_base_models=local_base_models,
|
||||||
)
|
)
|
||||||
|
|
||||||
async def update_in_library_versions(
|
async def update_in_library_versions(
|
||||||
@@ -1053,6 +1106,7 @@ class ModelUpdateService:
|
|||||||
force_refresh: bool = False,
|
force_refresh: bool = False,
|
||||||
prefetched_response: Optional[Mapping[str, Any]] = None,
|
prefetched_response: Optional[Mapping[str, Any]] = None,
|
||||||
all_local_version_ids: Optional[Sequence[int]] = None,
|
all_local_version_ids: Optional[Sequence[int]] = None,
|
||||||
|
local_base_models: Optional[Mapping[int, str]] = None,
|
||||||
) -> Optional[ModelUpdateRecord]:
|
) -> Optional[ModelUpdateRecord]:
|
||||||
normalized_local = self._normalize_sequence(local_versions)
|
normalized_local = self._normalize_sequence(local_versions)
|
||||||
# When folder-filtering, this carries the cross-folder version set
|
# When folder-filtering, this carries the cross-folder version set
|
||||||
@@ -1177,6 +1231,7 @@ class ModelUpdateService:
|
|||||||
existing,
|
existing,
|
||||||
now,
|
now,
|
||||||
all_local_version_ids=normalized_all,
|
all_local_version_ids=normalized_all,
|
||||||
|
local_base_models=local_base_models,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
record = self._merge_with_local_versions(
|
record = self._merge_with_local_versions(
|
||||||
@@ -1383,27 +1438,17 @@ class ModelUpdateService:
|
|||||||
await self._enrich_version_entries(metadata_provider, aggregated)
|
await self._enrich_version_entries(metadata_provider, aggregated)
|
||||||
return aggregated
|
return aggregated
|
||||||
|
|
||||||
async def _collect_local_versions(
|
@staticmethod
|
||||||
self,
|
def _iter_local_civitai_items(
|
||||||
scanner,
|
cache,
|
||||||
*,
|
*,
|
||||||
target_model_ids: Optional[Sequence[int]] = None,
|
target_set: Optional[set[int]] = None,
|
||||||
folder_path: Optional[str] = None,
|
normalized_folder: Optional[str] = None,
|
||||||
) -> Dict[int, List[int]]:
|
) -> Iterator[tuple[int, int, Any]]:
|
||||||
cache = await scanner.get_cached_data()
|
"""Yield ``(modelId, versionId, base_model)`` for each scannable item."""
|
||||||
mapping: Dict[int, set[int]] = {}
|
|
||||||
if not cache or not getattr(cache, "raw_data", None):
|
if not cache or not getattr(cache, "raw_data", None):
|
||||||
return {}
|
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("/")
|
|
||||||
|
|
||||||
for item in cache.raw_data:
|
for item in cache.raw_data:
|
||||||
# Apply folder filter first (cheapest check)
|
# Apply folder filter first (cheapest check)
|
||||||
@@ -1423,10 +1468,75 @@ class ModelUpdateService:
|
|||||||
continue
|
continue
|
||||||
if target_set is not None and model_id not in target_set:
|
if target_set is not None and model_id not in target_set:
|
||||||
continue
|
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)
|
mapping.setdefault(model_id, set()).add(version_id)
|
||||||
|
|
||||||
return {model_id: sorted(ids) for model_id, ids in mapping.items()}
|
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(
|
def _merge_with_local_versions(
|
||||||
self,
|
self,
|
||||||
existing: Optional[ModelUpdateRecord],
|
existing: Optional[ModelUpdateRecord],
|
||||||
@@ -1506,6 +1616,7 @@ class ModelUpdateService:
|
|||||||
timestamp: float,
|
timestamp: float,
|
||||||
*,
|
*,
|
||||||
all_local_version_ids: Optional[Sequence[int]] = None,
|
all_local_version_ids: Optional[Sequence[int]] = None,
|
||||||
|
local_base_models: Optional[Mapping[int, str]] = None,
|
||||||
) -> ModelUpdateRecord:
|
) -> ModelUpdateRecord:
|
||||||
local_set = set(local_versions)
|
local_set = set(local_versions)
|
||||||
# When folder-filtering, also consider versions in other folders
|
# When folder-filtering, also consider versions in other folders
|
||||||
@@ -1552,6 +1663,7 @@ class ModelUpdateService:
|
|||||||
|
|
||||||
missing_local = local_set - seen_ids
|
missing_local = local_set - seen_ids
|
||||||
if missing_local:
|
if missing_local:
|
||||||
|
item_base_models = local_base_models or {}
|
||||||
for version_id in sorted(missing_local):
|
for version_id in sorted(missing_local):
|
||||||
existing_version = existing_map.get(version_id)
|
existing_version = existing_map.get(version_id)
|
||||||
if existing_version:
|
if existing_version:
|
||||||
@@ -1566,7 +1678,7 @@ class ModelUpdateService:
|
|||||||
ModelVersionRecord(
|
ModelVersionRecord(
|
||||||
version_id=version_id,
|
version_id=version_id,
|
||||||
name=None,
|
name=None,
|
||||||
base_model=None,
|
base_model=item_base_models.get(version_id),
|
||||||
released_at=None,
|
released_at=None,
|
||||||
size_bytes=None,
|
size_bytes=None,
|
||||||
preview_url=None,
|
preview_url=None,
|
||||||
|
|||||||
@@ -233,16 +233,26 @@ async def test_refresh_model_updates_filters_records_without_updates():
|
|||||||
model_type="lora",
|
model_type="lora",
|
||||||
model_id=1,
|
model_id=1,
|
||||||
versions=[
|
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(
|
ModelVersionRecord(
|
||||||
version_id=10,
|
version_id=10,
|
||||||
name="v1",
|
name="v1",
|
||||||
base_model=None,
|
base_model="Pony",
|
||||||
released_at=None,
|
released_at=None,
|
||||||
size_bytes=None,
|
size_bytes=None,
|
||||||
preview_url=None,
|
preview_url=None,
|
||||||
is_in_library=False,
|
is_in_library=False,
|
||||||
should_ignore=False,
|
should_ignore=False,
|
||||||
)
|
),
|
||||||
],
|
],
|
||||||
last_checked_at=None,
|
last_checked_at=None,
|
||||||
should_ignore_model=False,
|
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
|
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
|
@pytest.mark.asyncio
|
||||||
async def test_refresh_model_updates_with_target_ids():
|
async def test_refresh_model_updates_with_target_ids():
|
||||||
cache = SimpleNamespace(version_index={})
|
cache = SimpleNamespace(version_index={})
|
||||||
|
|||||||
@@ -187,6 +187,144 @@ def test_has_update_for_base_rejects_other_base_models():
|
|||||||
assert record.has_update_for_base(10, "Flux") is False
|
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
|
@pytest.mark.asyncio
|
||||||
async def test_refresh_persists_versions_and_uses_cache(tmp_path):
|
async def test_refresh_persists_versions_and_uses_cache(tmp_path):
|
||||||
db_path = tmp_path / "updates.sqlite"
|
db_path = tmp_path / "updates.sqlite"
|
||||||
|
|||||||
Reference in New Issue
Block a user