diff --git a/app/assets/helpers.py b/app/assets/helpers.py index 05876972a2c..b83a6855ab0 100644 --- a/app/assets/helpers.py +++ b/app/assets/helpers.py @@ -37,14 +37,51 @@ def sql_path_under_prefix( and normalizes there. Normalizing the column in SQL is not an option anyway: it would need a per-row Python call and would defeat the index. """ - base = os.path.abspath(prefix) - stem = base if base.endswith(os.sep) else base + os.sep + base, stem = _base_and_stem(prefix) return sa.or_( column == base, sa.func.substr(column, 1, len(stem)) == stem, ) +def _base_and_stem(prefix: str) -> tuple[str, str]: + base = os.path.abspath(prefix) + return base, base if base.endswith(os.sep) else base + os.sep + + +def stored_path_under_prefixes(prefixes: list[str]) -> Callable[[str], bool]: + """The Python twin of sql_path_under_prefix OR'd over ``prefixes``: the same + case-sensitive string test on a stored path, for filtering rows already fetched. + + Unlike path_prefix_matcher, it neither normalizes nor normcases the path. + """ + pairs = [_base_and_stem(prefix) for prefix in prefixes] + bases = frozenset(base for base, _ in pairs) + stems = tuple(stem for _, stem in pairs) + return lambda path: path in bases or path.startswith(stems) + + +# Each prefix adds two terms to one flat OR, and SQLite rejects an expression deeper than +# 1000, so about 500 prefixes in one statement fail with "Expression tree is too large". +# SQLAlchemy flattens nested ORs, so more prefixes than this are split across statements, +# or filtered in Python with stored_path_under_prefixes. +PREFIX_BATCH_SIZE = 200 + + +def sql_path_under_prefix_batches( + column: ColumnElement[str], prefixes: list[str] +) -> list[ColumnElement[bool]]: + """sql_path_under_prefix OR'd over each run of at most PREFIX_BATCH_SIZE prefixes. + + Run one statement per predicate and merge. Nested or overlapping prefixes can put a + row in more than one batch, so the caller dedupes. + """ + return [ + sa.or_(*(sql_path_under_prefix(column, p) for p in prefixes[i:i + PREFIX_BATCH_SIZE])) + for i in range(0, len(prefixes), PREFIX_BATCH_SIZE) + ] + + def path_prefix_matcher(prefixes: Iterable[str]) -> Callable[[str], bool]: """Return ``path -> Path(path).is_relative_to()``, with the prefixes normalized once. diff --git a/app/assets/scanner.py b/app/assets/scanner.py index 152caae254a..ec886cb2cfd 100644 --- a/app/assets/scanner.py +++ b/app/assets/scanner.py @@ -12,6 +12,7 @@ import logging import os from dataclasses import dataclass +from itertools import islice from pathlib import Path from typing import Callable, Literal, NamedTuple, Protocol, TypedDict @@ -28,7 +29,14 @@ create_record, ) from app.assets.database.models import Asset, AssetContent -from app.assets.helpers import path_prefix_matcher, sql_path_under_prefix, to_stored_hash +from app.assets.helpers import ( + PREFIX_BATCH_SIZE, + path_prefix_matcher, + sql_path_under_prefix, + sql_path_under_prefix_batches, + stored_path_under_prefixes, + to_stored_hash, +) from app.assets.lifecycle import get_excluded_scan_roots from app.assets.scanner_changes import ( clear_pending_verifications, @@ -379,19 +387,21 @@ def live_references_safely(root: RootType) -> dict[str, list[_ReferenceObservati live: dict[str, list[_ReferenceObservation]] = {} if not prefixes: return live - stmt = sa.select( - AssetContent.id, AssetContent.path, AssetContent.size_bytes, AssetContent.mtime_ns - ).where( - AssetContent.is_missing.is_(False), - sa.or_(*(sql_path_under_prefix(AssetContent.path, p) for p in prefixes)), - ) + seen: set[str] = set() try: with create_session() as session: - for content_id, path, size_bytes, mtime_ns in session.execute(stmt): - yield_gil(run=RESCAN_YIELD_RUN) - live.setdefault(os.path.abspath(path), []).append( - _ReferenceObservation(content_id, size_bytes, mtime_ns, None) - ) + for under_prefixes in sql_path_under_prefix_batches(AssetContent.path, prefixes): + stmt = sa.select( + AssetContent.id, AssetContent.path, AssetContent.size_bytes, AssetContent.mtime_ns + ).where(AssetContent.is_missing.is_(False), under_prefixes) + for content_id, path, size_bytes, mtime_ns in session.execute(stmt): + yield_gil(run=RESCAN_YIELD_RUN) + if content_id in seen: + continue + seen.add(content_id) + live.setdefault(os.path.abspath(path), []).append( + _ReferenceObservation(content_id, size_bytes, mtime_ns, None) + ) except Exception as exc: logging.exception("fast DB scan failed for %s: %s", root, exc) emit("scanner.fast_scan_failed", root=root, error_type=error_type(exc)) @@ -745,12 +755,10 @@ def insert_asset_specs( return created, first_error -def build_unenriched_candidates_statement( - prefixes: list[str], - compute_hashes: bool, - last_seen_id: str | None, - limit: int = 1000, +def unenriched_candidates_query( + compute_hashes: bool, last_seen_id: str | None ) -> sa.Select[tuple[str, str, str]]: + """Every unenriched live candidate after ``last_seen_id``, in id order.""" query = ( sa.select(AssetContent.id, Asset.id, AssetContent.path) .join(Asset, Asset.content_id == AssetContent.id) @@ -767,11 +775,19 @@ def build_unenriched_candidates_statement( query = query.where(Asset.system_metadata.is_(None)) if last_seen_id is not None: query = query.where(Asset.id > last_seen_id) + return query.order_by(Asset.id.asc()) + + +def build_unenriched_candidates_statement( + prefixes: list[str], + compute_hashes: bool, + last_seen_id: str | None, + limit: int = 1000, +) -> sa.Select[tuple[str, str, str]]: + """The next page of candidates under at most PREFIX_BATCH_SIZE prefixes.""" return ( - query.where( - sa.or_(*(sql_path_under_prefix(AssetContent.path, p) for p in prefixes)) - ) - .order_by(Asset.id.asc()) + unenriched_candidates_query(compute_hashes, last_seen_id) + .where(sa.or_(*(sql_path_under_prefix(AssetContent.path, p) for p in prefixes))) .limit(limit) ) @@ -789,14 +805,24 @@ def get_unenriched_assets_for_roots( if not prefixes: return [] - query = build_unenriched_candidates_statement( - prefixes, - compute_hashes, - last_seen_id, - limit, - ) with create_session() as sess: - rows = sess.execute(query).all() + if len(prefixes) <= PREFIX_BATCH_SIZE: + statement = build_unenriched_candidates_statement( + prefixes, + compute_hashes, + last_seen_id, + limit, + ) + rows = sess.execute(statement).all() + else: + # Too many prefixes for one SQL predicate. Paging each batch separately + # would rescan to the end of the table on every page for any batch with + # few matches, so filter a single id-ordered pass here instead. + is_under = stored_path_under_prefixes(prefixes) + candidates = sess.execute( + unenriched_candidates_query(compute_hashes, last_seen_id).execution_options(yield_per=500) + ) + rows = list(islice((row for row in candidates if is_under(row[2])), limit)) return [ UnenrichedContent(content_id, record_id, file_path) diff --git a/app/assets/scanner_changes.py b/app/assets/scanner_changes.py index bed169d4abe..9f5ddd45202 100644 --- a/app/assets/scanner_changes.py +++ b/app/assets/scanner_changes.py @@ -9,7 +9,7 @@ from __future__ import annotations import os -from collections.abc import Iterable +from collections.abc import Iterator from typing import Literal import sqlalchemy as sa @@ -22,7 +22,7 @@ mark_content_missing, unset_content_missing, ) -from app.assets.helpers import path_prefix_matcher, sql_path_under_prefix, to_stored_hash +from app.assets.helpers import path_prefix_matcher, sql_path_under_prefix_batches, to_stored_hash from app.assets.services.file_utils import get_mtime_ns from app.assets.services.path_utils import compute_loader_path, get_name_and_tags_from_asset_path from app.assets.services.snapshot_hash import snapshot_hash @@ -259,12 +259,12 @@ def drain_pending_verifications(session: Session, limit: int | None = None) -> i return processed -def live_contents_under_prefixes(session: Session, prefixes: list[str]) -> Iterable[AssetContent]: +def live_contents_under_prefixes(session: Session, prefixes: list[str]) -> Iterator[AssetContent]: """Stream the live contents under the prefixes in batches; consume it inside the session.""" - if not prefixes: - return [] - stmt = sa.select(AssetContent).where( - AssetContent.is_missing.is_(False), - sa.or_(*(sql_path_under_prefix(AssetContent.path, prefix) for prefix in prefixes)), - ) - return session.scalars(stmt.execution_options(yield_per=500)) + seen: set[str] = set() + for under_prefixes in sql_path_under_prefix_batches(AssetContent.path, prefixes): + stmt = sa.select(AssetContent).where(AssetContent.is_missing.is_(False), under_prefixes) + for content in session.scalars(stmt.execution_options(yield_per=500)): + if content.id not in seen: + seen.add(content.id) + yield content diff --git a/tests-unit/assets_test/services/test_many_scan_prefixes.py b/tests-unit/assets_test/services/test_many_scan_prefixes.py new file mode 100644 index 00000000000..4c2308b29a5 --- /dev/null +++ b/tests-unit/assets_test/services/test_many_scan_prefixes.py @@ -0,0 +1,370 @@ +"""Prefix filters over hundreds of scan folders. + +Each prefix adds terms to one OR, and SQLite rejects an expression tree deeper than +1000, so a single statement over about 500 prefixes failed every scan. The filters now +run in batches; these pin that the batched results equal the single-statement ones. +""" + +from __future__ import annotations + +import logging +import ntpath +import os +import sqlite3 +from contextlib import contextmanager +from pathlib import Path +from types import SimpleNamespace +from unittest.mock import patch + +import pytest +import sqlalchemy as sa +from sqlalchemy.orm import Session, sessionmaker + +from app.assets import helpers, scanner, seeder as seeder_module +from app.assets.database.models import Asset, AssetContent, Base +from app.assets.database.queries import create_content, create_record, mark_content_missing +from app.assets.helpers import ( + PREFIX_BATCH_SIZE, + sql_path_under_prefix, + sql_path_under_prefix_batches, + stored_path_under_prefixes, +) +from app.assets.scanner import get_unenriched_assets_for_roots, live_references_safely +from app.assets.scanner_changes import live_contents_under_prefixes +from app.database import db + +from .path_prefix_cases import anchor_case_paths, expected_prefix_case_paths, prefix_case_paths + +PREFIX_COUNTS = [1, PREFIX_BATCH_SIZE, PREFIX_BATCH_SIZE + 1, 499, 500, 2000] + +# SQLite before 3.32 allows 999 bound variables per statement, and each prefix binds four. +OLD_SQLITE_VARIABLE_LIMIT = 999 + + +def _capped_engine(url: str) -> sa.Engine: + """An engine whose connections allow only as many bound variables as old SQLite.""" + if not hasattr(sqlite3.Connection, "setlimit"): + pytest.skip("sqlite3.Connection.setlimit needs Python 3.11") + engine = sa.create_engine(url) + + @sa.event.listens_for(engine, "connect") + def _cap_variables(dbapi_connection, _record): + dbapi_connection.setlimit(sqlite3.SQLITE_LIMIT_VARIABLE_NUMBER, OLD_SQLITE_VARIABLE_LIMIT) + + Base.metadata.create_all(engine) + return engine + + +@pytest.fixture +def db_engine(): + """Every test here runs under the old variable limit as well as the depth limit.""" + engine = _capped_engine("sqlite:///:memory:") + yield engine + engine.dispose() + + +@contextmanager +def _reuse_session(session: Session): + yield session + + +def _folders(root: Path, count: int) -> list[str]: + return [str(root / f"b{i:05d}" / "checkpoints") for i in range(count)] + + +def _seed(session: Session, prefixes: list[str]) -> set[str]: + """One live row under each prefix, plus decoys none of them match.""" + inside = set() + for prefix in dict.fromkeys(prefixes): + path = os.path.join(prefix, "model.safetensors") + create_record(session, create_content(session, path).id, "model.safetensors") + inside.add(path) + create_content(session, prefix + "-other" + os.sep + "decoy.safetensors") + missing = create_content(session, os.path.join(prefixes[-1], "gone.safetensors")) + mark_content_missing(session, missing.id) + session.commit() + return inside + + +def test_batches_split_the_prefixes_in_order(): + prefixes = [f"/p{i}" for i in range(2 * PREFIX_BATCH_SIZE + 1)] + slices = [prefixes[:PREFIX_BATCH_SIZE], prefixes[PREFIX_BATCH_SIZE:-1], prefixes[-1:]] + compile_kwargs = {"literal_binds": True} + + batches = sql_path_under_prefix_batches(AssetContent.path, prefixes) + + assert [str(b.compile(compile_kwargs=compile_kwargs)) for b in batches] == [ + str(sa.or_(*(sql_path_under_prefix(AssetContent.path, p) for p in s)).compile(compile_kwargs=compile_kwargs)) + for s in slices + ] + assert sql_path_under_prefix_batches(AssetContent.path, []) == [] + + +def _stored_path_cases(root: str) -> list[str]: + paths = [os.path.abspath(path) for path, _ in prefix_case_paths(root)] + return paths + [path for path, _, _ in anchor_case_paths()] + ["/", root + os.sep] + + +@pytest.mark.parametrize("prefix_kind", ["dir", "trailing_sep", "root", "double_anchor"]) +def test_stored_path_under_prefixes_agrees_with_the_sql_predicate(session, temp_dir, prefix_kind): + root = str(temp_dir / "root") + prefix = {"dir": root, "trailing_sep": root + os.sep, "root": "/", "double_anchor": "//server"}[prefix_kind] + paths = _stored_path_cases(root) + session.add_all(AssetContent(path=path, size_bytes=0) for path in dict.fromkeys(paths)) + session.commit() + + selected = set( + session.scalars(sa.select(AssetContent.path).where(sql_path_under_prefix(AssetContent.path, prefix))) + ) + + is_under = stored_path_under_prefixes([prefix]) + assert {path for path in paths if is_under(path)} == selected + + +@pytest.mark.parametrize("count", [1, 20, PREFIX_BATCH_SIZE]) +def test_up_to_one_batch_compiles_to_the_single_statement_predicate(count): + prefixes = [f"/p{i}" for i in range(count)] + single = sa.or_(*(sql_path_under_prefix(AssetContent.path, p) for p in prefixes)) + + [batch] = sql_path_under_prefix_batches(AssetContent.path, prefixes) + + compile_kwargs = {"literal_binds": True} + assert str(batch.compile(compile_kwargs=compile_kwargs)) == str( + single.compile(compile_kwargs=compile_kwargs) + ) + + +def test_a_single_statement_over_500_prefixes_exceeds_sqlite_limits(session, temp_dir): + """The limits the batching works around: expression depth on any SQLite, and first + the variable limit on old SQLite, which the capped engine here emulates.""" + prefixes = _folders(temp_dir, 500) + stmt = sa.select(AssetContent.id).where( + sa.or_(*(sql_path_under_prefix(AssetContent.path, p) for p in prefixes)) + ) + + with pytest.raises( + sa.exc.OperationalError, match="Expression tree is too large|too many SQL variables" + ): + session.execute(stmt).all() + + +def test_the_variable_cap_is_in_force(session): + """Guards the cap above: without it, these tests would not cover old SQLite.""" + too_many = sa.select(AssetContent.id).where(AssetContent.path.in_([str(i) for i in range(1000)])) + + with pytest.raises(sa.exc.OperationalError, match="too many SQL variables"): + session.execute(too_many).all() + + +@pytest.mark.parametrize("count", PREFIX_COUNTS) +def test_live_contents_under_many_prefixes(session, temp_dir, count): + prefixes = _folders(temp_dir, count) + inside = _seed(session, prefixes) + + returned = [content.path for content in live_contents_under_prefixes(session, prefixes)] + + assert sorted(returned) == sorted(inside) + + +@pytest.mark.parametrize("count", PREFIX_COUNTS) +def test_live_references_under_many_prefixes(session, temp_dir, count): + prefixes = _folders(temp_dir, count) + inside = _seed(session, prefixes) + + with ( + patch.object(scanner, "create_session", lambda: _reuse_session(session)), + patch.object(scanner, "get_scan_prefixes_for_root", lambda _root: prefixes), + ): + live = live_references_safely("models") + + assert set(live) == inside + assert all(len(observations) == 1 for observations in live.values()) + + +@pytest.mark.parametrize("count", PREFIX_COUNTS) +def test_unenriched_candidates_under_many_prefixes(session, temp_dir, count): + prefixes = _folders(temp_dir, count) + inside = _seed(session, prefixes) + + with ( + patch.object(scanner, "create_session", lambda: _reuse_session(session)), + patch.object(scanner, "get_scan_prefixes_for_root", lambda _root: prefixes), + ): + rows = get_unenriched_assets_for_roots(("models",), compute_hashes=False, limit=10_000) + + assert sorted(row.file_path for row in rows) == sorted(inside) + + +def _nested_prefixes(temp_dir: Path) -> tuple[list[str], set[str]]: + """Duplicate and nested prefixes whose matches straddle batch boundaries.""" + prefixes = _folders(temp_dir, 2 * PREFIX_BATCH_SIZE + 50) + outer = str(temp_dir / "shared") + inner = str(temp_dir / "shared" / "inner") + prefixes[0] = outer + prefixes[PREFIX_BATCH_SIZE + 3] = inner + prefixes[-1] = outer + return prefixes, {outer, inner} + + +def test_nested_prefixes_across_batches_yield_each_row_once(session, temp_dir): + prefixes, _ = _nested_prefixes(temp_dir) + inside = _seed(session, prefixes) + assert len(inside) < len(prefixes) + + contents = [content.id for content in live_contents_under_prefixes(session, prefixes)] + with ( + patch.object(scanner, "create_session", lambda: _reuse_session(session)), + patch.object(scanner, "get_scan_prefixes_for_root", lambda _root: prefixes), + ): + live = live_references_safely("models") + rows = get_unenriched_assets_for_roots(("models",), compute_hashes=False, limit=10_000) + + # The outer prefix also takes the inner prefix's "-other" decoy. + under = inside | {str(temp_dir / "shared" / "inner-other" / "decoy.safetensors")} + assert len(contents) == len(set(contents)) == len(under) + assert set(live) == under + assert all(len(observations) == 1 for observations in live.values()) + assert len(rows) == len({row.record_id for row in rows}) == len(inside) + + +@pytest.mark.parametrize("limit", [7, 150, 1000]) +def test_unenriched_paging_across_batches_matches_a_single_ordered_scan(session, temp_dir, limit): + """Keyset pages over many batches equal the pages of one ordered query.""" + prefixes, _ = _nested_prefixes(temp_dir) + _seed(session, prefixes) + expected = list( + session.execute( + sa.select(Asset.id) + .join(AssetContent, Asset.content_id == AssetContent.id) + .where(AssetContent.is_missing.is_(False), AssetContent.path.not_like("%-other%")) + .order_by(Asset.id) + ).scalars() + ) + + pages: list[list[str]] = [] + last_seen_id = None + with ( + patch.object(scanner, "create_session", lambda: _reuse_session(session)), + patch.object(scanner, "get_scan_prefixes_for_root", lambda _root: prefixes), + ): + while True: + rows = get_unenriched_assets_for_roots( + ("models",), compute_hashes=False, limit=limit, last_seen_id=last_seen_id + ) + if not rows: + break + pages.append([row.record_id for row in rows]) + last_seen_id = rows[-1].record_id + + assert pages == [expected[i:i + limit] for i in range(0, len(expected), limit)] + + +def test_path_semantics_hold_in_a_later_batch(session, temp_dir): + root = str(temp_dir / "root") + for path, _ in prefix_case_paths(root): + create_content(session, path) + session.commit() + prefixes = _folders(temp_dir / "elsewhere", 499) + [root] + + returned = {content.path for content in live_contents_under_prefixes(session, prefixes)} + + assert returned == expected_prefix_case_paths(root) + + +def test_windows_paths_in_a_later_batch(session, monkeypatch): + """Drive-letter paths keep exact-or-under, the separator bound and case sensitivity.""" + monkeypatch.setattr(helpers, "os", SimpleNamespace(path=ntpath, sep="\\")) + root = "C:\\models\\target" + stored = { + "C:\\models\\target": True, + "C:\\models\\target\\ckpt.safetensors": True, + "C:\\models\\target\\sub\\lora.safetensors": True, + "C:\\models\\targetx\\ckpt.safetensors": False, + "C:\\models\\target-other\\ckpt.safetensors": False, + "C:\\Models\\Target\\ckpt.safetensors": False, + "D:\\models\\target\\ckpt.safetensors": False, + } + # Stored as Windows' abspath writes them; this host's abspath would mangle them. + session.add_all(AssetContent(path=path, size_bytes=0) for path in stored) + session.commit() + prefixes = [f"C:\\other\\b{i:05d}" for i in range(2 * PREFIX_BATCH_SIZE + 5)] + ["C:\\models\\target\\"] + + returned = {content.path for content in live_contents_under_prefixes(session, prefixes)} + + assert returned == {path for path, inside in stored.items() if inside} + assert root in returned + is_under = stored_path_under_prefixes(prefixes) + assert {path for path in stored if is_under(path)} == returned + + +# --- end to end: a real scan over hundreds of model folders --- + + +MODEL_FOLDERS = 520 + + +@pytest.fixture +def model_files(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> list[Path]: + """MODEL_FOLDERS registered folders; a file in every 40th and the last. + + The prefix count is what broke the scan. Fewer files keep the fast phase's + per-file folder checks from dominating the test's runtime. + """ + engine = _capped_engine(f"sqlite:///{tmp_path / 'assets.db'}") + monkeypatch.setattr("app.database.db.Session", sessionmaker(bind=engine)) + monkeypatch.setattr("app.database.db.WriteSession", sessionmaker(bind=engine)) + + folders = [tmp_path / "models" / f"base{i:04d}" / "checkpoints" for i in range(MODEL_FOLDERS)] + files = [] + for index, folder in enumerate(folders): + folder.mkdir(parents=True) + if index % 40 == 0 or index == MODEL_FOLDERS - 1: + files.append(folder / f"model{index:04d}.safetensors") + files[-1].write_bytes(b"\0" * 16) + folder_strs = [str(f) for f in folders] + monkeypatch.setattr( + scanner, + "get_comfy_models_folders", + lambda: [("checkpoints", folder_strs, {".safetensors"})], + ) + monkeypatch.setattr( + "folder_paths.folder_names_and_paths", + {"checkpoints": (folder_strs, {".safetensors"})}, + ) + monkeypatch.setattr("folder_paths.filename_list_cache", {}) + monkeypatch.setattr(seeder_module, "dependencies_available", lambda: True) + yield files + engine.dispose() + + +def _full_scan(caplog: pytest.LogCaptureFixture) -> list[str]: + caplog.clear() + seeder = seeder_module._AssetSeeder() + with caplog.at_level(logging.INFO): + assert seeder.start(roots=("models",), phase=seeder_module.ScanPhase.FULL) + assert seeder.wait(timeout=120) + return [record.getMessage() for record in caplog.records] + + +def _live_model_paths() -> set[str]: + with db.create_session() as session: + return set( + session.scalars(sa.select(AssetContent.path).where(AssetContent.is_missing.is_(False))) + ) + + +def test_full_scan_over_520_model_folders_completes(model_files, caplog): + messages = _full_scan(caplog) + + assert [m for m in messages if "scan_failed" in m] == [] + assert any("seeder.scan_completed" in m for m in messages), messages + assert _live_model_paths() == {str(f) for f in model_files} + assert get_unenriched_assets_for_roots(("models",), compute_hashes=False) == [] + + # A rescan runs the per-prefix sync, which must see the file that went away. + gone = model_files[-1] + gone.unlink() + messages = _full_scan(caplog) + + assert [m for m in messages if "scan_failed" in m] == [] + assert _live_model_paths() == {str(f) for f in model_files} - {str(gone)}