项目文件夹

文件
2026-07-13 13:22:34 +08:00

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