diff --git a/litellm/router.py b/litellm/router.py index a58b3ce25e1..a0ff3e0fd31 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -58,6 +58,7 @@ from litellm.constants import ( DEFAULT_HEALTH_CHECK_INTERVAL, DEFAULT_HEALTH_CHECK_STALENESS_MULTIPLIER, DEFAULT_MAX_LRU_CACHE_SIZE, + MINIMUM_PROMPT_CACHE_TOKEN_COUNT, ) from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.asyncify import run_async_function @@ -307,6 +308,8 @@ class Router: ] = RouterGeneralSettings(), deployment_affinity_ttl_seconds: int = 3600, model_group_affinity_config: Optional[Dict[str, List[str]]] = None, + prompt_prefix_affinity_tokens: int = 2048, + prompt_prefix_affinity_min_tokens: int = MINIMUM_PROMPT_CACHE_TOKEN_COUNT, ignore_invalid_deployments: bool = False, enable_health_check_routing: bool = False, health_check_staleness_threshold: Optional[int] = None, @@ -344,6 +347,11 @@ class Router: alerting_config (AlertingConfig): Slack alerting configuration. Defaults to None. provider_budget_config (ProviderBudgetConfig): Provider budget configuration. Use this to set llm_provider budget limits. example $100/day to OpenAI, $100/day to Azure, etc. Defaults to None. deployment_affinity_ttl_seconds (int): TTL for user-key -> deployment affinity mapping. Defaults to 3600. + prompt_prefix_affinity_tokens (int): Number of canonical prompt-prefix + tokens used for deterministic prompt-prefix routing. Defaults to 2048. + prompt_prefix_affinity_min_tokens (int): Minimum canonical prompt token + count before prompt-prefix affinity applies. Defaults to + MINIMUM_PROMPT_CACHE_TOKEN_COUNT. ignore_invalid_deployments (bool): Ignores invalid deployments, and continues with other deployments. Default is to raise an error. Returns: Router: An instance of the litellm.Router class. @@ -636,6 +644,8 @@ class Router: self.routing_strategy_args = routing_strategy_args self.provider_budget_config = provider_budget_config self.deployment_affinity_ttl_seconds = deployment_affinity_ttl_seconds + self.prompt_prefix_affinity_tokens = prompt_prefix_affinity_tokens + self.prompt_prefix_affinity_min_tokens = prompt_prefix_affinity_min_tokens self.router_budget_logger: Optional[RouterBudgetLimiting] = None if RouterBudgetLimiting.should_init_router_budget_limiter( model_list=model_list, provider_budget_config=self.provider_budget_config @@ -1389,6 +1399,40 @@ class Router: self.optional_callbacks.append(ec_callback) litellm.logging_callback_manager.add_litellm_callback(ec_callback) + # --------------------------------------------------------------------- + # Prompt prefix affinity + # --------------------------------------------------------------------- + if "prompt_prefix_affinity" in optional_pre_call_checks: + from litellm.router_utils.pre_call_checks.prompt_prefix_affinity_check import ( + PromptPrefixAffinityCheck, + ) + + if self.optional_callbacks is None: + self.optional_callbacks = [] + + existing_prompt_prefix_callback: Optional[PromptPrefixAffinityCheck] = None + for cb in self.optional_callbacks: + if isinstance(cb, PromptPrefixAffinityCheck): + existing_prompt_prefix_callback = cb + break + + if existing_prompt_prefix_callback is not None: + existing_prompt_prefix_callback.prefix_tokens = ( + self.prompt_prefix_affinity_tokens + ) + existing_prompt_prefix_callback.min_tokens = ( + self.prompt_prefix_affinity_min_tokens + ) + else: + prompt_prefix_callback = PromptPrefixAffinityCheck( + prefix_tokens=self.prompt_prefix_affinity_tokens, + min_tokens=self.prompt_prefix_affinity_min_tokens, + ) + self.optional_callbacks.append(prompt_prefix_callback) + litellm.logging_callback_manager.add_litellm_callback( + prompt_prefix_callback + ) + # --------------------------------------------------------------------- # Remaining optional pre-call checks # --------------------------------------------------------------------- @@ -1399,6 +1443,7 @@ class Router: "responses_api_deployment_check", "session_affinity", "encrypted_content_affinity", + "prompt_prefix_affinity", ): continue if pre_call_check == "prompt_caching": diff --git a/litellm/router_utils/pre_call_checks/prompt_prefix_affinity_check.py b/litellm/router_utils/pre_call_checks/prompt_prefix_affinity_check.py new file mode 100644 index 00000000000..5fddc8a4c8d --- /dev/null +++ b/litellm/router_utils/pre_call_checks/prompt_prefix_affinity_check.py @@ -0,0 +1,206 @@ +""" +Prompt-prefix-aware deterministic deployment affinity for the Router. + +This is a stateless optimization for upstream implicit prompt caching. It +canonicalizes the prompt-bearing parts of a request, hashes the first N tokens, +then uses rendezvous hashing to choose a stable deployment from the current +healthy deployment set. + +Unlike deployment/session affinity, this does not store a prompt -> deployment +mapping in Redis. All Router instances with the same config and deployment IDs +will make the same routing decision. +""" + +import hashlib +import json +from typing import Any, Dict, List, Optional, Tuple, cast + +from litellm._logging import verbose_router_logger +from litellm.constants import MINIMUM_PROMPT_CACHE_TOKEN_COUNT +from litellm.integrations.custom_logger import CustomLogger, Span +from litellm.types.llms.openai import AllMessageValues +from litellm.utils import encode + + +class PromptPrefixAffinityCheck(CustomLogger): + """ + Routes requests with the same canonical prompt prefix to the same deployment. + + This is intended to improve prompt-cache hit rate for providers where prompt + caching is scoped to the account/key behind a deployment. + """ + + CACHE_KEY_EXCLUDED_FIELDS = frozenset({"encrypted_content"}) + + def __init__( + self, + prefix_tokens: int = 2048, + min_tokens: int = MINIMUM_PROMPT_CACHE_TOKEN_COUNT, + ) -> None: + super().__init__() + self.prefix_tokens = prefix_tokens + self.min_tokens = min_tokens + + @classmethod + def _json_safe(cls, value: Any) -> Any: + if value is None or isinstance(value, (str, int, float, bool)): + return value + + if hasattr(value, "model_dump"): + return cls._json_safe(value.model_dump()) + + if hasattr(value, "dict"): + return cls._json_safe(value.dict()) + + if isinstance(value, dict): + return { + str(k): cls._json_safe(v) + for k, v in sorted(value.items(), key=lambda item: str(item[0])) + if str(k) not in cls.CACHE_KEY_EXCLUDED_FIELDS + } + + if isinstance(value, (list, tuple)): + return [cls._json_safe(item) for item in value] + + return str(value) + + @classmethod + def _build_canonical_prompt( + cls, + messages: Optional[List[AllMessageValues]], + request_kwargs: Dict[str, Any], + ) -> Optional[str]: + prompt_parts: List[Tuple[str, Any]] = [] + + for key in ("instructions", "tools"): + value = request_kwargs.get(key) + if value is not None: + prompt_parts.append((key, cls._json_safe(value))) + + if messages is not None: + prompt_parts.append(("messages", cls._json_safe(messages))) + + for key in ("input",): + value = request_kwargs.get(key) + if value is not None: + prompt_parts.append((key, cls._json_safe(value))) + + if not prompt_parts: + return None + + return json.dumps( + prompt_parts, + sort_keys=True, + separators=(",", ":"), + ensure_ascii=False, + ) + + def _get_prefix_hash( + self, + model: str, + canonical_prompt: str, + ) -> Optional[str]: + if self.prefix_tokens <= 0: + return None + + try: + token_ids = encode(model=model, text=canonical_prompt) + except Exception as e: + verbose_router_logger.debug( + "PromptPrefixAffinityCheck: failed to tokenize prompt for model=%s; error=%s", + model, + e, + ) + return None + + if len(token_ids) < self.min_tokens: + return None + + prefix_token_ids = token_ids[: self.prefix_tokens] + prefix_payload = json.dumps(prefix_token_ids, separators=(",", ":")) + return hashlib.sha256(prefix_payload.encode("utf-8")).hexdigest() + + @staticmethod + def _get_deployment_model_id(deployment: dict) -> Optional[str]: + model_info = deployment.get("model_info") + if not isinstance(model_info, dict): + return None + + model_id = model_info.get("id") + if model_id is None: + return None + + return str(model_id) + + def _score_deployment( + self, + prefix_hash: str, + deployment_model_id: str, + ) -> int: + payload = f"{prefix_hash}:{deployment_model_id}" + return int(hashlib.sha256(payload.encode("utf-8")).hexdigest(), 16) + + def _select_deployment( + self, + prefix_hash: str, + healthy_deployments: List[dict], + ) -> Optional[dict]: + best: Optional[Tuple[int, dict]] = None + + for deployment in healthy_deployments: + deployment_model_id = self._get_deployment_model_id(deployment) + if deployment_model_id is None: + continue + + score = self._score_deployment( + prefix_hash=prefix_hash, + deployment_model_id=deployment_model_id, + ) + if best is None or score > best[0]: + best = (score, deployment) + + return best[1] if best is not None else None + + async def async_filter_deployments( + self, + model: str, + healthy_deployments: List, + messages: Optional[List[AllMessageValues]], + request_kwargs: Optional[dict] = None, + parent_otel_span: Optional[Span] = None, + ) -> List[dict]: + typed_healthy_deployments = cast(List[dict], healthy_deployments) + + if len(typed_healthy_deployments) <= 1: + return typed_healthy_deployments + + request_kwargs = request_kwargs or {} + canonical_prompt = self._build_canonical_prompt( + messages=messages, + request_kwargs=request_kwargs, + ) + if canonical_prompt is None: + return typed_healthy_deployments + + prefix_hash = self._get_prefix_hash( + model=model, + canonical_prompt=canonical_prompt, + ) + if prefix_hash is None: + return typed_healthy_deployments + + deployment = self._select_deployment( + prefix_hash=prefix_hash, + healthy_deployments=typed_healthy_deployments, + ) + if deployment is None: + return typed_healthy_deployments + + request_kwargs["_prompt_prefix_affinity_pinned"] = True + verbose_router_logger.debug( + "PromptPrefixAffinityCheck: pinning model=%s prefix_hash=%s deployment=%s", + model, + prefix_hash[:8], + self._get_deployment_model_id(deployment), + ) + return [deployment] diff --git a/litellm/types/router.py b/litellm/types/router.py index 4257628e7cb..964c6e6041e 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -734,6 +734,7 @@ OptionalPreCallChecks = List[ "forward_client_headers_by_model_group", "enforce_model_rate_limits", "encrypted_content_affinity", + "prompt_prefix_affinity", ] ] diff --git a/tests/test_litellm/router_utils/pre_call_checks/test_prompt_prefix_affinity_check.py b/tests/test_litellm/router_utils/pre_call_checks/test_prompt_prefix_affinity_check.py new file mode 100644 index 00000000000..d7a0155c5d7 --- /dev/null +++ b/tests/test_litellm/router_utils/pre_call_checks/test_prompt_prefix_affinity_check.py @@ -0,0 +1,129 @@ +import os +import sys + +import pytest + +sys.path.insert(0, os.path.abspath("../..")) + +from litellm.router_utils.pre_call_checks.prompt_prefix_affinity_check import ( + PromptPrefixAffinityCheck, +) + + +def _deployments(): + return [ + {"model_info": {"id": "deployment-a"}}, + {"model_info": {"id": "deployment-b"}}, + {"model_info": {"id": "deployment-c"}}, + ] + + +@pytest.mark.asyncio +async def test_same_prompt_prefix_routes_to_same_deployment_across_suffixes(): + check = PromptPrefixAffinityCheck( + prefix_tokens=64, + min_tokens=0, + ) + shared_prefix = "shared context " * 300 + + first_kwargs = {"input": shared_prefix + "question A"} + second_kwargs = {"input": shared_prefix + "question B"} + + first = await check.async_filter_deployments( + model="gpt-3.5-turbo", + healthy_deployments=_deployments(), + messages=None, + request_kwargs=first_kwargs, + ) + second = await check.async_filter_deployments( + model="gpt-3.5-turbo", + healthy_deployments=list(reversed(_deployments())), + messages=None, + request_kwargs=second_kwargs, + ) + + assert len(first) == 1 + assert len(second) == 1 + assert first[0]["model_info"]["id"] == second[0]["model_info"]["id"] + assert first_kwargs["_prompt_prefix_affinity_pinned"] is True + assert second_kwargs["_prompt_prefix_affinity_pinned"] is True + + +def test_different_prompt_prefixes_get_different_prefix_hashes(): + check = PromptPrefixAffinityCheck( + prefix_tokens=64, + min_tokens=0, + ) + + first_prompt = check._build_canonical_prompt( + messages=None, + request_kwargs={"input": "alpha " * 300}, + ) + second_prompt = check._build_canonical_prompt( + messages=None, + request_kwargs={"input": "beta " * 300}, + ) + + assert first_prompt is not None + assert second_prompt is not None + assert check._get_prefix_hash( + "gpt-3.5-turbo", first_prompt + ) != check._get_prefix_hash("gpt-3.5-turbo", second_prompt) + + +def test_encrypted_content_is_excluded_from_canonical_prompt_hash(): + check = PromptPrefixAffinityCheck( + prefix_tokens=64, + min_tokens=0, + ) + + first_prompt = check._build_canonical_prompt( + messages=None, + request_kwargs={ + "input": [ + { + "type": "reasoning", + "encrypted_content": "encrypted-content-a", + }, + {"role": "user", "content": "shared context " * 300}, + ] + }, + ) + second_prompt = check._build_canonical_prompt( + messages=None, + request_kwargs={ + "input": [ + { + "type": "reasoning", + "encrypted_content": "encrypted-content-b", + }, + {"role": "user", "content": "shared context " * 300}, + ] + }, + ) + + assert first_prompt is not None + assert second_prompt is not None + assert check._get_prefix_hash( + "gpt-3.5-turbo", first_prompt + ) == check._get_prefix_hash("gpt-3.5-turbo", second_prompt) + + +@pytest.mark.asyncio +async def test_prompt_prefix_affinity_does_not_filter_below_min_tokens(): + check = PromptPrefixAffinityCheck( + prefix_tokens=64, + min_tokens=10_000, + ) + deployments = _deployments() + request_kwargs = {"input": "short prompt"} + + result = await check.async_filter_deployments( + model="gpt-3.5-turbo", + healthy_deployments=deployments, + messages=None, + request_kwargs=request_kwargs, + ) + + assert result == deployments + assert "_prompt_prefix_affinity_pinned" not in request_kwargs