diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index d8e70d3b125..79e67ef36f9 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -1135,25 +1135,41 @@ class Logging(LiteLLMLoggingBaseClass): dynamic_callback_params=self.standard_callback_dynamic_params, ) - if custom_logger: - breakpoints_before: Final = AnthropicCacheControlHook.count_request_cache_breakpoints(messages) - ( - model, - messages, - non_default_params, - ) = await custom_logger.async_get_chat_completion_prompt( - model=model, - messages=messages, - non_default_params=non_default_params or {}, - prompt_id=prompt_id, - prompt_spec=prompt_spec, - prompt_variables=prompt_variables, - dynamic_callback_params=self.standard_callback_dynamic_params, - litellm_logging_obj=self, - tools=tools, - prompt_label=prompt_label, - prompt_version=prompt_version, + # The Anthropic cache-control hook and the vector store hook both claim the single + # prompt-management slot; a request that seeds cache_control_injection_points (e.g. via + # `enable_anthropic_prompt_caching`) would otherwise starve knowledge-base retrieval, so + # run the vector store hook first when both apply to the same request. + vector_store_logger: Final = ( + self._get_vector_store_pre_call_hook() + if isinstance(custom_logger, AnthropicCacheControlHook) + and litellm.vector_store_registry is not None + and litellm.vector_store_registry.get_vector_store_ids_to_run( + non_default_params=non_default_params, tools=tools ) + else None + ) + prompt_loggers: Final = tuple(logger for logger in (vector_store_logger, custom_logger) if logger is not None) + + if prompt_loggers: + breakpoints_before: Final = AnthropicCacheControlHook.count_request_cache_breakpoints(messages) + for prompt_logger in prompt_loggers: + ( + model, + messages, + non_default_params, + ) = await prompt_logger.async_get_chat_completion_prompt( + model=model, + messages=messages, + non_default_params=non_default_params or {}, + prompt_id=prompt_id, + prompt_spec=prompt_spec, + prompt_variables=prompt_variables, + dynamic_callback_params=self.standard_callback_dynamic_params, + litellm_logging_obj=self, + tools=tools, + prompt_label=prompt_label, + prompt_version=prompt_version, + ) if request_kwargs is not None: AnthropicCacheControlHook.record_gateway_injection( request_kwargs, @@ -1289,19 +1305,26 @@ class Logging(LiteLLMLoggingBaseClass): # Vector Store / Knowledge Base hooks ######################################################### if litellm.vector_store_registry is not None: - vector_store_custom_logger: Final = _init_custom_logger_compatible_class( - logging_integration="vector_store_pre_call_hook", - internal_usage_cache=None, - llm_router=None, - ) - self.model_call_details["prompt_integration"] = vector_store_custom_logger.__class__.__name__ - # Add to global callbacks so post-call hooks are invoked - if vector_store_custom_logger and vector_store_custom_logger not in litellm.callbacks: - litellm.logging_callback_manager.add_litellm_callback(vector_store_custom_logger) - return vector_store_custom_logger + vector_store_custom_logger: Final = self._get_vector_store_pre_call_hook() + if vector_store_custom_logger is not None: + return vector_store_custom_logger return None + def _get_vector_store_pre_call_hook(self) -> CustomLogger | None: + vector_store_custom_logger: Final = _init_custom_logger_compatible_class( + logging_integration="vector_store_pre_call_hook", + internal_usage_cache=None, + llm_router=None, + ) + if vector_store_custom_logger is None: + return None + self.model_call_details["prompt_integration"] = vector_store_custom_logger.__class__.__name__ + # Add to global callbacks so post-call hooks are invoked + if vector_store_custom_logger not in litellm.callbacks: + litellm.logging_callback_manager.add_litellm_callback(vector_store_custom_logger) + return vector_store_custom_logger + def get_custom_logger_for_anthropic_cache_control_hook(self, non_default_params: dict) -> CustomLogger | None: if non_default_params.get("cache_control_injection_points", None): custom_logger: Final = _init_custom_logger_compatible_class( diff --git a/tests/unit/litellm_core_utils/test_litellm_logging.py b/tests/unit/litellm_core_utils/test_litellm_logging.py index 4b25da2ff79..20ae22b2eac 100644 --- a/tests/unit/litellm_core_utils/test_litellm_logging.py +++ b/tests/unit/litellm_core_utils/test_litellm_logging.py @@ -7267,6 +7267,98 @@ def test_prompt_hooks_skip_prompt_managers_when_no_prompt_id(logging_obj, tmp_pa litellm.logging_callback_manager.remove_callback_from_list_by_object(litellm.callbacks, hook) +@pytest.mark.asyncio +async def test_prompt_hooks_compose_vector_search_with_anthropic_cache_control(logging_obj, monkeypatch): + """ + Regression for https://github.com/BerriAI/litellm/issues/40908: with + `enable_anthropic_prompt_caching` on, `get_custom_logger_for_prompt_management` selects + AnthropicCacheControlHook over VectorStorePreCallHook (first-match-wins), so a model + configured with `vector_store_ids` never retrieved from its knowledge base and + `vector_store_ids` leaked into the provider request. Both hooks must now run. + """ + from litellm.integrations.anthropic_cache_control_hook import AnthropicCacheControlHook + from litellm.integrations.vector_store_integrations.vector_store_pre_call_hook import ( + VectorStorePreCallHook, + ) + from litellm.litellm_core_utils import litellm_logging as logging_module + from litellm.types.vector_stores import ( + LiteLLM_ManagedVectorStore, + VectorStoreResultContent, + VectorStoreSearchResponse, + VectorStoreSearchResult, + ) + from litellm.vector_stores.vector_store_registry import VectorStoreRegistry + + model = "bedrock/anthropic.claude-3-5-haiku-20241022-v1:0" + messages = [{"role": "user", "content": "What does the handbook say about refunds?"}] + params = {"custom_llm_provider": "bedrock", "vector_store_ids": ["vs_123"]} + + monkeypatch.setattr( + litellm, + "vector_store_registry", + VectorStoreRegistry( + vector_stores=[LiteLLM_ManagedVectorStore(vector_store_id="vs_123", custom_llm_provider="bedrock")] + ), + ) + monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True) + + AnthropicCacheControlHook.maybe_seed_default_injection_points( + non_default_params=params, messages=messages, model=model, custom_llm_provider="bedrock" + ) + assert "cache_control_injection_points" in params + + selected = logging_obj.get_custom_logger_for_prompt_management( + model=model, non_default_params=params, tools=None, prompt_id=None + ) + assert isinstance(selected, AnthropicCacheControlHook) + + class _RecordingRouter: + def __init__(self) -> None: + self.calls: list[dict[str, object]] = [] + + async def avector_store_search(self, **kwargs: object) -> VectorStoreSearchResponse: + self.calls.append(kwargs) + return VectorStoreSearchResponse( + object="vector_store.search_results.page", + search_query="What does the handbook say about refunds?", + data=[ + VectorStoreSearchResult( + score=1.0, + content=[VectorStoreResultContent(text="Refunds take seven days", type="text")], + ) + ], + ) + + router = _RecordingRouter() + runtime = MagicMock() + runtime.llm_router.return_value = router + runtime.prisma_client.return_value = None + vector_store_hook = VectorStorePreCallHook(proxy_runtime=runtime) + previous_loggers = tuple(logging_module._in_memory_loggers) + logging_module._in_memory_loggers.clear() + logging_module._in_memory_loggers.append(vector_store_hook) + try: + _, result_messages, remaining_params = await logging_obj.async_get_chat_completion_prompt( + model=model, + messages=messages, + non_default_params=params, + prompt_variables=None, + ) + + assert len(router.calls) == 1 + assert result_messages[0]["content"] == "Context:\n\nRefunds take seven days\n\n" + assert result_messages[1]["content"] == "What does the handbook say about refunds?" + assert result_messages[1]["cache_control"] == {"type": "ephemeral"} + assert "vector_store_ids" not in remaining_params + assert "cache_control_injection_points" not in remaining_params + finally: + litellm.logging_callback_manager.remove_callback_from_list_by_object( + litellm.callbacks, vector_store_hook, require_self=False + ) + logging_module._in_memory_loggers.clear() + logging_module._in_memory_loggers.extend(previous_loggers) + + def test_newrelic_dispatch_prefers_otel_v2_when_flag_on(monkeypatch): """With LITELLM_OTEL_V2 on and operator credentials present, the "newrelic" callback builds the OTel v2 logger (per-team credential routing); with the