mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-28 01:32:17 +00:00
fix(prompt-templates): treat initial/final prompt values as optional
`litellm.completion()` stores `initial_prompt_value` and `final_prompt_value` in `custom_prompt_dict` only when they are truthy, but seven call sites read them back with direct indexing. A custom template that sets `roles` alone - which the docs describe as valid, since both string keys are documented as optional - raised `KeyError: 'initial_prompt_value'`, surfaced to the caller as an `APIConnectionError` before any request was attempted. Read both keys with `.get(key, "")`, matching the call sites that already do (bedrock, sagemaker, replicate, huggingface, watsonx) and the `""` defaults that `custom_prompt()` itself declares. Templates that set every key render exactly as before. Fixes #39759 Co-Authored-By: Claude Opus 5 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01J7Jr1zaqrMwsDJqpEfiCpx
This commit is contained in:
parent
f74bc9427b
commit
30ec4e4fce
8 changed files with 125 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)
|
||||
|
|
|
|||
109
tests/test_litellm/llms/test_partial_custom_prompt_template.py
Normal file
109
tests/test_litellm/llms/test_partial_custom_prompt_template.py
Normal file
|
|
@ -0,0 +1,109 @@
|
|||
"""
|
||||
Regression tests for partial custom prompt templates.
|
||||
|
||||
`initial_prompt_value` and `final_prompt_value` are documented as optional, and
|
||||
`litellm.completion()` only stores each key when its value is truthy. Call sites that
|
||||
read them back with direct indexing therefore raised `KeyError` for any template that
|
||||
set `roles` alone, surfacing as an `APIConnectionError` before a request was ever made.
|
||||
|
||||
See https://github.com/BerriAI/litellm/issues/39759
|
||||
"""
|
||||
|
||||
from typing import Final
|
||||
|
||||
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.ollama.completion.transformation import OllamaConfig
|
||||
from litellm.llms.predibase.chat.transformation import PredibaseConfig
|
||||
|
||||
MODEL: Final = "test-model"
|
||||
|
||||
ROLES: Final = {
|
||||
"system": {"pre_message": "<<SYS>>\n", "post_message": "\n<</SYS>>\n"},
|
||||
"user": {"pre_message": "[INST] ", "post_message": " [/INST]"},
|
||||
"assistant": {"pre_message": "", "post_message": "</s>"},
|
||||
}
|
||||
|
||||
MESSAGES: Final = [{"role": "user", "content": "hello"}]
|
||||
|
||||
RENDERED_MESSAGES: Final = "[INST] hello [/INST]"
|
||||
|
||||
|
||||
def _ollama_prompt(template: dict) -> str:
|
||||
request = OllamaConfig().transform_request(
|
||||
model=MODEL,
|
||||
messages=MESSAGES,
|
||||
optional_params={},
|
||||
litellm_params={"custom_prompt_dict": {MODEL: template}},
|
||||
headers={},
|
||||
)
|
||||
return request["prompt"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"template, expected",
|
||||
[
|
||||
pytest.param({"roles": ROLES}, RENDERED_MESSAGES, id="roles-only"),
|
||||
pytest.param(
|
||||
{"roles": ROLES, "initial_prompt_value": "BEGIN\n"},
|
||||
f"BEGIN\n{RENDERED_MESSAGES}",
|
||||
id="no-final-prompt-value",
|
||||
),
|
||||
pytest.param(
|
||||
{"roles": ROLES, "final_prompt_value": "\nEND"},
|
||||
f"{RENDERED_MESSAGES}\nEND",
|
||||
id="no-initial-prompt-value",
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_ollama_renders_partial_custom_prompt_template(template: dict, expected: str) -> None:
|
||||
"""An unset optional key renders as the empty string instead of raising KeyError."""
|
||||
assert _ollama_prompt(template) == expected
|
||||
|
||||
|
||||
def test_ollama_full_custom_prompt_template_is_unchanged() -> None:
|
||||
"""A template that sets every key keeps rendering exactly as before."""
|
||||
template = {
|
||||
"roles": ROLES,
|
||||
"initial_prompt_value": "BEGIN\n",
|
||||
"final_prompt_value": "\nEND",
|
||||
}
|
||||
|
||||
assert _ollama_prompt(template) == f"BEGIN\n{RENDERED_MESSAGES}\nEND"
|
||||
|
||||
|
||||
def test_predibase_renders_partial_custom_prompt_template() -> None:
|
||||
request = PredibaseConfig().transform_request(
|
||||
model=MODEL,
|
||||
messages=MESSAGES,
|
||||
optional_params={},
|
||||
litellm_params={"custom_prompt_dict": {MODEL: {"roles": ROLES}}},
|
||||
headers={},
|
||||
)
|
||||
|
||||
assert request["inputs"] == RENDERED_MESSAGES
|
||||
|
||||
|
||||
def test_anthropic_text_renders_partial_custom_prompt_template(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setattr(litellm, "custom_prompt_dict", {MODEL: {"roles": ROLES}})
|
||||
|
||||
prompt = AnthropicTextConfig()._get_anthropic_text_prompt_from_messages(
|
||||
messages=MESSAGES,
|
||||
model=MODEL,
|
||||
)
|
||||
|
||||
assert prompt == RENDERED_MESSAGES
|
||||
|
||||
|
||||
def test_response_schema_prompt_renders_partial_custom_prompt_template(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setattr(litellm, "custom_prompt_dict", {"response_schema_prompt": {"roles": ROLES}})
|
||||
response_schema: Final = {"type": "object", "properties": {"name": {"type": "string"}}}
|
||||
|
||||
prompt = response_schema_prompt(model=MODEL, response_schema=response_schema)
|
||||
|
||||
assert prompt == f"[INST] {response_schema} [/INST]"
|
||||
Loading…
Add table
Reference in a new issue