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
118 changes: 90 additions & 28 deletions ir/store.py
Original file line number Diff line number Diff line change
Expand Up @@ -33,6 +33,7 @@
import io
import json
import os
import uuid
from collections.abc import Iterator, Mapping, MutableMapping
from pathlib import Path
from typing import Any
Expand Down Expand Up @@ -381,38 +382,93 @@ def _invalidate_matrix(self) -> None:
self._clear_packed()
self._packed_stale = True

def _packed_paths(self):
# Legacy (pre-generation) flat file names. A cache written by an older
# ``ir`` is still read through these; new writes never use them.
_LEGACY_PACKED_FILES = {
"matrix": "matrix.npy",
"ids": "ids.json",
"metas": "metas.json",
}

def _packed_paths(self, generation: str | None = None):
"""Paths of one packed set: the shared ``sig.json`` plus its data files.

Each writer puts its data files under a fresh *generation* token
(``matrix-<gen>.npy``, ...), so no data file is ever written by two
writers, and ``sig.json`` -- replaced atomically, last -- names the one
generation readers should load. ``generation=None`` gives the legacy
flat names an older ``ir`` wrote.
"""
d = self._packed_dir
# ``sig`` is written last and removed first, so a half-written or
# half-cleared cache (no/sig-less dir) always reads as invalid.
return {
"sig": d / "sig.json",
"matrix": d / "matrix.npy",
"ids": d / "ids.json",
"metas": d / "metas.json",
}

def _clear_packed(self) -> None:
if generation is None:
files = self._LEGACY_PACKED_FILES
else:
files = {
"matrix": f"matrix-{generation}.npy",
"ids": f"ids-{generation}.json",
"metas": f"metas-{generation}.json",
}
return {"sig": d / "sig.json", **{k: d / v for k, v in files.items()}}

def _packed_data_files(self) -> list[Path]:
"""Every packed data file on disk (legacy names and all generations)."""
d = self._packed_dir
found = [d / name for name in self._LEGACY_PACKED_FILES.values()]
for pattern in ("matrix-*.npy", "ids-*.json", "metas-*.json", "sig-*.tmp"):
found.extend(d.glob(pattern))
return found

def _clear_packed(self, *, keep: str | None = None) -> None:
"""Remove the packed cache (``sig.json`` first, so it reads as invalid).

With ``keep``, only ``sig.json`` is kept and only that generation's data
files survive: this is the post-publish sweep of older generations.
Best-effort -- a file another process still has mapped may refuse to go
(Windows); it is then just left for a later sweep.
"""
if self._packed_dir is None:
return
paths = self._packed_paths()
for key in ("sig", "matrix", "ids", "metas"): # sig first
if keep is None:
try:
paths[key].unlink()
self._packed_paths()["sig"].unlink()
except OSError:
pass
if not self._packed_dir.is_dir():
return
kept: set[Path] = set()
if keep is not None:
kept.update(self._packed_paths(keep).values())
try: # another writer may have published after us: keep its set too
current = json.loads(self._packed_paths()["sig"].read_text("utf-8"))
if isinstance(current.get("generation"), str):
kept.update(self._packed_paths(current["generation"]).values())
except (OSError, ValueError):
pass
for path in self._packed_data_files():
if path in kept:
continue
try:
path.unlink()
except OSError:
pass

def _load_packed(self):
"""Load ``(ids, mmap_matrix, metas)`` from the packed cache, or ``None``."""
if self._packed_dir is None:
return None
paths = self._packed_paths()
if not paths["sig"].exists():
sig_path = self._packed_paths()["sig"]
if not sig_path.exists():
return None
try:
sig = json.loads(paths["sig"].read_text(encoding="utf-8"))
sig = json.loads(sig_path.read_text(encoding="utf-8"))
if sig.get("format") != _PACKED_FORMAT:
return None
generation = sig.get("generation")
if generation is not None and not (
isinstance(generation, str) and generation.isalnum()
):
return None
paths = self._packed_paths(generation)
mat = np.load(paths["matrix"], mmap_mode="r")
ids_json = paths["ids"].read_bytes()
metas_json = paths["metas"].read_bytes()
Expand All @@ -424,9 +480,7 @@ def _load_packed(self):
return None
# A cache written before ``sig.json`` carried these fields has neither;
# keep accepting it on the length checks alone rather than invalidating
# every existing cache. When they are present they must agree, or the
# set mixes two writers' files and reads as a miss (rebuild) instead of
# serving rows under the wrong ids.
# every existing cache. When they are present they must agree.
content_sig = sig.get("content_sig")
if content_sig is not None and content_sig != _packed_content_sig(
ids_json, metas_json
Expand All @@ -440,40 +494,48 @@ def _load_packed(self):
def _save_packed(self, result: tuple[list[str], np.ndarray, list[dict]]) -> None:
"""Persist a freshly built matrix to the packed cache (best-effort).

Skips empty corpora. Writes ``sig.json`` last so a crash mid-write
leaves the cache marked invalid (no sig) rather than torn. The sig also
carries a :func:`_packed_content_sig` digest of the ids/metas bytes and
the matrix shape, which is what lets :meth:`_load_packed` notice the
case a write order cannot defend against: a *second* writer replacing
some of the four files while this set is on disk.
Skips empty corpora. The matrix, ids and metas go to files named by a
fresh generation token that only this call writes, and ``sig.json``
(naming that generation, plus a :func:`_packed_content_sig` digest and
the matrix shape) is published last with an atomic ``os.replace``. So a
crash mid-write leaves the previous sig (or none) in charge, and two
concurrent writers can never leave one writer's matrix beside the
other's ids: a reader follows ``sig.json`` to exactly one writer's
complete set. Older generations are swept after publishing.
"""
if self._packed_dir is None:
return
ids, mat, metas = result
if not ids:
return
generation = uuid.uuid4().hex
try:
self._packed_dir.mkdir(parents=True, exist_ok=True)
paths = self._packed_paths()
paths = self._packed_paths(generation)
arr = np.asarray(mat, dtype=np.float32)
ids_json = json.dumps(ids).encode("utf-8")
metas_json = json.dumps(metas).encode("utf-8")
np.save(paths["matrix"], arr)
paths["ids"].write_bytes(ids_json)
paths["metas"].write_bytes(metas_json)
paths["sig"].write_text(
sig_tmp = self._packed_dir / f"sig-{generation}.tmp"
sig_tmp.write_text(
json.dumps(
{
"format": _PACKED_FORMAT,
"count": len(ids),
"generation": generation,
"content_sig": _packed_content_sig(ids_json, metas_json),
"shape": list(arr.shape),
}
),
encoding="utf-8",
)
os.replace(sig_tmp, paths["sig"])
self._packed_stale = False
except OSError:
# A read cache that can't be written is non-fatal: fall back to the
# in-process cache (already set by the caller) for this process.
self._clear_packed()
return
self._clear_packed(keep=generation)
111 changes: 104 additions & 7 deletions tests/test_store_packed.py
Original file line number Diff line number Diff line change
Expand Up @@ -41,6 +41,18 @@ def _file_store(tmp_path):
return store, root


def _packed_file(root, kind):
"""Path of the ``kind`` (matrix/ids/metas) file the current sig points at."""
packed = root / "matrix"
sig = json.loads((packed / "sig.json").read_text(encoding="utf-8"))
gen = sig.get("generation")
legacy = {"matrix": "matrix.npy", "ids": "ids.json", "metas": "metas.json"}
if gen is None:
return packed / legacy[kind]
ext = "npy" if kind == "matrix" else "json"
return packed / f"{kind}-{gen}.{ext}"


def test_packed_cache_written_then_reloaded_by_fresh_store(tmp_path):
store, root = _file_store(tmp_path)
store.put_record(_rec("r1", vec=(3.0, 0.0, 0.0)))
Expand Down Expand Up @@ -115,13 +127,13 @@ def test_torn_packed_cache_is_rejected_not_served(tmp_path):
store.put_record(_rec("r2", vec=(0.0, 4.0, 0.0)))
store.matrix() # builds from records + persists packed

packed = root / "matrix"
# Writer B rewrites ids/metas (same length, different order) after writer A
# wrote matrix.npy — the common same-size incremental case.
ids = json.loads((packed / "ids.json").read_text(encoding="utf-8"))
metas = json.loads((packed / "metas.json").read_text(encoding="utf-8"))
(packed / "ids.json").write_text(json.dumps(ids[::-1]), encoding="utf-8")
(packed / "metas.json").write_text(json.dumps(metas[::-1]), encoding="utf-8")
# wrote the matrix — the common same-size incremental case.
ids_path, metas_path = _packed_file(root, "ids"), _packed_file(root, "metas")
ids = json.loads(ids_path.read_text(encoding="utf-8"))
metas = json.loads(metas_path.read_text(encoding="utf-8"))
ids_path.write_text(json.dumps(ids[::-1]), encoding="utf-8")
metas_path.write_text(json.dumps(metas[::-1]), encoding="utf-8")

store2, _ = _file_store(tmp_path)
ids2, mat2, _metas2 = store2.matrix()
Expand All @@ -139,7 +151,7 @@ def test_packed_cache_with_mismatched_matrix_shape_is_rejected(tmp_path):
store.matrix()

# Writer B's matrix (a different embedding dim) over writer A's sig/ids.
np.save(root / "matrix" / "matrix.npy", np.zeros((2, 5), dtype=np.float32))
np.save(_packed_file(root, "matrix"), np.zeros((2, 5), dtype=np.float32))

store2, _ = _file_store(tmp_path)
_ids2, mat2, _metas2 = store2.matrix()
Expand Down Expand Up @@ -167,3 +179,88 @@ def _no_rebuild():
ids2, _mat2, metas2 = store2.matrix()
assert ids2 == ids
assert metas2 == metas


def test_interleaved_packed_writers_never_serve_mismatched_rows(tmp_path, monkeypatch):
"""Writer A saves its matrix, writer B saves a whole set, then A finishes.

With four independently written files, the disk ends up holding B's matrix
beside A's ids/metas/sig. A signature over ids/metas alone vouches for that
mixture (both writers' matrices have the same shape), so rows get served
under the wrong ids. Whatever the interleaving, a reload must either miss
(and rebuild) or return one writer's consistent set.
"""
import ir.store as ir_store

store_a, root = _file_store(tmp_path)
store_a.put_record(_rec("r1", vec=(3.0, 0.0, 0.0)))
store_a.put_record(_rec("r2", vec=(0.0, 4.0, 0.0)))
result_a = store_a._build_matrix()
ids_a, mat_a, metas_a = result_a
# Writer B: the same corpus listed in the other order (a legitimate build).
order = [ids_a.index(rid) for rid in reversed(ids_a)]
result_b = (
[ids_a[i] for i in order],
np.asarray(mat_a)[order],
[metas_a[i] for i in order],
)
store_b, _ = _file_store(tmp_path)

real_save = np.save
state = {"interleaved": False}

def save_then_let_b_run(*args, **kwargs):
real_save(*args, **kwargs)
if not state["interleaved"]:
state["interleaved"] = True
store_b._save_packed(result_b) # B runs start to finish here

monkeypatch.setattr(ir_store.np, "save", save_then_let_b_run)
store_a._save_packed(result_a)
monkeypatch.setattr(ir_store.np, "save", real_save)
assert state["interleaved"]

store2, _ = _file_store(tmp_path)
ids2, mat2, _metas2 = store2.matrix()
expected = {"r1": (1.0, 0.0, 0.0), "r2": (0.0, 1.0, 0.0)}
assert sorted(ids2) == ["r1", "r2"]
for i, rid in enumerate(ids2):
np.testing.assert_allclose(np.asarray(mat2[i]), expected[rid], atol=1e-6)


def test_legacy_flat_packed_layout_still_loads(tmp_path):
"""A cache in the pre-generation flat layout (``matrix.npy`` ...) is used."""
store, root = _file_store(tmp_path)
store.put_record(_rec("r1", vec=(3.0, 0.0, 0.0)))
ids, mat, metas = store._build_matrix()
packed = root / "matrix"
packed.mkdir(parents=True, exist_ok=True)
np.save(packed / "matrix.npy", np.asarray(mat, dtype=np.float32))
(packed / "ids.json").write_text(json.dumps(ids), encoding="utf-8")
(packed / "metas.json").write_text(json.dumps(metas), encoding="utf-8")
(packed / "sig.json").write_text(
json.dumps({"format": 1, "count": len(ids)}), encoding="utf-8"
)

store2, _ = _file_store(tmp_path)

def _no_rebuild():
raise AssertionError("legacy packed cache should be used, not rebuilt")

store2._build_matrix = _no_rebuild
ids2, _mat2, metas2 = store2.matrix()
assert ids2 == ids and metas2 == metas


def test_republishing_sweeps_older_generations(tmp_path):
"""Each save leaves exactly one generation (plus sig.json) on disk."""
store, root = _file_store(tmp_path)
store.put_record(_rec("r1"))
result = store._build_matrix()
for _ in range(3):
store._save_packed(result)
names = sorted(p.name for p in (root / "matrix").iterdir())
assert len(names) == 4 and "sig.json" in names

store.put_record(_rec("r2")) # a write clears the cache entirely
assert list((root / "matrix").iterdir()) == []
Loading