mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-13 23:11:40 +00:00
fix(responses): reduce cyclomatic complexity and satisfy type discipline gate
This commit is contained in:
parent
7268856285
commit
9c3593d94a
2 changed files with 117 additions and 53 deletions
|
|
@ -449,16 +449,39 @@ class LiteLLMCompletionResponsesConfig:
|
|||
|
||||
return LiteLLMCompletionResponsesConfig._normalize_system_messages(messages)
|
||||
|
||||
@staticmethod
|
||||
def _extract_system_content(message: object) -> tuple[str, ...]:
|
||||
raw: Final = (
|
||||
message.get("content")
|
||||
if isinstance(message, dict)
|
||||
else (message.content if hasattr(message, "content") else None)
|
||||
)
|
||||
if isinstance(raw, str):
|
||||
return (raw,) if raw else ()
|
||||
if isinstance(raw, list):
|
||||
|
||||
def _iter_blocks() -> Iterator[str]:
|
||||
for block in raw:
|
||||
if isinstance(block, str) and block:
|
||||
yield block
|
||||
elif isinstance(block, dict):
|
||||
text: Final = block.get("text")
|
||||
if isinstance(text, str) and text:
|
||||
yield text
|
||||
|
||||
return tuple(_iter_blocks())
|
||||
return ()
|
||||
|
||||
@staticmethod
|
||||
def _normalize_system_messages(
|
||||
messages: list[
|
||||
messages: list[ # mutable-ok: input chat completion messages list
|
||||
AllMessageValues
|
||||
| GenericChatCompletionMessage
|
||||
| ChatCompletionMessageToolCall
|
||||
| ChatCompletionResponseMessage
|
||||
| Message
|
||||
],
|
||||
) -> list[
|
||||
) -> list[ # mutable-ok: output chat completion messages list
|
||||
AllMessageValues
|
||||
| GenericChatCompletionMessage
|
||||
| ChatCompletionMessageToolCall
|
||||
|
|
@ -475,58 +498,26 @@ class LiteLLMCompletionResponsesConfig:
|
|||
def _is_system(msg: object) -> bool:
|
||||
if isinstance(msg, dict):
|
||||
return msg.get("role") == "system"
|
||||
elif hasattr(msg, "role"):
|
||||
return msg.role == "system"
|
||||
return False
|
||||
return bool(hasattr(msg, "role") and msg.role == "system")
|
||||
|
||||
system_messages: list[
|
||||
AllMessageValues
|
||||
| GenericChatCompletionMessage
|
||||
| ChatCompletionMessageToolCall
|
||||
| ChatCompletionResponseMessage
|
||||
| Message
|
||||
] = [m for m in messages if _is_system(m)]
|
||||
if not system_messages:
|
||||
system_indices: Final = tuple(i for i, m in enumerate(messages) if _is_system(m))
|
||||
if not system_indices or (len(system_indices) == 1 and system_indices[0] == 0):
|
||||
return messages
|
||||
|
||||
non_system_messages: list[
|
||||
AllMessageValues
|
||||
| GenericChatCompletionMessage
|
||||
| ChatCompletionMessageToolCall
|
||||
| ChatCompletionResponseMessage
|
||||
| Message
|
||||
] = [m for m in messages if not _is_system(m)]
|
||||
non_system: Final = tuple(m for i, m in enumerate(messages) if i not in system_indices)
|
||||
if len(system_indices) == 1:
|
||||
return [messages[system_indices[0]], *non_system] # mutable-ok: chat completion messages list
|
||||
|
||||
if len(system_messages) == 1:
|
||||
if messages and _is_system(messages[0]):
|
||||
return messages
|
||||
return [system_messages[0]] + non_system_messages
|
||||
|
||||
merged_content_parts: list[str] = []
|
||||
for sm in system_messages:
|
||||
raw_content: object = None
|
||||
if isinstance(sm, dict):
|
||||
raw_content = sm.get("content")
|
||||
elif hasattr(sm, "content"):
|
||||
raw_content = sm.content
|
||||
|
||||
if isinstance(raw_content, str):
|
||||
if raw_content:
|
||||
merged_content_parts.append(raw_content)
|
||||
elif isinstance(raw_content, list):
|
||||
for block in raw_content:
|
||||
if isinstance(block, str) and block:
|
||||
merged_content_parts.append(block)
|
||||
elif isinstance(block, dict):
|
||||
text = block.get("text")
|
||||
if isinstance(text, str) and text:
|
||||
merged_content_parts.append(text)
|
||||
|
||||
merged_system_message = ChatCompletionSystemMessage(
|
||||
role="system",
|
||||
content="\n\n".join(merged_content_parts),
|
||||
merged_parts: Final = tuple(
|
||||
part
|
||||
for idx in system_indices
|
||||
for part in LiteLLMCompletionResponsesConfig._extract_system_content(messages[idx])
|
||||
)
|
||||
return [merged_system_message] + non_system_messages
|
||||
merged_system: Final = ChatCompletionSystemMessage(
|
||||
role="system",
|
||||
content="\n\n".join(merged_parts),
|
||||
)
|
||||
return [merged_system, *non_system] # mutable-ok: chat completion messages list
|
||||
|
||||
@staticmethod
|
||||
async def async_responses_api_session_handler(
|
||||
|
|
|
|||
|
|
@ -40,10 +40,7 @@ def test_reproduce_issue_40693_non_leading_system_message() -> None:
|
|||
assert messages[1]["role"] == "user"
|
||||
|
||||
# 2. No non-leading system messages
|
||||
assert all(
|
||||
(m.get("role") if isinstance(m, dict) else getattr(m, "role", None)) != "system"
|
||||
for m in messages[1:]
|
||||
)
|
||||
assert all((m.get("role") if isinstance(m, dict) else getattr(m, "role", None)) != "system" for m in messages[1:])
|
||||
|
||||
# 3. Content from both instructions and subsequent system message are preserved
|
||||
system_content = messages[0]["content"]
|
||||
|
|
@ -168,3 +165,79 @@ def test_transform_responses_api_request_to_chat_completion_request_normalizes_s
|
|||
assert "Follow instructions" in messages[0]["content"]
|
||||
assert messages[1]["role"] == "user"
|
||||
assert messages[1]["content"] == "Hello"
|
||||
|
||||
|
||||
def test_system_message_with_list_of_strings_and_empty_content() -> None:
|
||||
"""
|
||||
Ensures list of strings and empty strings are handled properly in content extraction.
|
||||
"""
|
||||
input_items = [
|
||||
{"role": "system", "content": ["Line 1", "", "Line 2"]},
|
||||
{"role": "system", "content": ""},
|
||||
{"role": "system", "content": None},
|
||||
{"role": "user", "content": "Query"},
|
||||
]
|
||||
messages = LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages(
|
||||
input=input_items,
|
||||
responses_api_request={},
|
||||
)
|
||||
assert len(messages) == 2
|
||||
assert messages[0]["role"] == "system"
|
||||
assert messages[0]["content"] == "Line 1\n\nLine 2"
|
||||
assert messages[1]["role"] == "user"
|
||||
|
||||
|
||||
def test_system_message_object_with_attributes() -> None:
|
||||
"""
|
||||
Ensures messages that are objects with .role and .content attributes (not dicts) are handled.
|
||||
"""
|
||||
|
||||
class ObjMessage:
|
||||
def __init__(self, role: str, content: Any) -> None:
|
||||
self.role = role
|
||||
self.content = content
|
||||
|
||||
input_items = [
|
||||
ObjMessage(role="user", content="Hello from user"),
|
||||
ObjMessage(role="system", content="System instruction from obj"),
|
||||
]
|
||||
normalized = LiteLLMCompletionResponsesConfig._normalize_system_messages(input_items) # type: ignore[arg-type]
|
||||
assert len(normalized) == 2
|
||||
assert normalized[0].role == "system" # type: ignore[union-attr]
|
||||
assert normalized[0].content == "System instruction from obj" # type: ignore[union-attr]
|
||||
assert normalized[1].role == "user" # type: ignore[union-attr]
|
||||
|
||||
|
||||
def test_multiple_system_message_objects_merged() -> None:
|
||||
"""
|
||||
Ensures multiple object-based system messages are extracted and merged into a single system message.
|
||||
"""
|
||||
|
||||
class ObjMessage:
|
||||
def __init__(self, role: str, content: Any) -> None:
|
||||
self.role = role
|
||||
self.content = content
|
||||
|
||||
input_items = [
|
||||
ObjMessage(role="system", content="System part A"),
|
||||
ObjMessage(role="user", content="User prompt"),
|
||||
ObjMessage(role="system", content="System part B"),
|
||||
]
|
||||
normalized = LiteLLMCompletionResponsesConfig._normalize_system_messages(input_items) # type: ignore[arg-type]
|
||||
assert len(normalized) == 2
|
||||
assert normalized[0]["role"] == "system"
|
||||
assert normalized[0]["content"] == "System part A\n\nSystem part B"
|
||||
assert normalized[1].role == "user" # type: ignore[union-attr]
|
||||
|
||||
|
||||
def test_extract_system_content_edge_cases() -> None:
|
||||
"""
|
||||
Directly tests _extract_system_content edge cases including non-string/non-list content.
|
||||
"""
|
||||
assert LiteLLMCompletionResponsesConfig._extract_system_content({"content": None}) == ()
|
||||
assert LiteLLMCompletionResponsesConfig._extract_system_content({"content": 12345}) == ()
|
||||
assert LiteLLMCompletionResponsesConfig._extract_system_content({"content": "hello"}) == ("hello",)
|
||||
assert LiteLLMCompletionResponsesConfig._extract_system_content({"content": ["a", "b"]}) == ("a", "b")
|
||||
assert LiteLLMCompletionResponsesConfig._extract_system_content(
|
||||
{"content": [{"text": "t1"}, {"other": "none"}]}
|
||||
) == ("t1",)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue