fix(prompts): fix prompt info lookup and delete using correct IDs (#19358)

* fix(prompts): fix prompt info lookup and delete using correct IDs

* add regression tests cases
This commit is contained in:
Harshit Jain
2026-01-20 12:28:34 -08:00
committed by GitHub
parent 56bf6001e9
commit 20323feecc
3 changed files with 315 additions and 89 deletions
+97 -85
View File
@@ -36,13 +36,13 @@ router = APIRouter()
def get_base_prompt_id(prompt_id: str) -> str:
"""
Extract the base prompt ID by stripping the version suffix if present.
Args:
prompt_id: Prompt ID that may include version suffix (e.g., "jack_success.v1" or "jack_success_v1")
Returns:
Base prompt ID without version suffix (e.g., "jack_success")
Examples:
>>> get_base_prompt_id("jack_success.v1")
"jack_success"
@@ -63,13 +63,13 @@ def get_base_prompt_id(prompt_id: str) -> str:
def get_version_number(prompt_id: str) -> int:
"""
Extract the version number from a versioned prompt ID.
Args:
prompt_id: Prompt ID that may include version suffix (e.g., "jack_success.v2" or "jack_success_v2")
Returns:
Version number (defaults to 1 if no version suffix or invalid format)
Examples:
>>> get_version_number("jack_success.v2")
2
@@ -85,7 +85,7 @@ def get_version_number(prompt_id: str) -> int:
return int(version_str)
except ValueError:
pass
# Try underscore separator (_v)
if "_v" in prompt_id:
version_str = prompt_id.split("_v")[1]
@@ -93,21 +93,21 @@ def get_version_number(prompt_id: str) -> int:
return int(version_str)
except ValueError:
pass
return 1
def construct_versioned_prompt_id(prompt_id: str, version: Optional[int] = None) -> str:
"""
Construct a versioned prompt ID from a base prompt_id and version number.
Args:
prompt_id: Base prompt ID (e.g., "jack_success")
version: Version number (if None, returns the base prompt_id unchanged)
Returns:
Versioned prompt ID (e.g., "jack_success.v4")
Examples:
>>> construct_versioned_prompt_id("jack_success", 4)
"jack_success.v4"
@@ -118,7 +118,7 @@ def construct_versioned_prompt_id(prompt_id: str, version: Optional[int] = None)
"""
if version is None:
return prompt_id
# Strip any existing version suffix first
base_id = get_base_prompt_id(prompt_id)
return f"{base_id}.v{version}"
@@ -127,14 +127,14 @@ def construct_versioned_prompt_id(prompt_id: str, version: Optional[int] = None)
def get_latest_version_prompt_id(prompt_id: str, all_prompt_ids: Dict[str, Any]) -> str:
"""
Find the latest version of a prompt from available prompt IDs.
Args:
prompt_id: Base prompt ID or versioned prompt ID (e.g., "jack_success" or "jack_success.v2")
all_prompt_ids: Dictionary of all available prompt IDs (keys are prompt IDs)
Returns:
The prompt ID with the highest version number, or the original prompt_id if no versions exist
Examples:
>>> all_ids = {"jack.v1": {}, "jack.v2": {}, "jack.v3": {}}
>>> get_latest_version_prompt_id("jack", all_ids)
@@ -146,14 +146,14 @@ def get_latest_version_prompt_id(prompt_id: str, all_prompt_ids: Dict[str, Any])
"simple"
"""
base_id = get_base_prompt_id(prompt_id=prompt_id)
# Find all versions of this prompt
matching_versions = []
for stored_prompt_id in all_prompt_ids.keys():
if get_base_prompt_id(prompt_id=stored_prompt_id) == base_id:
version_num = get_version_number(prompt_id=stored_prompt_id)
matching_versions.append((version_num, stored_prompt_id))
# Use the highest version number
if matching_versions:
matching_versions.sort(reverse=True)
@@ -166,45 +166,47 @@ def get_latest_version_prompt_id(prompt_id: str, all_prompt_ids: Dict[str, Any])
def get_latest_prompt_versions(prompts: List[PromptSpec]) -> List[PromptSpec]:
"""
Filter a list of prompts to return only the latest version of each unique prompt.
Args:
prompts: List of PromptSpec objects
Returns:
List of PromptSpec objects with only the latest version of each prompt
"""
latest_prompts: Dict[str, PromptSpec] = {}
for prompt in prompts:
base_id = get_base_prompt_id(prompt_id=prompt.prompt_id)
version = get_version_number(prompt_id=prompt.prompt_id)
# Keep the prompt with the highest version number
if base_id not in latest_prompts:
latest_prompts[base_id] = prompt
else:
existing_version = get_version_number(prompt_id=latest_prompts[base_id].prompt_id)
existing_version = get_version_number(
prompt_id=latest_prompts[base_id].prompt_id
)
if version > existing_version:
latest_prompts[base_id] = prompt
return list(latest_prompts.values())
async def get_next_version_for_prompt(prisma_client, prompt_id: str) -> int:
"""
Get the next version number for a prompt.
Args:
prisma_client: Prisma database client
prompt_id: Base prompt ID
Returns:
Next version number (1 if no versions exist, max_version + 1 otherwise)
"""
existing_prompts = await prisma_client.db.litellm_prompttable.find_many(
where={"prompt_id": prompt_id}
)
if existing_prompts:
max_version = max(p.version for p in existing_prompts)
return max_version + 1
@@ -215,27 +217,27 @@ async def get_next_version_for_prompt(prisma_client, prompt_id: str) -> int:
def create_versioned_prompt_spec(db_prompt) -> PromptSpec:
"""
Helper function to create a PromptSpec with versioned prompt_id from a DB prompt entry.
Args:
db_prompt: The DB prompt object (from prisma)
Returns:
PromptSpec with versioned prompt_id (e.g., "chat_prompt.v1")
"""
import json
from litellm.types.prompts.init_prompts import PromptLiteLLMParams
prompt_dict = db_prompt.model_dump()
base_prompt_id = prompt_dict["prompt_id"]
version = prompt_dict.get("version", 1)
# Parse litellm_params
litellm_params_data = prompt_dict.get("litellm_params")
if isinstance(litellm_params_data, str):
litellm_params_data = json.loads(litellm_params_data)
litellm_params = PromptLiteLLMParams(**litellm_params_data)
# Parse prompt_info
prompt_info_data = prompt_dict.get("prompt_info")
if prompt_info_data:
@@ -244,10 +246,10 @@ def create_versioned_prompt_spec(db_prompt) -> PromptSpec:
prompt_info = PromptInfo(**prompt_info_data)
else:
prompt_info = PromptInfo(prompt_type="db")
# Create versioned prompt_id
versioned_prompt_id = f"{base_prompt_id}.v{version}"
return PromptSpec(
prompt_id=versioned_prompt_id,
litellm_params=litellm_params,
@@ -319,10 +321,14 @@ async def list_prompts(
prompt_list = []
for prompt_id in prompts:
if prompt_id in IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS:
original_prompt = IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS[prompt_id]
original_prompt = IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS[
prompt_id
]
# Create a copy with base prompt_id (without version suffix)
prompt_copy = PromptSpec(
prompt_id=get_base_prompt_id(prompt_id=original_prompt.prompt_id),
prompt_id=get_base_prompt_id(
prompt_id=original_prompt.prompt_id
),
litellm_params=original_prompt.litellm_params,
prompt_info=original_prompt.prompt_info,
created_at=original_prompt.created_at,
@@ -407,32 +413,33 @@ async def get_prompt_versions(
raise HTTPException(
status_code=403, detail="Only proxy admins can view prompt versions"
)
# Strip version suffix if provided (e.g., "jack_success.v1" -> "jack_success")
base_prompt_id = get_base_prompt_id(prompt_id=prompt_id)
# Get all prompts and filter by base_prompt_id
all_prompts = list(IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS.values())
prompt_versions = [
prompt for prompt in all_prompts
prompt
for prompt in all_prompts
if get_base_prompt_id(prompt_id=prompt.prompt_id) == base_prompt_id
]
if not prompt_versions:
raise HTTPException(
status_code=404, detail=f"No versions found for prompt ID {base_prompt_id}"
)
# Create response with explicit version field for each prompt
versioned_prompts = []
for prompt in prompt_versions:
# Extract version number from the root prompt_id which has version suffix
# (e.g., "jack-sparrow.v3" -> 3)
version_number = get_version_number(prompt_id=prompt.prompt_id)
# Strip version from prompt_id for clean display
base_prompt_id = get_base_prompt_id(prompt_id=prompt.prompt_id)
# Create a copy with explicit version field and clean prompt_id
versioned_prompt = PromptSpec(
prompt_id=base_prompt_id, # Clean ID without version (e.g., "jack-sparrow")
@@ -443,10 +450,10 @@ async def get_prompt_versions(
version=version_number, # Explicit version field (e.g., 3)
)
versioned_prompts.append(versioned_prompt)
# Sort by version number (descending - newest first)
versioned_prompts.sort(key=lambda p: p.version or 1, reverse=True)
return ListPromptsResponse(prompts=versioned_prompts)
@@ -518,21 +525,21 @@ async def get_prompt_info(
# Try to get prompt directly first
prompt_spec = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(prompt_id)
# If not found, try to find the latest version
if prompt_spec is None:
latest_prompt_id = get_latest_version_prompt_id(
prompt_id=prompt_id,
all_prompt_ids=IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS
all_prompt_ids=IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS,
)
prompt_spec = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(latest_prompt_id)
if prompt_spec is None:
raise HTTPException(status_code=400, detail=f"Prompt {prompt_id} not found")
# Extract version number from the prompt_id
version_number = get_version_number(prompt_id=prompt_spec.prompt_id)
# Create a copy of the prompt spec with the base prompt ID (stripped of version)
# and explicit version field for consistency with list_prompts and versions endpoints
prompt_spec_response = PromptSpec(
@@ -547,7 +554,9 @@ async def get_prompt_info(
# 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)
prompt_callback = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id(
prompt_spec.prompt_id
)
if prompt_callback is not None:
# Extract content based on integration type
integration_name = prompt_callback.integration_name
@@ -723,12 +732,12 @@ async def update_prompt(
try:
# Strip version suffix from prompt_id if present (e.g., "jack_success.v1" -> "jack_success")
base_prompt_id = get_base_prompt_id(prompt_id=prompt_id)
# Check if any version exists
existing_prompts = await prisma_client.db.litellm_prompttable.find_many(
where={"prompt_id": base_prompt_id}
)
if not existing_prompts:
raise HTTPException(
status_code=404, detail=f"Prompt with ID {base_prompt_id} not found"
@@ -736,7 +745,10 @@ async def update_prompt(
# Check if it's a config prompt
existing_in_memory = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(prompt_id)
if existing_in_memory and existing_in_memory.prompt_info.prompt_type == "config":
if (
existing_in_memory
and existing_in_memory.prompt_info.prompt_type == "config"
):
raise HTTPException(
status_code=400,
detail="Cannot update config prompts.",
@@ -828,17 +840,19 @@ async def delete_prompt(
try:
# Try to get prompt directly first
existing_prompt = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(prompt_id)
# If not found, try to find the latest version
if existing_prompt is None:
latest_prompt_id = get_latest_version_prompt_id(
prompt_id=prompt_id,
all_prompt_ids=IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS
all_prompt_ids=IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS,
)
existing_prompt = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(
latest_prompt_id
)
existing_prompt = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(latest_prompt_id)
# Use the resolved prompt_id for deletion
prompt_id = latest_prompt_id
if existing_prompt is None:
raise HTTPException(
status_code=404, detail=f"Prompt with ID {prompt_id} not found"
@@ -850,17 +864,18 @@ async def delete_prompt(
detail="Cannot delete config prompts.",
)
# Delete the prompt from the database
# Get the base prompt ID (without version suffix) for database deletion
base_prompt_id = get_base_prompt_id(prompt_id=prompt_id)
# Delete all versions of the prompt from the database
await prisma_client.db.litellm_prompttable.delete_many(
where={"prompt_id": prompt_id}
where={"prompt_id": base_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]
# Remove all versions of the prompt from memory
IN_MEMORY_PROMPT_REGISTRY.delete_prompts_by_base_id(base_prompt_id)
return {"message": f"Prompt {prompt_id} deleted successfully"}
return {"message": f"Prompt {base_prompt_id} deleted successfully"}
except HTTPException as e:
raise e
@@ -1036,68 +1051,66 @@ async def test_prompt(
user_temperature,
version,
)
try:
# Parse the dotprompt content and create PromptTemplate
prompt_manager = PromptManager()
frontmatter, template_content = prompt_manager._parse_frontmatter(
content=request.dotprompt_content
)
# Create PromptTemplate to leverage existing parameter extraction logic
template = PromptTemplate(
content=template_content,
metadata=frontmatter,
template_id="test_prompt"
content=template_content, metadata=frontmatter, template_id="test_prompt"
)
# Extract model from template
if not template.model:
raise HTTPException(
status_code=400,
detail="Model is required in dotprompt metadata"
status_code=400, detail="Model is required in dotprompt metadata"
)
# Always render the template to extract system messages and other metadata
variables = request.prompt_variables or {}
rendered_content = prompt_manager.jinja_env.from_string(
template_content
).render(**variables)
# Convert rendered content to messages using DotpromptManager's method
dotprompt_manager = DotpromptManager()
rendered_messages = dotprompt_manager._convert_to_messages(
rendered_content=rendered_content
)
if not rendered_messages:
raise HTTPException(
status_code=400,
detail="No messages found in rendered prompt"
status_code=400, detail="No messages found in rendered prompt"
)
# If conversation history is provided, use it but preserve system messages
if request.conversation_history:
# Extract system messages from rendered prompt
system_messages = [msg for msg in rendered_messages if msg.get("role") == "system"]
system_messages = [
msg for msg in rendered_messages if msg.get("role") == "system"
]
# Use conversation history for user/assistant messages
messages = system_messages + request.conversation_history
else:
messages = rendered_messages # type: ignore[assignment]
# Use PromptTemplate's optional_params which already extracts all parameters
optional_params = template.optional_params.copy()
# Always stream the response
optional_params["stream"] = True
# Build request data for chat completion
data = {
"model": template.model,
"messages": messages,
}
data.update(optional_params)
# Use ProxyBaseLLMRequestProcessing to go through all proxy logic
base_llm_response_processor = ProxyBaseLLMRequestProcessing(data=data)
result = await base_llm_response_processor.base_process_llm_request(
@@ -1118,12 +1131,12 @@ async def test_prompt(
user_api_base=user_api_base,
version=version,
)
if isinstance(result, BaseModel):
return result.model_dump(exclude_none=True, exclude_unset=True)
else:
return result
except HTTPException as e:
raise e
except Exception as e:
@@ -1192,4 +1205,3 @@ async def convert_prompt_file_to_json(
temp_file_path.parent.rmdir()
except OSError:
pass # Directory not empty or other error
+29 -4
View File
@@ -97,9 +97,9 @@ class InMemoryPromptRegistry:
Prompt id to Prompt object mapping
"""
self.prompt_id_to_custom_prompt: Dict[str, Optional[CustomPromptManagement]] = (
{}
)
self.prompt_id_to_custom_prompt: Dict[
str, Optional[CustomPromptManagement]
] = {}
"""
Guardrail id to CustomGuardrail object mapping
"""
@@ -174,5 +174,30 @@ class InMemoryPromptRegistry:
"""
return self.prompt_id_to_custom_prompt.get(prompt_id)
def delete_prompts_by_base_id(self, base_prompt_id: str) -> list[str]:
"""
Delete all prompts matching the given base prompt ID from memory.
IN_MEMORY_PROMPT_REGISTRY = InMemoryPromptRegistry()
Args:
base_prompt_id: The base prompt ID (without version suffix)
Returns:
List of prompt IDs that were deleted
"""
from litellm.proxy.prompts.prompt_endpoints import get_base_prompt_id
prompts_to_delete = [
pid
for pid in self.IN_MEMORY_PROMPTS.keys()
if get_base_prompt_id(prompt_id=pid) == base_prompt_id
]
for pid in prompts_to_delete:
del self.IN_MEMORY_PROMPTS[pid]
if pid in self.prompt_id_to_custom_prompt:
del self.prompt_id_to_custom_prompt[pid]
return prompts_to_delete
IN_MEMORY_PROMPT_REGISTRY = InMemoryPromptRegistry()
@@ -0,0 +1,189 @@
import pytest
from unittest.mock import MagicMock, AsyncMock, patch
from litellm.proxy._types import UserAPIKeyAuth, LitellmUserRoles
from litellm.types.prompts.init_prompts import (
PromptSpec,
PromptLiteLLMParams,
PromptInfo,
)
@pytest.mark.asyncio
async def test_delete_prompt_success():
"""
Test that delete_prompt correctly identifies the base prompt ID
and deletes all versions from DB and memory.
"""
from litellm.proxy.prompts.prompt_endpoints import delete_prompt
# Mock user auth
mock_user_auth = UserAPIKeyAuth(
api_key="sk-1234", user_role=LitellmUserRoles.PROXY_ADMIN
)
# Mock DB Client
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_prompttable.delete_many = AsyncMock(return_value=None)
# Mock In-Memory Registry
with patch(
"litellm.proxy.prompts.prompt_registry.IN_MEMORY_PROMPT_REGISTRY"
) as mock_registry:
# User passes "test_prompt.v2"
# We simulate that get_prompt_by_id returns the prompt spec for v2
prompt_spec = PromptSpec(
prompt_id="test_prompt.v2",
litellm_params=PromptLiteLLMParams(
prompt_id="test_prompt", prompt_integration="dotprompt"
),
prompt_info=PromptInfo(prompt_type="db"),
)
mock_registry.get_prompt_by_id.return_value = prompt_spec
# Patch the prisma client in the endpoint module
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client):
response = await delete_prompt(
prompt_id="test_prompt.v2", user_api_key_dict=mock_user_auth
)
# Assertions
expected_base_id = "test_prompt"
# 1. DB deletion should use base ID
mock_prisma_client.db.litellm_prompttable.delete_many.assert_called_once_with(
where={"prompt_id": expected_base_id}
)
# 2. Memory deletion should use base ID
mock_registry.delete_prompts_by_base_id.assert_called_once_with(
expected_base_id
)
assert response == {
"message": f"Prompt {expected_base_id} deleted successfully"
}
@pytest.mark.asyncio
async def test_delete_prompt_by_base_id_success():
"""
Test that delete_prompt works when passed a base ID directly,
finding the latest version to confirm existence, then deleting.
"""
from litellm.proxy.prompts.prompt_endpoints import delete_prompt
# Mock user auth
mock_user_auth = UserAPIKeyAuth(
api_key="sk-1234", user_role=LitellmUserRoles.PROXY_ADMIN
)
# Mock DB Client
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_prompttable.delete_many = AsyncMock(return_value=None)
# Mock In-Memory Registry
with patch(
"litellm.proxy.prompts.prompt_registry.IN_MEMORY_PROMPT_REGISTRY"
) as mock_registry:
# User passes "test_prompt" (base ID)
# 1. get_prompt_by_id("test_prompt") -> None (if it's not registered as base)
# 2. It calls get_latest_version_prompt_id -> returns "test_prompt.v3"
# 3. get_prompt_by_id("test_prompt.v3") -> returns Spec
# Setup mocks behavior
def get_prompt_side_effect(prompt_id):
if prompt_id == "test_prompt":
return None
if prompt_id == "test_prompt.v3":
return PromptSpec(
prompt_id="test_prompt.v3",
litellm_params=PromptLiteLLMParams(
prompt_id="test_prompt", prompt_integration="dotprompt"
),
prompt_info=PromptInfo(prompt_type="db"),
)
return None
mock_registry.get_prompt_by_id.side_effect = get_prompt_side_effect
mock_registry.IN_MEMORY_PROMPTS = {
"test_prompt.v1": {},
"test_prompt.v2": {},
"test_prompt.v3": {},
}
# Patch the prisma client in the endpoint module
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client):
response = await delete_prompt(
prompt_id="test_prompt", user_api_key_dict=mock_user_auth
)
# Assertions
expected_base_id = "test_prompt"
# 1. DB deletion should use base ID
mock_prisma_client.db.litellm_prompttable.delete_many.assert_called_once_with(
where={"prompt_id": expected_base_id}
)
# 2. Memory deletion should use base ID
mock_registry.delete_prompts_by_base_id.assert_called_once_with(
expected_base_id
)
assert response == {
"message": f"Prompt {expected_base_id} deleted successfully"
}
@pytest.mark.asyncio
async def test_get_prompt_info_by_base_id():
"""
Test that get_prompt_info correctly resolves a base ID to the latest version.
"""
from litellm.proxy.prompts.prompt_endpoints import get_prompt_info
# Mock user auth
mock_user_auth = UserAPIKeyAuth(
api_key="sk-1234", user_role=LitellmUserRoles.PROXY_ADMIN
)
# Mock In-Memory Registry
with patch(
"litellm.proxy.prompts.prompt_registry.IN_MEMORY_PROMPT_REGISTRY"
) as mock_registry:
# Setup mocks behavior
prompt_spec_v3 = PromptSpec(
prompt_id="test_prompt.v3",
litellm_params=PromptLiteLLMParams(
prompt_id="test_prompt", prompt_integration="dotprompt"
),
prompt_info=PromptInfo(prompt_type="db"),
)
# When get_prompt_by_id is called with "test_prompt", return None (so it searches versions)
# When called with "test_prompt.v3", return the spec
def get_prompt_side_effect(prompt_id):
if prompt_id == "test_prompt":
return None
if prompt_id == "test_prompt.v3":
return prompt_spec_v3
return None
mock_registry.get_prompt_by_id.side_effect = get_prompt_side_effect
mock_registry.IN_MEMORY_PROMPTS = {
"test_prompt.v1": {},
"test_prompt.v2": {},
"test_prompt.v3": {},
}
# We also need to mock get_prompt_callback_by_id to avoid content extraction errors/logic
mock_registry.get_prompt_callback_by_id.return_value = None
response = await get_prompt_info(
prompt_id="test_prompt", user_api_key_dict=mock_user_auth
)
assert (
response.prompt_spec.prompt_id == "test_prompt"
) # Should return base ID in spec response
assert response.prompt_spec.version == 3 # Should identify it as version 3