mirror of
https://github.com/langgenius/dify.git
synced 2026-09-28 06:13:22 +08:00
perf(api): reduce knowledge retrieval latency (#42673)
This commit is contained in:
@@ -135,7 +135,14 @@ class ProviderConfiguration(BaseModel):
|
||||
)
|
||||
and ConfigurateMethod.PREDEFINED_MODEL not in self.provider.configurate_methods
|
||||
):
|
||||
self.provider.configurate_methods.append(ConfigurateMethod.PREDEFINED_MODEL)
|
||||
self.provider = self.provider.model_copy(
|
||||
update={
|
||||
"configurate_methods": [
|
||||
*self.provider.configurate_methods,
|
||||
ConfigurateMethod.PREDEFINED_MODEL,
|
||||
]
|
||||
},
|
||||
)
|
||||
return self
|
||||
|
||||
def bind_model_runtime(self, model_runtime: ModelRuntime) -> None:
|
||||
|
||||
@@ -16,9 +16,12 @@ metadata.
|
||||
|
||||
import logging
|
||||
import time
|
||||
from collections import OrderedDict
|
||||
from collections.abc import Generator, Mapping, Sequence
|
||||
from contextlib import contextmanager
|
||||
from hashlib import sha256
|
||||
from mimetypes import guess_type
|
||||
from threading import Lock
|
||||
from typing import Literal, Protocol
|
||||
|
||||
import zstandard
|
||||
@@ -109,6 +112,15 @@ class PluginService:
|
||||
PLUGIN_MODEL_PROVIDERS_LOCK_WAIT_INTERVAL = 0.05
|
||||
PLUGIN_MODEL_PROVIDERS_CACHE_COMPRESSION_PREFIX = b"\x00dify-plugin-model-providers-zstd-v1:"
|
||||
PLUGIN_MODEL_PROVIDERS_CACHE_COMPRESSION_MIN_BYTES = 64 * 1024
|
||||
# Provider declarations are tenant-scoped but contain no tenant credentials. Cache the parsed tuple by payload
|
||||
# content so unchanged Redis data does not pay the Pydantic validation cost on every retrieval request. The cached
|
||||
# declarations are shared and consumers must treat them as read-only. A changed payload always has a different key,
|
||||
# while the small LRU bounds per-process memory usage.
|
||||
PLUGIN_MODEL_PROVIDERS_PARSED_CACHE_MAX_ENTRIES = 8
|
||||
_parsed_plugin_model_providers_cache: OrderedDict[tuple[int, bytes], tuple[PluginModelProviderDeclaration, ...]] = (
|
||||
OrderedDict()
|
||||
)
|
||||
_parsed_plugin_model_providers_cache_lock = Lock()
|
||||
PLUGIN_INSTALL_TASK_TERMINAL_STATUSES = (PluginInstallTaskStatus.Success, PluginInstallTaskStatus.Failed)
|
||||
# Mirror the detail-panel endpoint query size so list reconciliation and
|
||||
# the visible endpoint drawer exercise the same daemon pagination path.
|
||||
@@ -217,6 +229,54 @@ class PluginService:
|
||||
except zstandard.ZstdError as exc:
|
||||
raise ValueError("Invalid compressed plugin model providers cache payload.") from exc
|
||||
|
||||
@staticmethod
|
||||
def _plugin_model_providers_payload_cache_key(payload: bytes | bytearray | str) -> tuple[int, bytes]:
|
||||
payload_bytes = payload.encode("utf-8") if isinstance(payload, str) else bytes(payload)
|
||||
return len(payload_bytes), sha256(payload_bytes).digest()
|
||||
|
||||
@classmethod
|
||||
def _get_parsed_plugin_model_providers_from_local_cache(
|
||||
cls, payload: bytes | bytearray | str
|
||||
) -> tuple[PluginModelProviderDeclaration, ...] | None:
|
||||
cache_key = cls._plugin_model_providers_payload_cache_key(payload)
|
||||
with cls._parsed_plugin_model_providers_cache_lock:
|
||||
providers = cls._parsed_plugin_model_providers_cache.pop(cache_key, None)
|
||||
if providers is not None:
|
||||
cls._parsed_plugin_model_providers_cache[cache_key] = providers
|
||||
return providers
|
||||
|
||||
@classmethod
|
||||
def _store_parsed_plugin_model_providers_in_local_cache(
|
||||
cls,
|
||||
payload: bytes | bytearray | str,
|
||||
providers: tuple[PluginModelProviderDeclaration, ...],
|
||||
) -> None:
|
||||
cache_key = cls._plugin_model_providers_payload_cache_key(payload)
|
||||
with cls._parsed_plugin_model_providers_cache_lock:
|
||||
cls._parsed_plugin_model_providers_cache.pop(cache_key, None)
|
||||
cls._parsed_plugin_model_providers_cache[cache_key] = providers
|
||||
while len(cls._parsed_plugin_model_providers_cache) > cls.PLUGIN_MODEL_PROVIDERS_PARSED_CACHE_MAX_ENTRIES:
|
||||
cls._parsed_plugin_model_providers_cache.popitem(last=False)
|
||||
|
||||
@classmethod
|
||||
def _parse_plugin_model_providers_cache_payload(
|
||||
cls, payload: bytes | bytearray | str
|
||||
) -> tuple[PluginModelProviderDeclaration, ...]:
|
||||
decoded_payload = cls._decode_plugin_model_providers_cache_payload(payload)
|
||||
return tuple(_provider_entities_adapter.validate_json(decoded_payload))
|
||||
|
||||
@classmethod
|
||||
def _get_or_parse_plugin_model_providers_cache_payload(
|
||||
cls, payload: bytes | bytearray | str
|
||||
) -> tuple[PluginModelProviderDeclaration, ...]:
|
||||
providers = cls._get_parsed_plugin_model_providers_from_local_cache(payload)
|
||||
if providers is not None:
|
||||
return providers
|
||||
|
||||
providers = cls._parse_plugin_model_providers_cache_payload(payload)
|
||||
cls._store_parsed_plugin_model_providers_in_local_cache(payload, providers)
|
||||
return providers
|
||||
|
||||
@classmethod
|
||||
def _load_plugin_model_providers_generation(cls, tenant_id: str) -> int | None:
|
||||
cache_key = cls._get_plugin_model_providers_generation_cache_key(tenant_id)
|
||||
@@ -274,8 +334,7 @@ class PluginService:
|
||||
continue
|
||||
|
||||
try:
|
||||
payload = cls._decode_plugin_model_providers_cache_payload(cached_providers)
|
||||
providers = tuple(_provider_entities_adapter.validate_json(payload))
|
||||
providers = cls._get_or_parse_plugin_model_providers_cache_payload(cached_providers)
|
||||
return providers, True
|
||||
except (TypeError, ValueError, ValidationError):
|
||||
logger.warning(
|
||||
@@ -305,6 +364,7 @@ class PluginService:
|
||||
_provider_entities_adapter.dump_json(list(providers))
|
||||
)
|
||||
redis_client.setex(cache_key, dify_config.PLUGIN_MODEL_PROVIDERS_CACHE_TTL, payload)
|
||||
cls._store_parsed_plugin_model_providers_in_local_cache(payload, tuple(providers))
|
||||
except (RedisError, RuntimeError):
|
||||
logger.warning("Failed to cache plugin model providers for tenant %s.", tenant_id, exc_info=True)
|
||||
|
||||
|
||||
@@ -83,6 +83,12 @@ class _LazyEmbeddings(Embeddings):
|
||||
|
||||
@override
|
||||
def embed_query(self, text: str) -> list[float]:
|
||||
provider = self._dataset.embedding_model_provider
|
||||
model_name = self._dataset.embedding_model
|
||||
if provider and model_name:
|
||||
cached_embedding = CacheEmbedding.get_cached_query_embedding(provider, model_name, text)
|
||||
if cached_embedding is not None:
|
||||
return cached_embedding
|
||||
return self._ensure().embed_query(text)
|
||||
|
||||
@override
|
||||
|
||||
@@ -22,9 +22,28 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class CacheEmbedding(Embeddings):
|
||||
QUERY_CACHE_TTL = 600
|
||||
|
||||
def __init__(self, model_instance: ModelInstance):
|
||||
self._model_instance = model_instance
|
||||
|
||||
@staticmethod
|
||||
def _query_cache_key(provider: str, model_name: str, query_hash: str) -> str:
|
||||
return f"{provider}_{model_name}_{query_hash}"
|
||||
|
||||
@classmethod
|
||||
def get_cached_query_embedding(cls, provider: str, model_name: str, text: str) -> list[float] | None:
|
||||
"""Return a cached query vector without requiring a model instance."""
|
||||
query_hash = helper.generate_text_hash(text)
|
||||
embedding_cache_key = cls._query_cache_key(provider, model_name, query_hash)
|
||||
embedding = redis_client.get(embedding_cache_key)
|
||||
if not embedding:
|
||||
return None
|
||||
|
||||
redis_client.expire(embedding_cache_key, cls.QUERY_CACHE_TTL)
|
||||
decoded_embedding = np.frombuffer(base64.b64decode(embedding), dtype="float")
|
||||
return [float(x) for x in decoded_embedding]
|
||||
|
||||
@override
|
||||
def embed_documents(self, texts: list[str]) -> list[list[float]]:
|
||||
"""Embed search docs in batches of 10."""
|
||||
@@ -196,12 +215,14 @@ class CacheEmbedding(Embeddings):
|
||||
"""Embed query text."""
|
||||
# use doc embedding cache or store if not exists
|
||||
hash = helper.generate_text_hash(text)
|
||||
embedding_cache_key = f"{self._model_instance.provider}_{self._model_instance.model_name}_{hash}"
|
||||
embedding = redis_client.get(embedding_cache_key)
|
||||
if embedding:
|
||||
redis_client.expire(embedding_cache_key, 600)
|
||||
decoded_embedding = np.frombuffer(base64.b64decode(embedding), dtype="float")
|
||||
return [float(x) for x in decoded_embedding]
|
||||
embedding_cache_key = self._query_cache_key(
|
||||
self._model_instance.provider, self._model_instance.model_name, hash
|
||||
)
|
||||
cached_embedding = self.get_cached_query_embedding(
|
||||
self._model_instance.provider, self._model_instance.model_name, text
|
||||
)
|
||||
if cached_embedding is not None:
|
||||
return cached_embedding
|
||||
try:
|
||||
embedding_result = self._model_instance.invoke_text_embedding(
|
||||
texts=[text], input_type=EmbeddingInputType.QUERY
|
||||
@@ -225,7 +246,7 @@ class CacheEmbedding(Embeddings):
|
||||
encoded_vector = base64.b64encode(vector_bytes)
|
||||
# Transform to string
|
||||
encoded_str = encoded_vector.decode("utf-8")
|
||||
redis_client.setex(embedding_cache_key, 600, encoded_str)
|
||||
redis_client.setex(embedding_cache_key, self.QUERY_CACHE_TTL, encoded_str)
|
||||
except Exception as ex:
|
||||
if dify_config.DEBUG:
|
||||
logger.exception(
|
||||
|
||||
@@ -5,14 +5,14 @@ from sqlalchemy.orm import Session
|
||||
|
||||
from core.credit_usage import CreditUsageCreatedBy
|
||||
from core.model_context import with_credit_usage_created_by
|
||||
from core.model_manager import ModelInstance, ModelManager
|
||||
from core.model_manager import ModelInstance
|
||||
from core.rag.index_processor.constant.doc_type import DocType
|
||||
from core.rag.index_processor.constant.query_type import QueryType
|
||||
from core.rag.models.document import Document
|
||||
from core.rag.rerank.rerank_base import BaseRerankRunner
|
||||
from extensions.ext_storage import storage
|
||||
from extensions.otel import trace_span
|
||||
from graphon.model_runtime.entities.model_entities import ModelType
|
||||
from graphon.model_runtime.entities.model_entities import ModelFeature
|
||||
from graphon.model_runtime.entities.rerank_entities import MultimodalRerankInput, RerankResult
|
||||
from models.model import UploadFile
|
||||
|
||||
@@ -43,15 +43,7 @@ class RerankModelRunner(BaseRerankRunner):
|
||||
:param top_n: top n
|
||||
:return:
|
||||
"""
|
||||
model_manager = ModelManager.for_tenant(
|
||||
tenant_id=self.rerank_model_instance.provider_model_bundle.configuration.tenant_id
|
||||
)
|
||||
is_support_vision = model_manager.check_model_support_vision(
|
||||
tenant_id=self.rerank_model_instance.provider_model_bundle.configuration.tenant_id,
|
||||
provider=self.rerank_model_instance.provider,
|
||||
model=self.rerank_model_instance.model_name,
|
||||
model_type=ModelType.RERANK,
|
||||
)
|
||||
is_support_vision = self._check_model_support_vision()
|
||||
if not is_support_vision:
|
||||
if query_type == QueryType.TEXT_QUERY:
|
||||
rerank_result, unique_documents = self.fetch_text_rerank(query, documents, score_threshold, top_n)
|
||||
@@ -78,6 +70,11 @@ class RerankModelRunner(BaseRerankRunner):
|
||||
rerank_documents.sort(key=lambda x: x.metadata.get("score", 0.0), reverse=True)
|
||||
return rerank_documents[:top_n] if top_n else rerank_documents
|
||||
|
||||
def _check_model_support_vision(self) -> bool:
|
||||
"""Check capabilities on the model instance already resolved for this run."""
|
||||
model_schema = self.rerank_model_instance.get_model_schema()
|
||||
return ModelFeature.VISION in (model_schema.features or [])
|
||||
|
||||
def fetch_text_rerank(
|
||||
self,
|
||||
query: str,
|
||||
|
||||
@@ -241,6 +241,23 @@ def test_lazy_embeddings_defer_real_load_until_first_embed_call(vector_factory_m
|
||||
inner_model.embed_documents.assert_called_once_with(["world"])
|
||||
|
||||
|
||||
def test_lazy_embeddings_query_cache_hit_skips_model_resolution(vector_factory_module, monkeypatch: pytest.MonkeyPatch):
|
||||
"""A cached query vector should not construct ModelManager or a model instance."""
|
||||
proxy = vector_factory_module._LazyEmbeddings(_dataset())
|
||||
cached_vector = [0.1, 0.2]
|
||||
cache_lookup = MagicMock(return_value=cached_vector)
|
||||
for_tenant = MagicMock(side_effect=AssertionError("model resolution must be skipped on cache hit"))
|
||||
monkeypatch.setattr(vector_factory_module.CacheEmbedding, "get_cached_query_embedding", cache_lookup)
|
||||
monkeypatch.setattr(vector_factory_module.ModelManager, "for_tenant", for_tenant)
|
||||
|
||||
result = proxy.embed_query("hello")
|
||||
|
||||
assert result == cached_vector
|
||||
cache_lookup.assert_called_once_with("openai", "text-embedding-3-small", "hello")
|
||||
for_tenant.assert_not_called()
|
||||
assert proxy._real is None
|
||||
|
||||
|
||||
def test_init_vector_prefers_dataset_index_struct(
|
||||
vector_factory_module, monkeypatch: pytest.MonkeyPatch, unbound_session: Session
|
||||
):
|
||||
|
||||
@@ -253,6 +253,17 @@ class TestMultimodalDocumentCache:
|
||||
|
||||
|
||||
class TestQueryCache:
|
||||
@patch("core.rag.embedding.cached_embedding.redis_client")
|
||||
def test_cached_query_can_be_read_without_model_instance(self, redis: Mock) -> None:
|
||||
vector = np.array([0.25, 0.75], dtype=float)
|
||||
redis.get.return_value = base64.b64encode(vector.tobytes())
|
||||
|
||||
result = CacheEmbedding.get_cached_query_embedding("openai", "text-embedding-ada-002", "query")
|
||||
|
||||
assert result == vector.tolist()
|
||||
redis.get.assert_called_once_with(f"openai_text-embedding-ada-002_{helper.generate_text_hash('query')}")
|
||||
redis.expire.assert_called_once_with(f"openai_text-embedding-ada-002_{helper.generate_text_hash('query')}", 600)
|
||||
|
||||
@patch("core.rag.embedding.cached_embedding.redis_client")
|
||||
def test_query_cache_miss_normalizes_and_stores(self, redis: Mock, model_instance: Mock) -> None:
|
||||
redis.get.return_value = None
|
||||
|
||||
@@ -30,6 +30,7 @@ from core.rag.rerank.rerank_model import RerankModelRunner
|
||||
from core.rag.rerank.rerank_type import RerankMode
|
||||
from core.rag.rerank.weight_rerank import WeightRerankRunner
|
||||
from extensions.storage.storage_type import StorageType
|
||||
from graphon.model_runtime.entities.model_entities import ModelFeature
|
||||
from graphon.model_runtime.entities.rerank_entities import RerankDocument, RerankResult
|
||||
from models.enums import CreatorUserRole
|
||||
from models.model import UploadFile
|
||||
@@ -44,6 +45,7 @@ def create_mock_model_instance() -> ModelInstance:
|
||||
mock_instance.provider_model_bundle.configuration.tenant_id = "test-tenant-id"
|
||||
mock_instance.provider = "test-provider"
|
||||
mock_instance.model_name = "test-model"
|
||||
mock_instance.get_model_schema.return_value = Mock(features=[])
|
||||
return mock_instance
|
||||
|
||||
|
||||
@@ -66,13 +68,6 @@ class TestRerankModelRunner(_UsesSQLiteSession):
|
||||
- Metadata preservation and score injection
|
||||
"""
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def mock_model_manager(self):
|
||||
"""Auto-use fixture to patch ModelManager for all tests in this class."""
|
||||
with patch("core.rag.rerank.rerank_model.ModelManager.for_tenant", autospec=True) as mock_mm:
|
||||
mock_mm.return_value.check_model_support_vision.return_value = False
|
||||
yield mock_mm
|
||||
|
||||
@pytest.fixture
|
||||
def mock_model_instance(self):
|
||||
"""Create a mock ModelInstance for reranking."""
|
||||
@@ -374,14 +369,12 @@ class TestRerankModelRunner(_UsesSQLiteSession):
|
||||
# Assert: Empty result is returned
|
||||
assert len(result) == 0
|
||||
|
||||
def test_run_uses_bound_model_instance(
|
||||
self, rerank_runner, mock_model_instance, sample_documents, mock_model_manager
|
||||
):
|
||||
def test_run_uses_bound_model_instance(self, rerank_runner, mock_model_instance, sample_documents):
|
||||
"""Test that rerank uses the bound model instance directly.
|
||||
|
||||
Verifies:
|
||||
- The injected model instance is used for invocation
|
||||
- No late rebinding occurs through ModelManager.get_model_instance
|
||||
- Capability detection uses the already-bound model instance
|
||||
"""
|
||||
# Arrange: Mock rerank result
|
||||
mock_rerank_result = RerankResult(
|
||||
@@ -400,7 +393,7 @@ class TestRerankModelRunner(_UsesSQLiteSession):
|
||||
|
||||
# Assert: The injected model instance is invoked directly.
|
||||
assert len(result) == 1
|
||||
mock_model_manager.return_value.get_model_instance.assert_not_called()
|
||||
mock_model_instance.get_model_schema.assert_called_once_with()
|
||||
call_kwargs = mock_model_instance.invoke_rerank.call_args.kwargs
|
||||
assert call_kwargs["query"] == "test"
|
||||
assert "user" not in call_kwargs
|
||||
@@ -448,9 +441,7 @@ class TestRerankModelRunnerMultimodal(_UsesSQLiteSession):
|
||||
Document(page_content="doc", metadata={"doc_id": "doc1"}, provider="dify"),
|
||||
]
|
||||
|
||||
with patch("core.rag.rerank.rerank_model.ModelManager.for_tenant") as mock_mm:
|
||||
mock_mm.return_value.check_model_support_vision.return_value = False
|
||||
result = rerank_runner.run(query="image-file-id", documents=documents, query_type=QueryType.IMAGE_QUERY)
|
||||
result = rerank_runner.run(query="image-file-id", documents=documents, query_type=QueryType.IMAGE_QUERY)
|
||||
|
||||
assert result == documents
|
||||
mock_model_instance.invoke_rerank.assert_not_called()
|
||||
@@ -464,15 +455,12 @@ class TestRerankModelRunnerMultimodal(_UsesSQLiteSession):
|
||||
docs=[RerankDocument(index=0, text="doc", score=0.88)],
|
||||
)
|
||||
|
||||
with (
|
||||
patch("core.rag.rerank.rerank_model.ModelManager.for_tenant") as mock_mm,
|
||||
patch.object(
|
||||
rerank_runner,
|
||||
"fetch_multimodal_rerank",
|
||||
return_value=(rerank_result, documents),
|
||||
) as mock_multimodal,
|
||||
):
|
||||
mock_mm.return_value.check_model_support_vision.return_value = True
|
||||
rerank_runner.rerank_model_instance.get_model_schema.return_value = Mock(features=[ModelFeature.VISION])
|
||||
with patch.object(
|
||||
rerank_runner,
|
||||
"fetch_multimodal_rerank",
|
||||
return_value=(rerank_result, documents),
|
||||
) as mock_multimodal:
|
||||
result = rerank_runner.run(query="python", documents=documents, query_type=QueryType.TEXT_QUERY)
|
||||
|
||||
mock_multimodal.assert_called_once()
|
||||
@@ -1189,13 +1177,6 @@ class TestRerankIntegration(_UsesSQLiteSession):
|
||||
- Real-world usage scenarios
|
||||
"""
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def mock_model_manager(self):
|
||||
"""Auto-use fixture to patch ModelManager for all tests in this class."""
|
||||
with patch("core.rag.rerank.rerank_model.ModelManager.for_tenant", autospec=True) as mock_mm:
|
||||
mock_mm.return_value.check_model_support_vision.return_value = False
|
||||
yield mock_mm
|
||||
|
||||
def test_model_reranking_full_workflow(self):
|
||||
"""Test complete model-based reranking workflow.
|
||||
|
||||
@@ -1302,13 +1283,6 @@ class TestRerankEdgeCases(_UsesSQLiteSession):
|
||||
- Concurrent reranking scenarios
|
||||
"""
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def mock_model_manager(self):
|
||||
"""Auto-use fixture to patch ModelManager for all tests in this class."""
|
||||
with patch("core.rag.rerank.rerank_model.ModelManager.for_tenant", autospec=True) as mock_mm:
|
||||
mock_mm.return_value.check_model_support_vision.return_value = False
|
||||
yield mock_mm
|
||||
|
||||
def test_rerank_with_empty_metadata(self):
|
||||
"""Test reranking when documents have empty metadata.
|
||||
|
||||
@@ -1643,13 +1617,6 @@ class TestRerankPerformance(_UsesSQLiteSession):
|
||||
- Score calculation optimization
|
||||
"""
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def mock_model_manager(self):
|
||||
"""Auto-use fixture to patch ModelManager for all tests in this class."""
|
||||
with patch("core.rag.rerank.rerank_model.ModelManager.for_tenant", autospec=True) as mock_mm:
|
||||
mock_mm.return_value.check_model_support_vision.return_value = False
|
||||
yield mock_mm
|
||||
|
||||
def test_rerank_batch_processing(self):
|
||||
"""Test that documents are processed in a single batch.
|
||||
|
||||
@@ -1760,13 +1727,6 @@ class TestRerankErrorHandling(_UsesSQLiteSession):
|
||||
- Error propagation
|
||||
"""
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def mock_model_manager(self):
|
||||
"""Auto-use fixture to patch ModelManager for all tests in this class."""
|
||||
with patch("core.rag.rerank.rerank_model.ModelManager.for_tenant", autospec=True) as mock_mm:
|
||||
mock_mm.return_value.check_model_support_vision.return_value = False
|
||||
yield mock_mm
|
||||
|
||||
def test_rerank_model_invocation_error(self):
|
||||
"""Test handling of model invocation errors.
|
||||
|
||||
|
||||
@@ -254,6 +254,8 @@ class TestProviderConfiguration:
|
||||
|
||||
# Assert
|
||||
assert ConfigurateMethod.PREDEFINED_MODEL in config.provider.configurate_methods
|
||||
assert config.provider is not mock_provider_entity
|
||||
assert mock_provider_entity.configurate_methods == [ConfigurateMethod.CUSTOMIZABLE_MODEL]
|
||||
|
||||
def test_get_current_credentials_with_restricted_models(self, provider_configuration):
|
||||
"""Test getting credentials with model restrictions"""
|
||||
|
||||
@@ -40,6 +40,18 @@ def _plugin_config(config_overrides) -> None:
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _clear_parsed_provider_cache():
|
||||
"""Keep the process-local provider cache isolated between tests."""
|
||||
from core.plugin.plugin_service import PluginService
|
||||
|
||||
with PluginService._parsed_plugin_model_providers_cache_lock:
|
||||
PluginService._parsed_plugin_model_providers_cache.clear()
|
||||
yield
|
||||
with PluginService._parsed_plugin_model_providers_cache_lock:
|
||||
PluginService._parsed_plugin_model_providers_cache.clear()
|
||||
|
||||
|
||||
def _build_provider_entity(
|
||||
provider: str = "openai",
|
||||
installation_source: PluginInstallationSource | None = PluginInstallationSource.Marketplace,
|
||||
@@ -157,6 +169,57 @@ class TestFetchLatestPluginVersion:
|
||||
|
||||
|
||||
class TestPluginModelProviderCache:
|
||||
def test_reuses_parsed_provider_payload_without_revalidating_json(self) -> None:
|
||||
"""An unchanged Redis payload should pay Pydantic validation only once per process."""
|
||||
from core.plugin.plugin_service import PluginService, _provider_entities_adapter
|
||||
|
||||
payload = TypeAdapter(list[PluginModelProviderDeclaration]).dump_json([_build_provider_entity()])
|
||||
with patch.object(
|
||||
_provider_entities_adapter,
|
||||
"validate_json",
|
||||
wraps=_provider_entities_adapter.validate_json,
|
||||
) as validate_json:
|
||||
first = PluginService._get_or_parse_plugin_model_providers_cache_payload(payload)
|
||||
second = PluginService._get_or_parse_plugin_model_providers_cache_payload(payload)
|
||||
|
||||
assert second is first
|
||||
assert validate_json.call_count == 1
|
||||
|
||||
def test_changed_provider_payload_is_parsed_separately(self) -> None:
|
||||
"""Content-addressing must not return stale declarations after a Redis payload change."""
|
||||
from core.plugin.plugin_service import PluginService, _provider_entities_adapter
|
||||
|
||||
first_payload = TypeAdapter(list[PluginModelProviderDeclaration]).dump_json([_build_provider_entity()])
|
||||
second_payload = TypeAdapter(list[PluginModelProviderDeclaration]).dump_json(
|
||||
[_build_provider_entity(provider="anthropic")]
|
||||
)
|
||||
with patch.object(
|
||||
_provider_entities_adapter,
|
||||
"validate_json",
|
||||
wraps=_provider_entities_adapter.validate_json,
|
||||
) as validate_json:
|
||||
first = PluginService._get_or_parse_plugin_model_providers_cache_payload(first_payload)
|
||||
second = PluginService._get_or_parse_plugin_model_providers_cache_payload(second_payload)
|
||||
|
||||
assert first[0].provider == "langgenius/openai/openai"
|
||||
assert second[0].provider == "langgenius/anthropic/anthropic"
|
||||
assert validate_json.call_count == 2
|
||||
|
||||
def test_parsed_provider_cache_is_bounded(self) -> None:
|
||||
"""Distinct provider payloads should evict least-recently-used entries."""
|
||||
from core.plugin.plugin_service import PluginService
|
||||
|
||||
for index in range(PluginService.PLUGIN_MODEL_PROVIDERS_PARSED_CACHE_MAX_ENTRIES + 1):
|
||||
payload = TypeAdapter(list[PluginModelProviderDeclaration]).dump_json(
|
||||
[_build_provider_entity(provider=f"provider-{index}")]
|
||||
)
|
||||
PluginService._get_or_parse_plugin_model_providers_cache_payload(payload)
|
||||
|
||||
assert (
|
||||
len(PluginService._parsed_plugin_model_providers_cache)
|
||||
== PluginService.PLUGIN_MODEL_PROVIDERS_PARSED_CACHE_MAX_ENTRIES
|
||||
)
|
||||
|
||||
def test_store_cached_plugin_model_providers_compresses_large_payload(self) -> None:
|
||||
"""Large provider metadata payloads are compressed before being stored in Redis."""
|
||||
large_provider = _build_provider_entity()
|
||||
|
||||
Reference in New Issue
Block a user