mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(prompt-templates): treat initial_prompt_value and final_prompt_value as optional
completion() only stores a custom prompt template key when its value is truthy, and custom_prompt() already defaults both prompt values to "". Seven call sites read them back with direct indexing, so a template that sets only roles raises KeyError instead of rendering, on ollama, petals, vllm, predibase, codestral, anthropic text completion, and response_schema_prompt. Five other providers already read defensively, so the same config works on sagemaker and 500s on ollama. Read both keys with .get(key, ""), matching those five providers and the defaults custom_prompt() already declares. Templates that set all three keys render exactly as before.
This commit is contained in:
parent
7276caecd4
commit
7e4dc1d935
8 changed files with 214 additions and 16 deletions
|
|
@ -5165,8 +5165,8 @@ def response_schema_prompt(model: str, response_schema: dict) -> str:
|
|||
if custom_prompt_details is not None:
|
||||
return custom_prompt(
|
||||
role_dict=custom_prompt_details["roles"],
|
||||
initial_prompt_value=custom_prompt_details["initial_prompt_value"],
|
||||
final_prompt_value=custom_prompt_details["final_prompt_value"],
|
||||
initial_prompt_value=custom_prompt_details.get("initial_prompt_value", ""),
|
||||
final_prompt_value=custom_prompt_details.get("final_prompt_value", ""),
|
||||
messages=response_schema_as_message,
|
||||
)
|
||||
else:
|
||||
|
|
|
|||
|
|
@ -239,8 +239,8 @@ class AnthropicTextConfig(BaseConfig):
|
|||
model_prompt_details: Final = custom_prompt_dict[model]
|
||||
prompt = custom_prompt(
|
||||
role_dict=model_prompt_details["roles"],
|
||||
initial_prompt_value=model_prompt_details["initial_prompt_value"],
|
||||
final_prompt_value=model_prompt_details["final_prompt_value"],
|
||||
initial_prompt_value=model_prompt_details.get("initial_prompt_value", ""),
|
||||
final_prompt_value=model_prompt_details.get("final_prompt_value", ""),
|
||||
messages=messages,
|
||||
)
|
||||
else:
|
||||
|
|
|
|||
|
|
@ -263,8 +263,8 @@ class CodestralTextCompletion:
|
|||
model_prompt_details: Final = custom_prompt_dict[model]
|
||||
prompt = custom_prompt(
|
||||
role_dict=model_prompt_details["roles"],
|
||||
initial_prompt_value=model_prompt_details["initial_prompt_value"],
|
||||
final_prompt_value=model_prompt_details["final_prompt_value"],
|
||||
initial_prompt_value=model_prompt_details.get("initial_prompt_value", ""),
|
||||
final_prompt_value=model_prompt_details.get("final_prompt_value", ""),
|
||||
messages=messages,
|
||||
)
|
||||
else:
|
||||
|
|
|
|||
|
|
@ -360,8 +360,8 @@ class OllamaConfig(BaseConfig):
|
|||
model_prompt_details: Final = custom_prompt_dict[model]
|
||||
ollama_prompt = custom_prompt(
|
||||
role_dict=model_prompt_details["roles"],
|
||||
initial_prompt_value=model_prompt_details["initial_prompt_value"],
|
||||
final_prompt_value=model_prompt_details["final_prompt_value"],
|
||||
initial_prompt_value=model_prompt_details.get("initial_prompt_value", ""),
|
||||
final_prompt_value=model_prompt_details.get("final_prompt_value", ""),
|
||||
messages=messages,
|
||||
)
|
||||
elif text_completion_request: # handle `/completions` requests
|
||||
|
|
|
|||
|
|
@ -44,8 +44,8 @@ def completion(
|
|||
model_prompt_details: Final = litellm.custom_prompt_dict[model]
|
||||
prompt = custom_prompt(
|
||||
role_dict=model_prompt_details["roles"],
|
||||
initial_prompt_value=model_prompt_details["initial_prompt_value"],
|
||||
final_prompt_value=model_prompt_details["final_prompt_value"],
|
||||
initial_prompt_value=model_prompt_details.get("initial_prompt_value", ""),
|
||||
final_prompt_value=model_prompt_details.get("final_prompt_value", ""),
|
||||
messages=messages,
|
||||
)
|
||||
else:
|
||||
|
|
|
|||
|
|
@ -261,8 +261,8 @@ class PredibaseConfig(BaseConfig):
|
|||
model_prompt_details: Final = custom_prompt_dict[model]
|
||||
prompt = custom_prompt(
|
||||
role_dict=model_prompt_details["roles"],
|
||||
initial_prompt_value=model_prompt_details["initial_prompt_value"],
|
||||
final_prompt_value=model_prompt_details["final_prompt_value"],
|
||||
initial_prompt_value=model_prompt_details.get("initial_prompt_value", ""),
|
||||
final_prompt_value=model_prompt_details.get("final_prompt_value", ""),
|
||||
messages=messages,
|
||||
)
|
||||
else:
|
||||
|
|
|
|||
|
|
@ -58,8 +58,8 @@ def completion(
|
|||
model_prompt_details: Final = custom_prompt_dict[model]
|
||||
prompt = custom_prompt(
|
||||
role_dict=model_prompt_details["roles"],
|
||||
initial_prompt_value=model_prompt_details["initial_prompt_value"],
|
||||
final_prompt_value=model_prompt_details["final_prompt_value"],
|
||||
initial_prompt_value=model_prompt_details.get("initial_prompt_value", ""),
|
||||
final_prompt_value=model_prompt_details.get("final_prompt_value", ""),
|
||||
messages=messages,
|
||||
)
|
||||
else:
|
||||
|
|
@ -146,8 +146,8 @@ def batch_completions(model: str, messages: list, optional_params=None, custom_p
|
|||
for message in messages:
|
||||
prompt = custom_prompt(
|
||||
role_dict=model_prompt_details["roles"],
|
||||
initial_prompt_value=model_prompt_details["initial_prompt_value"],
|
||||
final_prompt_value=model_prompt_details["final_prompt_value"],
|
||||
initial_prompt_value=model_prompt_details.get("initial_prompt_value", ""),
|
||||
final_prompt_value=model_prompt_details.get("final_prompt_value", ""),
|
||||
messages=message,
|
||||
)
|
||||
prompts.append(prompt)
|
||||
|
|
|
|||
|
|
@ -0,0 +1,198 @@
|
|||
"""Every consumer of ``custom_prompt_dict`` must treat the two prompt-value keys as optional.
|
||||
|
||||
``litellm.completion()`` only stores ``initial_prompt_value`` / ``final_prompt_value`` when the
|
||||
caller passes a truthy value, and ``custom_prompt()`` already defaults both to ``""``. A template
|
||||
that sets ``roles`` alone therefore has to render, not raise ``KeyError``.
|
||||
"""
|
||||
|
||||
import sys
|
||||
import types
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.prompt_templates.factory import response_schema_prompt
|
||||
from litellm.llms.anthropic.completion.transformation import AnthropicTextConfig
|
||||
from litellm.llms.codestral.completion.handler import CodestralTextCompletion
|
||||
from litellm.llms.ollama.completion.transformation import OllamaConfig
|
||||
from litellm.llms.petals.completion import handler as petals_handler
|
||||
from litellm.llms.predibase.chat.transformation import PredibaseConfig
|
||||
from litellm.llms.vllm.completion import handler as vllm_handler
|
||||
from litellm.types.utils import ModelResponse, TextCompletionResponse
|
||||
|
||||
MODEL = "partial-template-model"
|
||||
ROLES = {"user": {"pre_message": "[INST] ", "post_message": " [/INST]"}}
|
||||
MESSAGES = [{"role": "user", "content": "hi"}]
|
||||
RENDERED_MESSAGES = "[INST] hi [/INST]"
|
||||
|
||||
RESPONSE_SCHEMA = {"type": "object"}
|
||||
RENDERED_SCHEMA = f"[INST] {RESPONSE_SCHEMA} [/INST]"
|
||||
|
||||
INITIAL = "<start>"
|
||||
FINAL = "<end>"
|
||||
|
||||
PARTIAL_TEMPLATE = {"roles": ROLES}
|
||||
FULL_TEMPLATE = {"roles": ROLES, "initial_prompt_value": INITIAL, "final_prompt_value": FINAL}
|
||||
|
||||
|
||||
class _PromptCaptured(Exception):
|
||||
"""Raised by the test doubles below to stop a provider call once it has built its prompt."""
|
||||
|
||||
def __init__(self, prompt):
|
||||
super().__init__(prompt)
|
||||
self.prompt = prompt
|
||||
|
||||
|
||||
class _CapturingLogging:
|
||||
"""Stands in for the logging object providers call with the prompt they are about to send."""
|
||||
|
||||
def pre_call(self, input, api_key, additional_args=None, **kwargs):
|
||||
raise _PromptCaptured(input)
|
||||
|
||||
def post_call(self, *args, **kwargs):
|
||||
pass
|
||||
|
||||
|
||||
def _captured_prompt(call):
|
||||
with pytest.raises(_PromptCaptured) as excinfo:
|
||||
call()
|
||||
return excinfo.value.prompt
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def stub_vllm_import(monkeypatch):
|
||||
"""``vllm`` is an optional heavyweight dependency; its handler only needs the two names."""
|
||||
|
||||
class _SamplingParams:
|
||||
def __init__(self, **kwargs):
|
||||
self.kwargs = kwargs
|
||||
|
||||
class _LLM:
|
||||
def __init__(self, model):
|
||||
self.model = model
|
||||
|
||||
def generate(self, prompts, sampling_params):
|
||||
raise _PromptCaptured(prompts)
|
||||
|
||||
stub = types.ModuleType("vllm")
|
||||
stub.LLM = _LLM
|
||||
stub.SamplingParams = _SamplingParams
|
||||
|
||||
monkeypatch.setitem(sys.modules, "vllm", stub)
|
||||
monkeypatch.setattr(vllm_handler, "llm", None)
|
||||
|
||||
|
||||
def _render_ollama(template, monkeypatch):
|
||||
request = OllamaConfig().transform_request(
|
||||
model=MODEL,
|
||||
messages=MESSAGES,
|
||||
optional_params={},
|
||||
litellm_params={"custom_prompt_dict": {MODEL: template}},
|
||||
headers={},
|
||||
)
|
||||
return request["prompt"]
|
||||
|
||||
|
||||
def _render_predibase(template, monkeypatch):
|
||||
request = PredibaseConfig().transform_request(
|
||||
model=MODEL,
|
||||
messages=MESSAGES,
|
||||
optional_params={},
|
||||
litellm_params={"custom_prompt_dict": {MODEL: template}},
|
||||
headers={},
|
||||
)
|
||||
return request["inputs"]
|
||||
|
||||
|
||||
def _render_anthropic_text(template, monkeypatch):
|
||||
monkeypatch.setattr(litellm, "custom_prompt_dict", {MODEL: template})
|
||||
return AnthropicTextConfig()._get_anthropic_text_prompt_from_messages(messages=MESSAGES, model=MODEL)
|
||||
|
||||
|
||||
def _render_response_schema_prompt(template, monkeypatch):
|
||||
monkeypatch.setattr(litellm, "custom_prompt_dict", {f"{MODEL}/response_schema_prompt": template})
|
||||
return response_schema_prompt(model=MODEL, response_schema=RESPONSE_SCHEMA)
|
||||
|
||||
|
||||
def _render_petals(template, monkeypatch):
|
||||
monkeypatch.setattr(litellm, "custom_prompt_dict", {MODEL: template})
|
||||
return _captured_prompt(
|
||||
lambda: petals_handler.completion(
|
||||
model=MODEL,
|
||||
messages=MESSAGES,
|
||||
api_base="http://petals.invalid",
|
||||
model_response=ModelResponse(),
|
||||
print_verbose=lambda *args, **kwargs: None,
|
||||
encoding=None,
|
||||
logging_obj=_CapturingLogging(),
|
||||
optional_params={},
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _render_codestral(template, monkeypatch):
|
||||
return _captured_prompt(
|
||||
lambda: CodestralTextCompletion().completion(
|
||||
model=MODEL,
|
||||
messages=MESSAGES,
|
||||
api_base="http://codestral.invalid",
|
||||
custom_prompt_dict={MODEL: template},
|
||||
model_response=TextCompletionResponse(),
|
||||
print_verbose=lambda *args, **kwargs: None,
|
||||
encoding=None,
|
||||
api_key="sk-not-used",
|
||||
logging_obj=_CapturingLogging(),
|
||||
optional_params={},
|
||||
timeout=1.0,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _render_vllm_completion(template, monkeypatch):
|
||||
return _captured_prompt(
|
||||
lambda: vllm_handler.completion(
|
||||
model=MODEL,
|
||||
messages=MESSAGES,
|
||||
model_response=ModelResponse(),
|
||||
print_verbose=lambda *args, **kwargs: None,
|
||||
encoding=None,
|
||||
logging_obj=_CapturingLogging(),
|
||||
optional_params={},
|
||||
custom_prompt_dict={MODEL: template},
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _render_vllm_batch(template, monkeypatch):
|
||||
prompts = _captured_prompt(
|
||||
lambda: vllm_handler.batch_completions(
|
||||
model=MODEL,
|
||||
messages=[MESSAGES],
|
||||
optional_params={},
|
||||
custom_prompt_dict={MODEL: template},
|
||||
)
|
||||
)
|
||||
assert len(prompts) == 1
|
||||
return prompts[0]
|
||||
|
||||
|
||||
CALL_SITES = [
|
||||
pytest.param(_render_ollama, RENDERED_MESSAGES, id="ollama"),
|
||||
pytest.param(_render_petals, RENDERED_MESSAGES, id="petals"),
|
||||
pytest.param(_render_vllm_completion, RENDERED_MESSAGES, id="vllm-completion"),
|
||||
pytest.param(_render_vllm_batch, RENDERED_MESSAGES, id="vllm-batch-completions"),
|
||||
pytest.param(_render_predibase, RENDERED_MESSAGES, id="predibase"),
|
||||
pytest.param(_render_codestral, RENDERED_MESSAGES, id="codestral"),
|
||||
pytest.param(_render_anthropic_text, RENDERED_MESSAGES, id="anthropic-text"),
|
||||
pytest.param(_render_response_schema_prompt, RENDERED_SCHEMA, id="response-schema-prompt"),
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(("render", "rendered_body"), CALL_SITES)
|
||||
def test_template_without_the_optional_prompt_values_renders(render, rendered_body, monkeypatch):
|
||||
assert render(PARTIAL_TEMPLATE, monkeypatch) == rendered_body
|
||||
|
||||
|
||||
@pytest.mark.parametrize(("render", "rendered_body"), CALL_SITES)
|
||||
def test_template_with_the_optional_prompt_values_still_applies_them(render, rendered_body, monkeypatch):
|
||||
assert render(FULL_TEMPLATE, monkeypatch) == INITIAL + rendered_body + FINAL
|
||||
Loading…
Add table
Reference in a new issue