Coverage for src/lilbee/data/store/shard_merge.py: 100%
73 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-09-28 17:20 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-09-28 17:20 +0000
1"""Fold the per-worker shard stores of a multi-GPU ingest into one index."""
3from __future__ import annotations
5import logging
6from typing import TYPE_CHECKING
8from lilbee.core.config import CHUNKS_TABLE, INGEST_SOURCE_COLUMNS, META_TABLE, SOURCES_TABLE
9from lilbee.data.store.lance_helpers import escape_sql_string, table_names
11if TYPE_CHECKING:
12 from pathlib import Path
14 from lancedb.db import LanceDBConnection
15 from lancedb.table import LanceTable
17 from lilbee.data.store.core import Store
19log = logging.getLogger(__name__)
21# Rows read from a shard per append. A chunks row carries its vector, so a whole
22# shard table does not fit in memory.
23_MERGE_BATCH_ROWS = 10_000
25# Source names per ``IN`` predicate when only part of a shard is merged. LanceDB
26# parses the predicate as one SQL string, so the names are chunked rather than
27# joined into a single clause per sync.
28_NAMES_PER_PREDICATE = 500
31def merge_shards(
32 store: Store, shard_dirs: list[Path], *, sources: set[str] | None = None
33) -> dict[str, int]:
34 """Append every shard's rows into *store*, returning the rows merged per table.
36 Shards own disjoint sources, so their union is an append with no dedup. A
37 whole-shard copy (*sources* None) is the fresh-index case. Naming the touched
38 sources is the re-sync case: the store's own rows for those keys are dropped
39 first, so re-merging replaces them instead of doubling them.
40 """
41 from lancedb.db import LanceDBConnection
43 if sources is not None:
44 store.remove_documents(sorted(sources))
45 merged: dict[str, int] = {}
46 adopted = _adopt_chunks(store, shard_dirs, sources)
47 if adopted is not None:
48 merged[CHUNKS_TABLE] = adopted
49 for shard_dir in shard_dirs:
50 database = LanceDBConnection(str(shard_dir))
51 for name in table_names(database):
52 # The merged store writes its own meta row from the running config;
53 # a shard's copy would land beside it as a second row.
54 if name == META_TABLE:
55 continue
56 if name == CHUNKS_TABLE and adopted is not None:
57 continue # already taken over whole, without reading a row
58 rows = _copy_table(database.open_table(name), store, name, sources)
59 merged[name] = merged.get(name, 0) + rows
60 log.info("Merged %d shard(s): %s", len(shard_dirs), merged)
61 _reconcile_sources(store, shard_dirs)
62 return merged
65def _adopt_chunks(store: Store, shard_dirs: list[Path], sources: set[str] | None) -> int | None:
66 """Take over every shard's chunk fragments whole; None when that cannot apply.
68 The chunks table carries the vectors, so it is the whole cost of the merge:
69 at 8.8M rows by 4096 dims the row copy rewrites about 144GB, and since the
70 shard stores stay as resume state the corpus then sits on disk twice.
71 Adopting the fragments is metadata only, and the hard links mean one physical
72 copy with two names.
74 Whole-fragment, so only a full merge qualifies: a scoped re-sync names its
75 sources, and a fragment there holds touched and untouched rows together.
76 Returns None when the caller should copy rows instead, which also covers a
77 shard on another filesystem or a data file whose name is already taken; the
78 merge is correct either way, only slower.
79 """
80 if sources is not None:
81 return None
82 tables = [shard_dir / f"{CHUNKS_TABLE}.lance" for shard_dir in shard_dirs]
83 present = [table for table in tables if table.exists()]
84 if not present:
85 return None
86 try:
87 return store.adopt_fragments(CHUNKS_TABLE, present)
88 except OSError as exc:
89 log.warning("Adopting shard fragments failed (%s); copying rows instead", exc)
90 return None
93def _reconcile_sources(store: Store, shard_dirs: list[Path]) -> None:
94 """Say so when the merged index tracks fewer sources than the workers hold.
96 A scoped merge only takes what the run touched, so a source a worker holds and
97 the index does not (an earlier merge that failed, a removal against the index
98 alone) would otherwise stay missing with nothing to show for it.
99 """
100 from lancedb.db import LanceDBConnection
102 held = sum(_source_count(LanceDBConnection(str(shard_dir))) for shard_dir in shard_dirs)
103 merged = _source_count(store.get_db())
104 if merged < held:
105 log.warning(
106 "The index tracks %d source(s) against %d across the ingest workers. "
107 "Re-run with --force to fold every worker's shard back in.",
108 merged,
109 held,
110 )
113def _source_count(database: LanceDBConnection) -> int:
114 """Rows in a store's source table, zero when it has none."""
115 if SOURCES_TABLE not in table_names(database):
116 return 0
117 return int(database.open_table(SOURCES_TABLE).count_rows())
120def _copy_table(table: LanceTable, store: Store, name: str, sources: set[str] | None) -> int:
121 """Append the rows of *table* that this merge wants into *store*."""
122 return sum(
123 _copy_rows(table, store, name, predicate) for predicate in _predicates(name, sources)
124 )
127def _predicates(name: str, sources: set[str] | None) -> list[str | None]:
128 """The where-clauses selecting the rows to merge from table *name*.
130 ``None`` is the whole table. A table with no source column holds corpus-level
131 aggregates that the post-merge passes rebuild, so a scoped merge skips it.
132 """
133 if sources is None:
134 return [None]
135 column = INGEST_SOURCE_COLUMNS.get(name)
136 if column is None:
137 return []
138 names = sorted(sources)
139 return [
140 _in_predicate(column, names[start : start + _NAMES_PER_PREDICATE])
141 for start in range(0, len(names), _NAMES_PER_PREDICATE)
142 ]
145def _in_predicate(column: str, names: list[str]) -> str:
146 """``column IN (...)`` over *names*."""
147 quoted = ", ".join(f"'{escape_sql_string(name)}'" for name in names)
148 return f"{column} IN ({quoted})"
151def _copy_rows(table: LanceTable, store: Store, name: str, predicate: str | None) -> int:
152 """Stream the rows *predicate* selects from *table* into *store*."""
153 import pyarrow as pa
155 query = table.search()
156 if predicate is not None:
157 query = query.where(predicate)
158 reader = query.limit(0).to_batches(_MERGE_BATCH_ROWS)
159 copied = 0
160 for batch in reader:
161 copied += store.absorb_rows(name, pa.Table.from_batches([batch], schema=reader.schema))
162 return copied