mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
Merge pull request #38275 from BerriAI/litellm_together_chat_template_kwargs
fix(together_ai): strip internal thinking fields from outbound messages, keep reasoning_content
This commit is contained in:
commit
1ff615c335
2 changed files with 143 additions and 2 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
import json
|
||||
import logging
|
||||
from collections.abc import Mapping, Sequence
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import httpx
|
||||
|
|
@ -268,6 +269,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: Sequence[Mapping[str, object]]):
|
||||
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: list[httpx.Request] = []
|
||||
|
||||
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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue