mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(caching): truncate semantic cache embedding input, send extra_body top-level
This commit is contained in:
parent
6d32d4081d
commit
ef2c30227a
15 changed files with 359 additions and 18 deletions
|
|
@ -12,8 +12,11 @@ This module is dependency-injected: callers pass the proxy ``llm_router`` and
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
import litellm
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.router import Router
|
||||
|
||||
|
|
@ -41,3 +44,28 @@ def build_router_embedding_metadata(
|
|||
metadata: Final[dict[str, Any]] = dict(request_metadata or {})
|
||||
metadata["semantic-cache-embedding"] = True
|
||||
return metadata
|
||||
|
||||
|
||||
def resolve_embedding_max_input_tokens(
|
||||
configured_max_input_tokens: int | None,
|
||||
embedding_model: str,
|
||||
router: Router | None,
|
||||
) -> int | None:
|
||||
"""Explicit cache setting first, else the Router deployment's configured ``max_input_tokens``."""
|
||||
if configured_max_input_tokens is not None:
|
||||
return configured_max_input_tokens
|
||||
if router is None:
|
||||
return None
|
||||
deployment_max_input_tokens, _ = router.get_configured_token_limits(embedding_model)
|
||||
return deployment_max_input_tokens
|
||||
|
||||
|
||||
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:
|
||||
return prompt
|
||||
tokens: Final[Sequence[int]] = litellm.encode(model=embedding_model, text=prompt)
|
||||
if len(tokens) <= max_input_tokens:
|
||||
return prompt
|
||||
truncated: Final[str] = litellm.decode(model=embedding_model, tokens=tokens[:max_input_tokens])
|
||||
return truncated
|
||||
|
|
|
|||
|
|
@ -97,6 +97,7 @@ class Cache:
|
|||
qdrant_quantization_config: str | None = None,
|
||||
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,
|
||||
# GCP IAM authentication parameters
|
||||
gcp_service_account: str | None = None,
|
||||
gcp_ssl_ca_certs: str | None = None,
|
||||
|
|
@ -122,6 +123,7 @@ class Cache:
|
|||
qdrant_api_key (str, optional): The api_key for the local or cloud qdrant cluster.
|
||||
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.
|
||||
|
||||
# Disk Cache Args
|
||||
disk_cache_dir (str, optional): The directory for the disk cache. Defaults to None.
|
||||
|
|
@ -192,6 +194,7 @@ class Cache:
|
|||
similarity_threshold=similarity_threshold,
|
||||
embedding_model=redis_semantic_cache_embedding_model,
|
||||
index_name=redis_semantic_cache_index_name,
|
||||
embedding_max_input_tokens=semantic_cache_embedding_max_input_tokens,
|
||||
**kwargs,
|
||||
)
|
||||
elif type == LiteLLMCacheType.VALKEY_SEMANTIC:
|
||||
|
|
@ -207,6 +210,7 @@ class Cache:
|
|||
embedding_model=valkey_semantic_cache_embedding_model,
|
||||
index_name=valkey_semantic_cache_index_name,
|
||||
startup_nodes=redis_startup_nodes,
|
||||
embedding_max_input_tokens=semantic_cache_embedding_max_input_tokens,
|
||||
**kwargs,
|
||||
)
|
||||
elif type == LiteLLMCacheType.QDRANT_SEMANTIC:
|
||||
|
|
@ -218,6 +222,7 @@ class Cache:
|
|||
quantization_config=qdrant_quantization_config,
|
||||
embedding_model=qdrant_semantic_cache_embedding_model,
|
||||
vector_size=qdrant_semantic_cache_vector_size,
|
||||
embedding_max_input_tokens=semantic_cache_embedding_max_input_tokens,
|
||||
)
|
||||
elif type == LiteLLMCacheType.LOCAL:
|
||||
self.cache = InMemoryCache()
|
||||
|
|
|
|||
|
|
@ -12,7 +12,7 @@ import ast
|
|||
import asyncio
|
||||
import json
|
||||
import os
|
||||
from typing import Any, Final, cast
|
||||
from typing import TYPE_CHECKING, Any, Final, cast
|
||||
|
||||
import litellm
|
||||
from litellm._logging import print_verbose
|
||||
|
|
@ -22,12 +22,21 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
|||
)
|
||||
from litellm.types.utils import EmbeddingResponse
|
||||
|
||||
from ._embedding_router import build_router_embedding_metadata, resolve_embedding_router
|
||||
from ._embedding_router import (
|
||||
build_router_embedding_metadata,
|
||||
resolve_embedding_max_input_tokens,
|
||||
resolve_embedding_router,
|
||||
truncate_embedding_input,
|
||||
)
|
||||
from .base_cache import BaseCache
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.router import Router
|
||||
|
||||
|
||||
class QdrantSemanticCache(BaseCache):
|
||||
CACHE_KEY_FIELD_NAME = "litellm_cache_key"
|
||||
embedding_max_input_tokens: int | None = None
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
|
|
@ -39,6 +48,7 @@ class QdrantSemanticCache(BaseCache):
|
|||
embedding_model="text-embedding-ada-002",
|
||||
host_type=None,
|
||||
vector_size=None,
|
||||
embedding_max_input_tokens: int | None = None,
|
||||
):
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
_get_httpx_client,
|
||||
|
|
@ -57,6 +67,7 @@ class QdrantSemanticCache(BaseCache):
|
|||
raise Exception("similarity_threshold must be provided, passed None")
|
||||
self.similarity_threshold = similarity_threshold
|
||||
self.embedding_model = embedding_model
|
||||
self.embedding_max_input_tokens = embedding_max_input_tokens
|
||||
self.vector_size = vector_size if vector_size is not None else QDRANT_VECTOR_SIZE
|
||||
headers = {}
|
||||
|
||||
|
|
@ -188,6 +199,13 @@ class QdrantSemanticCache(BaseCache):
|
|||
cached_key: Final = payload.get(self.CACHE_KEY_FIELD_NAME)
|
||||
return cached_key is not None and str(cached_key) == str(key)
|
||||
|
||||
def _embedding_input(self, prompt: str, router: "Router | None") -> str:
|
||||
return truncate_embedding_input(
|
||||
prompt,
|
||||
self.embedding_model,
|
||||
resolve_embedding_max_input_tokens(self.embedding_max_input_tokens, self.embedding_model, router),
|
||||
)
|
||||
|
||||
def _get_embedding(self, prompt: str, metadata: dict[str, Any] | None = None) -> EmbeddingResponse:
|
||||
"""Embed via the proxy Router when it serves the model, else direct."""
|
||||
try:
|
||||
|
|
@ -197,16 +215,17 @@ class QdrantSemanticCache(BaseCache):
|
|||
llm_router = None
|
||||
|
||||
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 router.embedding(
|
||||
model=self.embedding_model,
|
||||
input=prompt,
|
||||
input=embedding_input,
|
||||
cache={"no-store": True, "no-cache": True},
|
||||
metadata=build_router_embedding_metadata(metadata),
|
||||
)
|
||||
return litellm.embedding(
|
||||
model=self.embedding_model,
|
||||
input=prompt,
|
||||
input=embedding_input,
|
||||
cache={"no-store": True, "no-cache": True},
|
||||
)
|
||||
|
||||
|
|
@ -218,17 +237,18 @@ class QdrantSemanticCache(BaseCache):
|
|||
llm_router = None
|
||||
|
||||
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(
|
||||
model=self.embedding_model,
|
||||
input=prompt,
|
||||
input=embedding_input,
|
||||
cache={"no-store": True, "no-cache": True},
|
||||
metadata=build_router_embedding_metadata(metadata),
|
||||
)
|
||||
|
||||
return await litellm.aembedding(
|
||||
model=self.embedding_model,
|
||||
input=prompt,
|
||||
input=embedding_input,
|
||||
cache={"no-store": True, "no-cache": True},
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -14,7 +14,7 @@ import asyncio
|
|||
import json
|
||||
import os
|
||||
from collections.abc import Callable, Mapping
|
||||
from typing import Any, Final, cast
|
||||
from typing import TYPE_CHECKING, Any, Final, cast
|
||||
|
||||
import litellm
|
||||
from litellm._logging import print_verbose, verbose_logger
|
||||
|
|
@ -23,9 +23,17 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
|||
)
|
||||
from litellm.types.utils import EmbeddingResponse
|
||||
|
||||
from ._embedding_router import build_router_embedding_metadata, resolve_embedding_router
|
||||
from ._embedding_router import (
|
||||
build_router_embedding_metadata,
|
||||
resolve_embedding_max_input_tokens,
|
||||
resolve_embedding_router,
|
||||
truncate_embedding_input,
|
||||
)
|
||||
from .base_cache import BaseCache
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.router import Router
|
||||
|
||||
|
||||
class RedisSemanticCache(BaseCache):
|
||||
"""
|
||||
|
|
@ -38,6 +46,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
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
|
|
@ -48,6 +57,7 @@ class RedisSemanticCache(BaseCache):
|
|||
similarity_threshold: float | None = None,
|
||||
embedding_model: str = "text-embedding-ada-002",
|
||||
index_name: str | None = None,
|
||||
embedding_max_input_tokens: int | None = None,
|
||||
**kwargs: object,
|
||||
):
|
||||
"""
|
||||
|
|
@ -62,6 +72,8 @@ class RedisSemanticCache(BaseCache):
|
|||
where 1.0 requires exact matches and 0.0 accepts any match
|
||||
embedding_model: Model to use for generating embeddings
|
||||
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
|
||||
ttl: Default time-to-live for cache entries in seconds
|
||||
**kwargs: Additional arguments passed to the Redis client
|
||||
|
||||
|
|
@ -86,6 +98,7 @@ class RedisSemanticCache(BaseCache):
|
|||
# While similarity: 1 = most similar, 0 = least similar
|
||||
self.distance_threshold = 1 - similarity_threshold
|
||||
self.embedding_model = embedding_model
|
||||
self.embedding_max_input_tokens = embedding_max_input_tokens
|
||||
|
||||
# Set up Redis connection
|
||||
if redis_url is None:
|
||||
|
|
@ -307,6 +320,13 @@ class RedisSemanticCache(BaseCache):
|
|||
return dict_method()
|
||||
return value
|
||||
|
||||
def _embedding_input(self, prompt: str, router: "Router | None") -> str:
|
||||
return truncate_embedding_input(
|
||||
prompt,
|
||||
self.embedding_model,
|
||||
resolve_embedding_max_input_tokens(self.embedding_max_input_tokens, self.embedding_model, router),
|
||||
)
|
||||
|
||||
def _get_embedding(self, prompt: str, metadata: dict[str, Any] | None = None) -> list[float]:
|
||||
"""
|
||||
Routes through the proxy Router when the embedding model is a Router
|
||||
|
|
@ -320,12 +340,13 @@ class RedisSemanticCache(BaseCache):
|
|||
llm_router = None
|
||||
|
||||
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:
|
||||
embedding_response = cast(
|
||||
EmbeddingResponse,
|
||||
router.embedding(
|
||||
model=self.embedding_model,
|
||||
input=prompt,
|
||||
input=embedding_input,
|
||||
cache={"no-store": True, "no-cache": True},
|
||||
metadata=build_router_embedding_metadata(metadata),
|
||||
),
|
||||
|
|
@ -335,7 +356,7 @@ class RedisSemanticCache(BaseCache):
|
|||
EmbeddingResponse,
|
||||
litellm.embedding(
|
||||
model=self.embedding_model,
|
||||
input=prompt,
|
||||
input=embedding_input,
|
||||
cache={"no-store": True, "no-cache": True},
|
||||
),
|
||||
)
|
||||
|
|
@ -490,18 +511,19 @@ class RedisSemanticCache(BaseCache):
|
|||
llm_router = None
|
||||
|
||||
router: Final = resolve_embedding_router(self.embedding_model, llm_router, llm_model_list)
|
||||
embedding_input: Final = self._embedding_input(prompt, router)
|
||||
try:
|
||||
if router is not None:
|
||||
embedding_response = await router.aembedding(
|
||||
model=self.embedding_model,
|
||||
input=prompt,
|
||||
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=prompt,
|
||||
input=embedding_input,
|
||||
cache={"no-store": True, "no-cache": True},
|
||||
)
|
||||
return embedding_response["data"][0]["embedding"]
|
||||
|
|
|
|||
|
|
@ -61,6 +61,7 @@ class ValkeySemanticCache(RedisSemanticCache):
|
|||
startup_nodes: list | None = None,
|
||||
sync_client: Redis | None = None,
|
||||
async_client: AsyncRedis | None = None,
|
||||
embedding_max_input_tokens: int | None = None,
|
||||
**kwargs: Any,
|
||||
):
|
||||
if similarity_threshold is None:
|
||||
|
|
@ -78,6 +79,7 @@ class ValkeySemanticCache(RedisSemanticCache):
|
|||
|
||||
self.similarity_threshold = similarity_threshold
|
||||
self.embedding_model = embedding_model
|
||||
self.embedding_max_input_tokens = embedding_max_input_tokens
|
||||
self.index_name = index_name or self.DEFAULT_VALKEY_INDEX_NAME
|
||||
self.key_prefix = f"{self.index_name}:"
|
||||
self._index_dim: int | None = None
|
||||
|
|
|
|||
|
|
@ -893,6 +893,7 @@ class BaseLLMHTTPHandler:
|
|||
)
|
||||
if provider_config is None:
|
||||
raise ValueError(f"Provider {custom_llm_provider} does not support embedding")
|
||||
embedding_extra_body: Final[Mapping[str, object] | None] = optional_params.pop("extra_body", None)
|
||||
# get config from model, custom llm provider
|
||||
headers = provider_config.validate_environment(
|
||||
api_key=api_key,
|
||||
|
|
@ -917,6 +918,8 @@ class BaseLLMHTTPHandler:
|
|||
optional_params=optional_params,
|
||||
headers=headers,
|
||||
)
|
||||
if embedding_extra_body:
|
||||
data.update(embedding_extra_body)
|
||||
|
||||
# Some providers (e.g. OCI) require request signing after the body is built.
|
||||
# The default BaseConfig.sign_request returns (headers, None) — a no-op for
|
||||
|
|
|
|||
|
|
@ -2110,7 +2110,7 @@ def encode(model="", text="", custom_tokenizer: dict | None = None):
|
|||
|
||||
def decode(
|
||||
model="",
|
||||
tokens: list[int] = [],
|
||||
tokens: Sequence[int] = (),
|
||||
custom_tokenizer: dict | None = None,
|
||||
skip_special_tokens: bool = True,
|
||||
):
|
||||
|
|
@ -2132,7 +2132,7 @@ def decode(
|
|||
return dec
|
||||
|
||||
|
||||
def _strip_huggingface_special_token_ids(tokenizer: Tokenizer, tokens: list[int]) -> list[int]:
|
||||
def _strip_huggingface_special_token_ids(tokenizer: Tokenizer, tokens: Sequence[int]) -> Sequence[int]:
|
||||
try:
|
||||
added_tokens_decoder: Final = tokenizer.get_added_tokens_decoder()
|
||||
except Exception:
|
||||
|
|
|
|||
|
|
@ -33,7 +33,7 @@
|
|||
"limit": 2
|
||||
},
|
||||
"B006": {
|
||||
"limit": 178
|
||||
"limit": 177
|
||||
},
|
||||
"B008": {
|
||||
"limit": 503
|
||||
|
|
|
|||
|
|
@ -89,6 +89,20 @@ def _semantic_cache():
|
|||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"cache_type",
|
||||
[LiteLLMCacheType.REDIS_SEMANTIC, LiteLLMCacheType.VALKEY_SEMANTIC],
|
||||
)
|
||||
def test_semantic_cache_embedding_max_input_tokens_reaches_backend(cache_type):
|
||||
cache = Cache(
|
||||
type=cache_type,
|
||||
redis_url="redis://localhost:6379",
|
||||
similarity_threshold=0.8,
|
||||
semantic_cache_embedding_max_input_tokens=2048,
|
||||
)
|
||||
assert cache.cache.embedding_max_input_tokens == 2048
|
||||
|
||||
|
||||
def test_semantic_cache_key_excludes_prompt_so_paraphrases_share_a_bucket():
|
||||
cache = _semantic_cache()
|
||||
tenant = {"user_api_key": "hash-abc"}
|
||||
|
|
|
|||
|
|
@ -4,9 +4,12 @@ from unittest.mock import MagicMock
|
|||
|
||||
sys.path.insert(0, os.path.abspath("../../.."))
|
||||
|
||||
import litellm
|
||||
from litellm.caching._embedding_router import (
|
||||
build_router_embedding_metadata,
|
||||
resolve_embedding_max_input_tokens,
|
||||
resolve_embedding_router,
|
||||
truncate_embedding_input,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -65,3 +68,40 @@ def test_build_metadata_handles_none_and_does_not_mutate_input():
|
|||
assert md == {"user_api_key": "sk-x", "semantic-cache-embedding": True}
|
||||
assert original == {"user_api_key": "sk-x"}
|
||||
assert build_router_embedding_metadata(None) == {"semantic-cache-embedding": True}
|
||||
|
||||
|
||||
def test_resolve_max_input_tokens_prefers_configured_over_deployment():
|
||||
router = MagicMock()
|
||||
router.get_configured_token_limits.return_value = (8191, None)
|
||||
assert resolve_embedding_max_input_tokens(512, "sem-embed", router) == 512
|
||||
router.get_configured_token_limits.assert_not_called()
|
||||
|
||||
|
||||
def test_resolve_max_input_tokens_falls_back_to_deployment_limit():
|
||||
router = MagicMock()
|
||||
router.get_configured_token_limits.return_value = (8191, 4096)
|
||||
assert resolve_embedding_max_input_tokens(None, "sem-embed", router) == 8191
|
||||
router.get_configured_token_limits.assert_called_once_with("sem-embed")
|
||||
|
||||
|
||||
def test_resolve_max_input_tokens_is_none_without_router_or_deployment_limit():
|
||||
router = MagicMock()
|
||||
router.get_configured_token_limits.return_value = (None, None)
|
||||
assert resolve_embedding_max_input_tokens(None, "sem-embed", router) is None
|
||||
assert resolve_embedding_max_input_tokens(None, "sem-embed", None) is None
|
||||
|
||||
|
||||
def test_truncate_embedding_input_keeps_prompt_within_limit():
|
||||
prompt = "The quick brown fox jumps over the lazy dog"
|
||||
assert truncate_embedding_input(prompt, "sem-embed", None) == prompt
|
||||
assert truncate_embedding_input(prompt, "sem-embed", 100) == prompt
|
||||
token_count = len(litellm.encode(model="sem-embed", text=prompt))
|
||||
assert truncate_embedding_input(prompt, "sem-embed", token_count) == prompt
|
||||
|
||||
|
||||
def test_truncate_embedding_input_cuts_prompt_to_token_limit():
|
||||
prompt = " ".join(f"word{i}" for i in range(400))
|
||||
truncated = truncate_embedding_input(prompt, "sem-embed", 50)
|
||||
assert prompt.startswith(truncated)
|
||||
assert len(truncated) < len(prompt)
|
||||
assert len(litellm.encode(model="sem-embed", text=truncated)) == 50
|
||||
|
|
|
|||
|
|
@ -43,6 +43,7 @@ def test_qdrant_semantic_cache_initialization(monkeypatch):
|
|||
qdrant_api_base="http://test.qdrant.local",
|
||||
qdrant_api_key="test_key",
|
||||
similarity_threshold=0.8,
|
||||
embedding_max_input_tokens=512,
|
||||
)
|
||||
|
||||
# Verify the cache was initialized with correct parameters
|
||||
|
|
@ -50,6 +51,7 @@ def test_qdrant_semantic_cache_initialization(monkeypatch):
|
|||
assert qdrant_cache.qdrant_api_base == "http://test.qdrant.local"
|
||||
assert qdrant_cache.qdrant_api_key == "test_key"
|
||||
assert qdrant_cache.similarity_threshold == 0.8
|
||||
assert qdrant_cache.embedding_max_input_tokens == 512
|
||||
mock_sync_client_instance.put.assert_called_once_with(
|
||||
url="http://test.qdrant.local/collections/test_collection/index",
|
||||
headers={
|
||||
|
|
@ -832,6 +834,7 @@ def test_qdrant_sync_get_cache_routes_through_router(monkeypatch):
|
|||
cache.sync_client.post.return_value = search_response
|
||||
|
||||
router = MagicMock()
|
||||
router.get_configured_token_limits.return_value = (None, None)
|
||||
router.embedding = MagicMock(
|
||||
return_value={"data": [{"embedding": [0.3, 0.3, 0.3]}]}
|
||||
)
|
||||
|
|
@ -892,6 +895,7 @@ async def test_qdrant_async_embedding_forwards_full_metadata(monkeypatch):
|
|||
cache.embedding_model = "sem-embed"
|
||||
|
||||
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,
|
||||
|
|
@ -908,3 +912,57 @@ async def test_qdrant_async_embedding_forwards_full_metadata(monkeypatch):
|
|||
assert md["user_api_key"] == "sk-x"
|
||||
assert md["user_api_key_team_id"] == "team-1"
|
||||
assert md["semantic-cache-embedding"] is True
|
||||
|
||||
|
||||
LONG_PROMPT = " ".join(f"token{i}" for i in range(300))
|
||||
|
||||
|
||||
def _token_count(model, text):
|
||||
import litellm
|
||||
|
||||
return len(litellm.encode(model=model, text=text))
|
||||
|
||||
|
||||
def test_qdrant_get_embedding_truncates_to_deployment_max_input_tokens(monkeypatch):
|
||||
from litellm.caching.qdrant_semantic_cache import QdrantSemanticCache
|
||||
|
||||
cache = QdrantSemanticCache.__new__(QdrantSemanticCache)
|
||||
cache.embedding_model = "sem-embed"
|
||||
|
||||
router = MagicMock()
|
||||
router.get_configured_token_limits.return_value = (5, 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"),
|
||||
)
|
||||
|
||||
cache._get_embedding(LONG_PROMPT)
|
||||
|
||||
sent_input = router.embedding.call_args.kwargs["input"]
|
||||
assert LONG_PROMPT.startswith(sent_input)
|
||||
assert _token_count("sem-embed", sent_input) == 5
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_qdrant_async_embedding_explicit_limit_beats_deployment_limit(monkeypatch):
|
||||
from litellm.caching.qdrant_semantic_cache import QdrantSemanticCache
|
||||
|
||||
cache = QdrantSemanticCache.__new__(QdrantSemanticCache)
|
||||
cache.embedding_model = "sem-embed"
|
||||
cache.embedding_max_input_tokens = 3
|
||||
|
||||
router = MagicMock()
|
||||
router.get_configured_token_limits.return_value = (8191, 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(LONG_PROMPT)
|
||||
|
||||
sent_input = router.aembedding.call_args.kwargs["input"]
|
||||
assert _token_count("sem-embed", sent_input) == 3
|
||||
|
|
|
|||
|
|
@ -901,6 +901,7 @@ def test_redis_get_embedding_routes_through_router(monkeypatch):
|
|||
cache.embedding_model = "sem-embed"
|
||||
|
||||
router = MagicMock()
|
||||
router.get_configured_token_limits.return_value = (None, None)
|
||||
router.embedding = MagicMock(return_value={"data": [{"embedding": [0.5, 0.6]}]})
|
||||
fake_proxy = types.ModuleType("litellm.proxy.proxy_server")
|
||||
fake_proxy.llm_router = router
|
||||
|
|
@ -1145,6 +1146,7 @@ async def test_redis_async_embedding_forwards_full_metadata(monkeypatch):
|
|||
cache.embedding_model = "sem-embed"
|
||||
|
||||
router = MagicMock()
|
||||
router.get_configured_token_limits.return_value = (None, None)
|
||||
router.aembedding = AsyncMock(return_value={"data": [{"embedding": [0.1, 0.2]}]})
|
||||
fake_proxy = types.ModuleType("litellm.proxy.proxy_server")
|
||||
fake_proxy.llm_router = router
|
||||
|
|
@ -1162,6 +1164,99 @@ async def test_redis_async_embedding_forwards_full_metadata(monkeypatch):
|
|||
assert md["semantic-cache-embedding"] is True
|
||||
|
||||
|
||||
LONG_PROMPT = " ".join(f"token{i}" for i in range(300))
|
||||
|
||||
|
||||
def _proxy_with_router(monkeypatch, router, model_name):
|
||||
import sys
|
||||
import types
|
||||
|
||||
fake_proxy = types.ModuleType("litellm.proxy.proxy_server")
|
||||
fake_proxy.llm_router = router
|
||||
fake_proxy.llm_model_list = [{"model_name": model_name}]
|
||||
monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", fake_proxy)
|
||||
|
||||
|
||||
def _token_count(model, text):
|
||||
import litellm
|
||||
|
||||
return len(litellm.encode(model=model, text=text))
|
||||
|
||||
|
||||
def test_redis_get_embedding_truncates_to_deployment_max_input_tokens(monkeypatch):
|
||||
from litellm.caching.redis_semantic_cache import RedisSemanticCache
|
||||
|
||||
cache = RedisSemanticCache.__new__(RedisSemanticCache)
|
||||
cache.embedding_model = "sem-embed"
|
||||
|
||||
router = MagicMock()
|
||||
router.get_configured_token_limits.return_value = (5, None)
|
||||
router.embedding = MagicMock(return_value={"data": [{"embedding": [0.5, 0.6]}]})
|
||||
_proxy_with_router(monkeypatch, router, "sem-embed")
|
||||
|
||||
assert cache._get_embedding(LONG_PROMPT) == [0.5, 0.6]
|
||||
|
||||
sent_input = router.embedding.call_args.kwargs["input"]
|
||||
assert LONG_PROMPT.startswith(sent_input)
|
||||
assert _token_count("sem-embed", sent_input) == 5
|
||||
assert _token_count("sem-embed", LONG_PROMPT) > 5
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_redis_async_embedding_explicit_limit_beats_deployment_limit(monkeypatch):
|
||||
from litellm.caching.redis_semantic_cache import RedisSemanticCache
|
||||
|
||||
cache = RedisSemanticCache.__new__(RedisSemanticCache)
|
||||
cache.embedding_model = "sem-embed"
|
||||
cache.embedding_max_input_tokens = 3
|
||||
|
||||
router = MagicMock()
|
||||
router.get_configured_token_limits.return_value = (8191, None)
|
||||
router.aembedding = AsyncMock(return_value={"data": [{"embedding": [0.1, 0.2]}]})
|
||||
_proxy_with_router(monkeypatch, router, "sem-embed")
|
||||
|
||||
assert await cache._get_async_embedding(LONG_PROMPT) == [0.1, 0.2]
|
||||
|
||||
sent_input = router.aembedding.call_args.kwargs["input"]
|
||||
assert _token_count("sem-embed", sent_input) == 3
|
||||
|
||||
|
||||
def test_redis_get_embedding_truncates_direct_path_with_explicit_limit(monkeypatch):
|
||||
import sys
|
||||
import types
|
||||
|
||||
from litellm.caching.redis_semantic_cache import RedisSemanticCache
|
||||
|
||||
cache = RedisSemanticCache.__new__(RedisSemanticCache)
|
||||
cache.embedding_model = "text-embedding-3-small"
|
||||
cache.embedding_max_input_tokens = 4
|
||||
|
||||
fake_proxy = types.ModuleType("litellm.proxy.proxy_server")
|
||||
fake_proxy.llm_router = None
|
||||
fake_proxy.llm_model_list = None
|
||||
monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", fake_proxy)
|
||||
|
||||
with patch(
|
||||
"litellm.embedding", return_value={"data": [{"embedding": [0.1, 0.2]}]}
|
||||
) as direct_embed:
|
||||
cache._get_embedding(LONG_PROMPT)
|
||||
|
||||
sent_input = direct_embed.call_args.kwargs["input"]
|
||||
assert _token_count("text-embedding-3-small", sent_input) == 4
|
||||
|
||||
|
||||
def test_redis_semantic_cache_init_stores_embedding_max_input_tokens(monkeypatch):
|
||||
from litellm.caching.redis_semantic_cache import RedisSemanticCache
|
||||
|
||||
cache = RedisSemanticCache(
|
||||
redis_url="redis://localhost:6379",
|
||||
similarity_threshold=0.8,
|
||||
embedding_max_input_tokens=512,
|
||||
)
|
||||
assert cache.embedding_max_input_tokens == 512
|
||||
assert RedisSemanticCache(redis_url="redis://localhost:6379", similarity_threshold=0.8).embedding_max_input_tokens is None
|
||||
|
||||
|
||||
def test_redis_init_defers_redisvl_construction(monkeypatch):
|
||||
semantic_cache_mock = MagicMock()
|
||||
custom_vectorizer_mock = MagicMock()
|
||||
|
|
|
|||
|
|
@ -105,6 +105,17 @@ def test_init_requires_similarity_threshold():
|
|||
ValkeySemanticCache(sync_client=MagicMock(), async_client=AsyncMock())
|
||||
|
||||
|
||||
def test_init_stores_embedding_max_input_tokens():
|
||||
cache = ValkeySemanticCache(
|
||||
similarity_threshold=0.8,
|
||||
sync_client=MagicMock(),
|
||||
async_client=AsyncMock(),
|
||||
embedding_max_input_tokens=512,
|
||||
)
|
||||
assert cache.embedding_max_input_tokens == 512
|
||||
assert _make_cache().embedding_max_input_tokens is None
|
||||
|
||||
|
||||
def test_init_rejects_cluster_startup_nodes():
|
||||
with pytest.raises(ValueError, match="cluster-mode-enabled"):
|
||||
ValkeySemanticCache(
|
||||
|
|
|
|||
|
|
@ -8,7 +8,7 @@ especially ensuring that encoding_format is not included when not provided.
|
|||
import json
|
||||
import os
|
||||
import sys
|
||||
from unittest.mock import MagicMock, Mock, patch
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
|
@ -289,6 +289,49 @@ class TestHostedVLLMEmbeddingTransformation:
|
|||
assert sent_data["model"] == "BAAI/bge-small-en-v1.5"
|
||||
assert sent_data["input"] == ["Hello world"]
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"provider_params",
|
||||
[
|
||||
{"extra_body": {"truncate": "END", "input_type": "query"}},
|
||||
{"truncate": "END", "input_type": "query"},
|
||||
],
|
||||
)
|
||||
def test_provider_params_are_sent_at_the_top_level_of_the_request(self, provider_params):
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
|
||||
client = HTTPHandler()
|
||||
|
||||
with patch.object(HTTPHandler, "post") as mock_post:
|
||||
mock_response = Mock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.headers = {"content-type": "application/json"}
|
||||
mock_response.json.return_value = {
|
||||
"object": "list",
|
||||
"data": [{"object": "embedding", "index": 0, "embedding": [0.1, 0.2]}],
|
||||
"model": "nvidia/nv-embedqa-e5-v5",
|
||||
"usage": {"prompt_tokens": 2, "total_tokens": 2},
|
||||
}
|
||||
mock_response.text = json.dumps(mock_response.json.return_value)
|
||||
mock_post.return_value = mock_response
|
||||
|
||||
litellm.embedding(
|
||||
model="hosted_vllm/nvidia/nv-embedqa-e5-v5",
|
||||
input=["Hello world"],
|
||||
api_base="https://integrate.api.nvidia.com/v1",
|
||||
api_key="fake-key",
|
||||
client=client,
|
||||
caching=False,
|
||||
**provider_params,
|
||||
)
|
||||
|
||||
sent_data = json.loads(mock_post.call_args.kwargs["data"])
|
||||
|
||||
assert sent_data["truncate"] == "END"
|
||||
assert sent_data["input_type"] == "query"
|
||||
assert "extra_body" not in sent_data
|
||||
assert sent_data["model"] == "nvidia/nv-embedqa-e5-v5"
|
||||
assert sent_data["input"] == ["Hello world"]
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
pytest.main([__file__, "-v", "-s"])
|
||||
|
|
|
|||
|
|
@ -1,9 +1,9 @@
|
|||
{
|
||||
"LIT001": {
|
||||
"limit": 22900
|
||||
"limit": 22897
|
||||
},
|
||||
"LIT002": {
|
||||
"limit": 26889
|
||||
"limit": 26888
|
||||
},
|
||||
"LIT003": {
|
||||
"limit": 269
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue