diff --git a/litellm/litellm_core_utils/prompt_templates/common_utils.py b/litellm/litellm_core_utils/prompt_templates/common_utils.py index 4ba7c3966c0..1e96e20a03b 100644 --- a/litellm/litellm_core_utils/prompt_templates/common_utils.py +++ b/litellm/litellm_core_utils/prompt_templates/common_utils.py @@ -2261,6 +2261,37 @@ def system_messages_first( ] +def _system_content_as_text_parts(content: object) -> tuple[object, ...]: + if isinstance(content, str): + return (ChatCompletionTextObject(type="text", text=content),) + return tuple(cast(Sequence[object], content)) # cast-ok: non-str system content is a list of content parts + + +def _merge_system_message_run(run: Sequence[AllMessageValues]) -> AllMessageValues: + if len(run) == 1: + return run[0] + contents: Final = tuple(content for content in (message.get("content") for message in run) if content is not None) + if not contents: + return run[0] + if all(isinstance(content, str) for content in contents): + joined_text: Final = "\n\n".join(cast(tuple[str, ...], contents)) # cast-ok: every content is a str + return cast(AllMessageValues, {**run[0], "content": joined_text}) # cast-ok: dict spread keeps message shape + merged_parts: Final = [ # mutable-ok: chat message content must stay a json list + part for content in contents for part in _system_content_as_text_parts(content) + ] + return cast(AllMessageValues, {**run[0], "content": merged_parts}) # cast-ok: dict spread keeps message shape + + +def merge_consecutive_system_messages( + messages: list[AllMessageValues], # mutable-ok: message pipelines type messages as mutable lists +) -> list[AllMessageValues]: # mutable-ok: message pipelines type messages as mutable lists + return [ # mutable-ok: pipelines mutate message lists + merged + for is_system_run, run in groupby(messages, key=lambda message: message.get("role") == "system") + for merged in ((_merge_system_message_run(tuple(run)),) if is_system_run else run) + ] + + def _attempt_json_repair(s: str) -> object | None: """ Attempt to repair truncated JSON produced by LLM tool calls. diff --git a/litellm/llms/databricks/chat/transformation.py b/litellm/llms/databricks/chat/transformation.py index dd257cd68b0..30144d29f51 100644 --- a/litellm/llms/databricks/chat/transformation.py +++ b/litellm/llms/databricks/chat/transformation.py @@ -16,6 +16,7 @@ from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response impo ) from litellm.litellm_core_utils.prompt_templates.common_utils import ( _extract_reasoning_content, # pyright: ignore[reportPrivateUsage] # same import as the OpenAI transformation + merge_consecutive_system_messages, strip_litellm_internal_message_fields, strip_name_from_message, ) @@ -465,7 +466,9 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig): new_messages.append(_message) if "claude" not in model: - new_messages = _split_parallel_tool_calls(cast(list[AllMessageValues], new_messages)) + new_messages = _split_parallel_tool_calls( + merge_consecutive_system_messages(cast(list[AllMessageValues], new_messages)) + ) if is_async: return super()._transform_messages(messages=new_messages, model=model, is_async=cast(Literal[True], True)) diff --git a/litellm/responses/main.py b/litellm/responses/main.py index 5a4a08b760c..c5032536df4 100644 --- a/litellm/responses/main.py +++ b/litellm/responses/main.py @@ -433,13 +433,18 @@ def _bridges_to_chat_completions( return responses_api_provider_config is None or use_chat_completions_api is True +_RESPONSES_ONLY_REQUEST_FIELDS_NEVER_BRIDGED: Final = frozenset({"client_metadata"}) + + def _bridge_kwargs( kwargs: Mapping[str, object], responses_api_provider_config: BaseResponsesAPIConfig | None, allowed_openai_params: Sequence[str] | None, ) -> Mapping[str, object]: if responses_api_provider_config is None: - return kwargs + return MappingProxyType( + {key: value for key, value in kwargs.items() if key not in _RESPONSES_ONLY_REQUEST_FIELDS_NEVER_BRIDGED} + ) forwarded_keys: Final = frozenset( ( *litellm.OPENAI_CHAT_COMPLETION_PARAMS, @@ -448,7 +453,7 @@ def _bridge_kwargs( *GenericLiteLLMParams.model_fields, *(allowed_openai_params or ()), ) - ) + ).difference(_RESPONSES_ONLY_REQUEST_FIELDS_NEVER_BRIDGED) return MappingProxyType({key: value for key, value in kwargs.items() if key in forwarded_keys}) diff --git a/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py b/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py index c67f72680a8..62cf5680266 100644 --- a/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py +++ b/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py @@ -20,6 +20,7 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import ( handle_any_messages_to_chat_completion_str_messages_conversion, hoist_images_from_tool_messages, is_encrypted_reasoning_block, + merge_consecutive_system_messages, responses_reasoning_items_from_thinking_blocks, split_concatenated_json_objects, strip_encrypted_reasoning_from_messages, @@ -1846,3 +1847,95 @@ class TestEncryptedReasoningReplay: strip_encrypted_reasoning_from_messages(messages) assert messages == before + + +class TestMergeConsecutiveSystemMessages: + def test_merges_each_run_of_string_system_messages_with_a_blank_line(self): + messages = [ + {"role": "system", "content": "You are terse.", "cache_control": {"type": "ephemeral"}}, + {"role": "system", "content": "Skills: none."}, + {"role": "user", "content": "Hello"}, + {"role": "assistant", "content": "Hi"}, + {"role": "system", "content": "Reminder A"}, + {"role": "system", "content": "Reminder B"}, + {"role": "user", "content": "Bye"}, + ] + + merged = merge_consecutive_system_messages(messages) + + assert merged == [ + {"role": "system", "content": "You are terse.\n\nSkills: none.", "cache_control": {"type": "ephemeral"}}, + {"role": "user", "content": "Hello"}, + {"role": "assistant", "content": "Hi"}, + {"role": "system", "content": "Reminder A\n\nReminder B"}, + {"role": "user", "content": "Bye"}, + ] + + def test_merges_into_text_parts_when_any_system_content_is_a_list(self): + cached_part = {"type": "text", "text": "Skills: none.", "cache_control": {"type": "ephemeral"}} + messages = [ + {"role": "system", "content": "You are terse."}, + {"role": "system", "content": [cached_part, {"type": "text", "text": "Be brief."}]}, + {"role": "system", "content": "Answer in English."}, + {"role": "user", "content": "Hello"}, + ] + + merged = merge_consecutive_system_messages(messages) + + assert merged == [ + { + "role": "system", + "content": [ + {"type": "text", "text": "You are terse."}, + cached_part, + {"type": "text", "text": "Be brief."}, + {"type": "text", "text": "Answer in English."}, + ], + }, + {"role": "user", "content": "Hello"}, + ] + assert merged[0]["content"][1] is cached_part + + @pytest.mark.parametrize( + "messages", + [ + [{"role": "system", "content": "You are terse."}, {"role": "user", "content": "Hello"}], + [{"role": "user", "content": "Hello"}, {"role": "assistant", "content": "Hi"}], + [ + {"role": "system", "content": "You are terse."}, + {"role": "user", "content": "Hello"}, + {"role": "system", "content": "Reminder"}, + ], + [], + ], + ids=["single-system", "no-system", "separated-systems", "empty"], + ) + def test_leaves_messages_without_consecutive_system_messages_untouched(self, messages): + before = copy.deepcopy(messages) + + merged = merge_consecutive_system_messages(messages) + + assert merged == before + assert [message is original for message, original in zip(merged, messages)] == [True] * len(messages) + + @pytest.mark.parametrize( + ("messages", "expected_content"), + [ + ([{"role": "system"}, {"role": "system", "content": "Skills: none."}], "Skills: none."), + ([{"role": "system", "content": "You are terse."}, {"role": "system"}], "You are terse."), + ( + [{"role": "system"}, {"role": "system", "content": [{"type": "text", "text": "Be brief."}]}], + [{"type": "text", "text": "Be brief."}], + ), + ], + ids=["missing-then-str", "str-then-missing", "missing-then-list"], + ) + def test_skips_system_messages_without_content_when_merging(self, messages, expected_content): + merged = merge_consecutive_system_messages([*messages, {"role": "user", "content": "Hello"}]) + + assert merged == [{"role": "system", "content": expected_content}, {"role": "user", "content": "Hello"}] + + def test_keeps_the_first_message_when_no_system_message_in_the_run_has_content(self): + merged = merge_consecutive_system_messages([{"role": "system"}, {"role": "system"}, {"role": "user", "content": "Hi"}]) + + assert merged == [{"role": "system"}, {"role": "user", "content": "Hi"}] diff --git a/tests/test_litellm/llms/databricks/chat/__init__.py b/tests/test_litellm/llms/databricks/chat/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/llms/databricks/chat/test_databricks_chat_transformation.py b/tests/test_litellm/llms/databricks/chat/test_databricks_chat_transformation.py new file mode 100644 index 00000000000..a3391a2c585 --- /dev/null +++ b/tests/test_litellm/llms/databricks/chat/test_databricks_chat_transformation.py @@ -0,0 +1,79 @@ +import json +from typing import Final + +import httpx +import respx + +import litellm + + +def test_completion_merges_leading_system_and_developer_messages_for_chat_template_models( + respx_mock: respx.MockRouter, +): + upstream: Final = respx_mock.post("https://example.databricks.test/serving-endpoints/chat/completions").mock( + return_value=httpx.Response( + status_code=200, + json={ + "id": "chatcmpl-123", + "object": "chat.completion", + "created": 1677652288, + "model": "my-custom-model", + "choices": [{"index": 0, "message": {"role": "assistant", "content": "Answer"}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 9, "completion_tokens": 1, "total_tokens": 10}, + }, + ) + ) + + response: Final = litellm.completion( + model="databricks/my-custom-model", + messages=[ + {"role": "system", "content": "You are terse."}, + {"role": "developer", "content": "Skills: none."}, + {"role": "user", "content": "Hello"}, + ], + api_base="https://example.databricks.test/serving-endpoints", + api_key="fake-databricks-api-key", + num_retries=0, + ) + + assert upstream.call_count == 1 + request_body: Final = json.loads(upstream.calls[0].request.read()) + assert request_body["messages"] == [ + {"role": "system", "content": "You are terse.\n\nSkills: none."}, + {"role": "user", "content": "Hello"}, + ] + assert response.choices[0].message.content == "Answer" + + +def test_completion_merges_system_messages_when_one_has_empty_content(respx_mock: respx.MockRouter): + upstream: Final = respx_mock.post("https://example.databricks.test/serving-endpoints/chat/completions").mock( + return_value=httpx.Response( + status_code=200, + json={ + "id": "chatcmpl-123", + "object": "chat.completion", + "created": 1677652288, + "model": "my-custom-model", + "choices": [{"index": 0, "message": {"role": "assistant", "content": "Answer"}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 9, "completion_tokens": 1, "total_tokens": 10}, + }, + ) + ) + + litellm.completion( + model="databricks/my-custom-model", + messages=[ + {"role": "system", "content": "You are terse."}, + {"role": "system", "content": ""}, + {"role": "user", "content": "Hello"}, + ], + api_base="https://example.databricks.test/serving-endpoints", + api_key="fake-databricks-api-key", + num_retries=0, + ) + + request_body: Final = json.loads(upstream.calls[0].request.read()) + assert request_body["messages"] == [ + {"role": "system", "content": "You are terse."}, + {"role": "user", "content": "Hello"}, + ] diff --git a/tests/test_litellm/responses/test_responses_api_bridge_flag.py b/tests/test_litellm/responses/test_responses_api_bridge_flag.py index 2f64cc8debc..16135106b41 100644 --- a/tests/test_litellm/responses/test_responses_api_bridge_flag.py +++ b/tests/test_litellm/responses/test_responses_api_bridge_flag.py @@ -262,6 +262,126 @@ class TestUseResponsesApiBridgeFlag: assert request_body["messages"] == [{"role": "user", "content": "Hello"}] assert response.output[0].content[0].text == "Answer" + def test_bridge_drops_client_metadata_even_when_allowed_openai_params_names_it( + self, respx_mock: respx.MockRouter + ): + upstream: Final = respx_mock.post("https://api.openai.com/v1/chat/completions").mock( + return_value=httpx.Response( + status_code=200, + json={ + "id": "chatcmpl-123", + "object": "chat.completion", + "created": 1677652288, + "model": "my-custom-model", + "choices": [ + {"index": 0, "message": {"role": "assistant", "content": "Answer"}, "finish_reason": "stop"} + ], + "usage": {"prompt_tokens": 9, "completion_tokens": 1, "total_tokens": 10}, + }, + ) + ) + + response: Final = litellm.responses( + model="openai/my-custom-model", + input="Hello", + use_chat_completions_api=True, + allowed_openai_params=["client_metadata"], + client_metadata={"turn_id": "turn-1", "thread_id": "thread-1"}, + api_key="fake-provider-api-key", + num_retries=0, + ) + + assert upstream.call_count == 1 + request_body: Final = json.loads(upstream.calls[0].request.read()) + assert "client_metadata" not in request_body + assert request_body["messages"] == [{"role": "user", "content": "Hello"}] + assert response.output[0].content[0].text == "Answer" + + def test_bridge_merges_instructions_and_developer_input_for_databricks(self, respx_mock: respx.MockRouter): + upstream: Final = respx_mock.post("https://example.databricks.test/serving-endpoints/chat/completions").mock( + return_value=httpx.Response( + status_code=200, + json={ + "id": "chatcmpl-123", + "object": "chat.completion", + "created": 1677652288, + "model": "my-custom-model", + "choices": [ + {"index": 0, "message": {"role": "assistant", "content": "Answer"}, "finish_reason": "stop"} + ], + "usage": {"prompt_tokens": 9, "completion_tokens": 1, "total_tokens": 10}, + }, + ) + ) + + response: Final = litellm.responses( + model="databricks/my-custom-model", + instructions="You are terse.", + input=[ + {"role": "developer", "content": [{"type": "input_text", "text": "Skills: none."}]}, + {"role": "user", "content": [{"type": "input_text", "text": "Hello"}]}, + ], + client_metadata={"turn_id": "turn-1", "thread_id": "thread-1"}, + chat_template_kwargs={"thinking": True}, + api_base="https://example.databricks.test/serving-endpoints", + api_key="fake-databricks-api-key", + num_retries=0, + ) + + assert upstream.call_count == 1 + request_body: Final = json.loads(upstream.calls[0].request.read()) + assert request_body["messages"] == [ + { + "role": "system", + "content": [{"type": "text", "text": "You are terse."}, {"type": "text", "text": "Skills: none."}], + }, + {"role": "user", "content": [{"type": "text", "text": "Hello"}]}, + ] + assert "client_metadata" not in request_body + assert request_body["chat_template_kwargs"] == {"thinking": True} + assert response.output[0].content[0].text == "Answer" + + def test_bridge_drops_client_metadata_for_provider_without_native_config(self, respx_mock: respx.MockRouter): + upstream: Final = respx_mock.post("https://example.databricks.test/serving-endpoints/chat/completions").mock( + return_value=httpx.Response( + status_code=200, + json={ + "id": "chatcmpl-123", + "object": "chat.completion", + "created": 1677652288, + "model": "my-custom-model", + "choices": [ + {"index": 0, "message": {"role": "assistant", "content": "Answer"}, "finish_reason": "stop"} + ], + "usage": {"prompt_tokens": 9, "completion_tokens": 1, "total_tokens": 10}, + }, + ) + ) + + response: Final = litellm.responses( + model="databricks/my-custom-model", + input="Hello", + client_metadata={ + "turn_id": "turn-1", + "thread_id": "thread-1", + "session_id": "session-1", + "root_turn_id": "turn-1", + "x-codex-installation-id": "install-1", + "x-codex-turn-metadata": '{"turn_id":"turn-1"}', + }, + chat_template_kwargs={"thinking": True}, + api_base="https://example.databricks.test/serving-endpoints", + api_key="fake-databricks-api-key", + num_retries=0, + ) + + assert upstream.call_count == 1 + request_body: Final = json.loads(upstream.calls[0].request.read()) + assert "client_metadata" not in request_body + assert request_body["chat_template_kwargs"] == {"thinking": True} + assert request_body["messages"] == [{"role": "user", "content": "Hello"}] + assert response.output[0].content[0].text == "Answer" + def test_bridge_keeps_deployment_credentials_while_dropping_unknown_params(self, respx_mock: respx.MockRouter): upstream: Final = respx_mock.post( "https://example-resource.openai.azure.com/openai/deployments/my-deployment/chat/completions",