diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20250802162330_prompt_table/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20250802162330_prompt_table/migration.sql new file mode 100644 index 00000000000..e5c00ef4adb --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20250802162330_prompt_table/migration.sql @@ -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"); + diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 2bea0225d93..b8f2201d6b5 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -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 diff --git a/litellm/integrations/dotprompt/__init__.py b/litellm/integrations/dotprompt/__init__.py index a50e8f5e55e..3af7fbf6dd3 100644 --- a/litellm/integrations/dotprompt/__init__.py +++ b/litellm/integrations/dotprompt/__init__.py @@ -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 = { diff --git a/litellm/integrations/dotprompt/dotprompt_manager.py b/litellm/integrations/dotprompt/dotprompt_manager.py index cf3cfa374b9..0f0d7b938f3 100644 --- a/litellm/integrations/dotprompt/dotprompt_manager.py +++ b/litellm/integrations/dotprompt/dotprompt_manager.py @@ -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) diff --git a/litellm/integrations/dotprompt/prompt_manager.py b/litellm/integrations/dotprompt/prompt_manager.py index c8bfd6e68b9..9623ddab5fb 100644 --- a/litellm/integrations/dotprompt/prompt_manager.py +++ b/litellm/integrations/dotprompt/prompt_manager.py @@ -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) diff --git a/litellm/integrations/prompt_management_base.py b/litellm/integrations/prompt_management_base.py index 4a8bcd2e249..34b4455f564 100644 --- a/litellm/integrations/prompt_management_base.py +++ b/litellm/integrations/prompt_management_base.py @@ -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( diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index c7ee6ca0612..e1ea9918e68 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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 ] diff --git a/litellm/proxy/_new_secret_config.yaml b/litellm/proxy/_new_secret_config.yaml index cf634684ba4..ad4dabf9d04 100644 --- a/litellm/proxy/_new_secret_config.yaml +++ b/litellm/proxy/_new_secret_config.yaml @@ -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"] \ No newline at end of file diff --git a/litellm/proxy/management_endpoints/common_utils.py b/litellm/proxy/management_endpoints/common_utils.py index 4e22edd7eab..fcab831e3e3 100644 --- a/litellm/proxy/management_endpoints/common_utils.py +++ b/litellm/proxy/management_endpoints/common_utils.py @@ -1,7 +1,6 @@ from typing import Any, Optional, Union from litellm.proxy._types import ( - GenerateKeyRequest, KeyRequestBase, LiteLLM_ManagementEndpoint_MetadataFields_Premium, LiteLLM_TeamTable, diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 191815b78ed..88eed7fe2a3 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -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. diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 58084874f4f..e3a656f3209 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -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" diff --git a/litellm/proxy/prompts/init_prompts.py b/litellm/proxy/prompts/init_prompts.py index b2b7ca0cc47..a39f06b1242 100644 --- a/litellm/proxy/prompts/init_prompts.py +++ b/litellm/proxy/prompts/init_prompts.py @@ -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: diff --git a/litellm/proxy/prompts/prompt_endpoints.py b/litellm/proxy/prompts/prompt_endpoints.py index 2b8d533c70b..8a3b60ea8ad 100644 --- a/litellm/proxy/prompts/prompt_endpoints.py +++ b/litellm/proxy/prompts/prompt_endpoints.py @@ -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 " + ``` + + 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 " + ``` + + 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 " \\ + -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 " \\ + -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 " + ``` + + 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 " \\ + -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 diff --git a/litellm/proxy/prompts/prompt_registry.py b/litellm/proxy/prompts/prompt_registry.py index 49a25998d66..a6fc377c1a7 100644 --- a/litellm/proxy/prompts/prompt_registry.py +++ b/litellm/proxy/prompts/prompt_registry.py @@ -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 diff --git a/litellm/proxy/proxy_cli.py b/litellm/proxy/proxy_cli.py index d730b0621f5..80b2f69e67e 100644 --- a/litellm/proxy/proxy_cli.py +++ b/litellm/proxy/proxy_cli.py @@ -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 diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 7c82f5de709..af9f5d363da 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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 ( diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 2bea0225d93..b8f2201d6b5 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -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 diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 9e05a90b940..f1e113d8ded 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -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: \nArguments: " # 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: diff --git a/litellm/types/prompts/init_prompts.py b/litellm/types/prompts/init_prompts.py index 43489be131f..e4172aca4a1 100644 --- a/litellm/types/prompts/init_prompts.py +++ b/litellm/types/prompts/init_prompts.py @@ -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): diff --git a/schema.prisma b/schema.prisma index 2bea0225d93..b8f2201d6b5 100644 --- a/schema.prisma +++ b/schema.prisma @@ -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 diff --git a/tests/test_litellm/integrations/dotprompt/test_prompt_manager.py b/tests/test_litellm/integrations/dotprompt/test_prompt_manager.py index 3259eec571c..c19503641e1 100644 --- a/tests/test_litellm/integrations/dotprompt/test_prompt_manager.py +++ b/tests/test_litellm/integrations/dotprompt/test_prompt_manager.py @@ -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. diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index ca0a79485a4..04be4fd7c82 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -87,6 +87,17 @@ export interface PromptSpec { updated_at?: string } +export interface PromptTemplateBase { + litellm_prompt_id: string; + content: string; + metadata?: Record | 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 => { 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 => { +export const getPromptInfo = async (accessToken: String, promptId: string): Promise => { 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 diff --git a/ui/litellm-dashboard/src/components/prompts.tsx b/ui/litellm-dashboard/src/components/prompts.tsx index b2ccc4d7f80..477a7c93e6d 100644 --- a/ui/litellm-dashboard/src/components/prompts.tsx +++ b/ui/litellm-dashboard/src/components/prompts.tsx @@ -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 = ({ accessToken, userRole }) => { const [promptsList, setPromptsList] = useState([]) const [isLoading, setIsLoading] = useState(false) const [selectedPromptId, setSelectedPromptId] = useState(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 = ({ 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 (
{selectedPromptId ? ( @@ -50,20 +98,48 @@ const PromptsPanel: React.FC = ({ accessToken, userRole }) => { onClose={() => setSelectedPromptId(null)} accessToken={accessToken} isAdmin={isAdmin} + onDelete={fetchPrompts} /> ) : ( <>
- Prompts +
)} + + + + {promptToDelete && ( + +

Are you sure you want to delete prompt: {promptToDelete.name} ?

+

This action cannot be undone.

+
+ )}
) } diff --git a/ui/litellm-dashboard/src/components/prompts/add_prompt_form.tsx b/ui/litellm-dashboard/src/components/prompts/add_prompt_form.tsx new file mode 100644 index 00000000000..a9dd143fee4 --- /dev/null +++ b/ui/litellm-dashboard/src/components/prompts/add_prompt_form.tsx @@ -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 = ({ + visible, + onClose, + accessToken, + onSuccess, +}) => { + const [form] = Form.useForm() + const [loading, setLoading] = useState(false) + const [fileList, setFileList] = useState([]) + const [promptIntegration, setPromptIntegration] = useState("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 ( + + Cancel + , + , + ]} + width={600} + > +
+ + + + + + + + + {promptIntegration === "dotprompt" && ( + <> + + + + + + {fileList.length > 0 && ( +
+ Selected: {fileList[0].name} +
+ )} +
+ + )} + +
+ ) +} + +export default AddPromptForm \ No newline at end of file diff --git a/ui/litellm-dashboard/src/components/prompts/prompt_info.tsx b/ui/litellm-dashboard/src/components/prompts/prompt_info.tsx index aa309429334..07587891610 100644 --- a/ui/litellm-dashboard/src/components/prompts/prompt_info.tsx +++ b/ui/litellm-dashboard/src/components/prompts/prompt_info.tsx @@ -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 = ({ promptId, onClose, accessToken, isAdmin }) => { +const PromptInfoView: React.FC = ({ promptId, onClose, accessToken, isAdmin, onDelete }) => { const [promptData, setPromptData] = useState(null) + const [promptTemplate, setPromptTemplate] = useState(null) const [rawApiResponse, setRawApiResponse] = useState(null) const [loading, setLoading] = useState(true) const [copiedStates, setCopiedStates] = useState>({}) + 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 = ({ 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 (
Back to Prompts - Prompt Details -
- {promptData.prompt_id} -
+
+ {isAdmin && ( + + Delete Prompt + + )}
Overview + {promptTemplate ? Prompt Template : <>} {isAdmin ? Details : <>} Raw JSON @@ -147,6 +193,55 @@ const PromptInfoView: React.FC = ({ promptId, onClose, accessTo )} + {/* Prompt Template Panel */} + {promptTemplate && ( + + +
+ Prompt Template + +
+ +
+
+ Template ID +
{promptTemplate.litellm_prompt_id}
+
+ +
+ Content +
+
{promptTemplate.content}
+
+
+ + {promptTemplate.metadata && Object.keys(promptTemplate.metadata).length > 0 && ( +
+ Template Metadata +
+
+                          {JSON.stringify(promptTemplate.metadata, null, 2)}
+                        
+
+
+ )} +
+
+
+ )} + {/* Details Panel (only for admins) */} {isAdmin && ( @@ -224,6 +319,20 @@ const PromptInfoView: React.FC = ({ promptId, onClose, accessTo
+ + {/* Delete Confirmation Modal */} + +

Are you sure you want to delete prompt: {promptData?.prompt_id}?

+

This action cannot be undone.

+
) } diff --git a/ui/litellm-dashboard/src/components/prompts/prompt_table.tsx b/ui/litellm-dashboard/src/components/prompts/prompt_table.tsx index 6c2417cd387..31267b326ef 100644 --- a/ui/litellm-dashboard/src/components/prompts/prompt_table.tsx +++ b/ui/litellm-dashboard/src/components/prompts/prompt_table.tsx @@ -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 = ({ promptsList, isLoading, onPromptClick, + onDeleteClick, + accessToken, + isAdmin, }) => { const [sorting, setSorting] = useState([{ id: "created_at", desc: true }]) @@ -86,6 +92,33 @@ const PromptTable: React.FC = ({ ) }, }, + ...(isAdmin ? [{ + header: "Actions", + id: "actions", + enableSorting: false, + cell: ({ row }: any) => { + const prompt = row.original + const promptName = prompt.prompt_id || "Unknown Prompt" + + return ( +
+ +
+ ) + }, + }] : []), ] const table = useReactTable({