This commit is contained in:
Mateo Wang 2026-08-27 11:31:22 -07:00 committed by GitHub
commit 0ad50a39c3
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
22 changed files with 686 additions and 552 deletions

View file

@ -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.",

View file

@ -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()

View file

@ -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)

View file

@ -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."""

View file

@ -3541,6 +3541,7 @@ all_litellm_params = (
"litellm_system_prompt",
"provider_specific_header",
"prompt_version",
"prompt_environment",
"api_base",
"force_timeout",
"logger_fn",

View file

@ -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

View file

@ -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",

View file

@ -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"
)

View file

@ -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

View file

@ -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:

View file

@ -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()

View file

@ -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 () => {

View file

@ -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}

View file

@ -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
}
/>
),
},

View file

@ -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());
});
});

View file

@ -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>

View file

@ -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);

View file

@ -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>

View file

@ -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);

View file

@ -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);
}

View file

@ -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");

View file

@ -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);
}