mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(bedrock): preserve native reason through conversions
This commit is contained in:
parent
cf2faa852a
commit
2625c7bea3
3 changed files with 83 additions and 10 deletions
|
|
@ -1,7 +1,7 @@
|
|||
import base64
|
||||
import time
|
||||
from collections.abc import Callable, Iterator, Mapping, Sequence
|
||||
from itertools import groupby
|
||||
from itertools import chain, groupby
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, TypeAlias, TypedDict, Union, cast
|
||||
|
||||
|
|
@ -369,6 +369,14 @@ class ChunkProcessor:
|
|||
# Fall back to first chunk's model if no different model found
|
||||
return first_chunk_model
|
||||
|
||||
@staticmethod
|
||||
def _choice_provider_specific_fields(chunk: "_BaseChunk") -> Mapping[str, object]:
|
||||
choices: Final = chunk.get("choices")
|
||||
if not choices:
|
||||
return MappingProxyType({})
|
||||
fields: Final = choices[0].get("provider_specific_fields")
|
||||
return fields if isinstance(fields, dict) else MappingProxyType({})
|
||||
|
||||
def build_base_response(self, chunks: Sequence["_BaseChunk"]) -> ModelResponse:
|
||||
chunk = self.first_chunk
|
||||
id: Final = ChunkProcessor._get_chunk_id(chunks)
|
||||
|
|
@ -391,14 +399,9 @@ class ChunkProcessor:
|
|||
if chunk_finish_reason is not None:
|
||||
finish_reason = chunk_finish_reason
|
||||
|
||||
choice_provider_specific_fields: Final[dict[str, object]] = { # mutable-ok: response field requires a dict
|
||||
key: value
|
||||
for chunk in chunks
|
||||
if chunk.get("choices")
|
||||
for fields in (chunk["choices"][0].get("provider_specific_fields"),)
|
||||
if isinstance(fields, dict)
|
||||
for key, value in fields.items()
|
||||
}
|
||||
choice_provider_specific_fields: Final[dict[str, object]] = dict( # mutable-ok: response field requires a dict
|
||||
chain.from_iterable(ChunkProcessor._choice_provider_specific_fields(chunk).items() for chunk in chunks)
|
||||
)
|
||||
|
||||
# Initialize the response dictionary
|
||||
response = ModelResponse(
|
||||
|
|
|
|||
|
|
@ -2658,7 +2658,8 @@ class AmazonConverseConfig(BaseConfig):
|
|||
|
||||
## HANDLE TOOL CALLS
|
||||
_message: Final = Message(**chat_completion_message)
|
||||
initial_finish_reason = completion_response["stopReason"]
|
||||
raw_finish_reason: Final = completion_response["stopReason"]
|
||||
initial_finish_reason = raw_finish_reason
|
||||
|
||||
# When json_mode filtered out all synthetic tool calls the response
|
||||
# is plain content, not a pending tool invocation. Fix finish_reason
|
||||
|
|
@ -2674,11 +2675,17 @@ class AmazonConverseConfig(BaseConfig):
|
|||
tools=optional_params.get("tools"),
|
||||
initial_finish_reason=initial_finish_reason,
|
||||
)
|
||||
choice_provider_specific_fields: Final = (
|
||||
{"native_finish_reason": raw_finish_reason} # mutable-ok: Choices requires a dict
|
||||
if returned_finish_reason != raw_finish_reason
|
||||
else None
|
||||
)
|
||||
model_response.choices = [
|
||||
litellm.Choices(
|
||||
finish_reason=returned_finish_reason,
|
||||
index=0,
|
||||
message=returned_message,
|
||||
provider_specific_fields=choice_provider_specific_fields,
|
||||
)
|
||||
]
|
||||
model_response.created = int(time.time())
|
||||
|
|
|
|||
|
|
@ -4672,6 +4672,69 @@ def test_transform_response_preserves_raw_bedrock_stop_reason(stop_reason: str):
|
|||
}
|
||||
|
||||
|
||||
def test_transform_response_preserves_raw_stop_reason_when_content_becomes_tool_call():
|
||||
response_json = {
|
||||
"metrics": {"latencyMs": 1},
|
||||
"output": {
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{
|
||||
"text": json.dumps(
|
||||
{
|
||||
"type": "function",
|
||||
"name": "lookup_weather",
|
||||
"parameters": {"city": "Seoul"},
|
||||
}
|
||||
)
|
||||
}
|
||||
],
|
||||
}
|
||||
},
|
||||
"stopReason": "stop_sequence",
|
||||
"usage": {
|
||||
"inputTokens": 10,
|
||||
"outputTokens": 1,
|
||||
"totalTokens": 11,
|
||||
},
|
||||
}
|
||||
raw_response = httpx.Response(
|
||||
200,
|
||||
json=response_json,
|
||||
request=httpx.Request("POST", "https://bedrock.test/converse"),
|
||||
)
|
||||
|
||||
result = AmazonConverseConfig().transform_response(
|
||||
model="bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0",
|
||||
raw_response=raw_response,
|
||||
model_response=ModelResponse(),
|
||||
logging_obj=None,
|
||||
request_data={},
|
||||
messages=[],
|
||||
optional_params={
|
||||
"tools": [
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "lookup_weather",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {"city": {"type": "string"}},
|
||||
},
|
||||
},
|
||||
}
|
||||
]
|
||||
},
|
||||
litellm_params={},
|
||||
encoding=None,
|
||||
)
|
||||
|
||||
assert result.choices[0].finish_reason == "tool_calls"
|
||||
assert result.choices[0].provider_specific_fields == {
|
||||
"native_finish_reason": "stop_sequence"
|
||||
}
|
||||
|
||||
|
||||
# AWS Bedrock Converse stopReason values, accessed 2026-09-16:
|
||||
# https://docs.aws.amazon.com/bedrock/latest/APIReference/API_runtime_Converse.html
|
||||
@pytest.mark.parametrize("stop_reason", ("stop_sequence", "end_turn"))
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue