diff --git a/litellm/llms/together_ai/chat/transformation.py b/litellm/llms/together_ai/chat/transformation.py index 3162a34f1b9..6bf53ee348b 100644 --- a/litellm/llms/together_ai/chat/transformation.py +++ b/litellm/llms/together_ai/chat/transformation.py @@ -4,18 +4,25 @@ Translates from OpenAI's `/v1/chat/completions` to Together AI's `/v1/chat/compl Docs: https://docs.together.ai/docs/chat-overview """ -from collections.abc import Container +from collections.abc import Container, Coroutine from types import MappingProxyType -from typing import Final +from typing import ( + Final, + Literal, + cast, # noqa: TID251 # rebuilding a TypedDict minus keys has no checked spelling + overload, +) import litellm from litellm._logging import verbose_logger from litellm.exceptions import UnsupportedParamsError +from litellm.types.llms.openai import AllMessageValues from litellm.utils import supports_function_calling from ...openai.chat.gpt_transformation import OpenAIGPTConfig TOOL_CALLING_PARAMS: Final = ("tools", "tool_choice", "function_call") +LITELLM_INTERNAL_ASSISTANT_FIELDS: Final = frozenset({"thinking_blocks", "provider_specific_fields"}) PLAIN_TEXT_RESPONSE_FORMAT: Final = MappingProxyType({"type": "text"}) FUNCTION_CALLING_DOCS_URL: Final = "https://docs.together.ai/docs/function-calling" @@ -61,7 +68,50 @@ def _tool_params_to_drop(passed_params: Container[str], model: str, drop_params: ) +def _without_litellm_internal_fields(message: AllMessageValues) -> AllMessageValues: + if message["role"] != "assistant" or LITELLM_INTERNAL_ASSISTANT_FIELDS.isdisjoint(message): + return message + return cast( # cast-ok: rebuilding the same TypedDict minus internal keys loses the narrowed type + "AllMessageValues", + { # mutable-ok: TypedDict rebuild minus internal keys + key: value for key, value in message.items() if key not in LITELLM_INTERNAL_ASSISTANT_FIELDS + }, + ) + + class TogetherAIChatConfig(OpenAIGPTConfig): + @overload + def _transform_messages( + self, + messages: list[AllMessageValues], # mutable-ok: inherited contract + model: str, + is_async: Literal[True], + ) -> Coroutine[object, object, list[AllMessageValues]]: ... # mutable-ok: inherited contract + + @overload + def _transform_messages( + self, + messages: list[AllMessageValues], # mutable-ok: inherited contract + model: str, + is_async: Literal[False] = False, + ) -> list[AllMessageValues]: ... # mutable-ok: inherited contract + + def _transform_messages( + self, + messages: list[AllMessageValues], # mutable-ok: inherited contract + model: str, + is_async: bool = False, + ) -> list[AllMessageValues] | Coroutine[object, object, list[AllMessageValues]]: # mutable-ok: inherited contract + """Together consumes replayed assistant `reasoning_content` (preserved thinking via + `chat_template_kwargs: {"clear_thinking": false}`), so it must stay in the payload; + only litellm-internal fields are stripped before sending.""" + stripped: Final = [ # mutable-ok: super() requires a list + _without_litellm_internal_fields(message) for message in messages + ] + if is_async: + return super()._transform_messages(stripped, model, is_async=True) + return super()._transform_messages(stripped, model, is_async=False) + def get_supported_openai_params(self, model: str) -> list: supports_fc: Final = _function_calling_verdict(model) supported_params: Final = super().get_supported_openai_params(model) diff --git a/tests/test_litellm/llms/together_ai/chat/test_together_ai_chat_transformation.py b/tests/test_litellm/llms/together_ai/chat/test_together_ai_chat_transformation.py index 0b9fd5364f9..110d93cbc0d 100644 --- a/tests/test_litellm/llms/together_ai/chat/test_together_ai_chat_transformation.py +++ b/tests/test_litellm/llms/together_ai/chat/test_together_ai_chat_transformation.py @@ -268,6 +268,96 @@ def test_streaming_chunk_preserves_tool_call_index_and_id(): assert continuation["function"]["arguments"] == '{"city": "San' +REPLAYED_ASSISTANT_MESSAGE = { + "role": "assistant", + "content": "The digit sum is 11.", + "reasoning_content": "The secret number is 47. 4 + 7 = 11.", + "thinking_blocks": [{"type": "thinking", "thinking": "The secret number is 47.", "signature": ""}], + "provider_specific_fields": {"thinking_blocks": [{"type": "thinking", "thinking": "The secret number is 47."}]}, +} + +PRESERVED_THINKING_MESSAGES = [ + {"role": "user", "content": "Pick a secret two-digit number and tell me only its digit sum."}, + REPLAYED_ASSISTANT_MESSAGE, + {"role": "user", "content": "What was the secret number?"}, +] + + +def _assert_internal_fields_stripped_reasoning_kept(transformed_messages: list): + assistant_message = transformed_messages[1] + assert assistant_message["reasoning_content"] == REPLAYED_ASSISTANT_MESSAGE["reasoning_content"] + assert "thinking_blocks" not in assistant_message + assert "provider_specific_fields" not in assistant_message + assert assistant_message["content"] == REPLAYED_ASSISTANT_MESSAGE["content"] + assert transformed_messages[0] == PRESERVED_THINKING_MESSAGES[0] + assert transformed_messages[2] == PRESERVED_THINKING_MESSAGES[2] + + +def test_transform_request_keeps_reasoning_content_strips_internal_fields(): + request = TogetherAIChatConfig().transform_request( + model=REASONING_MODEL, + messages=[dict(message) for message in PRESERVED_THINKING_MESSAGES], + optional_params={}, + litellm_params={"custom_llm_provider": "together_ai"}, + headers={}, + ) + + _assert_internal_fields_stripped_reasoning_kept(request["messages"]) + + +async def test_async_transform_request_keeps_reasoning_content_strips_internal_fields(): + request = await TogetherAIChatConfig().async_transform_request( + model=REASONING_MODEL, + messages=[dict(message) for message in PRESERVED_THINKING_MESSAGES], + optional_params={}, + litellm_params={"custom_llm_provider": "together_ai"}, + headers={}, + ) + + _assert_internal_fields_stripped_reasoning_kept(request["messages"]) + + +def test_completion_sends_chat_template_kwargs_and_preserved_reasoning(): + from litellm.llms.custom_httpx.http_handler import HTTPHandler + + captured_requests = [] + + def respond(request: httpx.Request) -> httpx.Response: + captured_requests.append(request) + return httpx.Response( + 200, + json={ + "id": "chatcmpl-together-preserved", + "object": "chat.completion", + "created": 1234567890, + "model": REASONING_MODEL, + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "47"}, + "finish_reason": "stop", + } + ], + "usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}, + }, + ) + + client = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(respond))) + + litellm.completion( + model=f"together_ai/{REASONING_MODEL}", + messages=[dict(message) for message in PRESERVED_THINKING_MESSAGES], + chat_template_kwargs={"clear_thinking": False}, + api_key="fake-key", + client=client, + ) + + request_body = json.loads(captured_requests[0].content) + assert request_body["chat_template_kwargs"] == {"clear_thinking": False} + assert "extra_body" not in request_body + _assert_internal_fields_stripped_reasoning_kept(request_body["messages"]) + + def test_together_ai_config_alias_points_at_chat_config(): assert litellm.TogetherAIConfig is litellm.TogetherAIChatConfig config = litellm.TogetherAIConfig(max_tokens=10)