diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 6bad7ee29e2..e5c52ce48ae 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -585,7 +585,10 @@ class Logging(LiteLLMLoggingBaseClass): custom_logger = ( prompt_management_logger or self.get_custom_logger_for_prompt_management( - model=model, non_default_params=non_default_params + model=model, + non_default_params=non_default_params, + prompt_id=prompt_id, + dynamic_callback_params=self.standard_callback_dynamic_params, ) ) @@ -622,7 +625,11 @@ class Logging(LiteLLMLoggingBaseClass): custom_logger = ( prompt_management_logger or self.get_custom_logger_for_prompt_management( - model=model, tools=tools, non_default_params=non_default_params + model=model, + tools=tools, + non_default_params=non_default_params, + prompt_id=prompt_id, + dynamic_callback_params=self.standard_callback_dynamic_params, ) ) @@ -646,19 +653,69 @@ class Logging(LiteLLMLoggingBaseClass): self.messages = messages return model, messages, non_default_params + def _auto_detect_prompt_management_logger( + self, + prompt_id: str, + dynamic_callback_params: StandardCallbackDynamicParams, + ) -> Optional[CustomLogger]: + """ + Auto-detect which prompt management system owns the given prompt_id. + + This allows a user to just pass prompt_id in the completion call and it will be auto-detected which system owns this prompt. + + Args: + prompt_id: The prompt ID to check + dynamic_callback_params: Dynamic callback parameters for should_run_prompt_management checks + + Returns: + A CustomLogger instance if a matching prompt management system is found, None otherwise + """ + prompt_management_loggers = ( + litellm.logging_callback_manager.get_custom_loggers_for_type( + callback_type=CustomPromptManagement + ) + ) + + for logger in prompt_management_loggers: + if isinstance(logger, CustomPromptManagement): + try: + if logger.should_run_prompt_management( + prompt_id=prompt_id, + dynamic_callback_params=dynamic_callback_params, + ): + self.model_call_details["prompt_integration"] = ( + logger.__class__.__name__ + ) + return logger + except Exception: + # If check fails, continue to next logger + continue + + return None + def get_custom_logger_for_prompt_management( - self, model: str, non_default_params: Dict, tools: Optional[List[Dict]] = None + self, + model: str, + non_default_params: Dict, + tools: Optional[List[Dict]] = None, + prompt_id: Optional[str] = None, + dynamic_callback_params: Optional[StandardCallbackDynamicParams] = None, ) -> Optional[CustomLogger]: """ Get a custom logger for prompt management based on model name or available callbacks. Args: model: The model name to check for prompt management integration + non_default_params: Non-default parameters passed to the completion call + tools: Optional tools passed to the completion call + prompt_id: Optional prompt ID to auto-detect which system owns this prompt + dynamic_callback_params: Dynamic callback parameters for should_run_prompt_management checks Returns: A CustomLogger instance if one is found, None otherwise """ # First check if model starts with a known custom logger compatible callback + # This takes precedence for backward compatibility for callback_name in litellm._known_custom_logger_compatible_callbacks: if model.startswith(callback_name): custom_logger = _init_custom_logger_compatible_class( @@ -670,7 +727,16 @@ class Logging(LiteLLMLoggingBaseClass): self.model_call_details["prompt_integration"] = model.split("/")[0] return custom_logger - # Then check for any registered CustomPromptManagement loggers + # If prompt_id is provided, try to auto-detect which system has this prompt + if prompt_id and dynamic_callback_params is not None: + auto_detected_logger = self._auto_detect_prompt_management_logger( + prompt_id=prompt_id, + dynamic_callback_params=dynamic_callback_params, + ) + if auto_detected_logger is not None: + return auto_detected_logger + + # Then check for any registered CustomPromptManagement loggers (fallback) prompt_management_loggers = ( litellm.logging_callback_manager.get_custom_loggers_for_type( callback_type=CustomPromptManagement diff --git a/tests/test_litellm/integrations/dotprompt/test_prompt_manager.py b/tests/test_litellm/integrations/dotprompt/test_prompt_manager.py index c19503641e1..77e3e0d4a45 100644 --- a/tests/test_litellm/integrations/dotprompt/test_prompt_manager.py +++ b/tests/test_litellm/integrations/dotprompt/test_prompt_manager.py @@ -521,3 +521,79 @@ def test_prompt_main(): """ # TODO: Implement once PromptManager is integrated with litellm completion pass + + +@pytest.mark.asyncio +async def test_dotprompt_auto_detection_with_model_only(): + """ + Test that dotprompt prompts can be auto-detected when passing model="gpt-4" and prompt_id, + without needing to specify model="dotprompt/gpt-4". + """ + from litellm.integrations.dotprompt import DotpromptManager + + prompt_dir = Path(__file__).parent + dotprompt_manager = DotpromptManager(prompt_directory=str(prompt_dir)) + + # Register the dotprompt manager in callbacks + original_callbacks = litellm.callbacks.copy() + litellm.callbacks = [dotprompt_manager] + + try: + # Mock the HTTP handler to avoid actual API calls + with patch("litellm.llms.custom_httpx.llm_http_handler.AsyncHTTPHandler.post") as mock_post: + mock_response_data = litellm.ModelResponse( + choices=[ + litellm.Choices( + message=litellm.Message(content="Hello!"), + index=0, + finish_reason="stop", + ) + ] + ).model_dump() + + # Create a proper mock response + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.text = json.dumps(mock_response_data) + mock_response.headers = {"Content-Type": "application/json"} + mock_response.json.return_value = mock_response_data + + mock_post.return_value = mock_response + + # Call with model="gpt-4" (no "dotprompt/" prefix) and prompt_id + await litellm.acompletion( + model="gpt-4", + prompt_id="chat_prompt", + prompt_variables={"user_message": "Hello world"}, + messages=[{"role": "user", "content": "This will be ignored"}], + ) + + mock_post.assert_called_once() + + # Get request body from the call (it's passed as 'data' parameter as JSON string) + data_str = mock_post.call_args.kwargs.get("data", "{}") + request_body = json.loads(data_str) + + print(f"Request body: {json.dumps(request_body, indent=2)}") + + # Verify the prompt was auto-detected and used + # The chat_prompt.prompt has metadata: model: gpt-4, temperature: 0.7, max_tokens: 150 + assert request_body["model"] == "gpt-4" + + # Note: OpenAI API might strip out temperature/max_tokens if they're not in the request + # The key test is that the messages were transformed + + # Verify the messages were transformed using the prompt template + # chat_prompt template: "User: {{user_message}}" + messages = request_body["messages"] + assert len(messages) >= 1 + + # The first message should be from the prompt template with the variable substituted + # Template is: "User: {{user_message}}" with user_message="Hello world" + first_message_content = messages[0]["content"] + print(f"First message content: {first_message_content}") + assert "Hello world" in first_message_content + + finally: + # Restore original callbacks + litellm.callbacks = original_callbacks