Prompt Management - add local dotprompt file support

This commit is contained in:
Krrish Dholakia
2025-07-31 22:28:29 -07:00
parent 817b139420
commit fe24c270de
21 changed files with 1558 additions and 128 deletions
+70 -50
View File
@@ -61,6 +61,11 @@ from litellm.constants import (
DEFAULT_SOFT_BUDGET,
DEFAULT_ALLOWED_FAILS,
)
from litellm.integrations.dotprompt import (
global_prompt_manager,
global_prompt_directory,
set_global_prompt_directory,
)
from litellm.types.guardrails import GuardrailItem
from litellm.types.secret_managers.main import (
KeyManagementSystem,
@@ -83,7 +88,6 @@ if litellm_mode == "DEV":
# Register async client cleanup to prevent resource leaks
register_async_client_cleanup()
####################################################
if set_verbose == True:
_turn_on_debug()
@@ -130,6 +134,7 @@ _custom_logger_compatible_callbacks_literal = Literal[
"s3_v2",
"aws_sqs",
"vector_store_pre_call_hook",
"dotprompt",
]
logged_real_time_event_types: Optional[Union[List[str], Literal["*"]]] = None
_known_custom_logger_compatible_callbacks: List = list(
@@ -145,22 +150,22 @@ prometheus_initialize_budget_metrics: Optional[bool] = False
require_auth_for_metrics_endpoint: Optional[bool] = False
argilla_batch_size: Optional[int] = None
datadog_use_v1: Optional[bool] = False # if you want to use v1 datadog logged payload.
gcs_pub_sub_use_v1: Optional[
bool
] = False # if you want to use v1 gcs pubsub logged payload
generic_api_use_v1: Optional[
bool
] = False # if you want to use v1 generic api logged payload
gcs_pub_sub_use_v1: Optional[bool] = (
False # if you want to use v1 gcs pubsub logged payload
)
generic_api_use_v1: Optional[bool] = (
False # if you want to use v1 generic api logged payload
)
argilla_transformation_object: Optional[Dict[str, Any]] = None
_async_input_callback: List[
Union[str, Callable, CustomLogger]
] = [] # internal variable - async custom callbacks are routed here.
_async_success_callback: List[
Union[str, Callable, CustomLogger]
] = [] # internal variable - async custom callbacks are routed here.
_async_failure_callback: List[
Union[str, Callable, CustomLogger]
] = [] # internal variable - async custom callbacks are routed here.
_async_input_callback: List[Union[str, Callable, CustomLogger]] = (
[]
) # internal variable - async custom callbacks are routed here.
_async_success_callback: List[Union[str, Callable, CustomLogger]] = (
[]
) # internal variable - async custom callbacks are routed here.
_async_failure_callback: List[Union[str, Callable, CustomLogger]] = (
[]
) # internal variable - async custom callbacks are routed here.
pre_call_rules: List[Callable] = []
post_call_rules: List[Callable] = []
turn_off_message_logging: Optional[bool] = False
@@ -168,18 +173,18 @@ log_raw_request_response: bool = False
redact_messages_in_exceptions: Optional[bool] = False
redact_user_api_key_info: Optional[bool] = False
filter_invalid_headers: Optional[bool] = False
add_user_information_to_llm_headers: Optional[
bool
] = None # adds user_id, team_id, token hash (params from StandardLoggingMetadata) to request headers
add_user_information_to_llm_headers: Optional[bool] = (
None # adds user_id, team_id, token hash (params from StandardLoggingMetadata) to request headers
)
store_audit_logs = False # Enterprise feature, allow users to see audit logs
### end of callbacks #############
email: Optional[
str
] = None # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648
token: Optional[
str
] = None # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648
email: Optional[str] = (
None # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648
)
token: Optional[str] = (
None # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648
)
telemetry = True
max_tokens: int = DEFAULT_MAX_TOKENS # OpenAI Defaults
drop_params = bool(os.getenv("LITELLM_DROP_PARAMS", False))
@@ -267,11 +272,15 @@ enable_loadbalancing_on_batch_endpoints: Optional[bool] = None
enable_caching_on_provider_specific_optional_params: bool = (
False # feature-flag for caching on optional params - e.g. 'top_k'
)
caching: bool = False # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648
caching_with_models: bool = False # # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648
cache: Optional[
Cache
] = None # cache object <- use this - https://docs.litellm.ai/docs/caching
caching: bool = (
False # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648
)
caching_with_models: bool = (
False # # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648
)
cache: Optional[Cache] = (
None # cache object <- use this - https://docs.litellm.ai/docs/caching
)
default_in_memory_ttl: Optional[float] = None
default_redis_ttl: Optional[float] = None
default_redis_batch_cache_expiry: Optional[float] = None
@@ -279,9 +288,9 @@ model_alias_map: Dict[str, str] = {}
model_group_alias_map: Dict[str, str] = {}
model_group_settings: Optional["ModelGroupSettings"] = None
max_budget: float = 0.0 # set the max budget across all providers
budget_duration: Optional[
str
] = None # proxy only - resets budget after fixed duration. You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d").
budget_duration: Optional[str] = (
None # proxy only - resets budget after fixed duration. You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d").
)
default_soft_budget: float = (
DEFAULT_SOFT_BUDGET # by default all litellm proxy keys have a soft budget of 50.0
)
@@ -290,11 +299,15 @@ forward_traceparent_to_llm_provider: bool = False
_current_cost = 0.0 # private variable, used if max budget is set
error_logs: Dict = {}
add_function_to_prompt: bool = False # if function calling not supported by api, append function call details to system prompt
add_function_to_prompt: bool = (
False # if function calling not supported by api, append function call details to system prompt
)
client_session: Optional[httpx.Client] = None
aclient_session: Optional[httpx.AsyncClient] = None
model_fallbacks: Optional[List] = None # Deprecated for 'litellm.fallbacks'
model_cost_map_url: str = "https://raw.githubusercontent.com/BerriAI/litellm/main/model_prices_and_context_window.json"
model_cost_map_url: str = (
"https://raw.githubusercontent.com/BerriAI/litellm/main/model_prices_and_context_window.json"
)
suppress_debug_info = False
dynamodb_table_name: Optional[str] = None
s3_callback_params: Optional[Dict] = None
@@ -323,7 +336,9 @@ prometheus_metrics_config: Optional[List] = None
disable_add_prefix_to_prompt: bool = (
False # used by anthropic, to disable adding prefix to prompt
)
disable_copilot_system_to_assistant: bool = False # If false (default), converts all 'system' role messages to 'assistant' for GitHub Copilot compatibility. Set to true to disable this behavior.
disable_copilot_system_to_assistant: bool = (
False # If false (default), converts all 'system' role messages to 'assistant' for GitHub Copilot compatibility. Set to true to disable this behavior.
)
public_model_groups: Optional[List[str]] = None
public_model_groups_links: Dict[str, str] = {}
#### REQUEST PRIORITIZATION #####
@@ -331,13 +346,17 @@ priority_reservation: Optional[Dict[str, float]] = None
######## Networking Settings ########
use_aiohttp_transport: bool = True # Older variable, aiohttp is now the default. use disable_aiohttp_transport instead.
use_aiohttp_transport: bool = (
True # Older variable, aiohttp is now the default. use disable_aiohttp_transport instead.
)
aiohttp_trust_env: bool = False # set to true to use HTTP_ Proxy settings
disable_aiohttp_transport: bool = False # Set this to true to use httpx instead
disable_aiohttp_trust_env: bool = (
False # When False, aiohttp will respect HTTP(S)_PROXY env vars
)
force_ipv4: bool = False # when True, litellm will force ipv4 for all LLM requests. Some users have seen httpx ConnectionError when using ipv6.
force_ipv4: bool = (
False # when True, litellm will force ipv4 for all LLM requests. Some users have seen httpx ConnectionError when using ipv6.
)
module_level_aclient = AsyncHTTPHandler(
timeout=request_timeout, client_alias="module level aclient"
)
@@ -351,13 +370,13 @@ fallbacks: Optional[List] = None
context_window_fallbacks: Optional[List] = None
content_policy_fallbacks: Optional[List] = None
allowed_fails: int = 3
num_retries_per_request: Optional[
int
] = None # for the request overall (incl. fallbacks + model retries)
num_retries_per_request: Optional[int] = (
None # for the request overall (incl. fallbacks + model retries)
)
####### SECRET MANAGERS #####################
secret_manager_client: Optional[
Any
] = None # list of instantiated key management clients - e.g. azure kv, infisical, etc.
secret_manager_client: Optional[Any] = (
None # list of instantiated key management clients - e.g. azure kv, infisical, etc.
)
_google_kms_resource_name: Optional[str] = None
_key_management_system: Optional[KeyManagementSystem] = None
_key_management_settings: KeyManagementSettings = KeyManagementSettings()
@@ -496,6 +515,7 @@ lambda_ai_models: List = []
hyperbolic_models: List = []
recraft_models: List = []
def is_bedrock_pricing_only_model(key: str) -> bool:
"""
Excludes keys with the pattern 'bedrock/<region>/<model>'. These are in the model_prices_and_context_window.json file for pricing purposes only.
@@ -1225,12 +1245,12 @@ from .types.llms.custom_llm import CustomLLMItem
from .types.utils import GenericStreamingChunk
custom_provider_map: List[CustomLLMItem] = []
_custom_providers: List[
str
] = [] # internal helper util, used to track names of custom providers
disable_hf_tokenizer_download: Optional[
bool
] = None # disable huggingface tokenizer download. Defaults to openai clk100
_custom_providers: List[str] = (
[]
) # internal helper util, used to track names of custom providers
disable_hf_tokenizer_download: Optional[bool] = (
None # disable huggingface tokenizer download. Defaults to openai clk100
)
global_disable_no_log_param: bool = False
### PASSTHROUGH ###
+316
View File
@@ -0,0 +1,316 @@
# LiteLLM Dotprompt Manager
A powerful prompt management system for LiteLLM that supports [Google's Dotprompt specification](https://google.github.io/dotprompt/getting-started/). This allows you to manage your AI prompts in organized `.prompt` files with YAML frontmatter, Handlebars templating, and full integration with LiteLLM's completion API.
## Features
- **📁 File-based prompt management**: Organize prompts in `.prompt` files
- **🎯 YAML frontmatter**: Define model, parameters, and schemas in file headers
- **🔧 Handlebars templating**: Use `{{variable}}` syntax with Jinja2 backend
- **✅ Input validation**: Automatic validation against defined schemas
- **🔗 LiteLLM integration**: Works seamlessly with `litellm.completion()`
- **💬 Smart message parsing**: Converts prompts to proper chat messages
- **⚙️ Parameter extraction**: Automatically applies model settings from prompts
## Quick Start
### 1. Create a `.prompt` file
Create a file called `chat_assistant.prompt`:
```yaml
---
model: gpt-4
temperature: 0.7
max_tokens: 150
input:
schema:
user_message: string
system_context?: string
---
{% if system_context %}System: {{system_context}}
{% endif %}User: {{user_message}}
```
### 2. Use with LiteLLM
```python
import litellm
litellm.set_global_prompt_directory("path/to/your/prompts")
# Use with completion - the model prefix 'dotprompt/' tells LiteLLM to use prompt management
response = litellm.completion(
model="dotprompt/gpt-4", # The actual model comes from the .prompt file
prompt_id="chat_assistant",
prompt_variables={
"user_message": "What is machine learning?",
"system_context": "You are a helpful AI tutor."
},
# Any additional messages will be appended after the prompt
messages=[{"role": "user", "content": "Please explain it simply."}]
)
print(response.choices[0].message.content)
```
## Prompt File Format
### Basic Structure
```yaml
---
# Model configuration
model: gpt-4
temperature: 0.7
max_tokens: 500
# Input schema (optional)
input:
schema:
name: string
age: integer
preferences?: array
---
# Template content using Handlebars syntax
Hello {{name}}!
{% if age >= 18 %}
You're an adult, so here are some mature recommendations:
{% else %}
Here are some age-appropriate suggestions:
{% endif %}
{% for pref in preferences %}
- Based on your interest in {{pref}}, I recommend...
{% endfor %}
```
### Supported Frontmatter Fields
- **`model`**: The LLM model to use (e.g., `gpt-4`, `claude-3-sonnet`)
- **`input.schema`**: Define expected input variables and their types
- **`output.format`**: Expected output format (`json`, `text`, etc.)
- **`output.schema`**: Structure of expected output
### Additional Parameters
- **`temperature`**: Model temperature (0.0 to 1.0)
- **`max_tokens`**: Maximum tokens to generate
- **`top_p`**: Nucleus sampling parameter (0.0 to 1.0)
- **`frequency_penalty`**: Frequency penalty (0.0 to 1.0)
- **`presence_penalty`**: Presence penalty (0.0 to 1.0)
- any other parameters that are not model or schema-related will be treated as optional parameters to the model.
### Input Schema Types
- `string` or `str`: Text values
- `integer` or `int`: Whole numbers
- `float`: Decimal numbers
- `boolean` or `bool`: True/false values
- `array` or `list`: Lists of values
- `object` or `dict`: Key-value objects
Use `?` suffix for optional fields: `name?: string`
## Message Format Conversion
The dotprompt manager intelligently converts your rendered prompts into proper chat messages:
### Simple Text → User Message
```yaml
---
model: gpt-4
---
Tell me about {{topic}}.
```
Becomes: `[{"role": "user", "content": "Tell me about AI."}]`
### Role-Based Format → Multiple Messages
```yaml
---
model: gpt-4
---
System: You are a {{role}}.
User: {{question}}
```
Becomes:
```python
[
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": "What is AI?"}
]
```
## Example Prompts
### Data Extraction
```yaml
# extract_info.prompt
---
model: gemini/gemini-1.5-pro
input:
schema:
text: string
output:
format: json
schema:
title?: string
summary: string
tags: array
---
Extract the requested information from the given text. Return JSON format.
Text: {{text}}
```
### Code Assistant
```yaml
# code_helper.prompt
---
model: claude-3-5-sonnet-20241022
temperature: 0.2
max_tokens: 2000
input:
schema:
language: string
task: string
code?: string
---
You are an expert {{language}} programmer.
Task: {{task}}
{% if code %}
Current code:
```{{language}}
{{code}}
```
{% endif %}
Please provide a complete, well-documented solution.
```
### Multi-turn Conversation
```yaml
# conversation.prompt
---
model: gpt-4
temperature: 0.8
input:
schema:
personality: string
context: string
---
System: You are a {{personality}}. {{context}}
User: Let's start our conversation.
```
## API Reference
### PromptManager
The core class for managing `.prompt` files.
#### Methods
- **`__init__(prompt_directory: str)`**: Initialize with directory path
- **`render(prompt_id: str, variables: dict) -> str`**: Render prompt with variables
- **`list_prompts() -> List[str]`**: Get all available prompt IDs
- **`get_prompt(prompt_id: str) -> PromptTemplate`**: Get prompt template object
- **`get_prompt_metadata(prompt_id: str) -> dict`**: Get prompt metadata
- **`reload_prompts() -> None`**: Reload all prompts from directory
- **`add_prompt(prompt_id: str, content: str, metadata: dict)`**: Add prompt programmatically
### DotpromptManager
LiteLLM integration class extending `PromptManagementBase`.
#### Methods
- **`__init__(prompt_directory: str)`**: Initialize with directory path
- **`should_run_prompt_management(prompt_id: str, params: dict) -> bool`**: Check if prompt exists
- **`set_prompt_directory(directory: str)`**: Change prompt directory
- **`reload_prompts()`**: Reload prompts from directory
### PromptTemplate
Represents a single prompt with metadata.
#### Properties
- **`content: str`**: The prompt template content
- **`metadata: dict`**: Full metadata from frontmatter
- **`model: str`**: Specified model name
- **`temperature: float`**: Model temperature
- **`max_tokens: int`**: Token limit
- **`input_schema: dict`**: Input validation schema
- **`output_format: str`**: Expected output format
- **`output_schema: dict`**: Output structure schema
## Best Practices
1. **Organize by purpose**: Group related prompts in subdirectories
2. **Use descriptive names**: `extract_user_info.prompt` vs `prompt1.prompt`
3. **Define schemas**: Always specify input schemas for validation
4. **Version control**: Store `.prompt` files in git for change tracking
5. **Test prompts**: Use the test framework to validate prompt behavior
6. **Keep templates focused**: One prompt should do one thing well
7. **Use includes**: Break complex prompts into reusable components
## Troubleshooting
### Common Issues
**Prompt not found**: Ensure the `.prompt` file exists and has correct extension
```python
# Check available prompts
from litellm.prompts import get_dotprompt_manager
manager = get_dotprompt_manager()
print(manager.prompt_manager.list_prompts())
```
**Template errors**: Verify Handlebars syntax and variable names
```python
# Test rendering directly
manager.prompt_manager.render("my_prompt", {"test": "value"})
```
**Model not working**: Check that model name in frontmatter is correct
```python
# Check prompt metadata
metadata = manager.prompt_manager.get_prompt_metadata("my_prompt")
print(metadata)
```
### Validation Errors
Input validation failures show helpful error messages:
```
ValueError: Invalid type for field 'age': expected int, got str
```
Make sure your variables match the defined schema types.
## Contributing
The LiteLLM Dotprompt manager follows the [Dotprompt specification](https://google.github.io/dotprompt/) for maximum compatibility. When contributing:
1. Ensure compatibility with existing `.prompt` files
2. Add tests for new features
3. Update documentation
4. Follow the existing code style
## License
This prompt management system is part of LiteLLM and follows the same license terms.
@@ -0,0 +1,33 @@
from typing import TYPE_CHECKING, Optional
if TYPE_CHECKING:
from .prompt_manager import PromptManager, PromptTemplate
from .dotprompt_manager import DotpromptManager
# Global instances
global_prompt_directory: Optional[str] = None
global_prompt_manager: Optional["PromptManager"] = None
def set_global_prompt_directory(directory: str) -> None:
"""
Set the global prompt directory for dotprompt files.
Args:
directory: Path to directory containing .prompt files
"""
import litellm
litellm.global_prompt_directory = directory # type: ignore
# Export public API
__all__ = [
"PromptManager",
"DotpromptManager",
"PromptTemplate",
"set_global_prompt_directory",
"global_prompt_directory",
"global_prompt_manager",
]
@@ -0,0 +1,225 @@
"""
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 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
from .prompt_manager import PromptManager, PromptTemplate
class DotpromptManager(PromptManagementBase, CustomLogger):
"""
Dotprompt manager that integrates with LiteLLM's prompt management system.
This class enables using .prompt files with the litellm completion() function
by implementing the PromptManagementBase interface.
Usage:
# Set global prompt directory
litellm.prompt_directory = "path/to/prompts"
# Use with completion
response = litellm.completion(
model="dotprompt/gpt-4",
prompt_id="my_prompt",
prompt_variables={"variable": "value"},
messages=[{"role": "user", "content": "This will be combined with the prompt"}]
)
"""
def __init__(self, prompt_directory: Optional[str] = None):
import litellm
self.prompt_directory = prompt_directory or litellm.global_prompt_directory
self._prompt_manager: Optional[PromptManager] = None
@property
def integration_name(self) -> str:
"""Integration name used in model names like 'dotprompt/gpt-4'."""
return "dotprompt"
@property
def prompt_manager(self) -> PromptManager:
"""Lazy-load the prompt manager."""
if self._prompt_manager is None:
if self.prompt_directory is None:
raise ValueError(
"prompt_directory must be set before using dotprompt manager. "
"Set litellm.global_prompt_directory or initialize with prompt_directory parameter."
)
self._prompt_manager = PromptManager(self.prompt_directory)
return self._prompt_manager
def should_run_prompt_management(
self,
prompt_id: str,
dynamic_callback_params: StandardCallbackDynamicParams,
) -> bool:
"""
Determine if prompt management should run based on the prompt_id.
Returns True if the prompt_id exists in our prompt manager.
"""
try:
return prompt_id in self.prompt_manager.list_prompts()
except Exception:
# If there's any error accessing prompts, don't run prompt management
return False
def _compile_prompt_helper(
self,
prompt_id: str,
prompt_variables: Optional[dict],
dynamic_callback_params: StandardCallbackDynamicParams,
prompt_label: Optional[str] = None,
prompt_version: Optional[int] = None,
) -> PromptManagementClient:
"""
Compile a .prompt file into a PromptManagementClient structure.
This method:
1. Loads the prompt template from the .prompt file
2. Renders it with the provided variables
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)
if template is None:
raise ValueError(f"Prompt '{prompt_id}' not found in prompt directory")
# Render the template with variables
rendered_content = self.prompt_manager.render(prompt_id, prompt_variables)
# Convert rendered content to chat messages
messages = self._convert_to_messages(rendered_content)
# Extract model from metadata (if specified)
template_model = template.model
# Extract optional parameters from metadata
optional_params = self._extract_optional_params(template)
return PromptManagementClient(
prompt_id=prompt_id,
prompt_template=messages,
prompt_template_model=template_model,
prompt_template_optional_params=optional_params,
completed_messages=None,
)
except Exception as e:
raise ValueError(f"Error compiling prompt '{prompt_id}': {e}")
def _convert_to_messages(self, rendered_content: str) -> List[AllMessageValues]:
"""
Convert rendered prompt content to chat messages.
This method supports multiple formats:
1. Simple text -> converted to user message
2. Text with role prefixes (System:, User:, Assistant:) -> parsed into separate messages
3. Already formatted as a single message
"""
# Clean up the content
content = rendered_content.strip()
# Try to parse role-based format (System: ..., User: ..., etc.)
messages = []
current_role = None
current_content = []
lines = content.split("\n")
for line in lines:
line = line.strip()
# Check for role prefixes
if line.startswith("System:"):
if current_role and current_content:
messages.append(
self._create_message(
current_role, "\n".join(current_content).strip()
)
)
current_role = "system"
current_content = [line[7:].strip()] # Remove "System:" prefix
elif line.startswith("User:"):
if current_role and current_content:
messages.append(
self._create_message(
current_role, "\n".join(current_content).strip()
)
)
current_role = "user"
current_content = [line[5:].strip()] # Remove "User:" prefix
elif line.startswith("Assistant:"):
if current_role and current_content:
messages.append(
self._create_message(
current_role, "\n".join(current_content).strip()
)
)
current_role = "assistant"
current_content = [line[10:].strip()] # Remove "Assistant:" prefix
else:
# Continue current message content
if current_role:
current_content.append(line)
else:
# No role prefix found, treat as user message
current_role = "user"
current_content = [line]
# Add the last message
if current_role and current_content:
content_text = "\n".join(current_content).strip()
if content_text: # Only add if there's actual content
messages.append(self._create_message(current_role, content_text))
# If no messages were created, treat the entire content as a user message
if not messages and content:
messages.append(self._create_message("user", content))
return messages
def _create_message(self, role: str, content: str) -> AllMessageValues:
"""Create a message with the specified role and content."""
return {
"role": role, # type: ignore
"content": content,
}
def _extract_optional_params(self, template: PromptTemplate) -> dict:
"""
Extract optional parameters from the prompt template metadata.
Includes parameters like temperature, max_tokens, etc.
"""
optional_params = {}
# Extract common parameters from metadata
if template.optional_params is not None:
optional_params.update(template.optional_params)
return optional_params
def set_prompt_directory(self, prompt_directory: str) -> None:
"""Set the prompt directory and reload prompts."""
self.prompt_directory = prompt_directory
self._prompt_manager = None # Reset to force reload
def reload_prompts(self) -> None:
"""Reload all prompts from the directory."""
if self._prompt_manager:
self._prompt_manager.reload_prompts()
@@ -0,0 +1,220 @@
"""
Based on Google's GenAI Kit dotprompt implementation: https://google.github.io/dotprompt/reference/frontmatter/
"""
import re
from pathlib import Path
from typing import Any, Dict, List, Optional, Tuple, Union
import yaml
from jinja2 import DictLoader, Environment, select_autoescape
class PromptTemplate:
"""Represents a single prompt template with metadata and content."""
def __init__(
self,
content: str,
metadata: Optional[Dict[str, Any]] = None,
template_id: Optional[str] = None,
):
self.content = content
self.metadata = metadata or {}
self.template_id = template_id
# Extract common metadata fields
restricted_keys = ["model", "input", "output"]
self.model = self.metadata.get("model")
self.input_schema = self.metadata.get("input", {}).get("schema", {})
self.output_format = self.metadata.get("output", {}).get("format")
self.output_schema = self.metadata.get("output", {}).get("schema", {})
self.optional_params = {}
for key in self.metadata.keys():
if key not in restricted_keys:
self.optional_params[key] = self.metadata[key]
def __repr__(self):
return f"PromptTemplate(id='{self.template_id}', model='{self.model}')"
class PromptManager:
"""
Manager for loading and rendering .prompt files following the Dotprompt specification.
Supports:
- YAML frontmatter for metadata
- Handlebars-style templating (using Jinja2)
- Input/output schema validation
- Model configuration
"""
def __init__(self, prompt_directory: str):
self.prompt_directory = Path(prompt_directory)
self.prompts: Dict[str, PromptTemplate] = {}
self.jinja_env = Environment(
loader=DictLoader({}),
autoescape=select_autoescape(["html", "xml"]),
# Use Handlebars-style delimiters to match Dotprompt spec
variable_start_string="{{",
variable_end_string="}}",
block_start_string="{%",
block_end_string="%}",
comment_start_string="{#",
comment_end_string="#}",
)
# Load all prompts in the directory
self._load_prompts()
def _load_prompts(self) -> None:
"""Load all .prompt files from the prompt directory."""
if not self.prompt_directory.exists():
raise ValueError(
f"Prompt directory does not exist: {self.prompt_directory}"
)
prompt_files = list(self.prompt_directory.glob("*.prompt"))
for prompt_file in prompt_files:
try:
prompt_id = prompt_file.stem # filename without extension
template = self._load_prompt_file(prompt_file, prompt_id)
self.prompts[prompt_id] = template
# Optional: print(f"Loaded prompt: {prompt_id}")
except Exception:
# Optional: print(f"Error loading prompt file {prompt_file}")
pass
def _load_prompt_file(self, file_path: Path, prompt_id: str) -> PromptTemplate:
"""Load and parse a single .prompt file."""
content = file_path.read_text(encoding="utf-8")
# Split frontmatter and content
frontmatter, template_content = self._parse_frontmatter(content)
return PromptTemplate(
content=template_content.strip(),
metadata=frontmatter,
template_id=prompt_id,
)
def _parse_frontmatter(self, content: str) -> Tuple[Dict[str, Any], str]:
"""Parse YAML frontmatter from prompt content."""
# Match YAML frontmatter between --- delimiters
frontmatter_pattern = r"^---\s*\n(.*?)\n---\s*\n(.*)$"
match = re.match(frontmatter_pattern, content, re.DOTALL)
if match:
frontmatter_yaml = match.group(1)
template_content = match.group(2)
try:
frontmatter = yaml.safe_load(frontmatter_yaml) or {}
except yaml.YAMLError as e:
raise ValueError(f"Invalid YAML frontmatter: {e}")
else:
# No frontmatter found, treat entire content as template
frontmatter = {}
template_content = content
return frontmatter, template_content
def render(
self, prompt_id: str, prompt_variables: Optional[Dict[str, Any]] = None
) -> str:
"""
Render a prompt template with the given variables.
Args:
prompt_id: The ID of the prompt template to render
prompt_variables: Variables to substitute in the template
Returns:
The rendered prompt string
Raises:
KeyError: If prompt_id is not found
ValueError: If template rendering fails
"""
if prompt_id not in self.prompts:
available_prompts = list(self.prompts.keys())
raise KeyError(
f"Prompt '{prompt_id}' not found. Available prompts: {available_prompts}"
)
template = self.prompts[prompt_id]
variables = prompt_variables or {}
# Validate input variables against schema if defined
if template.input_schema:
self._validate_input(variables, template.input_schema)
try:
# Create Jinja2 template and render
jinja_template = self.jinja_env.from_string(template.content)
rendered = jinja_template.render(**variables)
return rendered
except Exception as e:
raise ValueError(f"Error rendering template '{prompt_id}': {e}")
def _validate_input(
self, variables: Dict[str, Any], schema: Dict[str, Any]
) -> None:
"""Basic validation of input variables against schema."""
for field_name, field_type in schema.items():
if field_name in variables:
value = variables[field_name]
expected_type = self._get_python_type(field_type)
if not isinstance(value, expected_type):
raise ValueError(
f"Invalid type for field '{field_name}': "
f"expected {getattr(expected_type, '__name__', str(expected_type))}, got {type(value).__name__}"
)
def _get_python_type(self, schema_type: str) -> Union[type, tuple]:
"""Convert schema type string to Python type."""
type_mapping: Dict[str, Union[type, tuple]] = {
"string": str,
"str": str,
"number": (int, float),
"integer": int,
"int": int,
"float": float,
"boolean": bool,
"bool": bool,
"array": list,
"list": list,
"object": dict,
"dict": dict,
}
return type_mapping.get(schema_type.lower(), str) # type: ignore
def get_prompt(self, prompt_id: str) -> Optional[PromptTemplate]:
"""Get a prompt template by ID."""
return self.prompts.get(prompt_id)
def list_prompts(self) -> List[str]:
"""Get a list of all available prompt IDs."""
return list(self.prompts.keys())
def get_prompt_metadata(self, prompt_id: str) -> Optional[Dict[str, Any]]:
"""Get metadata for a specific prompt."""
template = self.prompts.get(prompt_id)
return template.metadata if template else None
def reload_prompts(self) -> None:
"""Reload all prompts from the directory."""
self.prompts.clear()
self._load_prompts()
def add_prompt(
self, prompt_id: str, content: str, metadata: Optional[Dict[str, Any]] = None
) -> None:
"""Add a prompt template programmatically."""
template = PromptTemplate(
content=content, metadata=metadata or {}, template_id=prompt_id
)
self.prompts[prompt_id] = template
@@ -7,6 +7,7 @@ Example:
"datadog" -> DataDogLogger
"prometheus" -> PrometheusLogger
"""
from typing import Union
from litellm.integrations.agentops import AgentOps
@@ -31,10 +32,12 @@ from litellm.integrations.mlflow import MlflowLogger
from litellm.integrations.openmeter import OpenMeterLogger
from litellm.integrations.opentelemetry import OpenTelemetry
from litellm.integrations.opik.opik import OpikLogger
try:
from litellm_enterprise.integrations.prometheus import PrometheusLogger
except Exception:
PrometheusLogger = None
from litellm.integrations.dotprompt import DotpromptManager
from litellm.integrations.s3_v2 import S3Logger
from litellm.integrations.sqs import SQSLogger
from litellm.integrations.vector_store_integrations.vector_store_pre_call_hook import (
@@ -47,6 +50,7 @@ class CustomLoggerRegistry:
"""
Registry mapping the callback class string to the class type.
"""
CALLBACK_CLASS_STR_TO_CLASS_TYPE = {
"lago": LagoLogger,
"openmeter": OpenMeterLogger,
@@ -80,6 +84,7 @@ class CustomLoggerRegistry:
"aws_sqs": SQSLogger,
"dynamic_rate_limiter": _PROXY_DynamicRateLimitHandler,
"vector_store_pre_call_hook": VectorStorePreCallHook,
"dotprompt": DotpromptManager,
}
try:
@@ -110,14 +115,17 @@ class CustomLoggerRegistry:
def get_callback_str_from_class_type(cls, class_type: type) -> Union[str, None]:
"""
Get the callback string from the class type.
Args:
class_type: The class type to find the string for
Returns:
str: The callback string, or None if not found
"""
for callback_str, callback_class in cls.CALLBACK_CLASS_STR_TO_CLASS_TYPE.items():
for (
callback_str,
callback_class,
) in cls.CALLBACK_CLASS_STR_TO_CLASS_TYPE.items():
if callback_class == class_type:
return callback_str
return None
@@ -127,15 +135,18 @@ class CustomLoggerRegistry:
"""
Get all callback strings that map to the same class type.
Some class types (like OpenTelemetry) have multiple string mappings.
Args:
class_type: The class type to find all strings for
Returns:
list: List of callback strings that map to the class type
"""
callback_strs: list[str] = []
for callback_str, callback_class in cls.CALLBACK_CLASS_STR_TO_CLASS_TYPE.items():
for (
callback_str,
callback_class,
) in cls.CALLBACK_CLASS_STR_TO_CLASS_TYPE.items():
if callback_class == class_type:
callback_strs.append(callback_str)
return callback_strs
return callback_strs
+35 -33
View File
@@ -120,6 +120,7 @@ from ..integrations.azure_storage.azure_storage import AzureBlobStorageLogger
from ..integrations.custom_prompt_management import CustomPromptManagement
from ..integrations.datadog.datadog import DataDogLogger
from ..integrations.datadog.datadog_llm_obs import DataDogLLMObsLogger
from ..integrations.dotprompt import DotpromptManager
from ..integrations.dynamodb import DyanmoDBLogger
from ..integrations.galileo import GalileoObserve
from ..integrations.gcs_bucket.gcs_bucket import GCSBucketLogger
@@ -172,7 +173,6 @@ try:
StandardLoggingPayloadSetup as EnterpriseStandardLoggingPayloadSetup,
)
EnterpriseStandardLoggingPayloadSetupVAR: Optional[
Type[EnterpriseStandardLoggingPayloadSetup]
] = EnterpriseStandardLoggingPayloadSetup
@@ -599,9 +599,7 @@ 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
)
)
@@ -673,16 +671,16 @@ class Logging(LiteLLMLoggingBaseClass):
# Vector Store / Knowledge Base hooks
#########################################################
if litellm.vector_store_registry is not None:
vector_store_custom_logger = _init_custom_logger_compatible_class(
logging_integration="vector_store_pre_call_hook",
internal_usage_cache=None,
llm_router=None,
)
self.model_call_details["prompt_integration"] = (
vector_store_custom_logger.__class__.__name__
)
return vector_store_custom_logger
vector_store_custom_logger = _init_custom_logger_compatible_class(
logging_integration="vector_store_pre_call_hook",
internal_usage_cache=None,
llm_router=None,
)
self.model_call_details["prompt_integration"] = (
vector_store_custom_logger.__class__.__name__
)
return vector_store_custom_logger
return None
@@ -1315,9 +1313,9 @@ class Logging(LiteLLMLoggingBaseClass):
if (
EnterpriseCallbackControls is not None
and EnterpriseCallbackControls.is_callback_disabled_dynamically(
callback=callback,
callback=callback,
litellm_params=litellm_params,
standard_callback_dynamic_params = self.standard_callback_dynamic_params
standard_callback_dynamic_params=self.standard_callback_dynamic_params,
)
):
verbose_logger.debug(
@@ -2266,7 +2264,7 @@ class Logging(LiteLLMLoggingBaseClass):
start_time=start_time,
end_time=end_time,
)
if isinstance(callback, CustomLogger): # custom logger class
model_call_details: Dict = self.model_call_details
##################################
@@ -2276,10 +2274,7 @@ class Logging(LiteLLMLoggingBaseClass):
)
##################################
if self.stream is True:
if (
"async_complete_streaming_response"
in model_call_details
):
if "async_complete_streaming_response" in model_call_details:
await callback.async_log_success_event(
kwargs=model_call_details,
response_obj=model_call_details[
@@ -3217,11 +3212,10 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915
_literalai_logger = LiteralAILogger()
_in_memory_loggers.append(_literalai_logger)
return _literalai_logger # type: ignore
elif logging_integration == "prometheus":
if PrometheusLogger is not None:
for callback in _in_memory_loggers:
if isinstance(callback, PrometheusLogger):
return callback # type: ignore
elif logging_integration == "prometheus" and PrometheusLogger is not None:
for callback in _in_memory_loggers:
if isinstance(callback, PrometheusLogger):
return callback # type: ignore
_prometheus_logger = PrometheusLogger()
_in_memory_loggers.append(_prometheus_logger)
@@ -3493,7 +3487,7 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915
from litellm.integrations.vector_store_integrations.vector_store_pre_call_hook import (
VectorStorePreCallHook,
)
for callback in _in_memory_loggers:
if isinstance(callback, VectorStorePreCallHook):
return callback
@@ -3536,6 +3530,15 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915
humanloop_logger = HumanloopLogger()
_in_memory_loggers.append(humanloop_logger)
return humanloop_logger # type: ignore
elif logging_integration == "dotprompt":
for callback in _in_memory_loggers:
if isinstance(callback, DotpromptManager):
return callback
dotprompt_logger = DotpromptManager()
_in_memory_loggers.append(dotprompt_logger)
return dotprompt_logger # type: ignore
return None
except Exception as e:
verbose_logger.exception(
f"[Non-Blocking Error] Error initializing custom logger: {e}"
@@ -3582,11 +3585,10 @@ def get_custom_logger_compatible_class( # noqa: PLR0915
for callback in _in_memory_loggers:
if isinstance(callback, LiteralAILogger):
return callback
elif logging_integration == "prometheus":
if PrometheusLogger is not None:
for callback in _in_memory_loggers:
if isinstance(callback, PrometheusLogger):
return callback
elif logging_integration == "prometheus" and PrometheusLogger is not None:
for callback in _in_memory_loggers:
if isinstance(callback, PrometheusLogger):
return callback
elif logging_integration == "datadog":
for callback in _in_memory_loggers:
if isinstance(callback, DataDogLogger):
@@ -3686,7 +3688,7 @@ def get_custom_logger_compatible_class( # noqa: PLR0915
from litellm.integrations.vector_store_integrations.vector_store_pre_call_hook import (
VectorStorePreCallHook,
)
for callback in _in_memory_loggers:
if isinstance(callback, VectorStorePreCallHook):
return callback
@@ -1121,13 +1121,14 @@ def convert_to_gemini_tool_call_result(
}
"""
content_str: str = ""
if isinstance(message["content"], str):
content_str = message["content"]
elif isinstance(message["content"], List):
content_list = message["content"]
for content in content_list:
if content["type"] == "text":
content_str += content["text"]
if "content" in message:
if isinstance(message["content"], str):
content_str = message["content"]
elif isinstance(message["content"], List):
content_list = message["content"]
for content in content_list:
if content["type"] == "text":
content_str += content["text"]
name: Optional[str] = message.get("name", "") # type: ignore
# Recover name from last message with tool calls
+10 -5
View File
@@ -147,8 +147,8 @@ from .llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from .llms.custom_llm import CustomLLM, custom_chat_llm_router
from .llms.databricks.embed.handler import DatabricksEmbeddingHandler
from .llms.deprecated_providers import aleph_alpha, palm
from .llms.groq.chat.handler import GroqChatCompletion
from .llms.gemini.common_utils import get_api_key_from_env
from .llms.groq.chat.handler import GroqChatCompletion
from .llms.huggingface.embedding.handler import HuggingFaceEmbedding
from .llms.nlp_cloud.chat.handler import completion as nlp_cloud_chat_completion
from .llms.ollama.completion import handler as ollama
@@ -1049,11 +1049,13 @@ def completion( # type: ignore # noqa: PLR0915
non_default_params = get_non_default_completion_params(kwargs=kwargs)
litellm_params = {} # used to prevent unbound var errors
## PROMPT MANAGEMENT HOOKS ##
if isinstance(litellm_logging_obj, LiteLLMLoggingObj) and (
litellm_logging_obj.should_run_prompt_management_hooks(
prompt_id=prompt_id, non_default_params=non_default_params
)
):
(
model,
messages,
@@ -3999,9 +4001,7 @@ def embedding( # noqa: PLR0915
litellm_params={},
)
elif custom_llm_provider == "gemini":
gemini_api_key = (
api_key or get_api_key_from_env() or litellm.api_key
)
gemini_api_key = api_key or get_api_key_from_env() or litellm.api_key
api_base = api_base or litellm.api_base or get_secret_str("GEMINI_API_BASE")
@@ -5431,6 +5431,7 @@ def speech( # noqa: PLR0915
##### Health Endpoints #######################
async def ahealth_check(
model_params: dict,
mode: Optional[
@@ -5476,7 +5477,11 @@ async def ahealth_check(
log_raw_request_response=True,
)
model_params["litellm_logging_obj"] = litellm_logging_obj
model_params = HealthCheckHelpers._update_model_params_with_health_check_tracking_information(model_params=model_params)
model_params = (
HealthCheckHelpers._update_model_params_with_health_check_tracking_information(
model_params=model_params
)
)
#########################################################
try:
model: Optional[str] = model_params.get("model", None)
File diff suppressed because one or more lines are too long
+11
View File
@@ -1,4 +1,5 @@
model_list:
<<<<<<< HEAD
- model_name: "gpt-4o-mini-openai"
litellm_params:
model: gpt-4o-mini
@@ -6,3 +7,13 @@ model_list:
router_settings:
model_group_alias: {"gpt-4o": "gpt-4o-mini-openai"}
=======
- model_name: openai-test
litellm_params:
model: dotprompt/gpt-3.5-turbo
prompt_id: test_hello_world_prompt
api_key: os.environ/OPENAI_API_KEY
litellm_settings:
global_prompt_directory: /Users/krrishdholakia/Documents/litellm/litellm/proxy/test_prompts
>>>>>>> litellm_dev_07_31_2025_p1
+9
View File
@@ -1822,6 +1822,15 @@ 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,
)
set_global_prompt_directory(value)
verbose_proxy_logger.info(
f"{blue_color_code}Set Global Prompt Directory on LiteLLM Proxy{reset_color_code}"
)
elif key == "callbacks":
initialize_callbacks_on_proxy(
value=value,
@@ -0,0 +1,10 @@
---
model: gpt-3.5-turbo
input:
schema:
text: string
---
Extract the requested information from the given text. If a piece of information is not present, omit that field from the output.
Text: {{text}}
+1
View File
@@ -2319,6 +2319,7 @@ class LlmProviders(str, Enum):
HYPERBOLIC = "hyperbolic"
RECRAFT = "recraft"
AUTO_ROUTER = "auto_router"
DOTPROMPT = "dotprompt"
# Create a set of all provider values for quick lookup
@@ -0,0 +1,13 @@
---
model: gpt-4
temperature: 0.7
max_tokens: 150
input:
schema:
user_message: string
system_context?: string
---
{% if system_context %}System: {{system_context}}
{% endif %}User: {{user_message}}
@@ -0,0 +1,33 @@
---
model: claude-3-5-sonnet-20241022
temperature: 0.2
max_tokens: 2000
input:
schema:
language: string
task: string
code?: string
requirements?: array
output:
format: text
---
You are a helpful coding assistant. {% if language %}Focus on {{language}} programming.{% endif %}
Task: {{task}}
{% if code %}
Current code:
```{{language}}
{{code}}
```
{% endif %}
{% if requirements %}
Requirements:
{% for req in requirements %}
- {{req}}
{% endfor %}
{% endif %}
Please provide a clear and well-documented solution.
@@ -0,0 +1,16 @@
---
model: gemini/gemini-1.5-pro
input:
schema:
text: string
output:
format: json
schema:
title?: string, the title of the article if it has one
summary: string, a 3-sentence summary of the text
tags?(array, a list of string tag category for the text): string
---
Extract the requested information from the given text. If a piece of information is not present, omit that field from the output.
Text: {{text}}
@@ -0,0 +1,238 @@
import json
import os
import sys
import tempfile
from pathlib import Path
import pytest
from fastapi.testclient import TestClient
sys.path.insert(
0, os.path.abspath("../../..")
) # Adds the parent directory to the system path
from unittest.mock import MagicMock, patch
import litellm
from litellm.integrations.dotprompt import DotpromptManager
from litellm.types.utils import StandardCallbackDynamicParams
def test_dotprompt_manager_initialization():
"""Test basic DotpromptManager initialization."""
prompt_dir = "." # Current directory when running from tests/test_litellm/prompts
manager = DotpromptManager(prompt_dir)
assert manager.integration_name == "dotprompt"
assert manager.prompt_directory == prompt_dir
def test_should_run_prompt_management():
"""Test should_run_prompt_management method."""
prompt_dir = "."
manager = DotpromptManager(prompt_dir)
# Test with existing prompt
assert (
manager.should_run_prompt_management(
"sample_prompt", StandardCallbackDynamicParams()
)
== True
)
# Test with non-existing prompt
assert (
manager.should_run_prompt_management(
"nonexistent_prompt", StandardCallbackDynamicParams()
)
== False
)
def test_convert_to_messages_simple():
"""Test converting simple text to messages."""
prompt_dir = "."
manager = DotpromptManager(prompt_dir)
# Test simple text
messages = manager._convert_to_messages("Hello world!")
assert len(messages) == 1
assert messages[0]["role"] == "user"
assert messages[0]["content"] == "Hello world!"
def test_convert_to_messages_with_roles():
"""Test converting text with role prefixes to messages."""
prompt_dir = "."
manager = DotpromptManager(prompt_dir)
# Test text with role prefixes
content = """System: You are a helpful assistant.
User: What is the capital of France?"""
messages = manager._convert_to_messages(content)
assert len(messages) == 2
assert messages[0]["role"] == "system"
assert messages[0]["content"] == "You are a helpful assistant."
assert messages[1]["role"] == "user"
assert messages[1]["content"] == "What is the capital of France?"
def test_compile_prompt_helper():
"""Test the _compile_prompt_helper method."""
prompt_dir = "."
manager = DotpromptManager(prompt_dir)
# Test compiling a simple prompt
result = manager._compile_prompt_helper(
prompt_id="sample_prompt",
prompt_variables={"text": "This is a test article."},
dynamic_callback_params=StandardCallbackDynamicParams(),
)
assert result["prompt_id"] == "sample_prompt"
assert result["prompt_template_model"] == "gemini/gemini-1.5-pro"
assert len(result["prompt_template"]) >= 1
assert "This is a test article." in result["prompt_template"][0]["content"]
def test_compile_prompt_helper_with_chat_format():
"""Test compiling a prompt that generates role-based messages."""
prompt_dir = "."
manager = DotpromptManager(prompt_dir)
# Test with chat_prompt that has system context
result = manager._compile_prompt_helper(
prompt_id="chat_prompt",
prompt_variables={
"user_message": "Hello there!",
"system_context": "You are a helpful assistant.",
},
dynamic_callback_params=StandardCallbackDynamicParams(),
)
assert result["prompt_id"] == "chat_prompt"
assert result["prompt_template_model"] == "gpt-4"
assert len(result["prompt_template"]) == 2
# Should have system message first
assert result["prompt_template"][0]["role"] == "system"
assert "You are a helpful assistant." in result["prompt_template"][0]["content"]
# Then user message
assert result["prompt_template"][1]["role"] == "user"
assert "Hello there!" in result["prompt_template"][1]["content"]
def test_extract_optional_params():
"""Test extracting optional parameters from template metadata."""
prompt_dir = "."
manager = DotpromptManager(prompt_dir)
# Get a template with optional params
template = manager.prompt_manager.get_prompt("chat_prompt")
params = manager._extract_optional_params(template)
assert "temperature" in params
assert params["temperature"] == 0.7
assert "max_tokens" in params
assert params["max_tokens"] == 150
def test_error_handling():
"""Test error handling for invalid prompts."""
prompt_dir = "."
manager = DotpromptManager(prompt_dir)
# Test with non-existent prompt
with pytest.raises(ValueError, match="Prompt 'nonexistent' not found"):
manager._compile_prompt_helper(
prompt_id="nonexistent",
prompt_variables={},
dynamic_callback_params=StandardCallbackDynamicParams(),
)
def test_integration_with_prompt_management():
"""Test integration with the prompt management system."""
with tempfile.TemporaryDirectory() as temp_dir:
# Create a test prompt
prompt_file = Path(temp_dir) / "test_integration.prompt"
prompt_file.write_text(
"""---
model: gpt-3.5-turbo
temperature: 0.5
---
System: You are a {{role}}.
User: {{question}}"""
)
manager = DotpromptManager(temp_dir)
# Test should_run_prompt_management
assert (
manager.should_run_prompt_management(
"test_integration", StandardCallbackDynamicParams()
)
== True
)
# Test compile_prompt_helper
result = manager._compile_prompt_helper(
prompt_id="test_integration",
prompt_variables={"role": "helpful assistant", "question": "What is AI?"},
dynamic_callback_params=StandardCallbackDynamicParams(),
)
assert result["prompt_template_model"] == "gpt-3.5-turbo"
assert result["prompt_template_optional_params"]["temperature"] == 0.5
assert len(result["prompt_template"]) == 2
assert result["prompt_template"][0]["role"] == "system"
assert "helpful assistant" in result["prompt_template"][0]["content"]
assert result["prompt_template"][1]["role"] == "user"
assert "What is AI?" in result["prompt_template"][1]["content"]
def test_set_prompt_directory():
"""Test setting and changing prompt directory."""
with tempfile.TemporaryDirectory() as temp_dir:
manager = DotpromptManager(temp_dir)
# Initially should be empty
assert not manager.should_run_prompt_management(
"test_prompt", StandardCallbackDynamicParams()
)
# Create a prompt file
prompt_file = Path(temp_dir) / "test_prompt.prompt"
prompt_file.write_text("Hello {{name}}!")
# Set directory to force reload
manager.set_prompt_directory(temp_dir)
# Now should find the prompt
assert manager.should_run_prompt_management(
"test_prompt", StandardCallbackDynamicParams()
)
def test_no_prompt_directory_error():
"""Test error when no prompt directory is set."""
manager = DotpromptManager(None)
# should_run_prompt_management returns False when there's an error
result = manager.should_run_prompt_management(
"any_prompt", StandardCallbackDynamicParams()
)
assert result == False
# But accessing prompt_manager property should raise an error
with pytest.raises(ValueError, match="prompt_directory must be set"):
_ = manager.prompt_manager
@@ -0,0 +1,246 @@
import json
import os
import sys
import tempfile
from pathlib import Path
import pytest
from fastapi.testclient import TestClient
sys.path.insert(
0, os.path.abspath("../../..")
) # Adds the parent directory to the system path
from unittest.mock import MagicMock, patch
import litellm
from litellm.integrations.dotprompt.prompt_manager import PromptManager, PromptTemplate
def test_prompt_manager_initialization():
"""Test basic PromptManager initialization and loading."""
# Test with the existing prompts directory
prompt_dir = "." # Current directory when running from tests/test_litellm/prompts
manager = PromptManager(prompt_dir)
# Should have loaded at least the sample prompts
assert len(manager.prompts) >= 3
assert "sample_prompt" in manager.prompts
assert "chat_prompt" in manager.prompts
assert "coding_assistant" in manager.prompts
def test_prompt_template_creation():
"""Test PromptTemplate creation and metadata extraction."""
metadata = {
"model": "gpt-4",
"temperature": 0.7,
"input": {"schema": {"text": "string"}},
"output": {"format": "json"},
}
template = PromptTemplate(
content="Hello {{name}}!", metadata=metadata, template_id="test_template"
)
assert template.content == "Hello {{name}}!"
assert template.model == "gpt-4"
assert template.optional_params["temperature"] == 0.7
assert template.input_schema == {"text": "string"}
assert template.output_format == "json"
def test_render_simple_template():
"""Test rendering a simple template with variables."""
prompt_dir = "." # Current directory when running from tests/test_litellm/prompts
manager = PromptManager(prompt_dir)
# Test sample_prompt rendering
rendered = manager.render(
"sample_prompt", {"text": "This is a test article about AI."}
)
expected_content = "Extract the requested information from the given text. If a piece of information is not present, omit that field from the output.\n\nText: This is a test article about AI."
assert rendered == expected_content
def test_render_chat_prompt():
"""Test rendering the chat prompt with conditional content."""
prompt_dir = "." # Current directory when running from tests/test_litellm/prompts
manager = PromptManager(prompt_dir)
# Test with system context
rendered = manager.render(
"chat_prompt",
{
"user_message": "Hello there!",
"system_context": "You are a helpful assistant.",
},
)
assert "System: You are a helpful assistant." in rendered
assert "User: Hello there!" in rendered
# Test without system context
rendered_no_system = manager.render("chat_prompt", {"user_message": "Hello there!"})
assert "System:" not in rendered_no_system
assert "User: Hello there!" in rendered_no_system
def test_render_coding_assistant():
"""Test rendering the coding assistant prompt with complex logic."""
prompt_dir = "." # Current directory when running from tests/test_litellm/prompts
manager = PromptManager(prompt_dir)
rendered = manager.render(
"coding_assistant",
{
"language": "Python",
"task": "Create a function to calculate fibonacci numbers",
"code": "def fib(n):\n pass",
"requirements": ["Use recursion", "Handle edge cases", "Add documentation"],
},
)
assert "Focus on Python programming." in rendered
assert "Create a function to calculate fibonacci numbers" in rendered
assert "def fib(n):" in rendered
assert "Use recursion" in rendered
assert "Handle edge cases" in rendered
assert "Add documentation" in rendered
def test_input_validation():
"""Test input validation against schema."""
# Create a temporary directory with a test prompt
with tempfile.TemporaryDirectory() as temp_dir:
prompt_file = Path(temp_dir) / "test_validation.prompt"
prompt_file.write_text(
"""---
input:
schema:
name: string
age: integer
active: boolean
---
Hello {{name}}, you are {{age}} years old and {'active' if active else 'inactive'}."""
)
manager = PromptManager(temp_dir)
# Valid input should work
rendered = manager.render(
"test_validation", {"name": "Alice", "age": 30, "active": True}
)
assert "Hello Alice, you are 30 years old" in rendered
# Invalid type should raise error
with pytest.raises(ValueError, match="Invalid type for field 'age'"):
manager.render(
"test_validation",
{
"name": "Alice",
"age": "thirty", # string instead of int
"active": True,
},
)
def test_prompt_not_found():
"""Test error handling for non-existent prompts."""
prompt_dir = "." # Current directory when running from tests/test_litellm/prompts
manager = PromptManager(prompt_dir)
with pytest.raises(KeyError, match="Prompt 'nonexistent' not found"):
manager.render("nonexistent", {"some": "variable"})
def test_list_prompts():
"""Test listing available prompts."""
prompt_dir = "." # Current directory when running from tests/test_litellm/prompts
manager = PromptManager(prompt_dir)
prompts = manager.list_prompts()
assert isinstance(prompts, list)
assert "sample_prompt" in prompts
assert "chat_prompt" in prompts
assert "coding_assistant" in prompts
def test_get_prompt_metadata():
"""Test retrieving prompt metadata."""
prompt_dir = "." # Current directory when running from tests/test_litellm/prompts
manager = PromptManager(prompt_dir)
metadata = manager.get_prompt_metadata("sample_prompt")
assert metadata is not None
assert metadata["model"] == "gemini/gemini-1.5-pro"
assert "input" in metadata
assert "output" in metadata
def test_add_prompt_programmatically():
"""Test adding prompts programmatically."""
prompt_dir = "." # Current directory when running from tests/test_litellm/prompts
manager = PromptManager(prompt_dir)
initial_count = len(manager.prompts)
manager.add_prompt(
"dynamic_prompt",
"Hello {{name}}! Welcome to {{place}}.",
{"model": "gpt-3.5-turbo", "temperature": 0.5},
)
assert len(manager.prompts) == initial_count + 1
assert "dynamic_prompt" in manager.prompts
rendered = manager.render("dynamic_prompt", {"name": "World", "place": "Earth"})
assert rendered == "Hello World! Welcome to Earth."
def test_frontmatter_parsing():
"""Test YAML frontmatter parsing."""
# Create a temporary directory with a test prompt
with tempfile.TemporaryDirectory() as temp_dir:
# Test with frontmatter
prompt_with_frontmatter = Path(temp_dir) / "with_frontmatter.prompt"
prompt_with_frontmatter.write_text(
"""---
model: gpt-4
temperature: 0.8
input:
schema:
topic: string
---
Write about {{topic}}."""
)
# Test without frontmatter
prompt_without_frontmatter = Path(temp_dir) / "without_frontmatter.prompt"
prompt_without_frontmatter.write_text("Simple template: {{message}}")
manager = PromptManager(temp_dir)
# Check frontmatter was parsed correctly
with_meta = manager.get_prompt("with_frontmatter")
assert with_meta.model == "gpt-4"
assert with_meta.optional_params["temperature"] == 0.8
# Check template without frontmatter still works
without_meta = manager.get_prompt("without_frontmatter")
assert without_meta.metadata == {}
rendered = manager.render("without_frontmatter", {"message": "Hello!"})
assert rendered == "Simple template: Hello!"
def test_prompt_main():
"""
Integration test placeholder for litellm completion integration.
This would be implemented once the PromptManager is integrated with litellm.
"""
# TODO: Implement once PromptManager is integrated with litellm completion
pass
@@ -87,44 +87,33 @@ def test_convert_to_azure_openai_messages():
"""Test coverting image_url to azure_openai spec"""
from typing import List
from litellm.litellm_core_utils.prompt_templates.factory import (
convert_to_azure_openai_messages,
)
from litellm.types.llms.openai import AllMessageValues
from litellm.litellm_core_utils.prompt_templates.factory import convert_to_azure_openai_messages
input: List[AllMessageValues] = [
{
"role": "user",
"content": [
{
"type": "text",
"text": "What is in this image?"
},
{
"type": "image_url",
"image_url": "www.mock.com"
}
]
{"type": "text", "text": "What is in this image?"},
{"type": "image_url", "image_url": "www.mock.com"},
],
}
]
expected_content = [
{
"type": "text",
"text": "What is in this image?"
},
{
"type": "image_url",
"image_url": {"url": "www.mock.com"}
}
{"type": "text", "text": "What is in this image?"},
{"type": "image_url", "image_url": {"url": "www.mock.com"}},
]
output = convert_to_azure_openai_messages(input)
content = output[0].get('content')
content = output[0].get("content")
assert content == expected_content
def test_bedrock_validate_format_image_or_video():
"""Test the _validate_format method for images, videos, and documents"""
@@ -405,12 +394,44 @@ def test_unpack_defs_resolves_nested_ref_inside_anyof_items():
unpack_defs(schema, schema["$defs"])
# Extract the items schema after unpacking
items_schema = (
schema["properties"]["vatAmounts"]["anyOf"][0]["items"]
)
items_schema = schema["properties"]["vatAmounts"]["anyOf"][0]["items"]
# Assertions: items_schema should now be the resolved object, not an empty dict
assert isinstance(items_schema, dict), "Items schema should be a dict after unpacking"
assert isinstance(
items_schema, dict
), "Items schema should be a dict after unpacking"
assert items_schema.get("type") == "object"
# Ensure essential properties are present
assert set(items_schema.get("properties", {}).keys()) == {"vatRate", "vatAmount"}
def test_convert_gemini_messages():
"""
Handle 'content' not being present in the message - https://github.com/BerriAI/litellm/issues/13169
"""
from litellm.litellm_core_utils.prompt_templates.factory import (
convert_to_gemini_tool_call_result,
)
from litellm.types.llms.openai import ChatCompletionToolMessage
message = ChatCompletionToolMessage(
role="tool",
tool_call_id="call_d5b2e3fe-d2c0-451d-b034-cf4fbb22e66c",
)
last_message_with_tool_calls = {
"role": "assistant",
"content": "",
"tool_calls": [
{
"id": "call_d5b2e3fe-d2c0-451d-b034-cf4fbb22e66c",
"type": "function",
"index": 0,
"function": {"name": "tool_MAX_Data__get_issues", "arguments": "{}"},
}
],
}
convert_to_gemini_tool_call_result(
message=message,
last_message_with_tool_calls=last_message_with_tool_calls,
)