fix(together_ai): strip internal thinking fields from outbound messages, pin chat_template_kwargs passthrough

This commit is contained in:
mateo-berri 2026-08-25 15:48:58 -07:00
parent e4ff44f623
commit 1ba66b1555
2 changed files with 142 additions and 2 deletions

View file

@ -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)

View file

@ -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)