aether-context 0.3.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,819 @@
1
+ # aether-context (Unlimited Context)
2
+ # Copyright (c) 2026 Aether AI
3
+ # SPDX-License-Identifier: Apache-2.0
4
+ """B2 context pool — session-namespaced, mmap'd vector store + budget governor.
5
+
6
+ This is the **"disk" of the virtual-memory-for-attention design**. Encoded
7
+ :class:`Slice` payloads (256-dim retrieval vector + text + tokens + meta) land here; the
8
+ vectors live in an **mmap'd file** so the pool is disk-resident and survives a reopen, and
9
+ a **budget governor** evicts the lowest-retention slices (ranked by the :class:`Witness`)
10
+ so the pool never exceeds its byte ceiling.
11
+
12
+ What it is / is not
13
+ -------------------
14
+ * It *is* a cosine nearest-neighbor store with a session/namespace filter, so two
15
+ far-apart sessions never bleed into each other's results.
16
+ * It *is* fail-soft about the index: the ``flat`` numpy brute-force index **always
17
+ works**; if ``hnswlib`` is importable and ``config.index == "hnsw"`` it uses HNSW for
18
+ speed, otherwise it transparently falls back to flat. A missing optional dependency is
19
+ never a hard failure.
20
+ * A :class:`Slice` is a self-contained dataclass over the 256-dim retrieval embedding.
21
+
22
+ On-disk layout
23
+ --------------
24
+ Two files inside the pool dir (``config.dir``):
25
+
26
+ * ``vectors.f32`` — a flat ``float32`` mmap of ``[capacity, dim]`` rows; each live slice
27
+ owns one row by its ``row`` index. Grown by re-allocating to a larger capacity.
28
+ * ``pool.json`` — the sidecar header + per-slice metadata records. The vectors file is
29
+ the bulk store; the sidecar is the index of record (id → row + payload). A malformed
30
+ sidecar raises :class:`~aether_context.errors.PoolCorrupt`.
31
+
32
+ The in-RAM ANN index is rebuilt from the persisted vectors on open, so reopening never
33
+ requires the optional ``hnswlib`` to have been present when the pool was written.
34
+ """
35
+ from __future__ import annotations
36
+
37
+ import json
38
+ import math
39
+ import os
40
+ import warnings
41
+ from dataclasses import dataclass, field
42
+ from pathlib import Path
43
+ from typing import Any
44
+
45
+ import numpy as np
46
+
47
+ from aether_context._log import get_logger
48
+ from aether_context.config import PoolConfig, TOKENS_PER_GB
49
+ from aether_context.errors import PoolBudgetError, PoolCorrupt
50
+ from aether_context.quantize import dequantize, packed_bytes_per_row, quantize
51
+ from aether_context.witness import Witness
52
+
53
+ logger = get_logger(__name__)
54
+
55
+ # --- optional fast ANN backend (never a hard dependency) ---------------------
56
+ try: # pragma: no cover - import availability is environment-dependent
57
+ import hnswlib as _hnswlib # type: ignore[import-untyped, import-not-found]
58
+
59
+ _HNSWLIB_AVAILABLE = True
60
+ except ImportError: # pragma: no cover - the common CI path (flat fallback)
61
+ _hnswlib = None
62
+ _HNSWLIB_AVAILABLE = False
63
+
64
+ # --- persistence constants ---------------------------------------------------
65
+ #: On-disk format version for the sidecar/header. Bump on any layout change.
66
+ POOL_FORMAT_VERSION: int = 1
67
+ #: mmap'd vectors filename inside the pool dir.
68
+ VECTORS_FILENAME: str = "vectors.f32"
69
+ #: Sidecar metadata filename inside the pool dir.
70
+ METADATA_FILENAME: str = "pool.json"
71
+ #: Initial row capacity for a fresh vectors mmap (grown geometrically).
72
+ _INITIAL_CAPACITY: int = 64
73
+
74
+ # --- budget accounting -------------------------------------------------------
75
+ #: Fixed per-slice payload overhead charged on top of the vector bytes. Mirrors the
76
+ #: README pool math (~2.2 KB/slice at 512 tok/slice, 256-dim) so ``reach`` lines up:
77
+ #: 256 float32 = 1024 B of vector, leaving ~1.2 KB for text + meta + index bookkeeping.
78
+ SLICE_PAYLOAD_BYTES: int = 1200
79
+
80
+
81
+ def slice_cost_bytes(dim: int) -> int:
82
+ """Bytes the pool charges per resident slice: vector bytes + fixed payload overhead.
83
+
84
+ ``dim * 4`` (float32 vector) plus :data:`SLICE_PAYLOAD_BYTES` for the text/meta/index
85
+ bookkeeping. The governor uses this to translate the GB ceiling into a max slice count
86
+ (it stays in lockstep with :meth:`ContextPool.bytes_used`).
87
+ """
88
+ return dim * 4 + SLICE_PAYLOAD_BYTES
89
+
90
+
91
+ def _default_ceiling_bytes(pool_gb: int, dim: int) -> int:
92
+ """Byte ceiling for a pool sized ``pool_gb`` of reach.
93
+
94
+ Reach is ``pool_gb * TOKENS_PER_GB`` tokens; at 512 tok/slice that is a slice count,
95
+ and each slice costs :func:`slice_cost_bytes`. We derive the ceiling from *reach* (not
96
+ raw GB) so "pool size = reach" stays the honest mental model.
97
+ """
98
+ slices = (pool_gb * TOKENS_PER_GB) // 512
99
+ return int(slices) * slice_cost_bytes(dim)
100
+
101
+
102
+ @dataclass
103
+ class Slice:
104
+ """A self-contained encoded chunk of context — the pool's unit of storage.
105
+
106
+ Fields (all carried verbatim through persistence and search):
107
+
108
+ id stable identifier, unique within the pool
109
+ session namespace; ``search(session=...)`` isolates one session from the rest
110
+ vector the 256-dim float32 retrieval embedding
111
+ text the original text the vector encodes (paged back to the model on a hit)
112
+ tokens token count of ``text`` (for window/budget math)
113
+ meta arbitrary JSON-serializable tags (phase, source, etc.)
114
+ score retention salience in ``[0,1]`` — the witness uses this to rank for eviction
115
+ """
116
+
117
+ id: str
118
+ session: str
119
+ vector: np.ndarray
120
+ text: str
121
+ tokens: int
122
+ meta: dict[str, Any] = field(default_factory=dict)
123
+ score: float = 0.0
124
+
125
+
126
+ class _FlatIndex:
127
+ """Brute-force numpy cosine index — the always-available fallback.
128
+
129
+ Search is a single ``matrix @ query`` dot product (vectors are unit, so dot == cosine).
130
+ O(N) but correct on every platform with zero extra dependencies.
131
+ """
132
+
133
+ kind = "flat"
134
+
135
+ def __init__(self, dim: int) -> None:
136
+ self._dim = dim
137
+
138
+ def search(
139
+ self, matrix: np.ndarray, query: np.ndarray, k: int
140
+ ) -> list[tuple[int, float]]:
141
+ """Top-``k`` ``(row, cosine)`` pairs over ``matrix`` rows, highest cosine first.
142
+
143
+ ``matrix`` is ``(n, dim)`` unit rows; ``query`` is a ``(dim,)`` unit vector. Returns
144
+ at most ``k`` pairs. An empty matrix yields ``[]``.
145
+ """
146
+ n = matrix.shape[0]
147
+ if n == 0 or k <= 0:
148
+ return []
149
+ cosines = matrix @ query # unit rows -> dot product is cosine similarity
150
+ kk = min(k, n)
151
+ # argpartition for the top-kk, then sort just those descending (stable on ties).
152
+ top = np.argpartition(-cosines, kk - 1)[:kk]
153
+ top = top[np.argsort(-cosines[top], kind="stable")]
154
+ return [(int(r), float(cosines[r])) for r in top]
155
+
156
+
157
+ class _HnswIndex:
158
+ """Thin wrapper over ``hnswlib`` for fast approximate cosine search.
159
+
160
+ Used only when ``hnswlib`` is importable and ``config.index == "hnsw"``. Rebuilt from
161
+ the pool's vector matrix on add/open (cheap relative to retrieval over a long run).
162
+ """
163
+
164
+ kind = "hnsw"
165
+
166
+ def __init__(self, dim: int) -> None:
167
+ self._dim = dim
168
+ self._index: Any = None
169
+ self._rows: list[int] = []
170
+
171
+ @property
172
+ def is_built(self) -> bool:
173
+ """Whether a graph has been built yet (``False`` before the first add/rebuild)."""
174
+ return self._index is not None
175
+
176
+ def size(self) -> int:
177
+ """Number of vectors currently in the graph (``0`` before it is built)."""
178
+ if self._index is None:
179
+ return 0
180
+ return int(self._index.get_current_count())
181
+
182
+ def rebuild(self, matrix: np.ndarray, rows: list[int]) -> None:
183
+ """(Re)build the HNSW graph from ``matrix`` rows labelled by ``rows`` (from scratch)."""
184
+ n = matrix.shape[0]
185
+ index = _hnswlib.Index(space="cosine", dim=self._dim)
186
+ index.init_index(max_elements=max(1, n), ef_construction=200, M=16)
187
+ if n:
188
+ index.add_items(matrix, np.asarray(rows, dtype=np.int64))
189
+ index.set_ef(max(16, min(200, n)))
190
+ self._index = index
191
+ self._rows = list(rows)
192
+
193
+ def add_rows(self, new_matrix: np.ndarray, rows: list[int]) -> None:
194
+ """Incrementally insert ``new_matrix`` (labelled by ``rows``) into the existing graph.
195
+
196
+ This is the scaling fix: appending the few new rows since the last sync is ``O(new)``,
197
+ versus an ``O(N)`` full :meth:`rebuild` of the whole pool on every add. The graph is
198
+ grown via ``resize_index`` when it would overflow its current capacity. Falls back to a
199
+ fresh init when no graph exists yet.
200
+ """
201
+ n_new = new_matrix.shape[0]
202
+ if n_new == 0:
203
+ return
204
+ if self._index is None:
205
+ self._index = _hnswlib.Index(space="cosine", dim=self._dim)
206
+ self._index.init_index(max_elements=max(1, n_new), ef_construction=200, M=16)
207
+ cur = int(self._index.get_current_count())
208
+ if cur + n_new > int(self._index.get_max_elements()):
209
+ self._index.resize_index(cur + n_new)
210
+ self._index.add_items(new_matrix, np.asarray(rows, dtype=np.int64))
211
+ total = cur + n_new
212
+ self._index.set_ef(max(16, min(200, total)))
213
+ self._rows.extend(rows)
214
+
215
+ def search(
216
+ self, matrix: np.ndarray, query: np.ndarray, k: int
217
+ ) -> list[tuple[int, float]]:
218
+ """Top-``k`` ``(row, cosine)`` pairs via HNSW; builds lazily if not yet built."""
219
+ n = matrix.shape[0]
220
+ if n == 0 or k <= 0:
221
+ return []
222
+ if self._index is None:
223
+ self.rebuild(matrix, list(range(n)))
224
+ kk = min(k, n)
225
+ labels, distances = self._index.knn_query(query, k=kk)
226
+ # hnswlib 'cosine' returns distance = 1 - cosine; recover the similarity.
227
+ pairs = [
228
+ (int(lbl), float(1.0 - dist))
229
+ for lbl, dist in zip(labels[0], distances[0])
230
+ ]
231
+ pairs.sort(key=lambda p: p[1], reverse=True)
232
+ return pairs
233
+
234
+
235
+ class ContextPool:
236
+ """Session-namespaced, mmap'd, budget-governed vector store of :class:`Slice`.
237
+
238
+ Construct with a :class:`~aether_context.config.PoolConfig`. If the config's ``dir``
239
+ already holds a pool, it is reopened (vectors + sidecar restored, in-RAM index rebuilt);
240
+ otherwise a fresh empty pool is created lazily on first :meth:`add`.
241
+
242
+ Public surface:
243
+ * :meth:`add` — store a slice, then enforce the byte budget via the witness.
244
+ * :meth:`search` — cosine top-``k``, optionally filtered to one ``session``.
245
+ * :meth:`evict_to_budget` — drop lowest-retention slices until under the ceiling.
246
+ * :meth:`bytes_used`, :meth:`stats` — accounting.
247
+ * :meth:`close` — flush vectors + sidecar to disk (also called on GC).
248
+ """
249
+
250
+ def __init__(self, config: PoolConfig, *, ceiling_bytes: int | None = None) -> None:
251
+ self._config = config
252
+ self._dim = config.dim
253
+ # TurboVec: 0 = float32 (default; this whole class is byte-identical to before). >0 stores
254
+ # quantized CODES on disk (8-bit recall-safe) while keeping the float32 unit vector in RAM for
255
+ # exact search — so disk footprint drops ~4x with no recall loss; reload dequantizes.
256
+ self._qbits = int(getattr(config, "quantize_bits", 0) or 0)
257
+ self._mdtype = np.uint8 if self._qbits else np.float32
258
+ self._mcols = packed_bytes_per_row(self._dim, self._qbits) if self._qbits else self._dim
259
+ self._vname = f"vectors.q{self._qbits}" if self._qbits else VECTORS_FILENAME
260
+ self._scale_of: dict[str, float] = {} # per-slice quant scale (qbits>0 only)
261
+ self._dir = Path(config.dir)
262
+ # Resolve the index kind: honor 'hnsw' only when the lib is importable. 'tiered' is
263
+ # accepted by config but not yet built — rather than silently masquerade as a paged
264
+ # index (a silent capability claim), it announces the fallback and runs flat.
265
+ requested = config.index
266
+ if requested == "tiered":
267
+ warnings.warn(
268
+ "index='tiered' is not implemented yet and currently runs the flat index "
269
+ "(no graph paging). RAM/latency match --index flat for now. Use 'hnsw' for "
270
+ "approximate speed, or shrink --pool to fit the resident index.",
271
+ RuntimeWarning,
272
+ stacklevel=2,
273
+ )
274
+ requested = "flat"
275
+ if requested == "hnsw" and _HNSWLIB_AVAILABLE:
276
+ self._index: _FlatIndex | _HnswIndex = _HnswIndex(self._dim)
277
+ else:
278
+ if requested == "hnsw":
279
+ logger.debug(
280
+ "index='hnsw' requested but hnswlib unavailable; using flat fallback"
281
+ )
282
+ self._index = _FlatIndex(self._dim)
283
+ # Byte ceiling: explicit override (tests) else derived from pool reach.
284
+ self._ceiling_bytes = (
285
+ int(ceiling_bytes)
286
+ if ceiling_bytes is not None
287
+ else _default_ceiling_bytes(config.pool_gb, self._dim)
288
+ )
289
+ if self._ceiling_bytes <= 0:
290
+ raise PoolBudgetError(
291
+ f"computed pool ceiling is {self._ceiling_bytes} bytes (non-positive)",
292
+ hint="Raise pool_gb (floor is 5 GB) or pass a positive ceiling_bytes.",
293
+ )
294
+ # In-RAM state. The mmap is the bulk vector store; these mirror live slices.
295
+ self._slices: dict[str, Slice] = {} # id -> Slice (stored unit vector)
296
+ self._row_of: dict[str, int] = {} # id -> row index in the mmap
297
+ self._order: list[str] = [] # insertion order of live ids
298
+ self._witness = Witness() # retention scores for eviction
299
+ self._mmap: np.memmap | None = None
300
+ self._capacity = 0
301
+ self._dirty = False
302
+ # Monotone write clock: bumped on every add so the witness can tell a freshly paged-in
303
+ # slice from a stale one and apply its temporal lock-in (anti-thrash) on eviction.
304
+ self._tick = 0.0
305
+ # Two index-staleness signals (see _refresh_index): `_index_dirty` means rows were
306
+ # *appended* (an incremental add suffices); `_index_rebuild` means rows were
307
+ # *renumbered* (eviction/compaction/load) so the HNSW graph must be rebuilt wholesale.
308
+ self._index_dirty = False
309
+ self._index_rebuild = True
310
+ # Restore from disk if a pool already exists in this dir.
311
+ if self._metadata_file().exists():
312
+ self._load()
313
+
314
+ # -- paths -----------------------------------------------------------------
315
+ def _vectors_file(self) -> Path:
316
+ return self._dir / self._vname
317
+
318
+ def _write_row(self, row: int, sid: str, unit_vec: np.ndarray) -> None:
319
+ """Write a unit vector into mmap row ``row`` — quantized to codes when TurboVec is on."""
320
+ assert self._mmap is not None
321
+ if self._qbits:
322
+ codes, scale = quantize(unit_vec[None], self._qbits)
323
+ self._mmap[row] = codes[0]
324
+ self._scale_of[sid] = float(scale[0])
325
+ else:
326
+ self._mmap[row] = unit_vec
327
+
328
+ def _metadata_file(self) -> Path:
329
+ return self._dir / METADATA_FILENAME
330
+
331
+ @property
332
+ def metadata_path(self) -> Path:
333
+ """Path to the sidecar metadata file (``pool.json``) inside the pool dir."""
334
+ return self._metadata_file()
335
+
336
+ @property
337
+ def vectors_path(self) -> Path:
338
+ """Path to the mmap'd vectors file (``vectors.f32``) inside the pool dir."""
339
+ return self._vectors_file()
340
+
341
+ # -- introspection ---------------------------------------------------------
342
+ @property
343
+ def index_kind(self) -> str:
344
+ """The resolved index kind actually in use (``"flat"`` or ``"hnsw"``)."""
345
+ return self._index.kind
346
+
347
+ @property
348
+ def ceiling_bytes(self) -> int:
349
+ """The byte ceiling the governor holds the pool at or below."""
350
+ return self._ceiling_bytes
351
+
352
+ def __len__(self) -> int:
353
+ return len(self._slices)
354
+
355
+ # -- write -----------------------------------------------------------------
356
+ def add(self, sl: Slice) -> None:
357
+ """Store ``sl`` (overwriting any slice with the same id), then enforce the budget.
358
+
359
+ The vector is validated to be ``(dim,)`` and a normalized copy is written into the
360
+ mmap's row for this id. The slice's ``score`` is fed to the witness as its retention
361
+ salience. After the write the governor runs (:meth:`evict_to_budget`) so the pool is
362
+ *always* at or below its ceiling the instant ``add`` returns.
363
+ """
364
+ vec = np.asarray(sl.vector, dtype=np.float32)
365
+ if vec.shape != (self._dim,):
366
+ raise PoolBudgetError(
367
+ f"slice {sl.id!r} vector has shape {vec.shape}, expected {(self._dim,)}",
368
+ hint=f"Encode with the pool's dim ({self._dim}); the 256-dim retrieval "
369
+ f"embedding is the only vector the pool stores.",
370
+ )
371
+ self._tick += 1.0
372
+ row = self._row_of.get(sl.id)
373
+ if row is None:
374
+ row = len(self._order)
375
+ self._ensure_capacity(row + 1)
376
+ self._order.append(sl.id)
377
+ self._row_of[sl.id] = row
378
+ stored_vec = self._unit(vec)
379
+ assert self._mmap is not None
380
+ self._write_row(row, sl.id, stored_vec)
381
+ self._slices[sl.id] = Slice(
382
+ id=sl.id,
383
+ session=sl.session,
384
+ vector=stored_vec.copy(),
385
+ text=sl.text,
386
+ tokens=int(sl.tokens),
387
+ meta=dict(sl.meta),
388
+ score=float(sl.score),
389
+ )
390
+ # Touch at the current write tick so the witness's temporal lock-in can protect this
391
+ # freshly written slice from immediate eviction (the anti-thrash guarantee).
392
+ self._witness.touch(sl.id, salience=float(sl.score), now=self._tick)
393
+ self._dirty = True
394
+ self._index_dirty = True
395
+ self.evict_to_budget()
396
+
397
+ # -- read ------------------------------------------------------------------
398
+ def search(
399
+ self, query_vec: np.ndarray, k: int, session: str | None = None,
400
+ *, sources: set[str] | None = None,
401
+ ) -> list[Slice]:
402
+ """Top-``k`` slices by cosine similarity to ``query_vec``, highest first.
403
+
404
+ If ``session`` is given the search is **scoped to that namespace** — slices from
405
+ other sessions are invisible, so far-apart sessions never cross-contaminate (even
406
+ when another session has a closer vector).
407
+
408
+ If ``sources`` is given (a set like ``{"user", "tool"}``) the search is restricted to
409
+ slices whose ``meta["source"]`` is in the set — a **provenance filter** so a caller can
410
+ retrieve only trusted memory (e.g. exclude model-authored notes; see SAFETY.md). A slice
411
+ with no ``source`` tag is treated as source ``"user"`` (the conservative default).
412
+
413
+ Returns ``[]`` for an empty pool, an unknown session/source, or ``k <= 0``.
414
+ """
415
+ if k <= 0 or not self._order:
416
+ return []
417
+ query = self._unit(np.asarray(query_vec, dtype=np.float32))
418
+ if query.shape != (self._dim,):
419
+ raise PoolBudgetError(
420
+ f"query vector has shape {query.shape}, expected {(self._dim,)}",
421
+ hint=f"Search with a {self._dim}-dim vector (the encoder's output dim).",
422
+ )
423
+ if session is None and sources is None:
424
+ return self._search_global(query, k)
425
+ return self._search_filtered(query, k, session, sources)
426
+
427
+ def _search_global(self, query: np.ndarray, k: int) -> list[Slice]:
428
+ """Unfiltered top-``k`` over every live slice (uses the resolved ANN index)."""
429
+ matrix = self._live_matrix()
430
+ self._refresh_index(matrix)
431
+ pairs = self._index.search(matrix, query, k)
432
+ return [self._slices[self._order[row]] for row, _ in pairs]
433
+
434
+ def _search_filtered(
435
+ self, query: np.ndarray, k: int, session: str | None, sources: set[str] | None,
436
+ ) -> list[Slice]:
437
+ """Top-``k`` restricted by session and/or source via a brute-force masked dot product.
438
+
439
+ Both filters are exact (a numpy mask over the live matrix), so isolation/provenance hold
440
+ regardless of which index backend is active — correctness never rides on the ANN. An
441
+ untagged slice counts as source ``"user"``.
442
+ """
443
+ def _ok(sid: str) -> bool:
444
+ sl = self._slices[sid]
445
+ if session is not None and sl.session != session:
446
+ return False
447
+ if sources is not None and sl.meta.get("source", "user") not in sources:
448
+ return False
449
+ return True
450
+
451
+ rows = [i for i, sid in enumerate(self._order) if _ok(sid)]
452
+ if not rows:
453
+ return []
454
+ matrix = self._live_matrix()
455
+ sub = matrix[rows]
456
+ cosines = sub @ query
457
+ kk = min(k, len(rows))
458
+ top = np.argpartition(-cosines, kk - 1)[:kk]
459
+ top = top[np.argsort(-cosines[top], kind="stable")]
460
+ return [self._slices[self._order[rows[int(j)]]] for j in top]
461
+
462
+ # -- governor --------------------------------------------------------------
463
+ def evict_to_budget(self) -> list[str]:
464
+ """Evict the lowest-retention slices until the pool fits under its byte ceiling.
465
+
466
+ Ranking is by the :class:`Witness` (lowest score first), exactly as the witness's
467
+ own ``budget_evict`` does — so the survivors are always the most-retained slices.
468
+ Returns the evicted ids (``[]`` when the pool already fits). The governor is called
469
+ automatically after every :meth:`add`, so "never exceeds budget" is literally true.
470
+ """
471
+ cost = slice_cost_bytes(self._dim)
472
+ max_slices = max(0, self._ceiling_bytes // cost)
473
+ if len(self._order) <= max_slices:
474
+ return []
475
+ # Most-evictable first, honoring the witness's temporal lock-in at the current tick.
476
+ order = self._witness.eviction_order(now=self._tick)
477
+ known = set(order)
478
+ # Any live id with no witness entry has no retention signal -> evict it first.
479
+ unranked = [sid for sid in self._order if sid not in known]
480
+ full_order = unranked + order
481
+ n_evict = len(self._order) - max_slices
482
+ evict_ids = full_order[:n_evict]
483
+ for sid in evict_ids:
484
+ self._remove(sid)
485
+ if evict_ids:
486
+ self._compact_rows()
487
+ self._dirty = True
488
+ self._index_rebuild = True # rows renumbered by compaction -> full graph rebuild
489
+ logger.debug(
490
+ "evicted %d slice(s) to hold pool at <=%d bytes",
491
+ len(evict_ids), self._ceiling_bytes,
492
+ )
493
+ return evict_ids
494
+
495
+ def clear_session(self, session_id: str | None) -> int:
496
+ """Remove every live slice belonging to ``session_id`` and return the count removed.
497
+
498
+ This is the storage half of the engine's *clear* semantics: dropping the slices a
499
+ single session externalized into the pool. When ``session_id`` is ``None`` the whole
500
+ pool is emptied (the shared/global clear), so a ``shared`` pool clears all sessions.
501
+
502
+ Rows are compacted after removal so ``_live_matrix`` stays a dense prefix and the
503
+ sidecar row indices remain consistent; the byte accounting and stats follow the
504
+ surviving slices exactly. Returns ``0`` when nothing matched (idempotent, safe).
505
+ """
506
+ if session_id is None:
507
+ removed = list(self._order)
508
+ else:
509
+ removed = [
510
+ sid for sid in self._order
511
+ if self._slices[sid].session == session_id
512
+ ]
513
+ if not removed:
514
+ return 0
515
+ for sid in removed:
516
+ self._remove(sid)
517
+ self._compact_rows()
518
+ self._dirty = True
519
+ self._index_rebuild = True # rows renumbered by compaction -> full graph rebuild
520
+ logger.debug(
521
+ "cleared %d slice(s) for session %r", len(removed), session_id
522
+ )
523
+ return len(removed)
524
+
525
+ def _remove(self, sid: str) -> None:
526
+ """Drop a single slice id from the in-RAM structures (mmap row reclaimed on compact)."""
527
+ self._slices.pop(sid, None)
528
+ self._row_of.pop(sid, None)
529
+ self._scale_of.pop(sid, None)
530
+ self._witness.forget(sid)
531
+ try:
532
+ self._order.remove(sid)
533
+ except ValueError:
534
+ pass
535
+
536
+ def _compact_rows(self) -> None:
537
+ """Re-pack live slices into contiguous rows ``[0..len)`` after eviction.
538
+
539
+ Keeps the mmap dense so ``_live_matrix`` is a simple prefix slice and row indices
540
+ stay stable for the sidecar. The stored ``Slice.vector`` is the source of truth for
541
+ each row, so packing never depends on a stale mmap position.
542
+ """
543
+ if self._mmap is None:
544
+ return
545
+ n = len(self._order)
546
+ if n == 0:
547
+ self._row_of = {}
548
+ return
549
+ packed = np.empty((n, self._mcols), dtype=self._mdtype)
550
+ for new_row, sid in enumerate(self._order):
551
+ v = self._slices[sid].vector
552
+ if self._qbits:
553
+ codes, scale = quantize(v[None], self._qbits)
554
+ packed[new_row] = codes[0]
555
+ self._scale_of[sid] = float(scale[0])
556
+ else:
557
+ packed[new_row] = v
558
+ self._row_of[sid] = new_row
559
+ self._mmap[:n] = packed
560
+
561
+ # -- accounting ------------------------------------------------------------
562
+ def bytes_used(self) -> int:
563
+ """Total bytes the live slices occupy, by the same accounting the governor uses."""
564
+ return len(self._slices) * slice_cost_bytes(self._dim)
565
+
566
+ def stats(self) -> dict[str, Any]:
567
+ """A small dict snapshot: count, bytes, ceiling, dim, index kind, sessions."""
568
+ return {
569
+ "count": len(self._slices),
570
+ "bytes_used": self.bytes_used(),
571
+ "ceiling_bytes": self._ceiling_bytes,
572
+ "dim": self._dim,
573
+ "index": self.index_kind,
574
+ "sessions": sorted({s.session for s in self._slices.values()}),
575
+ "quantized": self._qbits > 0,
576
+ "quantize_bits": self._qbits,
577
+ }
578
+
579
+ # -- persistence -----------------------------------------------------------
580
+ def close(self) -> None:
581
+ """Flush vectors + sidecar to disk and release the mmap. Idempotent.
582
+
583
+ Releasing the mmap matters on Windows: a still-mapped file cannot be reopened, so
584
+ reopening the same pool dir would fail with ``[Errno 22]`` unless the prior handle
585
+ is dropped here.
586
+ """
587
+ self._flush()
588
+ self._release_mmap()
589
+
590
+ def __del__(self) -> None: # pragma: no cover - best-effort flush on GC
591
+ try:
592
+ self._flush()
593
+ self._release_mmap()
594
+ except Exception: # noqa: BLE001 - never raise from a finalizer
595
+ pass
596
+
597
+ def _flush(self) -> None:
598
+ """Write the mmap and sidecar metadata if anything changed since the last flush."""
599
+ if not self._dirty:
600
+ return
601
+ self._dir.mkdir(parents=True, exist_ok=True)
602
+ if self._mmap is not None:
603
+ self._mmap.flush()
604
+ records = [
605
+ {
606
+ "id": sid,
607
+ "session": self._slices[sid].session,
608
+ "row": self._row_of[sid],
609
+ "text": self._slices[sid].text,
610
+ "tokens": self._slices[sid].tokens,
611
+ "meta": self._slices[sid].meta,
612
+ "score": self._slices[sid].score,
613
+ "scale": self._scale_of.get(sid, 1.0),
614
+ }
615
+ for sid in self._order
616
+ ]
617
+ header = {
618
+ "version": POOL_FORMAT_VERSION,
619
+ "dim": self._dim,
620
+ "quantize_bits": self._qbits,
621
+ "count": len(self._order),
622
+ "capacity": self._capacity,
623
+ "index": self._config.index,
624
+ "ceiling_bytes": self._ceiling_bytes,
625
+ "slices": records,
626
+ }
627
+ tmp = self._metadata_file().with_suffix(".json.tmp")
628
+ tmp.write_text(json.dumps(header), encoding="utf-8")
629
+ os.replace(tmp, self._metadata_file())
630
+ self._dirty = False
631
+
632
+ def _load(self) -> None:
633
+ """Restore a pool from ``pool.json`` + ``vectors.f32``; raise PoolCorrupt on bad data."""
634
+ path = self._metadata_file()
635
+ try:
636
+ header = json.loads(path.read_text(encoding="utf-8"))
637
+ except (json.JSONDecodeError, OSError, ValueError) as exc:
638
+ raise PoolCorrupt(f"could not read pool metadata at {path}: {exc}") from exc
639
+ try:
640
+ disk_dim = int(header["dim"])
641
+ records = header["slices"]
642
+ capacity = int(header.get("capacity", len(records)))
643
+ except (KeyError, TypeError, ValueError) as exc:
644
+ raise PoolCorrupt(
645
+ f"pool metadata at {path} is missing required fields: {exc}"
646
+ ) from exc
647
+ if disk_dim != self._dim:
648
+ raise PoolCorrupt(
649
+ f"pool at {self._dir} was written with dim={disk_dim}, "
650
+ f"but this pool expects dim={self._dim}",
651
+ )
652
+ disk_qbits = int(header.get("quantize_bits", 0))
653
+ if disk_qbits != self._qbits:
654
+ raise PoolCorrupt(
655
+ f"pool at {self._dir} was written with quantize_bits={disk_qbits}, "
656
+ f"but this pool expects quantize_bits={self._qbits} (codes can't be reinterpreted)",
657
+ )
658
+ vfile = self._vectors_file()
659
+ if not vfile.exists():
660
+ raise PoolCorrupt(
661
+ f"pool metadata at {path} present but vectors file {vfile} is missing"
662
+ )
663
+ self._capacity = max(capacity, _INITIAL_CAPACITY)
664
+ self._open_mmap(self._capacity)
665
+ assert self._mmap is not None
666
+ try:
667
+ for rec in records:
668
+ sid = rec["id"]
669
+ row = int(rec["row"])
670
+ if self._qbits:
671
+ scale = float(rec.get("scale", 1.0))
672
+ self._scale_of[sid] = scale
673
+ vec = dequantize(np.asarray(self._mmap[row])[None],
674
+ np.array([scale], dtype=np.float32),
675
+ self._dim, self._qbits)[0]
676
+ else:
677
+ vec = np.asarray(self._mmap[row], dtype=np.float32).copy()
678
+ sl = Slice(
679
+ id=sid,
680
+ session=rec["session"],
681
+ vector=vec,
682
+ text=rec["text"],
683
+ tokens=int(rec["tokens"]),
684
+ meta=dict(rec.get("meta", {})),
685
+ score=float(rec.get("score", 0.0)),
686
+ )
687
+ self._order.append(sid)
688
+ self._row_of[sid] = row
689
+ self._slices[sid] = sl
690
+ # Advance the write tick per restored record so the on-disk insertion order
691
+ # becomes the recency gradient the temporal lock-in reads after a reopen.
692
+ self._tick += 1.0
693
+ self._witness.touch(sid, salience=sl.score, now=self._tick)
694
+ except (KeyError, TypeError, ValueError, IndexError) as exc:
695
+ raise PoolCorrupt(
696
+ f"pool metadata at {path} has a malformed slice record: {exc}"
697
+ ) from exc
698
+ self._index_rebuild = True
699
+
700
+ # -- mmap management -------------------------------------------------------
701
+ def _ensure_capacity(self, n_rows: int) -> None:
702
+ """Make sure the mmap can hold ``n_rows`` rows, growing geometrically if needed.
703
+
704
+ Windows cannot reopen (``w+``) a file that is still memory-mapped, so the old
705
+ mapping is copied into RAM and **fully released** before the larger mmap is opened.
706
+ """
707
+ if self._mmap is not None and n_rows <= self._capacity:
708
+ return
709
+ new_cap = max(_INITIAL_CAPACITY, self._capacity)
710
+ while new_cap < n_rows:
711
+ new_cap *= 2
712
+ carry: np.ndarray | None = None
713
+ if self._mmap is not None:
714
+ keep = min(self._capacity, len(self._order))
715
+ if keep > 0:
716
+ carry = np.asarray(self._mmap[:keep], dtype=self._mdtype).copy()
717
+ self._release_mmap()
718
+ self._dir.mkdir(parents=True, exist_ok=True)
719
+ self._open_mmap(new_cap, carry=carry)
720
+
721
+ def _release_mmap(self) -> None:
722
+ """Flush and drop the current mmap so the OS handle is released (Windows-safe)."""
723
+ if self._mmap is not None:
724
+ try:
725
+ self._mmap.flush()
726
+ except (ValueError, OSError): # already closed / detached
727
+ pass
728
+ self._mmap = None
729
+
730
+ def _open_mmap(self, capacity: int, *, carry: np.ndarray | None = None) -> None:
731
+ """Open the vectors mmap at ``capacity`` rows, optionally seeding it with ``carry``.
732
+
733
+ On a reopen (no ``carry`` provided, file already on disk) the existing rows are read
734
+ with :func:`numpy.fromfile` — which opens and closes the file cleanly, leaving no
735
+ lingering mapping before the ``w+`` open (the cause of an [Errno 22] on Windows).
736
+ """
737
+ vfile = self._vectors_file()
738
+ seed = carry
739
+ if seed is None and vfile.exists() and capacity > 0:
740
+ seed = self._read_existing_rows(vfile, capacity)
741
+ mm = np.memmap(vfile, dtype=self._mdtype, mode="w+", shape=(capacity, self._mcols))
742
+ if seed is not None and seed.shape[0] > 0:
743
+ rows = min(seed.shape[0], capacity)
744
+ mm[:rows] = seed[:rows]
745
+ self._mmap = mm
746
+ self._capacity = capacity
747
+
748
+ def _read_existing_rows(self, vfile: Path, capacity: int) -> np.ndarray | None:
749
+ """Read existing on-disk vector rows (up to ``capacity``) into RAM for a reopen.
750
+
751
+ Uses :func:`numpy.fromfile` (plain read, no mmap) so no file handle survives the
752
+ call — required so the subsequent ``w+`` mmap open succeeds on Windows.
753
+ """
754
+ size = vfile.stat().st_size
755
+ row_bytes = self._mcols * np.dtype(self._mdtype).itemsize
756
+ if row_bytes == 0:
757
+ return None
758
+ n_rows = size // row_bytes
759
+ if n_rows == 0:
760
+ return None
761
+ take = min(n_rows, capacity)
762
+ flat = np.fromfile(vfile, dtype=self._mdtype, count=take * self._mcols)
763
+ if flat.size < take * self._mcols:
764
+ return None
765
+ return flat.reshape(take, self._mcols)
766
+
767
+ # -- helpers ---------------------------------------------------------------
768
+ def _live_matrix(self) -> np.ndarray:
769
+ """The ``(n, dim)`` matrix of live slice vectors in current row order."""
770
+ n = len(self._order)
771
+ if n == 0 or self._mmap is None:
772
+ return np.empty((0, self._dim), dtype=np.float32)
773
+ if self._qbits:
774
+ # mmap holds quantized codes; search on the exact float32 RAM copies (recall parity).
775
+ return np.array([self._slices[sid].vector for sid in self._order], dtype=np.float32)
776
+ return np.asarray(self._mmap[:n])
777
+
778
+ def _refresh_index(self, matrix: np.ndarray) -> None:
779
+ """Sync the ANN index to ``matrix`` if it is stale (flat index is a no-op).
780
+
781
+ Two paths keep HNSW cheap over a long run:
782
+
783
+ * ``_index_rebuild`` — rows were renumbered (eviction/compaction/load) so labels are
784
+ stale; the graph is rebuilt wholesale.
785
+ * ``_index_dirty`` — rows were only *appended*; the new tail is inserted incrementally
786
+ (``O(new)``) instead of rebuilding the whole graph (``O(N)``) on every add.
787
+ """
788
+ if not isinstance(self._index, _HnswIndex):
789
+ self._index_dirty = False
790
+ self._index_rebuild = False
791
+ return
792
+ n = matrix.shape[0]
793
+ if self._index_rebuild or not self._index.is_built:
794
+ self._index.rebuild(matrix, list(range(n)))
795
+ elif self._index_dirty:
796
+ built = self._index.size()
797
+ if n > built:
798
+ self._index.add_rows(matrix[built:], list(range(built, n)))
799
+ self._index_dirty = False
800
+ self._index_rebuild = False
801
+
802
+ @staticmethod
803
+ def _unit(vec: np.ndarray) -> np.ndarray:
804
+ """L2-normalize ``vec`` to a float32 unit vector (zero vector passes through)."""
805
+ norm = float(np.linalg.norm(vec))
806
+ if norm < 1e-12 or math.isnan(norm):
807
+ return vec.astype(np.float32)
808
+ return (vec / norm).astype(np.float32)
809
+
810
+
811
+ __all__ = [
812
+ "ContextPool",
813
+ "Slice",
814
+ "slice_cost_bytes",
815
+ "POOL_FORMAT_VERSION",
816
+ "VECTORS_FILENAME",
817
+ "METADATA_FILENAME",
818
+ "SLICE_PAYLOAD_BYTES",
819
+ ]