diff --git a/reme/components/as_embedding/__init__.py b/reme/components/as_embedding/__init__.py index 2a66167d..1da9c167 100644 --- a/reme/components/as_embedding/__init__.py +++ b/reme/components/as_embedding/__init__.py @@ -1,5 +1,6 @@ """AgentScope embedding model wrappers.""" +import hashlib from typing import Any from agentscope.credential import ( @@ -38,6 +39,46 @@ class BaseAsEmbedding(BaseComponent): raise RuntimeError("Embedding dimensions are required before provider initialization.") return int(dimensions) + @property + def vector_space(self) -> tuple[str, ...]: + """Return the fields that make vectors from two embedding setups incompatible. + + The wrapper identifies the provider consistently before and after lazy model + construction. Model details come from the live provider when one has been + injected at runtime, otherwise they come from the configured kwargs. + """ + if self.model is not None: + return ( + self.backend or self.credential_cls.__name__, + str(getattr(self.model, "model", self.kwargs.get("model") or "")), + str(self.dimensions), + self._endpoint(getattr(self.model, "credential", self.kwargs.get("credential"))), + ) + return ( + self.backend or self.credential_cls.__name__, + str(self.kwargs.get("model") or ""), + str(self.dimensions), + self._endpoint(self.kwargs.get("credential")), + ) + + @property + def vector_space_id(self) -> str: + """Return a short digest of :attr:`vector_space` for naming persisted vectors. + + Cache consumers use this digest to avoid reusing vectors produced by a + different embedding setup. + """ + return hashlib.sha256("\x1f".join(self.vector_space).encode()).hexdigest()[:12] + + @staticmethod + def _endpoint(credential: Any) -> str: + """Read the provider endpoint from a credential object or a raw kwargs dict.""" + for field in ("base_url", "host"): + value = credential.get(field) if isinstance(credential, dict) else getattr(credential, field, None) + if value: + return str(value).rstrip("/") + return "" + async def __call__(self, inputs: list[Any], **kwargs) -> list[list[float]]: self._ensure_model() assert self.model is not None diff --git a/reme/components/embedding_store/local_embedding_store.py b/reme/components/embedding_store/local_embedding_store.py index 052a0851..3f6d137c 100644 --- a/reme/components/embedding_store/local_embedding_store.py +++ b/reme/components/embedding_store/local_embedding_store.py @@ -35,7 +35,8 @@ class LocalEmbeddingStore(BaseEmbeddingStore): self.enable_cache = enable_cache self.cache_version = cache_version self._cache: OrderedDict[str, np.ndarray] = OrderedDict() - self._key_suffix: bytes = b"" + self._cache_space: str = "" + self._cache_space_lock = asyncio.Lock() @property def dimensions(self) -> int: @@ -43,13 +44,25 @@ class LocalEmbeddingStore(BaseEmbeddingStore): assert self.as_embedding is not None, "embedding component not bound" return self.as_embedding.dimensions + @property + def vector_space_id(self) -> str: + """Return the digest of the vector space the bound provider currently produces.""" + assert self.as_embedding is not None, "embedding component not bound" + return self.as_embedding.vector_space_id + @property def cache_path(self) -> Path: - """Return the path to the disk cache file.""" - return self.component_metadata_path / f"{self.name}_{self.cache_version}.npz" + """Return the disk cache file for the current vector space. + + Each vector space owns its own file, so switching the embedding model cannot + read or overwrite vectors that belong to a different model. + """ + return self._cache_path(self.vector_space_id) + + def _cache_path(self, vector_space_id: str) -> Path: + return self.component_metadata_path / f"{self.name}_{self.cache_version}_{vector_space_id}.npz" async def _start(self) -> None: - self._key_suffix = f"|{self.dimensions}".encode() await self.load() async def _close(self) -> None: @@ -76,6 +89,7 @@ class LocalEmbeddingStore(BaseEmbeddingStore): # -- Public API -- async def get_embeddings(self, input_text: list[str], **kwargs) -> list[np.ndarray | None]: + await self._sync_cache_space() texts = [self._truncate(t) for t in input_text] results, misses = self._partition_by_cache(texts) if misses: @@ -97,12 +111,14 @@ class LocalEmbeddingStore(BaseEmbeddingStore): return results, misses async def _fill_misses(self, misses: list[Miss], results: list[np.ndarray | None], **kwargs) -> None: + vector_space_id = self._cache_space size = self.max_batch_size for start in range(0, len(misses), size): batch = misses[start : start + size] for idx, key, emb in await self._compute_batch(batch, **kwargs): results[idx] = emb - self._cache_put(key, emb) + if vector_space_id == self.vector_space_id: + self._cache_put(key, emb) async def _compute_batch(self, batch: list[Miss], **kwargs) -> list[tuple[int, str, np.ndarray]]: texts = [text for _, text, _ in batch] @@ -165,8 +181,25 @@ class LocalEmbeddingStore(BaseEmbeddingStore): # -- Cache -- + async def _sync_cache_space(self) -> None: + """Persist the previous space and restore the newly active space.""" + space = self.vector_space_id + if space == self._cache_space: + return + async with self._cache_space_lock: + space = self.vector_space_id + if space == self._cache_space: + return + previous = self._cache_space + if previous and self.enable_cache and self._cache: + await asyncio.to_thread(self._dump_sync, previous) + self._cache.clear() + self._cache_space = space + if self.enable_cache and self._cache_path(space).exists(): + await asyncio.to_thread(self._load_sync, space) + def _cache_key(self, text: str) -> str: - return hashlib.sha256(text.encode() + self._key_suffix).hexdigest() + return hashlib.sha256(text.encode()).hexdigest() def _cache_get(self, key: str) -> np.ndarray | None: if not self.enable_cache or key not in self._cache: @@ -190,13 +223,15 @@ class LocalEmbeddingStore(BaseEmbeddingStore): async def load(self) -> None: self._cache.clear() - if not self.enable_cache or not self.cache_path.exists(): + self._cache_space = self.vector_space_id + if not self.enable_cache or not self._cache_path(self._cache_space).exists(): return - await asyncio.to_thread(self._load_sync) + await asyncio.to_thread(self._load_sync, self._cache_space) - def _load_sync(self) -> None: + def _load_sync(self, vector_space_id: str) -> None: + path = self._cache_path(vector_space_id) try: - with np.load(self.cache_path) as data: + with np.load(path) as data: for key, emb in zip(data["keys"], data["embeddings"]): if len(emb) != self.dimensions: continue @@ -205,21 +240,23 @@ class LocalEmbeddingStore(BaseEmbeddingStore): self._cache[str(key)] = emb.astype(np.float16) except Exception: self.logger.exception("Failed to load embedding cache, removing") - self.cache_path.unlink(missing_ok=True) + path.unlink(missing_ok=True) return - self.logger.info(f"Loaded {len(self._cache)} embeddings from {self.cache_path}") + self.logger.info(f"Loaded {len(self._cache)} embeddings from {path}") async def dump(self) -> None: + await self._sync_cache_space() if not self.enable_cache or not self._cache: return - await asyncio.to_thread(self._dump_sync) + await asyncio.to_thread(self._dump_sync, self._cache_space) - def _dump_sync(self) -> None: - self.cache_path.parent.mkdir(parents=True, exist_ok=True) + def _dump_sync(self, vector_space_id: str) -> None: + path = self._cache_path(vector_space_id) + path.parent.mkdir(parents=True, exist_ok=True) keys = np.array(list(self._cache.keys()), dtype=str) embeddings = np.stack(list(self._cache.values())) try: - np.savez(self.cache_path, keys=keys, embeddings=embeddings) - self.logger.info(f"Saved {len(self._cache)} embeddings to {self.cache_path}") + np.savez(path, keys=keys, embeddings=embeddings) + self.logger.info(f"Saved {len(self._cache)} embeddings to {path}") except Exception: self.logger.exception("Failed to save embedding cache") diff --git a/tests/unit/test_as_embedding_lazy.py b/tests/unit/test_as_embedding_lazy.py index dc81f85f..14a4bb27 100644 --- a/tests/unit/test_as_embedding_lazy.py +++ b/tests/unit/test_as_embedding_lazy.py @@ -51,14 +51,22 @@ def test_provider_is_constructed_once_on_first_call(): async def go(): FakeModel.constructions = 0 - embedding = LazyAsEmbedding(dimensions=3, credential={"token": "test"}, parameters={"mode": "test"}) + embedding = LazyAsEmbedding( + backend="fake", + model="fake-model", + dimensions=3, + credential={"token": "test"}, + parameters={"mode": "test"}, + ) await embedding.start() assert embedding.model is None assert embedding.dimensions == 3 assert FakeModel.constructions == 0 + vector_space_id = embedding.vector_space_id assert await embedding(["first"]) == [[0.0, 0.0, 0.0]] + assert embedding.vector_space_id == vector_space_id assert await embedding(["second"]) == [[0.0, 0.0, 0.0]] assert FakeModel.constructions == 1 diff --git a/tests/unit/test_local_embedding_store.py b/tests/unit/test_local_embedding_store.py index 98d6b93a..fe432a8e 100644 --- a/tests/unit/test_local_embedding_store.py +++ b/tests/unit/test_local_embedding_store.py @@ -1,11 +1,13 @@ -"""Regression tests for LocalEmbeddingStore dimension handling.""" +"""Regression tests for LocalEmbeddingStore dimension and vector space handling.""" # pylint: disable=protected-access import asyncio +from types import SimpleNamespace import numpy as np +from reme.components.as_embedding import OpenAIAsEmbedding from reme.components.embedding_store.base_embedding_store import BaseEmbeddingStore from reme.components.embedding_store.local_embedding_store import LocalEmbeddingStore from reme.schema import EmbNode @@ -15,6 +17,7 @@ class FakeAsEmbedding: """Fake AgentScope embedding component.""" dimensions = 2 + vector_space_id = "fakespace000" async def __call__(self, texts: list[str], **_kwargs): return [[1.0] if text == "bad" else [1.0, 0.0] for text in texts] @@ -24,11 +27,21 @@ class BadHealthAsEmbedding: """Fake provider whose health probe returns the wrong dimension.""" dimensions = 2 + vector_space_id = "fakespace000" async def __call__(self, _texts: list[str], **_kwargs): return [[1.0]] +class FakeProviderModel: + """Stand-in for a constructed AgentScope embedding model object.""" + + def __init__(self, model: str, dimensions: int = 2, base_url: str = ""): + self.model = model + self.dimensions = dimensions + self.credential = SimpleNamespace(base_url=base_url) + + class InsufficientQuotaError(Exception): """OpenAI-compatible quota error used without importing the provider SDK.""" @@ -86,7 +99,6 @@ def test_compute_batch_rejects_embeddings_with_wrong_dimension(): async def go(): store = LocalEmbeddingStore(name="t_local_embedding_dim") store.as_embedding = FakeAsEmbedding() - store._key_suffix = f"|{store.dimensions}".encode() results = await store._compute_batch( [ @@ -178,3 +190,122 @@ def test_insufficient_quota_does_not_retry_without_opt_in(monkeypatch): assert not sleeps run(go()) + + +def test_vector_space_id_separates_models_of_equal_dimension(): + """Two models of the same width must not claim the same vector space.""" + common = {"backend": "openai", "dimensions": 1024, "credential": {"base_url": "https://example.com/v1"}} + v3 = OpenAIAsEmbedding(name="t_space_v3", model="text-embedding-v3", **common) + v4 = OpenAIAsEmbedding(name="t_space_v4", model="text-embedding-v4", **common) + + assert v3.dimensions == v4.dimensions + assert v3.vector_space_id != v4.vector_space_id + + +def test_vector_space_id_separates_endpoints_of_one_model_name(): + """The same model name served by two endpoints is two vector spaces.""" + common = {"backend": "openai", "model": "text-embedding-v4", "dimensions": 1024} + official = OpenAIAsEmbedding(name="t_space_a", credential={"base_url": "https://example.com/v1"}, **common) + self_hosted = OpenAIAsEmbedding(name="t_space_b", credential={"base_url": "http://127.0.0.1:8000/v1"}, **common) + + assert official.vector_space_id != self_hosted.vector_space_id + + +def test_vector_space_id_ignores_trailing_slash_and_api_key(): + """Cosmetic and secret credential changes must not invalidate stored vectors.""" + common = {"backend": "openai", "model": "text-embedding-v4", "dimensions": 1024} + first = OpenAIAsEmbedding( + name="t_space_c", + credential={"base_url": "https://example.com/v1", "api_key": "key-one"}, + **common, + ) + second = OpenAIAsEmbedding( + name="t_space_d", + credential={"base_url": "https://example.com/v1/", "api_key": "key-two"}, + **common, + ) + + assert first.vector_space_id == second.vector_space_id + + +def test_vector_space_id_follows_a_model_swapped_in_after_start(): + """A provider replaced at runtime must win over the original kwargs.""" + embedding = OpenAIAsEmbedding(name="t_space_swap", backend="openai", model="v3", dimensions=2) + before = embedding.vector_space_id + + # Mirrors Application.update_component("as_embedding", "default", model=). + embedding.model = FakeProviderModel("v4") + + assert embedding.vector_space_id != before + + +def test_vector_space_id_is_stable_across_lazy_provider_construction(): + """Constructing the configured provider must not look like a model switch.""" + embedding = OpenAIAsEmbedding( + name="t_space_lazy", + backend="openai", + model="v3", + dimensions=2, + credential={"base_url": "https://example.com/v1"}, + ) + before = embedding.vector_space_id + + embedding.model = FakeProviderModel("v3", base_url="https://example.com/v1") + + assert embedding.vector_space_id == before + + +def test_cache_is_saved_and_restored_per_vector_space(monkeypatch, tmp_path): + """Switching models persists the old cache and restores it when switched back.""" + + async def go(): + monkeypatch.setattr( + LocalEmbeddingStore, + "component_metadata_path", + property(lambda _self: tmp_path), + ) + embedding = OpenAIAsEmbedding(name="t_space_store", backend="openai", model="v3", dimensions=2) + store = LocalEmbeddingStore(name="t_local_space") + store.as_embedding = embedding + await store.load() + + key = store._cache_key("hello") + store._cache_put(key, np.array([1.0, 0.0], dtype=np.float16)) + v3_path = store.cache_path + + embedding.model = FakeProviderModel("v4") + await store._sync_cache_space() + + assert store._cache_key("hello") == key + assert store.cache_path != v3_path + assert store._cache_get(store._cache_key("hello")) is None + assert v3_path.exists() + + embedding.model = FakeProviderModel("v3") + await store._sync_cache_space() + + np.testing.assert_array_equal(store._cache_get(key), np.array([1.0, 0.0], dtype=np.float16)) + + run(go()) + + +def test_start_ignores_cache_file_without_vector_space_tag(monkeypatch, tmp_path): + """An unattributable legacy cache is ignored without deleting derived data.""" + + async def go(): + monkeypatch.setattr( + LocalEmbeddingStore, + "component_metadata_path", + property(lambda _self: tmp_path), + ) + store = LocalEmbeddingStore(name="t_local_untagged") + store.as_embedding = FakeAsEmbedding() + untagged = tmp_path / f"{store.name}_{store.cache_version}.npz" + untagged.write_bytes(b"vectors from an unknown model") + + await store._start() + + assert untagged.exists() + assert not store._cache + + run(go())