fix(prompts): validate a prompt replacement before swapping and isolate per-row sync failures

This commit is contained in:
mateo-berri 2026-08-26 15:49:31 -07:00
parent 5130bafda8
commit 6df307fef8
4 changed files with 98 additions and 19 deletions

View file

@ -118,7 +118,16 @@ class InMemoryPromptRegistry:
verbose_proxy_logger.debug("prompt_id already exists in IN_MEMORY_PROMPTS")
return self.IN_MEMORY_PROMPTS[prompt_id]
custom_prompt_callback: CustomPromptManagement | None = None
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
return parsed_prompt
def _build_prompt_callback(self, prompt: PromptSpec) -> tuple[PromptSpec, CustomPromptManagement]:
litellm_params_data: Final = prompt.litellm_params
verbose_proxy_logger.debug("litellm_params= %s", litellm_params_data)
@ -132,17 +141,17 @@ class InMemoryPromptRegistry:
raise ValueError("prompt_integration is required")
initializer: Final = prompt_initializer_registry.get(prompt_integration)
if initializer:
custom_prompt_callback = initializer(litellm_params, prompt)
if not isinstance(custom_prompt_callback, CustomPromptManagement):
raise ValueError(f"CustomPromptManagement is required, got {type(custom_prompt_callback)}")
litellm.logging_callback_manager.add_litellm_callback(custom_prompt_callback)
else:
if initializer is None:
raise ValueError(f"Unsupported prompt: {prompt_integration}")
custom_prompt_callback: Final = initializer(litellm_params, prompt)
if not isinstance(custom_prompt_callback, CustomPromptManagement):
raise ValueError( # noqa: TRY004 # prompt endpoints map ValueError to HTTP 400; keep the existing contract
f"CustomPromptManagement is required, got {type(custom_prompt_callback)}"
)
parsed_prompt: Final = PromptSpec(
prompt_id=prompt_id,
prompt_id=prompt.prompt_id,
litellm_params=litellm_params,
prompt_info=prompt.prompt_info or PromptInfo(prompt_type="config"),
created_at=prompt.created_at,
@ -151,21 +160,20 @@ class InMemoryPromptRegistry:
environment=prompt.environment,
created_by=prompt.created_by,
)
# 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
return parsed_prompt
return parsed_prompt, custom_prompt_callback
def reload_prompt(self, prompt: PromptSpec) -> PromptSpec | None:
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)
if stale_callback is not None:
litellm.logging_callback_manager.remove_callback_from_all_lists(stale_callback)
return self.initialize_prompt(prompt=prompt)
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
return parsed_prompt
def sync_prompt_from_db(self, prompt: PromptSpec) -> PromptSpec | None:
existing: Final = self.IN_MEMORY_PROMPTS.get(prompt.prompt_id)

View file

@ -7242,9 +7242,15 @@ class ProxyConfig:
try:
prompts_in_db: Final[Sequence[object]] = await PromptRepository(prisma_client).table.find_many()
for prompt in prompts_in_db:
# Convert DB object to dict and create versioned prompt_id
prompt_spec = self._get_prompt_spec_for_db_prompt(db_prompt=prompt)
IN_MEMORY_PROMPT_REGISTRY.sync_prompt_from_db(prompt=prompt_spec)
try:
prompt_spec = self._get_prompt_spec_for_db_prompt(db_prompt=prompt)
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
verbose_proxy_logger.exception(
"litellm.proxy.proxy_server.py::ProxyConfig:_init_prompts_in_db - failed to sync prompt %s: %s",
getattr(prompt, "prompt_id", None),
prompt_sync_error,
)
except Exception as e:
verbose_proxy_logger.debug("litellm.proxy.proxy_server.py::ProxyConfig:_init_prompts_in_db - %s", e)

View file

@ -65,3 +65,26 @@ def test_reload_prompt_replaces_callback_without_leaking_the_old_one(isolated_ca
assert _served_content(registry) == "begin every reply with HOWDY"
assert stale_callback not in isolated_callbacks
assert len(isolated_callbacks) == 1
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")
broken = PromptSpec(
prompt_id="greeting.v1",
litellm_params=PromptLiteLLMParams(
prompt_id="greeting",
prompt_integration="does_not_exist",
prompt_data={"content": "begin every reply with HOWDY", "metadata": {}},
),
prompt_info=PromptInfo(prompt_type="db"),
)
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 _served_content(registry) == "begin every reply with AHOY"
assert isolated_callbacks == [old_callback]

View file

@ -11327,6 +11327,48 @@ async def test_init_prompts_in_db_reloads_rows_patched_on_another_worker(monkeyp
IN_MEMORY_PROMPT_REGISTRY.delete_prompts_by_base_id("greeting_sync")
@pytest.mark.asyncio
async def test_init_prompts_in_db_syncs_remaining_rows_when_one_row_fails(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(prompt_id: str, integration: str) -> MagicMock:
row = MagicMock()
row.model_dump.return_value = {
"prompt_id": prompt_id,
"version": 1,
"environment": "development",
"created_by": None,
"litellm_params": json.dumps(
{
"prompt_id": prompt_id,
"prompt_integration": integration,
"prompt_data": {"content": "Begin every reply with AHOY", "metadata": {}},
}
),
"prompt_info": json.dumps({"prompt_type": "db"}),
"created_at": None,
"updated_at": None,
}
return row
prisma_client = MagicMock()
try:
prisma_client.db.litellm_prompttable.find_many = AsyncMock(
return_value=[db_row("broken_sync", "does_not_exist"), db_row("healthy_sync", "dotprompt")]
)
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")]
finally:
IN_MEMORY_PROMPT_REGISTRY.delete_prompts_by_base_id("healthy_sync")
IN_MEMORY_PROMPT_REGISTRY.delete_prompts_by_base_id("broken_sync")
class TestEmbeddingsFailureHookRequestData:
@pytest.mark.asyncio
async def test_failure_hook_gets_post_setup_data_with_logging_obj(self):