diff --git a/py/recipes/parsers/recipe_format.py b/py/recipes/parsers/recipe_format.py index 5ece07a7..4b378c1c 100644 --- a/py/recipes/parsers/recipe_format.py +++ b/py/recipes/parsers/recipe_format.py @@ -196,7 +196,7 @@ class RecipeFormatParser(RecipeMetadataParser): filtered_gen_params[key] = value return { - 'base_model': checkpoint['baseModel'] if checkpoint and checkpoint.get('baseModel') else recipe_metadata.get('base_model', ''), + 'base_model': checkpoint['baseModel'] if checkpoint and checkpoint.get('baseModel') else (recipe_metadata.get('base_model') or None), 'loras': loras, 'gen_params': filtered_gen_params, 'tags': recipe_metadata.get('tags', []), diff --git a/py/routes/handlers/recipe_handlers.py b/py/routes/handlers/recipe_handlers.py index fb520e16..8c035a1e 100644 --- a/py/routes/handlers/recipe_handlers.py +++ b/py/routes/handlers/recipe_handlers.py @@ -26,6 +26,7 @@ from ...services.recipes import ( RecipeValidationError, ) from ...services.metadata_service import get_default_metadata_provider +from ...services.recipe_scanner import UNKNOWN_BASE_MODEL_FILTER from ...utils.civitai_utils import ( build_civitai_image_page_url, extract_civitai_image_id, @@ -473,17 +474,32 @@ class RecipeQueryHandler: cache = await recipe_scanner.get_cached_data() base_model_counts: Dict[str, int] = {} + unknown_count = 0 for recipe in getattr(cache, "raw_data", []): base_model = recipe.get("base_model") if base_model: base_model_counts[base_model] = ( base_model_counts.get(base_model, 0) + 1 ) + else: + unknown_count += 1 sorted_models = [ {"name": model, "count": count} for model, count in base_model_counts.items() ] + if unknown_count: + # Synthetic "Unknown" bucket for recipes whose base model could + # not be determined. `value` carries the filter marker so the + # UI can display "Unknown" without colliding with real base + # model strings. + sorted_models.append( + { + "name": "Unknown", + "value": UNKNOWN_BASE_MODEL_FILTER, + "count": unknown_count, + } + ) sorted_models.sort(key=lambda entry: entry["count"], reverse=True) if limit > 0: sorted_models = sorted_models[:limit] diff --git a/py/services/recipe_scanner.py b/py/services/recipe_scanner.py index 217c89c0..968021e4 100644 --- a/py/services/recipe_scanner.py +++ b/py/services/recipe_scanner.py @@ -48,6 +48,12 @@ _CHECKPOINT_MODEL_TYPE_ALIASES = {"diffusionmodel": "diffusion_model"} # Valid LoRA availability statuses for the recipe listing filter. _VALID_LORA_AVAILABILITY_STATUSES = frozenset({"ready", "missing", "deleted"}) +# Filter marker for recipes whose base model could not be determined +# (base_model is None or empty). The UI displays "Unknown" for this bucket; +# the marker keeps the semantics explicit and disjoint from any real base +# model string. +UNKNOWN_BASE_MODEL_FILTER = "__unknown__" + class RecipeScanner: """Service for scanning and managing recipe images""" @@ -3530,11 +3536,23 @@ class RecipeScanner: if filters: # Filter by base model if "base_model" in filters and filters["base_model"]: - filtered_data = [ - item - for item in filtered_data - if item.get("base_model", "") in filters["base_model"] - ] + base_model_filter = filters["base_model"] + if UNKNOWN_BASE_MODEL_FILTER in base_model_filter: + # The unknown bucket matches recipes whose base model + # could not be determined (None/empty); real base + # models in the list still match by exact name. + filtered_data = [ + item + for item in filtered_data + if not item.get("base_model") + or item.get("base_model") in base_model_filter + ] + else: + filtered_data = [ + item + for item in filtered_data + if item.get("base_model", "") in base_model_filter + ] # Filter by favorite if "favorite" in filters and filters["favorite"]: diff --git a/static/js/managers/FilterManager.js b/static/js/managers/FilterManager.js index 5cc064e2..218bbfb9 100644 --- a/static/js/managers/FilterManager.js +++ b/static/js/managers/FilterManager.js @@ -511,18 +511,21 @@ export class FilterManager { filteredModels.forEach(model => { const tag = document.createElement('div'); tag.className = 'filter-tag base-model-tag'; - tag.dataset.baseModel = model.name; + // Display name may differ from the filter value (e.g. the "Unknown" + // bucket shows "Unknown" but filters via a dedicated marker). + const filterValue = model.value ?? model.name; + tag.dataset.baseModel = filterValue; tag.innerHTML = `${model.name} ${model.count}`; tag.addEventListener('click', async () => { tag.classList.toggle('active'); if (tag.classList.contains('active')) { - if (!this.filters.baseModel.includes(model.name)) { - this.filters.baseModel.push(model.name); + if (!this.filters.baseModel.includes(filterValue)) { + this.filters.baseModel.push(filterValue); } } else { - this.filters.baseModel = this.filters.baseModel.filter(m => m !== model.name); + this.filters.baseModel = this.filters.baseModel.filter(m => m !== filterValue); } this.updateActiveFiltersCount(); diff --git a/tests/frontend/components/pageControls.filtering.test.js b/tests/frontend/components/pageControls.filtering.test.js index 9ea8b0f5..408e0f19 100644 --- a/tests/frontend/components/pageControls.filtering.test.js +++ b/tests/frontend/components/pageControls.filtering.test.js @@ -350,6 +350,49 @@ describe('FilterManager tag and base model filters', () => { expect(baseModelChip.classList.contains('active')).toBe(false); }); + it('filters recipes by the unknown base model bucket via its marker value', async () => { + global.fetch = vi.fn().mockResolvedValue({ + ok: true, + json: async () => ({ + success: true, + base_models: [ + { name: 'Unknown', value: '__unknown__', count: 3 }, + { name: 'SDXL', count: 2 }, + ], + }), + }); + + renderControlsDom('recipes'); + const stateModule = await import('../../../static/js/state/index.js'); + stateModule.initPageState('recipes'); + const { getCurrentPageState } = stateModule; + const { FilterManager } = await import('../../../static/js/managers/FilterManager.js'); + + const loadRecipesMock = vi.fn().mockResolvedValue(undefined); + window.recipeManager = { loadRecipes: loadRecipesMock }; + + new FilterManager({ page: 'recipes' }); + + await vi.waitFor(() => { + const chip = document.querySelector('[data-base-model="__unknown__"]'); + expect(chip).not.toBeNull(); + }); + + const unknownChip = document.querySelector('[data-base-model="__unknown__"]'); + // Display label is "Unknown" even though the filter value is the marker + expect(unknownChip.textContent).toContain('Unknown'); + + unknownChip.dispatchEvent(new Event('click', { bubbles: true })); + await vi.waitFor(() => expect(loadRecipesMock).toHaveBeenCalledTimes(1)); + + expect(getCurrentPageState().filters.baseModel).toEqual(['__unknown__']); + expect(unknownChip.classList.contains('active')).toBe(true); + + const storageKey = 'lora_manager_recipes_filters'; + const storedFilters = JSON.parse(localStorage.getItem(storageKey)); + expect(storedFilters.baseModel).toEqual(['__unknown__']); + }); + it('filters base model chips locally without changing selected state', async () => { global.fetch = vi.fn().mockResolvedValue({ ok: true, diff --git a/tests/routes/test_recipe_query_handler.py b/tests/routes/test_recipe_query_handler.py index a897e9c4..96ed7ad1 100644 --- a/tests/routes/test_recipe_query_handler.py +++ b/tests/routes/test_recipe_query_handler.py @@ -5,6 +5,7 @@ from types import SimpleNamespace import pytest from py.routes.handlers.recipe_handlers import RecipeQueryHandler +from py.services.recipe_scanner import UNKNOWN_BASE_MODEL_FILTER async def _noop(): @@ -46,3 +47,42 @@ async def test_recipe_query_handler_base_models_limit_zero_returns_all(): {"name": "SDXL", "count": 2}, {"name": "LTXV 2.3", "count": 1}, ] + + +@pytest.mark.asyncio +async def test_recipe_query_handler_base_models_includes_unknown_bucket(): + cache = SimpleNamespace( + raw_data=[ + {"base_model": "SDXL"}, + {"base_model": None}, + {"base_model": ""}, + ] + ) + scanner = SimpleNamespace(get_cached_data=lambda: None) + + async def get_cached_data(): + return cache + + scanner.get_cached_data = get_cached_data + + handler = RecipeQueryHandler( + ensure_dependencies_ready=_noop, + recipe_scanner_getter=lambda: scanner, + format_recipe_file_url=lambda value: value, + logger=logging.getLogger(__name__), + ) + + response = await handler.get_base_models( + SimpleNamespace(query={"limit": "0"}) # pyright: ignore[reportArgumentType] + ) + text = response.text + assert text is not None + payload = json.loads(text) + + assert payload["success"] is True + # Unknown bucket carries a dedicated marker so the UI can show "Unknown" + # without colliding with real base model strings. + assert payload["base_models"] == [ + {"name": "Unknown", "value": UNKNOWN_BASE_MODEL_FILTER, "count": 2}, + {"name": "SDXL", "count": 1}, + ] diff --git a/tests/services/test_recipe_format_parser.py b/tests/services/test_recipe_format_parser.py index d1bb3e41..4fe11f72 100644 --- a/tests/services/test_recipe_format_parser.py +++ b/tests/services/test_recipe_format_parser.py @@ -111,6 +111,36 @@ async def test_recipe_format_parser_populates_checkpoint(monkeypatch): assert result["model"] == checkpoint +@pytest.mark.asyncio +async def test_recipe_format_parser_base_model_defaults_to_none_when_unknown(monkeypatch): + class _FakeScanner: + pass + + # No checkpoint and empty base_model -> unknown renders as None (not "") + result = await _parse( + monkeypatch, + {"title": "T", "base_model": "", "loras": [], "gen_params": {}}, + _FakeScanner(), + ) + assert result["base_model"] is None + + # Missing base_model key behaves the same + result = await _parse( + monkeypatch, + {"title": "T", "loras": [], "gen_params": {}}, + _FakeScanner(), + ) + assert result["base_model"] is None + + # A real base_model in recipe metadata is kept + result = await _parse( + monkeypatch, + {"title": "T", "base_model": "Illustrious", "loras": [], "gen_params": {}}, + _FakeScanner(), + ) + assert result["base_model"] == "Illustrious" + + @pytest.mark.asyncio async def test_recipe_format_parser_marks_lora_in_library_by_version(monkeypatch): async def fake_metadata_provider(): diff --git a/tests/services/test_recipe_scanner.py b/tests/services/test_recipe_scanner.py index 35229ca6..033e5fe3 100644 --- a/tests/services/test_recipe_scanner.py +++ b/tests/services/test_recipe_scanner.py @@ -12,7 +12,7 @@ from py.services import model_scanner as model_scanner_module from py.services.model_cache import ModelCache from py.services.model_hash_index import ModelHashIndex from py.services.model_scanner import CacheBuildResult, ModelScanner -from py.services.recipe_scanner import RecipeScanner +from py.services.recipe_scanner import RecipeScanner, UNKNOWN_BASE_MODEL_FILTER from py.services import settings_manager as settings_manager_module from py.utils.models import BaseModelMetadata from py.utils.utils import calculate_recipe_fingerprint @@ -1954,6 +1954,59 @@ async def test_get_paginated_data_filters_by_favorite(recipe_scanner): assert len(result_fav_false["items"]) == 2 +@pytest.mark.asyncio +async def test_get_paginated_data_filters_by_base_model_unknown_bucket(recipe_scanner): + scanner, _ = recipe_scanner + + await scanner.add_recipe( + { + "id": "known", + "file_path": "path/known.png", + "title": "Known Base Model", + "modified": 1.0, + "created_date": 1.0, + "base_model": "SDXL 1.0", + "loras": [], + } + ) + await scanner.add_recipe( + { + "id": "unknown", + "file_path": "path/unknown.png", + "title": "Unknown Base Model", + "modified": 2.0, + "created_date": 2.0, + "base_model": None, + "loras": [], + } + ) + + await asyncio.sleep(0) + await _wait_for_resort(scanner) + + # Exact-name filter matches only the recipe with that base model + result_known = await scanner.get_paginated_data( + page=1, page_size=10, filters={"base_model": ["SDXL 1.0"]} + ) + assert [item["id"] for item in result_known["items"]] == ["known"] + + # Unknown bucket matches recipes whose base model could not be determined + result_unknown = await scanner.get_paginated_data( + page=1, + page_size=10, + filters={"base_model": [UNKNOWN_BASE_MODEL_FILTER]}, + ) + assert [item["id"] for item in result_unknown["items"]] == ["unknown"] + + # Mixing known values with the unknown bucket matches both groups + result_both = await scanner.get_paginated_data( + page=1, + page_size=10, + filters={"base_model": ["SDXL 1.0", UNKNOWN_BASE_MODEL_FILTER]}, + ) + assert {item["id"] for item in result_both["items"]} == {"known", "unknown"} + + @pytest.mark.asyncio async def test_get_paginated_data_filters_by_prompt(recipe_scanner): scanner, _ = recipe_scanner