mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(core): merge mixed system message content
This commit is contained in:
parent
f6587faef5
commit
6396cc03f7
2 changed files with 62 additions and 1 deletions
|
|
@ -18,6 +18,7 @@ from litellm import verbose_logger
|
|||
from litellm._uuid import uuid
|
||||
from litellm.litellm_core_utils.url_utils import async_safe_get, safe_get
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler, get_async_httpx_client
|
||||
from litellm.types.completion import ChatCompletionContentPartTextParam
|
||||
from litellm.types.files import get_file_extension_from_mime_type
|
||||
from litellm.types.llms.anthropic import *
|
||||
from litellm.types.llms.bedrock import CachePointBlock
|
||||
|
|
@ -107,7 +108,20 @@ def map_system_message_pt(messages: list) -> list:
|
|||
next_role = next_m["role"]
|
||||
if next_role == "user" or next_role == "assistant": # Next message is a user or assistant message
|
||||
# Merge system prompt into the next message
|
||||
next_m["content"] = m["content"] + " " + next_m["content"]
|
||||
system_content = m["content"]
|
||||
next_content = next_m["content"]
|
||||
if isinstance(system_content, list):
|
||||
next_m["content"] = tuple(system_content) + (
|
||||
tuple(next_content)
|
||||
if isinstance(next_content, list)
|
||||
else (ChatCompletionContentPartTextParam(type="text", text=next_content),)
|
||||
)
|
||||
elif isinstance(next_content, list):
|
||||
next_m["content"] = (
|
||||
ChatCompletionContentPartTextParam(type="text", text=system_content),
|
||||
) + tuple(next_content)
|
||||
else:
|
||||
next_m["content"] = system_content + " " + next_content
|
||||
elif next_role == "system": # Next message is a system message
|
||||
# Append a user message instead of the system message
|
||||
new_message = {"role": "user", "content": m["content"]}
|
||||
|
|
|
|||
|
|
@ -55,6 +55,53 @@ def test_supports_system_message():
|
|||
assert isinstance(response, litellm.ModelResponse)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"system_content,user_content,expected_content",
|
||||
[
|
||||
(
|
||||
[{"type": "text", "text": "Follow these instructions."}],
|
||||
[{"type": "text", "text": "Hello there!"}],
|
||||
(
|
||||
{"type": "text", "text": "Follow these instructions."},
|
||||
{"type": "text", "text": "Hello there!"},
|
||||
),
|
||||
),
|
||||
(
|
||||
[{"type": "text", "text": "Follow these instructions."}],
|
||||
"Hello there!",
|
||||
(
|
||||
{"type": "text", "text": "Follow these instructions."},
|
||||
{"type": "text", "text": "Hello there!"},
|
||||
),
|
||||
),
|
||||
(
|
||||
"Follow these instructions.",
|
||||
[{"type": "text", "text": "Hello there!"}],
|
||||
(
|
||||
{"type": "text", "text": "Follow these instructions."},
|
||||
{"type": "text", "text": "Hello there!"},
|
||||
),
|
||||
),
|
||||
(
|
||||
"Follow these instructions.",
|
||||
"Hello there!",
|
||||
"Follow these instructions. Hello there!",
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_map_system_message_with_mixed_content_types(
|
||||
system_content, user_content, expected_content
|
||||
):
|
||||
messages = [
|
||||
{"role": "system", "content": system_content},
|
||||
{"role": "user", "content": user_content},
|
||||
]
|
||||
|
||||
new_messages = map_system_message_pt(messages=messages)
|
||||
|
||||
assert new_messages == [{"role": "user", "content": expected_content}]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"stop_sequence, expected_count", [("\n", 0), (["\n"], 0), (["finish_reason"], 1)]
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue