fix(prompts): key the in-memory prompt registry by environment

LiteLLM_PromptTable is unique on (prompt_id, version, environment) and
version numbering restarts at 1 per environment, but the in-memory
registry keyed prompts as {prompt_id}.v{version} with no environment, so
environments sharing a prompt id shadowed each other and only one
environment's template ever served.

Registry entries are now keyed {versioned_id}::{environment}, and serve
time resolution goes through resolve_prompt_spec(base_id, version,
environment): production > staging > development when no environment is
requested, latest version within the chosen environment when no version
is requested. Chat requests can pin an environment with a new optional
prompt_environment body param, filtered from provider-bound params like
prompt_id and prompt_version. The newest-updated_at dedupe in
_init_prompts_in_db is dropped since registry keys can no longer
collide, and the key-parsing serve helpers plus dead registry getters
are removed
This commit is contained in:
mateo-berri 2026-08-26 18:17:22 -07:00
parent f677292901
commit 7bac4a41af
11 changed files with 413 additions and 468 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,8 @@ 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)
prompt_spec = IN_MEMORY_PROMPT_REGISTRY.resolve_prompt_spec(prompt_id, version=requested_version)
if prompt_spec is None:
raise HTTPException(
@ -785,7 +642,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 +742,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 +754,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,52 +843,24 @@ 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)
# Remove matching prompts from memory — scope to environment if provided
if environment:
prompts_to_delete: Final = [
pid
for pid, prompt in IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS.items()
if get_base_prompt_id(prompt_id=pid) == base_prompt_id and prompt.environment == environment
]
for pid in prompts_to_delete:
del IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS[pid]
if pid in IN_MEMORY_PROMPT_REGISTRY.prompt_id_to_custom_prompt:
del IN_MEMORY_PROMPT_REGISTRY.prompt_id_to_custom_prompt[pid]
else:
IN_MEMORY_PROMPT_REGISTRY.delete_prompts_by_base_id(base_prompt_id)
IN_MEMORY_PROMPT_REGISTRY.delete_prompts_by_base_id(base_prompt_id=base_prompt_id, environment=environment)
env_msg: Final = f" from {environment}" if environment else ""
return {"message": f"Prompt {base_prompt_id} deleted successfully{env_msg}"}
@ -1105,7 +932,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
@ -1129,11 +956,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,57 +236,85 @@ 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
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 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 delete_prompts_by_base_id(self, base_prompt_id: str, environment: str | None = None) -> list[str]:
"""
return self.prompt_id_to_custom_prompt.get(prompt_id)
Delete matching prompts from memory, scoped to one environment when given.
def delete_prompts_by_base_id(self, base_prompt_id: str) -> list[str]:
Returns the registry keys that were deleted.
"""
Delete all prompts matching the given base prompt ID from memory.
Args:
base_prompt_id: The base prompt ID (without version suffix)
Returns:
List of prompt IDs that were deleted
"""
from litellm.proxy.prompts.prompt_endpoints import get_base_prompt_id
prompts_to_delete: Final = [
pid for pid in self.IN_MEMORY_PROMPTS if get_base_prompt_id(prompt_id=pid) == base_prompt_id
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:
del self.IN_MEMORY_PROMPTS[pid]
if pid in self.prompt_id_to_custom_prompt:
del self.prompt_id_to_custom_prompt[pid]
for key in keys_to_delete:
del self.IN_MEMORY_PROMPTS[key]
self.prompt_id_to_custom_prompt.pop(key, None)
return prompts_to_delete
return keys_to_delete
IN_MEMORY_PROMPT_REGISTRY: Final = InMemoryPromptRegistry()

View file

@ -7273,16 +7273,7 @@ class ProxyConfig:
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

View file

@ -1397,27 +1397,26 @@ 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.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:
(
@ -1444,6 +1443,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

@ -3522,6 +3522,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,8 +362,7 @@ 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)

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
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
base_prompt_id=expected_base_id, environment=None
)
assert response == {
@ -187,24 +169,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
@ -253,7 +219,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(
@ -294,7 +260,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(
@ -456,7 +422,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,11 +91,101 @@ 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 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"))
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 _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": "begin every reply with AHOY", "metadata": {}},
),
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

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

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"))
@ -818,3 +815,44 @@ async def test_process_prompt_template_async_get_prompt_error_raises(proxy_loggi
prompt_version=None,
call_type="completion",
)
@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"}]