mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix: feat: add litellm_system_prompt support
This commit is contained in:
parent
0888e17272
commit
32ce793587
5 changed files with 165 additions and 34 deletions
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue