Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
23 changes: 12 additions & 11 deletions src/shapepipe/modules/sextractor_package/sextractor_script.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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]])
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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()
Expand Down
146 changes: 146 additions & 0 deletions src/shapepipe/pipeline/sqlite_store.py
Original file line number Diff line number Diff line change
@@ -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())
20 changes: 20 additions & 0 deletions tests/module/test_sextractor_post_process.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)],
)
141 changes: 141 additions & 0 deletions tests/module/test_sqlite_store.py
Original file line number Diff line number Diff line change
@@ -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}}
Loading