From c7cf18cf67423d93ab171f7122c5b87f0bcfa5db Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Wed, 19 Nov 2025 13:19:56 -0800 Subject: [PATCH] [Feat] Prompt Management - Allow storing prompt version in DB (#16848) * test_dotprompt_auto_detection_with_model_only * fix _auto_detect_prompt_management_logger * test_dotprompt_with_prompt_version * add v1, v2 tests * add _compile_prompt_helper * fix _compile_prompt_helper * test_dotprompt_with_prompt_version * test_dotprompt_with_prompt_version, test_get_prompt_with_version * add version in schema * feat add _get_prompt_spec_for_db_prompt * add _get_prompt_spec_for_db_prompt * feat add _get_prompt_spec_for_db_prompt * update prompt table * add version in prompt DB * test_get_prompt_spec_for_db_prompt_with_versions --- .../migration.sql | 12 ++ .../litellm_proxy_extras/schema.prisma | 6 +- litellm/proxy/prompts/prompt_endpoints.py | 139 +++++++++++++----- litellm/proxy/proxy_server.py | 32 +++- litellm/proxy/schema.prisma | 6 +- schema.prisma | 6 +- tests/test_litellm/proxy/test_proxy_server.py | 44 ++++++ 7 files changed, 203 insertions(+), 42 deletions(-) create mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20251119131227_add_prompt_versioning/migration.sql diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20251119131227_add_prompt_versioning/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20251119131227_add_prompt_versioning/migration.sql new file mode 100644 index 0000000000..a9d9528bd2 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20251119131227_add_prompt_versioning/migration.sql @@ -0,0 +1,12 @@ +-- DropIndex +DROP INDEX "LiteLLM_PromptTable_prompt_id_key"; + +-- AlterTable +ALTER TABLE "LiteLLM_PromptTable" ADD COLUMN "version" INTEGER NOT NULL DEFAULT 1; + +-- CreateIndex +CREATE INDEX "LiteLLM_PromptTable_prompt_id_idx" ON "LiteLLM_PromptTable"("prompt_id"); + +-- CreateIndex +CREATE UNIQUE INDEX "LiteLLM_PromptTable_prompt_id_version_key" ON "LiteLLM_PromptTable"("prompt_id", "version"); + diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index d6b7cebbd1..6cfbb90c36 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -561,11 +561,15 @@ model LiteLLM_GuardrailsTable { // Prompt table for storing prompt configurations model LiteLLM_PromptTable { id String @id @default(uuid()) - prompt_id String @unique + prompt_id String + version Int @default(1) litellm_params Json prompt_info Json? created_at DateTime @default(now()) updated_at DateTime @updatedAt + + @@unique([prompt_id, version]) + @@index([prompt_id]) } model LiteLLM_HealthCheckTable { diff --git a/litellm/proxy/prompts/prompt_endpoints.py b/litellm/proxy/prompts/prompt_endpoints.py index 9ab63895ae..b4677ed374 100644 --- a/litellm/proxy/prompts/prompt_endpoints.py +++ b/litellm/proxy/prompts/prompt_endpoints.py @@ -21,10 +21,76 @@ from litellm.types.prompts.init_prompts import ( PromptTemplateBase, ) - router = APIRouter() +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 + else: + return 1 + + +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: + if isinstance(prompt_info_data, str): + prompt_info_data = json.loads(prompt_info_data) + 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, + prompt_info=prompt_info, + created_at=prompt_dict.get("created_at"), + updated_at=prompt_dict.get("updated_at"), + ) + + class Prompt(BaseModel): prompt_id: str litellm_params: PromptLiteLLMParams @@ -261,19 +327,16 @@ async def create_prompt( ) 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", - ) + # Get next version number + new_version = await get_next_version_for_prompt( + prisma_client=prisma_client, prompt_id=request.prompt_id + ) - # store prompt in db + # Store prompt in db with version prompt_db_entry = await prisma_client.db.litellm_prompttable.create( data={ "prompt_id": request.prompt_id, + "version": new_version, "litellm_params": request.litellm_params.model_dump_json(), "prompt_info": ( request.prompt_info.model_dump_json() @@ -283,7 +346,8 @@ async def create_prompt( } ) - prompt_spec = PromptSpec(**prompt_db_entry.model_dump()) + # Create versioned prompt spec + prompt_spec = create_versioned_prompt_spec(db_prompt=prompt_db_entry) # Initialize the prompt initialized_prompt = IN_MEMORY_PROMPT_REGISTRY.initialize_prompt( @@ -354,45 +418,50 @@ async def update_prompt( ) try: - # Check if prompt exists - existing_prompt = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(prompt_id) - if existing_prompt is None: + # Check if any version exists + existing_prompts = await prisma_client.db.litellm_prompttable.find_many( + where={"prompt_id": request.prompt_id} + ) + + if not existing_prompts: raise HTTPException( - status_code=404, detail=f"Prompt with ID {prompt_id} not found" + status_code=404, detail=f"Prompt with ID {request.prompt_id} not found" ) - if existing_prompt.prompt_info.prompt_type == "config": + # Check if it's a config prompt + base_prompt_id = request.prompt_id + existing_in_memory = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(base_prompt_id) + if existing_in_memory and existing_in_memory.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(), + # Get next version number (UPDATE creates a new version) + new_version = await get_next_version_for_prompt( + prisma_client=prisma_client, prompt_id=request.prompt_id ) - updated_prompt_db_entry = await prisma_client.db.litellm_prompttable.update( - where={"prompt_id": prompt_id}, + # Store new version in db + prompt_db_entry = await prisma_client.db.litellm_prompttable.create( data={ - "litellm_params": updated_prompt_spec.litellm_params.model_dump_json(), - "prompt_info": updated_prompt_spec.prompt_info.model_dump_json(), - }, + "prompt_id": request.prompt_id, + "version": new_version, + "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() + ), + } ) - # 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] + # Create versioned prompt spec + prompt_spec = create_versioned_prompt_spec(db_prompt=prompt_db_entry) - # Initialize the updated prompt + # Initialize the new version initialized_prompt = IN_MEMORY_PROMPT_REGISTRY.initialize_prompt( - prompt=PromptSpec(**updated_prompt_db_entry.model_dump()), - config_file_path=None, + prompt=prompt_spec, config_file_path=None ) if initialized_prompt is None: diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 1bc9655613..4040b0aa70 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -273,7 +273,9 @@ from litellm.proxy.management_endpoints.customer_endpoints import ( from litellm.proxy.management_endpoints.internal_user_endpoints import ( router as internal_user_router, ) -from litellm.proxy.management_endpoints.internal_user_endpoints import user_update +from litellm.proxy.management_endpoints.internal_user_endpoints import ( + user_update, +) from litellm.proxy.management_endpoints.key_management_endpoints import ( delete_verification_tokens, duration_in_seconds, @@ -327,7 +329,9 @@ from litellm.proxy.ocr_endpoints.endpoints import router as ocr_router from litellm.proxy.openai_files_endpoints.files_endpoints import ( router as openai_files_router, ) -from litellm.proxy.openai_files_endpoints.files_endpoints import set_files_config +from litellm.proxy.openai_files_endpoints.files_endpoints import ( + set_files_config, +) from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( passthrough_endpoint_router, ) @@ -417,7 +421,9 @@ from litellm.types.proxy.management_endpoints.ui_sso import ( LiteLLM_UpperboundKeyGenerateParams, ) from litellm.types.realtime import RealtimeQueryParams -from litellm.types.router import DeploymentTypedDict +from litellm.types.router import ( + DeploymentTypedDict, +) from litellm.types.router import ModelInfo as RouterModelInfo from litellm.types.router import ( RouterGeneralSettings, @@ -3576,14 +3582,32 @@ class ProxyConfig: verbose_proxy_logger.exception( f"Error in _check_and_reload_model_cost_map: {str(e)}" ) + def _get_prompt_spec_for_db_prompt(self, db_prompt): + """ + Convert a DB prompt object to a PromptSpec object. + + Handles the versioning of the prompt, if the DB prompt has a version, it will be used to create the versioned prompt_id. + + Args: + db_prompt: The DB prompt object + + Returns: + The PromptSpec object + """ + from litellm.proxy.prompts.prompt_endpoints import create_versioned_prompt_spec + + return create_versioned_prompt_spec(db_prompt=db_prompt) async def _init_prompts_in_db(self, prisma_client: PrismaClient): from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY + from litellm.types.prompts.init_prompts import PromptSpec try: 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) + # Convert DB object to dict and create versioned prompt_id + prompt_spec = self._get_prompt_spec_for_db_prompt(db_prompt=prompt) + IN_MEMORY_PROMPT_REGISTRY.initialize_prompt(prompt=prompt_spec) except Exception as e: verbose_proxy_logger.debug( "litellm.proxy.proxy_server.py::ProxyConfig:_init_prompts_in_db - {}".format( diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index d6b7cebbd1..6cfbb90c36 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -561,11 +561,15 @@ model LiteLLM_GuardrailsTable { // Prompt table for storing prompt configurations model LiteLLM_PromptTable { id String @id @default(uuid()) - prompt_id String @unique + prompt_id String + version Int @default(1) litellm_params Json prompt_info Json? created_at DateTime @default(now()) updated_at DateTime @updatedAt + + @@unique([prompt_id, version]) + @@index([prompt_id]) } model LiteLLM_HealthCheckTable { diff --git a/schema.prisma b/schema.prisma index d6b7cebbd1..6cfbb90c36 100644 --- a/schema.prisma +++ b/schema.prisma @@ -561,11 +561,15 @@ model LiteLLM_GuardrailsTable { // Prompt table for storing prompt configurations model LiteLLM_PromptTable { id String @id @default(uuid()) - prompt_id String @unique + prompt_id String + version Int @default(1) litellm_params Json prompt_info Json? created_at DateTime @default(now()) updated_at DateTime @updatedAt + + @@unique([prompt_id, version]) + @@index([prompt_id]) } model LiteLLM_HealthCheckTable { diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 865dc1b19a..8e87b67933 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -2450,3 +2450,47 @@ async def test_init_sso_settings_in_db_empty_settings(): # Verify empty dictionary assert uppercased_settings == {} + +def test_get_prompt_spec_for_db_prompt_with_versions(): + """ + Test that _get_prompt_spec_for_db_prompt correctly converts database prompts + to PromptSpec with versioned naming convention. + """ + from unittest.mock import MagicMock + + from litellm.proxy.proxy_server import ProxyConfig + + proxy_config = ProxyConfig() + + # Mock database prompt version 1 + mock_prompt_v1 = MagicMock() + mock_prompt_v1.model_dump.return_value = { + "id": "uuid-1", + "prompt_id": "chat_prompt", + "version": 1, + "litellm_params": '{"prompt_id": "chat_prompt", "prompt_integration": "dotprompt", "model": "gpt-3.5-turbo", "messages": [{"role": "user", "content": "v1 content"}]}', + "prompt_info": '{"prompt_type": "db"}', + "created_at": "2024-01-01T00:00:00", + "updated_at": "2024-01-01T00:00:00", + } + + # Mock database prompt version 2 + mock_prompt_v2 = MagicMock() + mock_prompt_v2.model_dump.return_value = { + "id": "uuid-2", + "prompt_id": "chat_prompt", + "version": 2, + "litellm_params": '{"prompt_id": "chat_prompt", "prompt_integration": "dotprompt", "model": "gpt-4", "messages": [{"role": "user", "content": "v2 content"}]}', + "prompt_info": '{"prompt_type": "db"}', + "created_at": "2024-01-02T00:00:00", + "updated_at": "2024-01-02T00:00:00", + } + + # Test version 1 + prompt_spec_v1 = proxy_config._get_prompt_spec_for_db_prompt(db_prompt=mock_prompt_v1) + assert prompt_spec_v1.prompt_id == "chat_prompt.v1" + + # Test version 2 + prompt_spec_v2 = proxy_config._get_prompt_spec_for_db_prompt(db_prompt=mock_prompt_v2) + assert prompt_spec_v2.prompt_id == "chat_prompt.v2" +