mirror of
https://github.com/BerriAI/litellm.git
synced 2026-08-28 05:25:59 +00:00
142 lines
5.9 KiB
Python
142 lines
5.9 KiB
Python
import pytest
|
|
|
|
import litellm
|
|
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:
|
|
return PromptSpec(
|
|
prompt_id="greeting.v1",
|
|
litellm_params=PromptLiteLLMParams(
|
|
prompt_id="greeting",
|
|
prompt_integration="dotprompt",
|
|
prompt_data={"content": content, "metadata": {}},
|
|
),
|
|
prompt_info=PromptInfo(prompt_type="db"),
|
|
)
|
|
|
|
|
|
def _served_content(registry: InMemoryPromptRegistry) -> str:
|
|
callback = registry.get_prompt_callback_by_id("greeting.v1")
|
|
assert callback is not None
|
|
return callback.prompt_manager.get_prompt("greeting").content
|
|
|
|
|
|
@pytest.fixture
|
|
def isolated_callbacks(monkeypatch: pytest.MonkeyPatch) -> list:
|
|
monkeypatch.setattr(litellm, "callbacks", [])
|
|
return litellm.callbacks
|
|
|
|
|
|
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")
|
|
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"
|
|
assert stale_callback not in isolated_callbacks
|
|
assert isolated_callbacks == [registry.get_prompt_callback_by_id("greeting.v1")]
|
|
|
|
|
|
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")
|
|
|
|
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 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")
|
|
|
|
reloaded = registry.reload_prompt(prompt=_db_prompt_spec("begin every reply with HOWDY"))
|
|
|
|
assert reloaded is not None
|
|
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]
|
|
|
|
|
|
def _versioned_prompt_spec(version: int, environment: str) -> PromptSpec:
|
|
return PromptSpec(
|
|
prompt_id=f"greeting.v{version}",
|
|
litellm_params=PromptLiteLLMParams(
|
|
prompt_id="greeting",
|
|
prompt_integration="dotprompt",
|
|
prompt_data={"content": f"begin every reply with AHOY v{version}", "metadata": {}},
|
|
),
|
|
prompt_info=PromptInfo(prompt_type="db", environment=environment),
|
|
version=version,
|
|
environment=environment,
|
|
)
|
|
|
|
|
|
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"))
|
|
assert len(isolated_callbacks) == 1
|
|
|
|
deleted = registry.delete_prompts_by_base_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 isolated_callbacks == []
|
|
|
|
|
|
def test_delete_prompts_by_base_id_environment_scope_keeps_other_environments(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")
|
|
|
|
deleted = registry.delete_prompts_by_base_id("greeting", environment="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 len(isolated_callbacks) == 1
|