fix(embedding): isolate caches by vector space

This commit is contained in:
jinli.yl 2026-08-10 20:11:27 +08:00
parent 5a5855f5ff
commit 2ef822b9ee
4 changed files with 237 additions and 20 deletions

View file

@ -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

View file

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

View file

@ -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

View file

@ -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=<new provider>).
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())