perf(api): reduce knowledge retrieval latency (#42673)

This commit is contained in:
呆萌闷油瓶
2026-09-24 06:50:36 +00:00
committed by GitHub
parent 989ba1a78f
commit 74af4f2308
10 changed files with 217 additions and 73 deletions
+8 -1
View File
@@ -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:
+62 -2
View File
@@ -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
+28 -7
View File
@@ -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(
+8 -11
View File
@@ -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()