Squashed commit of the following:

commit 3119064e94
Author: Krish Dholakia <krrishdholakia@gmail.com>
Date:   Sat Aug 2 22:33:37 2025 -0700

    Prompt Management - add prompts on UI  (#13240)

    * fix(create_key_button.tsx): add prompts on UI

    * feat(key_management_endpoints.py): support adding prompt to key via `/key/update`

    * fix(key_info_view.tsx): show existing prompts on key in key_info_view.tsx

    * fix(key_edit_view.tsx): UX - disable premium feature for non-premium users

    prevent accidental clicking

    * fix(create_key_button.tsx): disable premium features behind flag, prevent errors

    * feat(prompts.tsx): add new ui component to view created prompts

    enables viewing prompts created on config

    * feat(prompt_info.tsx): add component for viewing the prompt information

    * feat(prompt_endpoints.py): support converting dotprompt to json structure + accept json structure in promptmanager

    allows prompt manager to work with api endpoints

    * test(test_prompt_manager.py): add unit tests for json data input

    * feat(dotprompt/__init__.py): add prompt data to dotpromptmanager

    * fix(prompt_endpoints.py): working crud endpoints for prompt management

    * feat(prompts/): support `prompt_file` for dotprompt

    allows to precisely point to the prompt file a prompt should use

    * feat(proxy/utils.py): resolve prompt id correctly

    resolves user sent prompt id with internal prompt id

    * feat(schema.prisma): initial pr with db schema for prompt management table

    allows post endpoints to work with backend

    * feat(prompt_endpoints.py): use db in patch_prompt endpoint

    * feat(prompt_endpoints.py): use db for update_prompt endpoint

    * feat(prompt_endpoints.py): use db on prompt delete endpoint

    * build(schema.prisma): add prompt tale to schema.prisma in litellm-proxy-extras

    * build(migration.sql): add new sql migration file

    * fix(init_prompts.py): fix init

    * feat(prompt_info_view.tsx): show the raw prompt template on ui

    allows developer to know the prompt template they'll be calling

    * feat(add_prompt_form.tsx): working ui add prompt flow

    allows user to add prompts to litellm via ui

    * build(ui/): styling fixes

    * build(ui/): prompts.tsx

    styling improvements

    * fix(add_prompt_form.tsx): styling improvements

    * build(prompts.tsx): styling improvements

    * build(ui/): styling improvements

    * build(ui/): fix ui error

    * fix: fix ruff check

    * docs: document new api params

    * test: update tests

commit e47b30a76d
Author: Krish Dholakia <krrishdholakia@gmail.com>
Date:   Sat Aug 2 19:34:18 2025 -0700

    Prompt Management - Add table + prompt info page to UI  (#13232)

    * fix(create_key_button.tsx): add prompts on UI

    * feat(key_management_endpoints.py): support adding prompt to key via `/key/update`

    * fix(key_info_view.tsx): show existing prompts on key in key_info_view.tsx

    * fix(key_edit_view.tsx): UX - disable premium feature for non-premium users

    prevent accidental clicking

    * fix(create_key_button.tsx): disable premium features behind flag, prevent errors

    * feat(prompts.tsx): add new ui component to view created prompts

    enables viewing prompts created on config

    * feat(prompt_info.tsx): add component for viewing the prompt information

commit b79f55eec0
Author: Krish Dholakia <krrishdholakia@gmail.com>
Date:   Sat Aug 2 19:33:17 2025 -0700

    UI - Add giving keys prompt access (#13233)

    * fix(create_key_button.tsx): add prompts on UI

    * feat(key_management_endpoints.py): support adding prompt to key via `/key/update`

    * fix(key_info_view.tsx): show existing prompts on key in key_info_view.tsx

    * fix(key_edit_view.tsx): UX - disable premium feature for non-premium users

    prevent accidental clicking

    * fix(create_key_button.tsx): disable premium features behind flag, prevent errors

    * fix(key_management_endpoints.py): fix key update logic \

    * fix: fix check

    * docs: document new params

commit 4c217c66f5
Author: Ishaan Jaff <ishaanjaffer0324@gmail.com>
Date:   Sat Aug 2 17:26:40 2025 -0700

    docs User Agent Activity Tracking

commit 2ee4e84406
Author: Ishaan Jaff <ishaanjaffer0324@gmail.com>
Date:   Sat Aug 2 16:47:44 2025 -0700

    docs fix

commit 06856b4d37
Author: Ishaan Jaff <ishaanjaffer0324@gmail.com>
Date:   Sat Aug 2 15:47:09 2025 -0700

    docs fix

commit 0f9f5f7a6c
Author: Ishaan Jaff <ishaanjaffer0324@gmail.com>
Date:   Sat Aug 2 15:44:59 2025 -0700

    docs fix

commit 69a360429c
Author: Ishaan Jaff <ishaanjaffer0324@gmail.com>
Date:   Sat Aug 2 15:34:14 2025 -0700

    agent 4.png

commit e32169dc37
Author: Ishaan Jaff <ishaanjaffer0324@gmail.com>
Date:   Sat Aug 2 15:29:44 2025 -0700

    docs cost tracking coding

commit 340b64a46a
Author: Ishaan Jaff <ishaanjaffer0324@gmail.com>
Date:   Sat Aug 2 15:18:04 2025 -0700

    docs - Track Usage for Coding Tools

commit 9b029c35be
Author: Ishaan Jaff <ishaanjaffer0324@gmail.com>
Date:   Sat Aug 2 15:06:12 2025 -0700

    docs RC

commit 8d6b333909
Author: Ishaan Jaff <ishaanjaffer0324@gmail.com>
Date:   Sat Aug 2 15:00:48 2025 -0700

    docs computer use

commit e306fb6eee
Author: Ishaan Jaff <ishaanjaffer0324@gmail.com>
Date:   Sat Aug 2 14:15:10 2025 -0700

    [docs release notes] (#13237)

    * docs release notes

    * docs release notes

    * docs rnotes

    * docs api version

    * fixes docs

    * docs rn

commit 5dfc88473f
Author: Ishaan Jaff <ishaanjaffer0324@gmail.com>
Date:   Sat Aug 2 13:30:51 2025 -0700

    fixes MCP gateway docs

commit 6929767be2
Author: Ishaan Jaff <ishaanjaffer0324@gmail.com>
Date:   Sat Aug 2 12:56:00 2025 -0700

    docs release notes
This commit is contained in:
Krrish Dholakia 2025-08-05 21:59:54 -07:00
parent 51627ad543
commit 0d0a65a98f
26 changed files with 2164 additions and 324 deletions

View file

@ -0,0 +1,15 @@
-- CreateTable
CREATE TABLE "LiteLLM_PromptTable" (
"id" TEXT NOT NULL,
"prompt_id" TEXT NOT NULL,
"litellm_params" JSONB NOT NULL,
"prompt_info" JSONB,
"created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
"updated_at" TIMESTAMP(3) NOT NULL,
CONSTRAINT "LiteLLM_PromptTable_pkey" PRIMARY KEY ("id")
);
-- CreateIndex
CREATE UNIQUE INDEX "LiteLLM_PromptTable_prompt_id_key" ON "LiteLLM_PromptTable"("prompt_id");

View file

@ -520,6 +520,16 @@ model LiteLLM_GuardrailsTable {
updated_at DateTime @updatedAt
}
// Prompt table for storing prompt configurations
model LiteLLM_PromptTable {
id String @id @default(uuid())
prompt_id String @unique
litellm_params Json
prompt_info Json?
created_at DateTime @default(now())
updated_at DateTime @updatedAt
}
model LiteLLM_HealthCheckTable {
health_check_id String @id @default(uuid())
model_name String

View file

@ -33,10 +33,27 @@ def prompt_initializer(
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")
prompt_data = getattr(litellm_params, "prompt_data", None)
prompt_id = getattr(litellm_params, "prompt_id", None)
if prompt_directory:
raise ValueError(
"Cannot set prompt_directory when working with prompt_initializer. Needs to be a specific dotprompt file"
)
return DotpromptManager(prompt_directory)
prompt_file = getattr(litellm_params, "prompt_file", None)
try:
dot_prompt_manager = DotpromptManager(
prompt_directory=prompt_directory,
prompt_data=prompt_data,
prompt_file=prompt_file,
prompt_id=prompt_id,
)
return dot_prompt_manager
except Exception as e:
raise e
prompt_initializer_registry = {

View file

@ -3,7 +3,8 @@ 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, Tuple
import json
from typing import Any, Dict, List, Optional, Tuple, Union
from litellm.integrations.custom_prompt_management import CustomPromptManagement
from litellm.integrations.prompt_management_base import PromptManagementClient
@ -33,12 +34,25 @@ class DotpromptManager(CustomPromptManagement):
)
"""
def __init__(self, prompt_directory: Optional[str] = None):
def __init__(
self,
prompt_directory: Optional[str] = None,
prompt_file: Optional[str] = None,
prompt_data: Optional[Union[dict, str]] = None,
prompt_id: Optional[str] = None,
):
import litellm
self.prompt_directory = prompt_directory or litellm.global_prompt_directory
# Support for JSON-based prompts stored in memory/database
if isinstance(prompt_data, str):
self.prompt_data = json.loads(prompt_data)
else:
self.prompt_data = prompt_data or {}
self._prompt_manager: Optional[PromptManager] = None
self.prompt_file = prompt_file
self.prompt_id = prompt_id
@property
def integration_name(self) -> str:
@ -49,12 +63,21 @@ class DotpromptManager(CustomPromptManagement):
def prompt_manager(self) -> PromptManager:
"""Lazy-load the prompt manager."""
if self._prompt_manager is None:
if self.prompt_directory is None:
if (
self.prompt_directory is None
and not self.prompt_data
and not self.prompt_file
):
raise ValueError(
"prompt_directory must be set before using dotprompt manager. "
"Set litellm.global_prompt_directory or initialize with prompt_directory parameter."
"Either prompt_directory or prompt_data must be set before using dotprompt manager. "
"Set litellm.global_prompt_directory, initialize with prompt_directory parameter, or provide prompt_data."
)
self._prompt_manager = PromptManager(self.prompt_directory)
self._prompt_manager = PromptManager(
prompt_directory=self.prompt_directory,
prompt_data=self.prompt_data,
prompt_file=self.prompt_file,
prompt_id=self.prompt_id,
)
return self._prompt_manager
def should_run_prompt_management(
@ -92,6 +115,7 @@ class DotpromptManager(CustomPromptManagement):
"""
try:
# Get the prompt template
template = self.prompt_manager.get_prompt(prompt_id)
if template is None:
@ -131,6 +155,7 @@ class DotpromptManager(CustomPromptManagement):
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(
@ -246,3 +271,21 @@ class DotpromptManager(CustomPromptManagement):
"""Reload all prompts from the directory."""
if self._prompt_manager:
self._prompt_manager.reload_prompts()
def add_prompt_from_json(self, prompt_id: str, json_data: Dict[str, Any]) -> None:
"""Add a prompt from JSON data."""
content = json_data.get("content", "")
metadata = json_data.get("metadata", {})
self.prompt_manager.add_prompt(prompt_id, content, metadata)
def load_prompts_from_json(self, prompts_data: Dict[str, Dict[str, Any]]) -> None:
"""Load multiple prompts from JSON data."""
self.prompt_manager.load_prompts_from_json_data(prompts_data)
def get_prompts_as_json(self) -> Dict[str, Dict[str, Any]]:
"""Get all prompts in JSON format."""
return self.prompt_manager.get_all_prompts_as_json()
def convert_prompt_file_to_json(self, file_path: str) -> Dict[str, Any]:
"""Convert a .prompt file to JSON format."""
return self.prompt_manager.prompt_file_to_json(file_path)

View file

@ -49,9 +49,16 @@ class PromptManager:
- Model configuration
"""
def __init__(self, prompt_directory: str):
self.prompt_directory = Path(prompt_directory)
def __init__(
self,
prompt_id: Optional[str] = None,
prompt_directory: Optional[str] = None,
prompt_data: Optional[Dict[str, Dict[str, Any]]] = None,
prompt_file: Optional[str] = None,
):
self.prompt_directory = Path(prompt_directory) if prompt_directory else None
self.prompts: Dict[str, PromptTemplate] = {}
self.prompt_file = prompt_file
self.jinja_env = Environment(
loader=DictLoader({}),
autoescape=select_autoescape(["html", "xml"]),
@ -64,12 +71,24 @@ class PromptManager:
comment_end_string="#}",
)
# Load all prompts in the directory
self._load_prompts()
# Load prompts from directory if provided
if self.prompt_directory:
self._load_prompts()
if self.prompt_file:
if not prompt_id:
raise ValueError("prompt_id is required when prompt_file is provided")
template = self._load_prompt_file(self.prompt_file, prompt_id)
self.prompts[prompt_id] = template
# Load prompts from JSON data if provided
if prompt_data:
self._load_prompts_from_json(prompt_data, prompt_id)
def _load_prompts(self) -> None:
"""Load all .prompt files from the prompt directory."""
if not self.prompt_directory.exists():
if not self.prompt_directory or not self.prompt_directory.exists():
raise ValueError(
f"Prompt directory does not exist: {self.prompt_directory}"
)
@ -86,8 +105,51 @@ class PromptManager:
# Optional: print(f"Error loading prompt file {prompt_file}")
pass
def _load_prompt_file(self, file_path: Path, prompt_id: str) -> PromptTemplate:
def _load_prompts_from_json(
self, prompt_data: Dict[str, Dict[str, Any]], prompt_id: Optional[str] = None
) -> None:
"""Load prompts from JSON data structure.
Expected format:
{
"prompt_id": {
"content": "template content",
"metadata": {"model": "gpt-4", "temperature": 0.7, ...}
}
}
or
{
"content": "template content",
"metadata": {"model": "gpt-4", "temperature": 0.7, ...}
} + prompt_id
"""
if prompt_id:
prompt_data = {prompt_id: prompt_data}
for prompt_id, prompt_info in prompt_data.items():
try:
content = prompt_info.get("content", "")
metadata = prompt_info.get("metadata", {})
template = PromptTemplate(
content=content,
metadata=metadata,
template_id=prompt_id,
)
self.prompts[prompt_id] = template
except Exception:
# Optional: print(f"Error loading prompt from JSON: {prompt_id}")
pass
def _load_prompt_file(
self, file_path: Union[str, Path], prompt_id: str
) -> PromptTemplate:
"""Load and parse a single .prompt file."""
if isinstance(file_path, str):
file_path = Path(file_path)
content = file_path.read_text(encoding="utf-8")
# Split frontmatter and content
@ -206,9 +268,10 @@ class PromptManager:
return template.metadata if template else None
def reload_prompts(self) -> None:
"""Reload all prompts from the directory."""
"""Reload all prompts from the directory (if directory was provided)."""
self.prompts.clear()
self._load_prompts()
if self.prompt_directory:
self._load_prompts()
def add_prompt(
self, prompt_id: str, content: str, metadata: Optional[Dict[str, Any]] = None
@ -218,3 +281,63 @@ class PromptManager:
content=content, metadata=metadata or {}, template_id=prompt_id
)
self.prompts[prompt_id] = template
def prompt_file_to_json(self, file_path: Union[str, Path]) -> Dict[str, Any]:
"""Convert a .prompt file to JSON format.
Args:
file_path: Path to the .prompt file
Returns:
Dictionary with 'content' and 'metadata' keys
"""
file_path = Path(file_path)
content = file_path.read_text(encoding="utf-8")
# Parse frontmatter and content
frontmatter, template_content = self._parse_frontmatter(content)
return {"content": template_content.strip(), "metadata": frontmatter}
def json_to_prompt_file(self, prompt_data: Dict[str, Any]) -> str:
"""Convert JSON prompt data to .prompt file format.
Args:
prompt_data: Dictionary with 'content' and 'metadata' keys
Returns:
String content in .prompt file format
"""
content = prompt_data.get("content", "")
metadata = prompt_data.get("metadata", {})
if not metadata:
# No metadata, return just the content
return content
# Convert metadata to YAML frontmatter
import yaml
frontmatter_yaml = yaml.dump(metadata, default_flow_style=False)
return f"---\n{frontmatter_yaml}---\n{content}"
def get_all_prompts_as_json(self) -> Dict[str, Dict[str, Any]]:
"""Get all loaded prompts in JSON format.
Returns:
Dictionary mapping prompt_id to prompt data
"""
result = {}
for prompt_id, template in self.prompts.items():
result[prompt_id] = {
"content": template.content,
"metadata": template.metadata,
}
return result
def load_prompts_from_json_data(
self, prompt_data: Dict[str, Dict[str, Any]]
) -> None:
"""Load additional prompts from JSON data (merges with existing prompts)."""
self._load_prompts_from_json(prompt_data)

View file

@ -54,6 +54,7 @@ class PromptManagementBase(ABC):
prompt_label: Optional[str] = None,
prompt_version: Optional[int] = None,
) -> PromptManagementClient:
compiled_prompt_client = self._compile_prompt_helper(
prompt_id=prompt_id,
prompt_variables=prompt_variables,
@ -91,6 +92,7 @@ class PromptManagementBase(ABC):
prompt_label: Optional[str] = None,
prompt_version: Optional[int] = None,
) -> Tuple[str, List[AllMessageValues], dict]:
if prompt_id is None:
raise ValueError("prompt_id is required for Prompt Management Base class")
if not self.should_run_prompt_management(

View file

@ -49,16 +49,16 @@ def _deserialize_env_dict(env_data: Any) -> Optional[Dict[str, str]]:
"""
Helper function to deserialize environment dictionary from database storage.
Handles both JSON string and dictionary formats.
Args:
env_data: The environment data from database (could be JSON string or dict)
Returns:
Dict[str, str] or None: Deserialized environment dictionary
"""
if not env_data:
return None
if isinstance(env_data, str):
try:
return json.loads(env_data)
@ -70,31 +70,35 @@ def _deserialize_env_dict(env_data: Any) -> Optional[Dict[str, str]]:
return env_data
def _convert_protocol_version_to_enum(protocol_version: Optional[str | MCPSpecVersionType]) -> MCPSpecVersionType:
def _convert_protocol_version_to_enum(
protocol_version: Optional[str | MCPSpecVersionType],
) -> MCPSpecVersionType:
"""
Convert string protocol version to MCPSpecVersion enum.
Args:
protocol_version: String protocol version, enum, or None
Returns:
MCPSpecVersionType: The enum value
"""
if not protocol_version:
return cast(MCPSpecVersionType, MCPSpecVersion.jun_2025)
# If it's already an MCPSpecVersion enum, return it
if isinstance(protocol_version, MCPSpecVersion):
return cast(MCPSpecVersionType, protocol_version)
# If it's a string, try to match it to enum values
if isinstance(protocol_version, str):
for version in MCPSpecVersion:
if version.value == protocol_version:
return cast(MCPSpecVersionType, version)
# If no match found, return default
verbose_logger.warning(f"Unknown protocol version '{protocol_version}', using default")
verbose_logger.warning(
f"Unknown protocol version '{protocol_version}', using default"
)
return cast(MCPSpecVersionType, MCPSpecVersion.jun_2025)
@ -132,70 +136,86 @@ class MCPServerManager:
"""
return self.config_mcp_servers | self.registry
def load_servers_from_config(self, mcp_servers_config: Dict[str, Any], mcp_aliases: Optional[Dict[str, str]] = None):
def load_servers_from_config(
self,
mcp_servers_config: Dict[str, Any],
mcp_aliases: Optional[Dict[str, str]] = None,
):
"""
Load the MCP Servers from the config
Args:
mcp_servers_config: Dictionary of MCP server configurations
mcp_aliases: Optional dictionary mapping aliases to server names from litellm_settings
"""
verbose_logger.debug("Loading MCP Servers from config-----")
# Track which aliases have been used to ensure only first occurrence is used
used_aliases = set()
for server_name, server_config in mcp_servers_config.items():
validate_mcp_server_name(server_name)
_mcp_info: Dict[str, Any] = server_config.get("mcp_info", None) or {}
# Convert Dict[str, Any] to MCPInfo properly
mcp_info: MCPInfo = {
"server_name": _mcp_info.get("server_name", server_name),
"description": _mcp_info.get("description", server_config.get("description", None)),
"description": _mcp_info.get(
"description", server_config.get("description", None)
),
"logo_url": _mcp_info.get("logo_url", None),
"mcp_server_cost_info": _mcp_info.get("mcp_server_cost_info", None),
}
# Use alias for name if present, else server_name
alias = server_config.get("alias", None)
# Apply mcp_aliases mapping if provided
if mcp_aliases and alias is None:
# Check if this server_name has an alias in mcp_aliases
for alias_name, target_server_name in mcp_aliases.items():
if target_server_name == server_name and alias_name not in used_aliases:
if (
target_server_name == server_name
and alias_name not in used_aliases
):
alias = alias_name
used_aliases.add(alias_name)
verbose_logger.debug(f"Mapped alias '{alias_name}' to server '{server_name}'")
verbose_logger.debug(
f"Mapped alias '{alias_name}' to server '{server_name}'"
)
break
# Create a temporary server object to use with get_server_prefix utility
temp_server = type('TempServer', (), {
'alias': alias,
'server_name': server_name,
'server_id': None
})()
temp_server = type(
"TempServer",
(),
{"alias": alias, "server_name": server_name, "server_id": None},
)()
name_for_prefix = get_server_prefix(temp_server)
# Use alias for name if present, else server_name
alias = server_config.get("alias", None)
# Apply mcp_aliases mapping if provided
if mcp_aliases and alias is None:
# Check if this server_name has an alias in mcp_aliases
for alias_name, target_server_name in mcp_aliases.items():
if target_server_name == server_name and alias_name not in used_aliases:
if (
target_server_name == server_name
and alias_name not in used_aliases
):
alias = alias_name
used_aliases.add(alias_name)
verbose_logger.debug(f"Mapped alias '{alias_name}' to server '{server_name}'")
verbose_logger.debug(
f"Mapped alias '{alias_name}' to server '{server_name}'"
)
break
# Create a temporary server object to use with get_server_prefix utility
temp_server = type('TempServer', (), {
'alias': alias,
'server_name': server_name,
'server_id': None
})()
temp_server = type(
"TempServer",
(),
{"alias": alias, "server_name": server_name, "server_id": None},
)()
name_for_prefix = get_server_prefix(temp_server)
# Generate stable server ID based on parameters
@ -251,15 +271,17 @@ class MCPServerManager:
_mcp_info: MCPInfo = mcp_server.mcp_info or {}
# Use helper to deserialize environment dictionary
# Safely access env field which may not exist on Prisma model objects
env_data = getattr(mcp_server, 'env', None)
env_data = getattr(mcp_server, "env", None)
env_dict = _deserialize_env_dict(env_data)
# Use alias for name if present, else server_name
name_for_prefix = mcp_server.alias or mcp_server.server_name or mcp_server.server_id
name_for_prefix = (
mcp_server.alias or mcp_server.server_name or mcp_server.server_id
)
new_server = MCPServer(
server_id=mcp_server.server_id,
name=name_for_prefix,
alias=getattr(mcp_server, 'alias', None),
server_name=getattr(mcp_server, 'server_name', None),
alias=getattr(mcp_server, "alias", None),
server_name=getattr(mcp_server, "server_name", None),
url=mcp_server.url,
transport=cast(MCPTransportType, mcp_server.transport),
spec_version=_convert_protocol_version_to_enum(mcp_server.spec_version),
@ -270,14 +292,12 @@ class MCPServerManager:
mcp_server_cost_info=_mcp_info.get("mcp_server_cost_info", None),
),
# Stdio-specific fields
command=getattr(mcp_server, 'command', None),
args=getattr(mcp_server, 'args', None) or [],
command=getattr(mcp_server, "command", None),
args=getattr(mcp_server, "args", None) or [],
env=env_dict,
)
self.registry[mcp_server.server_id] = new_server
verbose_logger.debug(
f"Added MCP Server: {name_for_prefix}"
)
verbose_logger.debug(f"Added MCP Server: {name_for_prefix}")
async def get_allowed_mcp_servers(
self, user_api_key_auth: Optional[UserAPIKeyAuth] = None
@ -300,10 +320,11 @@ class MCPServerManager:
)
return list(self.get_registry().keys())
except Exception as e:
verbose_logger.warning(f"Failed to get allowed MCP servers: {str(e)}. Returning default registry servers.")
verbose_logger.warning(
f"Failed to get allowed MCP servers: {str(e)}. Returning default registry servers."
)
return list(self.get_registry().keys())
async def get_tools_for_server(self, server_id: str) -> List[MCPTool]:
"""
Get the tools for a given server
@ -315,12 +336,13 @@ class MCPServerManager:
return []
return await self._get_tools_from_server(server)
except Exception as e:
verbose_logger.warning(f"Failed to get tools from server {server_id}: {str(e)}")
verbose_logger.warning(
f"Failed to get tools from server {server_id}: {str(e)}"
)
return []
async def list_tools(
self,
self,
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
mcp_auth_header: Optional[str] = None,
mcp_server_auth_headers: Optional[Dict[str, str]] = None,
@ -348,18 +370,18 @@ class MCPServerManager:
if server is None:
verbose_logger.warning(f"MCP Server {server_id} not found")
continue
# Get server-specific auth header if available
server_auth_header = None
if mcp_server_auth_headers and server.alias:
server_auth_header = mcp_server_auth_headers.get(server.alias)
elif mcp_server_auth_headers and server.server_name:
server_auth_header = mcp_server_auth_headers.get(server.server_name)
# Fall back to deprecated mcp_auth_header if no server-specific header found
if server_auth_header is None:
server_auth_header = mcp_auth_header
try:
tools = await self._get_tools_from_server(
server=server,
@ -367,20 +389,29 @@ class MCPServerManager:
mcp_protocol_version=mcp_protocol_version,
)
list_tools_result.extend(tools)
verbose_logger.info(f"Successfully fetched {len(tools)} tools from server {server.name}")
verbose_logger.info(
f"Successfully fetched {len(tools)} tools from server {server.name}"
)
except Exception as e:
verbose_logger.warning(
f"Failed to list tools from server {server.name}: {str(e)}. Continuing with other servers."
)
# Continue with other servers instead of failing completely
verbose_logger.info(f"Successfully fetched {len(list_tools_result)} tools total from all servers")
verbose_logger.info(
f"Successfully fetched {len(list_tools_result)} tools total from all servers"
)
return list_tools_result
#########################################################
# Methods that call the upstream MCP servers
#########################################################
def _create_mcp_client(self, server: MCPServer, mcp_auth_header: Optional[str] = None, protocol_version: Optional[str] = None) -> MCPClient:
def _create_mcp_client(
self,
server: MCPServer,
mcp_auth_header: Optional[str] = None,
protocol_version: Optional[str] = None,
) -> MCPClient:
"""
Create an MCPClient instance for the given server.
@ -393,21 +424,21 @@ class MCPServerManager:
MCPClient: Configured MCP client instance
"""
transport = server.transport or MCPTransport.sse
# Convert protocol version string to enum
protocol_version_enum = _convert_protocol_version_to_enum(protocol_version or server.spec_version)
protocol_version_enum = _convert_protocol_version_to_enum(
protocol_version or server.spec_version
)
# Handle stdio transport
if transport == MCPTransport.stdio:
# For stdio, we need to get the stdio config from the server
stdio_config: Optional[MCPStdioConfig] = None
if server.command and server.args is not None:
stdio_config = MCPStdioConfig(
command=server.command,
args=server.args,
env=server.env or {}
command=server.command, args=server.args, env=server.env or {}
)
return MCPClient(
server_url="", # Not used for stdio
transport_type=transport,
@ -429,7 +460,12 @@ class MCPServerManager:
protocol_version=protocol_version_enum,
)
async def _get_tools_from_server(self, server: MCPServer, mcp_auth_header: Optional[str] = None, mcp_protocol_version: Optional[str] = None) -> List[MCPTool]:
async def _get_tools_from_server(
self,
server: MCPServer,
mcp_auth_header: Optional[str] = None,
mcp_protocol_version: Optional[str] = None,
) -> List[MCPTool]:
"""
Helper method to get tools from a single MCP server with prefixed names.
@ -443,9 +479,11 @@ class MCPServerManager:
verbose_logger.debug(f"Connecting to url: {server.url}")
verbose_logger.info(f"_get_tools_from_server for {server.name}...")
protocol_version = mcp_protocol_version if mcp_protocol_version else server.spec_version
protocol_version = (
mcp_protocol_version if mcp_protocol_version else server.spec_version
)
client = None
try:
client = self._create_mcp_client(
server=server,
@ -455,9 +493,11 @@ class MCPServerManager:
tools = await self._fetch_tools_with_timeout(client, server.name)
return self._create_prefixed_tools(tools, server)
except Exception as e:
verbose_logger.warning(f"Failed to get tools from server {server.name}: {str(e)}")
verbose_logger.warning(
f"Failed to get tools from server {server.name}: {str(e)}"
)
return []
finally:
if client:
@ -466,17 +506,20 @@ class MCPServerManager:
except Exception:
pass
async def _fetch_tools_with_timeout(self, client: MCPClient, server_name: str) -> List[MCPTool]:
async def _fetch_tools_with_timeout(
self, client: MCPClient, server_name: str
) -> List[MCPTool]:
"""
Fetch tools from MCP client with timeout and error handling.
Args:
client: MCP client instance
server_name: Name of the server for logging
Returns:
List of tools from the server
"""
async def _list_tools_task():
try:
await client.connect()
@ -487,7 +530,9 @@ class MCPServerManager:
verbose_logger.warning(f"Client operation cancelled for {server_name}")
return []
except Exception as e:
verbose_logger.warning(f"Client operation failed for {server_name}: {str(e)}")
verbose_logger.warning(
f"Client operation failed for {server_name}: {str(e)}"
)
return []
finally:
try:
@ -501,36 +546,42 @@ class MCPServerManager:
verbose_logger.warning(f"Timeout while listing tools from {server_name}")
return []
except asyncio.CancelledError:
verbose_logger.warning(f"Task cancelled while listing tools from {server_name}")
verbose_logger.warning(
f"Task cancelled while listing tools from {server_name}"
)
return []
except ConnectionError as e:
verbose_logger.warning(f"Connection error while listing tools from {server_name}: {str(e)}")
verbose_logger.warning(
f"Connection error while listing tools from {server_name}: {str(e)}"
)
return []
except Exception as e:
verbose_logger.warning(f"Error listing tools from {server_name}: {str(e)}")
return []
def _create_prefixed_tools(self, tools: List[MCPTool], server: MCPServer) -> List[MCPTool]:
def _create_prefixed_tools(
self, tools: List[MCPTool], server: MCPServer
) -> List[MCPTool]:
"""
Create prefixed tools and update tool mapping.
Args:
tools: List of original tools from server
server: Server instance
Returns:
List of tools with prefixed names
"""
prefixed_tools = []
prefix = get_server_prefix(server)
for tool in tools:
prefixed_name = add_server_prefix_to_tool_name(tool.name, prefix)
prefixed_tool = MCPTool(
name=prefixed_name,
description=tool.description,
inputSchema=tool.inputSchema
inputSchema=tool.inputSchema,
)
prefixed_tools.append(prefixed_tool)
@ -538,18 +589,20 @@ class MCPServerManager:
self.tool_name_to_mcp_server_name_mapping[tool.name] = prefix
self.tool_name_to_mcp_server_name_mapping[prefixed_name] = prefix
verbose_logger.info(f"Successfully fetched {len(prefixed_tools)} tools from server {server.name}")
verbose_logger.info(
f"Successfully fetched {len(prefixed_tools)} tools from server {server.name}"
)
return prefixed_tools
async def call_tool(
self,
name: str,
arguments: Dict[str, Any],
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
mcp_auth_header: Optional[str] = None,
mcp_server_auth_headers: Optional[Dict[str, str]] = None,
mcp_protocol_version: Optional[str] = None,
proxy_logging_obj: Optional[ProxyLogging] = None,
self,
name: str,
arguments: Dict[str, Any],
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
mcp_auth_header: Optional[str] = None,
mcp_server_auth_headers: Optional[Dict[str, str]] = None,
mcp_protocol_version: Optional[str] = None,
proxy_logging_obj: Optional[ProxyLogging] = None,
) -> CallToolResult:
"""
Call a tool with the given name and arguments (handles prefixed tool names)
@ -567,9 +620,11 @@ class MCPServerManager:
CallToolResult from the MCP server
"""
start_time = datetime.datetime.now()
# Remove prefix if present to get the original tool name
original_tool_name, server_name_from_prefix = get_server_name_prefix_tool_mcp(name)
original_tool_name, server_name_from_prefix = get_server_name_prefix_tool_mcp(
name
)
# Get the MCP server
mcp_server = self._get_mcp_server_from_tool_name(name)
@ -579,9 +634,12 @@ class MCPServerManager:
# Validate that the server from prefix matches the actual server (if prefix was used)
if server_name_from_prefix:
expected_prefix = get_server_prefix(mcp_server)
if normalize_server_name(server_name_from_prefix) != normalize_server_name(expected_prefix):
if normalize_server_name(server_name_from_prefix) != normalize_server_name(
expected_prefix
):
raise ValueError(
f"Tool {name} server prefix mismatch: expected {expected_prefix}, got {server_name_from_prefix}")
f"Tool {name} server prefix mismatch: expected {expected_prefix}, got {server_name_from_prefix}"
)
#########################################################
# Pre MCP Tool Call Hook
@ -601,14 +659,20 @@ class MCPServerManager:
start_time=start_time,
end_time=start_time,
)
if pre_hook_result:
if pre_hook_result:
# Apply any argument modifications
if pre_hook_result.get("modified_arguments"):
arguments = pre_hook_result["modified_arguments"]
except (BlockedPiiEntityError, GuardrailRaisedException, HTTPException) as e:
except (
BlockedPiiEntityError,
GuardrailRaisedException,
HTTPException,
) as e:
# Re-raise guardrail exceptions to properly fail the MCP call
verbose_logger.error(f"Guardrail blocked MCP tool call pre call: {str(e)}")
verbose_logger.error(
f"Guardrail blocked MCP tool call pre call: {str(e)}"
)
raise e
# Get server-specific auth header if available
@ -617,7 +681,7 @@ class MCPServerManager:
server_auth_header = mcp_server_auth_headers.get(mcp_server.alias)
elif mcp_server_auth_headers and mcp_server.server_name:
server_auth_header = mcp_server_auth_headers.get(mcp_server.server_name)
# Fall back to deprecated mcp_auth_header if no server-specific header found
if server_auth_header is None:
server_auth_header = mcp_auth_header
@ -627,7 +691,7 @@ class MCPServerManager:
mcp_auth_header=server_auth_header,
protocol_version=mcp_protocol_version,
)
async with client:
# Use the original tool name (without prefix) for the actual call
@ -635,7 +699,7 @@ class MCPServerManager:
name=original_tool_name,
arguments=arguments,
)
# Initialize during_hook_task as None
during_hook_task = None
tasks = []
@ -654,7 +718,6 @@ class MCPServerManager:
)
)
tasks.append(during_hook_task)
tasks.append(asyncio.create_task(client.call_tool(call_tool_params)))
try:
@ -665,18 +728,23 @@ class MCPServerManager:
# If proxy_logging_obj is not None, the tool call result is at index 1 (after the during hook task)
result_index = 1 if proxy_logging_obj else 0
result = mcp_responses[result_index]
return cast(CallToolResult, result)
except (BlockedPiiEntityError, GuardrailRaisedException, HTTPException) as e:
# Re-raise guardrail exceptions to properly fail the MCP call
verbose_logger.error(f"Guardrail blocked MCP tool call during result check: {str(e)}")
raise e
except (
BlockedPiiEntityError,
GuardrailRaisedException,
HTTPException,
) as e:
# Re-raise guardrail exceptions to properly fail the MCP call
verbose_logger.error(
f"Guardrail blocked MCP tool call during result check: {str(e)}"
)
raise e
#########################################################
# End of Methods that call the upstream MCP servers
#########################################################
def initialize_tool_name_to_mcp_server_name_mapping(self):
"""
On startup, initialize the tool name to MCP server name mapping
@ -719,14 +787,18 @@ class MCPServerManager:
if tool_name in self.tool_name_to_mcp_server_name_mapping:
server_name = self.tool_name_to_mcp_server_name_mapping[tool_name]
for server in self.get_registry().values():
if normalize_server_name(server.name) == normalize_server_name(server_name):
if normalize_server_name(server.name) == normalize_server_name(
server_name
):
return server
# If not found and tool name is prefixed, try extracting server name from prefix
if is_tool_name_prefixed(tool_name):
_, server_name_from_prefix = get_server_name_prefix_tool_mcp(tool_name)
for server in self.get_registry().values():
if normalize_server_name(server.name) == normalize_server_name(server_name_from_prefix):
if normalize_server_name(server.name) == normalize_server_name(
server_name_from_prefix
):
return server
return None
@ -784,9 +856,7 @@ class MCPServerManager:
A deterministic server ID string
"""
# Create a string from all the identifying parameters
params_string = (
f"{server_name}|{url}|{transport}|{spec_version}|{auth_type or ''}|{alias or ''}"
)
params_string = f"{server_name}|{url}|{transport}|{spec_version}|{auth_type or ''}|{alias or ''}"
# Generate SHA-256 hash
hash_object = hashlib.sha256(params_string.encode("utf-8"))
@ -795,20 +865,22 @@ class MCPServerManager:
# Take first 32 characters and format as UUID-like string
return hash_hex[:32]
async def health_check_server(self, server_id: str, mcp_auth_header: Optional[str] = None) -> Dict[str, Any]:
async def health_check_server(
self, server_id: str, mcp_auth_header: Optional[str] = None
) -> Dict[str, Any]:
"""
Perform a health check on a specific MCP server.
Args:
server_id: The ID of the server to health check
mcp_auth_header: Optional authentication header for the MCP server
Returns:
Dict containing health check results
"""
import time
from datetime import datetime
server = self.get_mcp_server_by_id(server_id)
if not server:
return {
@ -816,90 +888,96 @@ class MCPServerManager:
"status": "unknown",
"error": "Server not found",
"last_health_check": datetime.now().isoformat(),
"response_time_ms": None
"response_time_ms": None,
}
start_time = time.time()
try:
# Try to get tools from the server as a health check
tools = await self._get_tools_from_server(server, mcp_auth_header)
response_time = (time.time() - start_time) * 1000
return {
"server_id": server_id,
"status": "healthy",
"tools_count": len(tools),
"last_health_check": datetime.now().isoformat(),
"response_time_ms": round(response_time, 2),
"error": None
"error": None,
}
except Exception as e:
response_time = (time.time() - start_time) * 1000
error_message = str(e)
return {
"server_id": server_id,
"status": "unhealthy",
"last_health_check": datetime.now().isoformat(),
"response_time_ms": round(response_time, 2),
"error": error_message
"error": error_message,
}
async def health_check_all_servers(self, mcp_auth_header: Optional[str] = None) -> Dict[str, Any]:
async def health_check_all_servers(
self, mcp_auth_header: Optional[str] = None
) -> Dict[str, Any]:
"""
Perform health checks on all MCP servers.
Args:
mcp_auth_header: Optional authentication header for the MCP servers
Returns:
Dict containing health check results for all servers
"""
all_servers = self.get_registry()
results = {}
for server_id, server in all_servers.items():
results[server_id] = await self.health_check_server(server_id, mcp_auth_header)
results[server_id] = await self.health_check_server(
server_id, mcp_auth_header
)
return results
async def health_check_allowed_servers(
self,
self,
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
mcp_auth_header: Optional[str] = None
mcp_auth_header: Optional[str] = None,
) -> Dict[str, Any]:
"""
Perform health checks on all MCP servers that the user has access to.
Args:
user_api_key_auth: User authentication info for access control
mcp_auth_header: Optional authentication header for the MCP servers
Returns:
Dict containing health check results for accessible servers
"""
# Get allowed servers for the user
allowed_server_ids = await self.get_allowed_mcp_servers(user_api_key_auth)
# Perform health checks on allowed servers
results = {}
for server_id in allowed_server_ids:
results[server_id] = await self.health_check_server(server_id, mcp_auth_header)
results[server_id] = await self.health_check_server(
server_id, mcp_auth_header
)
return results
async def get_all_mcp_servers_with_health_and_teams(
self,
self,
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
include_health: bool = True
include_health: bool = True,
) -> List[LiteLLM_MCPServerTable]:
"""
Get all MCP servers that the user has access to, with health status and team information.
Args:
user_api_key_auth: User authentication info for access control
include_health: Whether to include health check information
Returns:
List of MCP server objects with health and team data
"""
@ -912,12 +990,12 @@ class MCPServerManager:
# Get allowed server IDs
allowed_server_ids = await self.get_allowed_mcp_servers(user_api_key_auth)
# Get servers from database
list_mcp_servers: List[LiteLLM_MCPServerTable] = []
if prisma_client is not None:
list_mcp_servers = await get_mcp_servers(prisma_client, allowed_server_ids)
# If admin, also get all servers from database
if user_api_key_auth and _user_has_admin_view(user_api_key_auth):
all_mcp_servers = await get_all_mcp_servers(prisma_client)
@ -941,19 +1019,23 @@ class MCPServerManager:
updated_at=datetime.datetime.now(),
mcp_info=_server_config.mcp_info,
# Stdio-specific fields
command=getattr(_server_config, 'command', None),
args=getattr(_server_config, 'args', None) or [],
env=getattr(_server_config, 'env', None) or {},
command=getattr(_server_config, "command", None),
args=getattr(_server_config, "args", None) or [],
env=getattr(_server_config, "env", None) or {},
)
)
# Get team information for non-admin users
server_to_teams_map: Dict[str, List[Dict[str, str]]] = {}
if user_api_key_auth and not _user_has_admin_view(user_api_key_auth) and prisma_client is not None:
if (
user_api_key_auth
and not _user_has_admin_view(user_api_key_auth)
and prisma_client is not None
):
teams = await prisma_client.db.litellm_teamtable.find_many(
include={"object_permission": True}
)
user_teams = []
for team in teams:
if team.members_with_roles:
@ -971,14 +1053,17 @@ class MCPServerManager:
for server_id in team.object_permission.mcp_servers:
if server_id not in server_to_teams_map:
server_to_teams_map[server_id] = []
server_to_teams_map[server_id].append({
"team_id": team.team_id,
"team_alias": team.team_alias,
"organization_id": team.organization_id
})
server_to_teams_map[server_id].append(
{
"team_id": team.team_id,
"team_alias": team.team_alias,
"organization_id": team.organization_id,
}
)
# Map servers to their teams and return with health data
from typing import cast
return [
LiteLLM_MCPServerTable(
server_id=server.server_id,
@ -993,13 +1078,20 @@ class MCPServerManager:
created_by=server.created_by,
updated_at=server.updated_at,
updated_by=server.updated_by,
mcp_access_groups=server.mcp_access_groups if server.mcp_access_groups is not None else [],
mcp_access_groups=(
server.mcp_access_groups
if server.mcp_access_groups is not None
else []
),
mcp_info=server.mcp_info,
teams=cast(List[Dict[str, str | None]], server_to_teams_map.get(server.server_id, [])),
teams=cast(
List[Dict[str, str | None]],
server_to_teams_map.get(server.server_id, []),
),
# Stdio-specific fields
command=getattr(server, 'command', None),
args=getattr(server, 'args', None) or [],
env=getattr(server, 'env', None) or {},
command=getattr(server, "command", None),
args=getattr(server, "args", None) or [],
env=getattr(server, "env", None) or {},
)
for server in list_mcp_servers
]

View file

@ -13,16 +13,16 @@ guardrails:
api_base: os.environ/AZURE_GUARDRAIL_API_BASE
prompts:
- prompt_id: test_hello_world_prompt
- prompt_id: test_my_json_prompt
litellm_params:
prompt_integration: dotprompt
prompt_id: test_hello_world_prompt
prompt_directory: /Users/krrishdholakia/Documents/litellm/litellm/proxy/test_prompts
prompt_file: /Users/krrishdholakia/Documents/litellm/litellm/proxy/test_prompts/test_hello_world_prompt.prompt
- prompt_id: test_hello_world_prompt_2
litellm_params:
prompt_integration: dotprompt
prompt_id: test_hello_world_prompt
prompt_directory: /Users/krrishdholakia/Documents/litellm/litellm/proxy/test_prompts
prompt_id: test_hello_world_prompt_2
prompt_file: /Users/krrishdholakia/Documents/litellm/litellm/proxy/test_prompts/test_hello_world_prompt.prompt
litellm_settings:
callbacks: ["datadog_llm_observability"]

View file

@ -1,7 +1,6 @@
from typing import Any, Optional, Union
from litellm.proxy._types import (
GenerateKeyRequest,
KeyRequestBase,
LiteLLM_ManagementEndpoint_MetadataFields_Premium,
LiteLLM_TeamTable,

View file

@ -585,6 +585,7 @@ async def generate_key_fn(
- tpm_limit: Optional[int] - Specify tpm limit for a given key (Tokens per minute)
- soft_budget: Optional[float] - Specify soft budget for a given key. Will trigger a slack alert when this soft budget is reached.
- tags: Optional[List[str]] - Tags for [tracking spend](https://litellm.vercel.app/docs/proxy/enterprise#tracking-spend-for-custom-tags) and/or doing [tag-based routing](https://litellm.vercel.app/docs/proxy/tag_routing).
- prompts: Optional[List[str]] - List of prompts that the key is allowed to use.
- enforced_params: Optional[List[str]] - List of enforced params for the key (Enterprise only). [Docs](https://docs.litellm.ai/docs/proxy/enterprise#enforce-required-params-for-llm-requests)
- prompts: Optional[List[str]] - List of prompts that the key is allowed to use.
- allowed_routes: Optional[list] - List of allowed routes for the key. Store the actual route or store a wildcard pattern for a set of routes. Example - ["/chat/completions", "/embeddings", "/keys/*"]
@ -993,6 +994,7 @@ async def update_key_fn(
- permissions: Optional[dict] - Key-specific permissions
- send_invite_email: Optional[bool] - Send invite email to user_id
- guardrails: Optional[List[str]] - List of active guardrails for the key
- prompts: Optional[List[str]] - List of prompts that the key is allowed to use.
- blocked: Optional[bool] - Whether the key is blocked
- aliases: Optional[dict] - Model aliases for the key - [Docs](https://litellm.vercel.app/docs/proxy/virtual_keys#model-aliases)
- config: Optional[dict] - [DEPRECATED PARAM] Key-specific config.

View file

@ -265,6 +265,7 @@ async def new_team( # noqa: PLR0915
- organization_id: Optional[str] - The organization id of the team. Default is None. Create via `/organization/new`.
- model_aliases: Optional[dict] - Model aliases for the team. [Docs](https://docs.litellm.ai/docs/proxy/team_based_routing#create-team-with-model-alias)
- guardrails: Optional[List[str]] - Guardrails for the team. [Docs](https://docs.litellm.ai/docs/proxy/guardrails)
- prompts: Optional[List[str]] - List of prompts that the team is allowed to use.
- object_permission: Optional[LiteLLM_ObjectPermissionBase] - team-specific object permission. Example - {"vector_stores": ["vector_store_1", "vector_store_2"]}. IF null or {} then no object permission.
- team_member_budget: Optional[float] - The maximum budget allocated to an individual team member.
- team_member_key_duration: Optional[str] - The duration for a team member's key. e.g. "1d", "1w", "1mo"
@ -701,6 +702,7 @@ async def update_team(
- organization_id: Optional[str] - The organization id of the team. Default is None. Create via `/organization/new`.
- model_aliases: Optional[dict] - Model aliases for the team. [Docs](https://docs.litellm.ai/docs/proxy/team_based_routing#create-team-with-model-alias)
- guardrails: Optional[List[str]] - Guardrails for the team. [Docs](https://docs.litellm.ai/docs/proxy/guardrails)
- prompts: Optional[List[str]] - List of prompts that the team is allowed to use.
- object_permission: Optional[LiteLLM_ObjectPermissionBase] - team-specific object permission. Example - {"vector_stores": ["vector_store_1", "vector_store_2"]}. IF null or {} then no object permission.
- team_member_budget: Optional[float] - The maximum budget allocated to an individual team member.
- team_member_key_duration: Optional[str] - The duration for a team member's key. e.g. "1d", "1w", "1mo"

View file

@ -2,7 +2,7 @@
Similar to init_guardrails.py, but for prompts.
"""
from typing import Dict, List, Optional, cast
from typing import Dict, List, Optional
from litellm._logging import verbose_proxy_logger
@ -19,7 +19,7 @@ def init_prompts(
for prompt in all_prompts:
initialized_prompt = IN_MEMORY_PROMPT_REGISTRY.initialize_prompt(
prompt=cast(PromptSpec, prompt),
prompt=PromptSpec(**prompt),
config_file_path=config_file_path,
)
if initialized_prompt:

View file

@ -2,19 +2,41 @@
CRUD ENDPOINTS FOR PROMPTS
"""
from typing import List, Optional, cast
import tempfile
from pathlib import Path
from typing import Any, Dict, List, Optional, cast
from fastapi import APIRouter, Depends, HTTPException
from fastapi import APIRouter, Depends, File, HTTPException, UploadFile
from pydantic import BaseModel
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import CommonProxyErrors, LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.types.prompts.init_prompts import ListPromptsResponse, PromptSpec
from litellm.types.prompts.init_prompts import (
ListPromptsResponse,
PromptInfo,
PromptInfoResponse,
PromptLiteLLMParams,
PromptSpec,
PromptTemplateBase,
)
router = APIRouter()
class Prompt(BaseModel):
prompt_id: str
litellm_params: PromptLiteLLMParams
prompt_info: Optional[PromptInfo] = None
class PatchPromptRequest(BaseModel):
litellm_params: Optional[PromptLiteLLMParams] = None
prompt_info: Optional[PromptInfo] = None
@router.get(
"/prompt/list",
"/prompts/list",
tags=["Prompt Management"],
dependencies=[Depends(user_api_key_auth)],
response_model=ListPromptsResponse,
@ -23,7 +45,35 @@ async def list_prompts(
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""
List of available prompts for a given key.
List the prompts that are available on the proxy server
👉 [Prompt docs](https://docs.litellm.ai/docs/proxy/prompt_management)
Example Request:
```bash
curl -X GET "http://localhost:4000/prompts/list" -H "Authorization: Bearer <your_api_key>"
```
Example Response:
```json
{
"prompts": [
{
"prompt_id": "my_prompt_id",
"litellm_params": {
"prompt_id": "my_prompt_id",
"prompt_integration": "dotprompt",
"prompt_directory": "/path/to/prompts"
},
"prompt_info": {
"prompt_type": "config"
},
"created_at": "2023-11-09T12:34:56.789Z",
"updated_at": "2023-11-09T12:34:56.789Z"
}
]
}
```
"""
from litellm.proxy._types import LitellmUserRoles
from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY
@ -53,17 +103,49 @@ async def list_prompts(
@router.get(
"/prompt/info",
"/prompts/{prompt_id}",
tags=["Prompt Management"],
dependencies=[Depends(user_api_key_auth)],
response_model=PromptSpec,
response_model=PromptInfoResponse,
)
async def get_prompt(
@router.get(
"/prompts/{prompt_id}/info",
tags=["Prompt Management"],
dependencies=[Depends(user_api_key_auth)],
response_model=PromptInfoResponse,
)
async def get_prompt_info(
prompt_id: str,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""
Get info about a prompt
Get detailed information about a specific prompt by ID, including prompt content
👉 [Prompt docs](https://docs.litellm.ai/docs/proxy/prompt_management)
Example Request:
```bash
curl -X GET "http://localhost:4000/prompts/my_prompt_id/info" \\
-H "Authorization: Bearer <your_api_key>"
```
Example Response:
```json
{
"prompt_id": "my_prompt_id",
"litellm_params": {
"prompt_id": "my_prompt_id",
"prompt_integration": "dotprompt",
"prompt_directory": "/path/to/prompts"
},
"prompt_info": {
"prompt_type": "config"
},
"created_at": "2023-11-09T12:34:56.789Z",
"updated_at": "2023-11-09T12:34:56.789Z",
"content": "System: You are a helpful assistant.\n\nUser: {{user_message}}"
}
```
"""
from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY
@ -89,4 +171,490 @@ async def get_prompt(
prompt_spec = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(prompt_id)
if prompt_spec is None:
raise HTTPException(status_code=400, detail=f"Prompt {prompt_id} not found")
return prompt_spec
# Get prompt content from the callback
prompt_template: Optional[PromptTemplateBase] = None
try:
prompt_callback = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id(prompt_id)
if prompt_callback is not None:
# Extract content based on integration type
integration_name = prompt_callback.integration_name
if integration_name == "dotprompt":
# For dotprompt integration, get content from the prompt manager
from litellm.integrations.dotprompt.dotprompt_manager import (
DotpromptManager,
)
if isinstance(prompt_callback, DotpromptManager):
template = prompt_callback.prompt_manager.get_all_prompts_as_json()
if template is not None and len(template) == 1:
template_id = list(template.keys())[0]
prompt_template = PromptTemplateBase(
litellm_prompt_id=template_id, # id sent to prompt management tool
content=template[template_id]["content"],
metadata=template[template_id]["metadata"],
)
except Exception:
# If content extraction fails, continue without content
pass
# Create response with content
return PromptInfoResponse(
prompt_spec=prompt_spec,
raw_prompt_template=prompt_template,
)
@router.post(
"/prompts",
tags=["Prompt Management"],
dependencies=[Depends(user_api_key_auth)],
)
async def create_prompt(
request: Prompt,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""
Create a new prompt
👉 [Prompt docs](https://docs.litellm.ai/docs/proxy/prompt_management)
Example Request:
```bash
curl -X POST "http://localhost:4000/prompts" \\
-H "Authorization: Bearer <your_api_key>" \\
-H "Content-Type: application/json" \\
-d '{
"prompt_id": "my_prompt",
"litellm_params": {
"prompt_id": "json_prompt",
"prompt_integration": "dotprompt",
### EITHER prompt_directory OR prompt_data MUST BE PROVIDED
"prompt_directory": "/path/to/dotprompt/folder",
"prompt_data": {"json_prompt": {"content": "This is a prompt", "metadata": {"model": "gpt-4"}}}
},
"prompt_info": {
"prompt_type": "config"
}
}'
```
"""
from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY
from litellm.proxy.proxy_server import prisma_client
# Only allow proxy admins to create prompts
if user_api_key_dict.user_role is None or (
user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN
and user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value
):
raise HTTPException(
status_code=403, detail="Only proxy admins can create prompts"
)
if prisma_client is None:
raise HTTPException(
status_code=500, detail=CommonProxyErrors.db_not_connected_error.value
)
try:
# Create the prompt spec
# Check if prompt exists and get current data
existing_prompt = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(request.prompt_id)
if existing_prompt is not None:
raise HTTPException(
status_code=404,
detail=f"Prompt with ID {request.prompt_id} already exists",
)
# store prompt in db
prompt_db_entry = await prisma_client.db.litellm_prompttable.create(
data={
"prompt_id": request.prompt_id,
"litellm_params": request.litellm_params.model_dump_json(),
"prompt_info": (
request.prompt_info.model_dump_json()
if request.prompt_info
else PromptInfo(prompt_type="db").model_dump_json()
),
}
)
prompt_spec = PromptSpec(**prompt_db_entry.model_dump())
# Initialize the prompt
initialized_prompt = IN_MEMORY_PROMPT_REGISTRY.initialize_prompt(
prompt=prompt_spec, config_file_path=None
)
if initialized_prompt is None:
raise HTTPException(status_code=500, detail="Failed to initialize prompt")
return initialized_prompt
except Exception as e:
verbose_proxy_logger.exception(f"Error creating prompt: {e}")
raise HTTPException(status_code=500, detail=str(e))
@router.put(
"/prompts/{prompt_id}",
tags=["Prompt Management"],
dependencies=[Depends(user_api_key_auth)],
)
async def update_prompt(
prompt_id: str,
request: Prompt,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""
Update an existing prompt
👉 [Prompt docs](https://docs.litellm.ai/docs/proxy/prompt_management)
Example Request:
```bash
curl -X PUT "http://localhost:4000/prompts/my_prompt_id" \\
-H "Authorization: Bearer <your_api_key>" \\
-H "Content-Type: application/json" \\
-d '{
"prompt_id": "my_prompt",
"litellm_params": {
"prompt_id": "my_prompt",
"prompt_integration": "dotprompt",
"prompt_directory": "/path/to/prompts"
},
"prompt_info": {
"prompt_type": "config"
}
}
}'
```
"""
from datetime import datetime
from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY
from litellm.proxy.proxy_server import prisma_client
# Only allow proxy admins to update prompts
if user_api_key_dict.user_role is None or (
user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN
and user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value
):
raise HTTPException(
status_code=403, detail="Only proxy admins can update prompts"
)
if prisma_client is None:
raise HTTPException(
status_code=500, detail=CommonProxyErrors.db_not_connected_error.value
)
try:
# Check if prompt exists
existing_prompt = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(prompt_id)
if existing_prompt is None:
raise HTTPException(
status_code=404, detail=f"Prompt with ID {prompt_id} not found"
)
if existing_prompt.prompt_info.prompt_type == "config":
raise HTTPException(
status_code=400,
detail="Cannot update config prompts.",
)
# Create updated prompt spec
updated_prompt_spec = PromptSpec(
prompt_id=prompt_id,
litellm_params=request.litellm_params,
prompt_info=request.prompt_info or PromptInfo(prompt_type="db"),
created_at=existing_prompt.created_at,
updated_at=datetime.now(),
)
updated_prompt_db_entry = await prisma_client.db.litellm_prompttable.update(
where={"prompt_id": prompt_id},
data={
"litellm_params": updated_prompt_spec.litellm_params.model_dump_json(),
"prompt_info": updated_prompt_spec.prompt_info.model_dump_json(),
},
)
# Remove the old prompt from memory
del IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS[prompt_id]
if prompt_id in IN_MEMORY_PROMPT_REGISTRY.prompt_id_to_custom_prompt:
del IN_MEMORY_PROMPT_REGISTRY.prompt_id_to_custom_prompt[prompt_id]
# Initialize the updated prompt
initialized_prompt = IN_MEMORY_PROMPT_REGISTRY.initialize_prompt(
prompt=PromptSpec(**updated_prompt_db_entry.model_dump()),
config_file_path=None,
)
if initialized_prompt is None:
raise HTTPException(status_code=500, detail="Failed to update prompt")
return initialized_prompt
except HTTPException as e:
raise e
except Exception as e:
verbose_proxy_logger.exception(f"Error updating prompt: {e}")
raise HTTPException(status_code=500, detail=str(e))
@router.delete(
"/prompts/{prompt_id}",
tags=["Prompt Management"],
dependencies=[Depends(user_api_key_auth)],
)
async def delete_prompt(
prompt_id: str,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""
Delete a prompt
👉 [Prompt docs](https://docs.litellm.ai/docs/proxy/prompt_management)
Example Request:
```bash
curl -X DELETE "http://localhost:4000/prompts/my_prompt_id" \\
-H "Authorization: Bearer <your_api_key>"
```
Example Response:
```json
{
"message": "Prompt my_prompt_id deleted successfully"
}
```
"""
from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY
from litellm.proxy.proxy_server import prisma_client
# Only allow proxy admins to delete prompts
if user_api_key_dict.user_role is None or (
user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN
and user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value
):
raise HTTPException(
status_code=403, detail="Only proxy admins can delete prompts"
)
if prisma_client is None:
raise HTTPException(
status_code=500, detail=CommonProxyErrors.db_not_connected_error.value
)
try:
# Check if prompt exists
existing_prompt = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(prompt_id)
if existing_prompt is None:
raise HTTPException(
status_code=404, detail=f"Prompt with ID {prompt_id} not found"
)
if existing_prompt.prompt_info.prompt_type == "config":
raise HTTPException(
status_code=400,
detail="Cannot delete config prompts.",
)
# Delete the prompt from the database
await prisma_client.db.litellm_prompttable.delete(
where={"prompt_id": prompt_id}
)
# Remove the prompt from memory
del IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS[prompt_id]
if prompt_id in IN_MEMORY_PROMPT_REGISTRY.prompt_id_to_custom_prompt:
del IN_MEMORY_PROMPT_REGISTRY.prompt_id_to_custom_prompt[prompt_id]
return {"message": f"Prompt {prompt_id} deleted successfully"}
except HTTPException as e:
raise e
except Exception as e:
verbose_proxy_logger.exception(f"Error deleting prompt: {e}")
raise HTTPException(status_code=500, detail=str(e))
@router.patch(
"/prompts/{prompt_id}",
tags=["Prompt Management"],
dependencies=[Depends(user_api_key_auth)],
)
async def patch_prompt(
prompt_id: str,
request: PatchPromptRequest,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""
Partially update an existing prompt
👉 [Prompt docs](https://docs.litellm.ai/docs/proxy/prompt_management)
This endpoint allows updating specific fields of a prompt without sending the entire object.
Only the following fields can be updated:
- litellm_params: LiteLLM parameters for the prompt
- prompt_info: Additional information about the prompt
Example Request:
```bash
curl -X PATCH "http://localhost:4000/prompts/my_prompt_id" \\
-H "Authorization: Bearer <your_api_key>" \\
-H "Content-Type: application/json" \\
-d '{
"prompt_info": {
"prompt_type": "db"
}
}'
```
"""
from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY
from litellm.proxy.proxy_server import prisma_client
# Only allow proxy admins to patch prompts
if user_api_key_dict.user_role is None or (
user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN
and user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value
):
raise HTTPException(
status_code=403, detail="Only proxy admins can patch prompts"
)
if prisma_client is None:
raise HTTPException(
status_code=500, detail=CommonProxyErrors.db_not_connected_error.value
)
try:
# Check if prompt exists and get current data
existing_prompt = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(prompt_id)
if existing_prompt is None:
raise HTTPException(
status_code=404, detail=f"Prompt with ID {prompt_id} not found"
)
if existing_prompt.prompt_info.prompt_type == "config":
raise HTTPException(
status_code=400,
detail="Cannot update config prompts.",
)
# Update fields if provided
updated_litellm_params = (
request.litellm_params
if request.litellm_params is not None
else existing_prompt.litellm_params
)
updated_prompt_info = (
request.prompt_info
if request.prompt_info is not None
else existing_prompt.prompt_info
)
# Ensure we have valid litellm_params
if updated_litellm_params is None:
raise HTTPException(status_code=400, detail="litellm_params cannot be None")
# Create updated prompt spec - cast to satisfy typing
updated_prompt_db_entry = await prisma_client.db.litellm_prompttable.update(
where={"prompt_id": prompt_id},
data={
"litellm_params": updated_litellm_params.model_dump_json(),
"prompt_info": updated_prompt_info.model_dump_json(),
},
)
updated_prompt_spec = PromptSpec(**updated_prompt_db_entry.model_dump())
# Remove the old prompt from memory
del IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS[prompt_id]
if prompt_id in IN_MEMORY_PROMPT_REGISTRY.prompt_id_to_custom_prompt:
del IN_MEMORY_PROMPT_REGISTRY.prompt_id_to_custom_prompt[prompt_id]
# Initialize the updated prompt
initialized_prompt = IN_MEMORY_PROMPT_REGISTRY.initialize_prompt(
prompt=updated_prompt_spec, config_file_path=None
)
if initialized_prompt is None:
raise HTTPException(status_code=500, detail="Failed to patch prompt")
return initialized_prompt
except HTTPException as e:
raise e
except Exception as e:
verbose_proxy_logger.exception(f"Error patching prompt: {e}")
raise HTTPException(status_code=500, detail=str(e))
@router.post(
"/utils/dotprompt_json_converter",
tags=["prompts", "utils"],
dependencies=[Depends(user_api_key_auth)],
)
async def convert_prompt_file_to_json(
file: UploadFile = File(...),
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
) -> Dict[str, Any]:
"""
Convert a .prompt file to JSON format.
This endpoint accepts a .prompt file upload and returns the equivalent JSON representation
that can be stored in a database or used programmatically.
Returns the JSON structure with 'content' and 'metadata' fields.
"""
global general_settings
from litellm.integrations.dotprompt.prompt_manager import PromptManager
# Validate file extension
if not file.filename or not file.filename.endswith(".prompt"):
raise HTTPException(status_code=400, detail="File must have .prompt extension")
temp_file_path = None
try:
# Read file content
file_content = await file.read()
# Create temporary file
temp_file_path = Path(tempfile.mkdtemp()) / file.filename
temp_file_path.write_bytes(file_content)
# Create a PromptManager instance just for conversion
prompt_manager = PromptManager()
# Convert to JSON
json_data = prompt_manager.prompt_file_to_json(temp_file_path)
# Extract prompt ID from filename
prompt_id = temp_file_path.stem
return {
"prompt_id": prompt_id,
"json_data": json_data,
}
except Exception as e:
raise HTTPException(
status_code=500, detail=f"Error converting prompt file: {str(e)}"
)
finally:
# Clean up temp file
if temp_file_path and temp_file_path.exists():
temp_file_path.unlink()
# Also try to remove the temp directory if it's empty
try:
temp_file_path.parent.rmdir()
except OSError:
pass # Directory not empty or other error

View file

@ -1,6 +1,5 @@
import importlib
import os
import uuid
from pathlib import Path
from typing import Callable, Dict, Optional
@ -117,14 +116,13 @@ class InMemoryPromptRegistry:
"""
import litellm
prompt_id = prompt.get("prompt_id") or str(uuid.uuid4())
prompt["prompt_id"] = prompt_id
prompt_id = prompt.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"]
litellm_params_data = prompt.litellm_params
verbose_proxy_logger.debug("litellm_params= %s", litellm_params_data)
if isinstance(litellm_params_data, dict):
@ -151,7 +149,9 @@ class InMemoryPromptRegistry:
parsed_prompt = PromptSpec(
prompt_id=prompt_id,
litellm_params=litellm_params,
prompt_info=PromptInfo(prompt_type="config"),
prompt_info=prompt.prompt_info or PromptInfo(prompt_type="config"),
created_at=prompt.created_at,
updated_at=prompt.updated_at,
)
# store references to the prompt in memory

View file

@ -687,6 +687,7 @@ def run_server( # noqa: PLR0915
if database_url is None and os.getenv("DATABASE_URL") is None:
# Use helper function to construct DATABASE_URL from individual variables
from litellm.proxy.utils import construct_database_url_from_env_vars
database_url = construct_database_url_from_env_vars()
if database_url:
os.environ["DATABASE_URL"] = database_url
@ -716,14 +717,19 @@ def run_server( # noqa: PLR0915
if config is None and os.getenv("DATABASE_URL") is None:
# Use helper function to construct DATABASE_URL from individual variables
from litellm.proxy.utils import construct_database_url_from_env_vars
database_url = construct_database_url_from_env_vars()
if database_url:
os.environ["DATABASE_URL"] = database_url
# Set default values for connection pool settings when no config is used
if config is None:
db_connection_pool_limit = LiteLLMDatabaseConnectionPool.database_connection_pool_limit.value
db_connection_timeout = LiteLLMDatabaseConnectionPool.database_connection_pool_timeout.value
db_connection_pool_limit = (
LiteLLMDatabaseConnectionPool.database_connection_pool_limit.value
)
db_connection_timeout = (
LiteLLMDatabaseConnectionPool.database_connection_pool_timeout.value
)
if (
os.getenv("DATABASE_URL", None) is not None

View file

@ -2924,6 +2924,14 @@ class ProxyConfig:
await self._init_vector_stores_in_db(prisma_client=prisma_client)
await self._init_mcp_servers_in_db()
await self._init_pass_through_endpoints_in_db()
await self._init_prompts_in_db(prisma_client=prisma_client)
async def _init_prompts_in_db(self, prisma_client: PrismaClient):
from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY
prompts_in_db = await prisma_client.db.litellm_prompttable.find_many()
for prompt in prompts_in_db:
IN_MEMORY_PROMPT_REGISTRY.initialize_prompt(prompt=prompt)
async def _init_guardrails_in_db(self, prisma_client: PrismaClient):
from litellm.proxy.guardrails.guardrail_registry import (

View file

@ -520,6 +520,16 @@ model LiteLLM_GuardrailsTable {
updated_at DateTime @updatedAt
}
// Prompt table for storing prompt configurations
model LiteLLM_PromptTable {
id String @id @default(uuid())
prompt_id String @unique
litellm_params Json
prompt_info Json?
created_at DateTime @default(now())
updated_at DateTime @updatedAt
}
model LiteLLM_HealthCheckTable {
health_check_id String @id @default(uuid())
model_name String

View file

@ -55,7 +55,11 @@ from litellm import (
from litellm._logging import verbose_proxy_logger
from litellm._service_logger import ServiceLogging, ServiceTypes
from litellm.caching.caching import DualCache, RedisCache
from litellm.exceptions import RejectedRequestError, BlockedPiiEntityError, GuardrailRaisedException
from litellm.exceptions import (
BlockedPiiEntityError,
GuardrailRaisedException,
RejectedRequestError,
)
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.integrations.custom_logger import CustomLogger
from litellm.integrations.SlackAlerting.slack_alerting import SlackAlerting
@ -472,8 +476,10 @@ class ProxyLogging:
# Create the request object if it's not already one
if not isinstance(request_obj, MCPPreCallRequestObject):
# Convert UserAPIKeyAuth object to dict if needed
user_api_key_auth_dict = self._convert_user_api_key_auth_to_dict(kwargs.get("user_api_key_auth"))
user_api_key_auth_dict = self._convert_user_api_key_auth_to_dict(
kwargs.get("user_api_key_auth")
)
request_obj = MCPPreCallRequestObject(
tool_name=kwargs.get("name", ""),
arguments=kwargs.get("arguments", {}),
@ -487,7 +493,9 @@ class ProxyLogging:
_callback: Optional[CustomLogger] = None
if isinstance(callback, str):
from typing import cast
from litellm import _custom_logger_compatible_callbacks_literal
_callback = litellm.litellm_core_utils.litellm_logging.get_custom_logger_compatible_class(
cast(_custom_logger_compatible_callbacks_literal, callback)
)
@ -507,26 +515,36 @@ class ProxyLogging:
continue
# Convert MCP tool call to LLM message format for existing guardrail logic
synthetic_llm_data = self._convert_mcp_to_llm_format(request_obj, kwargs)
synthetic_llm_data = self._convert_mcp_to_llm_format(
request_obj, kwargs
)
# Reuse existing LLM guardrail logic
user_api_key_auth_dict = self._convert_user_api_key_auth_to_dict(kwargs.get("user_api_key_auth"))
user_api_key_auth_dict = self._convert_user_api_key_auth_to_dict(
kwargs.get("user_api_key_auth")
)
result = await _callback.async_pre_call_hook(
user_api_key_dict=user_api_key_auth_dict,
user_api_key_dict=user_api_key_auth_dict, # type: ignore
cache=self.call_details["user_api_key_cache"],
data=synthetic_llm_data,
call_type="mcp_call"
call_type="mcp_call",
)
# Convert result back to MCP response format if blocked/modified
if result is not None:
mcp_response = self._convert_llm_result_to_mcp_response(result, request_obj)
mcp_response = self._convert_llm_result_to_mcp_response(
result, request_obj
)
if mcp_response is not None:
return self._parse_pre_mcp_call_hook_response(
response=mcp_response, original_request=request_obj
)
except (BlockedPiiEntityError, GuardrailRaisedException, HTTPException) as e:
except (
BlockedPiiEntityError,
GuardrailRaisedException,
HTTPException,
) as e:
# Re-raise guardrail exceptions so they can be properly handled
raise e
except Exception as e:
@ -543,10 +561,10 @@ class ProxyLogging:
Handles both Pydantic models and regular objects.
"""
if user_api_key_auth_obj is not None:
if hasattr(user_api_key_auth_obj, 'model_dump'):
if hasattr(user_api_key_auth_obj, "model_dump"):
# If it's a Pydantic model, convert to dict
return user_api_key_auth_obj.model_dump()
elif hasattr(user_api_key_auth_obj, '__dict__'):
elif hasattr(user_api_key_auth_obj, "__dict__"):
# If it's a regular object, convert to dict
return user_api_key_auth_obj.__dict__
return user_api_key_auth_obj
@ -556,15 +574,16 @@ class ProxyLogging:
Convert MCP tool call to LLM message format for existing guardrail validation.
"""
from litellm.types.llms.openai import ChatCompletionUserMessage
# Create a synthetic message that represents the tool call
tool_call_content = f"Tool: {request_obj.tool_name}\nArguments: {request_obj.arguments}"
synthetic_message = ChatCompletionUserMessage(
role="user",
content=tool_call_content
tool_call_content = (
f"Tool: {request_obj.tool_name}\nArguments: {request_obj.arguments}"
)
synthetic_message = ChatCompletionUserMessage(
role="user", content=tool_call_content
)
# Create synthetic LLM data that guardrails can process
synthetic_data = {
"messages": [synthetic_message],
@ -577,51 +596,65 @@ class ProxyLogging:
"mcp_tool_name": request_obj.tool_name, # Keep original for reference
"mcp_arguments": request_obj.arguments, # Keep original for reference
}
return synthetic_data
def _convert_llm_result_to_mcp_response(self, llm_result, request_obj) -> Optional[Any]:
def _convert_llm_result_to_mcp_response(
self, llm_result, request_obj
) -> Optional[Any]:
"""
Convert LLM guardrail result back to MCP response format.
"""
from litellm.types.mcp import MCPPreCallResponseObject
# If result is an exception, it means the guardrail blocked the request
if isinstance(llm_result, Exception):
return MCPPreCallResponseObject(
should_proceed=False,
error_message=str(llm_result),
modified_arguments=None
modified_arguments=None,
)
# If result is a dict with modified messages, check for content filtering
if isinstance(llm_result, dict):
modified_messages = llm_result.get("messages")
if modified_messages:
# Check if content was blocked/modified
original_content = f"Tool: {request_obj.tool_name}\nArguments: {request_obj.arguments}"
new_content = modified_messages[0].get("content", "") if modified_messages else ""
original_content = (
f"Tool: {request_obj.tool_name}\nArguments: {request_obj.arguments}"
)
new_content = (
modified_messages[0].get("content", "") if modified_messages else ""
)
if new_content != original_content:
# Content was modified - could be masking, redaction, or blocking
if not new_content or "blocked" in new_content.lower() or "violation" in new_content.lower():
if (
not new_content
or "blocked" in new_content.lower()
or "violation" in new_content.lower()
):
# Content was blocked completely
return MCPPreCallResponseObject(
should_proceed=False,
error_message="Content blocked by guardrail",
modified_arguments=None
modified_arguments=None,
)
else:
# Content was masked/redacted - extract the modified arguments
try:
# Try to parse the modified arguments from the masked content
modified_args = self._extract_modified_arguments_from_content(new_content, request_obj)
modified_args = (
self._extract_modified_arguments_from_content(
new_content, request_obj
)
)
if modified_args is not None:
# Return the masked/redacted arguments for the MCP call to use
return MCPPreCallResponseObject(
should_proceed=True,
error_message=None,
modified_arguments=modified_args
modified_arguments=modified_args,
)
else:
# Could not parse modified arguments, allow original call but warn
@ -630,129 +663,149 @@ class ProxyLogging:
)
return None
except Exception as e:
verbose_proxy_logger.error(f"Error parsing modified arguments: {e}")
verbose_proxy_logger.error(
f"Error parsing modified arguments: {e}"
)
# Fallback: allow original call
return None
# If result is a string, it's likely an error message
if isinstance(llm_result, str):
return MCPPreCallResponseObject(
should_proceed=False,
error_message=llm_result,
modified_arguments=None
should_proceed=False, error_message=llm_result, modified_arguments=None
)
return None
def _extract_modified_arguments_from_content(self, masked_content: str, request_obj) -> Optional[dict]:
def _extract_modified_arguments_from_content(
self, masked_content: str, request_obj
) -> Optional[dict]:
"""
Extract modified/masked arguments from the guardrail response content.
"""
import json
verbose_proxy_logger.debug(f"Extracting modified args from content: {masked_content}")
verbose_proxy_logger.debug(
f"Extracting modified args from content: {masked_content}"
)
try:
# The format should be: "Tool: <tool_name>\nArguments: <json_arguments>"
# Parse the arguments section
lines = masked_content.strip().split('\n')
lines = masked_content.strip().split("\n")
for i, line in enumerate(lines):
if line.startswith("Arguments:"):
# Get the arguments part - everything after "Arguments: "
args_text = line[len("Arguments:"):].strip()
args_text = line[len("Arguments:") :].strip()
verbose_proxy_logger.debug(f"Found arguments text: {args_text}")
# Try to parse as JSON first
try:
modified_args = json.loads(args_text)
verbose_proxy_logger.debug(f"Successfully parsed JSON args: {modified_args}")
verbose_proxy_logger.debug(
f"Successfully parsed JSON args: {modified_args}"
)
return modified_args
except json.JSONDecodeError as e:
# If JSON parsing fails, try to extract key-value pairs manually
verbose_proxy_logger.debug(f"Failed to parse JSON arguments: {args_text}, error: {e}")
return self._parse_arguments_manually(args_text, request_obj.arguments)
verbose_proxy_logger.debug(
f"Failed to parse JSON arguments: {args_text}, error: {e}"
)
return self._parse_arguments_manually(
args_text, request_obj.arguments
)
# If we can't find the Arguments: line, return None
verbose_proxy_logger.warning("Could not find 'Arguments:' line in masked content")
verbose_proxy_logger.warning(
"Could not find 'Arguments:' line in masked content"
)
return None
except Exception as e:
verbose_proxy_logger.error(f"Error extracting modified arguments: {e}")
return None
def _parse_arguments_manually(self, args_text: str, original_args: dict) -> Optional[dict]:
def _parse_arguments_manually(
self, args_text: str, original_args: dict
) -> Optional[dict]:
"""
Try to manually parse arguments when JSON parsing fails.
This is a fallback for cases where the guardrail modifies the format.
"""
import re
try:
# Start with original arguments and try to apply modifications
modified_args = original_args.copy()
# Look for simple key-value patterns
# This is a basic implementation - can be enhanced based on specific guardrail formats
for key, original_value in original_args.items():
if isinstance(original_value, str):
# Look for the key in the masked content and try to extract its value
pattern = rf"['\"]?{re.escape(key)}['\"]?\s*:\s*['\"]?([^,'\"]*)['\"]?"
pattern = (
rf"['\"]?{re.escape(key)}['\"]?\s*:\s*['\"]?([^,'\"]*)['\"]?"
)
match = re.search(pattern, args_text, re.IGNORECASE)
if match:
new_value = match.group(1).strip()
if new_value:
modified_args[key] = new_value
return modified_args
except Exception as e:
verbose_proxy_logger.error(f"Error in manual argument parsing: {e}")
return None
def _convert_llm_result_to_mcp_during_response(self, llm_result, request_obj) -> Optional[Any]:
def _convert_llm_result_to_mcp_during_response(
self, llm_result, request_obj
) -> Optional[Any]:
"""
Convert LLM guardrail result back to MCP during call response format.
"""
from litellm.types.mcp import MCPDuringCallResponseObject
# If result is an exception, it means the guardrail wants to stop execution
if isinstance(llm_result, Exception):
return MCPDuringCallResponseObject(
should_continue=False,
error_message=str(llm_result)
should_continue=False, error_message=str(llm_result)
)
# If result is a dict with modified messages, check for content filtering
if isinstance(llm_result, dict):
modified_messages = llm_result.get("messages")
if modified_messages:
# Check if content was blocked/modified
original_content = f"Tool: {request_obj.tool_name}\nArguments: {request_obj.arguments}"
new_content = modified_messages[0].get("content", "") if modified_messages else ""
original_content = (
f"Tool: {request_obj.tool_name}\nArguments: {request_obj.arguments}"
)
new_content = (
modified_messages[0].get("content", "") if modified_messages else ""
)
if new_content != original_content:
# Content was modified, could be masking or blocking
if not new_content or "blocked" in new_content.lower():
# Content was blocked
return MCPDuringCallResponseObject(
should_continue=False,
error_message="Content blocked by guardrail during execution"
error_message="Content blocked by guardrail during execution",
)
else:
# Content was masked/modified - for now, stop execution
# Content was masked/modified - for now, stop execution
return MCPDuringCallResponseObject(
should_continue=False,
error_message="Content modified by guardrail during execution"
error_message="Content modified by guardrail during execution",
)
# If result is a string, it's likely an error message
if isinstance(llm_result, str):
return MCPDuringCallResponseObject(
should_continue=False,
error_message=llm_result
should_continue=False, error_message=llm_result
)
return None
def get_combined_callback_list(
@ -799,7 +852,6 @@ class ProxyLogging:
from litellm.types.llms.base import HiddenParams
from litellm.types.mcp import MCPDuringCallRequestObject
callbacks = self.get_combined_callback_list(
dynamic_success_callbacks=getattr(self, "dynamic_success_callbacks", None),
global_callbacks=litellm.success_callback,
@ -820,7 +872,9 @@ class ProxyLogging:
_callback: Optional[CustomLogger] = None
if isinstance(callback, str):
from typing import cast
from litellm import _custom_logger_compatible_callbacks_literal
_callback = litellm.litellm_core_utils.litellm_logging.get_custom_logger_compatible_class(
cast(_custom_logger_compatible_callbacks_literal, callback)
)
@ -839,22 +893,29 @@ class ProxyLogging:
):
continue
# Convert MCP tool call to LLM message format for existing guardrail logic
synthetic_llm_data = self._convert_mcp_to_llm_format(request_obj, kwargs)
synthetic_llm_data = self._convert_mcp_to_llm_format(
request_obj, kwargs
)
# Reuse existing LLM guardrail logic for during call
user_api_key_auth_dict = self._convert_user_api_key_auth_to_dict(kwargs.get("user_api_key_auth"))
user_api_key_auth_dict = self._convert_user_api_key_auth_to_dict(
kwargs.get("user_api_key_auth")
)
result = await _callback.async_moderation_hook(
data=synthetic_llm_data,
user_api_key_dict=user_api_key_auth_dict,
call_type="mcp_call"
)
user_api_key_dict=user_api_key_auth_dict, # type: ignore
call_type="mcp_call",
)
# Convert result back to MCP response format if blocked/modified
if result is not None:
mcp_response = self._convert_llm_result_to_mcp_during_response(result, request_obj)
mcp_response = self._convert_llm_result_to_mcp_during_response(
result, request_obj
)
if mcp_response is not None:
return self._parse_during_mcp_call_hook_response(response=mcp_response)
return self._parse_during_mcp_call_hook_response(
response=mcp_response
)
except Exception as e:
raise e
@ -984,8 +1045,12 @@ class ProxyLogging:
custom_logger = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id(
prompt_id
)
prompt_spec = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(prompt_id)
litellm_prompt_id: Optional[str] = None
if prompt_spec is not None:
litellm_prompt_id = prompt_spec.litellm_params.prompt_id
if custom_logger:
if custom_logger and litellm_prompt_id is not None:
(
model,
messages,
@ -994,15 +1059,16 @@ class ProxyLogging:
model=data.get("model", ""),
messages=data.get("messages", []),
non_default_params=get_non_default_completion_params(kwargs=data),
prompt_id=prompt_id,
prompt_id=litellm_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.update(optional_params)
data["model"] = model
data["messages"] = messages
data.update(optional_params)
try:
for callback in litellm.callbacks:

View file

@ -1,6 +1,6 @@
from datetime import datetime
from enum import Enum
from typing import Dict, List, Literal, Optional
from typing import Any, Dict, List, Literal, Optional
from pydantic import BaseModel, ConfigDict
from typing_extensions import Required, TypedDict
@ -25,12 +25,34 @@ class PromptLiteLLMParams(BaseModel):
model_config = ConfigDict(extra="allow", protected_namespaces=())
class PromptSpec(TypedDict, total=False):
prompt_id: Required[str]
litellm_params: Required[PromptLiteLLMParams]
prompt_info: Optional[PromptInfo]
created_at: Optional[datetime]
updated_at: Optional[datetime]
class PromptSpec(BaseModel):
prompt_id: str
litellm_params: PromptLiteLLMParams
prompt_info: PromptInfo
created_at: Optional[datetime] = None
updated_at: Optional[datetime] = None
def __init__(self, **data):
if "prompt_info" not in data:
data["prompt_info"] = PromptInfo(prompt_type="config")
elif "prompt_info" in data:
if (
isinstance(data["prompt_info"], dict)
and data["prompt_info"].get("prompt_type") is None
):
data["prompt_info"]["prompt_type"] = "config"
super().__init__(**data)
class PromptTemplateBase(BaseModel):
litellm_prompt_id: str
content: str
metadata: Optional[Dict[str, Any]] = None
class PromptInfoResponse(BaseModel):
prompt_spec: PromptSpec
raw_prompt_template: Optional[PromptTemplateBase] = None
class ListPromptsResponse(BaseModel):

View file

@ -520,6 +520,16 @@ model LiteLLM_GuardrailsTable {
updated_at DateTime @updatedAt
}
// Prompt table for storing prompt configurations
model LiteLLM_PromptTable {
id String @id @default(uuid())
prompt_id String @unique
litellm_params Json
prompt_info Json?
created_at DateTime @default(now())
updated_at DateTime @updatedAt
}
model LiteLLM_HealthCheckTable {
health_check_id String @id @default(uuid())
model_name String

View file

@ -24,7 +24,7 @@ def test_prompt_manager_initialization():
prompt_dir = Path(
__file__
).parent # Current directory when running from tests/test_litellm/prompts
manager = PromptManager(prompt_dir)
manager = PromptManager(prompt_directory=str(prompt_dir))
# Should have loaded at least the sample prompts
assert len(manager.prompts) >= 3
@ -58,7 +58,7 @@ def test_render_simple_template():
prompt_dir = Path(
__file__
).parent # Current directory when running from tests/test_litellm/prompts
manager = PromptManager(prompt_dir)
manager = PromptManager(prompt_directory=str(prompt_dir))
# Test sample_prompt rendering
rendered = manager.render(
@ -74,7 +74,7 @@ def test_render_chat_prompt():
prompt_dir = Path(
__file__
).parent # Current directory when running from tests/test_litellm/prompts
manager = PromptManager(prompt_dir)
manager = PromptManager(prompt_directory=str(prompt_dir))
# Test with system context
rendered = manager.render(
@ -100,7 +100,7 @@ def test_render_coding_assistant():
prompt_dir = Path(
__file__
).parent # Current directory when running from tests/test_litellm/prompts
manager = PromptManager(prompt_dir)
manager = PromptManager(prompt_directory=str(prompt_dir))
rendered = manager.render(
"coding_assistant",
@ -136,7 +136,7 @@ input:
Hello {{name}}, you are {{age}} years old and {'active' if active else 'inactive'}."""
)
manager = PromptManager(temp_dir)
manager = PromptManager(prompt_directory=str(temp_dir))
# Valid input should work
rendered = manager.render(
@ -161,7 +161,7 @@ def test_prompt_not_found():
prompt_dir = Path(
__file__
).parent # Current directory when running from tests/test_litellm/prompts
manager = PromptManager(prompt_dir)
manager = PromptManager(prompt_directory=str(prompt_dir))
with pytest.raises(KeyError, match="Prompt 'nonexistent' not found"):
manager.render("nonexistent", {"some": "variable"})
@ -172,7 +172,7 @@ def test_list_prompts():
prompt_dir = Path(
__file__
).parent # Current directory when running from tests/test_litellm/prompts
manager = PromptManager(prompt_dir)
manager = PromptManager(prompt_directory=str(prompt_dir))
prompts = manager.list_prompts()
assert isinstance(prompts, list)
@ -186,7 +186,7 @@ def test_get_prompt_metadata():
prompt_dir = Path(
__file__
).parent # Current directory when running from tests/test_litellm/prompts
manager = PromptManager(prompt_dir)
manager = PromptManager(prompt_directory=str(prompt_dir))
metadata = manager.get_prompt_metadata("sample_prompt")
assert metadata is not None
@ -200,7 +200,7 @@ def test_add_prompt_programmatically():
prompt_dir = Path(
__file__
).parent # Current directory when running from tests/test_litellm/prompts
manager = PromptManager(prompt_dir)
manager = PromptManager(prompt_directory=str(prompt_dir))
initial_count = len(manager.prompts)
@ -238,21 +238,282 @@ Write about {{topic}}."""
prompt_without_frontmatter = Path(temp_dir) / "without_frontmatter.prompt"
prompt_without_frontmatter.write_text("Simple template: {{message}}")
manager = PromptManager(temp_dir)
manager = PromptManager(prompt_directory=str(temp_dir))
# Check frontmatter was parsed correctly
with_meta = manager.get_prompt("with_frontmatter")
assert with_meta is not None
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 is not None
assert without_meta.metadata == {}
rendered = manager.render("without_frontmatter", {"message": "Hello!"})
assert rendered == "Simple template: Hello!"
def test_prompt_manager_json_initialization():
"""Test PromptManager initialization with JSON data instead of directory."""
prompt_data = {
"json_test_prompt": {
"content": "Hello {{name}}! Welcome to {{service}}.",
"metadata": {"model": "gpt-4", "temperature": 0.8, "max_tokens": 150},
},
"simple_prompt": {
"content": "This is a simple prompt: {{message}}",
"metadata": {},
},
}
# Initialize PromptManager with JSON data only (no directory)
manager = PromptManager(prompt_data=prompt_data)
# Should have loaded the JSON prompts
assert len(manager.prompts) == 2
assert "json_test_prompt" in manager.prompts
assert "simple_prompt" in manager.prompts
# Test prompt properties
json_prompt = manager.get_prompt("json_test_prompt")
assert json_prompt.content == "Hello {{name}}! Welcome to {{service}}."
assert json_prompt.model == "gpt-4"
assert json_prompt.optional_params["temperature"] == 0.8
assert json_prompt.optional_params["max_tokens"] == 150
def test_prompt_manager_mixed_initialization():
"""Test PromptManager with both directory and JSON data."""
# Use existing directory
prompt_dir = Path(__file__).parent
# Add JSON data
json_data = {
"json_only_prompt": {
"content": "This prompt only exists in JSON: {{data}}",
"metadata": {"model": "gpt-3.5-turbo"},
}
}
manager = PromptManager(prompt_directory=str(prompt_dir), prompt_data=json_data)
# Should have prompts from both directory and JSON
assert "sample_prompt" in manager.prompts # From directory
assert "json_only_prompt" in manager.prompts # From JSON
# Test rendering both types
json_rendered = manager.render("json_only_prompt", {"data": "test"})
assert json_rendered == "This prompt only exists in JSON: test"
def test_load_prompts_from_json_data():
"""Test loading additional prompts from JSON data after initialization."""
# Start with directory-based manager
prompt_dir = Path(__file__).parent
manager = PromptManager(prompt_directory=str(prompt_dir))
initial_count = len(manager.prompts)
# Load additional prompts from JSON
additional_prompts = {
"dynamic_json_prompt": {
"content": "Dynamic prompt: {{dynamic_content}}",
"metadata": {"model": "claude-3", "temperature": 0.5},
},
"another_json_prompt": {
"content": "Another prompt with {{variable}}",
"metadata": {"model": "gpt-4"},
},
}
manager.load_prompts_from_json_data(additional_prompts)
# Should have added the new prompts
assert len(manager.prompts) == initial_count + 2
assert "dynamic_json_prompt" in manager.prompts
assert "another_json_prompt" in manager.prompts
# Test rendering the new prompts
rendered = manager.render("dynamic_json_prompt", {"dynamic_content": "test"})
assert rendered == "Dynamic prompt: test"
def test_prompt_file_to_json_conversion():
"""Test converting .prompt files to JSON format."""
# Create a temporary prompt file with frontmatter
with tempfile.TemporaryDirectory() as temp_dir:
prompt_file = Path(temp_dir) / "test_conversion.prompt"
prompt_file.write_text(
"""---
model: gpt-4
temperature: 0.7
max_tokens: 200
input:
schema:
user_input: string
context: string
output:
format: json
---
You are an AI assistant. Given the context: {{context}}
Please respond to: {{user_input}}"""
)
manager = PromptManager()
json_data = manager.prompt_file_to_json(prompt_file)
# Check the conversion
assert "content" in json_data
assert "metadata" in json_data
expected_content = """You are an AI assistant. Given the context: {{context}}
Please respond to: {{user_input}}"""
assert json_data["content"] == expected_content
metadata = json_data["metadata"]
assert metadata["model"] == "gpt-4"
assert metadata["temperature"] == 0.7
assert metadata["max_tokens"] == 200
assert metadata["input"]["schema"]["user_input"] == "string"
assert metadata["output"]["format"] == "json"
def test_json_to_prompt_file_conversion():
"""Test converting JSON data back to .prompt file format."""
json_data = {
"content": "Hello {{name}}! How can I help you with {{task}}?",
"metadata": {
"model": "gpt-3.5-turbo",
"temperature": 0.8,
"max_tokens": 100,
"input": {"schema": {"name": "string", "task": "string"}},
},
}
manager = PromptManager()
prompt_content = manager.json_to_prompt_file(json_data)
# Should have YAML frontmatter and content
assert prompt_content.startswith("---\n")
assert "---\n" in prompt_content[4:] # Second --- delimiter
assert "Hello {{name}}! How can I help you with {{task}}?" in prompt_content
assert "model: gpt-3.5-turbo" in prompt_content
assert "temperature: 0.8" in prompt_content
def test_json_to_prompt_file_without_metadata():
"""Test converting JSON with no metadata to .prompt format."""
json_data = {
"content": "Simple prompt without metadata: {{message}}",
"metadata": {},
}
manager = PromptManager()
prompt_content = manager.json_to_prompt_file(json_data)
# Should return just the content without frontmatter
assert prompt_content == "Simple prompt without metadata: {{message}}"
assert "---" not in prompt_content
def test_get_all_prompts_as_json():
"""Test exporting all prompts to JSON format."""
prompt_data = {
"prompt1": {
"content": "First prompt: {{var1}}",
"metadata": {"model": "gpt-4"},
},
"prompt2": {
"content": "Second prompt: {{var2}}",
"metadata": {"model": "claude-3", "temperature": 0.5},
},
}
manager = PromptManager(prompt_data=prompt_data)
all_prompts_json = manager.get_all_prompts_as_json()
assert len(all_prompts_json) == 2
assert "prompt1" in all_prompts_json
assert "prompt2" in all_prompts_json
# Check structure
prompt1_data = all_prompts_json["prompt1"]
assert prompt1_data["content"] == "First prompt: {{var1}}"
assert prompt1_data["metadata"]["model"] == "gpt-4"
prompt2_data = all_prompts_json["prompt2"]
assert prompt2_data["content"] == "Second prompt: {{var2}}"
assert prompt2_data["metadata"]["model"] == "claude-3"
def test_json_prompt_rendering_with_validation():
"""Test rendering JSON-based prompts with input validation."""
prompt_data = {
"validated_prompt": {
"content": "Process {{data}} for user {{user_id}}",
"metadata": {
"model": "gpt-4",
"input": {"schema": {"data": "string", "user_id": "integer"}},
},
}
}
manager = PromptManager(prompt_data=prompt_data)
# Valid input should work
rendered = manager.render("validated_prompt", {"data": "test data", "user_id": 123})
assert rendered == "Process test data for user 123"
# Invalid input should raise error
with pytest.raises(ValueError, match="Invalid type for field 'user_id'"):
manager.render(
"validated_prompt", {"data": "test data", "user_id": "not_an_int"}
)
def test_round_trip_conversion():
"""Test converting .prompt file to JSON and back to .prompt file."""
with tempfile.TemporaryDirectory() as temp_dir:
# Create original prompt file
original_file = Path(temp_dir) / "original.prompt"
original_content = """---
model: gpt-4
temperature: 0.6
---
Original prompt content: {{variable}}"""
original_file.write_text(original_content)
manager = PromptManager()
# Convert to JSON
json_data = manager.prompt_file_to_json(original_file)
# Convert back to prompt file format
converted_content = manager.json_to_prompt_file(json_data)
# Create new file with converted content
converted_file = Path(temp_dir) / "converted.prompt"
converted_file.write_text(converted_content)
# Load both files and compare
manager_original = PromptManager(prompt_directory=str(temp_dir))
original_template = manager_original.get_prompt("original")
converted_template = manager_original.get_prompt("converted")
# Content should be the same
assert original_template.content == converted_template.content
assert original_template.model == converted_template.model
assert (
original_template.optional_params["temperature"]
== converted_template.optional_params["temperature"]
)
def test_prompt_main():
"""
Integration test placeholder for litellm completion integration.

View file

@ -87,6 +87,17 @@ export interface PromptSpec {
updated_at?: string
}
export interface PromptTemplateBase {
litellm_prompt_id: string;
content: string;
metadata?: Record<string, any> | null;
}
export interface PromptInfoResponse {
prompt_spec: PromptSpec;
raw_prompt_template: PromptTemplateBase | null;
}
export interface ListPromptsResponse {
prompts: PromptSpec[];
}
@ -4899,8 +4910,8 @@ export const getGuardrailsList = async (accessToken: String) => {
export const getPromptsList = async (accessToken: String) : Promise<ListPromptsResponse> => {
try {
const url = proxyBaseUrl
? `${proxyBaseUrl}/prompt/list`
: `/prompt/list`;
? `${proxyBaseUrl}/prompts/list`
: `/prompts/list`;
const response = await fetch(url, {
method: "GET",
headers: {
@ -4923,9 +4934,9 @@ export const getPromptsList = async (accessToken: String) : Promise<ListPromptsR
}
};
export const getPromptInfo = async (accessToken: String, promptId: string): Promise<PromptSpec> => {
export const getPromptInfo = async (accessToken: String, promptId: string): Promise<PromptInfoResponse> => {
try {
const url = proxyBaseUrl ? `${proxyBaseUrl}/prompt/info?prompt_id=${promptId}` : `/prompt/info?prompt_id=${promptId}`;
const url = proxyBaseUrl ? `${proxyBaseUrl}/prompts/${promptId}/info` : `/prompts/${promptId}/info`;
const response = await fetch(url, {
method: "GET",
headers: {
@ -4948,6 +4959,160 @@ export const getPromptInfo = async (accessToken: String, promptId: string): Prom
}
};
export const createPromptCall = async (
accessToken: string,
promptData: any
) => {
try {
const url = proxyBaseUrl ? `${proxyBaseUrl}/prompts` : `/prompts`;
const response = await fetch(url, {
method: "POST",
headers: {
[globalLitellmHeaderName]: `Bearer ${accessToken}`,
"Content-Type": "application/json",
},
body: JSON.stringify(promptData),
});
if (!response.ok) {
const errorData = await response.text();
handleError(errorData);
throw new Error("Network response was not ok");
}
const data = await response.json();
return data;
} catch (error) {
console.error("Failed to create prompt:", error);
throw error;
}
};
export const updatePromptCall = async (
accessToken: string,
promptId: string,
promptData: any
) => {
try {
const url = proxyBaseUrl ? `${proxyBaseUrl}/prompts/${promptId}` : `/prompts/${promptId}`;
const response = await fetch(url, {
method: "PUT",
headers: {
[globalLitellmHeaderName]: `Bearer ${accessToken}`,
"Content-Type": "application/json",
},
body: JSON.stringify(promptData),
});
if (!response.ok) {
const errorData = await response.text();
handleError(errorData);
throw new Error("Network response was not ok");
}
const data = await response.json();
return data;
} catch (error) {
console.error("Failed to update prompt:", error);
throw error;
}
};
export const deletePromptCall = async (
accessToken: string,
promptId: string
) => {
try {
const url = proxyBaseUrl ? `${proxyBaseUrl}/prompts/${promptId}` : `/prompts/${promptId}`;
const response = await fetch(url, {
method: "DELETE",
headers: {
[globalLitellmHeaderName]: `Bearer ${accessToken}`,
"Content-Type": "application/json",
},
});
if (!response.ok) {
const errorData = await response.text();
handleError(errorData);
throw new Error("Network response was not ok");
}
const data = await response.json();
return data;
} catch (error) {
console.error("Failed to delete prompt:", error);
throw error;
}
};
export const convertPromptFileToJson = async (
accessToken: string,
file: File
): Promise<{ prompt_id: string; json_data: any }> => {
try {
const formData = new FormData();
formData.append("file", file);
const url = proxyBaseUrl
? `${proxyBaseUrl}/utils/dotprompt_json_converter`
: `/utils/dotprompt_json_converter`;
const response = await fetch(url, {
method: "POST",
headers: {
[globalLitellmHeaderName]: `Bearer ${accessToken}`,
},
body: formData,
});
if (!response.ok) {
const errorData = await response.text();
handleError(errorData);
throw new Error("Network response was not ok");
}
return await response.json();
} catch (error) {
console.error("Failed to convert prompt file:", error);
throw error;
}
};
export const patchPromptCall = async (
accessToken: string,
promptId: string,
promptData: any
) => {
try {
const url = proxyBaseUrl ? `${proxyBaseUrl}/prompts/${promptId}` : `/prompts/${promptId}`;
const response = await fetch(url, {
method: "PATCH",
headers: {
[globalLitellmHeaderName]: `Bearer ${accessToken}`,
"Content-Type": "application/json",
},
body: JSON.stringify(promptData),
});
if (!response.ok) {
const errorData = await response.text();
handleError(errorData);
throw new Error("Network response was not ok");
}
const data = await response.json();
return data;
} catch (error) {
console.error("Failed to patch prompt:", error);
throw error;
}
};
export const createGuardrailCall = async (
accessToken: string,
guardrailData: any

View file

@ -1,8 +1,12 @@
import React, { useState, useEffect } from "react"
import { Card, Text } from "@tremor/react"
import { getPromptsList, PromptSpec, ListPromptsResponse } from "./networking"
import { Card, Text, Button } from "@tremor/react"
import { Modal, message } from "antd"
import { getPromptsList, PromptSpec, ListPromptsResponse, deletePromptCall } from "./networking"
import PromptTable from "./prompts/prompt_table"
import PromptInfoView from "./prompts/prompt_info"
import AddPromptForm from "./prompts/add_prompt_form"
import { isAdminRole } from "@/utils/roles"
interface PromptsProps {
@ -14,6 +18,9 @@ const PromptsPanel: React.FC<PromptsProps> = ({ accessToken, userRole }) => {
const [promptsList, setPromptsList] = useState<PromptSpec[]>([])
const [isLoading, setIsLoading] = useState(false)
const [selectedPromptId, setSelectedPromptId] = useState<string | null>(null)
const [isAddModalVisible, setIsAddModalVisible] = useState(false)
const [isDeleting, setIsDeleting] = useState(false)
const [promptToDelete, setPromptToDelete] = useState<{ id: string; name: string } | null>(null)
const isAdmin = userRole ? isAdminRole(userRole) : false
@ -38,10 +45,51 @@ const PromptsPanel: React.FC<PromptsProps> = ({ accessToken, userRole }) => {
fetchPrompts()
}, [accessToken])
const handlePromptClick = (promptId: string) => {
const handlePromptClick = (promptId: string) => {
setSelectedPromptId(promptId)
}
const handleAddPrompt = () => {
if (selectedPromptId) {
setSelectedPromptId(null)
}
setIsAddModalVisible(true)
}
const handleCloseModal = () => {
setIsAddModalVisible(false)
}
const handleSuccess = () => {
fetchPrompts()
}
const handleDeleteClick = (promptId: string, promptName: string) => {
setPromptToDelete({ id: promptId, name: promptName })
}
const handleDeleteConfirm = async () => {
if (!promptToDelete || !accessToken) return
setIsDeleting(true)
try {
await deletePromptCall(accessToken, promptToDelete.id)
message.success(`Prompt "${promptToDelete.name}" deleted successfully`)
fetchPrompts() // Refresh the list
} catch (error) {
console.error("Error deleting prompt:", error)
message.error("Failed to delete prompt")
} finally {
setIsDeleting(false)
setPromptToDelete(null)
}
}
const handleDeleteCancel = () => {
setPromptToDelete(null)
}
return (
<div className="w-full mx-auto flex-auto overflow-y-auto m-8 p-2">
{selectedPromptId ? (
@ -50,20 +98,48 @@ const PromptsPanel: React.FC<PromptsProps> = ({ accessToken, userRole }) => {
onClose={() => setSelectedPromptId(null)}
accessToken={accessToken}
isAdmin={isAdmin}
onDelete={fetchPrompts}
/>
) : (
<>
<div className="flex justify-between items-center mb-4">
<Text className="text-lg font-semibold">Prompts</Text>
<Button onClick={handleAddPrompt} disabled={!accessToken}>
+ Add New Prompt
</Button>
</div>
<PromptTable
promptsList={promptsList}
isLoading={isLoading}
onPromptClick={handlePromptClick}
onDeleteClick={handleDeleteClick}
accessToken={accessToken}
isAdmin={isAdmin}
/>
</>
)}
<AddPromptForm
visible={isAddModalVisible}
onClose={handleCloseModal}
accessToken={accessToken}
onSuccess={handleSuccess}
/>
{promptToDelete && (
<Modal
title="Delete Prompt"
open={promptToDelete !== null}
onOk={handleDeleteConfirm}
onCancel={handleDeleteCancel}
confirmLoading={isDeleting}
okText="Delete"
okButtonProps={{ danger: true }}
>
<p>Are you sure you want to delete prompt: {promptToDelete.name} ?</p>
<p>This action cannot be undone.</p>
</Modal>
)}
</div>
)
}

View file

@ -0,0 +1,201 @@
import React, { useState } from "react"
import { Modal, Form, Select, Upload, Button, message, Divider } from "antd"
import { TextInput } from "@tremor/react"
import { UploadOutlined } from "@ant-design/icons"
import type { UploadFile, UploadProps } from "antd"
import { convertPromptFileToJson, createPromptCall } from "../networking"
const { Option } = Select
interface AddPromptFormProps {
visible: boolean
onClose: () => void
accessToken: string | null
onSuccess: () => void
}
interface PromptFormData {
prompt_id: string
prompt_integration: string
prompt_file?: File
}
const AddPromptForm: React.FC<AddPromptFormProps> = ({
visible,
onClose,
accessToken,
onSuccess,
}) => {
const [form] = Form.useForm()
const [loading, setLoading] = useState(false)
const [fileList, setFileList] = useState<UploadFile[]>([])
const [promptIntegration, setPromptIntegration] = useState<string>("dotprompt")
const handleCancel = () => {
form.resetFields()
setFileList([])
setPromptIntegration("dotprompt")
onClose()
}
const handleSubmit = async () => {
try {
const values = await form.validateFields()
console.log("values: ", values)
if (!accessToken) {
message.error("Access token is required")
return
}
if (promptIntegration === "dotprompt" && fileList.length === 0) {
message.error("Please upload a .prompt file")
return
}
setLoading(true)
let promptData: any = {}
if (promptIntegration === "dotprompt" && fileList.length > 0) {
// Convert the uploaded file to JSON
const file = fileList[0].originFileObj as File
try {
const conversionResult = await convertPromptFileToJson(accessToken, file)
console.log("Conversion result:", conversionResult)
// Prepare prompt data for creation
promptData = {
prompt_id: values.prompt_id,
litellm_params: {
prompt_integration: "dotprompt",
prompt_id: conversionResult.prompt_id,
prompt_data: conversionResult.json_data
},
prompt_info: {
prompt_type: "db"
}
}
} catch (conversionError) {
console.error("Error converting prompt file:", conversionError)
message.error("Failed to convert prompt file to JSON")
setLoading(false)
return
}
}
// Create the prompt
try {
await createPromptCall(accessToken, promptData)
message.success("Prompt created successfully!")
handleCancel()
onSuccess()
} catch (createError) {
console.error("Error creating prompt:", createError)
message.error("Failed to create prompt")
}
} catch (error) {
console.error("Form validation error:", error)
} finally {
setLoading(false)
}
}
const uploadProps: UploadProps = {
beforeUpload: (file) => {
if (!file.name.endsWith('.prompt')) {
message.error('Please upload a .prompt file')
return false
}
return false // Prevent automatic upload
},
fileList,
onChange: ({ fileList: newFileList }) => {
setFileList(newFileList.slice(-1)) // Keep only the last file
},
onRemove: () => {
setFileList([])
},
}
return (
<Modal
title="Add New Prompt"
open={visible}
onCancel={handleCancel}
footer={[
<Button key="cancel" onClick={handleCancel}>
Cancel
</Button>,
<Button
key="submit"
loading={loading}
onClick={handleSubmit}
>
Create Prompt
</Button>,
]}
width={600}
>
<Form
form={form}
layout="vertical"
requiredMark={false}
>
<Form.Item
label="Prompt ID"
name="prompt_id"
rules={[
{ required: true, message: "Please enter a prompt ID" },
{
pattern: /^[a-zA-Z0-9_-]+$/,
message: "Prompt ID can only contain letters, numbers, underscores, and hyphens"
}
]}
>
<TextInput
placeholder="Enter unique prompt ID (e.g., my_prompt_id)"
/>
</Form.Item>
<Form.Item
label="Prompt Integration"
name="prompt_integration"
initialValue="dotprompt"
>
<Select
value={promptIntegration}
onChange={setPromptIntegration}
>
<Option value="dotprompt">dotprompt</Option>
</Select>
</Form.Item>
{promptIntegration === "dotprompt" && (
<>
<Divider />
<Form.Item
label="Prompt File"
extra="Upload a .prompt file that follows the Dotprompt specification"
>
<Upload {...uploadProps}>
<Button icon={<UploadOutlined />}>
Select .prompt File
</Button>
</Upload>
{fileList.length > 0 && (
<div className="mt-2 text-sm text-gray-600">
Selected: {fileList[0].name}
</div>
)}
</Form.Item>
</>
)}
</Form>
</Modal>
)
}
export default AddPromptForm

View file

@ -12,9 +12,9 @@ import {
TabPanel,
TabPanels,
} from "@tremor/react"
import { Button, message, Tooltip } from "antd"
import { ArrowLeftIcon } from "@heroicons/react/outline"
import { getPromptInfo, PromptSpec } from "@/components/networking"
import { Button, message, Tooltip, Modal } from "antd"
import { ArrowLeftIcon, TrashIcon } from "@heroicons/react/outline"
import { getPromptInfo, PromptInfoResponse, PromptSpec, PromptTemplateBase, deletePromptCall } from "@/components/networking"
import { copyToClipboard as utilCopyToClipboard } from "@/utils/dataUtils"
import { CheckIcon, CopyIcon } from "lucide-react"
@ -23,20 +23,25 @@ export interface PromptInfoProps {
onClose: () => void
accessToken: string | null
isAdmin: boolean
onDelete?: () => void
}
const PromptInfoView: React.FC<PromptInfoProps> = ({ promptId, onClose, accessToken, isAdmin }) => {
const PromptInfoView: React.FC<PromptInfoProps> = ({ promptId, onClose, accessToken, isAdmin, onDelete }) => {
const [promptData, setPromptData] = useState<PromptSpec | null>(null)
const [promptTemplate, setPromptTemplate] = useState<PromptTemplateBase | null>(null)
const [rawApiResponse, setRawApiResponse] = useState<any>(null)
const [loading, setLoading] = useState(true)
const [copiedStates, setCopiedStates] = useState<Record<string, boolean>>({})
const [showDeleteConfirm, setShowDeleteConfirm] = useState(false)
const [isDeleting, setIsDeleting] = useState(false)
const fetchPromptInfo = async () => {
try {
setLoading(true)
if (!accessToken) return
const response = await getPromptInfo(accessToken, promptId)
setPromptData(response)
setPromptData(response.prompt_spec)
setPromptTemplate(response.raw_prompt_template)
setRawApiResponse(response) // Store the raw response for the Raw JSON tab
} catch (error) {
message.error("Failed to load prompt information")
@ -75,32 +80,73 @@ const PromptInfoView: React.FC<PromptInfoProps> = ({ promptId, onClose, accessTo
}
}
const handleDeleteClick = () => {
setShowDeleteConfirm(true)
}
const handleDeleteConfirm = async () => {
if (!accessToken || !promptData) return
setIsDeleting(true)
try {
await deletePromptCall(accessToken, promptData.prompt_id)
message.success(`Prompt "${promptData.prompt_id}" deleted successfully`)
onDelete?.() // Call the callback to refresh the parent component
onClose() // Close the info view
} catch (error) {
console.error("Error deleting prompt:", error)
message.error("Failed to delete prompt")
} finally {
setIsDeleting(false)
setShowDeleteConfirm(false)
}
}
const handleDeleteCancel = () => {
setShowDeleteConfirm(false)
}
return (
<div className="p-4">
<div>
<TremorButton icon={ArrowLeftIcon} variant="light" onClick={onClose} className="mb-4">
Back to Prompts
</TremorButton>
<Title>Prompt Details</Title>
<div className="flex items-center cursor-pointer">
<Text className="text-gray-500 font-mono">{promptData.prompt_id}</Text>
<Button
type="text"
size="small"
icon={copiedStates["prompt-id"] ? <CheckIcon size={12} /> : <CopyIcon size={12} />}
onClick={() => copyToClipboard(promptData.prompt_id, "prompt-id")}
className={`left-2 z-10 transition-all duration-200 ${
copiedStates["prompt-id"]
? "text-green-600 bg-green-50 border-green-200"
: "text-gray-500 hover:text-gray-700 hover:bg-gray-100"
}`}
/>
<div className="flex justify-between items-start mb-4">
<div>
<Title>Prompt Details</Title>
<div className="flex items-center cursor-pointer">
<Text className="text-gray-500 font-mono">{promptData.prompt_id}</Text>
<Button
type="text"
size="small"
icon={copiedStates["prompt-id"] ? <CheckIcon size={12} /> : <CopyIcon size={12} />}
onClick={() => copyToClipboard(promptData.prompt_id, "prompt-id")}
className={`left-2 z-10 transition-all duration-200 ${
copiedStates["prompt-id"]
? "text-green-600 bg-green-50 border-green-200"
: "text-gray-500 hover:text-gray-700 hover:bg-gray-100"
}`}
/>
</div>
</div>
{isAdmin && (
<TremorButton
icon={TrashIcon}
variant="secondary"
onClick={handleDeleteClick}
className="flex items-center"
>
Delete Prompt
</TremorButton>
)}
</div>
</div>
<TabGroup>
<TabList className="mb-4">
<Tab key="overview">Overview</Tab>
{promptTemplate ? <Tab key="prompt-template">Prompt Template</Tab> : <></>}
{isAdmin ? <Tab key="details">Details</Tab> : <></>}
<Tab key="raw-json">Raw JSON</Tab>
</TabList>
@ -147,6 +193,55 @@ const PromptInfoView: React.FC<PromptInfoProps> = ({ promptId, onClose, accessTo
)}
</TabPanel>
{/* Prompt Template Panel */}
{promptTemplate && (
<TabPanel>
<Card>
<div className="flex justify-between items-center mb-4">
<Title>Prompt Template</Title>
<Button
type="text"
size="small"
icon={copiedStates["prompt-content"] ? <CheckIcon size={16} /> : <CopyIcon size={16} />}
onClick={() => copyToClipboard(promptTemplate.content, "prompt-content")}
className={`transition-all duration-200 ${
copiedStates["prompt-content"]
? "text-green-600 bg-green-50 border-green-200"
: "text-gray-500 hover:text-gray-700 hover:bg-gray-100"
}`}
>
{copiedStates["prompt-content"] ? "Copied!" : "Copy Content"}
</Button>
</div>
<div className="space-y-4">
<div>
<Text className="font-medium">Template ID</Text>
<div className="font-mono text-sm bg-gray-50 p-2 rounded">{promptTemplate.litellm_prompt_id}</div>
</div>
<div>
<Text className="font-medium">Content</Text>
<div className="mt-2 p-4 bg-gray-50 rounded-md border overflow-auto max-h-96">
<pre className="text-sm text-gray-800 whitespace-pre-wrap">{promptTemplate.content}</pre>
</div>
</div>
{promptTemplate.metadata && Object.keys(promptTemplate.metadata).length > 0 && (
<div>
<Text className="font-medium">Template Metadata</Text>
<div className="mt-2 p-3 bg-gray-50 rounded-md border">
<pre className="text-xs text-gray-800 whitespace-pre-wrap overflow-auto max-h-64">
{JSON.stringify(promptTemplate.metadata, null, 2)}
</pre>
</div>
</div>
)}
</div>
</Card>
</TabPanel>
)}
{/* Details Panel (only for admins) */}
{isAdmin && (
<TabPanel>
@ -224,6 +319,20 @@ const PromptInfoView: React.FC<PromptInfoProps> = ({ promptId, onClose, accessTo
</TabPanel>
</TabPanels>
</TabGroup>
{/* Delete Confirmation Modal */}
<Modal
title="Delete Prompt"
open={showDeleteConfirm}
onOk={handleDeleteConfirm}
onCancel={handleDeleteCancel}
confirmLoading={isDeleting}
okText="Delete"
okButtonProps={{ danger: true }}
>
<p>Are you sure you want to delete prompt: <strong>{promptData?.prompt_id}</strong>?</p>
<p>This action cannot be undone.</p>
</Modal>
</div>
)
}

View file

@ -1,6 +1,6 @@
import React, { useState } from "react"
import { Table, TableBody, TableCell, TableHead, TableHeaderCell, TableRow, Button } from "@tremor/react"
import { SwitchVerticalIcon, ChevronUpIcon, ChevronDownIcon } from "@heroicons/react/outline"
import { SwitchVerticalIcon, ChevronUpIcon, ChevronDownIcon, TrashIcon } from "@heroicons/react/outline"
import { Tooltip } from "antd"
import {PromptSpec} from "@/components/networking"
import {
@ -17,12 +17,18 @@ interface PromptTableProps {
promptsList: PromptSpec[]
isLoading: boolean
onPromptClick?: (id: string) => void
onDeleteClick?: (id: string, name: string) => void
accessToken: string | null
isAdmin: boolean
}
const PromptTable: React.FC<PromptTableProps> = ({
promptsList,
isLoading,
onPromptClick,
onDeleteClick,
accessToken,
isAdmin,
}) => {
const [sorting, setSorting] = useState<SortingState>([{ id: "created_at", desc: true }])
@ -86,6 +92,33 @@ const PromptTable: React.FC<PromptTableProps> = ({
)
},
},
...(isAdmin ? [{
header: "Actions",
id: "actions",
enableSorting: false,
cell: ({ row }: any) => {
const prompt = row.original
const promptName = prompt.prompt_id || "Unknown Prompt"
return (
<div className="flex items-center gap-1">
<Tooltip title="Delete prompt">
<Button
size="xs"
variant="light"
color="red"
onClick={(e) => {
e.stopPropagation()
onDeleteClick?.(prompt.prompt_id, promptName)
}}
icon={TrashIcon}
className="text-red-500 hover:text-red-700 hover:bg-red-50"
/>
</Tooltip>
</div>
)
},
}] : []),
]
const table = useReactTable({