mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
4cb42f81b6
commit
363c30320f
9 changed files with 98 additions and 18 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
@ -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",
|
||||
]
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
52
litellm/proxy/prompts/prompt_endpoints.py
Normal file
52
litellm/proxy/prompts/prompt_endpoints.py
Normal 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=[])
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue