mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-08-28 05:25:04 +00:00
fix(embedding): isolate caches by vector space (#442)
Some checks are pending
Pre-commit / run (ubuntu-latest) (push) Waiting to run
Tests ReMe / Unit Tests - py3.11 (push) Waiting to run
Tests ReMe / Unit Tests - py3.12 (push) Waiting to run
Tests ReMe / Unit Tests - py3.13 (push) Waiting to run
Windows Smoke / CLI smoke - py3.11 (push) Waiting to run
Some checks are pending
Pre-commit / run (ubuntu-latest) (push) Waiting to run
Tests ReMe / Unit Tests - py3.11 (push) Waiting to run
Tests ReMe / Unit Tests - py3.12 (push) Waiting to run
Tests ReMe / Unit Tests - py3.13 (push) Waiting to run
Windows Smoke / CLI smoke - py3.11 (push) Waiting to run
* fix(embedding): isolate caches by vector space
* fix(embedding): stabilize cache space switching
* Revert "fix(embedding): stabilize cache space switching"
This reverts commit 74193c9a0a.
* fix(embedding): include resolved OpenAI endpoint in cache ID
* fix(embedding): stabilize cache space switching
* fix(embedding): isolate Ollama endpoint caches
This commit is contained in:
parent
5a5855f5ff
commit
21057931a9
4 changed files with 474 additions and 28 deletions
|
|
@ -1,5 +1,7 @@
|
|||
"""AgentScope embedding model wrappers."""
|
||||
|
||||
import hashlib
|
||||
import os
|
||||
from typing import Any
|
||||
|
||||
from agentscope.credential import (
|
||||
|
|
@ -23,6 +25,9 @@ class BaseAsEmbedding(BaseComponent):
|
|||
|
||||
component_type = ComponentEnum.AS_EMBEDDING
|
||||
credential_cls: type[CredentialBase]
|
||||
endpoint_fields = ("base_url", "host")
|
||||
endpoint_env: str | None = None
|
||||
default_endpoint = ""
|
||||
|
||||
def __init__(self, **kwargs) -> None:
|
||||
super().__init__(**kwargs)
|
||||
|
|
@ -38,6 +43,78 @@ 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._model_endpoint(),
|
||||
)
|
||||
return (
|
||||
self.backend or self.credential_cls.__name__,
|
||||
str(self.kwargs.get("model") or ""),
|
||||
str(self.dimensions),
|
||||
self._configured_endpoint(),
|
||||
)
|
||||
|
||||
@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 ""
|
||||
|
||||
def _configured_endpoint(self) -> str:
|
||||
"""Resolve the provider endpoint without constructing it eagerly."""
|
||||
credential = self.kwargs.get("credential")
|
||||
if isinstance(credential, dict):
|
||||
fields = getattr(self.credential_cls, "model_fields", {})
|
||||
for name in self.endpoint_fields:
|
||||
if name in credential:
|
||||
value = credential[name]
|
||||
if value is not None:
|
||||
return str(value).rstrip("/")
|
||||
continue
|
||||
field = fields.get(name)
|
||||
if field is not None and not field.is_required():
|
||||
value = field.get_default(call_default_factory=True)
|
||||
if value:
|
||||
return str(value).rstrip("/")
|
||||
else:
|
||||
endpoint = self._endpoint(credential)
|
||||
if endpoint:
|
||||
return endpoint
|
||||
if self.endpoint_env:
|
||||
endpoint = os.environ.get(self.endpoint_env)
|
||||
if endpoint:
|
||||
return endpoint.rstrip("/")
|
||||
return self.default_endpoint
|
||||
|
||||
def _model_endpoint(self) -> str:
|
||||
"""Return the endpoint used by the constructed provider."""
|
||||
assert self.model is not None
|
||||
endpoint = self._endpoint(getattr(self.model, "credential", self.kwargs.get("credential")))
|
||||
return endpoint or self._configured_endpoint()
|
||||
|
||||
async def __call__(self, inputs: list[Any], **kwargs) -> list[list[float]]:
|
||||
self._ensure_model()
|
||||
assert self.model is not None
|
||||
|
|
@ -78,6 +155,16 @@ class OpenAIAsEmbedding(BaseAsEmbedding):
|
|||
"""OpenAI embedding model wrapper."""
|
||||
|
||||
credential_cls = OpenAICredential
|
||||
endpoint_fields = ("base_url",)
|
||||
endpoint_env = "OPENAI_BASE_URL"
|
||||
default_endpoint = "https://api.openai.com/v1"
|
||||
|
||||
def _model_endpoint(self) -> str:
|
||||
"""Read the endpoint the OpenAI client actually resolved."""
|
||||
assert self.model is not None
|
||||
client = getattr(self.model, "client", None)
|
||||
base_url = getattr(client, "base_url", None)
|
||||
return str(base_url).rstrip("/") if base_url is not None else super()._model_endpoint()
|
||||
|
||||
|
||||
@R.register("dashscope")
|
||||
|
|
@ -106,6 +193,9 @@ class OllamaAsEmbedding(BaseAsEmbedding):
|
|||
"""Ollama embedding model wrapper."""
|
||||
|
||||
credential_cls = OllamaCredential
|
||||
endpoint_fields = ("host",)
|
||||
endpoint_env = "OLLAMA_HOST"
|
||||
default_endpoint = "http://127.0.0.1:11434"
|
||||
|
||||
|
||||
__all__ = [
|
||||
|
|
|
|||
|
|
@ -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_space:
|
||||
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,34 @@ 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:
|
||||
while True:
|
||||
space = self.vector_space_id
|
||||
if space == self._cache_space:
|
||||
return
|
||||
previous = self._cache_space
|
||||
snapshot = list(self._cache.items())
|
||||
if previous and self.enable_cache and snapshot:
|
||||
await asyncio.to_thread(self._dump_sync, previous, snapshot)
|
||||
if space != self.vector_space_id:
|
||||
continue
|
||||
dimensions = self.dimensions
|
||||
cache: OrderedDict[str, np.ndarray] = OrderedDict()
|
||||
if self.enable_cache and self._cache_path(space).exists():
|
||||
cache = await asyncio.to_thread(self._load_sync, space, dimensions)
|
||||
if space != self.vector_space_id:
|
||||
continue
|
||||
self._cache = cache
|
||||
self._cache_space = space
|
||||
return
|
||||
|
||||
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,36 +232,41 @@ class LocalEmbeddingStore(BaseEmbeddingStore):
|
|||
|
||||
async def load(self) -> None:
|
||||
self._cache.clear()
|
||||
if not self.enable_cache or not self.cache_path.exists():
|
||||
return
|
||||
await asyncio.to_thread(self._load_sync)
|
||||
self._cache_space = ""
|
||||
await self._sync_cache_space()
|
||||
|
||||
def _load_sync(self) -> None:
|
||||
def _load_sync(self, vector_space_id: str, dimensions: int) -> OrderedDict[str, np.ndarray]:
|
||||
path = self._cache_path(vector_space_id)
|
||||
cache: OrderedDict[str, np.ndarray] = OrderedDict()
|
||||
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:
|
||||
if len(emb) != dimensions:
|
||||
continue
|
||||
if len(self._cache) >= self.max_cache_size:
|
||||
if len(cache) >= self.max_cache_size:
|
||||
break
|
||||
self._cache[str(key)] = emb.astype(np.float16)
|
||||
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)
|
||||
return
|
||||
self.logger.info(f"Loaded {len(self._cache)} embeddings from {self.cache_path}")
|
||||
path.unlink(missing_ok=True)
|
||||
return cache
|
||||
self.logger.info(f"Loaded {len(cache)} embeddings from {path}")
|
||||
return cache
|
||||
|
||||
async def dump(self) -> None:
|
||||
if not self.enable_cache or not self._cache:
|
||||
await self._sync_cache_space()
|
||||
snapshot = list(self._cache.items())
|
||||
if not self.enable_cache or not snapshot:
|
||||
return
|
||||
await asyncio.to_thread(self._dump_sync)
|
||||
await asyncio.to_thread(self._dump_sync, self._cache_space, snapshot)
|
||||
|
||||
def _dump_sync(self) -> None:
|
||||
self.cache_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
keys = np.array(list(self._cache.keys()), dtype=str)
|
||||
embeddings = np.stack(list(self._cache.values()))
|
||||
def _dump_sync(self, vector_space_id: str, cache: list[tuple[str, np.ndarray]]) -> None:
|
||||
path = self._cache_path(vector_space_id)
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
keys = np.array([key for key, _ in cache], dtype=str)
|
||||
embeddings = np.stack([embedding for _, embedding in cache])
|
||||
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(cache)} embeddings to {path}")
|
||||
except Exception:
|
||||
self.logger.exception("Failed to save embedding cache")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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 DashScopeAsEmbedding, OllamaAsEmbedding, 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,26 @@ 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 = "https://api.openai.com/v1",
|
||||
):
|
||||
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 +104,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 +195,287 @@ 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_vector_space_id_resolves_default_endpoint_before_lazy_construction():
|
||||
"""Credential defaults must not change the cache namespace on the first request."""
|
||||
embedding = DashScopeAsEmbedding(
|
||||
name="t_space_default_endpoint",
|
||||
model="text-embedding-v3",
|
||||
dimensions=1024,
|
||||
credential={"api_key": "test"},
|
||||
)
|
||||
before = embedding.vector_space_id
|
||||
|
||||
embedding._ensure_model()
|
||||
|
||||
assert embedding.vector_space_id == before
|
||||
|
||||
|
||||
def test_openai_vector_space_id_uses_sdk_resolved_endpoint(monkeypatch):
|
||||
"""OPENAI_BASE_URL must separate caches and stay stable after lazy construction."""
|
||||
clients = []
|
||||
ids = []
|
||||
try:
|
||||
for endpoint in ("https://provider-a.example/v1", "https://provider-b.example/v1"):
|
||||
monkeypatch.setenv("OPENAI_BASE_URL", endpoint)
|
||||
embedding = OpenAIAsEmbedding(
|
||||
name="t_space_openai_env",
|
||||
backend="openai",
|
||||
model="text-embedding-3-small",
|
||||
dimensions=1536,
|
||||
credential={"api_key": "test"},
|
||||
)
|
||||
before = embedding.vector_space_id
|
||||
|
||||
embedding._ensure_model()
|
||||
clients.append(embedding.model.client)
|
||||
|
||||
assert str(embedding.model.client.base_url).rstrip("/") == endpoint
|
||||
assert embedding.vector_space_id == before
|
||||
ids.append(before)
|
||||
|
||||
assert ids[0] != ids[1]
|
||||
finally:
|
||||
for client in clients:
|
||||
run(client.close())
|
||||
|
||||
|
||||
def test_openai_vector_space_id_uses_sdk_default_endpoint(monkeypatch):
|
||||
"""The SDK default URL must not change the namespace on first construction."""
|
||||
monkeypatch.delenv("OPENAI_BASE_URL", raising=False)
|
||||
embedding = OpenAIAsEmbedding(
|
||||
name="t_space_openai_default",
|
||||
backend="openai",
|
||||
model="text-embedding-3-small",
|
||||
dimensions=1536,
|
||||
credential={"api_key": "test"},
|
||||
)
|
||||
before = embedding.vector_space_id
|
||||
|
||||
embedding._ensure_model()
|
||||
try:
|
||||
assert str(embedding.model.client.base_url).rstrip("/") == "https://api.openai.com/v1"
|
||||
assert embedding.vector_space_id == before
|
||||
finally:
|
||||
run(embedding.model.client.close())
|
||||
|
||||
|
||||
def test_ollama_vector_space_id_uses_sdk_resolved_endpoint(monkeypatch):
|
||||
"""OLLAMA_HOST must separate caches and stay stable after lazy construction."""
|
||||
ids = []
|
||||
for endpoint in ("http://provider-a.example:11434", "http://provider-b.example:11434"):
|
||||
monkeypatch.setenv("OLLAMA_HOST", endpoint)
|
||||
embedding = OllamaAsEmbedding(
|
||||
name="t_space_ollama_env",
|
||||
backend="ollama",
|
||||
model="nomic-embed-text",
|
||||
dimensions=768,
|
||||
credential={},
|
||||
)
|
||||
before = embedding.vector_space_id
|
||||
|
||||
embedding._ensure_model()
|
||||
|
||||
assert embedding.model.host is None
|
||||
assert embedding.vector_space_id == before
|
||||
ids.append(before)
|
||||
|
||||
assert ids[0] != ids[1]
|
||||
|
||||
|
||||
def test_ollama_vector_space_id_uses_sdk_default_endpoint(monkeypatch):
|
||||
"""The Ollama SDK default URL must remain stable after lazy construction."""
|
||||
monkeypatch.delenv("OLLAMA_HOST", raising=False)
|
||||
embedding = OllamaAsEmbedding(
|
||||
name="t_space_ollama_default",
|
||||
backend="ollama",
|
||||
model="nomic-embed-text",
|
||||
dimensions=768,
|
||||
credential={},
|
||||
)
|
||||
before = embedding.vector_space_id
|
||||
|
||||
embedding._ensure_model()
|
||||
|
||||
assert embedding.vector_space[-1] == "http://127.0.0.1:11434"
|
||||
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_cache_space_is_rechecked_after_async_load(monkeypatch, tmp_path):
|
||||
"""A provider switch during disk I/O must not publish the stale namespace."""
|
||||
|
||||
async def go():
|
||||
monkeypatch.setattr(
|
||||
LocalEmbeddingStore,
|
||||
"component_metadata_path",
|
||||
property(lambda _self: tmp_path),
|
||||
)
|
||||
embedding = OpenAIAsEmbedding(name="t_space_load_race", backend="openai", model="v3", dimensions=2)
|
||||
store = LocalEmbeddingStore(name="t_local_load_race")
|
||||
store.as_embedding = embedding
|
||||
await store.load()
|
||||
|
||||
embedding.model = FakeProviderModel("v4")
|
||||
v4_space = embedding.vector_space_id
|
||||
np.savez(
|
||||
store._cache_path(v4_space),
|
||||
keys=np.array([store._cache_key("hello")]),
|
||||
embeddings=np.array([[4.0, 0.0]], dtype=np.float16),
|
||||
)
|
||||
original_to_thread = asyncio.to_thread
|
||||
|
||||
async def switch_during_load(func, *args):
|
||||
if getattr(func, "__name__", "") == "_load_sync":
|
||||
embedding.model = FakeProviderModel("v3")
|
||||
return await original_to_thread(func, *args)
|
||||
|
||||
monkeypatch.setattr(asyncio, "to_thread", switch_during_load)
|
||||
await store._sync_cache_space()
|
||||
|
||||
assert store._cache_space == embedding.vector_space_id
|
||||
assert not store._cache
|
||||
|
||||
run(go())
|
||||
|
||||
|
||||
def test_completed_request_only_writes_to_its_active_cache_space():
|
||||
"""A v3 request must not populate v4 after the provider switches back to v3."""
|
||||
|
||||
async def go():
|
||||
embedding = OpenAIAsEmbedding(name="t_space_write_race", backend="openai", model="v3", dimensions=2)
|
||||
store = LocalEmbeddingStore(name="t_local_write_race")
|
||||
store.as_embedding = embedding
|
||||
store._cache_space = embedding.vector_space_id
|
||||
|
||||
async def compute_after_round_trip(_batch, **_kwargs):
|
||||
embedding.model = FakeProviderModel("v4")
|
||||
store._cache_space = embedding.vector_space_id
|
||||
embedding.model = FakeProviderModel("v3")
|
||||
return [(0, "key", np.array([3.0, 0.0], dtype=np.float16))]
|
||||
|
||||
store._compute_batch = compute_after_round_trip
|
||||
await store._fill_misses([(0, "text", "key")], [None])
|
||||
|
||||
assert "key" not in store._cache
|
||||
|
||||
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())
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue