diff --git a/src/application/ports/embedding.py b/src/application/ports/embedding.py index 8048ee9..0e4d114 100644 --- a/src/application/ports/embedding.py +++ b/src/application/ports/embedding.py @@ -19,6 +19,14 @@ class DenseEmbedder(Protocol): """ name: str + model_version: str + """Identifies the model that produced these vectors (ADR-0001). + + Written into every point's `embedding_model_version` payload field, which + exists so a future model swap can tell which chunks need re-embedding. The + embedder is what knows this, so it is reported here rather than + reconstructed from configuration at the call site. + """ async def embed_batch(self, texts: Sequence[str]) -> list[list[float]]: """Return one vector per input text, same order. Raises `EmbedderError` @@ -37,6 +45,12 @@ class SparseEmbedder(Protocol): """ name: str + model_version: str + """Identifies the analyzer/parameters that produced these vectors. + + Same purpose as `DenseEmbedder.model_version`; for BM25 the "model" is the + analyzer choice (ADR-0005), which is equally a re-embedding trigger. + """ def embed_batch(self, texts: Sequence[str], *, query: bool = False) -> list[SparseVector]: """Return one sparse vector per input text, same order. diff --git a/src/infrastructure/embedding/bm25.py b/src/infrastructure/embedding/bm25.py index a6ea684..ef74cc4 100644 --- a/src/infrastructure/embedding/bm25.py +++ b/src/infrastructure/embedding/bm25.py @@ -86,6 +86,7 @@ class Bm25SparseEmbedder: name = "sparse" def __init__(self, settings: SparseEmbeddingSettings) -> None: + self.model_version = f"bm25-{settings.analyzer}" self._settings = settings def embed_batch(self, texts: Sequence[str], *, query: bool = False) -> list[SparseVector]: diff --git a/src/infrastructure/embedding/openai_compatible.py b/src/infrastructure/embedding/openai_compatible.py index cb50daa..b8018b3 100644 --- a/src/infrastructure/embedding/openai_compatible.py +++ b/src/infrastructure/embedding/openai_compatible.py @@ -52,6 +52,7 @@ class OpenAICompatibleEmbedder: keep_alive: str | None = None, ) -> None: self.name = name + self.model_version = model self._client = client self._model = model self._dimensions = dimensions diff --git a/tests/fakes.py b/tests/fakes.py index c834d81..d67f183 100644 --- a/tests/fakes.py +++ b/tests/fakes.py @@ -29,6 +29,7 @@ class FakeDenseEmbedder: name: str dimensions: int = 4 + model_version: str = "fake-dense-v1" calls: list[list[str]] = field(default_factory=list) fail_next: bool = False delay_seconds: float = 0.0 @@ -49,6 +50,7 @@ class FakeSparseEmbedder: """A scripted `SparseEmbedder`. Returns an empty sparse vector per text.""" name: str = "sparse" + model_version: str = "fake-sparse-v1" calls: list[list[str]] = field(default_factory=list) fail_next: bool = False @@ -58,3 +60,4 @@ class FakeSparseEmbedder: self.fail_next = False raise RuntimeError("simulated embedder failure") return [SparseVector(indices=[], values=[]) for _ in texts] + diff --git a/tests/unit/application/test_embedding.py b/tests/unit/application/test_embedding.py index b3e179b..c08e517 100644 --- a/tests/unit/application/test_embedding.py +++ b/tests/unit/application/test_embedding.py @@ -37,6 +37,7 @@ class _TrackingDenseEmbedder: """ name: str + model_version: str = "stub-v1" dimensions: int = 3 batches: list[list[str]] = field(default_factory=list) in_flight: int = 0 @@ -54,6 +55,7 @@ class _TrackingDenseEmbedder: @dataclass class _FailingDenseEmbedder: name: str + model_version: str = "stub-v1" async def embed_batch(self, texts: Sequence[str]) -> list[list[float]]: raise RuntimeError("boom") @@ -62,6 +64,7 @@ class _FailingDenseEmbedder: @dataclass class _StubSparseEmbedder: name: str = "sparse" + model_version: str = "stub-sparse-v1" calls: list[list[str]] = field(default_factory=list) def embed_batch(self, texts: Sequence[str], *, query: bool = False) -> list[SparseVector]: @@ -72,6 +75,7 @@ class _StubSparseEmbedder: @dataclass class _FailingSparseEmbedder: name: str = "sparse" + model_version: str = "stub-sparse-v1" def embed_batch(self, texts: Sequence[str], *, query: bool = False) -> list[SparseVector]: raise RuntimeError("boom")