mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
[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:
parent
5f94b372f8
commit
1f8fe007a1
2 changed files with 146 additions and 4 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue