diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 829c1c9ca07..2a4a36d2c5e 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -97,6 +97,7 @@ from litellm.types.utils import ( CostBreakdown, CostResponseTypes, CustomPricingLiteLLMParams, + DYNAMIC_PROMPT_MANAGEMENT_PARAMS, DynamicPromptManagementParamLiteral, EmbeddingResponse, GuardrailStatus, @@ -684,7 +685,7 @@ class Logging(LiteLLMLoggingBaseClass): eg. AnthropicCacheControlHook and BedrockKnowledgeBaseHook both don't require a `prompt_id` to be passed in, they are triggered by dynamic params """ for param in non_default_params: - if param in DynamicPromptManagementParamLiteral.list_all_params(): + if param in DYNAMIC_PROMPT_MANAGEMENT_PARAMS: return True ############################################################################# diff --git a/litellm/types/utils.py b/litellm/types/utils.py index ed29d49fc29..fcfeec67e4a 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -3601,6 +3601,11 @@ class DynamicPromptManagementParamLiteral(str, Enum): return [param.value for param in cls] +DYNAMIC_PROMPT_MANAGEMENT_PARAMS: frozenset[str] = frozenset( + param.value for param in DynamicPromptManagementParamLiteral +) + + class CallbacksByType(TypedDict): success: List[str] failure: List[str] diff --git a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py index 1764d9c609f..6b568e22594 100644 --- a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py +++ b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py @@ -36,6 +36,49 @@ def test_get_masked_api_base(logging_obj): assert type(masked_api_base) == str +@pytest.mark.parametrize( + "non_default_params", + [ + {"cache_control_injection_points": [{"location": "message", "index": -1}]}, + {"knowledge_bases": ["kb-1"]}, + {"vector_store_ids": ["vs-1"]}, + ], +) +def test_should_run_prompt_management_hooks_without_prompt_id_for_dynamic_params( + logging_obj, non_default_params +): + assert ( + logging_obj.should_run_prompt_management_hooks( + non_default_params=non_default_params, + prompt_id=None, + tools=None, + ) + is True + ) + + +@pytest.mark.parametrize( + "non_default_params", + [ + {"temperature": 0.5}, + {"max_tokens": 100}, + {"temperature": 0.7, "top_p": 0.9, "max_tokens": 256}, + {}, + ], +) +def test_should_run_prompt_management_hooks_false_for_non_dynamic_params( + logging_obj, non_default_params +): + assert ( + logging_obj.should_run_prompt_management_hooks( + non_default_params=non_default_params, + prompt_id=None, + tools=None, + ) + is False + ) + + def test_sentry_sample_rate(): existing_sample_rate = os.getenv("SENTRY_API_SAMPLE_RATE") try: diff --git a/tests/test_litellm/types/test_types_utils.py b/tests/test_litellm/types/test_types_utils.py index c146847f391..ff71c4e325c 100644 --- a/tests/test_litellm/types/test_types_utils.py +++ b/tests/test_litellm/types/test_types_utils.py @@ -19,6 +19,27 @@ def test_hidden_params_response_ms(): assert hidden_params_dict.get("_response_ms") == 100 +def test_dynamic_prompt_management_list_all_params_preserves_enum_order(): + from litellm.types.utils import DynamicPromptManagementParamLiteral + + assert DynamicPromptManagementParamLiteral.list_all_params() == [ + "cache_control_injection_points", + "knowledge_bases", + "vector_store_ids", + ] + + +def test_dynamic_prompt_management_params_frozenset(): + from litellm.types.utils import DYNAMIC_PROMPT_MANAGEMENT_PARAMS + + assert "cache_control_injection_points" in DYNAMIC_PROMPT_MANAGEMENT_PARAMS + assert "knowledge_bases" in DYNAMIC_PROMPT_MANAGEMENT_PARAMS + assert "vector_store_ids" in DYNAMIC_PROMPT_MANAGEMENT_PARAMS + assert "temperature" not in DYNAMIC_PROMPT_MANAGEMENT_PARAMS + assert "max_tokens" not in DYNAMIC_PROMPT_MANAGEMENT_PARAMS + assert isinstance(DYNAMIC_PROMPT_MANAGEMENT_PARAMS, frozenset) + + def test_chat_completion_delta_tool_call(): from litellm.types.utils import ChatCompletionDeltaToolCall, Function