Prompt Management (2/2) - New /prompt/list endpoint + key-based access to prompt templates (#13218)

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

allows prompt templates to work before assigning specific models

* feat: initial logic for independent prompt management settings

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

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

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

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

easier management for prompt endpoints

* feat(prompt_endpoints.py): working `/prompt/list` endpoint

returns all available prompts on proxy

* feat(key_management_endpoints.py): support storing 'prompts' in key metadata

allows giving keys access to specific prompts

* feat(prompt_endpoints.py): enable key-based access to /prompts/list

ensures key can only see prompts it has access to

* fix(init_prompts.py): fix linting error

* fix: fix ruff check

* fix(proxy/_types.py): add 'prompts' to newteamrequest

* fix(litellm_logging.py): update logged message with scrubbed value
This commit is contained in:
Krish Dholakia 2025-08-02 10:36:38 -07:00 • committed by GitHub
parent 4cb42f81b6
commit 363c30320f
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
9 changed files with 98 additions and 18 deletions

View file

@ -504,6 +504,15 @@ class Logging(LiteLLMLoggingBaseClass):
if "custom_llm_provider" in self.model_call_details:
self.custom_llm_provider = self.model_call_details["custom_llm_provider"]
def update_messages(self, messages: List[AllMessageValues]):
"""
Update the logged value of the messages in the model_call_details
Allows pre-call hooks to update the messages before the call is made
"""
self.messages = messages
self.model_call_details["messages"] = messages
def should_run_prompt_management_hooks(
self,
non_default_params: Dict,
@ -562,9 +571,7 @@ class Logging(LiteLLMLoggingBaseClass):
custom_logger = (
prompt_management_logger
or self.get_custom_logger_for_prompt_management(
model=model,
non_default_params=non_default_params,
prompt_id=prompt_id,
model=model, non_default_params=non_default_params
)
)
@ -583,7 +590,6 @@ class Logging(LiteLLMLoggingBaseClass):
prompt_label=prompt_label,
prompt_version=prompt_version,
)
self.messages = messages
return model, messages, non_default_params
@ -602,10 +608,7 @@ class Logging(LiteLLMLoggingBaseClass):
custom_logger = (
prompt_management_logger
or self.get_custom_logger_for_prompt_management(
model=model,
tools=tools,
non_default_params=non_default_params,
prompt_id=prompt_id,
model=model, tools=tools, non_default_params=non_default_params
)
)
@ -630,11 +633,7 @@ class Logging(LiteLLMLoggingBaseClass):
return model, messages, non_default_params
def get_custom_logger_for_prompt_management(
self,
model: str,
non_default_params: Dict,
tools: Optional[List[Dict]] = None,
prompt_id: Optional[str] = None,
self, model: str, non_default_params: Dict, tools: Optional[List[Dict]] = None
) -> Optional[CustomLogger]:
"""
Get a custom logger for prompt management based on model name or available callbacks.
@ -645,7 +644,7 @@ class Logging(LiteLLMLoggingBaseClass):
Returns:
A CustomLogger instance if one is found, None otherwise
"""
# First check if model starts with a known custom logger compatible callback
for callback_name in litellm._known_custom_logger_compatible_callbacks:
if model.startswith(callback_name):
custom_logger = _init_custom_logger_compatible_class(

View file

@ -17,4 +17,12 @@ prompts:
litellm_params:
prompt_integration: dotprompt
prompt_id: test_hello_world_prompt
prompt_directory: /Users/krrishdholakia/Documents/litellm/litellm/proxy/test_prompts
prompt_directory: /Users/krrishdholakia/Documents/litellm/litellm/proxy/test_prompts
- 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
litellm_settings:
callbacks: ["datadog_llm_observability"]

View file

@ -513,6 +513,7 @@ class LiteLLMRoutes(enum.Enum):
"/model/delete",
"/user/daily/activity",
"/model/{model_id}/update",
"/prompt/list",
] # routes that manage their own allowed/disallowed logic
## Org Admin Routes ##
@ -683,6 +684,7 @@ class GenerateRequestBase(LiteLLMPydanticObjectBase):
model_rpm_limit: Optional[dict] = None
model_tpm_limit: Optional[dict] = None
guardrails: Optional[List[str]] = None
prompts: Optional[List[str]] = None
blocked: Optional[bool] = None
aliases: Optional[dict] = {}
object_permission: Optional[LiteLLM_ObjectPermissionBase] = None
@ -920,7 +922,10 @@ class LiteLLM_MCPServerTable(LiteLLMPydanticObjectBase):
mcp_access_groups: List[str] = Field(default_factory=list)
mcp_info: Optional[MCPInfo] = None
# Health check status
status: Optional[str] = Field(default="unknown", description="Health status: 'healthy', 'unhealthy', 'unknown'")
status: Optional[str] = Field(
default="unknown",
description="Health status: 'healthy', 'unhealthy', 'unknown'",
)
last_health_check: Optional[datetime] = None
health_check_error: Optional[str] = None
# Stdio-specific fields
@ -1178,6 +1183,7 @@ class NewTeamRequest(TeamBase):
model_aliases: Optional[dict] = None
tags: Optional[list] = None
guardrails: Optional[List[str]] = None
prompts: Optional[List[str]] = None
object_permission: Optional[LiteLLM_ObjectPermissionBase] = None
team_member_budget: Optional[float] = (
None # allow user to set a budget for all team members
@ -2891,6 +2897,7 @@ LiteLLM_ManagementEndpoint_MetadataFields_Premium = [
"guardrails",
"tags",
"team_member_key_duration",
"prompts",
]

View file

@ -343,6 +343,9 @@ class ProxyBaseLLMRequestProcessing:
user_api_key_dict=user_api_key_dict, data=self.data, call_type=route_type # type: ignore
)
if "messages" in self.data and self.data["messages"]:
logging_obj.update_messages(self.data["messages"])
return self.data, logging_obj
async def base_process_llm_request(

View file

@ -58,8 +58,8 @@ from litellm.proxy.utils import (
PrismaClient,
_hash_token_if_needed,
handle_exception_on_proxy,
jsonify_object,
is_valid_api_key,
jsonify_object,
)
from litellm.router import Router
from litellm.secret_managers.main import get_secret
@ -1502,6 +1502,7 @@ async def generate_key_helper_fn( # noqa: PLR0915
model_rpm_limit: Optional[dict] = None,
model_tpm_limit: Optional[dict] = None,
guardrails: Optional[list] = None,
prompts: Optional[list] = None,
teams: Optional[list] = None,
organization_id: Optional[str] = None,
table_name: Optional[Literal["key", "user"]] = None,
@ -1559,6 +1560,9 @@ async def generate_key_helper_fn( # noqa: PLR0915
if guardrails is not None:
metadata = metadata or {}
metadata["guardrails"] = guardrails
if prompts is not None:
metadata = metadata or {}
metadata["prompts"] = prompts
metadata_json = json.dumps(metadata)
validate_model_max_budget(model_max_budget)

View file

@ -0,0 +1,52 @@
"""
CRUD ENDPOINTS FOR PROMPTS
"""
from typing import List, Optional, cast
from fastapi import APIRouter, Depends
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.types.prompts.init_prompts import ListPromptsResponse
router = APIRouter()
@router.get(
"/prompt/list",
tags=["Prompt Management"],
dependencies=[Depends(user_api_key_auth)],
response_model=ListPromptsResponse,
)
async def list_prompts(
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""
List of available prompts for a given key.
"""
from litellm.proxy._types import LitellmUserRoles
from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY
# check key metadata for prompts
key_metadata = user_api_key_dict.metadata
if key_metadata is not None:
prompts = cast(Optional[List[str]], key_metadata.get("prompts", None))
if prompts is not None:
return ListPromptsResponse(
prompts=[
IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS[prompt]
for prompt in prompts
if prompt in IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS
]
)
# check if user is proxy admin - show all prompts
if user_api_key_dict.user_role is not None and (
user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN
or user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value
):
return ListPromptsResponse(
prompts=list(IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS.values())
)
else:
return ListPromptsResponse(prompts=[])

View file

@ -305,6 +305,7 @@ from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
router as pass_through_router,
)
from litellm.proxy.prompts.prompt_endpoints import router as prompts_router
from litellm.proxy.public_endpoints import router as public_endpoints_router
from litellm.proxy.rerank_endpoints.endpoints import router as rerank_router
from litellm.proxy.response_api_endpoints.endpoints import router as response_router
@ -8889,6 +8890,7 @@ app.include_router(cloudzero_router)
app.include_router(caching_router)
app.include_router(analytics_router)
app.include_router(guardrails_router)
app.include_router(prompts_router)
app.include_router(callback_management_endpoints_router)
app.include_router(debugging_endpoints_router)
app.include_router(ui_crud_endpoints_router)

View file

@ -3,6 +3,7 @@ model: gpt-3.5-turbo
input:
schema:
text: string
guardrails: ["azure-text-moderation"]
---
Extract the requested information from the given text. If a piece of information is not present, omit that field from the output.

View file

@ -1,6 +1,6 @@
from datetime import datetime
from enum import Enum
from typing import Dict, Optional
from typing import Dict, List, Optional
from pydantic import BaseModel, ConfigDict
from typing_extensions import Required, TypedDict
@ -25,3 +25,7 @@ class PromptSpec(TypedDict, total=False):
prompt_info: Optional[Dict]
created_at: Optional[datetime]
updated_at: Optional[datetime]
class ListPromptsResponse(BaseModel):
prompts: List[PromptSpec]