diff --git a/docs/my-website/docs/proxy/custom_prompt_management.md b/docs/my-website/docs/proxy/custom_prompt_management.md new file mode 100644 index 0000000000..72a7333276 --- /dev/null +++ b/docs/my-website/docs/proxy/custom_prompt_management.md @@ -0,0 +1,194 @@ +import Image from '@theme/IdealImage'; +import Tabs from '@theme/Tabs'; +import TabItem from '@theme/TabItem'; + +# Custom Prompt Management + +Connect LiteLLM to your prompt management system with custom hooks. + +## Overview + + + + + +## How it works + +## Quick Start + +### 1. Create Your Custom Prompt Manager + +Create a class that inherits from `CustomPromptManagement` to handle prompt retrieval and formatting: + +**Example Implementation** + +Create a new file called `custom_prompt.py` and add this code. The key method here is `get_chat_completion_prompt` you can implement custom logic to retrieve and format prompts based on the `prompt_id` and `prompt_variables`. + +```python +from typing import List, Tuple, Optional +from litellm.integrations.custom_prompt_management import CustomPromptManagement +from litellm.types.llms.openai import AllMessageValues +from litellm.types.utils import StandardCallbackDynamicParams + +class MyCustomPromptManagement(CustomPromptManagement): + def get_chat_completion_prompt( + self, + model: str, + messages: List[AllMessageValues], + non_default_params: dict, + prompt_id: str, + prompt_variables: Optional[dict], + dynamic_callback_params: StandardCallbackDynamicParams, + ) -> Tuple[str, List[AllMessageValues], dict]: + """ + Retrieve and format prompts based on prompt_id. + + Returns: + - model: The model to use + - messages: The formatted messages + - non_default_params: Optional parameters like temperature + """ + # Example matching the diagram: Add system message for prompt_id "1234" + if prompt_id == "1234": + # Prepend system message while preserving existing messages + new_messages = [ + {"role": "system", "content": "Be a good Bot!"}, + ] + messages + return model, new_messages, non_default_params + + # Default: Return original messages if no prompt_id match + return model, messages, non_default_params + +prompt_management = MyCustomPromptManagement() +``` + +### 2. Configure Your Prompt Manager in LiteLLM `config.yaml` + +```yaml +model_list: + - model_name: gpt-4 + litellm_params: + model: openai/gpt-4 + api_key: os.environ/OPENAI_API_KEY + +litellm_settings: + callbacks: custom_prompt.prompt_management # sets litellm.callbacks = [prompt_management] +``` + +### 3. Start LiteLLM Gateway + + + + +Mount your `custom_logger.py` on the LiteLLM Docker container. + +```shell +docker run -d \ + -p 4000:4000 \ + -e OPENAI_API_KEY=$OPENAI_API_KEY \ + --name my-app \ + -v $(pwd)/my_config.yaml:/app/config.yaml \ + -v $(pwd)/custom_logger.py:/app/custom_logger.py \ + my-app:latest \ + --config /app/config.yaml \ + --port 4000 \ + --detailed_debug \ +``` + + + + + +```shell +litellm --config config.yaml --detailed_debug +``` + + + + +### 4. Test Your Custom Prompt Manager + +When you pass `prompt_id="1234"`, the custom prompt manager will add a system message "Be a good Bot!" to your conversation: + + + + +```python +from openai import OpenAI + +client = OpenAI( + api_key="sk-1234", + base_url="http://0.0.0.0:4000" +) + +response = client.chat.completions.create( + model="gemini-1.5-pro", + messages=[{"role": "user", "content": "hi"}], + prompt_id="1234" +) + +print(response.choices[0].message.content) +``` + + + + +```python +from langchain.chat_models import ChatOpenAI +from langchain.schema import HumanMessage + +chat = ChatOpenAI( + model="gpt-4", + openai_api_key="sk-1234", + openai_api_base="http://0.0.0.0:4000", + extra_body={ + "prompt_id": "1234" + } +) + +messages = [] +response = chat(messages) + +print(response.content) +``` + + + + +```shell +curl -X POST http://0.0.0.0:4000/v1/chat/completions \ +-H "Content-Type: application/json" \ +-H "Authorization: Bearer sk-1234" \ +-d '{ + "model": "gemini-1.5-pro", + "messages": [{"role": "user", "content": "hi"}], + "prompt_id": "1234" +}' +``` + + + +The request will be transformed from: +```json +{ + "model": "gemini-1.5-pro", + "messages": [{"role": "user", "content": "hi"}], + "prompt_id": "1234" +} +``` + +To: +```json +{ + "model": "gemini-1.5-pro", + "messages": [ + {"role": "system", "content": "Be a good Bot!"}, + {"role": "user", "content": "hi"} + ] +} +``` + + diff --git a/docs/my-website/docs/proxy/prompt_management.md b/docs/my-website/docs/proxy/prompt_management.md index 980043f455..c09231dd59 100644 --- a/docs/my-website/docs/proxy/prompt_management.md +++ b/docs/my-website/docs/proxy/prompt_management.md @@ -2,7 +2,7 @@ import Image from '@theme/IdealImage'; import Tabs from '@theme/Tabs'; import TabItem from '@theme/TabItem'; -# [BETA] Prompt Management +# Prompt Management :::info @@ -12,9 +12,10 @@ This feature is currently in beta, and might change unexpectedly. We expect this Run experiments or change the specific model (e.g. from gpt-4o to gpt4o-mini finetune) from your prompt management tool (e.g. Langfuse) instead of making changes in the application. -Supported Integrations: -- [Langfuse](https://langfuse.com/docs/prompts/get-started) -- [Humanloop](../observability/humanloop) +| Supported Integrations | Link | +|------------------------|------| +| Langfuse | [Get Started](https://langfuse.com/docs/prompts/get-started) | +| Humanloop | [Get Started](../observability/humanloop) | ## Quick Start diff --git a/docs/my-website/img/custom_prompt_management.png b/docs/my-website/img/custom_prompt_management.png new file mode 100644 index 0000000000..2c96e0d116 Binary files /dev/null and b/docs/my-website/img/custom_prompt_management.png differ diff --git a/docs/my-website/sidebars.js b/docs/my-website/sidebars.js index 47d69e5d3f..a92ceb3b05 100644 --- a/docs/my-website/sidebars.js +++ b/docs/my-website/sidebars.js @@ -365,8 +365,12 @@ const sidebars = { ], }, { - type: "doc", - id: "proxy/prompt_management" + type: "category", + label: "[Beta] Prompt Management", + items: [ + "proxy/prompt_management", + "proxy/custom_prompt_management" + ], }, { type: "category", diff --git a/litellm/integrations/custom_prompt_management.py b/litellm/integrations/custom_prompt_management.py new file mode 100644 index 0000000000..5b34ef0c34 --- /dev/null +++ b/litellm/integrations/custom_prompt_management.py @@ -0,0 +1,49 @@ +from typing import List, Optional, Tuple + +from litellm.integrations.custom_logger import CustomLogger +from litellm.integrations.prompt_management_base import ( + PromptManagementBase, + PromptManagementClient, +) +from litellm.types.llms.openai import AllMessageValues +from litellm.types.utils import StandardCallbackDynamicParams + + +class CustomPromptManagement(CustomLogger, PromptManagementBase): + def get_chat_completion_prompt( + self, + model: str, + messages: List[AllMessageValues], + non_default_params: dict, + prompt_id: str, + prompt_variables: Optional[dict], + dynamic_callback_params: StandardCallbackDynamicParams, + ) -> Tuple[str, List[AllMessageValues], dict]: + """ + Returns: + - model: str - the model to use (can be pulled from prompt management tool) + - messages: List[AllMessageValues] - the messages to use (can be pulled from prompt management tool) + - non_default_params: dict - update with any optional params (e.g. temperature, max_tokens, etc.) to use (can be pulled from prompt management tool) + """ + return model, messages, non_default_params + + @property + def integration_name(self) -> str: + return "custom-prompt-management" + + def should_run_prompt_management( + self, + prompt_id: str, + dynamic_callback_params: StandardCallbackDynamicParams, + ) -> bool: + return True + + def _compile_prompt_helper( + self, + prompt_id: str, + prompt_variables: Optional[dict], + dynamic_callback_params: StandardCallbackDynamicParams, + ) -> PromptManagementClient: + raise NotImplementedError( + "Custom prompt management does not support compile prompt helper" + ) diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index f5afe69c74..3e694220a5 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -81,6 +81,7 @@ from ..integrations.arize.arize_phoenix import ArizePhoenixLogger from ..integrations.athina import AthinaLogger from ..integrations.azure_storage.azure_storage import AzureBlobStorageLogger from ..integrations.braintrust_logging import BraintrustLogger +from ..integrations.custom_prompt_management import CustomPromptManagement from ..integrations.datadog.datadog import DataDogLogger from ..integrations.datadog.datadog_llm_obs import DataDogLLMObsLogger from ..integrations.dynamodb import DyanmoDBLogger @@ -429,34 +430,58 @@ class Logging(LiteLLMLoggingBaseClass): prompt_variables: Optional[dict], ) -> Tuple[str, List[AllMessageValues], dict]: - for ( - custom_logger_compatible_callback - ) in litellm._known_custom_logger_compatible_callbacks: - if model.startswith(custom_logger_compatible_callback): + custom_logger = self.get_custom_logger_for_prompt_management(model) + if custom_logger: + model, messages, non_default_params = ( + custom_logger.get_chat_completion_prompt( + model=model, + messages=messages, + non_default_params=non_default_params, + prompt_id=prompt_id, + prompt_variables=prompt_variables, + dynamic_callback_params=self.standard_callback_dynamic_params, + ) + ) + self.messages = messages + return model, messages, non_default_params + + def get_custom_logger_for_prompt_management( + self, model: str + ) -> 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 + + Returns: + A CustomLogger instance if one is found, None otherwise + """ + # First check if model starts with a known custom logger compatible callback + for callback_name in litellm._known_custom_logger_compatible_callbacks: + if model.startswith(callback_name): custom_logger = _init_custom_logger_compatible_class( - logging_integration=custom_logger_compatible_callback, + logging_integration=callback_name, internal_usage_cache=None, llm_router=None, ) + if custom_logger is not None: + self.model_call_details["prompt_integration"] = model.split("/")[0] + return custom_logger - if custom_logger is None: - continue - old_name = model + # Then check for any registered CustomPromptManagement loggers + prompt_management_loggers = ( + litellm.logging_callback_manager.get_custom_loggers_for_type( + callback_type=CustomPromptManagement + ) + ) - model, messages, non_default_params = ( - custom_logger.get_chat_completion_prompt( - model=model, - messages=messages, - non_default_params=non_default_params, - prompt_id=prompt_id, - prompt_variables=prompt_variables, - dynamic_callback_params=self.standard_callback_dynamic_params, - ) - ) - self.model_call_details["prompt_integration"] = old_name.split("/")[0] - self.messages = messages + if prompt_management_loggers: + logger = prompt_management_loggers[0] + self.model_call_details["prompt_integration"] = logger.__class__.__name__ + return logger - return model, messages, non_default_params + return None def _get_raw_request_body(self, data: Optional[Union[dict, str]]) -> dict: if data is None: diff --git a/litellm/litellm_core_utils/logging_callback_manager.py b/litellm/litellm_core_utils/logging_callback_manager.py index a20e826c43..c57a2401b7 100644 --- a/litellm/litellm_core_utils/logging_callback_manager.py +++ b/litellm/litellm_core_utils/logging_callback_manager.py @@ -1,4 +1,4 @@ -from typing import Callable, List, Set, Union +from typing import Callable, List, Set, Type, Union import litellm from litellm._logging import verbose_logger @@ -86,21 +86,20 @@ class LoggingCallbackManager: callback=callback, parent_list=litellm._async_failure_callback ) - def remove_callback_from_list_by_object( - self, callback_list, obj - ): + def remove_callback_from_list_by_object(self, callback_list, obj): """ Remove callbacks that are methods of a particular object (e.g., router cleanup) """ - if not isinstance(callback_list, list): # Not list -> do nothing + if not isinstance(callback_list, list): # Not list -> do nothing return - - remove_list=[c for c in callback_list if hasattr(c, '__self__') and c.__self__ == obj] + + remove_list = [ + c for c in callback_list if hasattr(c, "__self__") and c.__self__ == obj + ] for c in remove_list: callback_list.remove(c) - def _add_string_callback_to_list( self, callback: str, parent_list: List[Union[CustomLogger, Callable, str]] ): @@ -254,3 +253,11 @@ class LoggingCallbackManager: ): matched_callbacks.add(callback) return matched_callbacks + + def get_custom_loggers_for_type( + self, callback_type: Type[CustomLogger] + ) -> List[CustomLogger]: + """ + Get all custom loggers that are instances of the given class type + """ + return [c for c in self._get_all_callbacks() if isinstance(c, callback_type)] diff --git a/litellm/proxy/custom_prompt_management.py b/litellm/proxy/custom_prompt_management.py new file mode 100644 index 0000000000..7f320ac00a --- /dev/null +++ b/litellm/proxy/custom_prompt_management.py @@ -0,0 +1,36 @@ +from typing import List, Optional, Tuple + +from litellm._logging import verbose_logger +from litellm.integrations.custom_prompt_management import CustomPromptManagement +from litellm.types.llms.openai import AllMessageValues +from litellm.types.utils import StandardCallbackDynamicParams + + +class X42PromptManagement(CustomPromptManagement): + def get_chat_completion_prompt( + self, + model: str, + messages: List[AllMessageValues], + non_default_params: dict, + prompt_id: str, + prompt_variables: Optional[dict], + dynamic_callback_params: StandardCallbackDynamicParams, + ) -> Tuple[str, List[AllMessageValues], dict]: + """ + Returns: + - model: str - the model to use (can be pulled from prompt management tool) + - messages: List[AllMessageValues] - the messages to use (can be pulled from prompt management tool) + - non_default_params: dict - update with any optional params (e.g. temperature, max_tokens, etc.) to use (can be pulled from prompt management tool) + """ + verbose_logger.debug( + f"in async get chat completion prompt. Prompt ID: {prompt_id}, Prompt Variables: {prompt_variables}, Dynamic Callback Params: {dynamic_callback_params}" + ) + + return model, messages, non_default_params + + @property + def integration_name(self) -> str: + return "x42-prompt-management" + + +x42_prompt_management = X42PromptManagement() diff --git a/litellm/proxy/proxy_config.yaml b/litellm/proxy/proxy_config.yaml index 6f37f0e140..0877a02a74 100644 --- a/litellm/proxy/proxy_config.yaml +++ b/litellm/proxy/proxy_config.yaml @@ -7,3 +7,5 @@ model_list: api_key: os.environ/AZURE_API_KEY +litellm_settings: + callbacks: ["custom_prompt_management.x42_prompt_management"] diff --git a/tests/litellm/integrations/test_custom_prompt_management.py b/tests/litellm/integrations/test_custom_prompt_management.py new file mode 100644 index 0000000000..00dc6fe96f --- /dev/null +++ b/tests/litellm/integrations/test_custom_prompt_management.py @@ -0,0 +1,132 @@ +import datetime +import json +import os +import sys +import unittest +from typing import List, Optional, Tuple +from unittest.mock import ANY, MagicMock, Mock, patch + +import httpx +import pytest + +sys.path.insert( + 0, os.path.abspath("../..") +) # Adds the parent directory to the system-path +import litellm +from litellm.integrations.custom_prompt_management import CustomPromptManagement +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler +from litellm.types.llms.openai import AllMessageValues +from litellm.types.utils import StandardCallbackDynamicParams + + +class TestCustomPromptManagement(CustomPromptManagement): + def get_chat_completion_prompt( + self, + model: str, + messages: List[AllMessageValues], + non_default_params: dict, + prompt_id: str, + prompt_variables: Optional[dict], + dynamic_callback_params: StandardCallbackDynamicParams, + ) -> Tuple[str, List[AllMessageValues], dict]: + print( + "TestCustomPromptManagement: running get_chat_completion_prompt for prompt_id: ", + prompt_id, + ) + if prompt_id == "test_prompt_id": + messages = [ + {"role": "user", "content": "This is the prompt for test_prompt_id"}, + ] + return model, messages, non_default_params + elif prompt_id == "prompt_with_variables": + content = "Hello, {name}! You are {age} years old and live in {city}." + content_with_variables = content.format(**(prompt_variables or {})) + messages = [ + {"role": "user", "content": content_with_variables}, + ] + return model, messages, non_default_params + else: + return model, messages, non_default_params + + +@pytest.mark.asyncio +async def test_custom_prompt_management_with_prompt_id(): + custom_prompt_management = TestCustomPromptManagement() + litellm.callbacks = [custom_prompt_management] + + # Mock AsyncHTTPHandler.post method + client = AsyncHTTPHandler() + with patch.object(client, "post", return_value=MagicMock()) as mock_post: + await litellm.acompletion( + model="anthropic/claude-3-5-sonnet", + messages=[{"role": "user", "content": "Hello, how are you?"}], + client=client, + prompt_id="test_prompt_id", + ) + + mock_post.assert_called_once() + print(mock_post.call_args.kwargs) + request_body = mock_post.call_args.kwargs["json"] + print("request_body: ", json.dumps(request_body, indent=4)) + + assert request_body["model"] == "claude-3-5-sonnet" + # the message gets applied to the prompt from the custom prompt management callback + assert ( + request_body["messages"][0]["content"][0]["text"] + == "This is the prompt for test_prompt_id" + ) + + +@pytest.mark.asyncio +async def test_custom_prompt_management_with_prompt_id_and_prompt_variables(): + custom_prompt_management = TestCustomPromptManagement() + litellm.callbacks = [custom_prompt_management] + + # Mock AsyncHTTPHandler.post method + client = AsyncHTTPHandler() + with patch.object(client, "post", return_value=MagicMock()) as mock_post: + await litellm.acompletion( + model="anthropic/claude-3-5-sonnet", + messages=[], + client=client, + prompt_id="prompt_with_variables", + prompt_variables={"name": "John", "age": 30, "city": "New York"}, + ) + + mock_post.assert_called_once() + print(mock_post.call_args.kwargs) + request_body = mock_post.call_args.kwargs["json"] + print("request_body: ", json.dumps(request_body, indent=4)) + + assert request_body["model"] == "claude-3-5-sonnet" + # the message gets applied to the prompt from the custom prompt management callback + assert ( + request_body["messages"][0]["content"][0]["text"] + == "Hello, John! You are 30 years old and live in New York." + ) + + +@pytest.mark.asyncio +async def test_custom_prompt_management_without_prompt_id(): + custom_prompt_management = TestCustomPromptManagement() + litellm.callbacks = [custom_prompt_management] + + # Mock AsyncHTTPHandler.post method + client = AsyncHTTPHandler() + with patch.object(client, "post", return_value=MagicMock()) as mock_post: + await litellm.acompletion( + model="anthropic/claude-3-5-sonnet", + messages=[{"role": "user", "content": "Hello, how are you?"}], + client=client, + ) + + mock_post.assert_called_once() + print(mock_post.call_args.kwargs) + request_body = mock_post.call_args.kwargs["json"] + print("request_body: ", json.dumps(request_body, indent=4)) + + assert request_body["model"] == "claude-3-5-sonnet" + # the message does not get applied to the prompt from the custom prompt management callback since we did not pass a prompt_id + assert ( + request_body["messages"][0]["content"][0]["text"] == "Hello, how are you?" + )