[Feat] Prompt Management - Allow specifying just prompt_id in a request to a model (#16834)

* test_dotprompt_auto_detection_with_model_only

* fix _auto_detect_prompt_management_logger

* test_dotprompt_auto_detection_with_model_only
This commit is contained in:
Ishaan Jaff 2025-11-19 10:20:58 -08:00 • committed by GitHub
parent 5f94b372f8
commit 1f8fe007a1
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 146 additions and 4 deletions

View file

@ -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

View file

@ -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