diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/mid_conversation_system.py b/litellm/llms/anthropic/experimental_pass_through/messages/mid_conversation_system.py index c4fd7bcd320..ddefec6bac9 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/mid_conversation_system.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/mid_conversation_system.py @@ -1,4 +1,5 @@ from collections.abc import Mapping, Sequence +from itertools import groupby from typing import Final CONVERTED_SYSTEM_NOTE: Final = ( @@ -39,39 +40,31 @@ def opens_with_tool_results(message: object) -> bool: ) -def system_run_before(messages: Sequence[Mapping[str, object]], index: int) -> Sequence[Mapping[str, object]]: - start: Final = next( - (j + 1 for j in range(index - 1, -1, -1) if not is_system_role_message(messages[j])), - 0, - ) - return messages[start:index] - - -def system_run_end(messages: Sequence[Mapping[str, object]], index: int) -> int: - return next( - (j for j in range(index, len(messages)) if not is_system_role_message(messages[j])), - len(messages), - ) - - -def reordered_around_tool_results( - messages: Sequence[Mapping[str, object]], index: int +def system_run_placed_after_tool_results( + system_run: Sequence[Mapping[str, object]], follower_run: Sequence[Mapping[str, object]] ) -> tuple[Mapping[str, object], ...]: - message: Final = messages[index] - if opens_with_tool_results(message): - return (message, *system_run_before(messages, index)) - if not is_system_role_message(message): - return (message,) - run_end: Final = system_run_end(messages, index) - follower: Final = messages[run_end] if run_end < len(messages) else None - return () if opens_with_tool_results(follower) else (message,) + if follower_run and opens_with_tool_results(follower_run[0]): + return (follower_run[0], *system_run, *follower_run[1:]) + return (*system_run, *follower_run) def system_turns_after_tool_results( messages: Sequence[Mapping[str, object]], ) -> tuple[Mapping[str, object], ...]: - return tuple( - message for index in range(len(messages)) for message in reordered_around_tool_results(messages, index) + runs: Final = tuple(tuple(run) for _, run in groupby(messages, key=is_system_role_message)) + if not runs: + return () + first_system_run: Final = 0 if is_system_role_message(runs[0][0]) else 1 + paired_runs: Final = tuple( + (runs[i], runs[i + 1] if i + 1 < len(runs) else ()) for i in range(first_system_run, len(runs), 2) + ) + return ( + *(runs[0] if first_system_run else ()), + *( + m + for system_run, follower_run in paired_runs + for m in system_run_placed_after_tool_results(system_run, follower_run) + ), ) diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_mid_conversation_system.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_mid_conversation_system.py index 776dbd98833..33f3f388995 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_mid_conversation_system.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_mid_conversation_system.py @@ -1,3 +1,5 @@ +import time + from litellm.llms.anthropic.experimental_pass_through.messages.mid_conversation_system import ( CONVERTED_SYSTEM_NOTE, convert_mid_conversation_system_turns, @@ -60,3 +62,19 @@ def test_convert_mid_conversation_system_turns_moves_system_after_tool_result(): assert result[1] is tool_result assert result[2]["role"] == "user" assert result[2]["content"][0]["text"] == CONVERTED_SYSTEM_NOTE + + +def test_convert_mid_conversation_system_turns_handles_long_system_run_in_linear_time(): + system_run = [{"role": "system", "content": f"reminder {i}"} for i in range(20_000)] + tool_result = { + "role": "user", + "content": [{"type": "tool_result", "tool_use_id": "toolu_1", "content": "Rainy"}], + } + + started = time.perf_counter() + result = convert_mid_conversation_system_turns([{"role": "user", "content": "hi"}, *system_run, tool_result]) + elapsed = time.perf_counter() - started + + assert elapsed < 5 + assert result[1] is tool_result + assert [m["content"][1]["text"] for m in result[2:]] == [m["content"] for m in system_run]