Updated bedrock invoke transform for mistral to support pixtral models (#10439)

This commit is contained in:
Anibal Angulo 2025-05-15 00:08:37 -06:00 committed by GitHub
parent 044f5f973b
commit 0434ca1781
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 60 additions and 5 deletions

View file

@ -1,10 +1,14 @@
import types
from typing import List, Optional
from typing import List, Optional, TYPE_CHECKING
from litellm.llms.base_llm.chat.transformation import BaseConfig
from litellm.llms.bedrock.chat.invoke_transformations.base_invoke_transformation import (
AmazonInvokeConfig,
)
from litellm.llms.bedrock.common_utils import BedrockError
if TYPE_CHECKING:
from litellm.types.utils import ModelResponse
class AmazonMistralConfig(AmazonInvokeConfig, BaseConfig):
@ -81,3 +85,27 @@ class AmazonMistralConfig(AmazonInvokeConfig, BaseConfig):
if k == "stream":
optional_params["stream"] = v
return optional_params
@staticmethod
def get_outputText(completion_response: dict, model_response: "ModelResponse") -> str:
"""This function extracts the output text from a bedrock mistral completion.
As a side effect, it updates the finish reason for a model response.
Args:
completion_response: JSON from the completion.
model_response: ModelResponse
Returns:
A string with the response of the LLM
"""
if "choices" in completion_response:
outputText = completion_response["choices"][0]["message"]["content"]
model_response.choices[0].finish_reason = completion_response["choices"][0]["finish_reason"]
elif "outputs" in completion_response:
outputText = completion_response["outputs"][0]["text"]
model_response.choices[0].finish_reason = completion_response["outputs"][0]["stop_reason"]
else:
raise BedrockError(message="Unexpected mistral completion response", status_code=400)
return outputText

View file

@ -323,10 +323,7 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
elif provider == "meta" or provider == "llama" or provider == "deepseek_r1":
outputText = completion_response["generation"]
elif provider == "mistral":
outputText = completion_response["outputs"][0]["text"]
model_response.choices[0].finish_reason = completion_response[
"outputs"
][0]["stop_reason"]
outputText = litellm.AmazonMistralConfig.get_outputText(completion_response, model_response)
else: # amazon titan
outputText = completion_response.get("results")[0].get("outputText")
except Exception as e:

View file

@ -0,0 +1,30 @@
from litellm.llms.bedrock.chat.invoke_transformations.amazon_mistral_transformation import AmazonMistralConfig
from litellm.types.utils import ModelResponse
def test_mistral_get_outputText():
# Set initial model response with arbitrary finish reason
model_response = ModelResponse()
model_response.choices[0].finish_reason = "None"
# Models like pixtral will return a completion with the openai format.
mock_json_with_choices = {"choices": [{"message": {"content": "Hello!"}, "finish_reason": "stop"}]}
outputText = AmazonMistralConfig.get_outputText(
completion_response=mock_json_with_choices, model_response=model_response
)
assert outputText == "Hello!"
assert model_response.choices[0].finish_reason == "stop"
# Other models might return a completion behind "outputs"
mock_json_with_output = {"outputs": [{"text": "Hi!", "stop_reason": "finish"}]}
outputText = AmazonMistralConfig.get_outputText(
completion_response=mock_json_with_output, model_response=model_response
)
assert outputText == "Hi!"
assert model_response.choices[0].finish_reason == "finish"