diff --git a/src/shapepipe/modules/sextractor_package/sextractor_script.py b/src/shapepipe/modules/sextractor_package/sextractor_script.py index 5b97fb07d..017333f37 100644 --- a/src/shapepipe/modules/sextractor_package/sextractor_script.py +++ b/src/shapepipe/modules/sextractor_package/sextractor_script.py @@ -11,9 +11,9 @@ import numpy as np from astropy.io import fits from astropy.wcs.wcs import InvalidCoordinateError -from sqlitedict import SqliteDict from shapepipe.pipeline import file_io +from shapepipe.pipeline.sqlite_store import read_sqlitedict def get_header_value(image_path, key): @@ -161,11 +161,14 @@ def make_post_process(cat_path, f_wcs_path, pos_params, ccd_size, w_log=None): ) cat.open() - f_wcs = SqliteDict(f_wcs_path) + # One lock-free read of the whole header log: merge_headers wrote and + # closed it in an earlier step, and keyed SqliteDict reads would take an + # NFS lock per access. + f_wcs = read_sqlitedict(f_wcs_path) # Tile-level logs from merge_headers carry a "TILE_ID" metadata entry # (inserted first); n_hdu must be derived from a real exposure entry, # otherwise it measures the tile ID string and truncates the CCD scan. - exp_keys = [key for key in f_wcs.keys() if key != "TILE_ID"] + exp_keys = [key for key in f_wcs if key != "TILE_ID"] if len(exp_keys) == 0: raise IOError(f"Could not read sql file '{f_wcs_path}'") n_hdu = len(f_wcs[exp_keys[0]]) @@ -194,14 +197,14 @@ def make_post_process(cat_path, f_wcs_path, pos_params, ccd_size, w_log=None): n_epoch = np.zeros(len(obj_id), dtype="int32") for idx, exp in enumerate(exp_list): + if exp not in f_wcs: + raise KeyError( + f"Exposure {exp} used in image {cat_path} but not" + + f" found in header file {f_wcs_path}. Make sure this" + + " file is complete." + ) pos_tmp = np.ones(len(obj_id), dtype="int32") * -1 for idx_j in range(n_hdu): - if exp not in f_wcs: - raise KeyError( - f"Exposure {exp} used in image {cat_path} but not" - + f" found in header file {f_wcs_path}. Make sure this" - + " file is complete." - ) w = f_wcs[exp][idx_j]["WCS"] # Only inverse-project objects near this CCD's footprint. # Positions far outside the distortion domain make the @@ -243,8 +246,6 @@ def make_post_process(cat_path, f_wcs_path, pos_params, ccd_size, w_log=None): cat.save_as_fits(data=a, ext_name=f"EPOCH_{idx}") cat.open() - f_wcs.close() - cat.add_col("N_EPOCH", n_epoch) cat.close() diff --git a/src/shapepipe/pipeline/sqlite_store.py b/src/shapepipe/pipeline/sqlite_store.py new file mode 100644 index 000000000..bc1240c9a --- /dev/null +++ b/src/shapepipe/pipeline/sqlite_store.py @@ -0,0 +1,146 @@ +"""SQLITE STORE. + +Lock-free, read-only access to the SqliteDict stores the pipeline writes +between modules (WCS header logs, vignet catalogues, PSF catalogues). + +Reading through :class:`sqlitedict.SqliteDict` costs one sqlite transaction +per key access, and each transaction takes and releases a POSIX lock on the +file. On a networked filesystem those locks are slow (tens to hundreds of +milliseconds, and far worse on a busy server), so a loop of keyed reads can +spend most of its time in lock traffic. :class:`ImmutableSqliteDict` opens the +file with sqlite's ``immutable=1`` URI parameter instead, which skips all +locking and change detection. + +``immutable=1`` is only correct for a file that no process writes while it +is open: sqlite neither sees concurrent changes nor replays a hot journal +left by a crashed writer, so it would return uncommitted rows. Every call +site of this module reads a store that an earlier module wrote and closed +(``make_post_process`` reads the header log ``merge_headers`` wrote). As a +guard, opening a store that has a rollback journal or write-ahead log next +to it raises instead of reading it. + +""" + +import sqlite3 +from collections.abc import Mapping +from pathlib import Path + +from sqlitedict import decode + + +class ImmutableSqliteDict(Mapping): + """Read-only, lock-free mapping over a SqliteDict file. + + Keys and values are those :class:`sqlitedict.SqliteDict` returns for a + store written with its default key and value encoding (identity keys, + pickled values); iteration follows insertion (``rowid``) order, as + SqliteDict's does. Each lookup is one indexed query; use + :func:`read_sqlitedict` to load a whole store at once. + + The file must not be written while it is open (see the module + docstring). + + Parameters + ---------- + path : str or os.PathLike + Path to an existing SqliteDict file + tablename : str, optional + SqliteDict table name, default ``"unnamed"`` (SqliteDict's default) + + Raises + ------ + FileNotFoundError + If ``path`` is not an existing file + RuntimeError + If a ``-journal`` or ``-wal`` file sits next to ``path``: a writer is + mid-transaction or crashed in one, and an immutable read would see + uncommitted data + + """ + + def __init__(self, path, tablename="unnamed"): + path = Path(path) + if not path.is_file(): + raise FileNotFoundError(f"SqliteDict file not found: '{path}'") + for suffix in ("-journal", "-wal"): + sidecar = path.with_name(path.name + suffix) + if sidecar.exists(): + raise RuntimeError( + f"SqliteDict file '{path}' has a '{suffix}' file next to" + + " it: it is being written, or a writer died" + + " mid-transaction. Open it with SqliteDict to roll the" + + " journal back, or rewrite it." + ) + self.path = path + self._table = '"' + tablename.replace('"', '""') + '"' + self._conn = sqlite3.connect( + f"{path.resolve().as_uri()}?immutable=1", + uri=True, + check_same_thread=False, + ) + + def __getitem__(self, key): + row = self._conn.execute( + f"SELECT value FROM {self._table} WHERE key = ?", (key,) + ).fetchone() + if row is None: + raise KeyError(key) + return decode(row[0]) + + def __contains__(self, key): + return ( + self._conn.execute( + f"SELECT 1 FROM {self._table} WHERE key = ?", (key,) + ).fetchone() + is not None + ) + + def __iter__(self): + for (key,) in self._conn.execute( + f"SELECT key FROM {self._table} ORDER BY rowid" + ): + yield key + + def __len__(self): + return self._conn.execute( + f"SELECT COUNT(*) FROM {self._table}" + ).fetchone()[0] + + def items(self): + """Iterate over ``(key, value)`` pairs with a single query.""" + for key, value in self._conn.execute( + f"SELECT key, value FROM {self._table} ORDER BY rowid" + ): + yield key, decode(value) + + def close(self): + """Close the underlying sqlite connection.""" + self._conn.close() + + def __enter__(self): + return self + + def __exit__(self, *exc): + self.close() + + +def read_sqlitedict(path, tablename="unnamed"): + """Load a whole SqliteDict file into a dict without taking locks. + + The file must not be written during the read (see the module docstring). + + Parameters + ---------- + path : str or os.PathLike + Path to an existing SqliteDict file + tablename : str, optional + SqliteDict table name, default ``"unnamed"`` + + Returns + ------- + dict + Every entry of the store, in insertion order + + """ + with ImmutableSqliteDict(path, tablename=tablename) as store: + return dict(store.items()) diff --git a/tests/module/test_sextractor_post_process.py b/tests/module/test_sextractor_post_process.py index 6797fdee1..446e1fbf6 100644 --- a/tests/module/test_sextractor_post_process.py +++ b/tests/module/test_sextractor_post_process.py @@ -175,3 +175,23 @@ def test_duplicate_history_cards_create_one_epoch_hdu(tmp_path): assert epoch_names == ["EPOCH_0"] np.testing.assert_array_equal(n_epoch, [1, 1]) + + +def test_exposure_missing_from_header_log_raises(tmp_path): + """A HISTORY exposure absent from the header log is a KeyError.""" + npy_path = tmp_path / "headers-123456.npy" + _write_exposure_headers(npy_path) + merge_headers([[str(npy_path)]], str(tmp_path), tile_number="53") + sqlite_path = tmp_path / "log_exp_headers53.sqlite" + + positions = _make_ccd_wcs(0)[0].all_pix2world([[50.0, 50.0]], 0) + cat_path = tmp_path / "sexcat.fits" + _write_sex_ldac(cat_path, ["123456", "999999"], positions) + + with pytest.raises(KeyError, match="999999"): + sextractor_script.make_post_process( + str(cat_path), + str(sqlite_path), + ["XWIN_WORLD", "YWIN_WORLD"], + ["0", str(CCD_NPIX), "0", str(CCD_NPIX)], + ) diff --git a/tests/module/test_sqlite_store.py b/tests/module/test_sqlite_store.py new file mode 100644 index 000000000..8a294bae5 --- /dev/null +++ b/tests/module/test_sqlite_store.py @@ -0,0 +1,141 @@ +"""UNIT TESTS FOR THE LOCK-FREE SQLITEDICT READER. + +``ImmutableSqliteDict`` and ``read_sqlitedict`` must return exactly what +``SqliteDict`` returns for the stores the pipeline writes: same keys, same +insertion order, same decoded values. + +""" + +import pickle +import subprocess +import sys + +import numpy as np +import pytest +from sqlitedict import SqliteDict + +from shapepipe.modules.merge_headers_package.merge_headers import merge_headers +from shapepipe.pipeline.sqlite_store import ImmutableSqliteDict, read_sqlitedict + +from .test_sextractor_post_process import _write_exposure_headers + + +def _sqlitedict_items(path): + with SqliteDict(str(path), flag="r") as db: + return list(db.items()) + + +def _assert_same_items(got, expected): + assert [key for key, _ in got] == [key for key, _ in expected] + for (_, got_value), (_, expected_value) in zip(got, expected): + assert pickle.dumps(got_value) == pickle.dumps(expected_value) + + +@pytest.fixture +def wcs_store(tmp_path): + """Tile header log written by merge_headers from split_exp-style files.""" + header_files = [] + for name in ("2104589", "2104590", "2366971"): + npy_path = tmp_path / f"headers-{name}.npy" + _write_exposure_headers(npy_path) + header_files.append([str(npy_path)]) + merge_headers(header_files, str(tmp_path), tile_number="270.283") + return tmp_path / "log_exp_headers270.283.sqlite" + + +@pytest.fixture +def object_store(tmp_path): + """Per-object store with overwritten keys, sentinels and numpy payloads.""" + path = tmp_path / "objects.sqlite" + with SqliteDict(str(path)) as db: + db["1"] = {"2104589-12": {"VIGNET": np.arange(9.0).reshape(3, 3)}} + db["2"] = "empty" + db["3"] = {} + db["1"] = {"2104589-13": {"SHAPES": {"HSM_FLAG_PSF": 0}}} + db.commit() + return path + + +def test_read_sqlitedict_matches_sqlitedict_on_wcs_store(wcs_store): + """The merge_headers WCS log reads back identically, TILE_ID first.""" + expected = _sqlitedict_items(wcs_store) + got = read_sqlitedict(wcs_store) + + assert isinstance(got, dict) + assert next(iter(got)) == "TILE_ID" + _assert_same_items(list(got.items()), expected) + assert got["2104590"][1]["WCS"].wcs.compare( + dict(expected)["2104590"][1]["WCS"].wcs + ) + + +def test_immutable_mapping_matches_sqlitedict(object_store): + """Lookups, membership, length and rowid order agree with SqliteDict.""" + expected = _sqlitedict_items(object_store) + with ImmutableSqliteDict(object_store) as store: + assert list(store) == ["2", "3", "1"] + assert len(store) == 3 + _assert_same_items(list(store.items()), expected) + _assert_same_items([(key, store[key]) for key in store], expected) + assert "2" in store + assert "4" not in store + with pytest.raises(KeyError): + store["4"] + + +def test_missing_file_raises_without_creating_it(tmp_path): + """A missing store is an error, and the read leaves no file behind.""" + path = tmp_path / "absent.sqlite" + with pytest.raises(FileNotFoundError): + read_sqlitedict(path) + assert not path.exists() + + +def test_read_leaves_no_journal(wcs_store): + """An immutable read writes nothing next to the store.""" + before = sorted(p.name for p in wcs_store.parent.iterdir()) + read_sqlitedict(wcs_store) + assert sorted(p.name for p in wcs_store.parent.iterdir()) == before + + +@pytest.mark.parametrize("suffix", ["-journal", "-wal"]) +def test_sidecar_journal_refuses_read(object_store, suffix): + """A rollback journal or WAL next to the store refuses the read.""" + object_store.with_name(object_store.name + suffix).write_bytes(b"") + with pytest.raises(RuntimeError, match=suffix): + ImmutableSqliteDict(object_store) + + +def test_hot_journal_from_killed_writer_refuses_read(tmp_path): + """A writer killed mid-transaction leaves a hot journal: refuse, not + return its uncommitted rows.""" + path = tmp_path / "store.sqlite" + with SqliteDict(str(path)) as db: + db["1"] = "committed" + db.commit() + script = ( + "import os, sqlitedict\n" + f"db = sqlitedict.SqliteDict({str(path)!r})\n" + "db['1'] = 'uncommitted'\n" + "db['2'] = 'uncommitted'\n" + "db.conn.select_one('SELECT 1')\n" + "os._exit(0)\n" + ) + subprocess.run([sys.executable, "-c", script], check=True) + if not path.with_name(path.name + "-journal").exists(): + pytest.skip("writer left no journal on this filesystem") + with pytest.raises(RuntimeError, match="-journal"): + read_sqlitedict(path) + with SqliteDict(str(path), flag="r") as db: + assert dict(db.items()) == {"1": "committed"} + + +def test_path_with_uri_special_characters(tmp_path): + """Space, '?', '#' and '%' in the path do not break the file: URI.""" + directory = tmp_path / "a dir?x=1#frag%20" + directory.mkdir() + path = directory / "st ore?#%.sqlite" + with SqliteDict(str(path)) as db: + db["k"] = {"v": 1} + db.commit() + assert read_sqlitedict(path) == {"k": {"v": 1}}