diff --git a/litellm/caching/_embedding_router.py b/litellm/caching/_embedding_router.py index 1073b34ef25..8dfcddf158a 100644 --- a/litellm/caching/_embedding_router.py +++ b/litellm/caching/_embedding_router.py @@ -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 diff --git a/litellm/caching/caching.py b/litellm/caching/caching.py index f0fb91b987f..6b68ae98111 100644 --- a/litellm/caching/caching.py +++ b/litellm/caching/caching.py @@ -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() diff --git a/litellm/caching/qdrant_semantic_cache.py b/litellm/caching/qdrant_semantic_cache.py index 8f8323550f3..8270c655d82 100644 --- a/litellm/caching/qdrant_semantic_cache.py +++ b/litellm/caching/qdrant_semantic_cache.py @@ -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}, ) diff --git a/litellm/caching/redis_semantic_cache.py b/litellm/caching/redis_semantic_cache.py index 604d6395ea1..d91260f4d9c 100644 --- a/litellm/caching/redis_semantic_cache.py +++ b/litellm/caching/redis_semantic_cache.py @@ -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"] diff --git a/litellm/caching/valkey_semantic_cache.py b/litellm/caching/valkey_semantic_cache.py index aa10d91fc66..e887861e1dc 100644 --- a/litellm/caching/valkey_semantic_cache.py +++ b/litellm/caching/valkey_semantic_cache.py @@ -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 diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index d67497dd4da..c2caa17fc57 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -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 diff --git a/litellm/utils.py b/litellm/utils.py index 68f4278c87a..d50871182b9 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -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: diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json index 4a6ff6af902..06bee4a054c 100644 --- a/ruff-strict-budget.json +++ b/ruff-strict-budget.json @@ -33,7 +33,7 @@ "limit": 2 }, "B006": { - "limit": 178 + "limit": 177 }, "B008": { "limit": 503 diff --git a/tests/test_litellm/caching/test_caching.py b/tests/test_litellm/caching/test_caching.py index b65e8773c85..955b0e531bc 100644 --- a/tests/test_litellm/caching/test_caching.py +++ b/tests/test_litellm/caching/test_caching.py @@ -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"} diff --git a/tests/test_litellm/caching/test_embedding_router.py b/tests/test_litellm/caching/test_embedding_router.py index 550095a112a..9ebe669d32d 100644 --- a/tests/test_litellm/caching/test_embedding_router.py +++ b/tests/test_litellm/caching/test_embedding_router.py @@ -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 diff --git a/tests/test_litellm/caching/test_qdrant_semantic_cache.py b/tests/test_litellm/caching/test_qdrant_semantic_cache.py index 67d4e2d9892..852bed4a9df 100644 --- a/tests/test_litellm/caching/test_qdrant_semantic_cache.py +++ b/tests/test_litellm/caching/test_qdrant_semantic_cache.py @@ -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 diff --git a/tests/test_litellm/caching/test_redis_semantic_cache.py b/tests/test_litellm/caching/test_redis_semantic_cache.py index 1d3129d6467..ad0ea6b7774 100644 --- a/tests/test_litellm/caching/test_redis_semantic_cache.py +++ b/tests/test_litellm/caching/test_redis_semantic_cache.py @@ -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() diff --git a/tests/test_litellm/caching/test_valkey_semantic_cache.py b/tests/test_litellm/caching/test_valkey_semantic_cache.py index d2df0a98e12..acf5a914e5c 100644 --- a/tests/test_litellm/caching/test_valkey_semantic_cache.py +++ b/tests/test_litellm/caching/test_valkey_semantic_cache.py @@ -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( diff --git a/tests/test_litellm/llms/hosted_vllm/embedding/test_hosted_vllm_embedding_transformation.py b/tests/test_litellm/llms/hosted_vllm/embedding/test_hosted_vllm_embedding_transformation.py index 35c0a63573f..29d58b35b84 100644 --- a/tests/test_litellm/llms/hosted_vllm/embedding/test_hosted_vllm_embedding_transformation.py +++ b/tests/test_litellm/llms/hosted_vllm/embedding/test_hosted_vllm_embedding_transformation.py @@ -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"]) diff --git a/type-discipline-budget.json b/type-discipline-budget.json index b5c07de8cd3..94c1f9b86f9 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -1,9 +1,9 @@ { "LIT001": { - "limit": 22900 + "limit": 22897 }, "LIT002": { - "limit": 26889 + "limit": 26888 }, "LIT003": { "limit": 269