mirror of
https://github.com/BerriAI/litellm.git
synced 2026-08-28 05:25:59 +00:00
Merge adfa42096d into 0fba05800d
This commit is contained in:
commit
0ad50a39c3
22 changed files with 686 additions and 552 deletions
|
|
@ -29,6 +29,12 @@ from litellm.proxy._types import (
|
|||
from litellm.proxy.auth.auth_utils import is_request_body_safe
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.common_utils.path_utils import safe_filename
|
||||
from litellm.proxy.prompts.prompt_registry import (
|
||||
DEFAULT_PROMPT_ENVIRONMENT,
|
||||
get_base_prompt_id,
|
||||
get_version_number,
|
||||
prompt_environment_or_default,
|
||||
)
|
||||
from litellm.repositories.table_repositories import PromptRepository
|
||||
from litellm.types.prompts.init_prompts import (
|
||||
ListPromptsResponse,
|
||||
|
|
@ -102,165 +108,20 @@ def _prompt_table(prisma_client: "PrismaClient") -> _PromptTableActions:
|
|||
return PromptRepository(prisma_client).table
|
||||
|
||||
|
||||
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"
|
||||
>>> get_base_prompt_id("jack_success_v1")
|
||||
"jack_success"
|
||||
>>> get_base_prompt_id("jack_success")
|
||||
"jack_success"
|
||||
"""
|
||||
# Try dot separator first (.v)
|
||||
if ".v" in prompt_id:
|
||||
return prompt_id.split(".v")[0]
|
||||
# Try underscore separator (_v)
|
||||
if "_v" in prompt_id:
|
||||
return prompt_id.split("_v")[0]
|
||||
return prompt_id
|
||||
|
||||
|
||||
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
|
||||
>>> get_version_number("jack_success_v2")
|
||||
2
|
||||
>>> get_version_number("jack_success")
|
||||
1
|
||||
"""
|
||||
# Try dot separator first (.v)
|
||||
if ".v" in prompt_id:
|
||||
version_str = prompt_id.split(".v")[1]
|
||||
try:
|
||||
return int(version_str)
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
# Try underscore separator (_v)
|
||||
if "_v" in prompt_id:
|
||||
version_str = prompt_id.split("_v")[1]
|
||||
try:
|
||||
return int(version_str)
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
return 1
|
||||
|
||||
|
||||
def construct_versioned_prompt_id(prompt_id: str, version: int | None = 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"
|
||||
>>> construct_versioned_prompt_id("jack_success", None)
|
||||
"jack_success"
|
||||
>>> construct_versioned_prompt_id("jack_success.v2", 4)
|
||||
"jack_success.v4"
|
||||
"""
|
||||
if version is None:
|
||||
return prompt_id
|
||||
|
||||
# Strip any existing version suffix first
|
||||
base_id: Final = get_base_prompt_id(prompt_id)
|
||||
return f"{base_id}.v{version}"
|
||||
|
||||
|
||||
def get_latest_version_prompt_id(prompt_id: str, all_prompt_ids: Mapping[str, object]) -> 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)
|
||||
"jack.v3"
|
||||
>>> get_latest_version_prompt_id("jack.v1", all_ids)
|
||||
"jack.v3"
|
||||
>>> all_ids = {"simple": {}}
|
||||
>>> get_latest_version_prompt_id("simple", all_ids)
|
||||
"simple"
|
||||
"""
|
||||
base_id: Final = get_base_prompt_id(prompt_id=prompt_id)
|
||||
|
||||
# Find all versions of this prompt
|
||||
matching_versions: Final = []
|
||||
for stored_prompt_id in all_prompt_ids:
|
||||
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)
|
||||
return matching_versions[0][1]
|
||||
else:
|
||||
# No versioned prompts found, use the base ID as-is
|
||||
return prompt_id
|
||||
|
||||
|
||||
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
|
||||
Filter prompts down to the latest version per (base prompt id, environment).
|
||||
"""
|
||||
latest_prompts: Final[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)
|
||||
if version > existing_version:
|
||||
latest_prompts[base_id] = prompt
|
||||
|
||||
sorted_prompts: Final = sorted(prompts, key=lambda prompt: get_version_number(prompt_id=prompt.prompt_id))
|
||||
latest_prompts: Final = {
|
||||
(get_base_prompt_id(prompt_id=prompt.prompt_id), prompt_environment_or_default(prompt.environment)): prompt
|
||||
for prompt in sorted_prompts
|
||||
}
|
||||
return list(latest_prompts.values())
|
||||
|
||||
|
||||
async def get_next_version_for_prompt(
|
||||
prisma_client: "PrismaClient", prompt_id: str, environment: str = "development"
|
||||
prisma_client: "PrismaClient", prompt_id: str, environment: str = DEFAULT_PROMPT_ENVIRONMENT
|
||||
) -> int:
|
||||
"""
|
||||
Get the next version number for a prompt in a specific environment.
|
||||
|
|
@ -403,11 +264,14 @@ async def list_prompts(
|
|||
if key_metadata is not None:
|
||||
prompts: Final = cast(list[str] | None, key_metadata.get("prompts", None))
|
||||
if prompts is not None:
|
||||
all_prompts = [
|
||||
IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS[prompt_id]
|
||||
for prompt_id in prompts
|
||||
if prompt_id in IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS
|
||||
allowed_prompt_ids: Final = frozenset(prompts)
|
||||
allowed_prompts: Final = [
|
||||
spec
|
||||
for spec in IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS.values()
|
||||
if spec.prompt_id in allowed_prompt_ids
|
||||
or get_base_prompt_id(prompt_id=spec.prompt_id) in allowed_prompt_ids
|
||||
]
|
||||
all_prompts = get_latest_prompt_versions(prompts=allowed_prompts)
|
||||
if environment:
|
||||
all_prompts = [p for p in all_prompts if p.environment == environment]
|
||||
prompt_list: Final = []
|
||||
|
|
@ -576,7 +440,7 @@ def _get_prompt_template(prompt_spec: PromptSpec, base_prompt_id: str) -> Prompt
|
|||
metadata=parsed.get("metadata"),
|
||||
)
|
||||
else:
|
||||
prompt_callback: Final = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id(prompt_spec.prompt_id)
|
||||
prompt_callback: Final = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_for_prompt(prompt=prompt_spec)
|
||||
if prompt_callback is not None:
|
||||
integration_name: Final = prompt_callback.integration_name
|
||||
if integration_name == "dotprompt":
|
||||
|
|
@ -690,15 +554,10 @@ async def get_prompt_info(
|
|||
if env_prompts:
|
||||
prompt_spec = create_versioned_prompt_spec(db_prompt=env_prompts[0])
|
||||
|
||||
# Fallback: use in-memory registry (no environment filter)
|
||||
if prompt_spec is None and environment is None:
|
||||
prompt_spec = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(prompt_id)
|
||||
if prompt_spec is None:
|
||||
latest_prompt_id: Final = get_latest_version_prompt_id(
|
||||
prompt_id=prompt_id,
|
||||
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:
|
||||
prompt_spec = IN_MEMORY_PROMPT_REGISTRY.resolve_prompt_spec(
|
||||
prompt_id, version=requested_version, environment=environment
|
||||
)
|
||||
|
||||
if prompt_spec is None:
|
||||
raise HTTPException(
|
||||
|
|
@ -785,7 +644,7 @@ async def create_prompt(
|
|||
environment: Final = (
|
||||
request.prompt_info.environment
|
||||
if request.prompt_info and request.prompt_info.environment
|
||||
else "development"
|
||||
else DEFAULT_PROMPT_ENVIRONMENT
|
||||
)
|
||||
|
||||
# Get next version number
|
||||
|
|
@ -885,7 +744,7 @@ async def update_prompt(
|
|||
environment: Final = (
|
||||
request.prompt_info.environment
|
||||
if request.prompt_info and request.prompt_info.environment
|
||||
else "development"
|
||||
else DEFAULT_PROMPT_ENVIRONMENT
|
||||
)
|
||||
|
||||
# Check if any version of this prompt exists (in any environment)
|
||||
|
|
@ -897,9 +756,7 @@ async def update_prompt(
|
|||
detail=f"Prompt with ID {base_prompt_id} not found",
|
||||
)
|
||||
|
||||
# Check if it's a config prompt
|
||||
existing_in_memory: Final = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(prompt_id)
|
||||
if existing_in_memory and existing_in_memory.prompt_info.prompt_type == "config":
|
||||
if IN_MEMORY_PROMPT_REGISTRY.has_config_prompt(base_prompt_id=base_prompt_id):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="Cannot update config prompts.",
|
||||
|
|
@ -988,40 +845,26 @@ async def delete_prompt(
|
|||
raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value)
|
||||
|
||||
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: Final = get_latest_version_prompt_id(
|
||||
prompt_id=prompt_id,
|
||||
all_prompt_ids=IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS,
|
||||
)
|
||||
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
|
||||
base_prompt_id: Final = get_base_prompt_id(prompt_id=prompt_id)
|
||||
existing_prompt: Final = IN_MEMORY_PROMPT_REGISTRY.resolve_prompt_spec(prompt_id, environment=environment)
|
||||
|
||||
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":
|
||||
if IN_MEMORY_PROMPT_REGISTRY.has_config_prompt(base_prompt_id=base_prompt_id):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="Cannot delete config prompts.",
|
||||
)
|
||||
|
||||
# Get the base prompt ID (without version suffix) for database deletion
|
||||
base_prompt_id: Final = get_base_prompt_id(prompt_id=prompt_id)
|
||||
|
||||
# Build delete filter; scope to environment if provided
|
||||
delete_where: Final[dict[str, str]] = {"prompt_id": base_prompt_id}
|
||||
if environment:
|
||||
delete_where["environment"] = environment
|
||||
|
||||
# Delete versions from the database (scoped to environment if provided)
|
||||
delete_where: Final[dict[str, str]] = {
|
||||
"prompt_id": base_prompt_id,
|
||||
**({"environment": environment} if environment else {}),
|
||||
}
|
||||
await _prompt_table(prisma_client).delete_many(where=delete_where)
|
||||
|
||||
IN_MEMORY_PROMPT_REGISTRY.delete_prompts_by_base_id(base_prompt_id, environment=environment or None)
|
||||
IN_MEMORY_PROMPT_REGISTRY.delete_prompts_by_base_id(
|
||||
base_prompt_id=base_prompt_id, environment=environment or None
|
||||
)
|
||||
|
||||
env_msg: Final = f" from {environment}" if environment else ""
|
||||
return {"message": f"Prompt {base_prompt_id} deleted successfully{env_msg}"}
|
||||
|
|
@ -1093,7 +936,7 @@ async def patch_prompt(
|
|||
try:
|
||||
# Resolve the target row: find the latest version in the given environment
|
||||
base_prompt_id: Final = get_base_prompt_id(prompt_id=prompt_id)
|
||||
env: Final = environment or "development"
|
||||
env: Final = prompt_environment_or_default(environment)
|
||||
requested_version: Final = get_version_number(prompt_id=prompt_id) if prompt_id != base_prompt_id else None
|
||||
|
||||
# Build query to find the exact row by composite unique key
|
||||
|
|
@ -1117,11 +960,7 @@ async def patch_prompt(
|
|||
|
||||
target_row: Final = db_rows[0]
|
||||
|
||||
# Check if prompt exists in memory
|
||||
versioned_id: Final = f"{base_prompt_id}.v{target_row.version}"
|
||||
existing_prompt: Final = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(versioned_id)
|
||||
|
||||
if existing_prompt and existing_prompt.prompt_info.prompt_type == "config":
|
||||
if IN_MEMORY_PROMPT_REGISTRY.has_config_prompt(base_prompt_id=base_prompt_id):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="Cannot update config prompts.",
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
import importlib
|
||||
import os
|
||||
from collections.abc import Callable
|
||||
from collections.abc import Callable, Sequence
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
|
|
@ -14,6 +14,77 @@ from litellm.types.prompts.init_prompts import (
|
|||
|
||||
prompt_initializer_registry = {}
|
||||
|
||||
DEFAULT_PROMPT_ENVIRONMENT: Final = "development"
|
||||
PROMPT_ENVIRONMENT_SERVE_PRECEDENCE: Final = ("production", "staging", "development")
|
||||
|
||||
|
||||
def get_base_prompt_id(prompt_id: str) -> str:
|
||||
"""
|
||||
Extract the base prompt ID by stripping the version suffix if present.
|
||||
|
||||
Examples:
|
||||
>>> get_base_prompt_id("jack_success.v1")
|
||||
"jack_success"
|
||||
>>> get_base_prompt_id("jack_success_v1")
|
||||
"jack_success"
|
||||
>>> get_base_prompt_id("jack_success")
|
||||
"jack_success"
|
||||
"""
|
||||
if ".v" in prompt_id:
|
||||
return prompt_id.split(".v")[0]
|
||||
if "_v" in prompt_id:
|
||||
return prompt_id.split("_v")[0]
|
||||
return prompt_id
|
||||
|
||||
|
||||
def get_version_number(prompt_id: str) -> int:
|
||||
"""
|
||||
Extract the version number from a versioned prompt ID (defaults to 1).
|
||||
|
||||
Examples:
|
||||
>>> get_version_number("jack_success.v2")
|
||||
2
|
||||
>>> get_version_number("jack_success_v2")
|
||||
2
|
||||
>>> get_version_number("jack_success")
|
||||
1
|
||||
"""
|
||||
if ".v" in prompt_id:
|
||||
version_str = prompt_id.split(".v")[1]
|
||||
try:
|
||||
return int(version_str)
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
if "_v" in prompt_id:
|
||||
version_str = prompt_id.split("_v")[1]
|
||||
try:
|
||||
return int(version_str)
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
return 1
|
||||
|
||||
|
||||
def prompt_environment_or_default(environment: str | None) -> str:
|
||||
return environment or DEFAULT_PROMPT_ENVIRONMENT
|
||||
|
||||
|
||||
def registry_key_for_prompt(prompt: PromptSpec) -> str:
|
||||
return f"{prompt.prompt_id}::{prompt_environment_or_default(prompt.environment)}"
|
||||
|
||||
|
||||
def _spec_version(prompt: PromptSpec) -> int:
|
||||
return prompt.version if prompt.version is not None else get_version_number(prompt_id=prompt.prompt_id)
|
||||
|
||||
|
||||
def _default_serve_environment(prompts: Sequence[PromptSpec]) -> str:
|
||||
present: Final = frozenset(prompt_environment_or_default(prompt.environment) for prompt in prompts)
|
||||
ladder_pick: Final = next((env for env in PROMPT_ENVIRONMENT_SERVE_PRECEDENCE if env in present), None)
|
||||
if ladder_pick is not None:
|
||||
return ladder_pick
|
||||
return min(present) if present else DEFAULT_PROMPT_ENVIRONMENT
|
||||
|
||||
|
||||
def get_prompt_initializer_from_integrations():
|
||||
"""
|
||||
|
|
@ -113,17 +184,16 @@ class InMemoryPromptRegistry:
|
|||
"""
|
||||
import litellm
|
||||
|
||||
prompt_id: Final = 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]
|
||||
registry_key: Final = registry_key_for_prompt(prompt)
|
||||
if registry_key in self.IN_MEMORY_PROMPTS:
|
||||
verbose_proxy_logger.debug("prompt already exists in IN_MEMORY_PROMPTS")
|
||||
return self.IN_MEMORY_PROMPTS[registry_key]
|
||||
|
||||
parsed_prompt, custom_prompt_callback = self._build_prompt_callback(prompt=prompt)
|
||||
litellm.logging_callback_manager.add_litellm_callback(custom_prompt_callback)
|
||||
|
||||
# store references to the prompt in memory
|
||||
self.IN_MEMORY_PROMPTS[prompt_id] = parsed_prompt
|
||||
self.prompt_id_to_custom_prompt[prompt_id] = custom_prompt_callback
|
||||
self.IN_MEMORY_PROMPTS[registry_key] = parsed_prompt
|
||||
self.prompt_id_to_custom_prompt[registry_key] = custom_prompt_callback
|
||||
|
||||
return parsed_prompt
|
||||
|
||||
|
|
@ -166,68 +236,93 @@ class InMemoryPromptRegistry:
|
|||
import litellm
|
||||
|
||||
parsed_prompt, new_callback = self._build_prompt_callback(prompt=prompt)
|
||||
stale_callback: Final = self.prompt_id_to_custom_prompt.pop(prompt.prompt_id, None)
|
||||
self.IN_MEMORY_PROMPTS.pop(prompt.prompt_id, None)
|
||||
registry_key: Final = registry_key_for_prompt(parsed_prompt)
|
||||
stale_callback: Final = self.prompt_id_to_custom_prompt.pop(registry_key, None)
|
||||
self.IN_MEMORY_PROMPTS.pop(registry_key, None)
|
||||
if stale_callback is not None:
|
||||
litellm.logging_callback_manager.remove_callback_from_all_lists(stale_callback)
|
||||
litellm.logging_callback_manager.add_litellm_callback(new_callback)
|
||||
self.IN_MEMORY_PROMPTS[prompt.prompt_id] = parsed_prompt
|
||||
self.prompt_id_to_custom_prompt[prompt.prompt_id] = new_callback
|
||||
self.IN_MEMORY_PROMPTS[registry_key] = parsed_prompt
|
||||
self.prompt_id_to_custom_prompt[registry_key] = new_callback
|
||||
return parsed_prompt
|
||||
|
||||
def sync_prompt_from_db(self, prompt: PromptSpec) -> PromptSpec | None:
|
||||
existing: Final = self.IN_MEMORY_PROMPTS.get(prompt.prompt_id)
|
||||
existing: Final = self.IN_MEMORY_PROMPTS.get(registry_key_for_prompt(prompt))
|
||||
if existing is None:
|
||||
return self.initialize_prompt(prompt=prompt)
|
||||
if existing.litellm_params == prompt.litellm_params and existing.prompt_info == prompt.prompt_info:
|
||||
return existing
|
||||
return self.reload_prompt(prompt=prompt)
|
||||
|
||||
def get_prompt_by_id(self, prompt_id: str) -> PromptSpec | None:
|
||||
def resolve_prompt_spec(
|
||||
self,
|
||||
prompt_id: str,
|
||||
version: int | None = None,
|
||||
environment: str | None = None,
|
||||
) -> PromptSpec | None:
|
||||
"""
|
||||
Get a prompt by its ID from memory
|
||||
"""
|
||||
return self.IN_MEMORY_PROMPTS.get(prompt_id)
|
||||
Resolve a prompt spec by base prompt id, optional version, and optional environment.
|
||||
|
||||
def get_prompt_callback_by_id(self, prompt_id: str) -> CustomPromptManagement | None:
|
||||
With no environment, resolves within the default serve environment
|
||||
(production > staging > development > alphabetical first present).
|
||||
With no version, resolves to the highest version in the chosen environment.
|
||||
"""
|
||||
Get a prompt callback by its ID from memory
|
||||
"""
|
||||
return self.prompt_id_to_custom_prompt.get(prompt_id)
|
||||
base_prompt_id: Final = get_base_prompt_id(prompt_id=prompt_id)
|
||||
base_matches: Final = tuple(
|
||||
spec
|
||||
for spec in self.IN_MEMORY_PROMPTS.values()
|
||||
if get_base_prompt_id(prompt_id=spec.prompt_id) == base_prompt_id
|
||||
)
|
||||
if not base_matches:
|
||||
return None
|
||||
resolved_environment: Final = (
|
||||
environment if environment is not None else _default_serve_environment(base_matches)
|
||||
)
|
||||
env_matches: Final = tuple(
|
||||
spec for spec in base_matches if prompt_environment_or_default(spec.environment) == resolved_environment
|
||||
)
|
||||
if not env_matches:
|
||||
return None
|
||||
if version is not None:
|
||||
return next((spec for spec in env_matches if _spec_version(spec) == version), None)
|
||||
return max(env_matches, key=_spec_version)
|
||||
|
||||
def remove_prompt(self, prompt_id: str) -> None:
|
||||
def get_prompt_callback_for_prompt(self, prompt: PromptSpec) -> CustomPromptManagement | None:
|
||||
return self.prompt_id_to_custom_prompt.get(registry_key_for_prompt(prompt))
|
||||
|
||||
def has_config_prompt(self, base_prompt_id: str) -> bool:
|
||||
return any(
|
||||
spec.prompt_info.prompt_type == "config"
|
||||
for spec in self.IN_MEMORY_PROMPTS.values()
|
||||
if get_base_prompt_id(prompt_id=spec.prompt_id) == base_prompt_id
|
||||
)
|
||||
|
||||
def remove_prompt(self, registry_key: str) -> None:
|
||||
import litellm
|
||||
|
||||
self.IN_MEMORY_PROMPTS.pop(prompt_id, None)
|
||||
stale_callback: Final = self.prompt_id_to_custom_prompt.pop(prompt_id, None)
|
||||
self.IN_MEMORY_PROMPTS.pop(registry_key, None)
|
||||
stale_callback: Final = self.prompt_id_to_custom_prompt.pop(registry_key, None)
|
||||
if stale_callback is not None:
|
||||
litellm.logging_callback_manager.remove_callback_from_all_lists(stale_callback)
|
||||
|
||||
def delete_prompts_by_base_id(self, base_prompt_id: str, environment: str | None = None) -> list[str]:
|
||||
"""
|
||||
Delete all prompts matching the given base prompt ID from memory, along with their
|
||||
registered callbacks; scoped to one environment when given.
|
||||
Delete matching prompts from memory, along with their registered callbacks,
|
||||
scoped to one environment when given.
|
||||
|
||||
Args:
|
||||
base_prompt_id: The base prompt ID (without version suffix)
|
||||
environment: When set, only delete prompts deployed to this environment
|
||||
|
||||
Returns:
|
||||
List of prompt IDs that were deleted
|
||||
Returns the registry keys that were deleted.
|
||||
"""
|
||||
from litellm.proxy.prompts.prompt_endpoints import get_base_prompt_id
|
||||
|
||||
prompts_to_delete: Final = [
|
||||
pid
|
||||
for pid, prompt in self.IN_MEMORY_PROMPTS.items()
|
||||
if get_base_prompt_id(prompt_id=pid) == base_prompt_id
|
||||
and (environment is None or prompt.environment == environment)
|
||||
keys_to_delete: Final = [
|
||||
key
|
||||
for key, spec in self.IN_MEMORY_PROMPTS.items()
|
||||
if get_base_prompt_id(prompt_id=spec.prompt_id) == base_prompt_id
|
||||
and (environment is None or prompt_environment_or_default(spec.environment) == environment)
|
||||
]
|
||||
|
||||
for pid in prompts_to_delete:
|
||||
self.remove_prompt(prompt_id=pid)
|
||||
for key in keys_to_delete:
|
||||
self.remove_prompt(registry_key=key)
|
||||
|
||||
return prompts_to_delete
|
||||
return keys_to_delete
|
||||
|
||||
|
||||
IN_MEMORY_PROMPT_REGISTRY: Final = InMemoryPromptRegistry()
|
||||
|
|
|
|||
|
|
@ -7253,7 +7253,7 @@ class ProxyConfig:
|
|||
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.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY, registry_key_for_prompt
|
||||
from litellm.types.prompts.init_prompts import PromptSpec
|
||||
|
||||
def parse_row(db_prompt: object) -> PromptSpec | None:
|
||||
|
|
@ -7268,21 +7268,12 @@ class ProxyConfig:
|
|||
return None
|
||||
|
||||
try:
|
||||
prompt_ids_loaded_before_db_read: Final = frozenset(IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS)
|
||||
registry_keys_loaded_before_db_read: Final = frozenset(IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS)
|
||||
prompts_in_db: Final[Sequence[object]] = await PromptRepository(prisma_client).table.find_many()
|
||||
parsed_specs: Final[tuple[PromptSpec, ...]] = tuple(
|
||||
spec for row in prompts_in_db if (spec := parse_row(row)) is not None
|
||||
)
|
||||
newest_spec_per_id: Final[Mapping[str, PromptSpec]] = MappingProxyType(
|
||||
{
|
||||
spec.prompt_id: spec
|
||||
for spec in sorted(
|
||||
parsed_specs,
|
||||
key=lambda s: s.updated_at.timestamp() if s.updated_at else float("-inf"),
|
||||
)
|
||||
}
|
||||
)
|
||||
for prompt_spec in newest_spec_per_id.values():
|
||||
for prompt_spec in parsed_specs:
|
||||
try:
|
||||
IN_MEMORY_PROMPT_REGISTRY.sync_prompt_from_db(prompt=prompt_spec)
|
||||
except Exception as prompt_sync_error: # noqa: BLE001 # one poisoned row must not block syncing the remaining prompts
|
||||
|
|
@ -7294,15 +7285,16 @@ class ProxyConfig:
|
|||
# An unparsable row still exists in the DB, so skip the sweep rather than unload its in-memory copy
|
||||
every_row_parsed: Final = len(parsed_specs) == len(prompts_in_db)
|
||||
if every_row_parsed:
|
||||
deleted_db_prompt_ids: Final = tuple(
|
||||
prompt_id
|
||||
for prompt_id in prompt_ids_loaded_before_db_read
|
||||
if (loaded_spec := IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS.get(prompt_id)) is not None
|
||||
db_registry_keys: Final = frozenset(registry_key_for_prompt(spec) for spec in parsed_specs)
|
||||
deleted_db_registry_keys: Final = tuple(
|
||||
registry_key
|
||||
for registry_key in registry_keys_loaded_before_db_read
|
||||
if (loaded_spec := IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS.get(registry_key)) is not None
|
||||
and loaded_spec.prompt_info.prompt_type == "db"
|
||||
and prompt_id not in newest_spec_per_id
|
||||
and registry_key not in db_registry_keys
|
||||
)
|
||||
for deleted_prompt_id in deleted_db_prompt_ids:
|
||||
IN_MEMORY_PROMPT_REGISTRY.remove_prompt(prompt_id=deleted_prompt_id)
|
||||
for deleted_registry_key in deleted_db_registry_keys:
|
||||
IN_MEMORY_PROMPT_REGISTRY.remove_prompt(registry_key=deleted_registry_key)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.debug("litellm.proxy.proxy_server.py::ProxyConfig:_init_prompts_in_db - %s", e)
|
||||
|
||||
|
|
|
|||
|
|
@ -1397,28 +1397,27 @@ class ProxyLogging:
|
|||
) -> None:
|
||||
"""Process prompt template if applicable."""
|
||||
|
||||
from litellm.proxy.prompts.prompt_endpoints import (
|
||||
construct_versioned_prompt_id,
|
||||
get_latest_version_prompt_id,
|
||||
)
|
||||
from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY
|
||||
from litellm.responses.utils import ResponsesAPIRequestUtils
|
||||
from litellm.utils import get_non_default_completion_params
|
||||
|
||||
if prompt_version is None:
|
||||
lookup_prompt_id = get_latest_version_prompt_id(
|
||||
prompt_id=prompt_id,
|
||||
all_prompt_ids=IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS,
|
||||
)
|
||||
else:
|
||||
lookup_prompt_id = construct_versioned_prompt_id(prompt_id=prompt_id, version=prompt_version)
|
||||
|
||||
custom_logger: Final = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id(lookup_prompt_id)
|
||||
prompt_spec: Final = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(lookup_prompt_id)
|
||||
raw_prompt_environment: Final = data.get("prompt_environment", None)
|
||||
prompt_environment: Final = raw_prompt_environment if isinstance(raw_prompt_environment, str) else None
|
||||
prompt_spec: Final = IN_MEMORY_PROMPT_REGISTRY.resolve_prompt_spec(
|
||||
prompt_id,
|
||||
version=prompt_version,
|
||||
environment=prompt_environment,
|
||||
)
|
||||
custom_logger: Final = (
|
||||
IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_for_prompt(prompt=prompt_spec)
|
||||
if prompt_spec is not None
|
||||
else None
|
||||
)
|
||||
litellm_prompt_id: str | None = None
|
||||
if prompt_spec is not None:
|
||||
litellm_prompt_id = prompt_spec.litellm_params.prompt_id
|
||||
data.pop("prompt_id", None)
|
||||
data.pop("prompt_environment", None)
|
||||
|
||||
if custom_logger and prompt_spec is not None:
|
||||
is_responses_call: Final = call_type == "aresponses"
|
||||
|
|
@ -1459,6 +1458,7 @@ class ProxyLogging:
|
|||
data.pop("prompt_variables", None)
|
||||
data.pop("prompt_label", None)
|
||||
data.pop("prompt_version", None)
|
||||
data.pop("prompt_environment", None)
|
||||
|
||||
def _process_guardrail_metadata(self, data: dict) -> None:
|
||||
"""Process guardrails from metadata and add to applied_guardrails."""
|
||||
|
|
|
|||
|
|
@ -3541,6 +3541,7 @@ all_litellm_params = (
|
|||
"litellm_system_prompt",
|
||||
"provider_specific_header",
|
||||
"prompt_version",
|
||||
"prompt_environment",
|
||||
"api_base",
|
||||
"force_timeout",
|
||||
"logger_fn",
|
||||
|
|
|
|||
|
|
@ -104,97 +104,6 @@ class TestPromptVersioning:
|
|||
assert get_base_prompt_id(prompt_id="jack") == "jack"
|
||||
assert get_base_prompt_id(prompt_id="my_prompt.v10") == "my_prompt"
|
||||
|
||||
def test_get_latest_version_prompt_id(self):
|
||||
"""
|
||||
Test that get_latest_version_prompt_id returns the highest version
|
||||
"""
|
||||
from litellm.proxy.prompts.prompt_endpoints import get_latest_version_prompt_id
|
||||
|
||||
# Mock prompt IDs dictionary
|
||||
all_prompt_ids = {
|
||||
"jack.v1": {},
|
||||
"jack.v2": {},
|
||||
"jack.v3": {},
|
||||
"jane.v1": {},
|
||||
"simple_prompt": {},
|
||||
}
|
||||
|
||||
# Test with base prompt ID - should return latest version
|
||||
assert (
|
||||
get_latest_version_prompt_id(
|
||||
prompt_id="jack", all_prompt_ids=all_prompt_ids
|
||||
)
|
||||
== "jack.v3"
|
||||
)
|
||||
|
||||
# Test with versioned prompt ID - should still return latest version
|
||||
assert (
|
||||
get_latest_version_prompt_id(
|
||||
prompt_id="jack.v1", all_prompt_ids=all_prompt_ids
|
||||
)
|
||||
== "jack.v3"
|
||||
)
|
||||
|
||||
# Test with single version
|
||||
assert (
|
||||
get_latest_version_prompt_id(
|
||||
prompt_id="jane", all_prompt_ids=all_prompt_ids
|
||||
)
|
||||
== "jane.v1"
|
||||
)
|
||||
|
||||
# Test with non-versioned prompt
|
||||
assert (
|
||||
get_latest_version_prompt_id(
|
||||
prompt_id="simple_prompt", all_prompt_ids=all_prompt_ids
|
||||
)
|
||||
== "simple_prompt"
|
||||
)
|
||||
|
||||
# Test with non-existent prompt
|
||||
assert (
|
||||
get_latest_version_prompt_id(
|
||||
prompt_id="nonexistent", all_prompt_ids=all_prompt_ids
|
||||
)
|
||||
== "nonexistent"
|
||||
)
|
||||
|
||||
def test_construct_versioned_prompt_id(self):
|
||||
"""
|
||||
Test that construct_versioned_prompt_id correctly builds versioned IDs
|
||||
"""
|
||||
from litellm.proxy.prompts.prompt_endpoints import construct_versioned_prompt_id
|
||||
|
||||
# Test with base prompt ID and version
|
||||
assert (
|
||||
construct_versioned_prompt_id(prompt_id="jack_success", version=4)
|
||||
== "jack_success.v4"
|
||||
)
|
||||
|
||||
# Test with None version - should return base ID unchanged
|
||||
assert (
|
||||
construct_versioned_prompt_id(prompt_id="jack_success", version=None)
|
||||
== "jack_success"
|
||||
)
|
||||
|
||||
# Test with existing versioned ID - should replace version
|
||||
assert (
|
||||
construct_versioned_prompt_id(prompt_id="jack_success.v2", version=4)
|
||||
== "jack_success.v4"
|
||||
)
|
||||
|
||||
# Test with hyphenated prompt ID
|
||||
assert (
|
||||
construct_versioned_prompt_id(prompt_id="my-prompt", version=1)
|
||||
== "my-prompt.v1"
|
||||
)
|
||||
|
||||
# Test with double-digit version
|
||||
assert (
|
||||
construct_versioned_prompt_id(prompt_id="test_prompt", version=10)
|
||||
== "test_prompt.v10"
|
||||
)
|
||||
|
||||
|
||||
class TestPromptVersionsEndpoint:
|
||||
"""
|
||||
|
|
@ -444,7 +353,7 @@ class TestAdminViewerReadAccess:
|
|||
"litellm.proxy.prompts.prompt_registry.IN_MEMORY_PROMPT_REGISTRY"
|
||||
) as mock_registry,
|
||||
):
|
||||
mock_registry.get_prompt_by_id.return_value = PromptSpec(
|
||||
mock_registry.resolve_prompt_spec.return_value = PromptSpec(
|
||||
prompt_id="jack.v2",
|
||||
litellm_params=PromptLiteLLMParams(
|
||||
prompt_id="jack",
|
||||
|
|
@ -453,10 +362,101 @@ class TestAdminViewerReadAccess:
|
|||
),
|
||||
prompt_info=PromptInfo(prompt_type="db"),
|
||||
)
|
||||
mock_registry.IN_MEMORY_PROMPTS = {"jack.v1": {}, "jack.v2": {}}
|
||||
mock_registry.get_prompt_callback_by_id.return_value = None
|
||||
mock_registry.get_prompt_callback_for_prompt.return_value = None
|
||||
|
||||
response = await get_prompt_info(prompt_id="jack", user_api_key_dict=viewer)
|
||||
|
||||
assert response.prompt_spec.prompt_id == "jack"
|
||||
assert response.prompt_spec.version == 2
|
||||
|
||||
|
||||
class TestConfigPromptInfoWithEnvironment:
|
||||
"""
|
||||
Regression: /prompts/{id}/info with an environment param must still resolve
|
||||
config-file (in-memory) prompts on a DB-backed proxy instead of 400ing.
|
||||
"""
|
||||
|
||||
def _registry_with_config_prompt(self):
|
||||
from litellm.proxy.prompts.prompt_registry import InMemoryPromptRegistry
|
||||
|
||||
registry = InMemoryPromptRegistry()
|
||||
registry.IN_MEMORY_PROMPTS["envgreet::development"] = PromptSpec(
|
||||
prompt_id="envgreet",
|
||||
litellm_params=PromptLiteLLMParams(
|
||||
prompt_id="envgreet",
|
||||
prompt_integration="dotprompt",
|
||||
dotprompt_content="AHOY {{user_message}}",
|
||||
),
|
||||
prompt_info=PromptInfo(prompt_type="config"),
|
||||
)
|
||||
return registry
|
||||
|
||||
def _prisma_client_with_empty_prompt_table(self):
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_prompttable.find_many = AsyncMock(return_value=[])
|
||||
return mock_prisma
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_prompt_info_with_environment_falls_back_to_registry(self):
|
||||
from unittest.mock import patch
|
||||
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy.prompts.prompt_endpoints import get_prompt_info
|
||||
|
||||
admin = UserAPIKeyAuth(
|
||||
api_key="test_key", user_role=LitellmUserRoles.PROXY_ADMIN
|
||||
)
|
||||
|
||||
with (
|
||||
patch( # test-quality-ok: endpoint reads prisma_client and the registry from module globals at call time; no injection seam
|
||||
"litellm.proxy.proxy_server.prisma_client",
|
||||
self._prisma_client_with_empty_prompt_table(),
|
||||
),
|
||||
patch( # test-quality-ok: endpoint reads prisma_client and the registry from module globals at call time; no injection seam
|
||||
"litellm.proxy.prompts.prompt_registry.IN_MEMORY_PROMPT_REGISTRY",
|
||||
self._registry_with_config_prompt(),
|
||||
),
|
||||
):
|
||||
response = await get_prompt_info(
|
||||
prompt_id="envgreet",
|
||||
environment="development",
|
||||
user_api_key_dict=admin,
|
||||
)
|
||||
|
||||
assert response.prompt_spec.prompt_id == "envgreet"
|
||||
assert response.prompt_spec.litellm_params.dotprompt_content == "AHOY {{user_message}}"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_prompt_info_with_wrong_environment_still_400s(self):
|
||||
from unittest.mock import patch
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy.prompts.prompt_endpoints import get_prompt_info
|
||||
|
||||
admin = UserAPIKeyAuth(
|
||||
api_key="test_key", user_role=LitellmUserRoles.PROXY_ADMIN
|
||||
)
|
||||
|
||||
with (
|
||||
patch( # test-quality-ok: endpoint reads prisma_client and the registry from module globals at call time; no injection seam
|
||||
"litellm.proxy.proxy_server.prisma_client",
|
||||
self._prisma_client_with_empty_prompt_table(),
|
||||
),
|
||||
patch( # test-quality-ok: endpoint reads prisma_client and the registry from module globals at call time; no injection seam
|
||||
"litellm.proxy.prompts.prompt_registry.IN_MEMORY_PROMPT_REGISTRY",
|
||||
self._registry_with_config_prompt(),
|
||||
),
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await get_prompt_info(
|
||||
prompt_id="envgreet",
|
||||
environment="production",
|
||||
user_api_key_dict=admin,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert "environment production" in exc_info.value.detail
|
||||
|
|
|
|||
|
|
@ -52,8 +52,6 @@ async def test_delete_prompt_success():
|
|||
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(
|
||||
|
|
@ -61,7 +59,8 @@ async def test_delete_prompt_success():
|
|||
),
|
||||
prompt_info=PromptInfo(prompt_type="db"),
|
||||
)
|
||||
mock_registry.get_prompt_by_id.return_value = prompt_spec
|
||||
mock_registry.resolve_prompt_spec.return_value = prompt_spec
|
||||
mock_registry.has_config_prompt.return_value = False
|
||||
|
||||
# Patch the prisma client in the endpoint module
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client):
|
||||
|
|
@ -79,7 +78,7 @@ async def test_delete_prompt_success():
|
|||
|
||||
# 2. Memory deletion should use base ID
|
||||
mock_registry.delete_prompts_by_base_id.assert_called_once_with(
|
||||
expected_base_id, environment=None
|
||||
base_prompt_id=expected_base_id, environment=None
|
||||
)
|
||||
|
||||
assert response == {
|
||||
|
|
@ -108,31 +107,14 @@ async def test_delete_prompt_by_base_id_success():
|
|||
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": {},
|
||||
}
|
||||
mock_registry.resolve_prompt_spec.return_value = PromptSpec(
|
||||
prompt_id="test_prompt.v3",
|
||||
litellm_params=PromptLiteLLMParams(
|
||||
prompt_id="test_prompt", prompt_integration="dotprompt"
|
||||
),
|
||||
prompt_info=PromptInfo(prompt_type="db"),
|
||||
)
|
||||
mock_registry.has_config_prompt.return_value = False
|
||||
|
||||
# Patch the prisma client in the endpoint module
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client):
|
||||
|
|
@ -150,7 +132,7 @@ async def test_delete_prompt_by_base_id_success():
|
|||
|
||||
# 2. Memory deletion should use base ID
|
||||
mock_registry.delete_prompts_by_base_id.assert_called_once_with(
|
||||
expected_base_id, environment=None
|
||||
base_prompt_id=expected_base_id, environment=None
|
||||
)
|
||||
|
||||
assert response == {
|
||||
|
|
@ -169,11 +151,12 @@ async def test_delete_prompt_environment_scope_reaches_db_and_registry():
|
|||
with patch( # test-quality-ok: stubs the collaborator so the test pins what the endpoint deletes
|
||||
"litellm.proxy.prompts.prompt_registry.IN_MEMORY_PROMPT_REGISTRY"
|
||||
) as mock_registry:
|
||||
mock_registry.get_prompt_by_id.return_value = PromptSpec(
|
||||
mock_registry.resolve_prompt_spec.return_value = PromptSpec(
|
||||
prompt_id="test_prompt.v2",
|
||||
litellm_params=PromptLiteLLMParams(prompt_id="test_prompt", prompt_integration="dotprompt"),
|
||||
prompt_info=PromptInfo(prompt_type="db"),
|
||||
prompt_info=PromptInfo(prompt_type="db", environment="production"),
|
||||
)
|
||||
mock_registry.has_config_prompt.return_value = False
|
||||
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client): # test-quality-ok: proxy_server module global is the endpoint's only injection point
|
||||
response = await delete_prompt(
|
||||
|
|
@ -185,7 +168,7 @@ async def test_delete_prompt_environment_scope_reaches_db_and_registry():
|
|||
mock_prisma_client.db.litellm_prompttable.delete_many.assert_called_once_with(
|
||||
where={"prompt_id": "test_prompt", "environment": "production"}
|
||||
)
|
||||
mock_registry.delete_prompts_by_base_id.assert_called_once_with("test_prompt", environment="production")
|
||||
mock_registry.delete_prompts_by_base_id.assert_called_once_with(base_prompt_id="test_prompt", environment="production")
|
||||
assert response == {"message": "Prompt test_prompt deleted successfully from production"}
|
||||
|
||||
|
||||
|
|
@ -218,24 +201,8 @@ async def test_get_prompt_info_by_base_id():
|
|||
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
|
||||
mock_registry.resolve_prompt_spec.return_value = prompt_spec_v3
|
||||
mock_registry.get_prompt_callback_for_prompt.return_value = None
|
||||
|
||||
response = await get_prompt_info(
|
||||
prompt_id="test_prompt", user_api_key_dict=mock_user_auth
|
||||
|
|
@ -284,7 +251,7 @@ async def test_patch_prompt_row_deleted_mid_update_returns_404():
|
|||
"litellm.proxy.prompts.prompt_registry.IN_MEMORY_PROMPT_REGISTRY"
|
||||
) as mock_registry,
|
||||
):
|
||||
mock_registry.get_prompt_by_id.return_value = existing_prompt
|
||||
mock_registry.has_config_prompt.return_value = False
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await patch_prompt(
|
||||
|
|
@ -325,7 +292,7 @@ async def test_patch_prompt_merges_unsent_fields_from_db_row_not_stale_memory():
|
|||
"litellm.proxy.prompts.prompt_registry.IN_MEMORY_PROMPT_REGISTRY"
|
||||
) as mock_registry,
|
||||
):
|
||||
mock_registry.get_prompt_by_id.return_value = stale_in_memory
|
||||
mock_registry.has_config_prompt.return_value = False
|
||||
mock_registry.reload_prompt.side_effect = lambda prompt: prompt
|
||||
|
||||
response = await patch_prompt(
|
||||
|
|
@ -487,7 +454,7 @@ async def test_patch_prompt_info_only_keeps_legacy_keyed_row_patchable():
|
|||
"litellm.proxy.prompts.prompt_registry.IN_MEMORY_PROMPT_REGISTRY"
|
||||
) as mock_registry,
|
||||
):
|
||||
mock_registry.get_prompt_by_id.return_value = existing_prompt
|
||||
mock_registry.has_config_prompt.return_value = False
|
||||
|
||||
await patch_prompt(
|
||||
prompt_id="agent-prompt",
|
||||
|
|
|
|||
|
|
@ -191,11 +191,7 @@ async def test_update_prompt_stores_environment_and_created_by():
|
|||
with patch(
|
||||
"litellm.proxy.prompts.prompt_registry.IN_MEMORY_PROMPT_REGISTRY"
|
||||
) as mock_registry:
|
||||
mock_registry.get_prompt_by_id.return_value = PromptSpec(
|
||||
prompt_id="my_prompt.v1",
|
||||
litellm_params=request.litellm_params,
|
||||
prompt_info=PromptInfo(prompt_type="db"),
|
||||
)
|
||||
mock_registry.has_config_prompt.return_value = False
|
||||
mock_registry.initialize_prompt.return_value = PromptSpec(
|
||||
prompt_id="my_prompt.v2",
|
||||
litellm_params=request.litellm_params,
|
||||
|
|
@ -239,7 +235,8 @@ async def test_delete_prompt_scoped_to_environment():
|
|||
prompt_info=PromptInfo(prompt_type="db"),
|
||||
environment="staging",
|
||||
)
|
||||
mock_registry.get_prompt_by_id.return_value = prompt_spec
|
||||
mock_registry.resolve_prompt_spec.return_value = prompt_spec
|
||||
mock_registry.has_config_prompt.return_value = False
|
||||
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client):
|
||||
await delete_prompt(
|
||||
|
|
@ -251,3 +248,6 @@ async def test_delete_prompt_scoped_to_environment():
|
|||
mock_prisma_client.db.litellm_prompttable.delete_many.assert_called_once_with(
|
||||
where={"prompt_id": "test_prompt", "environment": "staging"}
|
||||
)
|
||||
mock_registry.delete_prompts_by_base_id.assert_called_once_with(
|
||||
base_prompt_id="test_prompt", environment="staging"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,26 +1,35 @@
|
|||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.integrations.custom_prompt_management import CustomPromptManagement
|
||||
from litellm.proxy.prompts.prompt_registry import InMemoryPromptRegistry
|
||||
from litellm.types.prompts.init_prompts import PromptInfo, PromptLiteLLMParams, PromptSpec
|
||||
|
||||
|
||||
def _db_prompt_spec(content: str) -> PromptSpec:
|
||||
def _db_prompt_spec(content: str, environment: str = "development", version: int = 1) -> PromptSpec:
|
||||
return PromptSpec(
|
||||
prompt_id="greeting.v1",
|
||||
prompt_id=f"greeting.v{version}",
|
||||
litellm_params=PromptLiteLLMParams(
|
||||
prompt_id="greeting",
|
||||
prompt_integration="dotprompt",
|
||||
prompt_data={"content": content, "metadata": {}},
|
||||
),
|
||||
prompt_info=PromptInfo(prompt_type="db"),
|
||||
version=version,
|
||||
environment=environment,
|
||||
)
|
||||
|
||||
|
||||
def _served_content(registry: InMemoryPromptRegistry) -> str:
|
||||
callback = registry.get_prompt_callback_by_id("greeting.v1")
|
||||
def _resolved_callback(registry: InMemoryPromptRegistry, environment: str | None = None) -> CustomPromptManagement:
|
||||
spec = registry.resolve_prompt_spec("greeting", environment=environment)
|
||||
assert spec is not None
|
||||
callback = registry.get_prompt_callback_for_prompt(prompt=spec)
|
||||
assert callback is not None
|
||||
return callback.prompt_manager.get_prompt("greeting").content
|
||||
return callback
|
||||
|
||||
|
||||
def _served_content(registry: InMemoryPromptRegistry, environment: str | None = None) -> str:
|
||||
return _resolved_callback(registry, environment=environment).prompt_manager.get_prompt("greeting").content
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
|
|
@ -32,32 +41,34 @@ def isolated_callbacks(monkeypatch: pytest.MonkeyPatch) -> list:
|
|||
def test_sync_prompt_from_db_reloads_row_edited_elsewhere(isolated_callbacks: list) -> None:
|
||||
registry = InMemoryPromptRegistry()
|
||||
registry.sync_prompt_from_db(prompt=_db_prompt_spec("begin every reply with AHOY"))
|
||||
stale_callback = registry.get_prompt_callback_by_id("greeting.v1")
|
||||
stale_callback = _resolved_callback(registry)
|
||||
assert _served_content(registry) == "begin every reply with AHOY"
|
||||
|
||||
registry.sync_prompt_from_db(prompt=_db_prompt_spec("begin every reply with HOWDY"))
|
||||
|
||||
assert _served_content(registry) == "begin every reply with HOWDY"
|
||||
assert registry.get_prompt_by_id("greeting.v1").litellm_params.prompt_data["content"] == "begin every reply with HOWDY"
|
||||
reloaded_spec = registry.resolve_prompt_spec("greeting", environment="development")
|
||||
assert reloaded_spec is not None
|
||||
assert reloaded_spec.litellm_params.prompt_data["content"] == "begin every reply with HOWDY"
|
||||
assert stale_callback not in isolated_callbacks
|
||||
assert isolated_callbacks == [registry.get_prompt_callback_by_id("greeting.v1")]
|
||||
assert isolated_callbacks == [_resolved_callback(registry)]
|
||||
|
||||
|
||||
def test_sync_prompt_from_db_keeps_unchanged_row_in_place(isolated_callbacks: list) -> None:
|
||||
registry = InMemoryPromptRegistry()
|
||||
registry.sync_prompt_from_db(prompt=_db_prompt_spec("begin every reply with AHOY"))
|
||||
first_callback = registry.get_prompt_callback_by_id("greeting.v1")
|
||||
first_callback = _resolved_callback(registry)
|
||||
|
||||
registry.sync_prompt_from_db(prompt=_db_prompt_spec("begin every reply with AHOY"))
|
||||
|
||||
assert registry.get_prompt_callback_by_id("greeting.v1") is first_callback
|
||||
assert _resolved_callback(registry) is first_callback
|
||||
assert isolated_callbacks == [first_callback]
|
||||
|
||||
|
||||
def test_reload_prompt_replaces_callback_without_leaking_the_old_one(isolated_callbacks: list) -> None:
|
||||
registry = InMemoryPromptRegistry()
|
||||
registry.initialize_prompt(prompt=_db_prompt_spec("begin every reply with AHOY"))
|
||||
stale_callback = registry.get_prompt_callback_by_id("greeting.v1")
|
||||
stale_callback = _resolved_callback(registry)
|
||||
|
||||
reloaded = registry.reload_prompt(prompt=_db_prompt_spec("begin every reply with HOWDY"))
|
||||
|
||||
|
|
@ -70,7 +81,7 @@ def test_reload_prompt_replaces_callback_without_leaking_the_old_one(isolated_ca
|
|||
def test_reload_prompt_keeps_the_old_template_when_the_replacement_fails(isolated_callbacks: list) -> None:
|
||||
registry = InMemoryPromptRegistry()
|
||||
registry.initialize_prompt(prompt=_db_prompt_spec("begin every reply with AHOY"))
|
||||
old_callback = registry.get_prompt_callback_by_id("greeting.v1")
|
||||
old_callback = _resolved_callback(registry)
|
||||
|
||||
broken = PromptSpec(
|
||||
prompt_id="greeting.v1",
|
||||
|
|
@ -80,63 +91,126 @@ def test_reload_prompt_keeps_the_old_template_when_the_replacement_fails(isolate
|
|||
prompt_data={"content": "begin every reply with HOWDY", "metadata": {}},
|
||||
),
|
||||
prompt_info=PromptInfo(prompt_type="db"),
|
||||
version=1,
|
||||
environment="development",
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="Unsupported prompt"):
|
||||
registry.reload_prompt(prompt=broken)
|
||||
|
||||
assert registry.get_prompt_callback_by_id("greeting.v1") is old_callback
|
||||
assert _resolved_callback(registry) is old_callback
|
||||
assert _served_content(registry) == "begin every reply with AHOY"
|
||||
assert isolated_callbacks == [old_callback]
|
||||
|
||||
|
||||
def _versioned_prompt_spec(version: int, environment: str) -> PromptSpec:
|
||||
return PromptSpec(
|
||||
prompt_id=f"greeting.v{version}",
|
||||
def test_environments_sharing_a_prompt_id_keep_separate_templates(isolated_callbacks: list) -> None:
|
||||
registry = InMemoryPromptRegistry()
|
||||
registry.sync_prompt_from_db(prompt=_db_prompt_spec("begin every reply with AHOY", environment="development"))
|
||||
registry.sync_prompt_from_db(prompt=_db_prompt_spec("begin every reply with HOWDY", environment="production"))
|
||||
|
||||
assert _served_content(registry, environment="development") == "begin every reply with AHOY"
|
||||
assert _served_content(registry, environment="production") == "begin every reply with HOWDY"
|
||||
assert _resolved_callback(registry, environment="development") is not _resolved_callback(
|
||||
registry, environment="production"
|
||||
)
|
||||
|
||||
|
||||
def test_default_resolution_prefers_production(isolated_callbacks: list) -> None:
|
||||
registry = InMemoryPromptRegistry()
|
||||
registry.sync_prompt_from_db(prompt=_db_prompt_spec("begin every reply with AHOY", environment="development"))
|
||||
registry.sync_prompt_from_db(prompt=_db_prompt_spec("begin every reply with HOWDY", environment="production"))
|
||||
|
||||
assert _served_content(registry) == "begin every reply with HOWDY"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("environment", ["staging", "qa"])
|
||||
def test_default_resolution_serves_the_only_environment_present(isolated_callbacks: list, environment: str) -> None:
|
||||
registry = InMemoryPromptRegistry()
|
||||
registry.sync_prompt_from_db(prompt=_db_prompt_spec("begin every reply with AHOY", environment=environment))
|
||||
|
||||
assert _served_content(registry) == "begin every reply with AHOY"
|
||||
|
||||
|
||||
def test_resolution_picks_exact_version_and_latest_within_an_environment(isolated_callbacks: list) -> None:
|
||||
registry = InMemoryPromptRegistry()
|
||||
registry.sync_prompt_from_db(prompt=_db_prompt_spec("begin every reply with AHOY", environment="development", version=1))
|
||||
registry.sync_prompt_from_db(prompt=_db_prompt_spec("begin every reply with YO", environment="development", version=2))
|
||||
registry.sync_prompt_from_db(prompt=_db_prompt_spec("begin every reply with HOWDY", environment="production", version=1))
|
||||
|
||||
exact = registry.resolve_prompt_spec("greeting", version=1, environment="development")
|
||||
assert exact is not None
|
||||
assert exact.litellm_params.prompt_data["content"] == "begin every reply with AHOY"
|
||||
|
||||
latest = registry.resolve_prompt_spec("greeting", environment="development")
|
||||
assert latest is not None
|
||||
assert latest.litellm_params.prompt_data["content"] == "begin every reply with YO"
|
||||
|
||||
assert registry.resolve_prompt_spec("greeting", version=3, environment="development") is None
|
||||
|
||||
|
||||
def test_resolution_returns_none_for_unknown_environment_or_prompt(isolated_callbacks: list) -> None:
|
||||
registry = InMemoryPromptRegistry()
|
||||
registry.sync_prompt_from_db(prompt=_db_prompt_spec("begin every reply with AHOY", environment="development"))
|
||||
|
||||
assert registry.resolve_prompt_spec("greeting", environment="production") is None
|
||||
assert registry.resolve_prompt_spec("no_such_prompt") is None
|
||||
|
||||
|
||||
def test_delete_prompts_by_base_id_scoped_to_one_environment(isolated_callbacks: list) -> None:
|
||||
registry = InMemoryPromptRegistry()
|
||||
registry.sync_prompt_from_db(prompt=_db_prompt_spec("begin every reply with AHOY", environment="development"))
|
||||
registry.sync_prompt_from_db(prompt=_db_prompt_spec("begin every reply with HOWDY", environment="production"))
|
||||
production_callback = _resolved_callback(registry, environment="production")
|
||||
|
||||
deleted = registry.delete_prompts_by_base_id(base_prompt_id="greeting", environment="development")
|
||||
|
||||
assert deleted == ["greeting.v1::development"]
|
||||
assert registry.resolve_prompt_spec("greeting", environment="development") is None
|
||||
assert _resolved_callback(registry, environment="production") is production_callback
|
||||
assert _served_content(registry, environment="production") == "begin every reply with HOWDY"
|
||||
|
||||
deleted_rest = registry.delete_prompts_by_base_id(base_prompt_id="greeting")
|
||||
|
||||
assert deleted_rest == ["greeting.v1::production"]
|
||||
assert registry.resolve_prompt_spec("greeting") is None
|
||||
|
||||
|
||||
def test_has_config_prompt_matches_any_version_of_the_base_id(isolated_callbacks: list) -> None:
|
||||
registry = InMemoryPromptRegistry()
|
||||
config_spec = PromptSpec(
|
||||
prompt_id="greeting",
|
||||
litellm_params=PromptLiteLLMParams(
|
||||
prompt_id="greeting",
|
||||
prompt_integration="dotprompt",
|
||||
prompt_data={"content": f"begin every reply with AHOY v{version}", "metadata": {}},
|
||||
prompt_data={"content": "begin every reply with AHOY", "metadata": {}},
|
||||
),
|
||||
prompt_info=PromptInfo(prompt_type="db", environment=environment),
|
||||
version=version,
|
||||
environment=environment,
|
||||
prompt_info=PromptInfo(prompt_type="config"),
|
||||
)
|
||||
registry.initialize_prompt(prompt=config_spec)
|
||||
registry.sync_prompt_from_db(prompt=_db_prompt_spec("begin every reply with HOWDY", environment="production"))
|
||||
|
||||
assert registry.has_config_prompt(base_prompt_id="greeting") is True
|
||||
assert registry.has_config_prompt(base_prompt_id="other_prompt") is False
|
||||
|
||||
|
||||
def test_delete_prompts_by_base_id_removes_the_callbacks_from_litellm_callbacks(isolated_callbacks: list) -> None:
|
||||
registry = InMemoryPromptRegistry()
|
||||
registry.initialize_prompt(prompt=_versioned_prompt_spec(1, "development"))
|
||||
registry.initialize_prompt(prompt=_versioned_prompt_spec(2, "development"))
|
||||
registry.initialize_prompt(prompt=_db_prompt_spec("begin every reply with AHOY", version=1))
|
||||
registry.initialize_prompt(prompt=_db_prompt_spec("begin every reply with YO", version=2))
|
||||
assert len(isolated_callbacks) == 1
|
||||
|
||||
deleted = registry.delete_prompts_by_base_id("greeting")
|
||||
deleted = registry.delete_prompts_by_base_id(base_prompt_id="greeting")
|
||||
|
||||
assert sorted(deleted) == ["greeting.v1", "greeting.v2"]
|
||||
assert registry.get_prompt_by_id("greeting.v1") is None
|
||||
assert registry.get_prompt_callback_by_id("greeting.v2") is None
|
||||
assert sorted(deleted) == ["greeting.v1::development", "greeting.v2::development"]
|
||||
assert registry.resolve_prompt_spec("greeting") is None
|
||||
assert isolated_callbacks == []
|
||||
|
||||
|
||||
def test_delete_prompts_by_base_id_environment_scope_keeps_other_environments(isolated_callbacks: list) -> None:
|
||||
def test_remove_prompt_is_a_no_op_for_an_unknown_registry_key(isolated_callbacks: list) -> None:
|
||||
registry = InMemoryPromptRegistry()
|
||||
registry.initialize_prompt(prompt=_versioned_prompt_spec(1, "development"))
|
||||
registry.initialize_prompt(prompt=_versioned_prompt_spec(2, "production"))
|
||||
production_callback = registry.get_prompt_callback_by_id("greeting.v2")
|
||||
registry.initialize_prompt(prompt=_db_prompt_spec("begin every reply with AHOY"))
|
||||
|
||||
deleted = registry.delete_prompts_by_base_id("greeting", environment="development")
|
||||
registry.remove_prompt(registry_key="not_there.v1::development")
|
||||
|
||||
assert deleted == ["greeting.v1"]
|
||||
assert registry.get_prompt_by_id("greeting.v1") is None
|
||||
assert registry.get_prompt_by_id("greeting.v2") is not None
|
||||
assert registry.get_prompt_callback_by_id("greeting.v2") is production_callback
|
||||
|
||||
|
||||
def test_remove_prompt_is_a_no_op_for_an_unknown_id(isolated_callbacks: list) -> None:
|
||||
registry = InMemoryPromptRegistry()
|
||||
registry.initialize_prompt(prompt=_versioned_prompt_spec(1, "development"))
|
||||
|
||||
registry.remove_prompt(prompt_id="not_there.v1")
|
||||
|
||||
assert registry.get_prompt_by_id("greeting.v1") is not None
|
||||
assert registry.resolve_prompt_spec("greeting") is not None
|
||||
assert len(isolated_callbacks) == 1
|
||||
|
|
|
|||
|
|
@ -11379,10 +11379,15 @@ async def test_init_prompts_in_db_reloads_rows_patched_on_another_worker(monkeyp
|
|||
}
|
||||
return row
|
||||
|
||||
def served_content() -> str:
|
||||
callback = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id("greeting_sync.v1")
|
||||
def served_callback():
|
||||
spec = IN_MEMORY_PROMPT_REGISTRY.resolve_prompt_spec("greeting_sync")
|
||||
assert spec is not None
|
||||
callback = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_for_prompt(prompt=spec)
|
||||
assert callback is not None
|
||||
return callback.prompt_manager.get_prompt("greeting_sync").content
|
||||
return callback
|
||||
|
||||
def served_content() -> str:
|
||||
return served_callback().prompt_manager.get_prompt("greeting_sync").content
|
||||
|
||||
prisma_client = MagicMock()
|
||||
try:
|
||||
|
|
@ -11394,7 +11399,7 @@ async def test_init_prompts_in_db_reloads_rows_patched_on_another_worker(monkeyp
|
|||
await ProxyConfig()._init_prompts_in_db(prisma_client=prisma_client)
|
||||
|
||||
assert served_content() == "Begin every reply with HOWDY"
|
||||
assert litellm.callbacks == [IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id("greeting_sync.v1")]
|
||||
assert litellm.callbacks == [served_callback()]
|
||||
finally:
|
||||
IN_MEMORY_PROMPT_REGISTRY.delete_prompts_by_base_id("greeting_sync")
|
||||
|
||||
|
|
@ -11433,22 +11438,25 @@ async def test_init_prompts_in_db_syncs_remaining_rows_when_one_row_fails(monkey
|
|||
)
|
||||
await ProxyConfig()._init_prompts_in_db(prisma_client=prisma_client)
|
||||
|
||||
assert IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id("broken_sync.v1") is None
|
||||
assert IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id("healthy_sync.v1") is not None
|
||||
assert litellm.callbacks == [IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id("healthy_sync.v1")]
|
||||
assert IN_MEMORY_PROMPT_REGISTRY.resolve_prompt_spec("broken_sync") is None
|
||||
healthy_spec = IN_MEMORY_PROMPT_REGISTRY.resolve_prompt_spec("healthy_sync")
|
||||
assert healthy_spec is not None
|
||||
healthy_callback = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_for_prompt(prompt=healthy_spec)
|
||||
assert healthy_callback is not None
|
||||
assert litellm.callbacks == [healthy_callback]
|
||||
finally:
|
||||
IN_MEMORY_PROMPT_REGISTRY.delete_prompts_by_base_id("healthy_sync")
|
||||
IN_MEMORY_PROMPT_REGISTRY.delete_prompts_by_base_id("broken_sync")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_init_prompts_in_db_serves_the_newest_row_when_environments_collide_on_a_versioned_id(monkeypatch):
|
||||
async def test_init_prompts_in_db_syncs_every_environment_sharing_a_versioned_id(monkeypatch):
|
||||
from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY
|
||||
from litellm.proxy.proxy_server import ProxyConfig
|
||||
|
||||
monkeypatch.setattr(litellm, "callbacks", [])
|
||||
|
||||
def db_row(environment: str, content: str, updated_at: datetime) -> MagicMock:
|
||||
def db_row(environment: str, content: str) -> MagicMock:
|
||||
row = MagicMock()
|
||||
row.model_dump.return_value = {
|
||||
"prompt_id": "greeting_env",
|
||||
|
|
@ -11464,30 +11472,41 @@ async def test_init_prompts_in_db_serves_the_newest_row_when_environments_collid
|
|||
),
|
||||
"prompt_info": json.dumps({"prompt_type": "db"}),
|
||||
"created_at": None,
|
||||
"updated_at": updated_at,
|
||||
"updated_at": None,
|
||||
}
|
||||
return row
|
||||
|
||||
freshly_patched = db_row(
|
||||
"production", "Begin every reply with HOWDY", datetime(2026, 8, 26, 12, 0, tzinfo=timezone.utc)
|
||||
)
|
||||
stale_sibling = db_row(
|
||||
"development", "Begin every reply with AHOY", datetime(2026, 8, 26, 11, 0, tzinfo=timezone.utc)
|
||||
)
|
||||
def served_content(environment: str | None) -> str:
|
||||
spec = IN_MEMORY_PROMPT_REGISTRY.resolve_prompt_spec("greeting_env", environment=environment)
|
||||
assert spec is not None
|
||||
callback = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_for_prompt(prompt=spec)
|
||||
assert callback is not None
|
||||
return callback.prompt_manager.get_prompt("greeting_env").content
|
||||
|
||||
prisma_client = MagicMock()
|
||||
try:
|
||||
prisma_client.db.litellm_prompttable.find_many = AsyncMock(return_value=[freshly_patched, stale_sibling])
|
||||
prisma_client.db.litellm_prompttable.find_many = AsyncMock(
|
||||
return_value=[
|
||||
db_row("development", "Begin every reply with AHOY"),
|
||||
db_row("production", "Begin every reply with HOWDY"),
|
||||
]
|
||||
)
|
||||
await ProxyConfig()._init_prompts_in_db(prisma_client=prisma_client)
|
||||
|
||||
first_callback = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id("greeting_env.v1")
|
||||
assert first_callback is not None
|
||||
assert first_callback.prompt_manager.get_prompt("greeting_env").content == "Begin every reply with HOWDY"
|
||||
assert served_content("development") == "Begin every reply with AHOY"
|
||||
assert served_content("production") == "Begin every reply with HOWDY"
|
||||
assert served_content(None) == "Begin every reply with HOWDY"
|
||||
|
||||
prisma_client.db.litellm_prompttable.find_many = AsyncMock(
|
||||
return_value=[
|
||||
db_row("development", "Begin every reply with YO"),
|
||||
db_row("production", "Begin every reply with HOWDY"),
|
||||
]
|
||||
)
|
||||
await ProxyConfig()._init_prompts_in_db(prisma_client=prisma_client)
|
||||
|
||||
assert IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id("greeting_env.v1") is first_callback
|
||||
assert litellm.callbacks == [first_callback]
|
||||
assert served_content("development") == "Begin every reply with YO"
|
||||
assert served_content("production") == "Begin every reply with HOWDY"
|
||||
finally:
|
||||
IN_MEMORY_PROMPT_REGISTRY.delete_prompts_by_base_id("greeting_env")
|
||||
|
||||
|
|
@ -11530,13 +11549,14 @@ async def test_init_prompts_in_db_unloads_rows_deleted_on_another_worker(monkeyp
|
|||
return_value=[_prompt_db_row("greeting_del", _dotprompt_params("greeting_del"))]
|
||||
)
|
||||
await ProxyConfig()._init_prompts_in_db(prisma_client=prisma_client)
|
||||
assert IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id("greeting_del.v1") is not None
|
||||
loaded_spec = IN_MEMORY_PROMPT_REGISTRY.resolve_prompt_spec("greeting_del")
|
||||
assert loaded_spec is not None
|
||||
assert IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_for_prompt(prompt=loaded_spec) is not None
|
||||
|
||||
prisma_client.db.litellm_prompttable.find_many = AsyncMock(return_value=[])
|
||||
await ProxyConfig()._init_prompts_in_db(prisma_client=prisma_client)
|
||||
|
||||
assert IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id("greeting_del.v1") is None
|
||||
assert IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id("greeting_del.v1") is None
|
||||
assert IN_MEMORY_PROMPT_REGISTRY.resolve_prompt_spec("greeting_del") is None
|
||||
assert litellm.callbacks == []
|
||||
finally:
|
||||
IN_MEMORY_PROMPT_REGISTRY.delete_prompts_by_base_id("greeting_del")
|
||||
|
|
@ -11567,10 +11587,12 @@ async def test_init_prompts_in_db_keeps_config_prompts_when_their_id_has_no_db_r
|
|||
|
||||
await ProxyConfig()._init_prompts_in_db(prisma_client=prisma_client)
|
||||
|
||||
assert IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id("greeting_cfg") is not None
|
||||
surviving_spec = IN_MEMORY_PROMPT_REGISTRY.resolve_prompt_spec("greeting_cfg")
|
||||
assert surviving_spec is not None
|
||||
assert IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_for_prompt(prompt=surviving_spec) is not None
|
||||
assert len(litellm.callbacks) == 1
|
||||
finally:
|
||||
IN_MEMORY_PROMPT_REGISTRY.remove_prompt(prompt_id="greeting_cfg")
|
||||
IN_MEMORY_PROMPT_REGISTRY.delete_prompts_by_base_id("greeting_cfg")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -11586,7 +11608,9 @@ async def test_init_prompts_in_db_keeps_the_in_memory_copy_when_a_row_fails_to_p
|
|||
return_value=[_prompt_db_row("greeting_broken", _dotprompt_params("greeting_broken"))]
|
||||
)
|
||||
await ProxyConfig()._init_prompts_in_db(prisma_client=prisma_client)
|
||||
loaded_callback = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id("greeting_broken.v1")
|
||||
loaded_spec = IN_MEMORY_PROMPT_REGISTRY.resolve_prompt_spec("greeting_broken")
|
||||
assert loaded_spec is not None
|
||||
loaded_callback = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_for_prompt(prompt=loaded_spec)
|
||||
assert loaded_callback is not None
|
||||
|
||||
prisma_client.db.litellm_prompttable.find_many = AsyncMock(
|
||||
|
|
@ -11594,7 +11618,9 @@ async def test_init_prompts_in_db_keeps_the_in_memory_copy_when_a_row_fails_to_p
|
|||
)
|
||||
await ProxyConfig()._init_prompts_in_db(prisma_client=prisma_client)
|
||||
|
||||
assert IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id("greeting_broken.v1") is loaded_callback
|
||||
kept_spec = IN_MEMORY_PROMPT_REGISTRY.resolve_prompt_spec("greeting_broken")
|
||||
assert kept_spec is not None
|
||||
assert IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_for_prompt(prompt=kept_spec) is loaded_callback
|
||||
assert litellm.callbacks == [loaded_callback]
|
||||
finally:
|
||||
IN_MEMORY_PROMPT_REGISTRY.delete_prompts_by_base_id("greeting_broken")
|
||||
|
|
@ -11628,8 +11654,9 @@ async def test_init_prompts_in_db_keeps_a_prompt_created_while_the_sync_was_read
|
|||
prisma_client.db.litellm_prompttable.find_many = AsyncMock(side_effect=create_prompt_behind_the_select)
|
||||
await ProxyConfig()._init_prompts_in_db(prisma_client=prisma_client)
|
||||
|
||||
surviving_callback = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id("greeting_race.v1")
|
||||
assert IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id("greeting_race.v1") is not None
|
||||
surviving_spec = IN_MEMORY_PROMPT_REGISTRY.resolve_prompt_spec("greeting_race")
|
||||
assert surviving_spec is not None
|
||||
surviving_callback = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_for_prompt(prompt=surviving_spec)
|
||||
assert surviving_callback is not None
|
||||
assert litellm.callbacks == [surviving_callback]
|
||||
finally:
|
||||
|
|
|
|||
|
|
@ -726,10 +726,7 @@ async def test_process_prompt_template_no_op_when_no_prompt_spec(proxy_logging,
|
|||
from litellm.proxy.prompts import prompt_registry
|
||||
|
||||
monkeypatch.setattr(
|
||||
prompt_registry.IN_MEMORY_PROMPT_REGISTRY, "get_prompt_callback_by_id", lambda *a, **kw: None
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
prompt_registry.IN_MEMORY_PROMPT_REGISTRY, "get_prompt_by_id", lambda *a, **kw: None
|
||||
prompt_registry.IN_MEMORY_PROMPT_REGISTRY, "resolve_prompt_spec", lambda *a, **kw: None
|
||||
)
|
||||
data: Dict[str, Any] = {"messages": [{"role": "user"}], "model": "m", "temperature": 0.1}
|
||||
await proxy_logging._process_prompt_template(
|
||||
|
|
@ -752,11 +749,11 @@ async def test_process_prompt_template_applies_when_spec_resolves(proxy_logging,
|
|||
|
||||
monkeypatch.setattr(
|
||||
prompt_registry.IN_MEMORY_PROMPT_REGISTRY,
|
||||
"get_prompt_callback_by_id",
|
||||
"get_prompt_callback_for_prompt",
|
||||
lambda *a, **kw: custom_logger,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
prompt_registry.IN_MEMORY_PROMPT_REGISTRY, "get_prompt_by_id", lambda *a, **kw: prompt_spec
|
||||
prompt_registry.IN_MEMORY_PROMPT_REGISTRY, "resolve_prompt_spec", lambda *a, **kw: prompt_spec
|
||||
)
|
||||
|
||||
logging_obj = MagicMock()
|
||||
|
|
@ -802,11 +799,11 @@ async def test_process_prompt_template_async_get_prompt_error_raises(proxy_loggi
|
|||
prompt_spec.litellm_params = MagicMock(prompt_id="x")
|
||||
monkeypatch.setattr(
|
||||
prompt_registry.IN_MEMORY_PROMPT_REGISTRY,
|
||||
"get_prompt_callback_by_id",
|
||||
"get_prompt_callback_for_prompt",
|
||||
lambda *a, **kw: custom_logger,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
prompt_registry.IN_MEMORY_PROMPT_REGISTRY, "get_prompt_by_id", lambda *a, **kw: prompt_spec
|
||||
prompt_registry.IN_MEMORY_PROMPT_REGISTRY, "resolve_prompt_spec", lambda *a, **kw: prompt_spec
|
||||
)
|
||||
logging_obj = MagicMock()
|
||||
logging_obj.async_get_chat_completion_prompt = AsyncMock(side_effect=RuntimeError("bad prompt"))
|
||||
|
|
@ -820,6 +817,47 @@ async def test_process_prompt_template_async_get_prompt_error_raises(proxy_loggi
|
|||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_prompt_template_resolves_the_requested_environment(proxy_logging, monkeypatch):
|
||||
from litellm.proxy.prompts import prompt_registry
|
||||
|
||||
prompt_spec = MagicMock()
|
||||
prompt_spec.litellm_params = MagicMock(prompt_id="greeting")
|
||||
resolve_calls: list[dict] = []
|
||||
|
||||
def fake_resolve(prompt_id, version=None, environment=None):
|
||||
resolve_calls.append({"prompt_id": prompt_id, "version": version, "environment": environment})
|
||||
return prompt_spec
|
||||
|
||||
monkeypatch.setattr(prompt_registry.IN_MEMORY_PROMPT_REGISTRY, "resolve_prompt_spec", fake_resolve)
|
||||
monkeypatch.setattr(
|
||||
prompt_registry.IN_MEMORY_PROMPT_REGISTRY, "get_prompt_callback_for_prompt", lambda *a, **kw: MagicMock()
|
||||
)
|
||||
logging_obj = MagicMock()
|
||||
logging_obj.async_get_chat_completion_prompt = AsyncMock(
|
||||
return_value=("m", [{"role": "user", "content": "rendered"}], {})
|
||||
)
|
||||
data: Dict[str, Any] = {
|
||||
"messages": [{"role": "user", "content": "orig"}],
|
||||
"model": "m",
|
||||
"prompt_id": "greeting",
|
||||
"prompt_version": 1,
|
||||
"prompt_environment": "development",
|
||||
}
|
||||
await proxy_logging._process_prompt_template(
|
||||
data=data,
|
||||
litellm_logging_obj=logging_obj,
|
||||
prompt_id="greeting",
|
||||
prompt_version=1,
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
assert resolve_calls == [{"prompt_id": "greeting", "version": 1, "environment": "development"}]
|
||||
assert "prompt_environment" not in data
|
||||
assert "prompt_id" not in data
|
||||
assert data["messages"] == [{"role": "user", "content": "rendered"}]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_prompt_template_aresponses_swaps_model_and_merges_input(proxy_logging, monkeypatch):
|
||||
from litellm.proxy.prompts import prompt_registry
|
||||
|
|
@ -829,11 +867,11 @@ async def test_process_prompt_template_aresponses_swaps_model_and_merges_input(p
|
|||
prompt_spec.litellm_params = MagicMock(prompt_id="resolved-id")
|
||||
monkeypatch.setattr(
|
||||
prompt_registry.IN_MEMORY_PROMPT_REGISTRY,
|
||||
"get_prompt_callback_by_id",
|
||||
"get_prompt_callback_for_prompt",
|
||||
lambda *a, **kw: custom_logger,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
prompt_registry.IN_MEMORY_PROMPT_REGISTRY, "get_prompt_by_id", lambda *a, **kw: prompt_spec
|
||||
prompt_registry.IN_MEMORY_PROMPT_REGISTRY, "resolve_prompt_spec", lambda *a, **kw: prompt_spec
|
||||
)
|
||||
|
||||
logging_obj = MagicMock()
|
||||
|
|
|
|||
|
|
@ -65,11 +65,13 @@ describe("PromptTable", () => {
|
|||
expect(within(rows[1]).getByText("prompt-older")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should call onPromptClick when the prompt ID is clicked", async () => {
|
||||
it("should call onPromptClick with the row's environment, defaulting to development", async () => {
|
||||
const user = userEvent.setup();
|
||||
render(<PromptTable {...defaultProps} />);
|
||||
await user.click(screen.getByRole("button", { name: "prompt-newer" }));
|
||||
expect(mockOnPromptClick).toHaveBeenCalledWith("prompt-newer");
|
||||
expect(mockOnPromptClick).toHaveBeenCalledWith("prompt-newer", "production");
|
||||
await user.click(screen.getByRole("button", { name: "prompt-older" }));
|
||||
expect(mockOnPromptClick).toHaveBeenCalledWith("prompt-older", "development");
|
||||
});
|
||||
|
||||
it("should label the environment and default missing environments to development", () => {
|
||||
|
|
@ -83,7 +85,7 @@ describe("PromptTable", () => {
|
|||
render(<PromptTable {...defaultProps} />);
|
||||
await user.click(screen.getByTestId("prompt-actions-prompt-newer"));
|
||||
await user.click(await screen.findByTestId("prompt-action-delete"));
|
||||
expect(mockOnDeleteClick).toHaveBeenCalledWith("prompt-newer", "prompt-newer");
|
||||
expect(mockOnDeleteClick).toHaveBeenCalledWith("prompt-newer", "prompt-newer", "production");
|
||||
});
|
||||
|
||||
it("should copy the prompt ID through the actions menu", async () => {
|
||||
|
|
|
|||
|
|
@ -13,8 +13,8 @@ import { ModelGroupInfo } from "./prompt_utils";
|
|||
interface PromptTableProps {
|
||||
promptsList: PromptSpec[];
|
||||
isLoading: boolean;
|
||||
onPromptClick?: (id: string) => void;
|
||||
onDeleteClick?: (id: string, name: string) => void;
|
||||
onPromptClick?: (id: string, environment: string) => void;
|
||||
onDeleteClick?: (id: string, name: string, environment: string) => void;
|
||||
accessToken: string | null;
|
||||
isAdmin: boolean;
|
||||
}
|
||||
|
|
@ -74,7 +74,9 @@ const PromptTable: React.FC<PromptTableProps> = ({
|
|||
<DataTable
|
||||
data={promptsList}
|
||||
columns={columns}
|
||||
getRowId={(prompt, index) => prompt.prompt_id || String(index)}
|
||||
getRowId={(prompt, index) =>
|
||||
prompt.prompt_id ? `${prompt.prompt_id}::${prompt.environment || "development"}` : String(index)
|
||||
}
|
||||
sortingMode="client"
|
||||
sorting={sorting}
|
||||
onSortingChange={setSorting}
|
||||
|
|
|
|||
|
|
@ -64,7 +64,7 @@ function PromptModelCell({ prompt, modelHubData }: { prompt: PromptSpec; modelHu
|
|||
interface PromptRowActionsProps {
|
||||
prompt: PromptSpec;
|
||||
isAdmin: boolean;
|
||||
onDeleteClick?: (id: string, name: string) => void;
|
||||
onDeleteClick?: (id: string, name: string, environment: string) => void;
|
||||
}
|
||||
|
||||
function PromptRowActions({ prompt, isAdmin, onDeleteClick }: PromptRowActionsProps) {
|
||||
|
|
@ -91,7 +91,13 @@ function PromptRowActions({ prompt, isAdmin, onDeleteClick }: PromptRowActionsPr
|
|||
<DropdownMenuItem
|
||||
variant="destructive"
|
||||
data-testid="prompt-action-delete"
|
||||
onClick={() => onDeleteClick?.(prompt.prompt_id, prompt.prompt_id || "Unknown Prompt")}
|
||||
onClick={() =>
|
||||
onDeleteClick?.(
|
||||
prompt.prompt_id,
|
||||
prompt.prompt_id || "Unknown Prompt",
|
||||
prompt.environment || "development",
|
||||
)
|
||||
}
|
||||
>
|
||||
<Trash2 />
|
||||
Delete
|
||||
|
|
@ -106,8 +112,8 @@ function PromptRowActions({ prompt, isAdmin, onDeleteClick }: PromptRowActionsPr
|
|||
interface PromptTableColumnsDeps {
|
||||
modelHubData: Map<string, ModelGroupInfo>;
|
||||
isAdmin: boolean;
|
||||
onPromptClick?: (id: string) => void;
|
||||
onDeleteClick?: (id: string, name: string) => void;
|
||||
onPromptClick?: (id: string, environment: string) => void;
|
||||
onDeleteClick?: (id: string, name: string, environment: string) => void;
|
||||
}
|
||||
|
||||
export const getPromptTableColumns = ({
|
||||
|
|
@ -128,7 +134,11 @@ export const getPromptTableColumns = ({
|
|||
title={row.original.prompt_id}
|
||||
titleClassName="font-mono text-xs font-normal"
|
||||
className="max-w-60"
|
||||
onClick={onPromptClick ? () => onPromptClick(row.original.prompt_id) : undefined}
|
||||
onClick={
|
||||
onPromptClick
|
||||
? () => onPromptClick(row.original.prompt_id, row.original.environment || "development")
|
||||
: undefined
|
||||
}
|
||||
/>
|
||||
),
|
||||
},
|
||||
|
|
|
|||
|
|
@ -15,21 +15,31 @@ vi.mock("./PromptTable", () => ({
|
|||
__esModule: true,
|
||||
default: ({
|
||||
isLoading,
|
||||
onPromptClick,
|
||||
onDeleteClick,
|
||||
}: {
|
||||
isLoading: boolean;
|
||||
onDeleteClick: (id: string, name: string) => void;
|
||||
onPromptClick: (id: string, environment: string) => void;
|
||||
onDeleteClick: (id: string, name: string, environment: string) => void;
|
||||
}) => (
|
||||
<div data-testid="prompt-table">
|
||||
{isLoading ? "table-loading" : "table-loaded"}
|
||||
<button type="button" onClick={() => onDeleteClick("prompt-1", "my-prompt")}>
|
||||
<button type="button" onClick={() => onPromptClick("prompt-1", "staging")}>
|
||||
row-open
|
||||
</button>
|
||||
<button type="button" onClick={() => onDeleteClick("prompt-1", "my-prompt", "staging")}>
|
||||
row-delete
|
||||
</button>
|
||||
</div>
|
||||
),
|
||||
}));
|
||||
|
||||
vi.mock("./prompt_info", () => ({ __esModule: true, default: () => <div>prompt-info-view</div> }));
|
||||
vi.mock("./prompt_info", () => ({
|
||||
__esModule: true,
|
||||
default: ({ initialEnvironment }: { initialEnvironment?: string }) => (
|
||||
<div>prompt-info-view:{initialEnvironment ?? "none"}</div>
|
||||
),
|
||||
}));
|
||||
vi.mock("./add_prompt_form", () => ({
|
||||
__esModule: true,
|
||||
default: ({ visible }: { visible: boolean }) => (visible ? <div>add-prompt-form</div> : null),
|
||||
|
|
@ -143,6 +153,22 @@ describe("PromptsPanel toolbar", () => {
|
|||
});
|
||||
});
|
||||
|
||||
describe("PromptsPanel row navigation", () => {
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks();
|
||||
mockGetPromptsList.mockResolvedValue({ prompts: [] } as never);
|
||||
});
|
||||
|
||||
it("should open the info view preselected to the clicked row's environment", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderPanel("Admin");
|
||||
|
||||
await user.click(await screen.findByRole("button", { name: "row-open" }));
|
||||
|
||||
expect(screen.getByText("prompt-info-view:staging")).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
||||
describe("PromptsPanel delete confirmation", () => {
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks();
|
||||
|
|
@ -156,13 +182,13 @@ describe("PromptsPanel delete confirmation", () => {
|
|||
|
||||
await user.click(await screen.findByRole("button", { name: "row-delete" }));
|
||||
|
||||
expect(await screen.findByText(/delete prompt: my-prompt/i)).toBeInTheDocument();
|
||||
expect(await screen.findByText(/the staging copy of prompt: my-prompt/i)).toBeInTheDocument();
|
||||
expect(screen.getByText(/cannot be undone/i)).toBeInTheDocument();
|
||||
expect(mockDeletePromptCall).not.toHaveBeenCalled();
|
||||
|
||||
await user.click(screen.getByRole("button", { name: /^delete$/i }));
|
||||
|
||||
await waitFor(() => expect(mockDeletePromptCall).toHaveBeenCalledWith("sk-test", "prompt-1"));
|
||||
await waitFor(() => expect(mockDeletePromptCall).toHaveBeenCalledWith("sk-test", "prompt-1", "staging"));
|
||||
});
|
||||
|
||||
it("should abandon the delete when the confirmation is dismissed", async () => {
|
||||
|
|
@ -170,11 +196,11 @@ describe("PromptsPanel delete confirmation", () => {
|
|||
renderPanel("Admin");
|
||||
|
||||
await user.click(await screen.findByRole("button", { name: "row-delete" }));
|
||||
await screen.findByText(/delete prompt: my-prompt/i);
|
||||
await screen.findByText(/the staging copy of prompt: my-prompt/i);
|
||||
|
||||
await user.click(screen.getByRole("button", { name: /cancel/i }));
|
||||
|
||||
await waitFor(() => expect(screen.queryByText(/delete prompt: my-prompt/i)).not.toBeInTheDocument());
|
||||
await waitFor(() => expect(screen.queryByText(/the staging copy of prompt: my-prompt/i)).not.toBeInTheDocument());
|
||||
expect(mockDeletePromptCall).not.toHaveBeenCalled();
|
||||
});
|
||||
|
||||
|
|
@ -189,14 +215,14 @@ describe("PromptsPanel delete confirmation", () => {
|
|||
renderPanel("Admin");
|
||||
|
||||
await user.click(await screen.findByRole("button", { name: "row-delete" }));
|
||||
await screen.findByText(/delete prompt: my-prompt/i);
|
||||
await screen.findByText(/the staging copy of prompt: my-prompt/i);
|
||||
await user.click(screen.getByRole("button", { name: /^delete$/i }));
|
||||
await waitFor(() => expect(mockDeletePromptCall).toHaveBeenCalledWith("sk-test", "prompt-1"));
|
||||
await waitFor(() => expect(mockDeletePromptCall).toHaveBeenCalledWith("sk-test", "prompt-1", "staging"));
|
||||
|
||||
await user.keyboard("{Escape}");
|
||||
expect(screen.getByText(/delete prompt: my-prompt/i)).toBeInTheDocument();
|
||||
expect(screen.getByText(/the staging copy of prompt: my-prompt/i)).toBeInTheDocument();
|
||||
|
||||
finishDelete();
|
||||
await waitFor(() => expect(screen.queryByText(/delete prompt: my-prompt/i)).not.toBeInTheDocument());
|
||||
await waitFor(() => expect(screen.queryByText(/the staging copy of prompt: my-prompt/i)).not.toBeInTheDocument());
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -41,11 +41,12 @@ const PromptsPanel: React.FC<PromptsProps> = ({ accessToken, userRole }) => {
|
|||
const [isLoading, setIsLoading] = useState(true);
|
||||
const [selectedEnvironment, setSelectedEnvironment] = useState<string | undefined>(undefined);
|
||||
const [selectedPromptId, setSelectedPromptId] = useState<string | null>(null);
|
||||
const [selectedPromptEnvironment, setSelectedPromptEnvironment] = useState<string | undefined>(undefined);
|
||||
const [isAddModalVisible, setIsAddModalVisible] = useState(false);
|
||||
const [showEditorView, setShowEditorView] = useState(false);
|
||||
const [editPromptData, setEditPromptData] = useState<any>(null);
|
||||
const [isDeleting, setIsDeleting] = useState(false);
|
||||
const [promptToDelete, setPromptToDelete] = useState<{ id: string; name: string } | null>(null);
|
||||
const [promptToDelete, setPromptToDelete] = useState<{ id: string; name: string; environment: string } | null>(null);
|
||||
|
||||
// Admin Viewer follows the read-parity rule: see prompts, no writes.
|
||||
const canModify = userRole ? isProxyAdminRole(userRole) : false;
|
||||
|
|
@ -71,8 +72,9 @@ const PromptsPanel: React.FC<PromptsProps> = ({ accessToken, userRole }) => {
|
|||
fetchPrompts();
|
||||
}, [accessToken, selectedEnvironment]);
|
||||
|
||||
const handlePromptClick = (promptId: string) => {
|
||||
const handlePromptClick = (promptId: string, environment: string) => {
|
||||
setSelectedPromptId(promptId);
|
||||
setSelectedPromptEnvironment(environment);
|
||||
};
|
||||
|
||||
const handleAddPrompt = () => {
|
||||
|
|
@ -111,8 +113,8 @@ const PromptsPanel: React.FC<PromptsProps> = ({ accessToken, userRole }) => {
|
|||
setSelectedPromptId(null);
|
||||
};
|
||||
|
||||
const handleDeleteClick = (promptId: string, promptName: string) => {
|
||||
setPromptToDelete({ id: promptId, name: promptName });
|
||||
const handleDeleteClick = (promptId: string, promptName: string, environment: string) => {
|
||||
setPromptToDelete({ id: promptId, name: promptName, environment });
|
||||
};
|
||||
|
||||
const handleDeleteConfirm = async () => {
|
||||
|
|
@ -120,8 +122,8 @@ const PromptsPanel: React.FC<PromptsProps> = ({ accessToken, userRole }) => {
|
|||
|
||||
setIsDeleting(true);
|
||||
try {
|
||||
await deletePromptCall(accessToken, promptToDelete.id);
|
||||
toast.success(`Prompt "${promptToDelete.name}" deleted successfully`);
|
||||
await deletePromptCall(accessToken, promptToDelete.id, promptToDelete.environment);
|
||||
toast.success(`Prompt "${promptToDelete.name}" deleted successfully from ${promptToDelete.environment}`);
|
||||
fetchPrompts(); // Refresh the list
|
||||
} catch (error) {
|
||||
console.error("Error deleting prompt:", error);
|
||||
|
|
@ -148,6 +150,7 @@ const PromptsPanel: React.FC<PromptsProps> = ({ accessToken, userRole }) => {
|
|||
) : selectedPromptId ? (
|
||||
<PromptInfoView
|
||||
promptId={selectedPromptId}
|
||||
initialEnvironment={selectedPromptEnvironment}
|
||||
onClose={() => setSelectedPromptId(null)}
|
||||
accessToken={accessToken}
|
||||
isAdmin={canModify}
|
||||
|
|
@ -219,7 +222,8 @@ const PromptsPanel: React.FC<PromptsProps> = ({ accessToken, userRole }) => {
|
|||
<AlertDialogHeader>
|
||||
<AlertDialogTitle>Delete Prompt</AlertDialogTitle>
|
||||
<AlertDialogDescription>
|
||||
Are you sure you want to delete prompt: {promptToDelete.name} ? This action cannot be undone.
|
||||
Are you sure you want to delete the {promptToDelete.environment} copy of prompt: {promptToDelete.name}?
|
||||
This action cannot be undone.
|
||||
</AlertDialogDescription>
|
||||
</AlertDialogHeader>
|
||||
<AlertDialogFooter>
|
||||
|
|
|
|||
|
|
@ -29,6 +29,35 @@ const promptWithoutTemplate = {
|
|||
environments: [],
|
||||
};
|
||||
|
||||
describe("PromptInfoView environment scoping", () => {
|
||||
beforeEach(() => {
|
||||
vi.mocked(networking.getPromptInfo).mockReset().mockResolvedValue(promptWithoutTemplate);
|
||||
vi.mocked(networking.getPromptVersions).mockReset().mockResolvedValue({ prompts: [] });
|
||||
});
|
||||
|
||||
it("fetches the initial environment it was opened with", async () => {
|
||||
render(
|
||||
<PromptInfoView
|
||||
promptId="support-reply"
|
||||
initialEnvironment="staging"
|
||||
onClose={vi.fn()}
|
||||
accessToken="sk-test"
|
||||
isAdmin={true}
|
||||
/>,
|
||||
);
|
||||
|
||||
await screen.findByRole("tab", { name: "Raw JSON" });
|
||||
expect(networking.getPromptInfo).toHaveBeenCalledWith("sk-test", "support-reply", "staging");
|
||||
});
|
||||
|
||||
it("fetches the serve default when opened without an environment", async () => {
|
||||
render(<PromptInfoView promptId="support-reply" onClose={vi.fn()} accessToken="sk-test" isAdmin={true} />);
|
||||
|
||||
await screen.findByRole("tab", { name: "Raw JSON" });
|
||||
expect(networking.getPromptInfo).toHaveBeenCalledWith("sk-test", "support-reply", undefined);
|
||||
});
|
||||
});
|
||||
|
||||
describe("PromptInfoView tabs", () => {
|
||||
beforeEach(() => {
|
||||
vi.mocked(networking.getPromptInfo).mockReset().mockResolvedValue(promptWithoutTemplate);
|
||||
|
|
|
|||
|
|
@ -20,6 +20,7 @@ import { Dialog, DialogContent, DialogFooter, DialogHeader, DialogTitle } from "
|
|||
|
||||
export interface PromptInfoProps {
|
||||
promptId: string;
|
||||
initialEnvironment?: string;
|
||||
onClose: () => void;
|
||||
accessToken: string | null;
|
||||
isAdmin: boolean;
|
||||
|
|
@ -27,7 +28,15 @@ export interface PromptInfoProps {
|
|||
onEdit?: (promptData: any) => void;
|
||||
}
|
||||
|
||||
const PromptInfoView: React.FC<PromptInfoProps> = ({ promptId, onClose, accessToken, isAdmin, onDelete, onEdit }) => {
|
||||
const PromptInfoView: React.FC<PromptInfoProps> = ({
|
||||
promptId,
|
||||
initialEnvironment,
|
||||
onClose,
|
||||
accessToken,
|
||||
isAdmin,
|
||||
onDelete,
|
||||
onEdit,
|
||||
}) => {
|
||||
const [promptData, setPromptData] = useState<PromptSpec | null>(null);
|
||||
const [promptTemplate, setPromptTemplate] = useState<PromptTemplateBase | null>(null);
|
||||
const [rawApiResponse, setRawApiResponse] = useState<any>(null);
|
||||
|
|
@ -43,7 +52,7 @@ const PromptInfoView: React.FC<PromptInfoProps> = ({ promptId, onClose, accessTo
|
|||
const [selectedVersion, setSelectedVersion] = useState<number | null>(null);
|
||||
const [loadingVersions, setLoadingVersions] = useState(false);
|
||||
|
||||
// Initial fetch — no environment filter, gets default + all environments list
|
||||
// Fetches the requested environment (or the serve-time default when omitted) plus the environments list
|
||||
const fetchPromptInfo = async (environment?: string) => {
|
||||
try {
|
||||
setLoading(true);
|
||||
|
|
@ -89,7 +98,7 @@ const PromptInfoView: React.FC<PromptInfoProps> = ({ promptId, onClose, accessTo
|
|||
setSelectedEnv(null);
|
||||
setEnvironments([]);
|
||||
setVersionHistory([]);
|
||||
fetchPromptInfo();
|
||||
fetchPromptInfo(initialEnvironment);
|
||||
}, [promptId, accessToken]);
|
||||
|
||||
// When environment changes (user clicks tab), re-fetch — skip initial mount
|
||||
|
|
@ -493,7 +502,7 @@ const PromptInfoView: React.FC<PromptInfoProps> = ({ promptId, onClose, accessTo
|
|||
<DialogTitle>Delete Prompt</DialogTitle>
|
||||
</DialogHeader>
|
||||
<p>
|
||||
Are you sure you want to delete prompt: <strong>{basePromptId}</strong>?
|
||||
Are you sure you want to delete prompt: <strong>{basePromptId}</strong> from every environment?
|
||||
</p>
|
||||
<p>This action cannot be undone.</p>
|
||||
<DialogFooter>
|
||||
|
|
|
|||
|
|
@ -4629,9 +4629,12 @@ export const updatePromptCall = async (accessToken: string, promptId: string, pr
|
|||
}
|
||||
};
|
||||
|
||||
export const deletePromptCall = async (accessToken: string, promptId: string) => {
|
||||
export const deletePromptCall = async (accessToken: string, promptId: string, environment?: string) => {
|
||||
try {
|
||||
const data = await apiClient.delete(`/prompts/${promptId}`, { accessToken });
|
||||
const data = await apiClient.delete(`/prompts/${promptId}`, {
|
||||
accessToken,
|
||||
query: { environment: environment || undefined },
|
||||
});
|
||||
return data;
|
||||
} catch (error) {
|
||||
console.error("Failed to delete prompt:", error);
|
||||
|
|
|
|||
|
|
@ -349,7 +349,7 @@ const CreateKey: React.FC<CreateKeyProps> = ({ team, teams, data, addKey, autoOp
|
|||
const fetchPrompts = async () => {
|
||||
try {
|
||||
const response = await getPromptsList(accessToken);
|
||||
setPromptsList(response.prompts.map((prompt) => prompt.prompt_id));
|
||||
setPromptsList(Array.from(new Set(response.prompts.map((prompt) => prompt.prompt_id))));
|
||||
} catch (error) {
|
||||
console.error("Failed to fetch prompts:", error);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -425,6 +425,22 @@ describe("KeyEditView", () => {
|
|||
expect(screen.getByText("Policies")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("lists a prompt existing in several environments once in the dropdown", async () => {
|
||||
vi.mocked(getPromptsList).mockResolvedValueOnce({
|
||||
prompts: [
|
||||
{ prompt_id: "envgreet", litellm_params: {}, prompt_info: { prompt_type: "db" }, environment: "development" },
|
||||
{ prompt_id: "envgreet", litellm_params: {}, prompt_info: { prompt_type: "db" }, environment: "production" },
|
||||
],
|
||||
});
|
||||
|
||||
renderAs("Admin");
|
||||
|
||||
const prompts = await screen.findByLabelText(/Prompts/);
|
||||
await userEvent.type(prompts, "envgreet");
|
||||
|
||||
expect(await screen.findAllByRole("option", { name: "envgreet" })).toHaveLength(1);
|
||||
});
|
||||
|
||||
it("should omit both fields and fire neither admin-only request for an internal user", async () => {
|
||||
renderAs("Internal User");
|
||||
|
||||
|
|
|
|||
|
|
@ -165,7 +165,7 @@ export function KeyEditView({
|
|||
if (!accessToken) return;
|
||||
try {
|
||||
const response = await getPromptsList(accessToken);
|
||||
setPromptsList(response.prompts.map((prompt) => prompt.prompt_id));
|
||||
setPromptsList(Array.from(new Set(response.prompts.map((prompt) => prompt.prompt_id))));
|
||||
} catch (error) {
|
||||
console.error("Failed to fetch prompts:", error);
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue