embedflow 0.1.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.
Files changed (64) hide show
  1. embedflow/__init__.py +25 -0
  2. embedflow/__main__.py +3 -0
  3. embedflow/analysis.py +192 -0
  4. embedflow/cache/__init__.py +4 -0
  5. embedflow/cache/base.py +28 -0
  6. embedflow/cache/persistent_cache.py +198 -0
  7. embedflow/cli.py +1200 -0
  8. embedflow/compatibility/__init__.py +28 -0
  9. embedflow/compatibility/candidate_gap.py +105 -0
  10. embedflow/compatibility/containment.py +17 -0
  11. embedflow/compatibility/evaluate.py +319 -0
  12. embedflow/compatibility/metrics.py +75 -0
  13. embedflow/compatibility/migration_depth.py +67 -0
  14. embedflow/compatibility/probe.py +34 -0
  15. embedflow/compatibility/report.py +102 -0
  16. embedflow/compatibility/t2.py +64 -0
  17. embedflow/config.py +455 -0
  18. embedflow/data/__init__.py +1 -0
  19. embedflow/data/registry/__init__.py +1 -0
  20. embedflow/data/registry/benchmark_profiles.jsonl +3 -0
  21. embedflow/data/registry/checksums.sha256 +4 -0
  22. embedflow/data/registry/migrations.jsonl +15 -0
  23. embedflow/data/registry/registry_manifest.json +16 -0
  24. embedflow/data/registry/research_summaries.json +55 -0
  25. embedflow/data/registry/schema_version.json +5 -0
  26. embedflow/frozen/T2_V1_FROZEN_SPEC.md +71 -0
  27. embedflow/frozen/T2_V1_FROZEN_SPEC.sha256 +1 -0
  28. embedflow/indexes/__init__.py +5 -0
  29. embedflow/indexes/base.py +60 -0
  30. embedflow/indexes/faiss_backend.py +240 -0
  31. embedflow/indexes/qdrant_backend.py +225 -0
  32. embedflow/metrics/__init__.py +3 -0
  33. embedflow/metrics/latency.py +50 -0
  34. embedflow/migration/__init__.py +3 -0
  35. embedflow/migration/compatibility.py +156 -0
  36. embedflow/migration/facade.py +312 -0
  37. embedflow/migration/materializer.py +190 -0
  38. embedflow/migration/planner.py +78 -0
  39. embedflow/migration/state.py +81 -0
  40. embedflow/models/__init__.py +4 -0
  41. embedflow/models/base.py +31 -0
  42. embedflow/models/huggingface.py +226 -0
  43. embedflow/registry/__init__.py +47 -0
  44. embedflow/registry/loader.py +785 -0
  45. embedflow/registry/matcher.py +197 -0
  46. embedflow/registry/schema.py +266 -0
  47. embedflow/runtime.py +115 -0
  48. embedflow/serving/__init__.py +3 -0
  49. embedflow/serving/api.py +161 -0
  50. embedflow/serving/engine.py +222 -0
  51. embedflow/serving/factory.py +3 -0
  52. embedflow/serving/schemas.py +39 -0
  53. embedflow-0.1.0.dist-info/METADATA +210 -0
  54. embedflow-0.1.0.dist-info/RECORD +64 -0
  55. embedflow-0.1.0.dist-info/WHEEL +5 -0
  56. embedflow-0.1.0.dist-info/entry_points.txt +2 -0
  57. embedflow-0.1.0.dist-info/licenses/LICENSE +178 -0
  58. embedflow-0.1.0.dist-info/top_level.txt +2 -0
  59. src/__init__.py +1 -0
  60. src/embed.py +123 -0
  61. src/probe_features.py +24 -0
  62. src/storage.py +51 -0
  63. src/t2_v1.py +21 -0
  64. src/utils.py +53 -0
@@ -0,0 +1,312 @@
1
+ """Small public facade for starting a progressive model migration.
2
+
3
+ The research implementation deliberately exposes the lower-level config and
4
+ engine objects. This module adds the short path a user needs in an
5
+ application: point EmbedFlow at an existing FAISS/Qdrant index, name the old
6
+ and new embedding models, and receive a serving session. All retrieval,
7
+ cache, worker, and probe behavior is delegated to the existing implementation;
8
+ this is not a second query path.
9
+ """
10
+
11
+ from __future__ import annotations
12
+
13
+ import json
14
+ from collections.abc import Iterable
15
+ from pathlib import Path
16
+ from typing import Any
17
+
18
+ from ..cache import SQLiteVectorCache
19
+ from ..config import (
20
+ CacheConfig,
21
+ DocumentsConfig,
22
+ EmbedFlowConfig,
23
+ IndexConfig,
24
+ MigrationConfig,
25
+ ModelConfig,
26
+ hydrate_research_contract,
27
+ save_config,
28
+ )
29
+ from ..indexes.base import VectorIndex
30
+ from ..migration.compatibility import run_probe, save_probe
31
+ from ..migration.state import DocumentStore
32
+ from ..models import EmbeddingModel, load_embedding_model
33
+ from ..runtime import load_index
34
+ from ..serving.engine import MigrationEngine
35
+
36
+ _MODEL_REGISTRY = {
37
+ "sentence-transformers/all-MiniLM-L6-v2": "minilm_l6",
38
+ "Qwen/Qwen3-Embedding-0.6B": "qwen3_0_6b",
39
+ "Qwen/Qwen3-Embedding-4B": "qwen3_4b",
40
+ "Qwen/Qwen3-Embedding-8B": "qwen3_8b",
41
+ }
42
+
43
+
44
+ def _is_model(value: Any) -> bool:
45
+ return all(hasattr(value, name) for name in ("encode_queries", "encode_documents", "dimension", "fingerprint"))
46
+
47
+
48
+ def _model_config(value: str | EmbeddingModel, model_root: str | Path | None = None) -> ModelConfig:
49
+ """Build a contract for a model identifier or an already-loaded model."""
50
+ if _is_model(value):
51
+ # HuggingFaceEmbeddingModel retains its exact hydrated config. Reuse
52
+ # it when available so fingerprints and prompt contracts remain exact.
53
+ retained = getattr(value, "cfg", None)
54
+ if isinstance(retained, ModelConfig):
55
+ return retained
56
+ return ModelConfig(model=str(value.model_id), dimension=int(value.dimension))
57
+ if not isinstance(value, (str, Path)):
58
+ raise TypeError("old_model/new_model must be a model ID/path or EmbeddingModel instance")
59
+ identifier = str(value)
60
+ cfg = hydrate_research_contract(ModelConfig(identifier), project_root=Path(__file__).resolve().parents[2])
61
+ candidate = Path(identifier)
62
+ if candidate.exists():
63
+ cfg.local_path = str(candidate.resolve())
64
+ elif model_root is not None and identifier in _MODEL_REGISTRY:
65
+ staged = Path(model_root).expanduser().resolve() / _MODEL_REGISTRY[identifier]
66
+ if staged.exists():
67
+ cfg.local_path = str(staged)
68
+ return cfg
69
+
70
+
71
+ def _queries(value: str | Path | Iterable[tuple[str, str]]) -> list[tuple[str, str]]:
72
+ if isinstance(value, (str, Path)):
73
+ path = Path(value)
74
+ if not path.exists():
75
+ raise FileNotFoundError(path)
76
+ rows: list[tuple[str, str]] = []
77
+ with path.open() as handle:
78
+ seen: set[str] = set()
79
+ for i, line in enumerate(handle, 1):
80
+ if not line.strip():
81
+ continue
82
+ try:
83
+ row = json.loads(line)
84
+ except json.JSONDecodeError as exc:
85
+ raise ValueError(f"invalid query JSON at {path}:{i}") from exc
86
+ if not isinstance(row, dict):
87
+ raise ValueError(f"query row {i} in {path} must be a JSON object")
88
+ query_id = str(row.get("id", row.get("query_id", i)))
89
+ text = row.get("text", row.get("query"))
90
+ if query_id in seen:
91
+ raise ValueError(f"duplicate query ID {query_id!r} in {path}")
92
+ if not isinstance(text, str) or not text.strip():
93
+ raise ValueError(f"query {query_id!r} has no text")
94
+ seen.add(query_id)
95
+ rows.append((query_id, text))
96
+ if not rows:
97
+ raise ValueError(f"query file is empty: {value}")
98
+ return rows
99
+ try:
100
+ rows = [(str(query_id), text) for query_id, text in value]
101
+ except (TypeError, ValueError) as exc:
102
+ raise ValueError("probe_queries must be an iterable of (query_id, text) pairs") from exc
103
+ if len({query_id for query_id, _ in rows}) != len(rows):
104
+ raise ValueError("probe_queries must not contain duplicate query IDs")
105
+ if not rows or any(not isinstance(text, str) or not text.strip() for _, text in rows):
106
+ raise ValueError("probe_queries must contain non-empty query text")
107
+ return rows
108
+
109
+
110
+ class MigrationSession:
111
+ """Handle returned by :func:`migrate`.
112
+
113
+ ``search`` and ``status`` are synchronous convenience methods for an
114
+ application. ``serve`` mounts the same FastAPI dashboard/API used by the
115
+ CLI. Use ``close`` or a context manager to release models and the cache.
116
+ """
117
+
118
+ def __init__(self, engine: MigrationEngine, config: EmbedFlowConfig,
119
+ config_path: Path | None, owns_source: bool, owns_target: bool, owns_index: bool):
120
+ self.engine = engine
121
+ self.config = config
122
+ self.config_path = config_path
123
+ self._owns_source = owns_source
124
+ self._owns_target = owns_target
125
+ self._owns_index = owns_index
126
+ self._closed = False
127
+
128
+ @property
129
+ def plan(self):
130
+ return self.engine.plan
131
+
132
+ def search(self, query: str, top_k: int = 10, candidate_depth: int | None = None,
133
+ max_sync_misses: int | None = None) -> dict[str, Any]:
134
+ return self.engine.search(query, top_k, candidate_depth, max_sync_misses)
135
+
136
+ def status(self) -> dict[str, Any]:
137
+ return self.engine.status()
138
+
139
+ def prewarm(self, document_ids: list[str], asynchronous: bool = True) -> dict[str, Any]:
140
+ return self.engine.prewarm(document_ids, asynchronous=asynchronous)
141
+
142
+ def app(self):
143
+ from ..serving.api import create_app
144
+ return create_app(self.engine)
145
+
146
+ def serve(self, host: str = "127.0.0.1", port: int = 8000, log_level: str = "info") -> None:
147
+ """Run the API/dashboard until interrupted."""
148
+ try:
149
+ import uvicorn
150
+ except ImportError as exc:
151
+ raise RuntimeError("serving requires uvicorn; install `embedflow[dashboard]`") from exc
152
+ uvicorn.run(self.app(), host=host, port=int(port), log_level=log_level)
153
+
154
+ def close(self) -> None:
155
+ if self._closed:
156
+ return
157
+ self.engine.close(close_models=False, close_indexes=self._owns_index)
158
+ if self._owns_source:
159
+ self.engine.source_model.close()
160
+ if self._owns_target:
161
+ self.engine.target_model.close()
162
+ self._closed = True
163
+
164
+ def __enter__(self) -> MigrationSession:
165
+ return self
166
+
167
+ def __exit__(self, *_: Any) -> None:
168
+ self.close()
169
+
170
+
171
+ def migrate(
172
+ *,
173
+ index: str | Path | VectorIndex,
174
+ old_model: str | Path | EmbeddingModel,
175
+ new_model: str | Path | EmbeddingModel,
176
+ documents: str | Path | DocumentStore,
177
+ backend: str | None = None,
178
+ index_url: str | None = None,
179
+ collection: str = "embedflow",
180
+ vector_name: str | None = None,
181
+ api_key_env: str | None = "QDRANT_API_KEY",
182
+ metric: str = "cosine",
183
+ model_root: str | Path | None = None,
184
+ device: str | None = None,
185
+ cache_path: str | Path = "./embedflow_cache",
186
+ state_path: str | Path = "./embedflow_state.json",
187
+ config_path: str | Path | None = None,
188
+ candidate_depth: int = 50,
189
+ kmax_probe: int = 500,
190
+ max_sync_misses: int = 4,
191
+ background_batch_size: int = 32,
192
+ probe_queries: str | Path | Iterable[tuple[str, str]] | None = None,
193
+ probe_limit: int | None = None,
194
+ start_worker: bool = True,
195
+ ) -> MigrationSession:
196
+ """Start progressive migration over an existing index.
197
+
198
+ ``index`` may be an existing EmbedFlow ``VectorIndex`` instance or a path
199
+ to a FAISS index (or a Qdrant path/URL when ``backend="qdrant"``). Model
200
+ strings are revision-hydrated when they are registered by the research
201
+ contract; model objects can be supplied by an application directly.
202
+ ``probe_queries`` is optional so serving can start immediately. When
203
+ supplied, the existing frozen T2-v1 implementation is run before the
204
+ session is returned.
205
+ """
206
+ if isinstance(documents, DocumentStore):
207
+ document_store = documents
208
+ else:
209
+ document_store = DocumentStore(str(documents))
210
+ source_cfg = _model_config(old_model, model_root)
211
+ target_cfg = _model_config(new_model, model_root)
212
+ selected_device = device or target_cfg.device or source_cfg.device or "cpu"
213
+
214
+ owns_source = not _is_model(old_model)
215
+ owns_target = not _is_model(new_model)
216
+ owns_index = not (isinstance(index, VectorIndex) or all(hasattr(index, name) for name in ("search", "size", "dimension", "metadata")))
217
+ source_model = None
218
+ target_model = None
219
+ source_index = None
220
+ cache = None
221
+ engine = None
222
+ try:
223
+ source_model = old_model if not owns_source else load_embedding_model(source_cfg, model_root=model_root, device=selected_device)
224
+ target_model = new_model if not owns_target else load_embedding_model(target_cfg, model_root=model_root, device=selected_device)
225
+ source_cfg.dimension = int(source_model.dimension)
226
+ target_cfg.dimension = int(target_model.dimension)
227
+
228
+ if not owns_index:
229
+ source_index = index
230
+ index_backend = backend or ("qdrant" if str(source_index.metadata().get("backend", "")).lower() == "qdrant" else "faiss")
231
+ index_path = str(getattr(source_index, "path", "./legacy.index") or "./legacy.index")
232
+ else:
233
+ index_path = str(index_url or index)
234
+ index_backend = (backend or ("qdrant" if "://" in index_path else "faiss")).lower()
235
+ if index_backend not in {"faiss", "qdrant"}:
236
+ raise ValueError("backend must be faiss or qdrant")
237
+
238
+ cfg = EmbedFlowConfig(
239
+ source=source_cfg,
240
+ target=target_cfg,
241
+ index=IndexConfig(backend=index_backend, path=index_path, collection=collection,
242
+ url=index_url, metric=str(metric), vector_name=vector_name,
243
+ api_key_env=api_key_env),
244
+ documents=DocumentsConfig(path=str(document_store.path)),
245
+ migration=MigrationConfig(candidate_depth=int(candidate_depth), kmax_probe=max(int(kmax_probe), int(candidate_depth)),
246
+ max_sync_misses=int(max_sync_misses), background_batch_size=int(background_batch_size)),
247
+ cache=CacheConfig(path=str(cache_path)),
248
+ state_path=str(state_path),
249
+ dashboard_title=f"EmbedFlow — {source_cfg.model} → {target_cfg.model}",
250
+ )
251
+ base = Path(config_path).expanduser().resolve().parent if config_path else Path.cwd()
252
+ cfg.resolve_paths(base)
253
+ cfg.validate()
254
+
255
+ if owns_index:
256
+ source_index = load_index(cfg, document_store)
257
+ if int(source_index.dimension) != int(source_model.dimension):
258
+ raise ValueError(f"source model dimension {source_model.dimension} != existing index dimension {source_index.dimension}")
259
+ stored_fingerprint = source_index.metadata().get("model_fingerprint")
260
+ if stored_fingerprint and stored_fingerprint != source_model.fingerprint:
261
+ raise ValueError("source model fingerprint does not match the existing index contract")
262
+
263
+ cache = SQLiteVectorCache(cfg.cache.path, target_model.fingerprint, target_model.dimension)
264
+ probe: dict[str, Any] = {}
265
+ if probe_queries is not None:
266
+ probe = run_probe(source_model, target_model, source_index, document_store, _queries(probe_queries),
267
+ kmax=cfg.migration.kmax_probe, seed=42, limit=probe_limit)
268
+ if config_path:
269
+ save_probe(probe, Path(config_path).expanduser().resolve().with_name("probe_result.json"))
270
+ if config_path:
271
+ config_file = Path(config_path).expanduser().resolve()
272
+ save_config(cfg, config_file)
273
+ else:
274
+ config_file = None
275
+ engine = MigrationEngine(cfg, source_model, target_model, source_index, cache, document_store,
276
+ probe=probe, start_worker=start_worker)
277
+ return MigrationSession(engine, cfg, config_file, owns_source, owns_target, owns_index)
278
+ except Exception:
279
+ # A failed setup must not leak a loaded model, SQLite connection, or a
280
+ # local Qdrant file lock. Only close resources created by this call;
281
+ # caller-owned model/index objects remain their caller's responsibility.
282
+ if engine is not None:
283
+ try:
284
+ engine.close(close_models=False, close_indexes=owns_index)
285
+ except Exception:
286
+ pass
287
+ elif cache is not None:
288
+ try:
289
+ cache.close()
290
+ except Exception:
291
+ pass
292
+ if owns_index and source_index is not None:
293
+ try:
294
+ close_index = getattr(source_index, "close", None)
295
+ if callable(close_index):
296
+ close_index()
297
+ except Exception:
298
+ pass
299
+ if owns_source and source_model is not None:
300
+ try:
301
+ source_model.close()
302
+ except Exception:
303
+ pass
304
+ if owns_target and target_model is not None:
305
+ try:
306
+ target_model.close()
307
+ except Exception:
308
+ pass
309
+ raise
310
+
311
+
312
+ __all__ = ["MigrationSession", "migrate"]
@@ -0,0 +1,190 @@
1
+ from __future__ import annotations
2
+
3
+ import sqlite3
4
+ import threading
5
+ import time
6
+ from collections.abc import Iterable
7
+ from pathlib import Path
8
+ from typing import Any
9
+
10
+ import numpy as np
11
+
12
+ from ..cache import SQLiteVectorCache
13
+
14
+
15
+ class PersistentWorkQueue:
16
+ def __init__(self, path: str | Path, max_retries: int = 3):
17
+ p = Path(path); self.path = p if p.suffix else p / "materialization_queue.sqlite3"
18
+ try:
19
+ retries = int(max_retries)
20
+ exact = float(max_retries) == retries
21
+ except (TypeError, ValueError, OverflowError):
22
+ retries, exact = 0, False
23
+ if isinstance(max_retries, bool) or not exact or retries < 1:
24
+ raise ValueError("max_retries must be a positive integer")
25
+ self.path.parent.mkdir(parents=True, exist_ok=True); self.max_retries = retries
26
+ self.db = None; self.lock = threading.RLock()
27
+ try:
28
+ self.db = sqlite3.connect(str(self.path), check_same_thread=False, timeout=30)
29
+ self.db.execute("PRAGMA busy_timeout=30000")
30
+ self.db.execute("PRAGMA journal_mode=WAL")
31
+ self.db.execute("""CREATE TABLE IF NOT EXISTS work (
32
+ document_id TEXT PRIMARY KEY, status TEXT NOT NULL, attempts INTEGER NOT NULL DEFAULT 0,
33
+ error TEXT, enqueued_at REAL NOT NULL, updated_at REAL NOT NULL)""")
34
+ self.db.execute("UPDATE work SET status='pending',updated_at=? WHERE status='processing'", (time.time(),)); self.db.commit()
35
+ except sqlite3.DatabaseError as exc:
36
+ if self.db is not None:
37
+ self.db.close()
38
+ self.db = None
39
+ raise ValueError(f"invalid materialization queue database: {self.path}") from exc
40
+
41
+ def _ensure_open(self) -> None:
42
+ if self.db is None:
43
+ raise RuntimeError("materialization queue is closed")
44
+
45
+ def enqueue(self, ids: Iterable[str]) -> int:
46
+ now = time.time(); n = 0
47
+ with self.lock:
48
+ self._ensure_open()
49
+ for document_id in dict.fromkeys(str(x) for x in ids):
50
+ cur = self.db.execute("INSERT OR IGNORE INTO work(document_id,status,enqueued_at,updated_at) VALUES(?,?,?,?)",
51
+ (document_id, "pending", now, now)); n += cur.rowcount
52
+ self.db.commit()
53
+ return n
54
+
55
+ def claim(self, limit: int) -> list[str]:
56
+ if isinstance(limit, bool) or int(limit) != limit or int(limit) < 1:
57
+ raise ValueError("queue claim limit must be a positive integer")
58
+ with self.lock:
59
+ self._ensure_open()
60
+ # A write transaction makes SELECT+UPDATE atomic across multiple
61
+ # worker processes sharing the same SQLite queue. Without it,
62
+ # two workers can claim the same pending document before either
63
+ # marks it processing.
64
+ self.db.execute("BEGIN IMMEDIATE")
65
+ try:
66
+ rows = self.db.execute("SELECT document_id FROM work WHERE status='pending' AND attempts<? ORDER BY enqueued_at,document_id LIMIT ?",
67
+ (self.max_retries, int(limit))).fetchall()
68
+ ids = [str(x[0]) for x in rows]; now = time.time()
69
+ self.db.executemany("UPDATE work SET status='processing',attempts=attempts+1,updated_at=? WHERE document_id=?",
70
+ [(now, x) for x in ids]); self.db.commit(); return ids
71
+ except Exception:
72
+ self.db.rollback()
73
+ raise
74
+
75
+ def complete(self, ids: Iterable[str]) -> None:
76
+ with self.lock:
77
+ self._ensure_open()
78
+ self.db.executemany("UPDATE work SET status='done',updated_at=? WHERE document_id=?", [(time.time(), str(x)) for x in ids]); self.db.commit()
79
+
80
+ def fail(self, ids: Iterable[str], error: str) -> None:
81
+ if not str(error).strip():
82
+ error = "unknown materialization failure"
83
+ with self.lock:
84
+ self._ensure_open()
85
+ self.db.executemany("UPDATE work SET status=CASE WHEN attempts<? THEN 'pending' ELSE 'error' END,error=?,updated_at=? WHERE document_id=?",
86
+ [(self.max_retries, str(error), time.time(), str(x)) for x in ids]); self.db.commit()
87
+
88
+ def stats(self) -> dict[str, int]:
89
+ with self.lock:
90
+ self._ensure_open()
91
+ rows = self.db.execute("SELECT status,COUNT(*) FROM work GROUP BY status").fetchall()
92
+ out = {"pending": 0, "processing": 0, "done": 0, "error": 0}; out.update({str(k): int(v) for k, v in rows}); return out
93
+
94
+ def close(self) -> None:
95
+ with self.lock:
96
+ if self.db is not None:
97
+ self.db.close()
98
+ self.db = None
99
+
100
+
101
+ class MaterializationWorker:
102
+ def __init__(self, target_model: Any, documents: Any, cache: SQLiteVectorCache,
103
+ queue_path: str | Path, batch_size: int = 32, max_retries: int = 3,
104
+ state: Any | None = None):
105
+ self.target_model, self.documents, self.cache = target_model, documents, cache
106
+ try:
107
+ parsed_batch_size = int(batch_size)
108
+ exact = float(batch_size) == parsed_batch_size
109
+ except (TypeError, ValueError, OverflowError):
110
+ parsed_batch_size, exact = 0, False
111
+ if isinstance(batch_size, bool) or not exact or parsed_batch_size < 1:
112
+ raise ValueError("batch_size must be a positive integer")
113
+ self.queue = PersistentWorkQueue(queue_path, max_retries); self.batch_size = parsed_batch_size; self.state = state
114
+ self.stop_event = threading.Event(); self.thread: threading.Thread | None = None
115
+ self._closed = False
116
+ self._stats = {"materialized": 0, "errors": 0, "started_at": None, "last_throughput_docs_sec": 0.0}
117
+
118
+ def enqueue(self, ids: Iterable[str]) -> int:
119
+ if self._closed:
120
+ raise RuntimeError("materialization worker is closed")
121
+ return self.queue.enqueue(ids)
122
+
123
+ def start(self) -> None:
124
+ if self._closed:
125
+ raise RuntimeError("materialization worker is closed")
126
+ if self.thread and self.thread.is_alive(): return
127
+ self.stop_event.clear(); self._stats["started_at"] = time.time(); self.thread = threading.Thread(target=self._run, name="embedflow-materializer", daemon=True); self.thread.start()
128
+
129
+ def _run(self) -> None:
130
+ while not self.stop_event.is_set():
131
+ try:
132
+ ids = self.queue.claim(self.batch_size)
133
+ except RuntimeError:
134
+ if self.stop_event.is_set() or self._closed:
135
+ return
136
+ raise
137
+ except Exception as exc:
138
+ self._stats["errors"] += 1
139
+ if self.state:
140
+ self.state.add_error(f"materialization queue failed: {exc}")
141
+ time.sleep(0.1)
142
+ continue
143
+ if not ids:
144
+ time.sleep(0.05)
145
+ continue
146
+ try:
147
+ # Resolve documents individually so one missing/deleted row is
148
+ # marked as an error without poisoning the rest of a batch.
149
+ available: list[str] = []
150
+ texts: list[str] = []
151
+ missing: list[str] = []
152
+ for document_id in ids:
153
+ try:
154
+ value = self.documents.get([document_id]).get(document_id)
155
+ except Exception:
156
+ value = None
157
+ if value is None:
158
+ missing.append(document_id)
159
+ else:
160
+ available.append(document_id); texts.append(str(value))
161
+ if missing:
162
+ self.queue.fail(missing, "document text is unavailable")
163
+ self._stats["errors"] += len(missing)
164
+ if not available:
165
+ continue
166
+ vectors = self.target_model.encode_documents(texts, batch_size=self.batch_size)
167
+ self.cache.put(available, np.asarray(vectors, dtype="float32")); self.queue.complete(available)
168
+ self._stats["materialized"] += len(available)
169
+ elapsed = max(time.time() - float(self._stats["started_at"] or time.time()), 1e-6)
170
+ self._stats["last_throughput_docs_sec"] = self._stats["materialized"] / elapsed
171
+ if self.state: self.state.update(materialized_documents=self._stats["materialized"], materializer_throughput_docs_sec=self._stats["last_throughput_docs_sec"], queue=self.queue.stats())
172
+ except Exception as exc:
173
+ self.queue.fail(ids, repr(exc)); self._stats["errors"] += len(ids)
174
+ if self.state: self.state.add_error(f"materialization failed for {len(ids)} docs: {exc}")
175
+
176
+ def stop(self, timeout: float = 5.0) -> None:
177
+ self.stop_event.set()
178
+ if self.thread:
179
+ self.thread.join(timeout=timeout)
180
+ if self.thread.is_alive():
181
+ raise RuntimeError("materialization worker did not stop before timeout")
182
+
183
+ def stats(self) -> dict[str, Any]: return {**self._stats, "queue": self.queue.stats()}
184
+
185
+ def close(self, timeout: float = 5.0) -> None:
186
+ if self._closed:
187
+ return
188
+ self.stop(timeout=timeout)
189
+ self.queue.close()
190
+ self._closed = True
@@ -0,0 +1,78 @@
1
+ from __future__ import annotations
2
+
3
+ from dataclasses import asdict, dataclass
4
+ from typing import Any
5
+
6
+
7
+ @dataclass
8
+ class MigrationPlan:
9
+ status: str
10
+ target_model: str
11
+ candidate_depth: int
12
+ diagnostic: str
13
+ ann_status: str
14
+ corpus_size: int
15
+ cached_target_vectors: int
16
+ cache_fraction: float
17
+ migration_strategy: str
18
+ warnings: list[str]
19
+ rationale: str = ""
20
+ source_model: str = ""
21
+
22
+ def to_dict(self) -> dict[str, Any]: return asdict(self)
23
+
24
+
25
+ def make_plan(*, target_model: str, corpus_size: int, diagnostic: str,
26
+ recommended_k: int | None, ann_status: str, cached_target_vectors: int,
27
+ default_k: int = 50, throughput_docs_per_second: float | None = None,
28
+ gpu_price_per_hour: float | None = None, source_model: str = "") -> MigrationPlan:
29
+ diagnostic = str(diagnostic or "UNKNOWN").upper()
30
+ if diagnostic not in {"SAFE", "EXPAND", "UNSAFE_OR_UNCERTAIN", "UNKNOWN"}:
31
+ raise ValueError(f"unsupported diagnostic: {diagnostic}")
32
+ ann_value = str(ann_status or "UNKNOWN").upper()
33
+ if ann_value not in {"PASS", "WARNING", "UNKNOWN"}:
34
+ raise ValueError(f"unsupported ANN status: {ann_status}")
35
+ try:
36
+ corpus_int = int(corpus_size)
37
+ corpus_exact = float(corpus_size) == corpus_int
38
+ cached_int = int(cached_target_vectors)
39
+ cached_exact = float(cached_target_vectors) == cached_int
40
+ except (TypeError, ValueError, OverflowError) as exc:
41
+ raise ValueError("corpus_size and cached_target_vectors must be integers") from exc
42
+ if isinstance(corpus_size, bool) or not corpus_exact or corpus_int < 0:
43
+ raise ValueError("corpus_size must be a non-negative integer")
44
+ if isinstance(cached_target_vectors, bool) or not cached_exact or cached_int < 0:
45
+ raise ValueError("cached_target_vectors must be a non-negative integer")
46
+ # ``candidate_depth: auto`` is a configuration convenience. The actual
47
+ # serving plan always resolves to an integer, using the probe recommendation
48
+ # when available and the conservative default otherwise.
49
+ if str(default_k).lower() == "auto":
50
+ fallback_k = 50
51
+ else:
52
+ try:
53
+ fallback_k = int(default_k)
54
+ if isinstance(default_k, bool) or float(default_k) != fallback_k or fallback_k < 1:
55
+ raise ValueError
56
+ except (TypeError, ValueError, OverflowError) as exc:
57
+ raise ValueError("default_k must be auto or a positive integer") from exc
58
+ raw_k = fallback_k if recommended_k is None else recommended_k
59
+ try:
60
+ k = int(raw_k)
61
+ exact_k = float(raw_k) == k
62
+ except (TypeError, ValueError, OverflowError) as exc:
63
+ raise ValueError("recommended_k must be a positive integer") from exc
64
+ if isinstance(raw_k, bool) or not exact_k or k < 1:
65
+ raise ValueError("recommended_k must be a positive integer")
66
+ warnings: list[str] = []
67
+ if diagnostic != "SAFE":
68
+ warnings.append("Compatibility is not a SAFE T2-v1 diagnostic; validate manually before relying on progressive migration.")
69
+ if ann_value in {"WARNING", "UNKNOWN"}:
70
+ warnings.append("ANN fidelity is not confirmed; T2-v1 does not diagnose ANN approximation error.")
71
+ if cached_int > corpus_int:
72
+ warnings.append("Cache count exceeds corpus size; check document identity and cache contract.")
73
+ fraction = (cached_int / corpus_int) if corpus_int else 0.0
74
+ strategy = "PROGRESSIVE" if diagnostic == "SAFE" else "PROGRESSIVE_WITH_REVIEW"
75
+ status = "READY" if diagnostic == "SAFE" and ann_value in {"PASS", "UNKNOWN"} else "REVIEW"
76
+ rationale = "Finite-pool diagnostic supports the selected starting K; this is an empirical recommendation, not a guarantee." if diagnostic == "SAFE" else "Use source-only/full-backfill fallback until compatibility and ANN health are reviewed."
77
+ return MigrationPlan(status, target_model, k, diagnostic, ann_value, corpus_int,
78
+ cached_int, float(fraction), strategy, warnings, rationale, source_model)
@@ -0,0 +1,81 @@
1
+ from __future__ import annotations
2
+
3
+ import json
4
+ import threading
5
+ from collections.abc import Iterable
6
+ from pathlib import Path
7
+ from typing import Any
8
+
9
+ from ..config import EmbedFlowConfig
10
+
11
+
12
+ class DocumentStore:
13
+ """Memory-indexed JSONL document store for progressive materialization."""
14
+
15
+ def __init__(self, path: str | Path, id_field: str = "id", text_field: str = "text"):
16
+ self.path = Path(path)
17
+ if not self.path.exists(): raise FileNotFoundError(self.path)
18
+ if not isinstance(id_field, str) or not id_field.strip() or not isinstance(text_field, str) or not text_field.strip():
19
+ raise ValueError("document id_field and text_field must be non-empty strings")
20
+ self.documents: dict[str, str] = {}
21
+ self.id_field, self.text_field = id_field, text_field
22
+ with self.path.open() as handle:
23
+ for line_no, line in enumerate(handle, 1):
24
+ if not line.strip(): continue
25
+ try: row = json.loads(line)
26
+ except json.JSONDecodeError as exc: raise ValueError(f"invalid JSON at {self.path}:{line_no}") from exc
27
+ if not isinstance(row, dict): raise ValueError(f"document row {line_no} in {self.path} must be a JSON object")
28
+ if id_field not in row or text_field not in row: raise ValueError(f"missing {id_field}/{text_field} at {self.path}:{line_no}")
29
+ document_id = str(row[id_field])
30
+ if not document_id.strip(): raise ValueError(f"empty document ID at {self.path}:{line_no}")
31
+ if document_id in self.documents: raise ValueError(f"duplicate document ID {document_id!r}")
32
+ text = row[text_field]
33
+ if text is None or not str(text).strip(): raise ValueError(f"missing document text for {document_id!r}")
34
+ self.documents[document_id] = str(text)
35
+ if not self.documents: raise ValueError(f"document file is empty: {self.path}")
36
+
37
+ def get(self, document_ids: Iterable[str]) -> dict[str, str]:
38
+ ids = [str(x) for x in document_ids]
39
+ missing = [x for x in ids if x not in self.documents]
40
+ if missing: raise KeyError(f"document text missing for IDs: {missing[:5]}")
41
+ return {x: self.documents[x] for x in ids}
42
+
43
+ def size(self) -> int: return len(self.documents)
44
+
45
+
46
+ class MigrationState:
47
+ def __init__(self, path: str | Path, cfg: EmbedFlowConfig):
48
+ self.path = Path(path); self._lock = threading.RLock()
49
+ self.data: dict[str, Any] = {
50
+ "source_model": cfg.source.model, "target_model": cfg.target.model,
51
+ "target_model_fingerprint": cfg.target.fingerprint, "status": "INITIALIZING",
52
+ "diagnostic": "UNKNOWN", "ann_status": "UNKNOWN", "candidate_depth": cfg.migration.candidate_depth,
53
+ "migration_strategy": "PROGRESSIVE", "queries": 0, "errors": [],
54
+ }
55
+ if self.path.exists():
56
+ try:
57
+ loaded = json.loads(self.path.read_text())
58
+ except (OSError, json.JSONDecodeError) as exc: raise ValueError(f"invalid state file {self.path}") from exc
59
+ if not isinstance(loaded, dict):
60
+ raise ValueError(f"state file {self.path} must contain a JSON object")
61
+ stored_fingerprint = loaded.get("target_model_fingerprint")
62
+ if stored_fingerprint and str(stored_fingerprint) != str(cfg.target.fingerprint):
63
+ raise ValueError("state file target model fingerprint does not match the configured target contract")
64
+ self.data.update(loaded)
65
+ self.save()
66
+
67
+ def update(self, **values: Any) -> None:
68
+ with self._lock:
69
+ self.data.update(values); self.save()
70
+
71
+ def add_error(self, message: str) -> None:
72
+ with self._lock:
73
+ errors = self.data.setdefault("errors", []); errors.append(str(message)); self.data["errors"] = errors[-100:]; self.save()
74
+
75
+ def save(self) -> None:
76
+ self.path.parent.mkdir(parents=True, exist_ok=True)
77
+ tmp = self.path.with_suffix(self.path.suffix + ".tmp")
78
+ tmp.write_text(json.dumps(self.data, indent=2, sort_keys=True, ensure_ascii=False)); tmp.replace(self.path)
79
+
80
+ def snapshot(self) -> dict[str, Any]:
81
+ with self._lock: return json.loads(json.dumps(self.data))
@@ -0,0 +1,4 @@
1
+ from .base import EmbeddingModel
2
+ from .huggingface import HashEmbeddingModel, HuggingFaceEmbeddingModel, load_embedding_model
3
+
4
+ __all__ = ["EmbeddingModel", "HashEmbeddingModel", "HuggingFaceEmbeddingModel", "load_embedding_model"]