Prompt Management - abstract prompt templates away from model list (enables permission management on prompt templates) (#13219)

* feat: initial commit with prompt management support on pre-call hooks

allows prompt templates to work before assigning specific models

* feat: initial logic for independent prompt management settings

* feat(proxy_server.py): working logic for loading in the prompt templates from config yaml

allows creating an independent 'prompts' section in the config yaml

* feat(prompt_registry.py): working e2e custom prompt templates with guardrails and models

* refactor(prompts/): move folder inside proxy folder

easier management for prompt endpoints

* fix: fix linting error

* fix: fix check
This commit is contained in:
Krish Dholakia 2025-08-02 09:39:45 -07:00 • committed by GitHub
parent 825923e7be
commit a107a4bdba
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
17 changed files with 412 additions and 56 deletions

View file

@ -5,7 +5,17 @@ warnings.filterwarnings("ignore", message=".*conflict with protected namespace.*
### INIT VARIABLES ####################
import threading
import os
from typing import Callable, List, Optional, Dict, Union, Any, Literal, get_args, TYPE_CHECKING
from typing import (
Callable,
List,
Optional,
Dict,
Union,
Any,
Literal,
get_args,
TYPE_CHECKING,
)
from litellm.types.integrations.datadog_llm_obs import DatadogLLMObsInitParams
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.caching.caching import Cache, DualCache, RedisCache, InMemoryCache
@ -259,6 +269,11 @@ blocked_user_list: Optional[Union[str, List]] = None
banned_keywords_list: Optional[Union[str, List]] = None
llm_guard_mode: Literal["all", "key-specific", "request-specific"] = "all"
guardrail_name_config_map: Dict[str, GuardrailItem] = {}
### PROMPTS ###
from litellm.types.prompts.init_prompts import PromptSpec
prompt_name_config_map: Dict[str, PromptSpec] = {}
##################
### PREVIEW FEATURES ###
enable_preview_features: bool = False

View file

@ -275,7 +275,7 @@ Represents a single prompt with metadata.
**Prompt not found**: Ensure the `.prompt` file exists and has correct extension
```python
# Check available prompts
from litellm.prompts import get_dotprompt_manager
from litellm.integrations.dotprompt import get_dotprompt_manager
manager = get_dotprompt_manager()
print(manager.prompt_manager.list_prompts())
```

View file

@ -2,6 +2,10 @@ from typing import TYPE_CHECKING, Optional
if TYPE_CHECKING:
from .prompt_manager import PromptManager, PromptTemplate
from litellm.types.prompts.init_prompts import PromptLiteLLMParams, PromptSpec
from litellm.integrations.custom_prompt_management import CustomPromptManagement
from litellm.types.prompts.init_prompts import SupportedPromptIntegrations
from .dotprompt_manager import DotpromptManager
@ -22,6 +26,23 @@ def set_global_prompt_directory(directory: str) -> None:
litellm.global_prompt_directory = directory # type: ignore
def prompt_initializer(
litellm_params: "PromptLiteLLMParams", prompt_spec: "PromptSpec"
) -> "CustomPromptManagement":
"""
Initialize a prompt from a .prompt file.
"""
prompt_directory = getattr(litellm_params, "prompt_directory", None)
if not prompt_directory:
raise ValueError("prompt_directory is required for dotprompt")
return DotpromptManager(prompt_directory)
prompt_initializer_registry = {
SupportedPromptIntegrations.DOT_PROMPT.value: prompt_initializer,
}
# Export public API
__all__ = [
"PromptManager",

View file

@ -3,20 +3,17 @@ Dotprompt manager that integrates with LiteLLM's prompt management system.
Builds on top of PromptManagementBase to provide .prompt file support.
"""
from typing import List, Optional
from typing import List, Optional, Tuple
from litellm.integrations.custom_logger import CustomLogger
from litellm.integrations.prompt_management_base import (
PromptManagementBase,
PromptManagementClient,
)
from litellm.integrations.custom_prompt_management import CustomPromptManagement
from litellm.integrations.prompt_management_base import PromptManagementClient
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import StandardCallbackDynamicParams
from .prompt_manager import PromptManager, PromptTemplate
class DotpromptManager(PromptManagementBase, CustomLogger):
class DotpromptManager(CustomPromptManagement):
"""
Dotprompt manager that integrates with LiteLLM's prompt management system.
@ -93,6 +90,7 @@ class DotpromptManager(PromptManagementBase, CustomLogger):
3. Converts the rendered text into chat messages
4. Extracts model and optional parameters from metadata
"""
try:
# Get the prompt template
template = self.prompt_manager.get_prompt(prompt_id)
@ -122,6 +120,31 @@ class DotpromptManager(PromptManagementBase, CustomLogger):
except Exception as e:
raise ValueError(f"Error compiling prompt '{prompt_id}': {e}")
def get_chat_completion_prompt(
self,
model: str,
messages: List[AllMessageValues],
non_default_params: dict,
prompt_id: Optional[str],
prompt_variables: Optional[dict],
dynamic_callback_params: StandardCallbackDynamicParams,
prompt_label: Optional[str] = None,
prompt_version: Optional[int] = None,
) -> Tuple[str, List[AllMessageValues], dict]:
from litellm.integrations.prompt_management_base import PromptManagementBase
return PromptManagementBase.get_chat_completion_prompt(
self,
model,
messages,
non_default_params,
prompt_id,
prompt_variables,
dynamic_callback_params,
prompt_label,
prompt_version,
)
def _convert_to_messages(self, rendered_content: str) -> List[AllMessageValues]:
"""
Convert rendered prompt content to chat messages.

View file

@ -562,7 +562,9 @@ 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,
)
)
@ -581,6 +583,7 @@ class Logging(LiteLLMLoggingBaseClass):
prompt_label=prompt_label,
prompt_version=prompt_version,
)
self.messages = messages
return model, messages, non_default_params
@ -599,7 +602,10 @@ 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,
)
)
@ -624,7 +630,11 @@ class Logging(LiteLLMLoggingBaseClass):
return model, messages, non_default_params
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,
) -> Optional[CustomLogger]:
"""
Get a custom logger for prompt management based on model name or available callbacks.
@ -635,7 +645,7 @@ class Logging(LiteLLMLoggingBaseClass):
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(

File diff suppressed because one or more lines are too long

View file

@ -1,9 +1,20 @@
model_list:
- model_name: openai-test
litellm_params:
model: dotprompt/gpt-3.5-turbo
prompt_id: test_hello_world_prompt
model: gpt-3.5-turbo
api_key: os.environ/OPENAI_API_KEY
litellm_settings:
global_prompt_directory: /Users/krrishdholakia/Documents/litellm/litellm/proxy/test_prompts
guardrails:
- guardrail_name: azure-text-moderation
litellm_params:
guardrail: azure/text_moderations
mode: "post_call"
api_key: os.environ/AZURE_GUARDRAIL_API_KEY
api_base: os.environ/AZURE_GUARDRAIL_API_BASE
prompts:
- prompt_id: test_hello_world_prompt
litellm_params:
prompt_integration: dotprompt
prompt_id: test_hello_world_prompt
prompt_directory: /Users/krrishdholakia/Documents/litellm/litellm/proxy/test_prompts

View file

@ -328,10 +328,6 @@ class ProxyBaseLLMRequestProcessing:
)
### CALL HOOKS ### - modify/reject incoming data before calling the model
self.data = await proxy_logging_obj.pre_call_hook( # type: ignore
user_api_key_dict=user_api_key_dict, data=self.data, call_type=route_type # type: ignore
)
## LOGGING OBJECT ## - initialize logging object for logging success/failure events for call
## IMPORTANT Note: - initialize this before running pre-call checks. Ensures we log rejected requests to langfuse.
logging_obj, self.data = litellm.utils.function_setup(
@ -343,6 +339,10 @@ class ProxyBaseLLMRequestProcessing:
self.data["litellm_logging_obj"] = logging_obj
self.data = await proxy_logging_obj.pre_call_hook( # type: ignore
user_api_key_dict=user_api_key_dict, data=self.data, call_type=route_type # type: ignore
)
return self.data, logging_obj
async def base_process_llm_request(

View file

@ -415,6 +415,7 @@ class InMemoryGuardrailHandler:
litellm_params.api_base = str(get_secret(litellm_params.api_base))
guardrail_type = litellm_params.guardrail
if guardrail_type is None:
raise ValueError("guardrail_type is required")

View file

@ -9,7 +9,32 @@ from litellm.types.guardrails import Guardrail, GuardrailItem, GuardrailItemSpec
all_guardrails: List[GuardrailItem] = []
"""
Map guardrail_name: <pre_call>, <post_call>, during_call
"""
def init_guardrails_v2(
all_guardrails: List[Dict],
config_file_path: Optional[str] = None,
):
from litellm.proxy.guardrails.guardrail_registry import IN_MEMORY_GUARDRAIL_HANDLER
guardrail_list: List[Guardrail] = []
for guardrail in all_guardrails:
initialized_guardrail = IN_MEMORY_GUARDRAIL_HANDLER.initialize_guardrail(
guardrail=cast(Guardrail, guardrail),
config_file_path=config_file_path,
)
if initialized_guardrail:
guardrail_list.append(initialized_guardrail)
verbose_proxy_logger.debug(f"\nGuardrail List:{guardrail_list}\n")
### LEGACY IMPLEMENTATION ###
def initialize_guardrails(
guardrails_config: List[Dict[str, GuardrailItemSpec]],
premium_user: bool,
@ -65,28 +90,3 @@ def initialize_guardrails(
"error initializing guardrails {}".format(str(e))
)
raise e
"""
Map guardrail_name: <pre_call>, <post_call>, during_call
"""
def init_guardrails_v2(
all_guardrails: List[Dict],
config_file_path: Optional[str] = None,
):
from litellm.proxy.guardrails.guardrail_registry import IN_MEMORY_GUARDRAIL_HANDLER
guardrail_list: List[Guardrail] = []
for guardrail in all_guardrails:
initialized_guardrail = IN_MEMORY_GUARDRAIL_HANDLER.initialize_guardrail(
guardrail=cast(Guardrail, guardrail),
config_file_path=config_file_path,
)
if initialized_guardrail:
guardrail_list.append(initialized_guardrail)
verbose_proxy_logger.debug(f"\nGuardrail List:{guardrail_list}\n")

View file

View file

@ -0,0 +1,28 @@
"""
Similar to init_guardrails.py, but for prompts.
"""
from typing import Dict, List, Optional, cast
from litellm._logging import verbose_proxy_logger
def init_prompts(
all_prompts: List[Dict],
config_file_path: Optional[str] = None,
):
from litellm.types.prompts.init_prompts import PromptSpec
from .prompt_registry import IN_MEMORY_PROMPT_REGISTRY
prompt_list: List[PromptSpec] = []
for prompt in all_prompts:
initialized_prompt = IN_MEMORY_PROMPT_REGISTRY.initialize_prompt(
prompt=cast(PromptSpec, prompt),
config_file_path=config_file_path,
)
if initialized_prompt:
prompt_list.append(initialized_prompt)
verbose_proxy_logger.debug(f"\nPrompt List:{prompt_list}\n")

View file

@ -0,0 +1,173 @@
import importlib
import os
import uuid
from pathlib import Path
from typing import Callable, Dict, Optional
from litellm._logging import verbose_proxy_logger
from litellm.integrations.custom_prompt_management import CustomPromptManagement
from litellm.types.prompts.init_prompts import PromptLiteLLMParams, PromptSpec
prompt_initializer_registry = {}
def get_prompt_initializer_from_integrations():
"""
Get prompt initializers by discovering them from the prompt_integrations directory structure.
Scans the integrations directory for subdirectories containing __init__.py files
with either prompt_initializer_registry or initialize_prompt functions.
Returns:
Dict[str, Callable]: A dictionary mapping guardrail types to their initializer functions
"""
discovered_initializers: Dict[str, Callable] = {}
try:
# Get the path to the prompt_integrations directory
current_dir = Path(__file__).parent.parent.parent
integrations_dir = os.path.join(current_dir, "integrations")
if not os.path.exists(integrations_dir):
verbose_proxy_logger.debug("integrations directory not found")
return discovered_initializers
# Scan each subdirectory in prompt_integrations
for item in os.listdir(integrations_dir):
item_path = os.path.join(integrations_dir, item)
# Skip files and __pycache__ directories
if not os.path.isdir(item_path) or item.startswith("__"):
continue
# Check if the directory has an __init__.py file
init_file = os.path.join(item_path, "__init__.py")
if not os.path.exists(init_file):
continue
module_path = f"litellm.integrations.{item}"
try:
# Import the module
verbose_proxy_logger.debug(
f"Discovering prompt integrations in: {module_path}"
)
module = importlib.import_module(module_path)
# Check for prompt_initializer_registry dictionary
if hasattr(module, "prompt_initializer_registry"):
registry = getattr(module, "prompt_initializer_registry")
if isinstance(registry, dict):
discovered_initializers.update(registry)
verbose_proxy_logger.debug(
f"Found prompt_initializer_registry in {module_path}: {list(registry.keys())}"
)
except ImportError as e:
verbose_proxy_logger.error(f"Could not import {module_path}: {e}")
continue
except Exception as e:
verbose_proxy_logger.error(f"Error processing {module_path}: {e}")
continue
verbose_proxy_logger.debug(
f"Discovered {len(discovered_initializers)} prompt initializers: {list(discovered_initializers.keys())}"
)
except Exception as e:
verbose_proxy_logger.error(f"Error discovering prompt initializers: {e}")
return discovered_initializers
prompt_initializer_registry = get_prompt_initializer_from_integrations()
class InMemoryPromptRegistry:
"""
Class that handles adding prompt callbacks to the CallbacksManager.
"""
def __init__(self):
self.IN_MEMORY_PROMPTS: Dict[str, PromptSpec] = {}
"""
Prompt id to Prompt object mapping
"""
self.prompt_id_to_custom_prompt: Dict[str, Optional[CustomPromptManagement]] = (
{}
)
"""
Guardrail id to CustomGuardrail object mapping
"""
def initialize_prompt(
self,
prompt: PromptSpec,
config_file_path: Optional[str] = None,
) -> Optional[PromptSpec]:
"""
Initialize a guardrail from a dictionary and add it to the litellm callback manager
Returns a Guardrail object if the guardrail is initialized successfully
"""
import litellm
prompt_id = prompt.get("prompt_id") or str(uuid.uuid4())
prompt["prompt_id"] = prompt_id
if prompt_id in self.IN_MEMORY_PROMPTS:
verbose_proxy_logger.debug("prompt_id already exists in IN_MEMORY_PROMPTS")
return self.IN_MEMORY_PROMPTS[prompt_id]
custom_prompt_callback: Optional[CustomPromptManagement] = None
litellm_params_data = prompt["litellm_params"]
verbose_proxy_logger.debug("litellm_params= %s", litellm_params_data)
if isinstance(litellm_params_data, dict):
litellm_params = PromptLiteLLMParams(**litellm_params_data)
else:
litellm_params = litellm_params_data
prompt_integration = litellm_params.prompt_integration
if prompt_integration is None:
raise ValueError("prompt_integration is required")
initializer = prompt_initializer_registry.get(prompt_integration)
if initializer:
custom_prompt_callback = initializer(litellm_params, prompt)
if not isinstance(custom_prompt_callback, CustomPromptManagement):
raise ValueError(
f"CustomPromptManagement is required, got {type(custom_prompt_callback)}"
)
litellm.logging_callback_manager.add_litellm_callback(custom_prompt_callback) # type: ignore
else:
raise ValueError(f"Unsupported prompt: {prompt_integration}")
parsed_prompt = PromptSpec(
prompt_id=prompt_id,
litellm_params=litellm_params,
)
# store references to the prompt in memory
self.IN_MEMORY_PROMPTS[prompt_id] = parsed_prompt
self.prompt_id_to_custom_prompt[prompt_id] = custom_prompt_callback
return parsed_prompt
def get_prompt_by_id(self, prompt_id: str) -> Optional[PromptSpec]:
"""
Get a prompt by its ID from memory
"""
return self.IN_MEMORY_PROMPTS.get(prompt_id)
def get_prompt_callback_by_id(
self, prompt_id: str
) -> Optional[CustomPromptManagement]:
"""
Get a prompt callback by its ID from memory
"""
return self.prompt_id_to_custom_prompt.get(prompt_id)
IN_MEMORY_PROMPT_REGISTRY = InMemoryPromptRegistry()

View file

@ -1822,6 +1822,7 @@ class ProxyConfig:
)
litellm.guardrail_name_config_map = guardrail_name_config_map
elif key == "global_prompt_directory":
from litellm.integrations.dotprompt import (
set_global_prompt_directory,
@ -2201,6 +2202,15 @@ class ProxyConfig:
all_guardrails=guardrails_v2, config_file_path=config_file_path
)
## Prompt settings
prompts: Optional[List[Dict]] = None
if config is not None:
prompts = config.get("prompts", None)
if prompts:
from litellm.proxy.prompts.init_prompts import init_prompts
init_prompts(all_prompts=prompts, config_file_path=config_file_path)
## CREDENTIALS
credential_list_dict = self.load_credential_list(config=config)
litellm.credential_list = credential_list_dict

View file

@ -98,6 +98,8 @@ from litellm.types.utils import CallTypes, LLMResponseTypes, LoggedLiteLLMParams
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
Span = Union[_Span, Any]
else:
Span = Any
@ -957,6 +959,8 @@ class ProxyLogging:
2. /embeddings
3. /image/generation
"""
from litellm.utils import get_non_default_completion_params
verbose_proxy_logger.debug("Inside Proxy Logging Pre-call hook!")
self._init_response_taking_too_long_task(data=data)
@ -964,6 +968,42 @@ class ProxyLogging:
if data is None:
return None
litellm_logging_obj = cast(
Optional["LiteLLMLoggingObj"], data.get("litellm_logging_obj", None)
)
prompt_id = data.get("prompt_id", None)
## PROMPT TEMPLATE CHECK ##
if (
litellm_logging_obj is not None
and prompt_id is not None
and (call_type == "completion" or call_type == "acompletion")
):
from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY
custom_logger = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id(
prompt_id
)
if custom_logger:
(
model,
messages,
optional_params,
) = litellm_logging_obj.get_chat_completion_prompt(
model=data.get("model", ""),
messages=data.get("messages", []),
non_default_params=get_non_default_completion_params(kwargs=data),
prompt_id=prompt_id,
prompt_management_logger=custom_logger,
prompt_variables=data.get("prompt_variables", None),
prompt_label=data.get("prompt_label", None),
prompt_version=data.get("prompt_version", None),
)
data["model"] = model
data["messages"] = messages
data.update(optional_params)
try:
for callback in litellm.callbacks:
_callback = None
@ -3696,26 +3736,24 @@ def is_valid_api_key(key: str) -> bool:
def construct_database_url_from_env_vars() -> Optional[str]:
"""
Construct a DATABASE_URL from individual environment variables.
Returns:
Optional[str]: The constructed DATABASE_URL or None if required variables are missing
"""
import urllib.parse
# Check if all required variables are provided
database_host = os.getenv("DATABASE_HOST")
database_username = os.getenv("DATABASE_USERNAME")
database_password = os.getenv("DATABASE_PASSWORD")
database_name = os.getenv("DATABASE_NAME")
if (
database_host
and database_username
and database_name
):
if database_host and database_username and database_name:
# Handle the problem of special character escaping in the database URL
database_username_enc = urllib.parse.quote_plus(database_username)
database_password_enc = urllib.parse.quote_plus(database_password) if database_password else ""
database_password_enc = (
urllib.parse.quote_plus(database_password) if database_password else ""
)
database_name_enc = urllib.parse.quote_plus(database_name)
# Construct DATABASE_URL from the provided variables
@ -3725,5 +3763,5 @@ def construct_database_url_from_env_vars() -> Optional[str]:
database_url = f"postgresql://{database_username_enc}@{database_host}/{database_name_enc}"
return database_url
return None

View file

@ -0,0 +1,27 @@
from datetime import datetime
from enum import Enum
from typing import Dict, Optional
from pydantic import BaseModel, ConfigDict
from typing_extensions import Required, TypedDict
class SupportedPromptIntegrations(str, Enum):
DOT_PROMPT = "dotprompt"
LANGFUSE = "langfuse"
CUSTOM = "custom"
class PromptLiteLLMParams(BaseModel):
prompt_id: str
prompt_integration: str
model_config = ConfigDict(extra="allow", protected_namespaces=())
class PromptSpec(TypedDict, total=False):
prompt_id: Required[str]
litellm_params: Required[PromptLiteLLMParams]
prompt_info: Optional[Dict]
created_at: Optional[datetime]
updated_at: Optional[datetime]