Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 9 additions & 1 deletion strands-py/src/strands/models/anthropic.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,14 @@

T = TypeVar("T", bound=BaseModel)

_IMAGE_FORMAT_MIME_TYPES: dict[str, str] = {
"png": "image/png",
"jpeg": "image/jpeg",
"jpg": "image/jpeg",
"gif": "image/gif",
"webp": "image/webp",
}


class AnthropicModel(Model):
"""Anthropic model provider implementation."""
Expand Down Expand Up @@ -134,7 +142,7 @@ def _format_request_message_content(self, content: ContentBlock) -> dict[str, An
return {
"source": {
"data": base64.b64encode(content["image"]["source"]["bytes"]).decode("utf-8"),
"media_type": mimetypes.types_map.get(f".{content['image']['format']}", "application/octet-stream"),
"media_type": _IMAGE_FORMAT_MIME_TYPES.get(content["image"]["format"], "application/octet-stream"),
"type": "base64",
},
"type": "image",
Expand Down
10 changes: 6 additions & 4 deletions strands-py/src/strands/models/bedrock.py
Original file line number Diff line number Diff line change
Expand Up @@ -498,11 +498,12 @@ def _format_bedrock_messages(self, messages: Messages) -> list[dict[str, Any]]:
if formatted_content is None:
continue

# Wrap text or image content in guardContent if this is the last user text/image message
if idx == last_user_text_idx and ("text" in formatted_content or "image" in formatted_content):
# Wrap text or image content in guardContent if this is the last user text/image message.
# Bedrock guardContent only supports png and jpeg images; skip wrapping for other formats.
if idx == last_user_text_idx:
if "text" in formatted_content:
formatted_content = {"guardContent": {"text": {"text": formatted_content["text"]}}}
elif "image" in formatted_content:
elif "image" in formatted_content and formatted_content["image"].get("format") in ("png", "jpeg"):
formatted_content = {"guardContent": {"image": formatted_content["image"]}}

cleaned_content.append(formatted_content)
Expand Down Expand Up @@ -710,7 +711,8 @@ def _format_request_message_content(self, content: ContentBlock) -> dict[str, An
return None
elif "bytes" in source:
formatted_video_source = {"bytes": source["bytes"]}
result = {"format": video["format"], "source": formatted_video_source}
video_format = "three_gp" if video["format"] == "3gp" else video["format"]
result = {"format": video_format, "source": formatted_video_source}
return {"video": result}

# https://docs.aws.amazon.com/bedrock/latest/APIReference/API_runtime_CitationsContentBlock.html
Expand Down
10 changes: 9 additions & 1 deletion strands-py/src/strands/models/gemini.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,14 @@

T = TypeVar("T", bound=pydantic.BaseModel)

_IMAGE_FORMAT_MIME_TYPES: dict[str, str] = {
"png": "image/png",
"jpeg": "image/jpeg",
"jpg": "image/jpeg",
"gif": "image/gif",
"webp": "image/webp",
}


class GeminiModel(Model):
"""Google Gemini model provider implementation.
Expand Down Expand Up @@ -170,7 +178,7 @@ def _format_request_content_part(
return genai.types.Part(
inline_data=genai.types.Blob(
data=content["image"]["source"]["bytes"],
mime_type=mimetypes.types_map.get(f".{content['image']['format']}", "application/octet-stream"),
mime_type=_IMAGE_FORMAT_MIME_TYPES.get(content["image"]["format"], "application/octet-stream"),
),
)

Expand Down
11 changes: 9 additions & 2 deletions strands-py/src/strands/models/llamaapi.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,6 @@
import base64
import json
import logging
import mimetypes
from collections.abc import AsyncGenerator
from typing import Any, TypeVar, cast

Expand All @@ -27,6 +26,14 @@

T = TypeVar("T", bound=BaseModel)

_IMAGE_FORMAT_MIME_TYPES: dict[str, str] = {
"png": "image/png",
"jpeg": "image/jpeg",
"jpg": "image/jpeg",
"gif": "image/gif",
"webp": "image/webp",
}


class LlamaAPIModel(Model):
"""Llama API model provider implementation."""
Expand Down Expand Up @@ -105,7 +112,7 @@ def _format_request_message_content(self, content: ContentBlock) -> dict[str, An
TypeError: If the content block type cannot be converted to a LlamaAPI-compatible format.
"""
if "image" in content:
mime_type = mimetypes.types_map.get(f".{content['image']['format']}", "application/octet-stream")
mime_type = _IMAGE_FORMAT_MIME_TYPES.get(content["image"]["format"], "application/octet-stream")
image_data = base64.b64encode(content["image"]["source"]["bytes"]).decode("utf-8")

return {
Expand Down
10 changes: 9 additions & 1 deletion strands-py/src/strands/models/llamacpp.py
Original file line number Diff line number Diff line change
Expand Up @@ -36,6 +36,14 @@

T = TypeVar("T", bound=BaseModel)

_IMAGE_FORMAT_MIME_TYPES: dict[str, str] = {
"png": "image/png",
"jpeg": "image/jpeg",
"jpg": "image/jpeg",
"gif": "image/gif",
"webp": "image/webp",
}


class LlamaCppModel(Model):
"""llama.cpp model provider implementation.
Expand Down Expand Up @@ -219,7 +227,7 @@ def _format_message_content(self, content: ContentBlock | dict[str, Any]) -> dic
}

if "image" in content:
mime_type = mimetypes.types_map.get(f".{content['image']['format']}", "application/octet-stream")
mime_type = _IMAGE_FORMAT_MIME_TYPES.get(content["image"]["format"], "application/octet-stream")
image_data = base64.b64encode(content["image"]["source"]["bytes"]).decode("utf-8")
return {
"image_url": {
Expand Down
10 changes: 9 additions & 1 deletion strands-py/src/strands/models/openai.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,14 @@

T = TypeVar("T", bound=BaseModel)

_IMAGE_FORMAT_MIME_TYPES: dict[str, str] = {
"png": "image/png",
"jpeg": "image/jpeg",
"jpg": "image/jpeg",
"gif": "image/gif",
"webp": "image/webp",
}

# Alternative context overflow error messages
# These are commonly returned by OpenAI-compatible endpoints wrapping other providers
# (e.g., Databricks serving Bedrock models)
Expand Down Expand Up @@ -183,7 +191,7 @@ def format_request_message_content(cls, content: ContentBlock, **kwargs: Any) ->
}

if "image" in content:
mime_type = mimetypes.types_map.get(f".{content['image']['format']}", "application/octet-stream")
mime_type = _IMAGE_FORMAT_MIME_TYPES.get(content["image"]["format"], "application/octet-stream")
image_data = base64.b64encode(content["image"]["source"]["bytes"]).decode("utf-8")

return {
Expand Down
13 changes: 12 additions & 1 deletion strands-py/src/strands/models/openai_responses.py
Original file line number Diff line number Diff line change
Expand Up @@ -75,6 +75,14 @@
_CONTEXT_WINDOW_OVERFLOW_MSG = "OpenAI Responses API threw context window overflow error"
_RATE_LIMIT_MSG = "OpenAI Responses API threw rate limit error"

_IMAGE_FORMAT_MIME_TYPES: dict[str, str] = {
"png": "image/png",
"jpeg": "image/jpeg",
"jpg": "image/jpeg",
"gif": "image/gif",
"webp": "image/webp",
}


def _encode_media_to_data_url(data: bytes, format_ext: str, media_type: str = "image") -> str:
"""Encode media bytes to a base64 data URL with size validation.
Expand All @@ -95,7 +103,10 @@ def _encode_media_to_data_url(data: bytes, format_ext: str, media_type: str = "i
f"{media_type.capitalize()} size {len(data)} bytes exceeds maximum of"
f" {_MAX_MEDIA_SIZE_BYTES} bytes ({_MAX_MEDIA_SIZE_LABEL})"
)
mime_type = mimetypes.types_map.get(f".{format_ext}", _DEFAULT_MIME_TYPE)
if media_type == "image":
mime_type = _IMAGE_FORMAT_MIME_TYPES.get(format_ext, _DEFAULT_MIME_TYPE)
else:
mime_type = mimetypes.types_map.get(f".{format_ext}", _DEFAULT_MIME_TYPE)
encoded_data = base64.b64encode(data).decode("utf-8")
return f"data:{mime_type};base64,{encoded_data}"

Expand Down
16 changes: 12 additions & 4 deletions strands-py/src/strands/models/writer.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,6 @@
import base64
import json
import logging
import mimetypes
from collections.abc import AsyncGenerator
from typing import Any, TypeVar, cast

Expand All @@ -25,6 +24,14 @@

T = TypeVar("T", bound=BaseModel)

_IMAGE_FORMAT_MIME_TYPES: dict[str, str] = {
"png": "image/png",
"jpeg": "image/jpeg",
"jpg": "image/jpeg",
"gif": "image/gif",
"webp": "image/webp",
}


class WriterModel(Model):
"""Writer API model provider implementation."""
Expand Down Expand Up @@ -101,7 +108,7 @@ def _format_content_vision(content: ContentBlock) -> dict[str, Any]:
return {"text": content["text"], "type": "text"}

if "image" in content:
mime_type = mimetypes.types_map.get(f".{content['image']['format']}", "application/octet-stream")
mime_type = _IMAGE_FORMAT_MIME_TYPES.get(content["image"]["format"], "application/octet-stream")
image_data = base64.b64encode(content["image"]["source"]["bytes"]).decode("utf-8")

return {
Expand Down Expand Up @@ -141,8 +148,9 @@ def _format_content(content: ContentBlock) -> str:

content_blocks = list(
filter(
lambda content: content.get("text")
and not any(block_type in content for block_type in ["toolResult", "toolUse"]),
lambda content: (
content.get("text") and not any(block_type in content for block_type in ["toolResult", "toolUse"])
),
contents,
)
)
Expand Down
4 changes: 3 additions & 1 deletion strands-py/src/strands/types/content.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,14 +28,16 @@ class GuardContentText(TypedDict):
text: str


class GuardContent(TypedDict):
class GuardContent(TypedDict, total=False):
"""Content block to be evaluated by guardrails.

Attributes:
text: Text within content block to be evaluated by the guardrail.
image: Image within content block to be evaluated by the guardrail.
"""

text: GuardContentText
image: ImageContent


class ReasoningTextBlock(TypedDict, total=False):
Expand Down
2 changes: 1 addition & 1 deletion strands-py/src/strands/types/media.py
Original file line number Diff line number Diff line change
Expand Up @@ -107,7 +107,7 @@ class ImageContent(TypedDict):
source: ImageSource


VideoFormat = Literal["flv", "mkv", "mov", "mpeg", "mpg", "mp4", "three_gp", "webm", "wmv"]
VideoFormat = Literal["flv", "mkv", "mov", "mpeg", "mpg", "mp4", "3gp", "three_gp", "webm", "wmv"]
"""Supported video formats."""


Expand Down
16 changes: 13 additions & 3 deletions strands-py/tests/strands/models/test_anthropic.py
Original file line number Diff line number Diff line change
Expand Up @@ -214,14 +214,24 @@ def test_format_request_with_document(content, formatted_content, model, model_i
assert tru_request == exp_request


def test_format_request_with_image(model, model_id, max_tokens):
@pytest.mark.parametrize(
("image_format", "expected_media_type"),
[
("jpg", "image/jpeg"),
("jpeg", "image/jpeg"),
("png", "image/png"),
("gif", "image/gif"),
("webp", "image/webp"),
],
)
def test_format_request_with_image(model, model_id, max_tokens, image_format, expected_media_type):
messages = [
{
"role": "user",
"content": [
{
"image": {
"format": "jpg",
"format": image_format,
"source": {"bytes": b"base64encodedimage"},
},
},
Expand All @@ -239,7 +249,7 @@ def test_format_request_with_image(model, model_id, max_tokens):
{
"source": {
"data": "YmFzZTY0ZW5jb2RlZGltYWdl",
"media_type": "image/jpeg",
"media_type": expected_media_type,
"type": "base64",
},
"type": "image",
Expand Down
46 changes: 46 additions & 0 deletions strands-py/tests/strands/models/test_bedrock.py
Original file line number Diff line number Diff line change
Expand Up @@ -2289,6 +2289,28 @@ def test_format_request_video_s3_location(model, model_id):
assert video_source == {"s3Location": {"uri": "s3://my-bucket/video.mp4"}}


@pytest.mark.parametrize(
("input_format", "expected_format"),
[
("3gp", "three_gp"),
("three_gp", "three_gp"),
("mp4", "mp4"),
],
)
def test_format_request_video_format_normalizes_3gp(model, model_id, input_format, expected_format):
"""Bedrock expects the three_gp enum; accept the 3gp alias and map it through."""
messages = [
{
"role": "user",
"content": [{"video": {"format": input_format, "source": {"bytes": b"video_data"}}}],
}
]

formatted_video = model.format_request(messages)["messages"][0]["content"][0]["video"]

assert formatted_video["format"] == expected_format


def test_format_request_filters_document_content_blocks(model, model_id):
"""Test that format_request filters extra fields from document content blocks."""
messages = [
Expand Down Expand Up @@ -2774,6 +2796,30 @@ async def test_format_request_with_guardrail_latest_message(model):
assert formatted_messages[2]["content"][1]["guardContent"]["image"]["format"] == "png"


@pytest.mark.asyncio
@pytest.mark.parametrize("unsupported_format", ["webp", "gif"])
async def test_format_request_guardrail_skips_unsupported_image_format(model, unsupported_format):
"""Bedrock guardContent only supports png and jpeg; other formats pass through unwrapped."""
model.update_config(
guardrail_id="test-guardrail",
guardrail_version="DRAFT",
guardrail_latest_message=True,
)

messages = [
{
"role": "user",
"content": [{"image": {"format": unsupported_format, "source": {"bytes": b"fake_image_data"}}}],
},
]

request = model.format_request(messages)
formatted_image = request["messages"][0]["content"][0]

assert "guardContent" not in formatted_image
assert formatted_image["image"]["format"] == unsupported_format


@pytest.mark.asyncio
async def test_format_request_with_guardrail_latest_message_after_tool_use(model):
"""Test that guardContent wraps the last user text message even when a toolResult follows it."""
Expand Down
Loading