fix: feat: add litellm_system_prompt support

This commit is contained in:
Krrish Dholakia 2026-02-20 20:53:38 -08:00
parent 0888e17272
commit 32ce793587
5 changed files with 165 additions and 34 deletions

View file

@ -452,7 +452,7 @@ def update_responses_input_with_model_file_ids(
For managed files (unified file IDs), uses model_file_id_mapping if provided,
otherwise decodes the base64-encoded unified file ID and extracts the llm_output_file_id directly.
Args:
input: The responses API input parameter
model_id: The model ID to use for looking up provider-specific file IDs
@ -488,9 +488,13 @@ def update_responses_input_with_model_file_ids(
file_id = content_item.get("file_id")
if file_id:
provider_file_id = file_id # Default to original
# Check if we have a mapping for this file ID
if model_file_id_mapping and model_id and file_id in model_file_id_mapping:
if (
model_file_id_mapping
and model_id
and file_id in model_file_id_mapping
):
# Use the model-specific file ID from mapping
provider_file_id = (
model_file_id_mapping.get(file_id, {}).get(model_id)
@ -501,15 +505,19 @@ def update_responses_input_with_model_file_ids(
updated_content.append(updated_content_item)
else:
# Check if this is a base64-encoded unified file ID without mapping
is_unified_file_id = _is_base64_encoded_unified_file_id(file_id)
is_unified_file_id = _is_base64_encoded_unified_file_id(
file_id
)
if is_unified_file_id:
# Fallback: decode unified file ID
unified_file_id = convert_b64_uid_to_unified_uid(file_id)
unified_file_id = convert_b64_uid_to_unified_uid(
file_id
)
if "llm_output_file_id," in unified_file_id:
provider_file_id = unified_file_id.split(
"llm_output_file_id,"
)[1].split(";")[0]
updated_content_item = content_item.copy()
updated_content_item["file_id"] = provider_file_id
updated_content.append(updated_content_item)
@ -534,9 +542,9 @@ def update_responses_tools_with_model_file_ids(
) -> Optional[List[Dict[str, Any]]]:
"""
Updates responses API tools with provider-specific file IDs.
Handles code_interpreter tools with container.file_ids.
Args:
tools: The responses API tools parameter
model_id: The model ID to use for looking up provider-specific file IDs
@ -545,18 +553,18 @@ def update_responses_tools_with_model_file_ids(
"""
if not tools or not isinstance(tools, list):
return tools
if not model_file_id_mapping or not model_id:
return tools
updated_tools = []
for tool in tools:
if not isinstance(tool, dict):
updated_tools.append(tool)
continue
updated_tool = tool.copy()
# Handle code_interpreter with container file_ids
if tool.get("type") == "code_interpreter":
container = tool.get("container")
@ -578,14 +586,14 @@ def update_responses_tools_with_model_file_ids(
updated_file_ids.append(file_id)
else:
updated_file_ids.append(file_id)
# Update the tool with new file IDs
updated_container = container.copy()
updated_container["file_ids"] = updated_file_ids
updated_tool["container"] = updated_container
updated_tools.append(updated_tool)
return updated_tools
@ -1104,6 +1112,45 @@ def set_last_user_message(
return messages
def add_system_prompt_to_messages(
messages: List[AllMessageValues],
system_prompt: str,
merge_with_first_system: bool = False,
) -> List[AllMessageValues]:
"""
Add a system prompt to the messages list.
Args:
messages: List of chat completion messages
system_prompt: The system prompt content to add. If empty or None, returns messages unchanged.
merge_with_first_system: If True and the first message is already a system message,
prepends the new prompt to that message's content. If False, adds a new system
message at the beginning.
Returns:
New list of messages with the system prompt added
"""
if not system_prompt:
return list(messages)
if merge_with_first_system and messages and messages[0].get("role") == "system":
first = dict(messages[0])
existing_content = first.get("content", "")
if isinstance(existing_content, str):
merged_content = f"{system_prompt.strip()}\n\n{existing_content}"
elif isinstance(existing_content, list):
merged_content = [{"type": "text", "text": system_prompt.strip()}] + list(
existing_content
)
else:
merged_content = [{"type": "text", "text": system_prompt.strip()}]
first["content"] = merged_content
return [cast(AllMessageValues, first)] + list(messages[1:])
system_message: AllMessageValues = {"role": "system", "content": system_prompt}
return [system_message, *messages]
def convert_prefix_message_to_non_prefix_messages(
messages: List[AllMessageValues],
) -> List[AllMessageValues]:

View file

@ -159,6 +159,7 @@ from .litellm_core_utils.fallback_utils import (
completion_with_fallbacks,
)
from .litellm_core_utils.prompt_templates.common_utils import (
add_system_prompt_to_messages,
get_completion_messages,
update_messages_with_model_file_ids,
)
@ -599,7 +600,7 @@ async def acompletion( # noqa: PLR0915
# Add the context to the function
ctx = contextvars.copy_context()
func_with_context = partial(ctx.run, func)
init_response = await loop.run_in_executor(None, func_with_context)
if isinstance(init_response, dict) or isinstance(
init_response, ModelResponse
@ -939,7 +940,7 @@ def responses_api_bridge_check(
model = model.replace("responses/", "")
mode = "responses"
model_info["mode"] = mode
if web_search_options is not None and custom_llm_provider == "xai":
model_info["mode"] = "responses"
model = model.replace("responses/", "")
@ -1108,9 +1109,7 @@ def completion( # type: ignore # noqa: PLR0915
skip_mcp_handler = kwargs.pop("_skip_mcp_handler", False)
if not skip_mcp_handler and tools:
from litellm.responses.mcp.chat_completions_handler import (
acompletion_with_mcp,
)
from litellm.responses.mcp.chat_completions_handler import acompletion_with_mcp
from litellm.responses.mcp.litellm_proxy_mcp_handler import (
LiteLLM_Proxy_MCP_Handler,
)
@ -1245,6 +1244,7 @@ def completion( # type: ignore # noqa: PLR0915
### PROMPT MANAGEMENT ###
prompt_id = cast(Optional[str], kwargs.get("prompt_id", None))
prompt_variables = cast(Optional[dict], kwargs.get("prompt_variables", None))
litellm_system_prompt = kwargs.get("litellm_system_prompt", None)
### COPY MESSAGES ### - related issue https://github.com/BerriAI/litellm/discussions/4489
messages = get_completion_messages(
messages=messages,
@ -1276,6 +1276,14 @@ def completion( # type: ignore # noqa: PLR0915
prompt_version=kwargs.get("prompt_version", None),
)
### LITELLM SYSTEM PROMPT ###
if litellm_system_prompt:
messages = add_system_prompt_to_messages(
messages=messages,
system_prompt=litellm_system_prompt,
merge_with_first_system=True,
)
try:
if base_url is not None:
api_base = base_url
@ -1558,7 +1566,9 @@ def completion( # type: ignore # noqa: PLR0915
## RESPONSES API BRIDGE LOGIC ## - check if model has 'mode: responses' in litellm.model_cost map
model_info, model = responses_api_bridge_check(
model=model, custom_llm_provider=custom_llm_provider, web_search_options=web_search_options
model=model,
custom_llm_provider=custom_llm_provider,
web_search_options=web_search_options,
)
if model_info.get("mode") == "responses":
@ -2209,17 +2219,19 @@ def completion( # type: ignore # noqa: PLR0915
elif custom_llm_provider == "a2a":
# A2A (Agent-to-Agent) Protocol
# Resolve agent configuration from registry if model format is "a2a/<agent-name>"
api_base, api_key, headers = litellm.A2AConfig.resolve_agent_config_from_registry(
model=model,
api_base=api_base,
api_key=api_key,
headers=headers,
optional_params=optional_params,
api_base, api_key, headers = (
litellm.A2AConfig.resolve_agent_config_from_registry(
model=model,
api_base=api_base,
api_key=api_key,
headers=headers,
optional_params=optional_params,
)
)
# Fall back to environment variables and defaults
api_base = api_base or litellm.api_base or get_secret_str("A2A_API_BASE")
if api_base is None:
raise Exception(
"api_base is required for A2A provider. "
@ -4783,7 +4795,10 @@ def embedding( # noqa: PLR0915
or custom_llm_provider == "together_ai"
or custom_llm_provider == "nvidia_nim"
or custom_llm_provider == "litellm_proxy"
or (model in litellm.open_ai_embedding_models and custom_llm_provider is None)
or (
model in litellm.open_ai_embedding_models
and custom_llm_provider is None
)
):
api_base = (
api_base
@ -7239,7 +7254,11 @@ def stream_chunk_builder( # noqa: PLR0915
continue
choice = chunk["choices"][0]
delta_obj = choice.get("delta", {}) if isinstance(choice, dict) else getattr(choice, "delta", {})
delta_obj = (
choice.get("delta", {})
if isinstance(choice, dict)
else getattr(choice, "delta", {})
)
if isinstance(delta_obj, dict):
delta = delta_obj
elif hasattr(delta_obj, "model_dump"):
@ -7266,7 +7285,9 @@ def stream_chunk_builder( # noqa: PLR0915
if is_simple_text_stream:
if simple_content_parts:
response["choices"][0]["message"]["content"] = "".join(simple_content_parts)
response["choices"][0]["message"]["content"] = "".join(
simple_content_parts
)
completion_output = get_content_from_model_response(response)
usage = processor.calculate_usage(
chunks=chunks,
@ -7291,7 +7312,9 @@ def stream_chunk_builder( # noqa: PLR0915
if litellm.include_cost_in_streaming_usage and logging_obj is not None:
setattr(
usage, "cost", logging_obj._response_cost_calculator(result=response)
usage,
"cost",
logging_obj._response_cost_calculator(result=response),
)
return response
@ -7504,6 +7527,7 @@ def __getattr__(name: str) -> Any:
# before loading tiktoken, ensuring the local cache is used
# instead of downloading from the internet
from litellm._lazy_imports import _get_default_encoding
_encoding = _get_default_encoding()
# Cache it in the module's __dict__ for subsequent accesses
import sys

View file

@ -16,6 +16,10 @@ model_list:
- model_name: gpt-5-mini
litellm_params:
model: openai/gpt-5-mini
- model_name: custom_litellm_model
litellm_params:
model: litellm_agent/claude-sonnet-4-5-20250929
litellm_system_prompt: "Be a helpful assistant."
guardrails:

View file

@ -2917,8 +2917,9 @@ all_litellm_params = (
"api_key",
"api_version",
"prompt_id",
"provider_specific_header",
"prompt_variables",
"litellm_system_prompt",
"provider_specific_header",
"prompt_version",
"api_base",
"force_timeout",

View file

@ -10,6 +10,7 @@ sys.path.insert(
) # Adds the parent directory to the system path
from litellm.litellm_core_utils.prompt_templates.common_utils import (
add_system_prompt_to_messages,
get_format_from_file_id,
handle_any_messages_to_chat_completion_str_messages_conversion,
split_concatenated_json_objects,
@ -128,6 +129,60 @@ def test_handle_any_messages_to_chat_completion_str_messages_conversion_complex(
assert result[0]["input"] == json.dumps(message)
def test_add_system_prompt_to_messages_prepend():
"""Adds system prompt at beginning when no system message exists."""
messages = [
{"role": "user", "content": "Hello"},
{"role": "assistant", "content": "Hi there"},
]
result = add_system_prompt_to_messages(messages, "You are a helpful assistant.")
assert result == [
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": "Hello"},
{"role": "assistant", "content": "Hi there"},
]
def test_add_system_prompt_to_messages_empty_prompt_unchanged():
"""Returns messages unchanged when system_prompt is empty."""
messages = [{"role": "user", "content": "Hello"}]
assert add_system_prompt_to_messages(messages, "") == messages
assert add_system_prompt_to_messages(messages, None) == messages
def test_add_system_prompt_to_messages_merge_with_first_system():
"""Merges new prompt into first system message when merge_with_first_system=True."""
messages = [
{"role": "system", "content": "Existing system prompt."},
{"role": "user", "content": "Hello"},
]
result = add_system_prompt_to_messages(
messages, "You are helpful.", merge_with_first_system=True
)
assert result == [
{"role": "system", "content": "You are helpful.\n\nExisting system prompt."},
{"role": "user", "content": "Hello"},
]
def test_add_system_prompt_to_messages_merge_with_first_system_adds_new_when_no_system():
"""When merge_with_first_system=True but no system message, adds new one at start."""
messages = [{"role": "user", "content": "Hello"}]
result = add_system_prompt_to_messages(
messages, "You are helpful.", merge_with_first_system=True
)
assert result == [
{"role": "system", "content": "You are helpful."},
{"role": "user", "content": "Hello"},
]
def test_add_system_prompt_to_messages_empty_list():
"""Adds system prompt to empty messages list."""
result = add_system_prompt_to_messages([], "You are helpful.")
assert result == [{"role": "system", "content": "You are helpful."}]
def test_convert_prefix_message_to_non_prefix_messages():
from litellm.litellm_core_utils.prompt_templates.common_utils import (
convert_prefix_message_to_non_prefix_messages,