mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-09 22:31:41 +00:00
Merge pull request #37742 from BerriAI/litellm_lit5879_semantic_cache_embedding_timeout
fix(caching): bound the semantic cache embedding lookup so a dead embedding endpoint can't block requests
This commit is contained in:
commit
4e02e7e404
9 changed files with 296 additions and 24 deletions
|
|
@ -16,6 +16,7 @@ from collections.abc import Sequence
|
|||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
import litellm
|
||||
from litellm.constants import SEMANTIC_CACHE_EMBEDDING_TIMEOUT_SECONDS
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.router import Router
|
||||
|
|
@ -60,6 +61,13 @@ def resolve_embedding_max_input_tokens(
|
|||
return deployment_max_input_tokens
|
||||
|
||||
|
||||
def resolve_embedding_timeout(configured_timeout: float | None) -> float:
|
||||
"""Explicit cache setting first, else the short semantic-cache default."""
|
||||
if configured_timeout is not None:
|
||||
return configured_timeout
|
||||
return SEMANTIC_CACHE_EMBEDDING_TIMEOUT_SECONDS
|
||||
|
||||
|
||||
def truncate_embedding_input(prompt: str, embedding_model: str, max_input_tokens: int | None) -> str:
|
||||
"""Keep only the first ``max_input_tokens`` tokens of ``prompt`` for the embedding call."""
|
||||
if max_input_tokens is None:
|
||||
|
|
|
|||
|
|
@ -98,6 +98,7 @@ class Cache:
|
|||
qdrant_semantic_cache_embedding_model: str = "text-embedding-ada-002",
|
||||
qdrant_semantic_cache_vector_size: int | None = None,
|
||||
semantic_cache_embedding_max_input_tokens: int | None = None,
|
||||
semantic_cache_embedding_timeout: float | None = None,
|
||||
# GCP IAM authentication parameters
|
||||
gcp_service_account: str | None = None,
|
||||
gcp_ssl_ca_certs: str | None = None,
|
||||
|
|
@ -124,6 +125,7 @@ class Cache:
|
|||
qdrant_collection_name (str, optional): The name for your qdrant collection. Required if type is "qdrant-semantic".
|
||||
similarity_threshold (float, optional): The similarity threshold for semantic-caching, Required if type is "redis-semantic" or "qdrant-semantic".
|
||||
semantic_cache_embedding_max_input_tokens (int, optional): Truncate prompts to this many tokens before embedding them for semantic caching. Defaults to the embedding deployment's configured max_input_tokens.
|
||||
semantic_cache_embedding_timeout (float, optional): Seconds a semantic-cache lookup may spend embedding the prompt before it gives up and lets the request continue to the LLM. Defaults to SEMANTIC_CACHE_EMBEDDING_TIMEOUT_SECONDS.
|
||||
|
||||
# Disk Cache Args
|
||||
disk_cache_dir (str, optional): The directory for the disk cache. Defaults to None.
|
||||
|
|
@ -195,6 +197,7 @@ class Cache:
|
|||
embedding_model=redis_semantic_cache_embedding_model,
|
||||
index_name=redis_semantic_cache_index_name,
|
||||
embedding_max_input_tokens=semantic_cache_embedding_max_input_tokens,
|
||||
embedding_timeout=semantic_cache_embedding_timeout,
|
||||
**kwargs,
|
||||
)
|
||||
elif type == LiteLLMCacheType.VALKEY_SEMANTIC:
|
||||
|
|
@ -211,6 +214,7 @@ class Cache:
|
|||
index_name=valkey_semantic_cache_index_name,
|
||||
startup_nodes=redis_startup_nodes,
|
||||
embedding_max_input_tokens=semantic_cache_embedding_max_input_tokens,
|
||||
embedding_timeout=semantic_cache_embedding_timeout,
|
||||
**kwargs,
|
||||
)
|
||||
elif type == LiteLLMCacheType.QDRANT_SEMANTIC:
|
||||
|
|
@ -223,6 +227,7 @@ class Cache:
|
|||
embedding_model=qdrant_semantic_cache_embedding_model,
|
||||
vector_size=qdrant_semantic_cache_vector_size,
|
||||
embedding_max_input_tokens=semantic_cache_embedding_max_input_tokens,
|
||||
embedding_timeout=semantic_cache_embedding_timeout,
|
||||
)
|
||||
elif type == LiteLLMCacheType.LOCAL:
|
||||
self.cache = InMemoryCache()
|
||||
|
|
|
|||
|
|
@ -16,7 +16,11 @@ from typing import TYPE_CHECKING, Any, Final, cast
|
|||
|
||||
import litellm
|
||||
from litellm._logging import print_verbose
|
||||
from litellm.constants import QDRANT_SCALAR_QUANTILE, QDRANT_VECTOR_SIZE
|
||||
from litellm.constants import (
|
||||
QDRANT_SCALAR_QUANTILE,
|
||||
QDRANT_VECTOR_SIZE,
|
||||
SEMANTIC_CACHE_EMBEDDING_TIMEOUT_SECONDS,
|
||||
)
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
get_str_from_messages,
|
||||
)
|
||||
|
|
@ -26,6 +30,7 @@ from ._embedding_router import (
|
|||
build_router_embedding_metadata,
|
||||
resolve_embedding_max_input_tokens,
|
||||
resolve_embedding_router,
|
||||
resolve_embedding_timeout,
|
||||
truncate_embedding_input,
|
||||
)
|
||||
from .base_cache import BaseCache
|
||||
|
|
@ -37,6 +42,7 @@ if TYPE_CHECKING:
|
|||
class QdrantSemanticCache(BaseCache):
|
||||
CACHE_KEY_FIELD_NAME = "litellm_cache_key"
|
||||
embedding_max_input_tokens: int | None = None
|
||||
embedding_timeout: float = SEMANTIC_CACHE_EMBEDDING_TIMEOUT_SECONDS
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
|
|
@ -49,6 +55,7 @@ class QdrantSemanticCache(BaseCache):
|
|||
host_type=None,
|
||||
vector_size=None,
|
||||
embedding_max_input_tokens: int | None = None,
|
||||
embedding_timeout: float | None = None,
|
||||
):
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
_get_httpx_client,
|
||||
|
|
@ -68,6 +75,7 @@ class QdrantSemanticCache(BaseCache):
|
|||
self.similarity_threshold = similarity_threshold
|
||||
self.embedding_model = embedding_model
|
||||
self.embedding_max_input_tokens = embedding_max_input_tokens
|
||||
self.embedding_timeout = resolve_embedding_timeout(embedding_timeout)
|
||||
self.vector_size = vector_size if vector_size is not None else QDRANT_VECTOR_SIZE
|
||||
headers = {}
|
||||
|
||||
|
|
@ -222,11 +230,15 @@ class QdrantSemanticCache(BaseCache):
|
|||
input=embedding_input,
|
||||
cache={"no-store": True, "no-cache": True},
|
||||
metadata=build_router_embedding_metadata(metadata),
|
||||
timeout=self.embedding_timeout,
|
||||
num_retries=0,
|
||||
)
|
||||
return litellm.embedding(
|
||||
model=self.embedding_model,
|
||||
input=embedding_input,
|
||||
cache={"no-store": True, "no-cache": True},
|
||||
timeout=self.embedding_timeout,
|
||||
num_retries=0,
|
||||
)
|
||||
|
||||
async def _get_async_embedding(self, prompt: str, metadata: dict[str, Any] | None = None) -> EmbeddingResponse:
|
||||
|
|
@ -238,19 +250,25 @@ class QdrantSemanticCache(BaseCache):
|
|||
|
||||
router: Final = resolve_embedding_router(self.embedding_model, llm_router, llm_model_list)
|
||||
embedding_input: Final = self._embedding_input(prompt, router)
|
||||
if router is not None:
|
||||
return await router.aembedding(
|
||||
embedding_call: Final = (
|
||||
router.aembedding(
|
||||
model=self.embedding_model,
|
||||
input=embedding_input,
|
||||
cache={"no-store": True, "no-cache": True},
|
||||
metadata=build_router_embedding_metadata(metadata),
|
||||
timeout=self.embedding_timeout,
|
||||
num_retries=0,
|
||||
)
|
||||
if router is not None
|
||||
else litellm.aembedding(
|
||||
model=self.embedding_model,
|
||||
input=embedding_input,
|
||||
cache={"no-store": True, "no-cache": True},
|
||||
timeout=self.embedding_timeout,
|
||||
num_retries=0,
|
||||
)
|
||||
|
||||
return await litellm.aembedding(
|
||||
model=self.embedding_model,
|
||||
input=embedding_input,
|
||||
cache={"no-store": True, "no-cache": True},
|
||||
)
|
||||
return await asyncio.wait_for(embedding_call, self.embedding_timeout)
|
||||
|
||||
def set_cache(self, key, value, **kwargs):
|
||||
print_verbose(f"qdrant semantic-cache set_cache, kwargs: {kwargs}")
|
||||
|
|
|
|||
|
|
@ -18,6 +18,7 @@ from typing import TYPE_CHECKING, Any, Final, cast
|
|||
|
||||
import litellm
|
||||
from litellm._logging import print_verbose, verbose_logger
|
||||
from litellm.constants import SEMANTIC_CACHE_EMBEDDING_TIMEOUT_SECONDS
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
get_str_from_messages,
|
||||
)
|
||||
|
|
@ -27,6 +28,7 @@ from ._embedding_router import (
|
|||
build_router_embedding_metadata,
|
||||
resolve_embedding_max_input_tokens,
|
||||
resolve_embedding_router,
|
||||
resolve_embedding_timeout,
|
||||
truncate_embedding_input,
|
||||
)
|
||||
from .base_cache import BaseCache
|
||||
|
|
@ -47,6 +49,7 @@ class RedisSemanticCache(BaseCache):
|
|||
DEFAULT_REDIS_INDEX_NAME: str = "litellm_semantic_cache_index"
|
||||
CACHE_KEY_FIELD_NAME: str = "litellm_cache_key"
|
||||
embedding_max_input_tokens: int | None = None
|
||||
embedding_timeout: float = SEMANTIC_CACHE_EMBEDDING_TIMEOUT_SECONDS
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
|
|
@ -58,6 +61,7 @@ class RedisSemanticCache(BaseCache):
|
|||
embedding_model: str = "text-embedding-ada-002",
|
||||
index_name: str | None = None,
|
||||
embedding_max_input_tokens: int | None = None,
|
||||
embedding_timeout: float | None = None,
|
||||
**kwargs: object,
|
||||
):
|
||||
"""
|
||||
|
|
@ -74,6 +78,8 @@ class RedisSemanticCache(BaseCache):
|
|||
index_name: Name for the Redis index
|
||||
embedding_max_input_tokens: Truncate prompts to this many tokens before
|
||||
embedding; defaults to the Router deployment's configured max_input_tokens
|
||||
embedding_timeout: Seconds a cache lookup may spend embedding the prompt before it
|
||||
gives up and lets the request continue to the LLM
|
||||
ttl: Default time-to-live for cache entries in seconds
|
||||
**kwargs: Additional arguments passed to the Redis client
|
||||
|
||||
|
|
@ -99,6 +105,7 @@ class RedisSemanticCache(BaseCache):
|
|||
self.distance_threshold = 1 - similarity_threshold
|
||||
self.embedding_model = embedding_model
|
||||
self.embedding_max_input_tokens = embedding_max_input_tokens
|
||||
self.embedding_timeout = resolve_embedding_timeout(embedding_timeout)
|
||||
|
||||
# Set up Redis connection
|
||||
if redis_url is None:
|
||||
|
|
@ -349,6 +356,8 @@ class RedisSemanticCache(BaseCache):
|
|||
input=embedding_input,
|
||||
cache={"no-store": True, "no-cache": True},
|
||||
metadata=build_router_embedding_metadata(metadata),
|
||||
timeout=self.embedding_timeout,
|
||||
num_retries=0,
|
||||
),
|
||||
)
|
||||
else:
|
||||
|
|
@ -358,6 +367,8 @@ class RedisSemanticCache(BaseCache):
|
|||
model=self.embedding_model,
|
||||
input=embedding_input,
|
||||
cache={"no-store": True, "no-cache": True},
|
||||
timeout=self.embedding_timeout,
|
||||
num_retries=0,
|
||||
),
|
||||
)
|
||||
return embedding_response["data"][0]["embedding"]
|
||||
|
|
@ -512,20 +523,26 @@ class RedisSemanticCache(BaseCache):
|
|||
|
||||
router: Final = resolve_embedding_router(self.embedding_model, llm_router, llm_model_list)
|
||||
embedding_input: Final = self._embedding_input(prompt, router)
|
||||
embedding_call: Final = (
|
||||
router.aembedding(
|
||||
model=self.embedding_model,
|
||||
input=embedding_input,
|
||||
cache={"no-store": True, "no-cache": True},
|
||||
metadata=build_router_embedding_metadata(metadata),
|
||||
timeout=self.embedding_timeout,
|
||||
num_retries=0,
|
||||
)
|
||||
if router is not None
|
||||
else litellm.aembedding(
|
||||
model=self.embedding_model,
|
||||
input=embedding_input,
|
||||
cache={"no-store": True, "no-cache": True},
|
||||
timeout=self.embedding_timeout,
|
||||
num_retries=0,
|
||||
)
|
||||
)
|
||||
try:
|
||||
if router is not None:
|
||||
embedding_response = await router.aembedding(
|
||||
model=self.embedding_model,
|
||||
input=embedding_input,
|
||||
cache={"no-store": True, "no-cache": True},
|
||||
metadata=build_router_embedding_metadata(metadata),
|
||||
)
|
||||
else:
|
||||
embedding_response = await litellm.aembedding(
|
||||
model=self.embedding_model,
|
||||
input=embedding_input,
|
||||
cache={"no-store": True, "no-cache": True},
|
||||
)
|
||||
embedding_response: Final = await asyncio.wait_for(embedding_call, self.embedding_timeout)
|
||||
return embedding_response["data"][0]["embedding"]
|
||||
except Exception as e:
|
||||
print_verbose(f"Error generating async embedding: {e}")
|
||||
|
|
|
|||
|
|
@ -30,6 +30,7 @@ from litellm._logging import print_verbose
|
|||
from litellm._uuid import uuid
|
||||
from litellm.llms.valkey.common_utils import build_valkey_url, pack_vector
|
||||
|
||||
from ._embedding_router import resolve_embedding_timeout
|
||||
from .redis_semantic_cache import RedisSemanticCache
|
||||
|
||||
|
||||
|
|
@ -62,6 +63,7 @@ class ValkeySemanticCache(RedisSemanticCache):
|
|||
sync_client: Redis | None = None,
|
||||
async_client: AsyncRedis | None = None,
|
||||
embedding_max_input_tokens: int | None = None,
|
||||
embedding_timeout: float | None = None,
|
||||
**kwargs: Any,
|
||||
):
|
||||
if similarity_threshold is None:
|
||||
|
|
@ -80,6 +82,7 @@ class ValkeySemanticCache(RedisSemanticCache):
|
|||
self.similarity_threshold = similarity_threshold
|
||||
self.embedding_model = embedding_model
|
||||
self.embedding_max_input_tokens = embedding_max_input_tokens
|
||||
self.embedding_timeout = resolve_embedding_timeout(embedding_timeout)
|
||||
self.index_name = index_name or self.DEFAULT_VALKEY_INDEX_NAME
|
||||
self.key_prefix = f"{self.index_name}:"
|
||||
self._index_dim: int | None = None
|
||||
|
|
|
|||
|
|
@ -436,6 +436,9 @@ DEFAULT_REQUEST_TIMEOUT_SECONDS: Final[float] = 6000.0
|
|||
# deadline and connect handshake (see ``http_handler`` cached handler paths).
|
||||
COMPLETION_HTTP_FALLBACK_SECONDS: Final[float] = 600.0
|
||||
HTTP_HANDLER_CONNECT_TIMEOUT_SECONDS: Final[float] = 5.0
|
||||
SEMANTIC_CACHE_EMBEDDING_TIMEOUT_SECONDS: Final[float] = float(
|
||||
os.getenv("SEMANTIC_CACHE_EMBEDDING_TIMEOUT_SECONDS", "5.0")
|
||||
)
|
||||
request_timeout: float = float(os.getenv("REQUEST_TIMEOUT", str(int(DEFAULT_REQUEST_TIMEOUT_SECONDS))))
|
||||
request_timeout_explicitly_set: bool = "REQUEST_TIMEOUT" in os.environ
|
||||
DEFAULT_A2A_AGENT_TIMEOUT: Final[float] = float(os.getenv("DEFAULT_A2A_AGENT_TIMEOUT", 6000)) # 10 minutes
|
||||
|
|
|
|||
|
|
@ -5974,7 +5974,7 @@ def embedding(
|
|||
# Optional params
|
||||
dimensions: int | None = None,
|
||||
encoding_format: str | None = None,
|
||||
timeout=600, # default to 10 minutes
|
||||
timeout: float = 600, # default to 10 minutes
|
||||
# set api_base, api_version, api_key
|
||||
api_base: str | None = None,
|
||||
api_version: str | None = None,
|
||||
|
|
@ -6000,7 +6000,7 @@ def embedding(
|
|||
# Optional params
|
||||
dimensions: int | None = None,
|
||||
encoding_format: str | None = None,
|
||||
timeout=600, # default to 10 minutes
|
||||
timeout: float = 600, # default to 10 minutes
|
||||
# set api_base, api_version, api_key
|
||||
api_base: str | None = None,
|
||||
api_version: str | None = None,
|
||||
|
|
@ -6027,7 +6027,7 @@ def embedding(
|
|||
# Optional params
|
||||
dimensions: int | None = None,
|
||||
encoding_format: str | None = None,
|
||||
timeout=600, # default to 10 minutes
|
||||
timeout: float = 600, # default to 10 minutes
|
||||
# set api_base, api_version, api_key
|
||||
api_base: str | None = None,
|
||||
api_version: str | None = None,
|
||||
|
|
|
|||
|
|
@ -966,3 +966,67 @@ async def test_qdrant_async_embedding_explicit_limit_beats_deployment_limit(monk
|
|||
|
||||
sent_input = router.aembedding.call_args.kwargs["input"]
|
||||
assert _token_count("sem-embed", sent_input) == 3
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_qdrant_async_embedding_call_is_bounded(monkeypatch):
|
||||
from litellm.caching.qdrant_semantic_cache import QdrantSemanticCache
|
||||
|
||||
cache = QdrantSemanticCache.__new__(QdrantSemanticCache)
|
||||
cache.embedding_model = "sem-embed"
|
||||
cache.embedding_max_input_tokens = None
|
||||
cache.embedding_timeout = 1.5
|
||||
|
||||
router = MagicMock()
|
||||
router.get_configured_token_limits.return_value = (None, None)
|
||||
router.aembedding = AsyncMock(return_value={"data": [{"embedding": [0.1, 0.2]}]})
|
||||
monkeypatch.setitem(
|
||||
sys.modules,
|
||||
"litellm.proxy.proxy_server",
|
||||
_router_proxy_module(router, "sem-embed"),
|
||||
)
|
||||
|
||||
await cache._get_async_embedding("What is the capital of France?")
|
||||
|
||||
assert router.aembedding.call_args.kwargs["timeout"] == 1.5
|
||||
assert router.aembedding.call_args.kwargs["num_retries"] == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_qdrant_async_embedding_gives_up_on_unresponsive_endpoint(monkeypatch):
|
||||
import asyncio
|
||||
import time
|
||||
|
||||
from litellm.caching.qdrant_semantic_cache import QdrantSemanticCache
|
||||
|
||||
cache = QdrantSemanticCache.__new__(QdrantSemanticCache)
|
||||
cache.embedding_model = "sem-embed"
|
||||
cache.embedding_max_input_tokens = None
|
||||
cache.embedding_timeout = 0.05
|
||||
|
||||
async def never_responds(**kwargs):
|
||||
await asyncio.sleep(3)
|
||||
return {"data": [{"embedding": [0.1, 0.2]}]}
|
||||
|
||||
router = MagicMock()
|
||||
router.get_configured_token_limits.return_value = (None, None)
|
||||
router.aembedding = never_responds
|
||||
monkeypatch.setitem(
|
||||
sys.modules,
|
||||
"litellm.proxy.proxy_server",
|
||||
_router_proxy_module(router, "sem-embed"),
|
||||
)
|
||||
|
||||
started = time.monotonic()
|
||||
with pytest.raises(asyncio.TimeoutError):
|
||||
await cache._get_async_embedding("What is the capital of France?")
|
||||
assert time.monotonic() - started < 1.0
|
||||
|
||||
|
||||
def test_qdrant_semantic_cache_defaults_embedding_timeout():
|
||||
from litellm.caching.qdrant_semantic_cache import QdrantSemanticCache
|
||||
from litellm.constants import SEMANTIC_CACHE_EMBEDDING_TIMEOUT_SECONDS
|
||||
|
||||
cache = QdrantSemanticCache.__new__(QdrantSemanticCache)
|
||||
assert cache.embedding_timeout == SEMANTIC_CACHE_EMBEDDING_TIMEOUT_SECONDS
|
||||
assert SEMANTIC_CACHE_EMBEDDING_TIMEOUT_SECONDS < 60
|
||||
|
|
|
|||
|
|
@ -1329,3 +1329,157 @@ def test_redis_llmcache_setter_supported():
|
|||
sentinel = MagicMock()
|
||||
cache.llmcache = sentinel
|
||||
assert cache.llmcache is sentinel
|
||||
|
||||
|
||||
def _router_proxy_module(router, model_name):
|
||||
import types
|
||||
|
||||
fake_proxy = types.ModuleType("litellm.proxy.proxy_server")
|
||||
fake_proxy.llm_router = router
|
||||
fake_proxy.llm_model_list = [{"model_name": model_name}]
|
||||
return fake_proxy
|
||||
|
||||
|
||||
def test_redis_sync_embedding_call_is_bounded(monkeypatch):
|
||||
import sys
|
||||
|
||||
from litellm.caching.redis_semantic_cache import RedisSemanticCache
|
||||
|
||||
cache = RedisSemanticCache.__new__(RedisSemanticCache)
|
||||
cache.embedding_model = "sem-embed"
|
||||
cache.embedding_timeout = 1.5
|
||||
|
||||
router = MagicMock()
|
||||
router.get_configured_token_limits.return_value = (None, None)
|
||||
router.embedding = MagicMock(return_value={"data": [{"embedding": [0.5, 0.6]}]})
|
||||
monkeypatch.setitem(
|
||||
sys.modules,
|
||||
"litellm.proxy.proxy_server",
|
||||
_router_proxy_module(router, "sem-embed"),
|
||||
)
|
||||
|
||||
assert cache._get_embedding("hello") == [0.5, 0.6]
|
||||
assert router.embedding.call_args.kwargs["timeout"] == 1.5
|
||||
assert router.embedding.call_args.kwargs["num_retries"] == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_redis_async_embedding_call_is_bounded(monkeypatch):
|
||||
import sys
|
||||
|
||||
from litellm.caching.redis_semantic_cache import RedisSemanticCache
|
||||
|
||||
cache = RedisSemanticCache.__new__(RedisSemanticCache)
|
||||
cache.embedding_model = "sem-embed"
|
||||
cache.embedding_timeout = 1.5
|
||||
|
||||
router = MagicMock()
|
||||
router.get_configured_token_limits.return_value = (None, None)
|
||||
router.aembedding = AsyncMock(return_value={"data": [{"embedding": [0.5, 0.6]}]})
|
||||
monkeypatch.setitem(
|
||||
sys.modules,
|
||||
"litellm.proxy.proxy_server",
|
||||
_router_proxy_module(router, "sem-embed"),
|
||||
)
|
||||
|
||||
assert await cache._get_async_embedding("hello") == [0.5, 0.6]
|
||||
assert router.aembedding.call_args.kwargs["timeout"] == 1.5
|
||||
assert router.aembedding.call_args.kwargs["num_retries"] == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_redis_async_embedding_gives_up_on_unresponsive_endpoint(monkeypatch):
|
||||
import asyncio
|
||||
import sys
|
||||
import time
|
||||
|
||||
from litellm.caching.redis_semantic_cache import RedisSemanticCache
|
||||
|
||||
cache = RedisSemanticCache.__new__(RedisSemanticCache)
|
||||
cache.embedding_model = "sem-embed"
|
||||
cache.embedding_timeout = 0.05
|
||||
|
||||
async def never_responds(**kwargs):
|
||||
await asyncio.sleep(3)
|
||||
return {"data": [{"embedding": [0.1, 0.2]}]}
|
||||
|
||||
router = MagicMock()
|
||||
router.get_configured_token_limits.return_value = (None, None)
|
||||
router.aembedding = never_responds
|
||||
monkeypatch.setitem(
|
||||
sys.modules,
|
||||
"litellm.proxy.proxy_server",
|
||||
_router_proxy_module(router, "sem-embed"),
|
||||
)
|
||||
|
||||
started = time.monotonic()
|
||||
with pytest.raises(ValueError, match="Failed to generate embedding"):
|
||||
await cache._get_async_embedding("hello")
|
||||
assert time.monotonic() - started < 1.0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_redis_async_get_cache_fails_open_when_embedding_hangs(monkeypatch):
|
||||
import asyncio
|
||||
import sys
|
||||
import time
|
||||
|
||||
from litellm.caching.redis_semantic_cache import RedisSemanticCache
|
||||
|
||||
cache = RedisSemanticCache.__new__(RedisSemanticCache)
|
||||
cache.embedding_model = "sem-embed"
|
||||
cache.embedding_timeout = 0.05
|
||||
cache.similarity_threshold = 0.8
|
||||
cache.distance_threshold = 0.2
|
||||
cache.llmcache = MagicMock()
|
||||
|
||||
async def never_responds(**kwargs):
|
||||
await asyncio.sleep(3)
|
||||
return {"data": [{"embedding": [0.1, 0.2]}]}
|
||||
|
||||
router = MagicMock()
|
||||
router.get_configured_token_limits.return_value = (None, None)
|
||||
router.aembedding = never_responds
|
||||
monkeypatch.setitem(
|
||||
sys.modules,
|
||||
"litellm.proxy.proxy_server",
|
||||
_router_proxy_module(router, "sem-embed"),
|
||||
)
|
||||
|
||||
metadata = {}
|
||||
started = time.monotonic()
|
||||
result = await cache.async_get_cache(
|
||||
key="test_key",
|
||||
messages=[{"role": "user", "content": "What is the capital of France?"}],
|
||||
metadata=metadata,
|
||||
)
|
||||
elapsed = time.monotonic() - started
|
||||
|
||||
assert result is None
|
||||
assert metadata["semantic-similarity"] == 0.0
|
||||
assert elapsed < 1.0
|
||||
cache.llmcache.acheck.assert_not_called()
|
||||
|
||||
|
||||
def test_cache_forwards_semantic_cache_embedding_timeout():
|
||||
from litellm.caching.caching import Cache
|
||||
from litellm.types.caching import LiteLLMCacheType
|
||||
|
||||
with patch("litellm.caching.caching.RedisSemanticCache") as backend:
|
||||
Cache(
|
||||
type=LiteLLMCacheType.REDIS_SEMANTIC,
|
||||
similarity_threshold=0.8,
|
||||
redis_url="redis://localhost:6379",
|
||||
semantic_cache_embedding_timeout=2.5,
|
||||
)
|
||||
|
||||
assert backend.call_args.kwargs["embedding_timeout"] == 2.5
|
||||
|
||||
|
||||
def test_redis_semantic_cache_defaults_embedding_timeout():
|
||||
from litellm.caching.redis_semantic_cache import RedisSemanticCache
|
||||
from litellm.constants import SEMANTIC_CACHE_EMBEDDING_TIMEOUT_SECONDS
|
||||
|
||||
cache = RedisSemanticCache.__new__(RedisSemanticCache)
|
||||
assert cache.embedding_timeout == SEMANTIC_CACHE_EMBEDDING_TIMEOUT_SECONDS
|
||||
assert SEMANTIC_CACHE_EMBEDDING_TIMEOUT_SECONDS < 60
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue