mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
Squashed commit of the following:
commit3119064e94Author: Krish Dholakia <krrishdholakia@gmail.com> Date: Sat Aug 2 22:33:37 2025 -0700 Prompt Management - add prompts on UI (#13240) * fix(create_key_button.tsx): add prompts on UI * feat(key_management_endpoints.py): support adding prompt to key via `/key/update` * fix(key_info_view.tsx): show existing prompts on key in key_info_view.tsx * fix(key_edit_view.tsx): UX - disable premium feature for non-premium users prevent accidental clicking * fix(create_key_button.tsx): disable premium features behind flag, prevent errors * feat(prompts.tsx): add new ui component to view created prompts enables viewing prompts created on config * feat(prompt_info.tsx): add component for viewing the prompt information * feat(prompt_endpoints.py): support converting dotprompt to json structure + accept json structure in promptmanager allows prompt manager to work with api endpoints * test(test_prompt_manager.py): add unit tests for json data input * feat(dotprompt/__init__.py): add prompt data to dotpromptmanager * fix(prompt_endpoints.py): working crud endpoints for prompt management * feat(prompts/): support `prompt_file` for dotprompt allows to precisely point to the prompt file a prompt should use * feat(proxy/utils.py): resolve prompt id correctly resolves user sent prompt id with internal prompt id * feat(schema.prisma): initial pr with db schema for prompt management table allows post endpoints to work with backend * feat(prompt_endpoints.py): use db in patch_prompt endpoint * feat(prompt_endpoints.py): use db for update_prompt endpoint * feat(prompt_endpoints.py): use db on prompt delete endpoint * build(schema.prisma): add prompt tale to schema.prisma in litellm-proxy-extras * build(migration.sql): add new sql migration file * fix(init_prompts.py): fix init * feat(prompt_info_view.tsx): show the raw prompt template on ui allows developer to know the prompt template they'll be calling * feat(add_prompt_form.tsx): working ui add prompt flow allows user to add prompts to litellm via ui * build(ui/): styling fixes * build(ui/): prompts.tsx styling improvements * fix(add_prompt_form.tsx): styling improvements * build(prompts.tsx): styling improvements * build(ui/): styling improvements * build(ui/): fix ui error * fix: fix ruff check * docs: document new api params * test: update tests commite47b30a76dAuthor: Krish Dholakia <krrishdholakia@gmail.com> Date: Sat Aug 2 19:34:18 2025 -0700 Prompt Management - Add table + prompt info page to UI (#13232) * fix(create_key_button.tsx): add prompts on UI * feat(key_management_endpoints.py): support adding prompt to key via `/key/update` * fix(key_info_view.tsx): show existing prompts on key in key_info_view.tsx * fix(key_edit_view.tsx): UX - disable premium feature for non-premium users prevent accidental clicking * fix(create_key_button.tsx): disable premium features behind flag, prevent errors * feat(prompts.tsx): add new ui component to view created prompts enables viewing prompts created on config * feat(prompt_info.tsx): add component for viewing the prompt information commitb79f55eec0Author: Krish Dholakia <krrishdholakia@gmail.com> Date: Sat Aug 2 19:33:17 2025 -0700 UI - Add giving keys prompt access (#13233) * fix(create_key_button.tsx): add prompts on UI * feat(key_management_endpoints.py): support adding prompt to key via `/key/update` * fix(key_info_view.tsx): show existing prompts on key in key_info_view.tsx * fix(key_edit_view.tsx): UX - disable premium feature for non-premium users prevent accidental clicking * fix(create_key_button.tsx): disable premium features behind flag, prevent errors * fix(key_management_endpoints.py): fix key update logic \ * fix: fix check * docs: document new params commit4c217c66f5Author: Ishaan Jaff <ishaanjaffer0324@gmail.com> Date: Sat Aug 2 17:26:40 2025 -0700 docs User Agent Activity Tracking commit2ee4e84406Author: Ishaan Jaff <ishaanjaffer0324@gmail.com> Date: Sat Aug 2 16:47:44 2025 -0700 docs fix commit06856b4d37Author: Ishaan Jaff <ishaanjaffer0324@gmail.com> Date: Sat Aug 2 15:47:09 2025 -0700 docs fix commit0f9f5f7a6cAuthor: Ishaan Jaff <ishaanjaffer0324@gmail.com> Date: Sat Aug 2 15:44:59 2025 -0700 docs fix commit69a360429cAuthor: Ishaan Jaff <ishaanjaffer0324@gmail.com> Date: Sat Aug 2 15:34:14 2025 -0700 agent 4.png commite32169dc37Author: Ishaan Jaff <ishaanjaffer0324@gmail.com> Date: Sat Aug 2 15:29:44 2025 -0700 docs cost tracking coding commit340b64a46aAuthor: Ishaan Jaff <ishaanjaffer0324@gmail.com> Date: Sat Aug 2 15:18:04 2025 -0700 docs - Track Usage for Coding Tools commit9b029c35beAuthor: Ishaan Jaff <ishaanjaffer0324@gmail.com> Date: Sat Aug 2 15:06:12 2025 -0700 docs RC commit8d6b333909Author: Ishaan Jaff <ishaanjaffer0324@gmail.com> Date: Sat Aug 2 15:00:48 2025 -0700 docs computer use commite306fb6eeeAuthor: Ishaan Jaff <ishaanjaffer0324@gmail.com> Date: Sat Aug 2 14:15:10 2025 -0700 [docs release notes] (#13237) * docs release notes * docs release notes * docs rnotes * docs api version * fixes docs * docs rn commit5dfc88473fAuthor: Ishaan Jaff <ishaanjaffer0324@gmail.com> Date: Sat Aug 2 13:30:51 2025 -0700 fixes MCP gateway docs commit6929767be2Author: Ishaan Jaff <ishaanjaffer0324@gmail.com> Date: Sat Aug 2 12:56:00 2025 -0700 docs release notes
This commit is contained in:
parent
51627ad543
commit
0d0a65a98f
26 changed files with 2164 additions and 324 deletions
|
|
@ -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");
|
||||
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 = {
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
@ -1,7 +1,6 @@
|
|||
from typing import Any, Optional, Union
|
||||
|
||||
from litellm.proxy._types import (
|
||||
GenerateKeyRequest,
|
||||
KeyRequestBase,
|
||||
LiteLLM_ManagementEndpoint_MetadataFields_Premium,
|
||||
LiteLLM_TeamTable,
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -2,19 +2,41 @@
|
|||
CRUD ENDPOINTS FOR PROMPTS
|
||||
"""
|
||||
|
||||
from typing import List, Optional, cast
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, List, Optional, cast
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
from fastapi import APIRouter, Depends, File, HTTPException, UploadFile
|
||||
from pydantic import BaseModel
|
||||
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._types import CommonProxyErrors, LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.types.prompts.init_prompts import ListPromptsResponse, PromptSpec
|
||||
from litellm.types.prompts.init_prompts import (
|
||||
ListPromptsResponse,
|
||||
PromptInfo,
|
||||
PromptInfoResponse,
|
||||
PromptLiteLLMParams,
|
||||
PromptSpec,
|
||||
PromptTemplateBase,
|
||||
)
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
class Prompt(BaseModel):
|
||||
prompt_id: str
|
||||
litellm_params: PromptLiteLLMParams
|
||||
prompt_info: Optional[PromptInfo] = None
|
||||
|
||||
|
||||
class PatchPromptRequest(BaseModel):
|
||||
litellm_params: Optional[PromptLiteLLMParams] = None
|
||||
prompt_info: Optional[PromptInfo] = None
|
||||
|
||||
|
||||
@router.get(
|
||||
"/prompt/list",
|
||||
"/prompts/list",
|
||||
tags=["Prompt Management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=ListPromptsResponse,
|
||||
|
|
@ -23,7 +45,35 @@ async def list_prompts(
|
|||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
List of available prompts for a given key.
|
||||
List the prompts that are available on the proxy server
|
||||
|
||||
👉 [Prompt docs](https://docs.litellm.ai/docs/proxy/prompt_management)
|
||||
|
||||
Example Request:
|
||||
```bash
|
||||
curl -X GET "http://localhost:4000/prompts/list" -H "Authorization: Bearer <your_api_key>"
|
||||
```
|
||||
|
||||
Example Response:
|
||||
```json
|
||||
{
|
||||
"prompts": [
|
||||
{
|
||||
"prompt_id": "my_prompt_id",
|
||||
"litellm_params": {
|
||||
"prompt_id": "my_prompt_id",
|
||||
"prompt_integration": "dotprompt",
|
||||
"prompt_directory": "/path/to/prompts"
|
||||
},
|
||||
"prompt_info": {
|
||||
"prompt_type": "config"
|
||||
},
|
||||
"created_at": "2023-11-09T12:34:56.789Z",
|
||||
"updated_at": "2023-11-09T12:34:56.789Z"
|
||||
}
|
||||
]
|
||||
}
|
||||
```
|
||||
"""
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY
|
||||
|
|
@ -53,17 +103,49 @@ async def list_prompts(
|
|||
|
||||
|
||||
@router.get(
|
||||
"/prompt/info",
|
||||
"/prompts/{prompt_id}",
|
||||
tags=["Prompt Management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=PromptSpec,
|
||||
response_model=PromptInfoResponse,
|
||||
)
|
||||
async def get_prompt(
|
||||
@router.get(
|
||||
"/prompts/{prompt_id}/info",
|
||||
tags=["Prompt Management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=PromptInfoResponse,
|
||||
)
|
||||
async def get_prompt_info(
|
||||
prompt_id: str,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Get info about a prompt
|
||||
Get detailed information about a specific prompt by ID, including prompt content
|
||||
|
||||
👉 [Prompt docs](https://docs.litellm.ai/docs/proxy/prompt_management)
|
||||
|
||||
Example Request:
|
||||
```bash
|
||||
curl -X GET "http://localhost:4000/prompts/my_prompt_id/info" \\
|
||||
-H "Authorization: Bearer <your_api_key>"
|
||||
```
|
||||
|
||||
Example Response:
|
||||
```json
|
||||
{
|
||||
"prompt_id": "my_prompt_id",
|
||||
"litellm_params": {
|
||||
"prompt_id": "my_prompt_id",
|
||||
"prompt_integration": "dotprompt",
|
||||
"prompt_directory": "/path/to/prompts"
|
||||
},
|
||||
"prompt_info": {
|
||||
"prompt_type": "config"
|
||||
},
|
||||
"created_at": "2023-11-09T12:34:56.789Z",
|
||||
"updated_at": "2023-11-09T12:34:56.789Z",
|
||||
"content": "System: You are a helpful assistant.\n\nUser: {{user_message}}"
|
||||
}
|
||||
```
|
||||
"""
|
||||
from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY
|
||||
|
||||
|
|
@ -89,4 +171,490 @@ async def get_prompt(
|
|||
prompt_spec = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(prompt_id)
|
||||
if prompt_spec is None:
|
||||
raise HTTPException(status_code=400, detail=f"Prompt {prompt_id} not found")
|
||||
return prompt_spec
|
||||
|
||||
# Get prompt content from the callback
|
||||
prompt_template: Optional[PromptTemplateBase] = None
|
||||
try:
|
||||
prompt_callback = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id(prompt_id)
|
||||
if prompt_callback is not None:
|
||||
# Extract content based on integration type
|
||||
integration_name = prompt_callback.integration_name
|
||||
|
||||
if integration_name == "dotprompt":
|
||||
# For dotprompt integration, get content from the prompt manager
|
||||
from litellm.integrations.dotprompt.dotprompt_manager import (
|
||||
DotpromptManager,
|
||||
)
|
||||
|
||||
if isinstance(prompt_callback, DotpromptManager):
|
||||
template = prompt_callback.prompt_manager.get_all_prompts_as_json()
|
||||
if template is not None and len(template) == 1:
|
||||
template_id = list(template.keys())[0]
|
||||
prompt_template = PromptTemplateBase(
|
||||
litellm_prompt_id=template_id, # id sent to prompt management tool
|
||||
content=template[template_id]["content"],
|
||||
metadata=template[template_id]["metadata"],
|
||||
)
|
||||
|
||||
except Exception:
|
||||
# If content extraction fails, continue without content
|
||||
pass
|
||||
|
||||
# Create response with content
|
||||
return PromptInfoResponse(
|
||||
prompt_spec=prompt_spec,
|
||||
raw_prompt_template=prompt_template,
|
||||
)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/prompts",
|
||||
tags=["Prompt Management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
)
|
||||
async def create_prompt(
|
||||
request: Prompt,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Create a new prompt
|
||||
|
||||
👉 [Prompt docs](https://docs.litellm.ai/docs/proxy/prompt_management)
|
||||
|
||||
Example Request:
|
||||
```bash
|
||||
curl -X POST "http://localhost:4000/prompts" \\
|
||||
-H "Authorization: Bearer <your_api_key>" \\
|
||||
-H "Content-Type: application/json" \\
|
||||
-d '{
|
||||
"prompt_id": "my_prompt",
|
||||
"litellm_params": {
|
||||
"prompt_id": "json_prompt",
|
||||
"prompt_integration": "dotprompt",
|
||||
### EITHER prompt_directory OR prompt_data MUST BE PROVIDED
|
||||
"prompt_directory": "/path/to/dotprompt/folder",
|
||||
"prompt_data": {"json_prompt": {"content": "This is a prompt", "metadata": {"model": "gpt-4"}}}
|
||||
},
|
||||
"prompt_info": {
|
||||
"prompt_type": "config"
|
||||
}
|
||||
}'
|
||||
```
|
||||
"""
|
||||
|
||||
from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
# Only allow proxy admins to create prompts
|
||||
if user_api_key_dict.user_role is None or (
|
||||
user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN
|
||||
and user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=403, detail="Only proxy admins can create prompts"
|
||||
)
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(
|
||||
status_code=500, detail=CommonProxyErrors.db_not_connected_error.value
|
||||
)
|
||||
|
||||
try:
|
||||
# Create the prompt spec
|
||||
# Check if prompt exists and get current data
|
||||
existing_prompt = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(request.prompt_id)
|
||||
if existing_prompt is not None:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail=f"Prompt with ID {request.prompt_id} already exists",
|
||||
)
|
||||
|
||||
# store prompt in db
|
||||
prompt_db_entry = await prisma_client.db.litellm_prompttable.create(
|
||||
data={
|
||||
"prompt_id": request.prompt_id,
|
||||
"litellm_params": request.litellm_params.model_dump_json(),
|
||||
"prompt_info": (
|
||||
request.prompt_info.model_dump_json()
|
||||
if request.prompt_info
|
||||
else PromptInfo(prompt_type="db").model_dump_json()
|
||||
),
|
||||
}
|
||||
)
|
||||
|
||||
prompt_spec = PromptSpec(**prompt_db_entry.model_dump())
|
||||
|
||||
# Initialize the prompt
|
||||
initialized_prompt = IN_MEMORY_PROMPT_REGISTRY.initialize_prompt(
|
||||
prompt=prompt_spec, config_file_path=None
|
||||
)
|
||||
|
||||
if initialized_prompt is None:
|
||||
raise HTTPException(status_code=500, detail="Failed to initialize prompt")
|
||||
|
||||
return initialized_prompt
|
||||
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(f"Error creating prompt: {e}")
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
@router.put(
|
||||
"/prompts/{prompt_id}",
|
||||
tags=["Prompt Management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
)
|
||||
async def update_prompt(
|
||||
prompt_id: str,
|
||||
request: Prompt,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Update an existing prompt
|
||||
|
||||
👉 [Prompt docs](https://docs.litellm.ai/docs/proxy/prompt_management)
|
||||
|
||||
Example Request:
|
||||
```bash
|
||||
curl -X PUT "http://localhost:4000/prompts/my_prompt_id" \\
|
||||
-H "Authorization: Bearer <your_api_key>" \\
|
||||
-H "Content-Type: application/json" \\
|
||||
-d '{
|
||||
"prompt_id": "my_prompt",
|
||||
"litellm_params": {
|
||||
"prompt_id": "my_prompt",
|
||||
"prompt_integration": "dotprompt",
|
||||
"prompt_directory": "/path/to/prompts"
|
||||
},
|
||||
"prompt_info": {
|
||||
"prompt_type": "config"
|
||||
}
|
||||
}
|
||||
}'
|
||||
```
|
||||
"""
|
||||
from datetime import datetime
|
||||
|
||||
from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
# Only allow proxy admins to update prompts
|
||||
if user_api_key_dict.user_role is None or (
|
||||
user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN
|
||||
and user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=403, detail="Only proxy admins can update prompts"
|
||||
)
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(
|
||||
status_code=500, detail=CommonProxyErrors.db_not_connected_error.value
|
||||
)
|
||||
|
||||
try:
|
||||
# Check if prompt exists
|
||||
existing_prompt = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(prompt_id)
|
||||
if existing_prompt is None:
|
||||
raise HTTPException(
|
||||
status_code=404, detail=f"Prompt with ID {prompt_id} not found"
|
||||
)
|
||||
|
||||
if existing_prompt.prompt_info.prompt_type == "config":
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="Cannot update config prompts.",
|
||||
)
|
||||
|
||||
# Create updated prompt spec
|
||||
updated_prompt_spec = PromptSpec(
|
||||
prompt_id=prompt_id,
|
||||
litellm_params=request.litellm_params,
|
||||
prompt_info=request.prompt_info or PromptInfo(prompt_type="db"),
|
||||
created_at=existing_prompt.created_at,
|
||||
updated_at=datetime.now(),
|
||||
)
|
||||
|
||||
updated_prompt_db_entry = await prisma_client.db.litellm_prompttable.update(
|
||||
where={"prompt_id": prompt_id},
|
||||
data={
|
||||
"litellm_params": updated_prompt_spec.litellm_params.model_dump_json(),
|
||||
"prompt_info": updated_prompt_spec.prompt_info.model_dump_json(),
|
||||
},
|
||||
)
|
||||
|
||||
# Remove the old prompt from memory
|
||||
del IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS[prompt_id]
|
||||
if prompt_id in IN_MEMORY_PROMPT_REGISTRY.prompt_id_to_custom_prompt:
|
||||
del IN_MEMORY_PROMPT_REGISTRY.prompt_id_to_custom_prompt[prompt_id]
|
||||
|
||||
# Initialize the updated prompt
|
||||
initialized_prompt = IN_MEMORY_PROMPT_REGISTRY.initialize_prompt(
|
||||
prompt=PromptSpec(**updated_prompt_db_entry.model_dump()),
|
||||
config_file_path=None,
|
||||
)
|
||||
|
||||
if initialized_prompt is None:
|
||||
raise HTTPException(status_code=500, detail="Failed to update prompt")
|
||||
|
||||
return initialized_prompt
|
||||
|
||||
except HTTPException as e:
|
||||
raise e
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(f"Error updating prompt: {e}")
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
@router.delete(
|
||||
"/prompts/{prompt_id}",
|
||||
tags=["Prompt Management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
)
|
||||
async def delete_prompt(
|
||||
prompt_id: str,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Delete a prompt
|
||||
|
||||
👉 [Prompt docs](https://docs.litellm.ai/docs/proxy/prompt_management)
|
||||
|
||||
Example Request:
|
||||
```bash
|
||||
curl -X DELETE "http://localhost:4000/prompts/my_prompt_id" \\
|
||||
-H "Authorization: Bearer <your_api_key>"
|
||||
```
|
||||
|
||||
Example Response:
|
||||
```json
|
||||
{
|
||||
"message": "Prompt my_prompt_id deleted successfully"
|
||||
}
|
||||
```
|
||||
"""
|
||||
from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
# Only allow proxy admins to delete prompts
|
||||
if user_api_key_dict.user_role is None or (
|
||||
user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN
|
||||
and user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=403, detail="Only proxy admins can delete prompts"
|
||||
)
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(
|
||||
status_code=500, detail=CommonProxyErrors.db_not_connected_error.value
|
||||
)
|
||||
|
||||
try:
|
||||
# Check if prompt exists
|
||||
existing_prompt = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(prompt_id)
|
||||
if existing_prompt is None:
|
||||
raise HTTPException(
|
||||
status_code=404, detail=f"Prompt with ID {prompt_id} not found"
|
||||
)
|
||||
|
||||
if existing_prompt.prompt_info.prompt_type == "config":
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="Cannot delete config prompts.",
|
||||
)
|
||||
|
||||
# Delete the prompt from the database
|
||||
await prisma_client.db.litellm_prompttable.delete(
|
||||
where={"prompt_id": prompt_id}
|
||||
)
|
||||
|
||||
# Remove the prompt from memory
|
||||
del IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS[prompt_id]
|
||||
if prompt_id in IN_MEMORY_PROMPT_REGISTRY.prompt_id_to_custom_prompt:
|
||||
del IN_MEMORY_PROMPT_REGISTRY.prompt_id_to_custom_prompt[prompt_id]
|
||||
|
||||
return {"message": f"Prompt {prompt_id} deleted successfully"}
|
||||
|
||||
except HTTPException as e:
|
||||
raise e
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(f"Error deleting prompt: {e}")
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
@router.patch(
|
||||
"/prompts/{prompt_id}",
|
||||
tags=["Prompt Management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
)
|
||||
async def patch_prompt(
|
||||
prompt_id: str,
|
||||
request: PatchPromptRequest,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Partially update an existing prompt
|
||||
|
||||
👉 [Prompt docs](https://docs.litellm.ai/docs/proxy/prompt_management)
|
||||
|
||||
This endpoint allows updating specific fields of a prompt without sending the entire object.
|
||||
Only the following fields can be updated:
|
||||
- litellm_params: LiteLLM parameters for the prompt
|
||||
- prompt_info: Additional information about the prompt
|
||||
|
||||
Example Request:
|
||||
```bash
|
||||
curl -X PATCH "http://localhost:4000/prompts/my_prompt_id" \\
|
||||
-H "Authorization: Bearer <your_api_key>" \\
|
||||
-H "Content-Type: application/json" \\
|
||||
-d '{
|
||||
"prompt_info": {
|
||||
"prompt_type": "db"
|
||||
}
|
||||
}'
|
||||
```
|
||||
"""
|
||||
|
||||
from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
# Only allow proxy admins to patch prompts
|
||||
if user_api_key_dict.user_role is None or (
|
||||
user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN
|
||||
and user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=403, detail="Only proxy admins can patch prompts"
|
||||
)
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(
|
||||
status_code=500, detail=CommonProxyErrors.db_not_connected_error.value
|
||||
)
|
||||
|
||||
try:
|
||||
# Check if prompt exists and get current data
|
||||
existing_prompt = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(prompt_id)
|
||||
if existing_prompt is None:
|
||||
raise HTTPException(
|
||||
status_code=404, detail=f"Prompt with ID {prompt_id} not found"
|
||||
)
|
||||
|
||||
if existing_prompt.prompt_info.prompt_type == "config":
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="Cannot update config prompts.",
|
||||
)
|
||||
|
||||
# Update fields if provided
|
||||
updated_litellm_params = (
|
||||
request.litellm_params
|
||||
if request.litellm_params is not None
|
||||
else existing_prompt.litellm_params
|
||||
)
|
||||
|
||||
updated_prompt_info = (
|
||||
request.prompt_info
|
||||
if request.prompt_info is not None
|
||||
else existing_prompt.prompt_info
|
||||
)
|
||||
|
||||
# Ensure we have valid litellm_params
|
||||
if updated_litellm_params is None:
|
||||
raise HTTPException(status_code=400, detail="litellm_params cannot be None")
|
||||
|
||||
# Create updated prompt spec - cast to satisfy typing
|
||||
updated_prompt_db_entry = await prisma_client.db.litellm_prompttable.update(
|
||||
where={"prompt_id": prompt_id},
|
||||
data={
|
||||
"litellm_params": updated_litellm_params.model_dump_json(),
|
||||
"prompt_info": updated_prompt_info.model_dump_json(),
|
||||
},
|
||||
)
|
||||
|
||||
updated_prompt_spec = PromptSpec(**updated_prompt_db_entry.model_dump())
|
||||
|
||||
# Remove the old prompt from memory
|
||||
del IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS[prompt_id]
|
||||
if prompt_id in IN_MEMORY_PROMPT_REGISTRY.prompt_id_to_custom_prompt:
|
||||
del IN_MEMORY_PROMPT_REGISTRY.prompt_id_to_custom_prompt[prompt_id]
|
||||
|
||||
# Initialize the updated prompt
|
||||
initialized_prompt = IN_MEMORY_PROMPT_REGISTRY.initialize_prompt(
|
||||
prompt=updated_prompt_spec, config_file_path=None
|
||||
)
|
||||
|
||||
if initialized_prompt is None:
|
||||
raise HTTPException(status_code=500, detail="Failed to patch prompt")
|
||||
|
||||
return initialized_prompt
|
||||
|
||||
except HTTPException as e:
|
||||
raise e
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(f"Error patching prompt: {e}")
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
@router.post(
|
||||
"/utils/dotprompt_json_converter",
|
||||
tags=["prompts", "utils"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
)
|
||||
async def convert_prompt_file_to_json(
|
||||
file: UploadFile = File(...),
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Convert a .prompt file to JSON format.
|
||||
|
||||
This endpoint accepts a .prompt file upload and returns the equivalent JSON representation
|
||||
that can be stored in a database or used programmatically.
|
||||
|
||||
Returns the JSON structure with 'content' and 'metadata' fields.
|
||||
"""
|
||||
global general_settings
|
||||
from litellm.integrations.dotprompt.prompt_manager import PromptManager
|
||||
|
||||
# Validate file extension
|
||||
if not file.filename or not file.filename.endswith(".prompt"):
|
||||
raise HTTPException(status_code=400, detail="File must have .prompt extension")
|
||||
|
||||
temp_file_path = None
|
||||
try:
|
||||
# Read file content
|
||||
file_content = await file.read()
|
||||
|
||||
# Create temporary file
|
||||
temp_file_path = Path(tempfile.mkdtemp()) / file.filename
|
||||
temp_file_path.write_bytes(file_content)
|
||||
|
||||
# Create a PromptManager instance just for conversion
|
||||
prompt_manager = PromptManager()
|
||||
|
||||
# Convert to JSON
|
||||
json_data = prompt_manager.prompt_file_to_json(temp_file_path)
|
||||
|
||||
# Extract prompt ID from filename
|
||||
prompt_id = temp_file_path.stem
|
||||
|
||||
return {
|
||||
"prompt_id": prompt_id,
|
||||
"json_data": json_data,
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
raise HTTPException(
|
||||
status_code=500, detail=f"Error converting prompt file: {str(e)}"
|
||||
)
|
||||
|
||||
finally:
|
||||
# Clean up temp file
|
||||
if temp_file_path and temp_file_path.exists():
|
||||
temp_file_path.unlink()
|
||||
# Also try to remove the temp directory if it's empty
|
||||
try:
|
||||
temp_file_path.parent.rmdir()
|
||||
except OSError:
|
||||
pass # Directory not empty or other error
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -55,7 +55,11 @@ from litellm import (
|
|||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._service_logger import ServiceLogging, ServiceTypes
|
||||
from litellm.caching.caching import DualCache, RedisCache
|
||||
from litellm.exceptions import RejectedRequestError, BlockedPiiEntityError, GuardrailRaisedException
|
||||
from litellm.exceptions import (
|
||||
BlockedPiiEntityError,
|
||||
GuardrailRaisedException,
|
||||
RejectedRequestError,
|
||||
)
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.integrations.SlackAlerting.slack_alerting import SlackAlerting
|
||||
|
|
@ -472,8 +476,10 @@ class ProxyLogging:
|
|||
# Create the request object if it's not already one
|
||||
if not isinstance(request_obj, MCPPreCallRequestObject):
|
||||
# Convert UserAPIKeyAuth object to dict if needed
|
||||
user_api_key_auth_dict = self._convert_user_api_key_auth_to_dict(kwargs.get("user_api_key_auth"))
|
||||
|
||||
user_api_key_auth_dict = self._convert_user_api_key_auth_to_dict(
|
||||
kwargs.get("user_api_key_auth")
|
||||
)
|
||||
|
||||
request_obj = MCPPreCallRequestObject(
|
||||
tool_name=kwargs.get("name", ""),
|
||||
arguments=kwargs.get("arguments", {}),
|
||||
|
|
@ -487,7 +493,9 @@ class ProxyLogging:
|
|||
_callback: Optional[CustomLogger] = None
|
||||
if isinstance(callback, str):
|
||||
from typing import cast
|
||||
|
||||
from litellm import _custom_logger_compatible_callbacks_literal
|
||||
|
||||
_callback = litellm.litellm_core_utils.litellm_logging.get_custom_logger_compatible_class(
|
||||
cast(_custom_logger_compatible_callbacks_literal, callback)
|
||||
)
|
||||
|
|
@ -507,26 +515,36 @@ class ProxyLogging:
|
|||
continue
|
||||
|
||||
# Convert MCP tool call to LLM message format for existing guardrail logic
|
||||
synthetic_llm_data = self._convert_mcp_to_llm_format(request_obj, kwargs)
|
||||
synthetic_llm_data = self._convert_mcp_to_llm_format(
|
||||
request_obj, kwargs
|
||||
)
|
||||
# Reuse existing LLM guardrail logic
|
||||
user_api_key_auth_dict = self._convert_user_api_key_auth_to_dict(kwargs.get("user_api_key_auth"))
|
||||
user_api_key_auth_dict = self._convert_user_api_key_auth_to_dict(
|
||||
kwargs.get("user_api_key_auth")
|
||||
)
|
||||
|
||||
result = await _callback.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_auth_dict,
|
||||
user_api_key_dict=user_api_key_auth_dict, # type: ignore
|
||||
cache=self.call_details["user_api_key_cache"],
|
||||
data=synthetic_llm_data,
|
||||
call_type="mcp_call"
|
||||
call_type="mcp_call",
|
||||
)
|
||||
|
||||
# Convert result back to MCP response format if blocked/modified
|
||||
if result is not None:
|
||||
mcp_response = self._convert_llm_result_to_mcp_response(result, request_obj)
|
||||
mcp_response = self._convert_llm_result_to_mcp_response(
|
||||
result, request_obj
|
||||
)
|
||||
if mcp_response is not None:
|
||||
return self._parse_pre_mcp_call_hook_response(
|
||||
response=mcp_response, original_request=request_obj
|
||||
)
|
||||
|
||||
except (BlockedPiiEntityError, GuardrailRaisedException, HTTPException) as e:
|
||||
|
||||
except (
|
||||
BlockedPiiEntityError,
|
||||
GuardrailRaisedException,
|
||||
HTTPException,
|
||||
) as e:
|
||||
# Re-raise guardrail exceptions so they can be properly handled
|
||||
raise e
|
||||
except Exception as e:
|
||||
|
|
@ -543,10 +561,10 @@ class ProxyLogging:
|
|||
Handles both Pydantic models and regular objects.
|
||||
"""
|
||||
if user_api_key_auth_obj is not None:
|
||||
if hasattr(user_api_key_auth_obj, 'model_dump'):
|
||||
if hasattr(user_api_key_auth_obj, "model_dump"):
|
||||
# If it's a Pydantic model, convert to dict
|
||||
return user_api_key_auth_obj.model_dump()
|
||||
elif hasattr(user_api_key_auth_obj, '__dict__'):
|
||||
elif hasattr(user_api_key_auth_obj, "__dict__"):
|
||||
# If it's a regular object, convert to dict
|
||||
return user_api_key_auth_obj.__dict__
|
||||
return user_api_key_auth_obj
|
||||
|
|
@ -556,15 +574,16 @@ class ProxyLogging:
|
|||
Convert MCP tool call to LLM message format for existing guardrail validation.
|
||||
"""
|
||||
from litellm.types.llms.openai import ChatCompletionUserMessage
|
||||
|
||||
|
||||
# Create a synthetic message that represents the tool call
|
||||
tool_call_content = f"Tool: {request_obj.tool_name}\nArguments: {request_obj.arguments}"
|
||||
|
||||
synthetic_message = ChatCompletionUserMessage(
|
||||
role="user",
|
||||
content=tool_call_content
|
||||
tool_call_content = (
|
||||
f"Tool: {request_obj.tool_name}\nArguments: {request_obj.arguments}"
|
||||
)
|
||||
|
||||
|
||||
synthetic_message = ChatCompletionUserMessage(
|
||||
role="user", content=tool_call_content
|
||||
)
|
||||
|
||||
# Create synthetic LLM data that guardrails can process
|
||||
synthetic_data = {
|
||||
"messages": [synthetic_message],
|
||||
|
|
@ -577,51 +596,65 @@ class ProxyLogging:
|
|||
"mcp_tool_name": request_obj.tool_name, # Keep original for reference
|
||||
"mcp_arguments": request_obj.arguments, # Keep original for reference
|
||||
}
|
||||
|
||||
|
||||
return synthetic_data
|
||||
|
||||
def _convert_llm_result_to_mcp_response(self, llm_result, request_obj) -> Optional[Any]:
|
||||
def _convert_llm_result_to_mcp_response(
|
||||
self, llm_result, request_obj
|
||||
) -> Optional[Any]:
|
||||
"""
|
||||
Convert LLM guardrail result back to MCP response format.
|
||||
"""
|
||||
from litellm.types.mcp import MCPPreCallResponseObject
|
||||
|
||||
|
||||
# If result is an exception, it means the guardrail blocked the request
|
||||
if isinstance(llm_result, Exception):
|
||||
return MCPPreCallResponseObject(
|
||||
should_proceed=False,
|
||||
error_message=str(llm_result),
|
||||
modified_arguments=None
|
||||
modified_arguments=None,
|
||||
)
|
||||
|
||||
|
||||
# If result is a dict with modified messages, check for content filtering
|
||||
if isinstance(llm_result, dict):
|
||||
modified_messages = llm_result.get("messages")
|
||||
if modified_messages:
|
||||
# Check if content was blocked/modified
|
||||
original_content = f"Tool: {request_obj.tool_name}\nArguments: {request_obj.arguments}"
|
||||
new_content = modified_messages[0].get("content", "") if modified_messages else ""
|
||||
|
||||
original_content = (
|
||||
f"Tool: {request_obj.tool_name}\nArguments: {request_obj.arguments}"
|
||||
)
|
||||
new_content = (
|
||||
modified_messages[0].get("content", "") if modified_messages else ""
|
||||
)
|
||||
|
||||
if new_content != original_content:
|
||||
# Content was modified - could be masking, redaction, or blocking
|
||||
if not new_content or "blocked" in new_content.lower() or "violation" in new_content.lower():
|
||||
if (
|
||||
not new_content
|
||||
or "blocked" in new_content.lower()
|
||||
or "violation" in new_content.lower()
|
||||
):
|
||||
# Content was blocked completely
|
||||
return MCPPreCallResponseObject(
|
||||
should_proceed=False,
|
||||
error_message="Content blocked by guardrail",
|
||||
modified_arguments=None
|
||||
modified_arguments=None,
|
||||
)
|
||||
else:
|
||||
# Content was masked/redacted - extract the modified arguments
|
||||
try:
|
||||
# Try to parse the modified arguments from the masked content
|
||||
modified_args = self._extract_modified_arguments_from_content(new_content, request_obj)
|
||||
modified_args = (
|
||||
self._extract_modified_arguments_from_content(
|
||||
new_content, request_obj
|
||||
)
|
||||
)
|
||||
if modified_args is not None:
|
||||
# Return the masked/redacted arguments for the MCP call to use
|
||||
return MCPPreCallResponseObject(
|
||||
should_proceed=True,
|
||||
error_message=None,
|
||||
modified_arguments=modified_args
|
||||
modified_arguments=modified_args,
|
||||
)
|
||||
else:
|
||||
# Could not parse modified arguments, allow original call but warn
|
||||
|
|
@ -630,129 +663,149 @@ class ProxyLogging:
|
|||
)
|
||||
return None
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error(f"Error parsing modified arguments: {e}")
|
||||
verbose_proxy_logger.error(
|
||||
f"Error parsing modified arguments: {e}"
|
||||
)
|
||||
# Fallback: allow original call
|
||||
return None
|
||||
|
||||
|
||||
# If result is a string, it's likely an error message
|
||||
if isinstance(llm_result, str):
|
||||
return MCPPreCallResponseObject(
|
||||
should_proceed=False,
|
||||
error_message=llm_result,
|
||||
modified_arguments=None
|
||||
should_proceed=False, error_message=llm_result, modified_arguments=None
|
||||
)
|
||||
|
||||
|
||||
return None
|
||||
|
||||
def _extract_modified_arguments_from_content(self, masked_content: str, request_obj) -> Optional[dict]:
|
||||
def _extract_modified_arguments_from_content(
|
||||
self, masked_content: str, request_obj
|
||||
) -> Optional[dict]:
|
||||
"""
|
||||
Extract modified/masked arguments from the guardrail response content.
|
||||
"""
|
||||
import json
|
||||
|
||||
verbose_proxy_logger.debug(f"Extracting modified args from content: {masked_content}")
|
||||
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
f"Extracting modified args from content: {masked_content}"
|
||||
)
|
||||
|
||||
try:
|
||||
# The format should be: "Tool: <tool_name>\nArguments: <json_arguments>"
|
||||
# Parse the arguments section
|
||||
lines = masked_content.strip().split('\n')
|
||||
lines = masked_content.strip().split("\n")
|
||||
for i, line in enumerate(lines):
|
||||
if line.startswith("Arguments:"):
|
||||
# Get the arguments part - everything after "Arguments: "
|
||||
args_text = line[len("Arguments:"):].strip()
|
||||
|
||||
args_text = line[len("Arguments:") :].strip()
|
||||
|
||||
verbose_proxy_logger.debug(f"Found arguments text: {args_text}")
|
||||
|
||||
|
||||
# Try to parse as JSON first
|
||||
try:
|
||||
modified_args = json.loads(args_text)
|
||||
verbose_proxy_logger.debug(f"Successfully parsed JSON args: {modified_args}")
|
||||
verbose_proxy_logger.debug(
|
||||
f"Successfully parsed JSON args: {modified_args}"
|
||||
)
|
||||
return modified_args
|
||||
except json.JSONDecodeError as e:
|
||||
# If JSON parsing fails, try to extract key-value pairs manually
|
||||
verbose_proxy_logger.debug(f"Failed to parse JSON arguments: {args_text}, error: {e}")
|
||||
return self._parse_arguments_manually(args_text, request_obj.arguments)
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
f"Failed to parse JSON arguments: {args_text}, error: {e}"
|
||||
)
|
||||
return self._parse_arguments_manually(
|
||||
args_text, request_obj.arguments
|
||||
)
|
||||
|
||||
# If we can't find the Arguments: line, return None
|
||||
verbose_proxy_logger.warning("Could not find 'Arguments:' line in masked content")
|
||||
verbose_proxy_logger.warning(
|
||||
"Could not find 'Arguments:' line in masked content"
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error(f"Error extracting modified arguments: {e}")
|
||||
return None
|
||||
|
||||
def _parse_arguments_manually(self, args_text: str, original_args: dict) -> Optional[dict]:
|
||||
def _parse_arguments_manually(
|
||||
self, args_text: str, original_args: dict
|
||||
) -> Optional[dict]:
|
||||
"""
|
||||
Try to manually parse arguments when JSON parsing fails.
|
||||
This is a fallback for cases where the guardrail modifies the format.
|
||||
"""
|
||||
import re
|
||||
|
||||
|
||||
try:
|
||||
# Start with original arguments and try to apply modifications
|
||||
modified_args = original_args.copy()
|
||||
|
||||
|
||||
# Look for simple key-value patterns
|
||||
# This is a basic implementation - can be enhanced based on specific guardrail formats
|
||||
for key, original_value in original_args.items():
|
||||
if isinstance(original_value, str):
|
||||
# Look for the key in the masked content and try to extract its value
|
||||
pattern = rf"['\"]?{re.escape(key)}['\"]?\s*:\s*['\"]?([^,'\"]*)['\"]?"
|
||||
pattern = (
|
||||
rf"['\"]?{re.escape(key)}['\"]?\s*:\s*['\"]?([^,'\"]*)['\"]?"
|
||||
)
|
||||
match = re.search(pattern, args_text, re.IGNORECASE)
|
||||
if match:
|
||||
new_value = match.group(1).strip()
|
||||
if new_value:
|
||||
modified_args[key] = new_value
|
||||
|
||||
|
||||
return modified_args
|
||||
|
||||
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error(f"Error in manual argument parsing: {e}")
|
||||
return None
|
||||
|
||||
def _convert_llm_result_to_mcp_during_response(self, llm_result, request_obj) -> Optional[Any]:
|
||||
def _convert_llm_result_to_mcp_during_response(
|
||||
self, llm_result, request_obj
|
||||
) -> Optional[Any]:
|
||||
"""
|
||||
Convert LLM guardrail result back to MCP during call response format.
|
||||
"""
|
||||
from litellm.types.mcp import MCPDuringCallResponseObject
|
||||
|
||||
|
||||
# If result is an exception, it means the guardrail wants to stop execution
|
||||
if isinstance(llm_result, Exception):
|
||||
return MCPDuringCallResponseObject(
|
||||
should_continue=False,
|
||||
error_message=str(llm_result)
|
||||
should_continue=False, error_message=str(llm_result)
|
||||
)
|
||||
|
||||
|
||||
# If result is a dict with modified messages, check for content filtering
|
||||
if isinstance(llm_result, dict):
|
||||
modified_messages = llm_result.get("messages")
|
||||
if modified_messages:
|
||||
# Check if content was blocked/modified
|
||||
original_content = f"Tool: {request_obj.tool_name}\nArguments: {request_obj.arguments}"
|
||||
new_content = modified_messages[0].get("content", "") if modified_messages else ""
|
||||
|
||||
original_content = (
|
||||
f"Tool: {request_obj.tool_name}\nArguments: {request_obj.arguments}"
|
||||
)
|
||||
new_content = (
|
||||
modified_messages[0].get("content", "") if modified_messages else ""
|
||||
)
|
||||
|
||||
if new_content != original_content:
|
||||
# Content was modified, could be masking or blocking
|
||||
if not new_content or "blocked" in new_content.lower():
|
||||
# Content was blocked
|
||||
return MCPDuringCallResponseObject(
|
||||
should_continue=False,
|
||||
error_message="Content blocked by guardrail during execution"
|
||||
error_message="Content blocked by guardrail during execution",
|
||||
)
|
||||
else:
|
||||
# Content was masked/modified - for now, stop execution
|
||||
# Content was masked/modified - for now, stop execution
|
||||
return MCPDuringCallResponseObject(
|
||||
should_continue=False,
|
||||
error_message="Content modified by guardrail during execution"
|
||||
error_message="Content modified by guardrail during execution",
|
||||
)
|
||||
|
||||
|
||||
# If result is a string, it's likely an error message
|
||||
if isinstance(llm_result, str):
|
||||
return MCPDuringCallResponseObject(
|
||||
should_continue=False,
|
||||
error_message=llm_result
|
||||
should_continue=False, error_message=llm_result
|
||||
)
|
||||
|
||||
|
||||
return None
|
||||
|
||||
def get_combined_callback_list(
|
||||
|
|
@ -799,7 +852,6 @@ class ProxyLogging:
|
|||
from litellm.types.llms.base import HiddenParams
|
||||
from litellm.types.mcp import MCPDuringCallRequestObject
|
||||
|
||||
|
||||
callbacks = self.get_combined_callback_list(
|
||||
dynamic_success_callbacks=getattr(self, "dynamic_success_callbacks", None),
|
||||
global_callbacks=litellm.success_callback,
|
||||
|
|
@ -820,7 +872,9 @@ class ProxyLogging:
|
|||
_callback: Optional[CustomLogger] = None
|
||||
if isinstance(callback, str):
|
||||
from typing import cast
|
||||
|
||||
from litellm import _custom_logger_compatible_callbacks_literal
|
||||
|
||||
_callback = litellm.litellm_core_utils.litellm_logging.get_custom_logger_compatible_class(
|
||||
cast(_custom_logger_compatible_callbacks_literal, callback)
|
||||
)
|
||||
|
|
@ -839,22 +893,29 @@ class ProxyLogging:
|
|||
):
|
||||
continue
|
||||
# Convert MCP tool call to LLM message format for existing guardrail logic
|
||||
synthetic_llm_data = self._convert_mcp_to_llm_format(request_obj, kwargs)
|
||||
|
||||
synthetic_llm_data = self._convert_mcp_to_llm_format(
|
||||
request_obj, kwargs
|
||||
)
|
||||
|
||||
# Reuse existing LLM guardrail logic for during call
|
||||
user_api_key_auth_dict = self._convert_user_api_key_auth_to_dict(kwargs.get("user_api_key_auth"))
|
||||
user_api_key_auth_dict = self._convert_user_api_key_auth_to_dict(
|
||||
kwargs.get("user_api_key_auth")
|
||||
)
|
||||
|
||||
result = await _callback.async_moderation_hook(
|
||||
data=synthetic_llm_data,
|
||||
user_api_key_dict=user_api_key_auth_dict,
|
||||
call_type="mcp_call"
|
||||
)
|
||||
user_api_key_dict=user_api_key_auth_dict, # type: ignore
|
||||
call_type="mcp_call",
|
||||
)
|
||||
# Convert result back to MCP response format if blocked/modified
|
||||
if result is not None:
|
||||
mcp_response = self._convert_llm_result_to_mcp_during_response(result, request_obj)
|
||||
mcp_response = self._convert_llm_result_to_mcp_during_response(
|
||||
result, request_obj
|
||||
)
|
||||
if mcp_response is not None:
|
||||
return self._parse_during_mcp_call_hook_response(response=mcp_response)
|
||||
|
||||
return self._parse_during_mcp_call_hook_response(
|
||||
response=mcp_response
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
raise e
|
||||
|
|
@ -984,8 +1045,12 @@ class ProxyLogging:
|
|||
custom_logger = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id(
|
||||
prompt_id
|
||||
)
|
||||
prompt_spec = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(prompt_id)
|
||||
litellm_prompt_id: Optional[str] = None
|
||||
if prompt_spec is not None:
|
||||
litellm_prompt_id = prompt_spec.litellm_params.prompt_id
|
||||
|
||||
if custom_logger:
|
||||
if custom_logger and litellm_prompt_id is not None:
|
||||
(
|
||||
model,
|
||||
messages,
|
||||
|
|
@ -994,15 +1059,16 @@ class ProxyLogging:
|
|||
model=data.get("model", ""),
|
||||
messages=data.get("messages", []),
|
||||
non_default_params=get_non_default_completion_params(kwargs=data),
|
||||
prompt_id=prompt_id,
|
||||
prompt_id=litellm_prompt_id,
|
||||
prompt_management_logger=custom_logger,
|
||||
prompt_variables=data.get("prompt_variables", None),
|
||||
prompt_label=data.get("prompt_label", None),
|
||||
prompt_version=data.get("prompt_version", None),
|
||||
)
|
||||
|
||||
data.update(optional_params)
|
||||
data["model"] = model
|
||||
data["messages"] = messages
|
||||
data.update(optional_params)
|
||||
|
||||
try:
|
||||
for callback in litellm.callbacks:
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -87,6 +87,17 @@ export interface PromptSpec {
|
|||
updated_at?: string
|
||||
}
|
||||
|
||||
export interface PromptTemplateBase {
|
||||
litellm_prompt_id: string;
|
||||
content: string;
|
||||
metadata?: Record<string, any> | null;
|
||||
}
|
||||
|
||||
export interface PromptInfoResponse {
|
||||
prompt_spec: PromptSpec;
|
||||
raw_prompt_template: PromptTemplateBase | null;
|
||||
}
|
||||
|
||||
export interface ListPromptsResponse {
|
||||
prompts: PromptSpec[];
|
||||
}
|
||||
|
|
@ -4899,8 +4910,8 @@ export const getGuardrailsList = async (accessToken: String) => {
|
|||
export const getPromptsList = async (accessToken: String) : Promise<ListPromptsResponse> => {
|
||||
try {
|
||||
const url = proxyBaseUrl
|
||||
? `${proxyBaseUrl}/prompt/list`
|
||||
: `/prompt/list`;
|
||||
? `${proxyBaseUrl}/prompts/list`
|
||||
: `/prompts/list`;
|
||||
const response = await fetch(url, {
|
||||
method: "GET",
|
||||
headers: {
|
||||
|
|
@ -4923,9 +4934,9 @@ export const getPromptsList = async (accessToken: String) : Promise<ListPromptsR
|
|||
}
|
||||
};
|
||||
|
||||
export const getPromptInfo = async (accessToken: String, promptId: string): Promise<PromptSpec> => {
|
||||
export const getPromptInfo = async (accessToken: String, promptId: string): Promise<PromptInfoResponse> => {
|
||||
try {
|
||||
const url = proxyBaseUrl ? `${proxyBaseUrl}/prompt/info?prompt_id=${promptId}` : `/prompt/info?prompt_id=${promptId}`;
|
||||
const url = proxyBaseUrl ? `${proxyBaseUrl}/prompts/${promptId}/info` : `/prompts/${promptId}/info`;
|
||||
const response = await fetch(url, {
|
||||
method: "GET",
|
||||
headers: {
|
||||
|
|
@ -4948,6 +4959,160 @@ export const getPromptInfo = async (accessToken: String, promptId: string): Prom
|
|||
}
|
||||
};
|
||||
|
||||
export const createPromptCall = async (
|
||||
accessToken: string,
|
||||
promptData: any
|
||||
) => {
|
||||
try {
|
||||
const url = proxyBaseUrl ? `${proxyBaseUrl}/prompts` : `/prompts`;
|
||||
|
||||
const response = await fetch(url, {
|
||||
method: "POST",
|
||||
headers: {
|
||||
[globalLitellmHeaderName]: `Bearer ${accessToken}`,
|
||||
"Content-Type": "application/json",
|
||||
},
|
||||
body: JSON.stringify(promptData),
|
||||
});
|
||||
|
||||
if (!response.ok) {
|
||||
const errorData = await response.text();
|
||||
handleError(errorData);
|
||||
throw new Error("Network response was not ok");
|
||||
}
|
||||
|
||||
const data = await response.json();
|
||||
return data;
|
||||
} catch (error) {
|
||||
console.error("Failed to create prompt:", error);
|
||||
throw error;
|
||||
}
|
||||
};
|
||||
|
||||
export const updatePromptCall = async (
|
||||
accessToken: string,
|
||||
promptId: string,
|
||||
promptData: any
|
||||
) => {
|
||||
try {
|
||||
const url = proxyBaseUrl ? `${proxyBaseUrl}/prompts/${promptId}` : `/prompts/${promptId}`;
|
||||
|
||||
const response = await fetch(url, {
|
||||
method: "PUT",
|
||||
headers: {
|
||||
[globalLitellmHeaderName]: `Bearer ${accessToken}`,
|
||||
"Content-Type": "application/json",
|
||||
},
|
||||
body: JSON.stringify(promptData),
|
||||
});
|
||||
|
||||
if (!response.ok) {
|
||||
const errorData = await response.text();
|
||||
handleError(errorData);
|
||||
throw new Error("Network response was not ok");
|
||||
}
|
||||
|
||||
const data = await response.json();
|
||||
return data;
|
||||
} catch (error) {
|
||||
console.error("Failed to update prompt:", error);
|
||||
throw error;
|
||||
}
|
||||
};
|
||||
|
||||
export const deletePromptCall = async (
|
||||
accessToken: string,
|
||||
promptId: string
|
||||
) => {
|
||||
try {
|
||||
const url = proxyBaseUrl ? `${proxyBaseUrl}/prompts/${promptId}` : `/prompts/${promptId}`;
|
||||
|
||||
const response = await fetch(url, {
|
||||
method: "DELETE",
|
||||
headers: {
|
||||
[globalLitellmHeaderName]: `Bearer ${accessToken}`,
|
||||
"Content-Type": "application/json",
|
||||
},
|
||||
});
|
||||
|
||||
if (!response.ok) {
|
||||
const errorData = await response.text();
|
||||
handleError(errorData);
|
||||
throw new Error("Network response was not ok");
|
||||
}
|
||||
|
||||
const data = await response.json();
|
||||
return data;
|
||||
} catch (error) {
|
||||
console.error("Failed to delete prompt:", error);
|
||||
throw error;
|
||||
}
|
||||
};
|
||||
|
||||
export const convertPromptFileToJson = async (
|
||||
accessToken: string,
|
||||
file: File
|
||||
): Promise<{ prompt_id: string; json_data: any }> => {
|
||||
try {
|
||||
const formData = new FormData();
|
||||
formData.append("file", file);
|
||||
|
||||
const url = proxyBaseUrl
|
||||
? `${proxyBaseUrl}/utils/dotprompt_json_converter`
|
||||
: `/utils/dotprompt_json_converter`;
|
||||
|
||||
const response = await fetch(url, {
|
||||
method: "POST",
|
||||
headers: {
|
||||
[globalLitellmHeaderName]: `Bearer ${accessToken}`,
|
||||
},
|
||||
body: formData,
|
||||
});
|
||||
|
||||
if (!response.ok) {
|
||||
const errorData = await response.text();
|
||||
handleError(errorData);
|
||||
throw new Error("Network response was not ok");
|
||||
}
|
||||
|
||||
return await response.json();
|
||||
} catch (error) {
|
||||
console.error("Failed to convert prompt file:", error);
|
||||
throw error;
|
||||
}
|
||||
};
|
||||
|
||||
export const patchPromptCall = async (
|
||||
accessToken: string,
|
||||
promptId: string,
|
||||
promptData: any
|
||||
) => {
|
||||
try {
|
||||
const url = proxyBaseUrl ? `${proxyBaseUrl}/prompts/${promptId}` : `/prompts/${promptId}`;
|
||||
|
||||
const response = await fetch(url, {
|
||||
method: "PATCH",
|
||||
headers: {
|
||||
[globalLitellmHeaderName]: `Bearer ${accessToken}`,
|
||||
"Content-Type": "application/json",
|
||||
},
|
||||
body: JSON.stringify(promptData),
|
||||
});
|
||||
|
||||
if (!response.ok) {
|
||||
const errorData = await response.text();
|
||||
handleError(errorData);
|
||||
throw new Error("Network response was not ok");
|
||||
}
|
||||
|
||||
const data = await response.json();
|
||||
return data;
|
||||
} catch (error) {
|
||||
console.error("Failed to patch prompt:", error);
|
||||
throw error;
|
||||
}
|
||||
};
|
||||
|
||||
export const createGuardrailCall = async (
|
||||
accessToken: string,
|
||||
guardrailData: any
|
||||
|
|
|
|||
|
|
@ -1,8 +1,12 @@
|
|||
import React, { useState, useEffect } from "react"
|
||||
import { Card, Text } from "@tremor/react"
|
||||
import { getPromptsList, PromptSpec, ListPromptsResponse } from "./networking"
|
||||
|
||||
import { Card, Text, Button } from "@tremor/react"
|
||||
import { Modal, message } from "antd"
|
||||
import { getPromptsList, PromptSpec, ListPromptsResponse, deletePromptCall } from "./networking"
|
||||
import PromptTable from "./prompts/prompt_table"
|
||||
import PromptInfoView from "./prompts/prompt_info"
|
||||
import AddPromptForm from "./prompts/add_prompt_form"
|
||||
|
||||
import { isAdminRole } from "@/utils/roles"
|
||||
|
||||
interface PromptsProps {
|
||||
|
|
@ -14,6 +18,9 @@ const PromptsPanel: React.FC<PromptsProps> = ({ accessToken, userRole }) => {
|
|||
const [promptsList, setPromptsList] = useState<PromptSpec[]>([])
|
||||
const [isLoading, setIsLoading] = useState(false)
|
||||
const [selectedPromptId, setSelectedPromptId] = useState<string | null>(null)
|
||||
const [isAddModalVisible, setIsAddModalVisible] = useState(false)
|
||||
const [isDeleting, setIsDeleting] = useState(false)
|
||||
const [promptToDelete, setPromptToDelete] = useState<{ id: string; name: string } | null>(null)
|
||||
|
||||
const isAdmin = userRole ? isAdminRole(userRole) : false
|
||||
|
||||
|
|
@ -38,10 +45,51 @@ const PromptsPanel: React.FC<PromptsProps> = ({ accessToken, userRole }) => {
|
|||
fetchPrompts()
|
||||
}, [accessToken])
|
||||
|
||||
const handlePromptClick = (promptId: string) => {
|
||||
const handlePromptClick = (promptId: string) => {
|
||||
setSelectedPromptId(promptId)
|
||||
}
|
||||
|
||||
const handleAddPrompt = () => {
|
||||
if (selectedPromptId) {
|
||||
setSelectedPromptId(null)
|
||||
}
|
||||
setIsAddModalVisible(true)
|
||||
}
|
||||
|
||||
const handleCloseModal = () => {
|
||||
setIsAddModalVisible(false)
|
||||
}
|
||||
|
||||
const handleSuccess = () => {
|
||||
fetchPrompts()
|
||||
}
|
||||
|
||||
const handleDeleteClick = (promptId: string, promptName: string) => {
|
||||
setPromptToDelete({ id: promptId, name: promptName })
|
||||
}
|
||||
|
||||
const handleDeleteConfirm = async () => {
|
||||
if (!promptToDelete || !accessToken) return
|
||||
|
||||
setIsDeleting(true)
|
||||
try {
|
||||
await deletePromptCall(accessToken, promptToDelete.id)
|
||||
message.success(`Prompt "${promptToDelete.name}" deleted successfully`)
|
||||
fetchPrompts() // Refresh the list
|
||||
} catch (error) {
|
||||
console.error("Error deleting prompt:", error)
|
||||
message.error("Failed to delete prompt")
|
||||
} finally {
|
||||
setIsDeleting(false)
|
||||
setPromptToDelete(null)
|
||||
}
|
||||
}
|
||||
|
||||
const handleDeleteCancel = () => {
|
||||
setPromptToDelete(null)
|
||||
}
|
||||
|
||||
|
||||
return (
|
||||
<div className="w-full mx-auto flex-auto overflow-y-auto m-8 p-2">
|
||||
{selectedPromptId ? (
|
||||
|
|
@ -50,20 +98,48 @@ const PromptsPanel: React.FC<PromptsProps> = ({ accessToken, userRole }) => {
|
|||
onClose={() => setSelectedPromptId(null)}
|
||||
accessToken={accessToken}
|
||||
isAdmin={isAdmin}
|
||||
onDelete={fetchPrompts}
|
||||
/>
|
||||
) : (
|
||||
<>
|
||||
<div className="flex justify-between items-center mb-4">
|
||||
<Text className="text-lg font-semibold">Prompts</Text>
|
||||
<Button onClick={handleAddPrompt} disabled={!accessToken}>
|
||||
+ Add New Prompt
|
||||
</Button>
|
||||
</div>
|
||||
|
||||
<PromptTable
|
||||
promptsList={promptsList}
|
||||
isLoading={isLoading}
|
||||
onPromptClick={handlePromptClick}
|
||||
onDeleteClick={handleDeleteClick}
|
||||
accessToken={accessToken}
|
||||
isAdmin={isAdmin}
|
||||
/>
|
||||
</>
|
||||
)}
|
||||
|
||||
<AddPromptForm
|
||||
visible={isAddModalVisible}
|
||||
onClose={handleCloseModal}
|
||||
accessToken={accessToken}
|
||||
onSuccess={handleSuccess}
|
||||
/>
|
||||
|
||||
{promptToDelete && (
|
||||
<Modal
|
||||
title="Delete Prompt"
|
||||
open={promptToDelete !== null}
|
||||
onOk={handleDeleteConfirm}
|
||||
onCancel={handleDeleteCancel}
|
||||
confirmLoading={isDeleting}
|
||||
okText="Delete"
|
||||
okButtonProps={{ danger: true }}
|
||||
>
|
||||
<p>Are you sure you want to delete prompt: {promptToDelete.name} ?</p>
|
||||
<p>This action cannot be undone.</p>
|
||||
</Modal>
|
||||
)}
|
||||
</div>
|
||||
)
|
||||
}
|
||||
|
|
|
|||
201
ui/litellm-dashboard/src/components/prompts/add_prompt_form.tsx
Normal file
201
ui/litellm-dashboard/src/components/prompts/add_prompt_form.tsx
Normal file
|
|
@ -0,0 +1,201 @@
|
|||
import React, { useState } from "react"
|
||||
import { Modal, Form, Select, Upload, Button, message, Divider } from "antd"
|
||||
import { TextInput } from "@tremor/react"
|
||||
import { UploadOutlined } from "@ant-design/icons"
|
||||
import type { UploadFile, UploadProps } from "antd"
|
||||
import { convertPromptFileToJson, createPromptCall } from "../networking"
|
||||
|
||||
const { Option } = Select
|
||||
|
||||
interface AddPromptFormProps {
|
||||
visible: boolean
|
||||
onClose: () => void
|
||||
accessToken: string | null
|
||||
onSuccess: () => void
|
||||
}
|
||||
|
||||
interface PromptFormData {
|
||||
prompt_id: string
|
||||
prompt_integration: string
|
||||
prompt_file?: File
|
||||
}
|
||||
|
||||
const AddPromptForm: React.FC<AddPromptFormProps> = ({
|
||||
visible,
|
||||
onClose,
|
||||
accessToken,
|
||||
onSuccess,
|
||||
}) => {
|
||||
const [form] = Form.useForm()
|
||||
const [loading, setLoading] = useState(false)
|
||||
const [fileList, setFileList] = useState<UploadFile[]>([])
|
||||
const [promptIntegration, setPromptIntegration] = useState<string>("dotprompt")
|
||||
|
||||
const handleCancel = () => {
|
||||
form.resetFields()
|
||||
setFileList([])
|
||||
setPromptIntegration("dotprompt")
|
||||
onClose()
|
||||
}
|
||||
|
||||
const handleSubmit = async () => {
|
||||
try {
|
||||
const values = await form.validateFields()
|
||||
|
||||
console.log("values: ", values)
|
||||
if (!accessToken) {
|
||||
message.error("Access token is required")
|
||||
return
|
||||
}
|
||||
|
||||
if (promptIntegration === "dotprompt" && fileList.length === 0) {
|
||||
message.error("Please upload a .prompt file")
|
||||
return
|
||||
}
|
||||
|
||||
setLoading(true)
|
||||
|
||||
let promptData: any = {}
|
||||
|
||||
if (promptIntegration === "dotprompt" && fileList.length > 0) {
|
||||
// Convert the uploaded file to JSON
|
||||
const file = fileList[0].originFileObj as File
|
||||
|
||||
try {
|
||||
const conversionResult = await convertPromptFileToJson(accessToken, file)
|
||||
console.log("Conversion result:", conversionResult)
|
||||
|
||||
// Prepare prompt data for creation
|
||||
promptData = {
|
||||
prompt_id: values.prompt_id,
|
||||
litellm_params: {
|
||||
prompt_integration: "dotprompt",
|
||||
prompt_id: conversionResult.prompt_id,
|
||||
prompt_data: conversionResult.json_data
|
||||
},
|
||||
prompt_info: {
|
||||
prompt_type: "db"
|
||||
}
|
||||
}
|
||||
} catch (conversionError) {
|
||||
console.error("Error converting prompt file:", conversionError)
|
||||
message.error("Failed to convert prompt file to JSON")
|
||||
setLoading(false)
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// Create the prompt
|
||||
try {
|
||||
await createPromptCall(accessToken, promptData)
|
||||
message.success("Prompt created successfully!")
|
||||
handleCancel()
|
||||
onSuccess()
|
||||
} catch (createError) {
|
||||
console.error("Error creating prompt:", createError)
|
||||
message.error("Failed to create prompt")
|
||||
}
|
||||
|
||||
} catch (error) {
|
||||
console.error("Form validation error:", error)
|
||||
} finally {
|
||||
setLoading(false)
|
||||
}
|
||||
}
|
||||
|
||||
const uploadProps: UploadProps = {
|
||||
beforeUpload: (file) => {
|
||||
if (!file.name.endsWith('.prompt')) {
|
||||
message.error('Please upload a .prompt file')
|
||||
return false
|
||||
}
|
||||
return false // Prevent automatic upload
|
||||
},
|
||||
fileList,
|
||||
onChange: ({ fileList: newFileList }) => {
|
||||
setFileList(newFileList.slice(-1)) // Keep only the last file
|
||||
},
|
||||
onRemove: () => {
|
||||
setFileList([])
|
||||
},
|
||||
}
|
||||
|
||||
return (
|
||||
<Modal
|
||||
title="Add New Prompt"
|
||||
open={visible}
|
||||
onCancel={handleCancel}
|
||||
footer={[
|
||||
<Button key="cancel" onClick={handleCancel}>
|
||||
Cancel
|
||||
</Button>,
|
||||
<Button
|
||||
key="submit"
|
||||
loading={loading}
|
||||
onClick={handleSubmit}
|
||||
>
|
||||
Create Prompt
|
||||
</Button>,
|
||||
]}
|
||||
width={600}
|
||||
>
|
||||
<Form
|
||||
form={form}
|
||||
layout="vertical"
|
||||
requiredMark={false}
|
||||
>
|
||||
<Form.Item
|
||||
label="Prompt ID"
|
||||
name="prompt_id"
|
||||
rules={[
|
||||
{ required: true, message: "Please enter a prompt ID" },
|
||||
{
|
||||
pattern: /^[a-zA-Z0-9_-]+$/,
|
||||
message: "Prompt ID can only contain letters, numbers, underscores, and hyphens"
|
||||
}
|
||||
]}
|
||||
>
|
||||
<TextInput
|
||||
placeholder="Enter unique prompt ID (e.g., my_prompt_id)"
|
||||
/>
|
||||
</Form.Item>
|
||||
|
||||
<Form.Item
|
||||
label="Prompt Integration"
|
||||
name="prompt_integration"
|
||||
initialValue="dotprompt"
|
||||
>
|
||||
<Select
|
||||
value={promptIntegration}
|
||||
onChange={setPromptIntegration}
|
||||
>
|
||||
<Option value="dotprompt">dotprompt</Option>
|
||||
</Select>
|
||||
</Form.Item>
|
||||
|
||||
{promptIntegration === "dotprompt" && (
|
||||
<>
|
||||
<Divider />
|
||||
<Form.Item
|
||||
label="Prompt File"
|
||||
extra="Upload a .prompt file that follows the Dotprompt specification"
|
||||
>
|
||||
<Upload {...uploadProps}>
|
||||
<Button icon={<UploadOutlined />}>
|
||||
Select .prompt File
|
||||
</Button>
|
||||
</Upload>
|
||||
{fileList.length > 0 && (
|
||||
<div className="mt-2 text-sm text-gray-600">
|
||||
Selected: {fileList[0].name}
|
||||
</div>
|
||||
)}
|
||||
</Form.Item>
|
||||
</>
|
||||
)}
|
||||
</Form>
|
||||
</Modal>
|
||||
)
|
||||
}
|
||||
|
||||
export default AddPromptForm
|
||||
|
|
@ -12,9 +12,9 @@ import {
|
|||
TabPanel,
|
||||
TabPanels,
|
||||
} from "@tremor/react"
|
||||
import { Button, message, Tooltip } from "antd"
|
||||
import { ArrowLeftIcon } from "@heroicons/react/outline"
|
||||
import { getPromptInfo, PromptSpec } from "@/components/networking"
|
||||
import { Button, message, Tooltip, Modal } from "antd"
|
||||
import { ArrowLeftIcon, TrashIcon } from "@heroicons/react/outline"
|
||||
import { getPromptInfo, PromptInfoResponse, PromptSpec, PromptTemplateBase, deletePromptCall } from "@/components/networking"
|
||||
import { copyToClipboard as utilCopyToClipboard } from "@/utils/dataUtils"
|
||||
import { CheckIcon, CopyIcon } from "lucide-react"
|
||||
|
||||
|
|
@ -23,20 +23,25 @@ export interface PromptInfoProps {
|
|||
onClose: () => void
|
||||
accessToken: string | null
|
||||
isAdmin: boolean
|
||||
onDelete?: () => void
|
||||
}
|
||||
|
||||
const PromptInfoView: React.FC<PromptInfoProps> = ({ promptId, onClose, accessToken, isAdmin }) => {
|
||||
const PromptInfoView: React.FC<PromptInfoProps> = ({ promptId, onClose, accessToken, isAdmin, onDelete }) => {
|
||||
const [promptData, setPromptData] = useState<PromptSpec | null>(null)
|
||||
const [promptTemplate, setPromptTemplate] = useState<PromptTemplateBase | null>(null)
|
||||
const [rawApiResponse, setRawApiResponse] = useState<any>(null)
|
||||
const [loading, setLoading] = useState(true)
|
||||
const [copiedStates, setCopiedStates] = useState<Record<string, boolean>>({})
|
||||
const [showDeleteConfirm, setShowDeleteConfirm] = useState(false)
|
||||
const [isDeleting, setIsDeleting] = useState(false)
|
||||
|
||||
const fetchPromptInfo = async () => {
|
||||
try {
|
||||
setLoading(true)
|
||||
if (!accessToken) return
|
||||
const response = await getPromptInfo(accessToken, promptId)
|
||||
setPromptData(response)
|
||||
setPromptData(response.prompt_spec)
|
||||
setPromptTemplate(response.raw_prompt_template)
|
||||
setRawApiResponse(response) // Store the raw response for the Raw JSON tab
|
||||
} catch (error) {
|
||||
message.error("Failed to load prompt information")
|
||||
|
|
@ -75,32 +80,73 @@ const PromptInfoView: React.FC<PromptInfoProps> = ({ promptId, onClose, accessTo
|
|||
}
|
||||
}
|
||||
|
||||
const handleDeleteClick = () => {
|
||||
setShowDeleteConfirm(true)
|
||||
}
|
||||
|
||||
const handleDeleteConfirm = async () => {
|
||||
if (!accessToken || !promptData) return
|
||||
|
||||
setIsDeleting(true)
|
||||
try {
|
||||
await deletePromptCall(accessToken, promptData.prompt_id)
|
||||
message.success(`Prompt "${promptData.prompt_id}" deleted successfully`)
|
||||
onDelete?.() // Call the callback to refresh the parent component
|
||||
onClose() // Close the info view
|
||||
} catch (error) {
|
||||
console.error("Error deleting prompt:", error)
|
||||
message.error("Failed to delete prompt")
|
||||
} finally {
|
||||
setIsDeleting(false)
|
||||
setShowDeleteConfirm(false)
|
||||
}
|
||||
}
|
||||
|
||||
const handleDeleteCancel = () => {
|
||||
setShowDeleteConfirm(false)
|
||||
}
|
||||
|
||||
return (
|
||||
<div className="p-4">
|
||||
<div>
|
||||
<TremorButton icon={ArrowLeftIcon} variant="light" onClick={onClose} className="mb-4">
|
||||
Back to Prompts
|
||||
</TremorButton>
|
||||
<Title>Prompt Details</Title>
|
||||
<div className="flex items-center cursor-pointer">
|
||||
<Text className="text-gray-500 font-mono">{promptData.prompt_id}</Text>
|
||||
<Button
|
||||
type="text"
|
||||
size="small"
|
||||
icon={copiedStates["prompt-id"] ? <CheckIcon size={12} /> : <CopyIcon size={12} />}
|
||||
onClick={() => copyToClipboard(promptData.prompt_id, "prompt-id")}
|
||||
className={`left-2 z-10 transition-all duration-200 ${
|
||||
copiedStates["prompt-id"]
|
||||
? "text-green-600 bg-green-50 border-green-200"
|
||||
: "text-gray-500 hover:text-gray-700 hover:bg-gray-100"
|
||||
}`}
|
||||
/>
|
||||
<div className="flex justify-between items-start mb-4">
|
||||
<div>
|
||||
<Title>Prompt Details</Title>
|
||||
<div className="flex items-center cursor-pointer">
|
||||
<Text className="text-gray-500 font-mono">{promptData.prompt_id}</Text>
|
||||
<Button
|
||||
type="text"
|
||||
size="small"
|
||||
icon={copiedStates["prompt-id"] ? <CheckIcon size={12} /> : <CopyIcon size={12} />}
|
||||
onClick={() => copyToClipboard(promptData.prompt_id, "prompt-id")}
|
||||
className={`left-2 z-10 transition-all duration-200 ${
|
||||
copiedStates["prompt-id"]
|
||||
? "text-green-600 bg-green-50 border-green-200"
|
||||
: "text-gray-500 hover:text-gray-700 hover:bg-gray-100"
|
||||
}`}
|
||||
/>
|
||||
</div>
|
||||
</div>
|
||||
{isAdmin && (
|
||||
<TremorButton
|
||||
icon={TrashIcon}
|
||||
variant="secondary"
|
||||
onClick={handleDeleteClick}
|
||||
className="flex items-center"
|
||||
>
|
||||
Delete Prompt
|
||||
</TremorButton>
|
||||
)}
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<TabGroup>
|
||||
<TabList className="mb-4">
|
||||
<Tab key="overview">Overview</Tab>
|
||||
{promptTemplate ? <Tab key="prompt-template">Prompt Template</Tab> : <></>}
|
||||
{isAdmin ? <Tab key="details">Details</Tab> : <></>}
|
||||
<Tab key="raw-json">Raw JSON</Tab>
|
||||
</TabList>
|
||||
|
|
@ -147,6 +193,55 @@ const PromptInfoView: React.FC<PromptInfoProps> = ({ promptId, onClose, accessTo
|
|||
)}
|
||||
</TabPanel>
|
||||
|
||||
{/* Prompt Template Panel */}
|
||||
{promptTemplate && (
|
||||
<TabPanel>
|
||||
<Card>
|
||||
<div className="flex justify-between items-center mb-4">
|
||||
<Title>Prompt Template</Title>
|
||||
<Button
|
||||
type="text"
|
||||
size="small"
|
||||
icon={copiedStates["prompt-content"] ? <CheckIcon size={16} /> : <CopyIcon size={16} />}
|
||||
onClick={() => copyToClipboard(promptTemplate.content, "prompt-content")}
|
||||
className={`transition-all duration-200 ${
|
||||
copiedStates["prompt-content"]
|
||||
? "text-green-600 bg-green-50 border-green-200"
|
||||
: "text-gray-500 hover:text-gray-700 hover:bg-gray-100"
|
||||
}`}
|
||||
>
|
||||
{copiedStates["prompt-content"] ? "Copied!" : "Copy Content"}
|
||||
</Button>
|
||||
</div>
|
||||
|
||||
<div className="space-y-4">
|
||||
<div>
|
||||
<Text className="font-medium">Template ID</Text>
|
||||
<div className="font-mono text-sm bg-gray-50 p-2 rounded">{promptTemplate.litellm_prompt_id}</div>
|
||||
</div>
|
||||
|
||||
<div>
|
||||
<Text className="font-medium">Content</Text>
|
||||
<div className="mt-2 p-4 bg-gray-50 rounded-md border overflow-auto max-h-96">
|
||||
<pre className="text-sm text-gray-800 whitespace-pre-wrap">{promptTemplate.content}</pre>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{promptTemplate.metadata && Object.keys(promptTemplate.metadata).length > 0 && (
|
||||
<div>
|
||||
<Text className="font-medium">Template Metadata</Text>
|
||||
<div className="mt-2 p-3 bg-gray-50 rounded-md border">
|
||||
<pre className="text-xs text-gray-800 whitespace-pre-wrap overflow-auto max-h-64">
|
||||
{JSON.stringify(promptTemplate.metadata, null, 2)}
|
||||
</pre>
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
</Card>
|
||||
</TabPanel>
|
||||
)}
|
||||
|
||||
{/* Details Panel (only for admins) */}
|
||||
{isAdmin && (
|
||||
<TabPanel>
|
||||
|
|
@ -224,6 +319,20 @@ const PromptInfoView: React.FC<PromptInfoProps> = ({ promptId, onClose, accessTo
|
|||
</TabPanel>
|
||||
</TabPanels>
|
||||
</TabGroup>
|
||||
|
||||
{/* Delete Confirmation Modal */}
|
||||
<Modal
|
||||
title="Delete Prompt"
|
||||
open={showDeleteConfirm}
|
||||
onOk={handleDeleteConfirm}
|
||||
onCancel={handleDeleteCancel}
|
||||
confirmLoading={isDeleting}
|
||||
okText="Delete"
|
||||
okButtonProps={{ danger: true }}
|
||||
>
|
||||
<p>Are you sure you want to delete prompt: <strong>{promptData?.prompt_id}</strong>?</p>
|
||||
<p>This action cannot be undone.</p>
|
||||
</Modal>
|
||||
</div>
|
||||
)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
import React, { useState } from "react"
|
||||
import { Table, TableBody, TableCell, TableHead, TableHeaderCell, TableRow, Button } from "@tremor/react"
|
||||
import { SwitchVerticalIcon, ChevronUpIcon, ChevronDownIcon } from "@heroicons/react/outline"
|
||||
import { SwitchVerticalIcon, ChevronUpIcon, ChevronDownIcon, TrashIcon } from "@heroicons/react/outline"
|
||||
import { Tooltip } from "antd"
|
||||
import {PromptSpec} from "@/components/networking"
|
||||
import {
|
||||
|
|
@ -17,12 +17,18 @@ interface PromptTableProps {
|
|||
promptsList: PromptSpec[]
|
||||
isLoading: boolean
|
||||
onPromptClick?: (id: string) => void
|
||||
onDeleteClick?: (id: string, name: string) => void
|
||||
accessToken: string | null
|
||||
isAdmin: boolean
|
||||
}
|
||||
|
||||
const PromptTable: React.FC<PromptTableProps> = ({
|
||||
promptsList,
|
||||
isLoading,
|
||||
onPromptClick,
|
||||
onDeleteClick,
|
||||
accessToken,
|
||||
isAdmin,
|
||||
}) => {
|
||||
const [sorting, setSorting] = useState<SortingState>([{ id: "created_at", desc: true }])
|
||||
|
||||
|
|
@ -86,6 +92,33 @@ const PromptTable: React.FC<PromptTableProps> = ({
|
|||
)
|
||||
},
|
||||
},
|
||||
...(isAdmin ? [{
|
||||
header: "Actions",
|
||||
id: "actions",
|
||||
enableSorting: false,
|
||||
cell: ({ row }: any) => {
|
||||
const prompt = row.original
|
||||
const promptName = prompt.prompt_id || "Unknown Prompt"
|
||||
|
||||
return (
|
||||
<div className="flex items-center gap-1">
|
||||
<Tooltip title="Delete prompt">
|
||||
<Button
|
||||
size="xs"
|
||||
variant="light"
|
||||
color="red"
|
||||
onClick={(e) => {
|
||||
e.stopPropagation()
|
||||
onDeleteClick?.(prompt.prompt_id, promptName)
|
||||
}}
|
||||
icon={TrashIcon}
|
||||
className="text-red-500 hover:text-red-700 hover:bg-red-50"
|
||||
/>
|
||||
</Tooltip>
|
||||
</div>
|
||||
)
|
||||
},
|
||||
}] : []),
|
||||
]
|
||||
|
||||
const table = useReactTable({
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue