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

1"""Fold the per-worker shard stores of a multi-GPU ingest into one index.""" 

2 

3from __future__ import annotations 

4 

5import logging 

6from typing import TYPE_CHECKING 

7 

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 

10 

11if TYPE_CHECKING: 

12 from pathlib import Path 

13 

14 from lancedb.db import LanceDBConnection 

15 from lancedb.table import LanceTable 

16 

17 from lilbee.data.store.core import Store 

18 

19log = logging.getLogger(__name__) 

20 

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 

24 

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 

29 

30 

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. 

35 

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 

42 

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 

63 

64 

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. 

67 

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. 

73 

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 

91 

92 

93def _reconcile_sources(store: Store, shard_dirs: list[Path]) -> None: 

94 """Say so when the merged index tracks fewer sources than the workers hold. 

95 

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 

101 

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 ) 

111 

112 

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()) 

118 

119 

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 ) 

125 

126 

127def _predicates(name: str, sources: set[str] | None) -> list[str | None]: 

128 """The where-clauses selecting the rows to merge from table *name*. 

129 

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 ] 

143 

144 

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})" 

149 

150 

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 

154 

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