mlflow--mlflow
728 行
25 KiB
Python
728 行
25 KiB
Python
import io
|
|
from unittest import mock
|
|
|
|
import pytest
|
|
from fastapi.encoders import jsonable_encoder
|
|
|
|
from mlflow.gateway.config import (
|
|
AmazonBedrockConfig,
|
|
AWSBaseConfig,
|
|
AWSIdAndKey,
|
|
AWSRole,
|
|
EndpointConfig,
|
|
)
|
|
from mlflow.gateway.exceptions import AIGatewayException
|
|
from mlflow.gateway.providers.bedrock import AmazonBedrockModelProvider, AmazonBedrockProvider
|
|
from mlflow.gateway.schemas import chat, completions, embeddings
|
|
|
|
from tests.gateway.providers.test_anthropic import (
|
|
completions_response as anthropic_completions_response,
|
|
)
|
|
from tests.gateway.providers.test_anthropic import (
|
|
parsed_completions_response as anthropic_parsed_completions_response,
|
|
)
|
|
from tests.gateway.providers.test_cohere import completions_response as cohere_completions_response
|
|
|
|
|
|
def ai21_completion_response():
|
|
return {
|
|
"id": 1234,
|
|
"prompt": {
|
|
"text": "This is a test",
|
|
"tokens": [
|
|
{
|
|
"generatedToken": {
|
|
"token": "▁This▁is▁a",
|
|
"logprob": -7.127955436706543,
|
|
"raw_logprob": -7.127955436706543,
|
|
},
|
|
"topTokens": None,
|
|
"textRange": {"start": 0, "end": 9},
|
|
},
|
|
{
|
|
"generatedToken": {
|
|
"token": "▁test",
|
|
"logprob": -4.926638126373291,
|
|
"raw_logprob": -4.926638126373291,
|
|
},
|
|
"topTokens": None,
|
|
"textRange": {"start": 9, "end": 14},
|
|
},
|
|
],
|
|
},
|
|
"completions": [
|
|
{
|
|
"data": {
|
|
"text": "\nIt looks like",
|
|
"tokens": [
|
|
{
|
|
"generatedToken": {
|
|
"token": "<|newline|>",
|
|
"logprob": -0.021781044080853462,
|
|
"raw_logprob": -0.021781044080853462,
|
|
},
|
|
"topTokens": None,
|
|
"textRange": {"start": 0, "end": 1},
|
|
},
|
|
{
|
|
"generatedToken": {
|
|
"token": "▁It▁looks▁like",
|
|
"logprob": -3.2340049743652344,
|
|
"raw_logprob": -3.2340049743652344,
|
|
},
|
|
"topTokens": None,
|
|
"textRange": {"start": 1, "end": 14},
|
|
},
|
|
{
|
|
"generatedToken": {
|
|
"token": "<|endoftext|>",
|
|
"logprob": -0.01683046855032444,
|
|
"raw_logprob": -0.01683046855032444,
|
|
},
|
|
"topTokens": None,
|
|
"textRange": {"start": 14, "end": 14},
|
|
},
|
|
],
|
|
},
|
|
"finishReason": {"reason": "endoftext"},
|
|
}
|
|
],
|
|
}
|
|
|
|
|
|
def ai21_parsed_completion_response(mdl):
|
|
return {
|
|
"id": None,
|
|
"object": "text_completion",
|
|
"created": 1677858242,
|
|
"model": mdl,
|
|
"choices": [
|
|
{
|
|
"text": "\nIt looks like",
|
|
"index": 0,
|
|
"finish_reason": None,
|
|
}
|
|
],
|
|
"usage": {"prompt_tokens": None, "completion_tokens": None, "total_tokens": None},
|
|
}
|
|
|
|
|
|
bedrock_model_provider_fixtures = [
|
|
{
|
|
"provider": AmazonBedrockModelProvider.ANTHROPIC,
|
|
"config": {
|
|
"name": "completions",
|
|
"endpoint_type": "llm/v1/completions",
|
|
"model": {
|
|
"provider": "bedrock",
|
|
"name": "anthropic.claude-v1",
|
|
},
|
|
},
|
|
"response": anthropic_completions_response(),
|
|
"expected": anthropic_parsed_completions_response(),
|
|
"request": {"prompt": "How does a car work?", "max_tokens": 200},
|
|
"model_request": {
|
|
"max_tokens_to_sample": 200,
|
|
"prompt": "\n\nHuman: How does a car work?\n\nAssistant:",
|
|
"stop_sequences": ["\n\nHuman:"],
|
|
"anthropic_version": "bedrock-2023-05-31",
|
|
},
|
|
},
|
|
{
|
|
"provider": AmazonBedrockModelProvider.ANTHROPIC,
|
|
"config": {
|
|
"name": "completions",
|
|
"endpoint_type": "llm/v1/completions",
|
|
"model": {
|
|
"provider": "bedrock",
|
|
"name": "anthropic.claude-v2",
|
|
},
|
|
},
|
|
"response": anthropic_completions_response(),
|
|
"expected": anthropic_parsed_completions_response(),
|
|
"request": {"prompt": "How does a car work?", "max_tokens": 200},
|
|
"model_request": {
|
|
"max_tokens_to_sample": 200,
|
|
"prompt": "\n\nHuman: How does a car work?\n\nAssistant:",
|
|
"stop_sequences": ["\n\nHuman:"],
|
|
"anthropic_version": "bedrock-2023-05-31",
|
|
},
|
|
},
|
|
{
|
|
"provider": AmazonBedrockModelProvider.ANTHROPIC,
|
|
"config": {
|
|
"name": "completions",
|
|
"endpoint_type": "llm/v1/completions",
|
|
"model": {
|
|
"provider": "bedrock",
|
|
"name": "anthropic.claude-instant-v1",
|
|
},
|
|
},
|
|
"response": anthropic_completions_response(),
|
|
"expected": anthropic_parsed_completions_response(),
|
|
"request": {"prompt": "How does a car work?", "max_tokens": 200},
|
|
"model_request": {
|
|
"max_tokens_to_sample": 200,
|
|
"prompt": "\n\nHuman: How does a car work?\n\nAssistant:",
|
|
"stop_sequences": ["\n\nHuman:"],
|
|
"anthropic_version": "bedrock-2023-05-31",
|
|
},
|
|
},
|
|
{
|
|
"provider": AmazonBedrockModelProvider.AMAZON,
|
|
"config": {
|
|
"name": "completions",
|
|
"endpoint_type": "llm/v1/completions",
|
|
"model": {
|
|
"provider": "bedrock",
|
|
"name": "amazon.titan-tg1-large",
|
|
},
|
|
},
|
|
"request": {
|
|
"prompt": "This is a test",
|
|
"n": 1,
|
|
"temperature": 0.5,
|
|
"stop": ["foobar"],
|
|
"max_tokens": 1000,
|
|
},
|
|
"response": {
|
|
"results": [
|
|
{
|
|
"tokenCount": 5,
|
|
"outputText": "\nThis is a test",
|
|
"completionReason": "FINISH",
|
|
}
|
|
],
|
|
"inputTextTokenCount": 4,
|
|
},
|
|
"expected": {
|
|
"id": None,
|
|
"object": "text_completion",
|
|
"created": 1677858242,
|
|
"model": "amazon.titan-tg1-large",
|
|
"choices": [
|
|
{
|
|
"text": "\nThis is a test",
|
|
"index": 0,
|
|
"finish_reason": None,
|
|
}
|
|
],
|
|
"usage": {"prompt_tokens": None, "completion_tokens": None, "total_tokens": None},
|
|
},
|
|
"model_request": {
|
|
"inputText": "This is a test",
|
|
"textGenerationConfig": {
|
|
"temperature": 0.25,
|
|
"stopSequences": ["foobar"],
|
|
"maxTokenCount": 1000,
|
|
},
|
|
},
|
|
},
|
|
{
|
|
"provider": AmazonBedrockModelProvider.AI21,
|
|
"config": {
|
|
"name": "completions",
|
|
"endpoint_type": "llm/v1/completions",
|
|
"model": {
|
|
"provider": "bedrock",
|
|
"name": "ai21.j2-ultra",
|
|
},
|
|
},
|
|
"request": {
|
|
"prompt": "This is a test",
|
|
},
|
|
"response": ai21_completion_response(),
|
|
"expected": ai21_parsed_completion_response("ai21.j2-ultra"),
|
|
"model_request": {"prompt": "This is a test"},
|
|
},
|
|
{
|
|
"provider": AmazonBedrockModelProvider.AI21,
|
|
"config": {
|
|
"name": "completions",
|
|
"endpoint_type": "llm/v1/completions",
|
|
"model": {
|
|
"provider": "bedrock",
|
|
"name": "ai21.j2-mid",
|
|
},
|
|
},
|
|
"request": {"prompt": "This is a test", "n": 2, "max_tokens": 1000, "stop": ["foobar"]},
|
|
"response": ai21_completion_response(),
|
|
"expected": ai21_parsed_completion_response("ai21.j2-mid"),
|
|
"model_request": {
|
|
"prompt": "This is a test",
|
|
"stopSequences": ["foobar"],
|
|
"maxTokens": 1000,
|
|
"numResults": 2,
|
|
},
|
|
},
|
|
{
|
|
"provider": AmazonBedrockModelProvider.COHERE,
|
|
"config": {
|
|
"name": "completions",
|
|
"endpoint_type": "llm/v1/completions",
|
|
"model": {
|
|
"provider": "bedrock",
|
|
"name": "cohere.command",
|
|
},
|
|
},
|
|
"request": {
|
|
"prompt": "This is a test",
|
|
},
|
|
"response": cohere_completions_response(),
|
|
"expected": {},
|
|
"model_request": {},
|
|
},
|
|
]
|
|
|
|
bedrock_aws_configs = [
|
|
({"aws_region": "us-east-1"}, AWSBaseConfig),
|
|
(
|
|
{
|
|
"aws_region": "us-east-1",
|
|
"aws_access_key_id": "test-access-key-id",
|
|
"aws_secret_access_key": "test-secret-access-key",
|
|
"aws_session_token": "test-session-token",
|
|
},
|
|
AWSIdAndKey,
|
|
),
|
|
(
|
|
{
|
|
"aws_region": "us-east-1",
|
|
"aws_access_key_id": "test-access-key-id",
|
|
"aws_secret_access_key": "test-secret-access-key",
|
|
},
|
|
AWSIdAndKey,
|
|
),
|
|
({"aws_region": "us-east-1", "aws_role_arn": "test-aws-role-arn"}, AWSRole),
|
|
]
|
|
|
|
|
|
def _merge_model_and_aws_config(config, aws_config):
|
|
return {
|
|
**config,
|
|
"model": {
|
|
**config["model"],
|
|
"config": {**config["model"].get("config", {}), "aws_config": aws_config},
|
|
},
|
|
}
|
|
|
|
|
|
def _assert_any_call_at_least(mobj, *args, **kwargs):
|
|
if not mobj.call_args_list:
|
|
raise AssertionError(f"no calls to {mobj=}")
|
|
for call in mobj.call_args_list:
|
|
if all(call.kwargs.get(k) == v for k, v in kwargs.items()) and all(
|
|
call.args[i] == v for i, v in enumerate(args)
|
|
):
|
|
return
|
|
else:
|
|
raise AssertionError(f"No valid call to {mobj=} with {args=} and {kwargs=}")
|
|
|
|
|
|
def test_get_provider_name():
|
|
provider = AmazonBedrockProvider.__new__(AmazonBedrockProvider)
|
|
assert provider.DISPLAY_NAME == "Amazon Bedrock"
|
|
assert provider.get_provider_name() == "bedrock"
|
|
|
|
|
|
@pytest.mark.parametrize(("aws_config", "expected"), bedrock_aws_configs)
|
|
def test_bedrock_aws_config(aws_config, expected):
|
|
assert isinstance(
|
|
AmazonBedrockConfig.model_validate({"aws_config": aws_config}).aws_config, expected
|
|
)
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
("provider", "config"),
|
|
[(fix["provider"], fix["config"]) for fix in bedrock_model_provider_fixtures][:1],
|
|
)
|
|
@pytest.mark.parametrize("aws_config", [c for c, _ in bedrock_aws_configs])
|
|
def test_bedrock_aws_client(provider, config, aws_config):
|
|
with mock.patch("boto3.Session") as mock_session:
|
|
mock_client = mock.Mock()
|
|
mock_assume_role = mock.Mock()
|
|
mock_assume_role.return_value = mock.MagicMock()
|
|
|
|
mock_session.return_value.client = mock_client
|
|
mock_client.return_value.assume_role = mock_assume_role
|
|
|
|
provider = AmazonBedrockProvider(
|
|
EndpointConfig(**_merge_model_and_aws_config(config, aws_config))
|
|
)
|
|
provider.get_bedrock_client()
|
|
|
|
if "aws_region" in aws_config:
|
|
_assert_any_call_at_least(mock_session, region_name=aws_config["aws_region"])
|
|
|
|
if "aws_role_arn" in aws_config:
|
|
_assert_any_call_at_least(mock_client, service_name="sts")
|
|
_assert_any_call_at_least(mock_assume_role, RoleArn=aws_config["aws_role_arn"])
|
|
_assert_any_call_at_least(mock_client, service_name="bedrock-runtime")
|
|
|
|
elif {"aws_secret_access_key", "aws_access_key_id"} <= set(aws_config):
|
|
_assert_any_call_at_least(mock_client, service_name="bedrock-runtime")
|
|
_assert_any_call_at_least(
|
|
mock_client,
|
|
**{
|
|
k: v
|
|
for k, v in aws_config.items()
|
|
if k in {"aws_secret_access_key", "aws_access_key_id"}
|
|
},
|
|
)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
@pytest.mark.parametrize("aws_config", [c[0] for c in bedrock_aws_configs])
|
|
@pytest.mark.parametrize(
|
|
("provider", "config", "payload", "response", "expected", "model_request"),
|
|
[
|
|
pytest.param(
|
|
fix["provider"],
|
|
fix["config"],
|
|
fix["request"],
|
|
fix["response"],
|
|
fix["expected"],
|
|
fix["model_request"],
|
|
marks=[]
|
|
if fix["provider"] is not AmazonBedrockModelProvider.COHERE
|
|
else pytest.mark.skip("Cohere isn't available on Amazon Bedrock yet"),
|
|
)
|
|
for fix in bedrock_model_provider_fixtures
|
|
],
|
|
)
|
|
async def test_bedrock_request_response(
|
|
provider, config, payload, response, expected, model_request, aws_config
|
|
):
|
|
with (
|
|
mock.patch("time.time", return_value=1677858242),
|
|
mock.patch(
|
|
"mlflow.gateway.providers.bedrock.AmazonBedrockProvider._request", return_value=response
|
|
) as mock_request,
|
|
):
|
|
if not expected:
|
|
pytest.skip("no expected value")
|
|
|
|
expected["model"] = config["model"]["name"]
|
|
|
|
provider = AmazonBedrockProvider(
|
|
EndpointConfig(**_merge_model_and_aws_config(config, aws_config))
|
|
)
|
|
response = await provider.completions(completions.RequestPayload(**payload))
|
|
assert jsonable_encoder(response) == expected
|
|
|
|
mock_request.assert_called_once()
|
|
mock_request.assert_called_once_with(model_request)
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
("model_name", "expected"),
|
|
[
|
|
("us.anthropic.claude-3-sonnet", AmazonBedrockModelProvider.ANTHROPIC),
|
|
("apac.anthropic.claude-3-haiku", AmazonBedrockModelProvider.ANTHROPIC),
|
|
("anthropic.claude-3-5-sonnet", AmazonBedrockModelProvider.ANTHROPIC),
|
|
("ai21.jamba-1-5-large-v1:0", AmazonBedrockModelProvider.AI21),
|
|
("cohere.embed-multilingual-v3", AmazonBedrockModelProvider.COHERE),
|
|
("us.amazon.nova-premier-v1:0", AmazonBedrockModelProvider.AMAZON),
|
|
],
|
|
)
|
|
def test_amazon_bedrock_model_provider(model_name, expected):
|
|
provider = AmazonBedrockModelProvider.of_str(model_name)
|
|
assert provider == expected
|
|
|
|
|
|
# ---- Converse API tests ----
|
|
|
|
|
|
def _make_converse_provider():
|
|
"""Create a provider with a mock boto3 client for Converse API tests."""
|
|
|
|
config = {
|
|
"name": "chat",
|
|
"endpoint_type": "llm/v1/chat",
|
|
"model": {
|
|
"provider": "bedrock",
|
|
"name": "us.anthropic.claude-3-5-sonnet-20241022-v2:0",
|
|
"config": {"aws_config": {"aws_region": "us-east-1"}},
|
|
},
|
|
}
|
|
return AmazonBedrockProvider(EndpointConfig(**config))
|
|
|
|
|
|
def _converse_response():
|
|
return {
|
|
"output": {
|
|
"message": {
|
|
"role": "assistant",
|
|
"content": [{"text": "Hello from Bedrock!"}],
|
|
}
|
|
},
|
|
"stopReason": "end_turn",
|
|
"usage": {
|
|
"inputTokens": 10,
|
|
"outputTokens": 20,
|
|
"totalTokens": 30,
|
|
},
|
|
}
|
|
|
|
|
|
def _converse_response_with_tool_use():
|
|
return {
|
|
"output": {
|
|
"message": {
|
|
"role": "assistant",
|
|
"content": [
|
|
{
|
|
"toolUse": {
|
|
"toolUseId": "tool_abc123",
|
|
"name": "add",
|
|
"input": {"a": 17, "b": 25},
|
|
}
|
|
}
|
|
],
|
|
}
|
|
},
|
|
"stopReason": "tool_use",
|
|
"usage": {
|
|
"inputTokens": 30,
|
|
"outputTokens": 10,
|
|
"totalTokens": 40,
|
|
},
|
|
}
|
|
|
|
|
|
def _converse_stream_response():
|
|
return {
|
|
"stream": iter([
|
|
{"contentBlockDelta": {"delta": {"text": "Hello"}}},
|
|
{"contentBlockDelta": {"delta": {"text": " from Bedrock!"}}},
|
|
{"messageStop": {"stopReason": "end_turn"}},
|
|
{"metadata": {"usage": {"inputTokens": 10, "outputTokens": 20, "totalTokens": 30}}},
|
|
])
|
|
}
|
|
|
|
|
|
def _embeddings_invoke_response():
|
|
body = io.BytesIO(b'{"embedding": [0.1, 0.2, 0.3], "inputTextTokenCount": 5}')
|
|
return {"body": body}
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_bedrock_converse_chat():
|
|
|
|
provider = _make_converse_provider()
|
|
mock_client = mock.Mock()
|
|
mock_client.converse.return_value = _converse_response()
|
|
|
|
with mock.patch.object(provider, "get_bedrock_client", return_value=mock_client):
|
|
payload = chat.RequestPayload(
|
|
messages=[{"role": "user", "content": "Hello"}],
|
|
)
|
|
response = await provider.chat(payload)
|
|
|
|
result = jsonable_encoder(response)
|
|
assert result["choices"][0]["message"]["content"] == "Hello from Bedrock!"
|
|
assert result["choices"][0]["message"]["role"] == "assistant"
|
|
assert result["usage"]["prompt_tokens"] == 10
|
|
assert result["usage"]["completion_tokens"] == 20
|
|
mock_client.converse.assert_called_once()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_bedrock_converse_chat_stream():
|
|
|
|
provider = _make_converse_provider()
|
|
mock_client = mock.Mock()
|
|
mock_client.converse_stream.return_value = _converse_stream_response()
|
|
|
|
with mock.patch.object(provider, "get_bedrock_client", return_value=mock_client):
|
|
payload = chat.RequestPayload(
|
|
messages=[{"role": "user", "content": "Hello"}],
|
|
)
|
|
chunks = [jsonable_encoder(chunk) async for chunk in provider.chat_stream(payload)]
|
|
|
|
# Should have: 2 text deltas + 1 stop + 1 usage
|
|
assert len(chunks) == 4
|
|
assert chunks[0]["choices"][0]["delta"]["content"] == "Hello"
|
|
assert chunks[1]["choices"][0]["delta"]["content"] == " from Bedrock!"
|
|
assert chunks[2]["choices"][0]["finish_reason"] == "stop"
|
|
assert chunks[3]["usage"]["prompt_tokens"] == 10
|
|
mock_client.converse_stream.assert_called_once()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_bedrock_embeddings():
|
|
|
|
config = {
|
|
"name": "embeddings",
|
|
"endpoint_type": "llm/v1/embeddings",
|
|
"model": {
|
|
"provider": "bedrock",
|
|
"name": "amazon.titan-embed-text-v1",
|
|
"config": {"aws_config": {"aws_region": "us-east-1"}},
|
|
},
|
|
}
|
|
provider = AmazonBedrockProvider(EndpointConfig(**config))
|
|
mock_client = mock.Mock()
|
|
mock_client.invoke_model.return_value = _embeddings_invoke_response()
|
|
|
|
with mock.patch.object(provider, "get_bedrock_client", return_value=mock_client):
|
|
payload = embeddings.RequestPayload(input="Test text")
|
|
response = await provider.embeddings(payload)
|
|
|
|
result = jsonable_encoder(response)
|
|
assert result["data"][0]["embedding"] == [0.1, 0.2, 0.3]
|
|
assert result["usage"]["prompt_tokens"] == 5
|
|
mock_client.invoke_model.assert_called_once()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_bedrock_converse_with_system_message():
|
|
|
|
provider = _make_converse_provider()
|
|
mock_client = mock.Mock()
|
|
mock_client.converse.return_value = _converse_response()
|
|
|
|
with mock.patch.object(provider, "get_bedrock_client", return_value=mock_client):
|
|
payload = chat.RequestPayload(
|
|
messages=[
|
|
{"role": "system", "content": "You are helpful"},
|
|
{"role": "user", "content": "Hello"},
|
|
],
|
|
)
|
|
await provider.chat(payload)
|
|
|
|
call_kwargs = mock_client.converse.call_args.kwargs
|
|
assert call_kwargs["system"] == [{"text": "You are helpful"}]
|
|
assert len(call_kwargs["messages"]) == 1 # only user message
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_bedrock_converse_chat_with_tool_call():
|
|
|
|
provider = _make_converse_provider()
|
|
mock_client = mock.Mock()
|
|
mock_client.converse.return_value = _converse_response_with_tool_use()
|
|
|
|
with mock.patch.object(provider, "get_bedrock_client", return_value=mock_client):
|
|
payload = chat.RequestPayload(messages=[{"role": "user", "content": "add 17 and 25"}])
|
|
response = await provider.chat(payload)
|
|
|
|
result = jsonable_encoder(response)
|
|
tool_calls = result["choices"][0]["message"]["tool_calls"]
|
|
assert tool_calls[0]["function"]["name"] == "add"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_bedrock_converse_serializes_assistant_tool_call_history():
|
|
provider = _make_converse_provider()
|
|
mock_client = mock.Mock()
|
|
mock_client.converse.return_value = _converse_response()
|
|
|
|
with mock.patch.object(provider, "get_bedrock_client", return_value=mock_client):
|
|
payload = chat.RequestPayload(
|
|
messages=[
|
|
{"role": "user", "content": "Compute 17+25"},
|
|
{
|
|
"role": "assistant",
|
|
"content": None,
|
|
"tool_calls": [
|
|
{
|
|
"id": "tool_abc123",
|
|
"type": "function",
|
|
"function": {"name": "add", "arguments": '{"a": 17, "b": 25}'},
|
|
}
|
|
],
|
|
},
|
|
{"role": "tool", "tool_call_id": "tool_abc123", "content": "42"},
|
|
],
|
|
tools=[
|
|
{
|
|
"type": "function",
|
|
"function": {
|
|
"name": "add",
|
|
"description": "Add two integers.",
|
|
"parameters": {
|
|
"type": "object",
|
|
"properties": {"a": {"type": "integer"}, "b": {"type": "integer"}},
|
|
"required": ["a", "b"],
|
|
},
|
|
},
|
|
}
|
|
],
|
|
)
|
|
await provider.chat(payload)
|
|
|
|
call_kwargs = mock_client.converse.call_args.kwargs
|
|
assistant_blocks = call_kwargs["messages"][1]["content"]
|
|
tool_uses = [b["toolUse"] for b in assistant_blocks if "toolUse" in b]
|
|
assert tool_uses == [{"toolUseId": "tool_abc123", "name": "add", "input": {"a": 17, "b": 25}}]
|
|
mock_client.converse.assert_called_once()
|
|
|
|
|
|
@pytest.mark.parametrize("arguments", ["not-json", "", " "])
|
|
@pytest.mark.asyncio
|
|
async def test_bedrock_converse_rejects_invalid_assistant_tool_call_arguments(arguments):
|
|
provider = _make_converse_provider()
|
|
mock_client = mock.Mock()
|
|
mock_client.converse.return_value = _converse_response()
|
|
|
|
with mock.patch.object(provider, "get_bedrock_client", return_value=mock_client):
|
|
payload = chat.RequestPayload(
|
|
messages=[
|
|
{"role": "user", "content": "Compute 17+25"},
|
|
{
|
|
"role": "assistant",
|
|
"content": None,
|
|
"tool_calls": [
|
|
{
|
|
"id": "tool_bad_args",
|
|
"type": "function",
|
|
"function": {"name": "add", "arguments": arguments},
|
|
}
|
|
],
|
|
},
|
|
]
|
|
)
|
|
with pytest.raises(
|
|
AIGatewayException, match="Invalid assistant tool call arguments: not valid JSON"
|
|
) as exc_info:
|
|
await provider.chat(payload)
|
|
|
|
assert exc_info.value.status_code == 422
|
|
assert "tool_call_id=tool_bad_args" in exc_info.value.detail
|
|
assert "tool_name=add" in exc_info.value.detail
|
|
mock_client.converse.assert_not_called()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_bedrock_converse_rejects_assistant_tool_call_with_missing_name():
|
|
provider = _make_converse_provider()
|
|
mock_client = mock.Mock()
|
|
mock_client.converse.return_value = _converse_response()
|
|
|
|
with mock.patch.object(provider, "get_bedrock_client", return_value=mock_client):
|
|
payload = chat.RequestPayload(
|
|
messages=[
|
|
{"role": "user", "content": "Compute 17+25"},
|
|
{
|
|
"role": "assistant",
|
|
"content": None,
|
|
"tool_calls": [
|
|
{
|
|
"id": "tool_missing_name",
|
|
"type": "function",
|
|
"function": {"name": None, "arguments": '{"a": 17, "b": 25}'},
|
|
}
|
|
],
|
|
},
|
|
]
|
|
)
|
|
with pytest.raises(
|
|
AIGatewayException, match="Invalid assistant tool call: missing function name"
|
|
) as exc_info:
|
|
await provider.chat(payload)
|
|
|
|
assert exc_info.value.status_code == 422
|
|
assert "tool_call_id=tool_missing_name" in exc_info.value.detail
|
|
mock_client.converse.assert_not_called()
|