fix(caching): truncate semantic cache embedding input, send extra_body top-level

This commit is contained in:
mateo-berri 2026-08-18 15:03:51 -07:00
parent 6d32d4081d
commit ef2c30227a
15 changed files with 359 additions and 18 deletions

View file

@ -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

View file

@ -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()

View file

@ -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},
)

View file

@ -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"]

View file

@ -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

View file

@ -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

View file

@ -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:

View file

@ -33,7 +33,7 @@
"limit": 2
},
"B006": {
"limit": 178
"limit": 177
},
"B008": {
"limit": 503

View file

@ -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"}

View file

@ -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

View file

@ -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

View file

@ -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()

View file

@ -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(

View file

@ -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"])

View file

@ -1,9 +1,9 @@
{
"LIT001": {
"limit": 22900
"limit": 22897
},
"LIT002": {
"limit": 26889
"limit": 26888
},
"LIT003": {
"limit": 269